use crate::{Result, Subscripts};
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct EinsumSubscripts {
pub inputs: Vec<Vec<u32>>,
pub output: Vec<u32>,
}
impl EinsumSubscripts {
pub fn new(inputs: &[&[u32]], output: &[u32]) -> Self {
Self {
inputs: inputs.iter().map(|labels| labels.to_vec()).collect(),
output: output.to_vec(),
}
}
#[must_use]
pub fn input_count(&self) -> usize {
self.inputs.len()
}
}
impl From<Subscripts> for EinsumSubscripts {
fn from(subscripts: Subscripts) -> Self {
Self {
inputs: subscripts.inputs,
output: subscripts.output,
}
}
}
impl From<&Subscripts> for EinsumSubscripts {
fn from(subscripts: &Subscripts) -> Self {
Self {
inputs: subscripts.inputs.clone(),
output: subscripts.output.clone(),
}
}
}
impl From<EinsumSubscripts> for Subscripts {
fn from(subscripts: EinsumSubscripts) -> Self {
Self {
inputs: subscripts.inputs,
output: subscripts.output,
}
}
}
impl From<&EinsumSubscripts> for Subscripts {
fn from(subscripts: &EinsumSubscripts) -> Self {
Self {
inputs: subscripts.inputs.clone(),
output: subscripts.output.clone(),
}
}
}
pub fn parse_einsum_subscripts(notation: &str) -> Result<EinsumSubscripts> {
Subscripts::parse(notation).map(EinsumSubscripts::from)
}
#[cfg(test)]
mod tests;