Skip to main content

tenferro_einsum/
traced.rs

1use std::collections::hash_map::DefaultHasher;
2use std::hash::{Hash, Hasher};
3use std::sync::Arc;
4
5use tenferro_ops::dim_expr::DimExpr;
6use tenferro_runtime::extension::{ExtensionCacheKey, ExtensionCacheStore};
7use tenferro_runtime::program::ProgramBuildError;
8use tenferro_runtime::{TraceContext, TraceValue, TracedTensor};
9use tenferro_tensor::{ErrorKind, ValidationKind};
10
11use crate::cache::{
12    einsum_notation_retained_bytes, saturating_sum, ParsedEinsum, EINSUM_EXTENSION_FAMILY_ID,
13    EINSUM_PARSE_CACHE,
14};
15use crate::ellipsis::{notation_without_ellipsis, resolve_einsum_notation};
16use crate::extension::EinsumExtensionOp;
17use crate::optimize::{plan_spec_from_optimize, resolve_einsum_strategy_with_spec};
18use crate::{
19    parse_einsum_notation, EinsumAxis, EinsumNotation, EinsumOptimize, EinsumSubscripts, Error,
20    Result, Subscripts, TensorDotAxes,
21};
22
23/// Backend-neutral einsum tracing methods for [`TraceContext`].
24///
25/// Each method records one semantic einsum extension operation. Contraction
26/// path materialization and provider selection remain compiler/runtime work.
27///
28/// # Examples
29///
30/// ```
31/// use tenferro_einsum::TraceContextEinsumExt;
32/// use tenferro_ops::dim_expr::DimExpr;
33/// use tenferro_runtime::program::ProgramInputSpec;
34/// use tenferro_runtime::TraceContext;
35/// use tenferro_tensor::DType;
36///
37/// let matrix = || {
38///     ProgramInputSpec::new(
39///         DType::F64,
40///         [DimExpr::Const(2), DimExpr::Const(2)],
41///     )
42/// };
43/// let mut trace = TraceContext::new();
44/// let lhs = trace.input(matrix()).unwrap();
45/// let rhs = trace.input(matrix()).unwrap();
46/// let output = trace.einsum(&[lhs, rhs], "ij,jk->ik").unwrap();
47/// let graph = trace.finish(&[output]).unwrap();
48/// assert_eq!(graph.program().operations().count(), 1);
49/// ```
50pub trait TraceContextEinsumExt {
51    /// Trace textual einsum notation using the default optimizer policy.
52    ///
53    /// # Errors
54    ///
55    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
56    /// [`Error::Validation`] for invalid input metadata, [`Error::Planning`]
57    /// for an invalid optimizer policy, or [`Error::Runtime`] when semantic
58    /// program construction fails.
59    fn einsum(&mut self, inputs: &[TraceValue], subscripts: &str) -> Result<TraceValue>;
60
61    /// Trace rank-unresolved notation using the default optimizer policy.
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    ///
71    /// # Errors
72    ///
73    /// Returns [`Error::InvalidSubscripts`] for invalid axis tokens,
74    /// [`Error::Validation`] for invalid input metadata, [`Error::Planning`]
75    /// for symbolic ellipsis or an invalid optimizer policy, or [`Error::Runtime`]
76    /// for graph-build failures.
77    fn einsum_notation(
78        &mut self,
79        inputs: &[TraceValue],
80        notation: &EinsumNotation,
81    ) -> Result<TraceValue>;
82
83    /// Trace parsed einsum notation using the default optimizer policy.
84    ///
85    /// # Errors
86    ///
87    /// Returns [`Error::Validation`] for invalid input metadata,
88    /// [`Error::Planning`] for an invalid optimizer policy, or
89    /// [`Error::Runtime`] when semantic program construction fails.
90    fn einsum_subscripts(
91        &mut self,
92        inputs: &[TraceValue],
93        subscripts: &EinsumSubscripts,
94    ) -> Result<TraceValue>;
95
96    /// Trace textual einsum notation with an explicit optimizer policy.
97    ///
98    /// # Errors
99    ///
100    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
101    /// [`Error::Validation`] for invalid input metadata, [`Error::Planning`]
102    /// for an invalid optimizer policy, or [`Error::Runtime`] when semantic
103    /// program construction fails.
104    fn einsum_with(
105        &mut self,
106        inputs: &[TraceValue],
107        subscripts: &str,
108        optimize: EinsumOptimize,
109    ) -> Result<TraceValue>;
110
111    /// Trace rank-unresolved notation with an explicit optimizer policy.
112    ///
113    /// # Examples
114    ///
115    /// ```
116    /// use tenferro_einsum::{EinsumAxis, EinsumNotation};
117    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[]);
118    /// assert_eq!(notation.input_count(), 1);
119    /// ```
120    ///
121    /// # Errors
122    ///
123    /// Returns [`Error::InvalidSubscripts`] for invalid axis tokens,
124    /// [`Error::Validation`] for invalid input metadata, [`Error::Planning`]
125    /// for symbolic ellipsis or an invalid optimizer policy, or [`Error::Runtime`]
126    /// for graph-build failures.
127    fn einsum_notation_with(
128        &mut self,
129        inputs: &[TraceValue],
130        notation: &EinsumNotation,
131        optimize: EinsumOptimize,
132    ) -> Result<TraceValue>;
133
134    /// Trace parsed einsum notation with an explicit optimizer policy.
135    ///
136    /// # Errors
137    ///
138    /// Returns [`Error::Validation`] for invalid input metadata,
139    /// [`Error::Planning`] for an invalid optimizer policy, or
140    /// [`Error::Runtime`] when semantic program construction fails.
141    fn einsum_subscripts_with(
142        &mut self,
143        inputs: &[TraceValue],
144        subscripts: &EinsumSubscripts,
145        optimize: EinsumOptimize,
146    ) -> Result<TraceValue>;
147}
148
149impl TraceContextEinsumExt for TraceContext {
150    fn einsum(&mut self, inputs: &[TraceValue], subscripts: &str) -> Result<TraceValue> {
151        self.einsum_with(inputs, subscripts, EinsumOptimize::default())
152    }
153
154    fn einsum_notation(
155        &mut self,
156        inputs: &[TraceValue],
157        notation: &EinsumNotation,
158    ) -> Result<TraceValue> {
159        self.einsum_notation_with(inputs, notation, EinsumOptimize::default())
160    }
161
162    fn einsum_subscripts(
163        &mut self,
164        inputs: &[TraceValue],
165        subscripts: &EinsumSubscripts,
166    ) -> Result<TraceValue> {
167        self.einsum_subscripts_with(inputs, subscripts, EinsumOptimize::default())
168    }
169
170    fn einsum_with(
171        &mut self,
172        inputs: &[TraceValue],
173        subscripts: &str,
174        optimize: EinsumOptimize,
175    ) -> Result<TraceValue> {
176        let parsed = cached_subscripts(self.extension_caches_mut(), subscripts)?;
177        self.einsum_notation_with(inputs, &parsed.notation, optimize)
178    }
179
180    fn einsum_notation_with(
181        &mut self,
182        inputs: &[TraceValue],
183        notation: &EinsumNotation,
184        optimize: EinsumOptimize,
185    ) -> Result<TraceValue> {
186        trace_context_einsum_notation_with(self, inputs, notation, optimize)
187    }
188
189    fn einsum_subscripts_with(
190        &mut self,
191        inputs: &[TraceValue],
192        subscripts: &EinsumSubscripts,
193        optimize: EinsumOptimize,
194    ) -> Result<TraceValue> {
195        trace_context_einsum_subscripts_with(self, inputs, subscripts, optimize, false)
196    }
197}
198
199fn trace_context_einsum_notation_with(
200    trace: &mut TraceContext,
201    inputs: &[TraceValue],
202    notation: &EinsumNotation,
203    optimize: EinsumOptimize,
204) -> Result<TraceValue> {
205    let allow_broadcast = notation
206        .inputs
207        .iter()
208        .chain(std::iter::once(&notation.output))
209        .any(|term| term.contains(&EinsumAxis::Ellipsis));
210    let subscripts = if allow_broadcast {
211        let shapes = trace_concrete_shapes(trace, inputs)?.ok_or_else(|| {
212            Error::planning("ellipsis einsum requires concrete input dimensions during tracing")
213        })?;
214        let shape_refs: Vec<_> = shapes.iter().map(Vec::as_slice).collect();
215        resolve_einsum_notation(notation, &shape_refs)?.into()
216    } else {
217        notation_without_ellipsis(notation)?.into()
218    };
219    trace_context_einsum_subscripts_with(trace, inputs, &subscripts, optimize, allow_broadcast)
220}
221
222fn trace_context_einsum_subscripts_with(
223    trace: &mut TraceContext,
224    inputs: &[TraceValue],
225    subscripts: &EinsumSubscripts,
226    optimize: EinsumOptimize,
227    allow_broadcast: bool,
228) -> Result<TraceValue> {
229    if inputs.is_empty() {
230        return Err(Error::invalid_argument(
231            "einsum",
232            "inputs",
233            "einsum requires at least one input tensor",
234        ));
235    }
236    if subscripts.inputs.len() != inputs.len() {
237        return Err(Error::invalid_argument(
238            "einsum",
239            "inputs",
240            format!(
241                "einsum subscripts expect {} inputs, got {}",
242                subscripts.inputs.len(),
243                inputs.len()
244            ),
245        ));
246    }
247
248    let subs = Subscripts::from(subscripts);
249    let op = match optimize {
250        EinsumOptimize::Tree(tree) => {
251            let shapes = trace_concrete_shapes(trace, inputs)?.ok_or_else(|| {
252                Error::planning("precomputed contraction tree requires concrete input shapes")
253            })?;
254            let shape_refs: Vec<_> = shapes.iter().map(Vec::as_slice).collect();
255            let (plan_spec, _tree) =
256                resolve_einsum_strategy_with_spec(EinsumOptimize::Tree(tree), &subs, &shape_refs)?;
257            EinsumExtensionOp::with_plan_spec_and_broadcast(
258                subscripts.clone(),
259                plan_spec,
260                allow_broadcast,
261            )
262        }
263        optimize => EinsumExtensionOp::with_plan_spec_and_broadcast(
264            subscripts.clone(),
265            plan_spec_from_optimize(optimize, &subs)?,
266            allow_broadcast,
267        ),
268    };
269    let outputs = trace
270        .add_extension(Arc::new(op), inputs)
271        .map_err(semantic_trace_error)?;
272    outputs.first().copied().ok_or_else(|| {
273        Error::Runtime(tenferro_runtime::Error::Internal(
274            "einsum semantic extension produced no output".into(),
275        ))
276    })
277}
278
279fn trace_concrete_shapes(
280    trace: &TraceContext,
281    inputs: &[TraceValue],
282) -> Result<Option<Vec<Vec<usize>>>> {
283    let mut shapes = Vec::with_capacity(inputs.len());
284    for &value in inputs {
285        let metadata = trace.value_metadata(value).map_err(semantic_trace_error)?;
286        let Some(shape) = metadata
287            .shape()
288            .iter()
289            .map(|extent| match extent.as_exact() {
290                Some(DimExpr::Const(value)) => Some(*value),
291                _ => None,
292            })
293            .collect::<Option<Vec<_>>>()
294        else {
295            return Ok(None);
296        };
297        shapes.push(shape);
298    }
299    Ok(Some(shapes))
300}
301
302fn semantic_trace_error(source: ProgramBuildError) -> Error {
303    Error::Runtime(tenferro_runtime::Error::extension(
304        "einsum",
305        tenferro_runtime::ErrorPhase::GraphBuild,
306        EINSUM_EXTENSION_FAMILY_ID,
307        ErrorKind::Validation(ValidationKind::InvalidArgument),
308        source,
309    ))
310}
311
312/// Traced tensor contraction-sugar methods.
313pub trait TracedTensorEinsumExt {
314    ///
315    /// # Errors
316    ///
317    /// Returns [`Error::Validation`] with rank, axis, duplicate-axis, or
318    /// contracted-dimension payloads for invalid axes, or [`Error::Runtime`]
319    /// for graph-build failures.
320    ///
321    /// # Deferred errors
322    ///
323    /// Symbolic contracted-dimension equalities are checked during compilation
324    /// or execution and retain the runtime
325    /// [`ErrorPhase`](tenferro_runtime::ErrorPhase).
326    fn tensordot(&self, rhs: &TracedTensor, axes: TensorDotAxes<'_>) -> Result<TracedTensor>;
327}
328
329impl TracedTensorEinsumExt for TracedTensor {
330    fn tensordot(&self, rhs: &TracedTensor, axes: TensorDotAxes<'_>) -> Result<TracedTensor> {
331        tensordot(self, rhs, axes)
332    }
333}
334
335fn tensordot(
336    lhs: &TracedTensor,
337    rhs: &TracedTensor,
338    axes: TensorDotAxes<'_>,
339) -> Result<TracedTensor> {
340    let config = crate::tensordot::dot_general_config(axes, lhs.rank, rhs.rank)?;
341    crate::tensordot::validate_traced_contract_dims(lhs, rhs, &config)?;
342    lhs.dot_general(rhs, config).map_err(Error::Runtime)
343}
344
345struct ParsedEinsumCacheEntry {
346    notation: String,
347    parsed: Arc<ParsedEinsum>,
348}
349
350impl ParsedEinsumCacheEntry {
351    fn matches_notation(&self, notation: &str) -> bool {
352        self.notation == notation
353    }
354}
355
356fn cached_subscripts(
357    caches: &mut ExtensionCacheStore,
358    notation: &str,
359) -> Result<Arc<ParsedEinsum>> {
360    let key = ExtensionCacheKey::new(
361        EINSUM_EXTENSION_FAMILY_ID,
362        EINSUM_PARSE_CACHE,
363        hash_value(notation),
364    );
365    if let Some(cached) = caches.get::<ParsedEinsumCacheEntry>(&key) {
366        if cached.matches_notation(notation) {
367            return Ok(Arc::clone(&cached.parsed));
368        }
369    }
370
371    let parsed = Arc::new(ParsedEinsum {
372        notation: parse_einsum_notation(notation)?,
373    });
374    let entry = ParsedEinsumCacheEntry {
375        notation: notation.to_owned(),
376        parsed: Arc::clone(&parsed),
377    };
378    let retained_bytes = saturating_sum([
379        entry.notation.len(),
380        einsum_notation_retained_bytes(&parsed.notation),
381    ]);
382    caches.put(key, entry, retained_bytes);
383    Ok(parsed)
384}
385
386fn hash_value<T: Hash + ?Sized>(value: &T) -> u64 {
387    let mut hasher = DefaultHasher::new();
388    value.hash(&mut hasher);
389    hasher.finish()
390}
391
392#[cfg(test)]
393mod tests;