use crate::{Result, Subscripts};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum EinsumAxis {
Label(u32),
Ellipsis,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct EinsumNotation {
pub inputs: Vec<Vec<EinsumAxis>>,
pub output: Vec<EinsumAxis>,
}
impl EinsumNotation {
pub fn new(inputs: &[&[EinsumAxis]], output: &[EinsumAxis]) -> Self {
Self {
inputs: inputs.iter().map(|axes| axes.to_vec()).collect(),
output: output.to_vec(),
}
}
#[must_use]
pub fn input_count(&self) -> usize {
self.inputs.len()
}
pub fn parse(notation: &str) -> Result<Self> {
let (inputs_str, output_str) =
crate::syntax::notation::split_and_validate_notation(notation)?;
if inputs_str.contains(['(', ')']) || output_str.contains(['(', ')']) {
return Err(crate::Error::invalid_subscripts(
"EinsumNotation::parse does not accept parentheses; use NestedEinsum::parse for parenthesized contraction order",
));
}
let inputs = inputs_str
.split(',')
.map(parse_axis_term)
.collect::<Result<Vec<_>>>()?;
let output = parse_axis_term(output_str)?;
Ok(Self { inputs, output })
}
}
fn parse_axis_term(term: &str) -> Result<Vec<EinsumAxis>> {
let mut chars = term.chars();
let mut axes = Vec::new();
while let Some(c) = chars.next() {
if c == '.' {
if chars.next() != Some('.') || chars.next() != Some('.') {
return Err(crate::Error::invalid_subscripts(
"einsum ellipsis must be written as exactly three dots",
));
}
if axes.contains(&EinsumAxis::Ellipsis) {
return Err(crate::Error::invalid_subscripts(
"each einsum term may contain at most one ellipsis",
));
}
axes.push(EinsumAxis::Ellipsis);
} else {
axes.push(EinsumAxis::Label(crate::syntax::notation::char_to_label(
c,
)?));
}
}
Ok(axes)
}
#[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)
}
pub fn parse_einsum_notation(notation: &str) -> Result<EinsumNotation> {
EinsumNotation::parse(notation)
}
#[cfg(test)]
mod tests;