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_subscripts_retained_bytes, saturating_sum, ParsedEinsum, EINSUM_EXTENSION_FAMILY_ID,
13    EINSUM_PARSE_CACHE,
14};
15use crate::extension::EinsumExtensionOp;
16use crate::optimize::{plan_spec_from_optimize, resolve_einsum_strategy_with_spec};
17use crate::{
18    parse_einsum_subscripts, EinsumOptimize, EinsumSubscripts, Error, Result, Subscripts,
19    TensorDotAxes,
20};
21
22/// Backend-neutral einsum tracing methods for [`TraceContext`].
23///
24/// Each method records one semantic einsum extension operation. Contraction
25/// path materialization and provider selection remain compiler/runtime work.
26///
27/// # Examples
28///
29/// ```
30/// use tenferro_einsum::TraceContextEinsumExt;
31/// use tenferro_ops::dim_expr::DimExpr;
32/// use tenferro_runtime::program::ProgramInputSpec;
33/// use tenferro_runtime::TraceContext;
34/// use tenferro_tensor::DType;
35///
36/// let matrix = || {
37///     ProgramInputSpec::new(
38///         DType::F64,
39///         [DimExpr::Const(2), DimExpr::Const(2)],
40///     )
41/// };
42/// let mut trace = TraceContext::new();
43/// let lhs = trace.input(matrix()).unwrap();
44/// let rhs = trace.input(matrix()).unwrap();
45/// let output = trace.einsum(&[lhs, rhs], "ij,jk->ik").unwrap();
46/// let graph = trace.finish(&[output]).unwrap();
47/// assert_eq!(graph.program().operations().count(), 1);
48/// ```
49pub trait TraceContextEinsumExt {
50    /// Trace textual einsum notation using the default optimizer policy.
51    ///
52    /// # Errors
53    ///
54    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
55    /// [`Error::Validation`] for invalid input metadata, [`Error::Planning`]
56    /// for an invalid optimizer policy, or [`Error::Runtime`] when semantic
57    /// program construction fails.
58    fn einsum(&mut self, inputs: &[TraceValue], subscripts: &str) -> Result<TraceValue>;
59
60    /// Trace parsed einsum notation using the default optimizer policy.
61    ///
62    /// # Errors
63    ///
64    /// Returns [`Error::Validation`] for invalid input metadata,
65    /// [`Error::Planning`] for an invalid optimizer policy, or
66    /// [`Error::Runtime`] when semantic program construction fails.
67    fn einsum_subscripts(
68        &mut self,
69        inputs: &[TraceValue],
70        subscripts: &EinsumSubscripts,
71    ) -> Result<TraceValue>;
72
73    /// Trace textual einsum notation with an explicit optimizer policy.
74    ///
75    /// # Errors
76    ///
77    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
78    /// [`Error::Validation`] for invalid input metadata, [`Error::Planning`]
79    /// for an invalid optimizer policy, or [`Error::Runtime`] when semantic
80    /// program construction fails.
81    fn einsum_with(
82        &mut self,
83        inputs: &[TraceValue],
84        subscripts: &str,
85        optimize: EinsumOptimize,
86    ) -> Result<TraceValue>;
87
88    /// Trace parsed einsum notation with an explicit optimizer policy.
89    ///
90    /// # Errors
91    ///
92    /// Returns [`Error::Validation`] for invalid input metadata,
93    /// [`Error::Planning`] for an invalid optimizer policy, or
94    /// [`Error::Runtime`] when semantic program construction fails.
95    fn einsum_subscripts_with(
96        &mut self,
97        inputs: &[TraceValue],
98        subscripts: &EinsumSubscripts,
99        optimize: EinsumOptimize,
100    ) -> Result<TraceValue>;
101}
102
103impl TraceContextEinsumExt for TraceContext {
104    fn einsum(&mut self, inputs: &[TraceValue], subscripts: &str) -> Result<TraceValue> {
105        self.einsum_with(inputs, subscripts, EinsumOptimize::default())
106    }
107
108    fn einsum_subscripts(
109        &mut self,
110        inputs: &[TraceValue],
111        subscripts: &EinsumSubscripts,
112    ) -> Result<TraceValue> {
113        self.einsum_subscripts_with(inputs, subscripts, EinsumOptimize::default())
114    }
115
116    fn einsum_with(
117        &mut self,
118        inputs: &[TraceValue],
119        subscripts: &str,
120        optimize: EinsumOptimize,
121    ) -> Result<TraceValue> {
122        let parsed = cached_subscripts(self.extension_caches_mut(), subscripts)?;
123        self.einsum_subscripts_with(inputs, &parsed.subscripts, optimize)
124    }
125
126    fn einsum_subscripts_with(
127        &mut self,
128        inputs: &[TraceValue],
129        subscripts: &EinsumSubscripts,
130        optimize: EinsumOptimize,
131    ) -> Result<TraceValue> {
132        trace_context_einsum_subscripts_with(self, inputs, subscripts, optimize)
133    }
134}
135
136fn trace_context_einsum_subscripts_with(
137    trace: &mut TraceContext,
138    inputs: &[TraceValue],
139    subscripts: &EinsumSubscripts,
140    optimize: EinsumOptimize,
141) -> Result<TraceValue> {
142    if inputs.is_empty() {
143        return Err(Error::invalid_argument(
144            "einsum",
145            "inputs",
146            "einsum requires at least one input tensor",
147        ));
148    }
149    if subscripts.inputs.len() != inputs.len() {
150        return Err(Error::invalid_argument(
151            "einsum",
152            "inputs",
153            format!(
154                "einsum subscripts expect {} inputs, got {}",
155                subscripts.inputs.len(),
156                inputs.len()
157            ),
158        ));
159    }
160
161    let subs = Subscripts::from(subscripts);
162    let op = match optimize {
163        EinsumOptimize::Tree(tree) => {
164            let shapes = trace_concrete_shapes(trace, inputs)?.ok_or_else(|| {
165                Error::planning("precomputed contraction tree requires concrete input shapes")
166            })?;
167            let shape_refs: Vec<_> = shapes.iter().map(Vec::as_slice).collect();
168            let (plan_spec, _tree) =
169                resolve_einsum_strategy_with_spec(EinsumOptimize::Tree(tree), &subs, &shape_refs)?;
170            EinsumExtensionOp::with_plan_spec(subscripts.clone(), plan_spec)
171        }
172        optimize => EinsumExtensionOp::with_plan_spec(
173            subscripts.clone(),
174            plan_spec_from_optimize(optimize, &subs)?,
175        ),
176    };
177    let outputs = trace
178        .add_extension(Arc::new(op), inputs)
179        .map_err(semantic_trace_error)?;
180    outputs.first().copied().ok_or_else(|| {
181        Error::Runtime(tenferro_runtime::Error::Internal(
182            "einsum semantic extension produced no output".into(),
183        ))
184    })
185}
186
187fn trace_concrete_shapes(
188    trace: &TraceContext,
189    inputs: &[TraceValue],
190) -> Result<Option<Vec<Vec<usize>>>> {
191    let mut shapes = Vec::with_capacity(inputs.len());
192    for &value in inputs {
193        let metadata = trace.value_metadata(value).map_err(semantic_trace_error)?;
194        let Some(shape) = metadata
195            .shape()
196            .iter()
197            .map(|extent| match extent.as_exact() {
198                Some(DimExpr::Const(value)) => Some(*value),
199                _ => None,
200            })
201            .collect::<Option<Vec<_>>>()
202        else {
203            return Ok(None);
204        };
205        shapes.push(shape);
206    }
207    Ok(Some(shapes))
208}
209
210fn semantic_trace_error(source: ProgramBuildError) -> Error {
211    Error::Runtime(tenferro_runtime::Error::extension(
212        "einsum",
213        tenferro_runtime::ErrorPhase::GraphBuild,
214        EINSUM_EXTENSION_FAMILY_ID,
215        ErrorKind::Validation(ValidationKind::InvalidArgument),
216        source,
217    ))
218}
219
220/// Traced tensor contraction-sugar methods.
221pub trait TracedTensorEinsumExt {
222    ///
223    /// # Errors
224    ///
225    /// Returns [`Error::Validation`] with rank, axis, duplicate-axis, or
226    /// contracted-dimension payloads for invalid axes, or [`Error::Runtime`]
227    /// for graph-build failures.
228    ///
229    /// # Deferred errors
230    ///
231    /// Symbolic contracted-dimension equalities are checked during compilation
232    /// or execution and retain the runtime
233    /// [`ErrorPhase`](tenferro_runtime::ErrorPhase).
234    fn tensordot(&self, rhs: &TracedTensor, axes: TensorDotAxes<'_>) -> Result<TracedTensor>;
235}
236
237impl TracedTensorEinsumExt for TracedTensor {
238    fn tensordot(&self, rhs: &TracedTensor, axes: TensorDotAxes<'_>) -> Result<TracedTensor> {
239        tensordot(self, rhs, axes)
240    }
241}
242
243fn tensordot(
244    lhs: &TracedTensor,
245    rhs: &TracedTensor,
246    axes: TensorDotAxes<'_>,
247) -> Result<TracedTensor> {
248    let config = crate::tensordot::dot_general_config(axes, lhs.rank, rhs.rank)?;
249    crate::tensordot::validate_traced_contract_dims(lhs, rhs, &config)?;
250    lhs.dot_general(rhs, config).map_err(Error::Runtime)
251}
252
253struct ParsedEinsumCacheEntry {
254    notation: String,
255    parsed: Arc<ParsedEinsum>,
256}
257
258impl ParsedEinsumCacheEntry {
259    fn matches_notation(&self, notation: &str) -> bool {
260        self.notation == notation
261    }
262}
263
264fn cached_subscripts(
265    caches: &mut ExtensionCacheStore,
266    notation: &str,
267) -> Result<Arc<ParsedEinsum>> {
268    let key = ExtensionCacheKey::new(
269        EINSUM_EXTENSION_FAMILY_ID,
270        EINSUM_PARSE_CACHE,
271        hash_value(notation),
272    );
273    if let Some(cached) = caches.get::<ParsedEinsumCacheEntry>(&key) {
274        if cached.matches_notation(notation) {
275            return Ok(Arc::clone(&cached.parsed));
276        }
277    }
278
279    let parsed = Arc::new(ParsedEinsum {
280        subscripts: parse_einsum_subscripts(notation)?,
281    });
282    let entry = ParsedEinsumCacheEntry {
283        notation: notation.to_owned(),
284        parsed: Arc::clone(&parsed),
285    };
286    let retained_bytes = saturating_sum([
287        entry.notation.len(),
288        einsum_subscripts_retained_bytes(&parsed.subscripts),
289    ]);
290    caches.put(key, entry, retained_bytes);
291    Ok(parsed)
292}
293
294fn hash_value<T: Hash + ?Sized>(value: &T) -> u64 {
295    let mut hasher = DefaultHasher::new();
296    value.hash(&mut hasher);
297    hasher.finish()
298}
299
300#[cfg(test)]
301mod tests;