Skip to main content

tenferro_einsum/
subscripts.rs

1use crate::{Result, Subscripts};
2
3/// One unresolved axis token in rank-polymorphic einsum notation.
4///
5/// # Examples
6///
7/// ```
8/// use tenferro_einsum::EinsumAxis;
9/// assert_eq!(EinsumAxis::Ellipsis, EinsumAxis::Ellipsis);
10/// ```
11#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
12pub enum EinsumAxis {
13    /// An explicit integer label.
14    Label(u32),
15    /// A NumPy-style ellipsis whose rank is resolved from the inputs.
16    Ellipsis,
17}
18
19/// Rank-unresolved einsum notation.
20///
21/// This is the programmatic counterpart of string notation such as
22/// `"...ij,...jk->...ik"`. Resolve it through an einsum operation, after input
23/// ranks are known; [`EinsumSubscripts`] remains the rank-resolved runtime form.
24///
25/// # Examples
26///
27/// ```
28/// use tenferro_einsum::{EinsumAxis, EinsumNotation};
29/// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
30/// assert_eq!(notation.input_count(), 1);
31/// ```
32#[derive(Clone, Debug, PartialEq, Eq, Hash)]
33pub struct EinsumNotation {
34    /// Axis tokens for each input tensor.
35    pub inputs: Vec<Vec<EinsumAxis>>,
36    /// Axis tokens for the output tensor.
37    pub output: Vec<EinsumAxis>,
38}
39
40impl EinsumNotation {
41    /// Create rank-unresolved notation from axis-token arrays.
42    ///
43    /// # Examples
44    ///
45    /// ```
46    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
47    ///
48    /// let notation = EinsumNotation::new(
49    ///     &[&[EinsumAxis::Ellipsis, EinsumAxis::Label(0)], &[EinsumAxis::Ellipsis, EinsumAxis::Label(0)]],
50    ///     &[EinsumAxis::Ellipsis],
51    /// );
52    /// assert_eq!(notation.input_count(), 2);
53    /// ```
54    pub fn new(inputs: &[&[EinsumAxis]], output: &[EinsumAxis]) -> Self {
55        Self {
56            inputs: inputs.iter().map(|axes| axes.to_vec()).collect(),
57            output: output.to_vec(),
58        }
59    }
60
61    /// Number of input operands described by this notation.
62    ///
63    /// # Examples
64    ///
65    /// ```
66    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
67    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
68    /// assert_eq!(notation.input_count(), 1);
69    /// ```
70    #[must_use]
71    pub fn input_count(&self) -> usize {
72        self.inputs.len()
73    }
74
75    /// Parse flat string notation, retaining ellipsis tokens.
76    ///
77    /// # Errors
78    ///
79    /// Returns [`crate::Error::InvalidSubscripts`] for malformed separators,
80    /// labels, ellipses, or parenthesized contraction order.
81    pub fn parse(notation: &str) -> Result<Self> {
82        let (inputs_str, output_str) =
83            crate::syntax::notation::split_and_validate_notation(notation)?;
84        if inputs_str.contains(['(', ')']) || output_str.contains(['(', ')']) {
85            return Err(crate::Error::invalid_subscripts(
86                "EinsumNotation::parse does not accept parentheses; use NestedEinsum::parse for parenthesized contraction order",
87            ));
88        }
89        let inputs = inputs_str
90            .split(',')
91            .map(parse_axis_term)
92            .collect::<Result<Vec<_>>>()?;
93        let output = parse_axis_term(output_str)?;
94        Ok(Self { inputs, output })
95    }
96}
97
98fn parse_axis_term(term: &str) -> Result<Vec<EinsumAxis>> {
99    let mut chars = term.chars();
100    let mut axes = Vec::new();
101    while let Some(c) = chars.next() {
102        if c == '.' {
103            if chars.next() != Some('.') || chars.next() != Some('.') {
104                return Err(crate::Error::invalid_subscripts(
105                    "einsum ellipsis must be written as exactly three dots",
106                ));
107            }
108            if axes.contains(&EinsumAxis::Ellipsis) {
109                return Err(crate::Error::invalid_subscripts(
110                    "each einsum term may contain at most one ellipsis",
111                ));
112            }
113            axes.push(EinsumAxis::Ellipsis);
114        } else {
115            axes.push(EinsumAxis::Label(crate::syntax::notation::char_to_label(
116                c,
117            )?));
118        }
119    }
120    Ok(axes)
121}
122
123/// Canonical N-ary einsum subscripts using integer labels.
124///
125/// String notation is a user-facing convenience. Runtime integration layers can
126/// carry this representation in extension payloads so execution, shape
127/// inference, and AD do not need to parse strings.
128#[derive(Clone, Debug, PartialEq, Eq, Hash)]
129pub struct EinsumSubscripts {
130    /// Index labels for each input tensor.
131    pub inputs: Vec<Vec<u32>>,
132    /// Index labels for the output tensor.
133    pub output: Vec<u32>,
134}
135
136impl EinsumSubscripts {
137    /// Create subscripts from integer label arrays.
138    ///
139    /// # Examples
140    ///
141    /// ```
142    /// use tenferro_einsum::EinsumSubscripts;
143    ///
144    /// let subscripts = EinsumSubscripts::new(&[&[0, 1], &[1, 2]], &[0, 2]);
145    ///
146    /// assert_eq!(subscripts.inputs, vec![vec![0, 1], vec![1, 2]]);
147    /// assert_eq!(subscripts.output, vec![0, 2]);
148    /// ```
149    pub fn new(inputs: &[&[u32]], output: &[u32]) -> Self {
150        Self {
151            inputs: inputs.iter().map(|labels| labels.to_vec()).collect(),
152            output: output.to_vec(),
153        }
154    }
155
156    /// Number of input operands described by this specification.
157    ///
158    /// # Examples
159    ///
160    /// ```
161    /// use tenferro_einsum::EinsumSubscripts;
162    ///
163    /// let subscripts = EinsumSubscripts::new(&[&[0], &[0]], &[]);
164    ///
165    /// assert_eq!(subscripts.input_count(), 2);
166    /// ```
167    #[must_use]
168    pub fn input_count(&self) -> usize {
169        self.inputs.len()
170    }
171}
172
173impl From<Subscripts> for EinsumSubscripts {
174    fn from(subscripts: Subscripts) -> Self {
175        Self {
176            inputs: subscripts.inputs,
177            output: subscripts.output,
178        }
179    }
180}
181
182impl From<&Subscripts> for EinsumSubscripts {
183    fn from(subscripts: &Subscripts) -> Self {
184        Self {
185            inputs: subscripts.inputs.clone(),
186            output: subscripts.output.clone(),
187        }
188    }
189}
190
191impl From<EinsumSubscripts> for Subscripts {
192    fn from(subscripts: EinsumSubscripts) -> Self {
193        Self {
194            inputs: subscripts.inputs,
195            output: subscripts.output,
196        }
197    }
198}
199
200impl From<&EinsumSubscripts> for Subscripts {
201    fn from(subscripts: &EinsumSubscripts) -> Self {
202        Self {
203            inputs: subscripts.inputs.clone(),
204            output: subscripts.output.clone(),
205        }
206    }
207}
208
209/// Parse string einsum notation into canonical integer labels.
210///
211/// # Examples
212///
213/// ```
214/// use tenferro_einsum::parse_einsum_subscripts;
215///
216/// let subscripts = parse_einsum_subscripts("ij,jk->ik").unwrap();
217///
218/// assert_eq!(subscripts.inputs.len(), 2);
219/// assert_eq!(subscripts.output, vec![b'i' as u32, b'k' as u32]);
220/// ```
221///
222/// # Errors
223///
224/// Returns [`crate::Error::InvalidSubscripts`] when the notation is malformed,
225/// contains an invalid label, or has an invalid input/output separator.
226pub fn parse_einsum_subscripts(notation: &str) -> Result<EinsumSubscripts> {
227    Subscripts::parse(notation).map(EinsumSubscripts::from)
228}
229
230/// Parse string notation into rank-unresolved axis tokens.
231///
232/// # Examples
233///
234/// ```
235/// use tenferro_einsum::{parse_einsum_notation, EinsumAxis};
236///
237/// let notation = parse_einsum_notation("...ij,...jk->...ik").unwrap();
238/// assert_eq!(notation.inputs[0][0], EinsumAxis::Ellipsis);
239/// ```
240///
241/// # Errors
242///
243/// Returns [`crate::Error::InvalidSubscripts`] for malformed notation,
244/// multiple ellipses in one term, or parenthesized contraction order.
245pub fn parse_einsum_notation(notation: &str) -> Result<EinsumNotation> {
246    EinsumNotation::parse(notation)
247}
248
249#[cfg(test)]
250mod tests;