use super::ReadOptions;
use super::wire::{self, KEY_EDGES, KEY_METADATA, KEY_NODE, KEY_NODES, KEY_TYPE, KEY_VERSION};
use crate::error::{NirError, Result};
use crate::graph::NirGraph;
use crate::nodes::{
Affine, AvgPool2d, Conv1d, Conv2d, CubaLi, CubaLif, Delay, Flatten, I, If, Input, Li, Lif,
Linear, NirNode, Output, Padding, Scale, SumPool2d, Threshold,
};
use crate::types::{MetadataMap, MetadataValue, Tensor, TensorData};
use hdf5::plist::dataset_create::Layout;
use hdf5::types::{
FixedAscii, FixedUnicode, FloatSize, IntSize, TypeDescriptor as Td, VarLenAscii, VarLenUnicode,
};
use hdf5::{Dataset, File, Group, LocationToken};
use std::cell::Cell;
use std::path::Path;
pub(super) const MAX_NESTED_GRAPHS: usize = 1024;
struct ReadBudget {
limit: Option<usize>,
used: Cell<usize>,
}
impl ReadBudget {
fn new(opts: &ReadOptions) -> Self {
Self {
limit: opts.max_bytes,
used: Cell::new(0),
}
}
fn charge(&self, context: &str, requested: Option<usize>) -> Result<()> {
let Some(limit) = self.limit else {
return Ok(());
};
let used = self.used.get();
let Some(requested) = requested else {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested: usize::MAX,
});
};
let Some(next) = used.checked_add(requested) else {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested,
});
};
if next > limit {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested,
});
}
self.used.set(next);
Ok(())
}
fn would_fit(&self, context: &str, requested: Option<usize>) -> Result<()> {
let Some(limit) = self.limit else {
return Ok(());
};
let used = self.used.get();
let Some(requested) = requested else {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested: usize::MAX,
});
};
let Some(next) = used.checked_add(requested) else {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested,
});
};
if next > limit {
return Err(NirError::ReadLimitExceeded {
context: context.to_owned(),
limit,
used,
requested,
});
}
Ok(())
}
}
const FIXED_STRING_CAPS: [usize; 3] = [64, 256, 4096];
const _: () = assert!(
FIXED_STRING_CAPS[0] == 64 && FIXED_STRING_CAPS[1] == 256 && FIXED_STRING_CAPS[2] == 4096,
"FIXED_STRING_CAPS and the read_fixed! rungs in read_strings_unchecked must match"
);
pub(super) fn read(path: &Path, opts: &ReadOptions) -> Result<NirGraph> {
let budget = ReadBudget::new(opts);
let file = open(path)?;
validate_group_links(&file, "/")?;
let root = file.group(KEY_NODE).map_err(|_| {
NirError::MissingField(format!(
"/{KEY_NODE} (not a NIR graph file: {})",
path.display()
))
})?;
validate_group_links(&root, &format!("/{KEY_NODE}"))?;
let root_type = read_string_scalar(
&root
.dataset(KEY_TYPE)
.map_err(|_| NirError::MissingField(format!("/{KEY_NODE}/{KEY_TYPE}")))?,
KEY_NODE,
&budget,
)?;
if root_type != "NIRGraph" {
return Err(NirError::InvalidGraph(format!(
"/{KEY_NODE} must be a NIRGraph, found {root_type:?}"
)));
}
let mut graph = read_graph_body(&root, &format!("/{KEY_NODE}"), &mut Vec::new(), &budget)?;
graph.metadata = read_metadata(&root, &budget)?;
graph.version = match version_dataset(&file)? {
Some(ds) => Some(read_string_scalar(&ds, KEY_VERSION, &budget)?),
None => None,
};
Ok(graph)
}
fn version_dataset(file: &File) -> Result<Option<Dataset>> {
if !file.link_exists(KEY_VERSION) {
return Ok(None);
}
file.dataset(KEY_VERSION).map(Some).map_err(|e| {
NirError::Io(format!(
"/{KEY_VERSION}: expected a dataset, found another link kind: {e}"
))
})
}
pub(super) fn read_version(path: &Path, opts: &ReadOptions) -> Result<String> {
let budget = ReadBudget::new(opts);
let file = open(path)?;
validate_group_links(&file, "/")?;
let ds =
version_dataset(&file)?.ok_or_else(|| NirError::MissingField(format!("/{KEY_VERSION}")))?;
read_string_scalar(&ds, KEY_VERSION, &budget)
}
fn open(path: &Path) -> Result<File> {
File::open(path).map_err(|e| NirError::Io(format!("cannot open {}: {e}", path.display())))
}
fn read_graph_body(
group: &Group,
context: &str,
visited: &mut Vec<LocationToken>,
budget: &ReadBudget,
) -> Result<NirGraph> {
if visited.len() >= MAX_NESTED_GRAPHS {
return Err(NirError::InvalidGraph(format!(
"{context}: more than {MAX_NESTED_GRAPHS} nested NIRGraph groups"
)));
}
let token = group.loc_info()?.token;
if visited.contains(&token) {
return Err(NirError::InvalidGraph(format!(
"{context}: nested NIRGraph group is already being decoded (hard-link cycle or alias)"
)));
}
visited.push(token);
read_graph_body_inner(group, context, visited, budget)
}
fn read_graph_body_inner(
group: &Group,
context: &str,
visited: &mut Vec<LocationToken>,
budget: &ReadBudget,
) -> Result<NirGraph> {
let mut graph = NirGraph::new();
validate_group_links(group, context)?;
let nodes = group
.group(KEY_NODES)
.map_err(|_| NirError::MissingField(format!("{context}/{KEY_NODES}")))?;
validate_group_links(&nodes, &format!("{context}/{KEY_NODES}"))?;
let mut names = nodes.member_names()?;
names.sort();
for name in names {
let node_group = nodes
.group(&name)
.map_err(|e| NirError::Io(format!("node {name:?} is not a group: {e}")))?;
validate_group_links(&node_group, &name)?;
let node = read_node(&node_group, &name, visited, budget)?;
graph.insert_node(name, node)?;
}
let edges = group
.dataset(KEY_EDGES)
.map_err(|_| NirError::MissingField(format!("{context}/{KEY_EDGES}")))?;
graph.edges = read_edges(&edges, budget)?;
Ok(graph)
}
fn read_edges(ds: &Dataset, budget: &ReadBudget) -> Result<Vec<(String, String)>> {
validate_dataset_security(ds, KEY_EDGES)?;
if ds.size() == 0 {
return Ok(Vec::new());
}
let shape = ds.shape();
if shape.len() != 2 || shape[1] != 2 {
return Err(NirError::InvalidGraph(format!(
"{KEY_EDGES} must have shape (E, 2), found {shape:?}"
)));
}
let flat = read_strings_unchecked(ds, KEY_EDGES, budget)?;
let edge_count = shape[0];
budget.charge(
KEY_EDGES,
edge_count.checked_mul(std::mem::size_of::<(String, String)>()),
)?;
let mut strings = flat.into_iter();
Ok(std::iter::from_fn(|| Some((strings.next()?, strings.next()?))).collect())
}
fn read_node(
group: &Group,
name: &str,
visited: &mut Vec<LocationToken>,
budget: &ReadBudget,
) -> Result<NirNode> {
let type_ds = group
.dataset(KEY_TYPE)
.map_err(|_| NirError::MissingField(format!("{name}.{KEY_TYPE}")))?;
let ty = read_string_scalar(&type_ds, name, budget)?;
let metadata = read_metadata(group, budget)?;
let r = NodeReader {
group,
name,
budget,
};
let node = match ty.as_str() {
"Input" => NirNode::Input(read_input(&r, metadata)?),
"Output" => NirNode::Output(read_output(&r, metadata)?),
"Affine" => NirNode::Affine(read_affine(&r, metadata)?),
"Linear" => NirNode::Linear(read_linear(&r, metadata)?),
"Scale" => NirNode::Scale(read_scale(&r, metadata)?),
"Conv1d" => NirNode::Conv1d(read_conv1d(&r, metadata)?),
"Conv2d" => NirNode::Conv2d(read_conv2d(&r, metadata)?),
"CubaLI" => NirNode::CubaLi(read_cuba_li(&r, metadata)?),
"CubaLIF" => NirNode::CubaLif(read_cuba_lif(&r, metadata)?),
"Delay" => NirNode::Delay(read_delay(&r, metadata)?),
"Flatten" => NirNode::Flatten(read_flatten(&r, metadata)?),
"I" => NirNode::I(read_i(&r, metadata)?),
"IF" => NirNode::If(read_if(&r, metadata)?),
"LI" => NirNode::Li(read_li(&r, metadata)?),
"LIF" => NirNode::Lif(read_lif(&r, metadata)?),
"SumPool2d" => NirNode::SumPool2d(read_sum_pool2d(&r, metadata)?),
"AvgPool2d" => NirNode::AvgPool2d(read_avg_pool2d(&r, metadata)?),
"Threshold" => NirNode::Threshold(read_threshold(&r, metadata)?),
"NIRGraph" => {
let mut sub = read_graph_body(group, name, visited, budget)?;
sub.metadata = metadata;
NirNode::Graph(Box::new(sub))
}
other => return Err(NirError::UnknownNodeType(other.to_owned())),
};
debug_assert!(
wire::is_wire_type(node.type_name()),
"decoded a node whose type is not in WIRE_TYPES"
);
Ok(node)
}
fn read_input(r: &NodeReader, metadata: MetadataMap) -> Result<Input> {
Ok(Input {
shape: r.usizes("shape")?,
metadata,
})
}
fn read_output(r: &NodeReader, metadata: MetadataMap) -> Result<Output> {
Ok(Output {
shape: r.usizes("shape")?,
metadata,
})
}
fn read_affine(r: &NodeReader, metadata: MetadataMap) -> Result<Affine> {
Ok(Affine {
weight: r.tensor("weight")?,
bias: r.tensor("bias")?,
metadata,
})
}
fn read_linear(r: &NodeReader, metadata: MetadataMap) -> Result<Linear> {
Ok(Linear {
weight: r.tensor("weight")?,
metadata,
})
}
fn read_scale(r: &NodeReader, metadata: MetadataMap) -> Result<Scale> {
Ok(Scale {
scale: r.tensor("scale")?,
metadata,
})
}
struct ConvGeometry {
weight: Tensor,
stride: Vec<i64>,
padding: Padding,
dilation: Vec<i64>,
groups: i64,
bias: Tensor,
}
impl ConvGeometry {
fn read(r: &NodeReader) -> Result<Self> {
Ok(Self {
weight: r.tensor("weight")?,
stride: r.ints("stride")?,
padding: r.padding()?,
dilation: r.ints("dilation")?,
groups: r.int_scalar("groups")?,
bias: r.tensor("bias")?,
})
}
fn into_conv1d(self, input_shape: Option<usize>, metadata: MetadataMap) -> Conv1d {
let Self {
weight,
stride,
padding,
dilation,
groups,
bias,
} = self;
Conv1d {
weight,
stride,
padding,
dilation,
groups,
bias,
input_shape,
metadata,
}
}
fn into_conv2d(self, input_shape: Option<Vec<usize>>, metadata: MetadataMap) -> Conv2d {
Conv2d {
weight: self.weight,
stride: self.stride,
padding: self.padding,
dilation: self.dilation,
groups: self.groups,
bias: self.bias,
input_shape,
metadata,
}
}
}
fn read_conv1d(r: &NodeReader, metadata: MetadataMap) -> Result<Conv1d> {
Ok(ConvGeometry::read(r)?.into_conv1d(r.opt_dim("input_shape")?, metadata))
}
fn read_conv2d(r: &NodeReader, metadata: MetadataMap) -> Result<Conv2d> {
Ok(ConvGeometry::read(r)?.into_conv2d(r.opt_usizes("input_shape")?, metadata))
}
fn threshold_and_reset(r: &NodeReader) -> Result<(Tensor, Option<Tensor>)> {
let v_threshold = r.tensor("v_threshold")?;
let v_reset = r.v_reset(&v_threshold)?;
Ok((v_threshold, v_reset))
}
struct LeakDynamics {
tau: Tensor,
r: Tensor,
v_leak: Tensor,
}
impl LeakDynamics {
fn read(r: &NodeReader) -> Result<Self> {
Ok(Self {
tau: r.tensor("tau")?,
r: r.tensor("r")?,
v_leak: r.tensor("v_leak")?,
})
}
fn into_li(self, metadata: MetadataMap) -> Li {
let Self { tau, r, v_leak } = self;
Li {
tau,
r,
v_leak,
metadata,
}
}
fn into_lif(self, v_threshold: Tensor, v_reset: Option<Tensor>, metadata: MetadataMap) -> Lif {
let Self { tau, r, v_leak } = self;
Lif {
tau,
r,
v_leak,
v_reset,
v_threshold,
metadata,
}
}
}
struct CubaDynamics {
tau_syn: Tensor,
tau_mem: Tensor,
r: Tensor,
w_in: Option<Tensor>,
v_leak: Tensor,
}
impl CubaDynamics {
fn read(r: &NodeReader) -> Result<Self> {
let v_leak = r.tensor("v_leak")?;
Ok(Self {
tau_syn: r.tensor("tau_syn")?,
tau_mem: r.tensor("tau_mem")?,
r: r.tensor("r")?,
w_in: r.w_in(&v_leak)?,
v_leak,
})
}
fn into_li(self, metadata: MetadataMap) -> CubaLi {
let Self {
tau_syn,
tau_mem,
r,
w_in,
v_leak,
} = self;
CubaLi {
tau_syn,
tau_mem,
r,
w_in,
v_leak,
metadata,
}
}
fn into_lif(
self,
v_threshold: Tensor,
v_reset: Option<Tensor>,
metadata: MetadataMap,
) -> CubaLif {
let Self {
tau_syn,
tau_mem,
r,
w_in,
v_leak,
} = self;
CubaLif {
tau_syn,
tau_mem,
r,
v_reset,
w_in,
v_leak,
v_threshold,
metadata,
}
}
}
fn read_cuba_li(r: &NodeReader, metadata: MetadataMap) -> Result<CubaLi> {
Ok(CubaDynamics::read(r)?.into_li(metadata))
}
fn read_cuba_lif(r: &NodeReader, metadata: MetadataMap) -> Result<CubaLif> {
let (v_threshold, v_reset) = threshold_and_reset(r)?;
Ok(CubaDynamics::read(r)?.into_lif(v_threshold, v_reset, metadata))
}
fn read_i(r: &NodeReader, metadata: MetadataMap) -> Result<I> {
Ok(I {
r: r.tensor("r")?,
metadata,
})
}
fn read_if(r: &NodeReader, metadata: MetadataMap) -> Result<If> {
let (v_threshold, v_reset) = threshold_and_reset(r)?;
Ok(If {
r: r.tensor("r")?,
v_reset,
v_threshold,
metadata,
})
}
fn read_li(r: &NodeReader, metadata: MetadataMap) -> Result<Li> {
Ok(LeakDynamics::read(r)?.into_li(metadata))
}
fn read_lif(r: &NodeReader, metadata: MetadataMap) -> Result<Lif> {
let (v_threshold, v_reset) = threshold_and_reset(r)?;
Ok(LeakDynamics::read(r)?.into_lif(v_threshold, v_reset, metadata))
}
fn read_pool_window(r: &NodeReader) -> Result<(Tensor, Tensor, Tensor)> {
Ok((
r.tensor("kernel_size")?,
r.tensor("stride")?,
r.tensor("padding")?,
))
}
fn read_sum_pool2d(r: &NodeReader, metadata: MetadataMap) -> Result<SumPool2d> {
let (kernel_size, stride, padding) = read_pool_window(r)?;
Ok(SumPool2d {
kernel_size,
stride,
padding,
metadata,
})
}
fn read_avg_pool2d(r: &NodeReader, metadata: MetadataMap) -> Result<AvgPool2d> {
let (kernel_size, stride, padding) = read_pool_window(r)?;
Ok(AvgPool2d {
kernel_size,
stride,
padding,
metadata,
})
}
fn read_delay(r: &NodeReader, metadata: MetadataMap) -> Result<Delay> {
Ok(Delay {
delay: r.tensor("delay")?,
metadata,
})
}
fn read_flatten(r: &NodeReader, metadata: MetadataMap) -> Result<Flatten> {
Ok(Flatten {
start_dim: r.opt_int_scalar("start_dim")?.unwrap_or(1),
end_dim: r.opt_int_scalar("end_dim")?.unwrap_or(-1),
input_type: r.opt_usizes("input_type")?,
metadata,
})
}
fn read_threshold(r: &NodeReader, metadata: MetadataMap) -> Result<Threshold> {
Ok(Threshold {
threshold: r.tensor("threshold")?,
metadata,
})
}
struct NodeReader<'a> {
group: &'a Group,
name: &'a str,
budget: &'a ReadBudget,
}
impl NodeReader<'_> {
fn context(&self, field: &str) -> String {
format!("{}.{field}", self.name)
}
fn optional(&self, field: &str) -> Result<Option<Dataset>> {
if !self.group.link_exists(field) {
return Ok(None);
}
self.group.dataset(field).map(Some).map_err(|e| {
NirError::Io(format!(
"{}: expected a dataset, found another link kind: {e}",
self.context(field)
))
})
}
fn required(&self, field: &str) -> Result<Dataset> {
self.optional(field)?
.ok_or_else(|| NirError::MissingField(self.context(field)))
}
fn tensor(&self, field: &str) -> Result<Tensor> {
read_tensor(&self.required(field)?, &self.context(field), self.budget)
}
fn opt_tensor(&self, field: &str) -> Result<Option<Tensor>> {
match self.optional(field)? {
Some(ds) => read_tensor(&ds, &self.context(field), self.budget).map(Some),
None => Ok(None),
}
}
fn ints(&self, field: &str) -> Result<Vec<i64>> {
read_ints(&self.required(field)?, &self.context(field), self.budget)
}
fn opt_ints(&self, field: &str) -> Result<Option<Vec<i64>>> {
match self.optional(field)? {
Some(ds) => read_ints(&ds, &self.context(field), self.budget).map(Some),
None => Ok(None),
}
}
fn int_scalar(&self, field: &str) -> Result<i64> {
single_int(self.ints(field)?, &self.context(field))
}
fn opt_int_scalar(&self, field: &str) -> Result<Option<i64>> {
match self.opt_ints(field)? {
Some(values) => single_int(values, &self.context(field)).map(Some),
None => Ok(None),
}
}
fn usizes(&self, field: &str) -> Result<Vec<usize>> {
to_usizes(self.ints(field)?, &self.context(field), self.budget)
}
fn opt_usizes(&self, field: &str) -> Result<Option<Vec<usize>>> {
match self.opt_ints(field)? {
Some(values) => to_usizes(values, &self.context(field), self.budget).map(Some),
None => Ok(None),
}
}
fn opt_dim(&self, field: &str) -> Result<Option<usize>> {
let Some(values) = self.opt_usizes(field)? else {
return Ok(None);
};
match values.as_slice() {
[only] => Ok(Some(*only)),
other => Err(NirError::InvalidTensor(format!(
"{}: expected a single extent, found {} values",
self.context(field),
other.len()
))),
}
}
fn padding(&self) -> Result<Padding> {
let ds = self.required("padding")?;
if is_string(&ds)? {
wire::padding_from_wire_str(&read_string_scalar(&ds, self.name, self.budget)?)
} else {
Ok(Padding::Explicit(read_ints(
&ds,
&self.context("padding"),
self.budget,
)?))
}
}
fn v_reset(&self, v_threshold: &Tensor) -> Result<Option<Tensor>> {
match self.opt_tensor("v_reset")? {
Some(tensor) => Ok(Some(tensor)),
None => {
self.budget.charge(
&self.context("v_reset (synthesized)"),
v_threshold
.data()
.len()
.checked_mul(v_threshold.dtype().size_of()),
)?;
Ok(Some(v_threshold.zeros_like()))
}
}
}
fn w_in(&self, v_leak: &Tensor) -> Result<Option<Tensor>> {
match self.opt_tensor("w_in")? {
Some(tensor) => Ok(Some(tensor)),
None => {
self.budget.charge(
&self.context("w_in (synthesized)"),
v_leak.data().len().checked_mul(v_leak.dtype().size_of()),
)?;
Ok(Some(v_leak.ones_like()))
}
}
}
}
fn to_usizes(values: Vec<i64>, context: &str, budget: &ReadBudget) -> Result<Vec<usize>> {
budget.charge(context, values.len().checked_mul(size_of::<usize>()))?;
values
.into_iter()
.map(|v| {
usize::try_from(v).map_err(|_| {
NirError::InvalidTensor(format!("{context}: negative axis length {v}"))
})
})
.collect()
}
fn single_int(values: Vec<i64>, context: &str) -> Result<i64> {
match values.as_slice() {
[only] => Ok(*only),
other => Err(NirError::InvalidTensor(format!(
"{context}: expected a single integer, found {} values",
other.len()
))),
}
}
fn validate_group_links(group: &Group, context: &str) -> Result<()> {
let bad_link = group
.iter_visit_default(
None,
|_group, name, info, found: &mut Option<(String, &'static str)>| match info.link_type {
hdf5::LinkType::External => {
*found = Some((name.to_owned(), "external"));
false
}
hdf5::LinkType::Soft => {
*found = Some((name.to_owned(), "soft"));
false
}
hdf5::LinkType::Hard => true,
},
)
.map_err(|e| NirError::Io(format!("{context}: cannot inspect group links: {e}")))?;
if let Some((name, kind)) = bad_link {
return Err(NirError::InvalidGraph(format!(
"{context}: {kind} link '{name}' is not allowed"
)));
}
Ok(())
}
fn validate_dataset_security(ds: &Dataset, context: &str) -> Result<()> {
let dcpl = ds.dcpl().map_err(|e| {
NirError::Io(format!(
"{context}: cannot read dataset creation property list: {e}"
))
})?;
let external = dcpl.external();
if !external.is_empty() {
return Err(NirError::InvalidGraph(format!(
"{context}: external storage is not allowed ({} external file(s))",
external.len()
)));
}
if !matches!(
ds.layout(),
Layout::Compact | Layout::Contiguous | Layout::Chunked
) {
return Err(NirError::InvalidGraph(format!(
"{context}: virtual dataset layouts are not allowed"
)));
}
Ok(())
}
fn read_metadata(group: &Group, budget: &ReadBudget) -> Result<MetadataMap> {
if !group.link_exists(KEY_METADATA) {
return Ok(MetadataMap::new());
}
let md = group.group(KEY_METADATA).map_err(|e| {
NirError::Io(format!(
"{KEY_METADATA}: expected a group, found another link kind: {e}"
))
})?;
validate_group_links(&md, KEY_METADATA)?;
let mut out = MetadataMap::new();
for key in md.member_names()? {
let ds = md.dataset(&key).map_err(|e| {
NirError::Io(format!(
"metadata {key:?} must be a dataset, not a group: {e}"
))
})?;
let value = read_metadata_value(&ds, &key, budget)?;
out.insert(key, value);
}
Ok(out)
}
fn read_string_metadata(
ds: &Dataset,
key: &str,
context: &str,
budget: &ReadBudget,
) -> Result<MetadataValue> {
match ds.shape().as_slice() {
[] => Ok(MetadataValue::String(read_string_scalar_validated(
ds, key, budget,
)?)),
[_] => Ok(MetadataValue::StringList(read_strings_unchecked(
ds, context, budget,
)?)),
shape => Err(NirError::InvalidTensor(format!(
"{context}: string metadata must be scalar or rank-1, found shape {shape:?}"
))),
}
}
fn read_metadata_value(ds: &Dataset, key: &str, budget: &ReadBudget) -> Result<MetadataValue> {
validate_dataset_security(ds, &format!("{KEY_METADATA}.{key}"))?;
let scalar = ds.shape().is_empty();
let context = format!("{KEY_METADATA}.{key}");
let value = match ds.dtype()?.to_descriptor()? {
Td::VarLenUnicode | Td::VarLenAscii | Td::FixedAscii(_) | Td::FixedUnicode(_) => {
read_string_metadata(ds, key, &context, budget)?
}
Td::Boolean if scalar => {
budget.charge(&context, Some(size_of::<bool>()))?;
MetadataValue::Bool(ds.read_scalar::<bool>()?)
}
Td::Float(_) if scalar => {
budget.charge(&context, Some(size_of::<f64>()))?;
MetadataValue::F64(ds.read_scalar::<f64>()?)
}
Td::Integer(_) if scalar => {
budget.charge(&context, Some(size_of::<i64>()))?;
MetadataValue::I64(ds.read_scalar::<i64>()?)
}
Td::Unsigned(IntSize::U8) if scalar => {
budget.charge(&context, Some(size_of::<i64>()))?;
let v = ds.read_scalar::<u64>()?;
MetadataValue::I64(i64::try_from(v).map_err(|_| {
NirError::InvalidTensor(format!("{context}: u64 value {v} does not fit in i64"))
})?)
}
Td::Unsigned(IntSize::U1 | IntSize::U2 | IntSize::U4) if scalar => {
budget.charge(&context, Some(size_of::<i64>()))?;
MetadataValue::I64(ds.read_scalar::<i64>()?)
}
_ => MetadataValue::Tensor(read_tensor(ds, &context, budget)?),
};
Ok(value)
}
fn read_tensor(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<Tensor> {
validate_dataset_security(ds, context)?;
let descriptor = ds.dtype()?.to_descriptor()?;
let data = match descriptor {
Td::Float(FloatSize::U4) => {
charge_elements(budget, ds, context, size_of::<f32>())?;
TensorData::F32(ds.read_raw::<f32>()?)
}
Td::Float(FloatSize::U8) => {
charge_elements(budget, ds, context, size_of::<f64>())?;
TensorData::F64(ds.read_raw::<f64>()?)
}
Td::Integer(_) | Td::Unsigned(IntSize::U1 | IntSize::U2 | IntSize::U4) => {
charge_elements(budget, ds, context, size_of::<i64>())?;
TensorData::I64(ds.read_raw::<i64>()?)
}
Td::Unsigned(IntSize::U8) => {
let two_buffers = size_of::<u64>().checked_add(size_of::<i64>());
budget.charge(
context,
two_buffers.and_then(|width| ds.size().checked_mul(width)),
)?;
TensorData::I64(read_u64_as_i64(ds, context)?)
}
Td::Boolean => {
charge_elements(budget, ds, context, size_of::<bool>())?;
TensorData::Bool(ds.read_raw::<bool>()?)
}
other => {
return Err(NirError::InvalidTensor(format!(
"{context}: element type {other} has no NIR dtype"
)));
}
};
Tensor::new(ds.shape(), data)
}
fn charge_elements(
budget: &ReadBudget,
ds: &Dataset,
context: &str,
decoded_width: usize,
) -> Result<()> {
budget.charge(context, ds.size().checked_mul(decoded_width))
}
fn read_u64_as_i64(ds: &Dataset, context: &str) -> Result<Vec<i64>> {
ds.read_raw::<u64>()?
.into_iter()
.map(|v| {
i64::try_from(v).map_err(|_| {
NirError::InvalidTensor(format!("{context}: u64 value {v} does not fit in i64"))
})
})
.collect()
}
fn read_ints(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<Vec<i64>> {
match read_tensor(ds, context, budget)?.into_data() {
TensorData::I64(values) => Ok(values),
other => Err(NirError::InvalidTensor(format!(
"{context}: expected integer data, found {:?}",
other.dtype()
))),
}
}
fn is_string(ds: &Dataset) -> Result<bool> {
Ok(matches!(
ds.dtype()?.to_descriptor()?,
Td::VarLenUnicode | Td::VarLenAscii | Td::FixedAscii(_) | Td::FixedUnicode(_)
))
}
fn read_string_scalar(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<String> {
validate_dataset_security(ds, context)?;
read_string_scalar_validated(ds, context, budget)
}
fn read_string_scalar_validated(
ds: &Dataset,
context: &str,
budget: &ReadBudget,
) -> Result<String> {
let size = ds.size();
if size != 1 {
return Err(NirError::Io(format!(
"{context}: expected a single string, found {size} elements"
)));
}
let mut values = read_strings_unchecked(ds, context, budget)?;
match values.len() {
1 => Ok(values.remove(0)),
n => Err(NirError::Io(format!(
"{context}: expected a single string, found {n}"
))),
}
}
fn read_strings_unchecked(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<Vec<String>> {
charge_strings(ds, context, budget)?;
macro_rules! read_fixed {
($ty:ident, $width:expr, $($cap:literal),+) => {
$(if $width <= $cap {
return Ok(ds
.read_raw::<$ty<$cap>>()?
.iter()
.map(ToString::to_string)
.collect());
})+
};
}
match ds.dtype()?.to_descriptor()? {
Td::VarLenUnicode => Ok(ds
.read_raw::<VarLenUnicode>()?
.iter()
.map(ToString::to_string)
.collect()),
Td::VarLenAscii => Ok(ds
.read_raw::<VarLenAscii>()?
.iter()
.map(ToString::to_string)
.collect()),
Td::FixedAscii(width) => {
read_fixed!(FixedAscii, width, 64, 256, 4096);
Err(too_wide(context, width))
}
Td::FixedUnicode(width) => {
read_fixed!(FixedUnicode, width, 64, 256, 4096);
Err(too_wide(context, width))
}
other => Err(NirError::Io(format!(
"{context}: expected a string dataset, found {other}"
))),
}
}
fn charge_strings(ds: &Dataset, context: &str, budget: &ReadBudget) -> Result<()> {
if budget.limit.is_none() {
return Ok(());
}
let count = ds.size();
let headers = count.checked_mul(size_of::<String>());
let requested = match ds.dtype()?.to_descriptor()? {
Td::VarLenUnicode => {
let descriptors = count.checked_mul(size_of::<VarLenUnicode>());
let min = checked_sum([descriptors, headers]);
budget.would_fit(context, min)?;
let payload = vlen_payload_bytes(ds, context)?;
checked_sum([descriptors, Some(payload), headers, Some(payload)])
}
Td::VarLenAscii => {
let descriptors = count.checked_mul(size_of::<VarLenAscii>());
let min = checked_sum([descriptors, headers]);
budget.would_fit(context, min)?;
let payload = vlen_payload_bytes(ds, context)?;
checked_sum([descriptors, Some(payload), headers, Some(payload)])
}
Td::FixedAscii(width) | Td::FixedUnicode(width) => {
let Some(capacity) = FIXED_STRING_CAPS.iter().copied().find(|cap| width <= *cap) else {
return Ok(());
};
checked_sum([
count.checked_mul(capacity),
headers,
count.checked_mul(width),
])
}
_ => return Ok(()),
};
budget.charge(context, requested)
}
fn checked_sum<const N: usize>(parts: [Option<usize>; N]) -> Option<usize> {
parts
.into_iter()
.try_fold(0usize, |sum, part| sum.checked_add(part?))
}
#[allow(deprecated)]
fn vlen_payload_bytes(ds: &Dataset, context: &str) -> Result<usize> {
if ds.shape().is_empty() {
return usize::try_from(ds.file()?.size()).map_err(|_| {
NirError::Io(format!(
"{context}: containing file size does not fit usize"
))
});
}
let dtype = ds.dtype()?;
let space = ds.space()?;
let mut bytes: hdf5_sys::h5::hsize_t = 0;
let status = hdf5::sync::sync(|| {
unsafe { hdf5_sys::h5d::H5Dvlen_get_buf_size(ds.id(), dtype.id(), space.id(), &mut bytes) }
});
hdf5::h5check(status).map_err(|e| {
NirError::Io(format!(
"{context}: cannot determine variable-length string allocation: {e}"
))
})?;
usize::try_from(bytes).map_err(|_| {
NirError::Io(format!(
"{context}: variable-length string allocation does not fit usize"
))
})
}
fn too_wide(context: &str, width: usize) -> NirError {
NirError::Io(format!(
"{context}: fixed-length string of {width} bytes exceeds the supported maximum of 4096"
))
}