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;