use std::collections::{HashMap, HashSet};
use crate::syntax::subscripts::Subscripts;
use crate::{Error, Result};
pub(crate) fn build_size_dict(
subscripts: &Subscripts,
shapes: &[&[usize]],
extra: Option<&HashMap<u32, usize>>,
) -> Result<HashMap<u32, usize>> {
if subscripts.inputs.len() != shapes.len() {
return Err(Error::InvalidArgument(format!(
"expected {} input shapes, got {}",
subscripts.inputs.len(),
shapes.len()
)));
}
let mut size_dict: HashMap<u32, usize> = HashMap::new();
for (i, input_subs) in subscripts.inputs.iter().enumerate() {
if input_subs.len() != shapes[i].len() {
return Err(Error::InvalidArgument(format!(
"input {} has {} subscript labels but shape has {} dimensions",
i,
input_subs.len(),
shapes[i].len()
)));
}
for (j, &label) in input_subs.iter().enumerate() {
let size = shapes[i][j];
if let Some(&existing) = size_dict.get(&label) {
if existing != size {
return Err(Error::ShapeMismatch {
expected: vec![existing],
got: vec![size],
});
}
} else {
size_dict.insert(label, size);
}
}
}
if let Some(sd) = extra {
for (&label, &size) in sd {
size_dict.entry(label).or_insert(size);
}
}
Ok(size_dict)
}
pub(crate) fn compute_output_shape(
output_subs: &[u32],
size_dict: &HashMap<u32, usize>,
) -> Result<Vec<usize>> {
output_subs
.iter()
.map(|&label| {
size_dict
.get(&label)
.copied()
.ok_or_else(|| Error::InvalidArgument(format!("unknown size for label {label}")))
})
.collect()
}
pub(crate) fn intermediate_subs(
subs_left: &[u32],
subs_right: &[u32],
needed: &HashSet<u32>,
) -> Vec<u32> {
let mut seen = HashSet::new();
let mut output = Vec::new();
for &l in subs_left.iter().chain(subs_right.iter()) {
if needed.contains(&l) && seen.insert(l) {
output.push(l);
}
}
output
}