Skip to main content

tenferro_einsum/
eager_ad.rs

1//! Eager einsum and tensordot on a borrowed [`EagerSession`].
2
3use std::collections::hash_map::DefaultHasher;
4use std::error::Error as StdError;
5use std::hash::{Hash, Hasher};
6use std::mem::size_of;
7use std::sync::Arc;
8
9use computegraph::compile::{compile, CompiledProgram, Instruction};
10use computegraph::graph::GraphBuilder;
11use computegraph::materialize::materialize_merge;
12use computegraph::resolve::resolve;
13use computegraph::types::{ValueKey, ValueRef};
14use tenferro_ad::extension::{
15    adopt_untracked_eager_value, apply_eager_with_targeted_extension_in_session,
16    EagerExtensionBackendKind, EagerExtensionTarget,
17};
18use tenferro_ad::{EagerSession, EagerTensor};
19use tenferro_cpu::CpuBackend;
20#[cfg(feature = "cuda")]
21use tenferro_gpu::cuda::CudaBackend;
22#[cfg(feature = "webgpu")]
23use tenferro_gpu::webgpu::WebGpuBackend;
24use tenferro_ops::dim_expr::DimExpr;
25use tenferro_ops::input_key::TensorInputKey;
26use tenferro_ops::std_tensor_op::StdTensorOp;
27use tenferro_runtime::{ErrorPhase, ExtensionCacheKey, ExtensionModule};
28use tenferro_tensor::{ErrorKind, ShapeMismatch, ValidationError, ValidationKind};
29
30use crate::binary_dot::{try_build_exact_output_binary_dot_plan, BinaryDotOperandOrder};
31use crate::builder::build_einsum_graph;
32use crate::cache::{
33    saturating_sum, vec_retained_bytes, EINSUM_EAGER_EXPANDED_PROGRAMS_CACHE,
34    EINSUM_EXTENSION_FAMILY_ID,
35};
36use crate::ellipsis::resolve_einsum_notation;
37use crate::extension::EinsumExtensionOp;
38use crate::optimize::{
39    default_auto_options, hash_einsum_plan_spec, plan_specs_equal, resolve_plan_spec,
40    EinsumPlanSpec,
41};
42use crate::{
43    parse_einsum_notation, EinsumNotation, EinsumSubscripts, Error, Result, Subscripts,
44    TensorDotAxes,
45};
46
47/// Eager einsum and tensordot on a runtime-bound borrowed eager session.
48///
49/// Every operation runs inside the caller's session, so it composes with other
50/// eager operations in the same [`tenferro_ad::EagerRuntime::with_eager_session`]
51/// callback and never reopens the runtime. The calling thread's `no_grad` and
52/// `capture_trace` modes govern it like any other eager operation.
53///
54/// # Examples
55///
56/// ```rust
57/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
58/// use tenferro_einsum::EagerSessionEinsumExt;
59///
60/// let ctx = EagerRuntime::new()?;
61/// let a = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6])?, ctx.clone())?;
62/// let b = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12])?, ctx.clone())?;
63/// let product = ctx.with_eager_session(|session| {
64///     let c = session.einsum(&[&a, &b], "ij,jk->ik")?;
65///     session.einsum(&[&c], "ij->")
66/// })?;
67/// assert_eq!(product.value()?.as_slice::<f64>()?, &[24.0]);
68/// # Ok::<(), Box<dyn std::error::Error>>(())
69/// ```
70#[cfg_attr(docsrs, doc(cfg(feature = "autodiff")))]
71pub trait EagerSessionEinsumExt {
72    /// Execute an einsum from string notation.
73    ///
74    /// # Examples
75    ///
76    /// ```rust
77    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
78    /// use tenferro_einsum::EagerSessionEinsumExt;
79    ///
80    /// let ctx = EagerRuntime::new()?;
81    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?, ctx.clone())?;
82    /// let dot = ctx.with_eager_session(|session| session.einsum(&[&x, &x], "i,i->"))?;
83    /// assert_eq!(dot.value()?.as_slice::<f64>()?, &[13.0]);
84    /// # Ok::<(), Box<dyn std::error::Error>>(())
85    /// ```
86    ///
87    /// # Errors
88    ///
89    /// Returns [`Error::InvalidSubscripts`] for malformed notation,
90    /// [`Error::Validation`] for rank/shape/dtype mismatches, or
91    /// [`Error::Planning`] / [`Error::Runtime`] for contraction planning and
92    /// execution failures, including inputs owned by another runtime.
93    fn einsum(&mut self, inputs: &[&EagerTensor], subscripts: &str) -> Result<EagerTensor>;
94
95    /// Execute an einsum from rank-unresolved notation.
96    ///
97    /// # Examples
98    ///
99    /// ```rust
100    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
101    /// use tenferro_einsum::{EagerSessionEinsumExt, EinsumAxis, EinsumNotation};
102    ///
103    /// let ctx = EagerRuntime::new()?;
104    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?, ctx.clone())?;
105    /// // `...->...` keeps every axis the ellipsis covers.
106    /// let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[EinsumAxis::Ellipsis]);
107    /// let same = ctx.with_eager_session(|session| session.einsum_notation(&[&x], &notation))?;
108    /// assert_eq!(same.value()?.as_slice::<f64>()?, &[2.0, 3.0]);
109    /// # Ok::<(), Box<dyn std::error::Error>>(())
110    /// ```
111    ///
112    /// # Errors
113    ///
114    /// Returns a typed validation, planning, or runtime error when notation or
115    /// execution is invalid.
116    fn einsum_notation(
117        &mut self,
118        inputs: &[&EagerTensor],
119        notation: &EinsumNotation,
120    ) -> Result<EagerTensor>;
121
122    /// Execute an einsum from parsed integer labels.
123    ///
124    /// # Examples
125    ///
126    /// ```rust
127    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
128    /// use tenferro_einsum::{EagerSessionEinsumExt, EinsumSubscripts};
129    ///
130    /// let ctx = EagerRuntime::new()?;
131    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?, ctx.clone())?;
132    /// let subscripts = EinsumSubscripts::new(&[&[0], &[0]], &[]);
133    /// let dot = ctx.with_eager_session(|session| session.einsum_subscripts(&[&x, &x], &subscripts))?;
134    /// assert_eq!(dot.value()?.as_slice::<f64>()?, &[13.0]);
135    /// # Ok::<(), Box<dyn std::error::Error>>(())
136    /// ```
137    ///
138    /// # Errors
139    ///
140    /// Returns [`Error::Validation`] for rank/shape/dtype mismatches,
141    /// [`Error::Planning`] for an invalid contraction plan, or
142    /// [`Error::Runtime`] for extension registration or backend execution
143    /// failures.
144    fn einsum_subscripts(
145        &mut self,
146        inputs: &[&EagerTensor],
147        subscripts: &EinsumSubscripts,
148    ) -> Result<EagerTensor>;
149
150    /// Contract two eager tensors over the requested axes.
151    ///
152    /// # Examples
153    ///
154    /// ```rust
155    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
156    /// use tenferro_einsum::{EagerSessionEinsumExt, TensorDotAxes};
157    ///
158    /// let ctx = EagerRuntime::new()?;
159    /// let a = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6])?, ctx.clone())?;
160    /// let b = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12])?, ctx.clone())?;
161    /// let c = ctx.with_eager_session(|session| session.tensordot(&a, &b, TensorDotAxes::Count(1)))?;
162    /// assert_eq!(c.shape(), &[2, 4]);
163    /// # Ok::<(), Box<dyn std::error::Error>>(())
164    /// ```
165    ///
166    /// # Errors
167    ///
168    /// Returns [`Error::Validation`] for invalid axes or mismatched contracted
169    /// extents, or [`Error::Runtime`] for execution failures.
170    fn tensordot(
171        &mut self,
172        lhs: &EagerTensor,
173        rhs: &EagerTensor,
174        axes: TensorDotAxes<'_>,
175    ) -> Result<EagerTensor>;
176}
177
178impl EagerSessionEinsumExt for EagerSession<'_> {
179    fn einsum(&mut self, inputs: &[&EagerTensor], subscripts: &str) -> Result<EagerTensor> {
180        einsum(self, inputs, subscripts)
181    }
182
183    fn einsum_notation(
184        &mut self,
185        inputs: &[&EagerTensor],
186        notation: &EinsumNotation,
187    ) -> Result<EagerTensor> {
188        einsum_notation(self, inputs, notation)
189    }
190
191    fn einsum_subscripts(
192        &mut self,
193        inputs: &[&EagerTensor],
194        subscripts: &EinsumSubscripts,
195    ) -> Result<EagerTensor> {
196        einsum_subscripts_with_broadcast(self, inputs, subscripts, false)
197    }
198
199    fn tensordot(
200        &mut self,
201        lhs: &EagerTensor,
202        rhs: &EagerTensor,
203        axes: TensorDotAxes<'_>,
204    ) -> Result<EagerTensor> {
205        tensordot(self, lhs, rhs, axes)
206    }
207}
208
209fn eager_extension_module(
210    target: EagerExtensionTarget,
211) -> tenferro_runtime::Result<Arc<dyn ExtensionModule>> {
212    let EagerExtensionTarget {
213        engine_id,
214        backend_kind,
215    } = target;
216    match backend_kind {
217        EagerExtensionBackendKind::Cpu => {
218            crate::extension::extension_module::<CpuBackend>(engine_id)
219                .map_err(eager_runtime_config_error)
220        }
221        #[cfg(feature = "cuda")]
222        EagerExtensionBackendKind::Cuda => {
223            crate::extension::extension_module::<CudaBackend>(engine_id)
224                .map_err(eager_runtime_config_error)
225        }
226        #[cfg(feature = "webgpu")]
227        EagerExtensionBackendKind::WebGpu => {
228            crate::extension::extension_module::<WebGpuBackend>(engine_id)
229                .map_err(eager_runtime_config_error)
230        }
231    }
232}
233
234fn eager_runtime_config_error(
235    source: tenferro_runtime::RuntimeConfigError,
236) -> tenferro_runtime::Error {
237    tenferro_runtime::Error::runtime_state_source(
238        "tenferro_einsum::eager_extension_module",
239        ErrorPhase::Execution,
240        source,
241    )
242}
243
244fn einsum(
245    session: &mut EagerSession<'_>,
246    inputs: &[&EagerTensor],
247    subscripts: &str,
248) -> Result<EagerTensor> {
249    let notation = parse_einsum_notation(subscripts)?;
250    einsum_notation(session, inputs, &notation)
251}
252
253fn einsum_notation(
254    session: &mut EagerSession<'_>,
255    inputs: &[&EagerTensor],
256    notation: &EinsumNotation,
257) -> Result<EagerTensor> {
258    let shapes: Vec<&[usize]> = inputs.iter().map(|tensor| tensor.shape()).collect();
259    let subscripts = resolve_einsum_notation(notation, &shapes)?;
260    let subscripts = EinsumSubscripts::from(subscripts);
261    let allow_broadcast = notation
262        .inputs
263        .iter()
264        .chain(std::iter::once(&notation.output))
265        .any(|term| term.contains(&crate::EinsumAxis::Ellipsis))
266        || requires_broadcast(inputs, &subscripts);
267    einsum_subscripts_with_broadcast(session, inputs, &subscripts, allow_broadcast)
268}
269
270fn einsum_subscripts_with_broadcast(
271    session: &mut EagerSession<'_>,
272    inputs: &[&EagerTensor],
273    subscripts: &EinsumSubscripts,
274    allow_broadcast: bool,
275) -> Result<EagerTensor> {
276    if let Some(result) = try_direct_binary_dot_general(session, inputs, subscripts) {
277        return result;
278    }
279
280    let output_shape_hint = infer_eager_output_shape(subscripts, inputs)?;
281    if !requires_broadcast(inputs, subscripts) {
282        if let Some(result) = try_expand_eager_einsum(session, inputs, subscripts)? {
283            return Ok(result);
284        }
285    }
286
287    let plan_spec = EinsumPlanSpec::Auto(default_auto_options());
288    let op = Arc::new(if allow_broadcast {
289        EinsumExtensionOp::with_output_shape_hint_and_broadcast(
290            subscripts.clone(),
291            output_shape_hint,
292            plan_spec,
293            true,
294        )
295    } else {
296        EinsumExtensionOp::with_output_shape_hint(subscripts.clone(), output_shape_hint, plan_spec)
297    });
298    let mut outputs = apply_eager_with_targeted_extension_in_session(
299        session,
300        op,
301        inputs,
302        eager_extension_module,
303    )?;
304    outputs.pop().ok_or_else(|| {
305        Error::Runtime(tenferro_runtime::Error::MissingInput(
306            "einsum extension produced no eager output".into(),
307        ))
308    })
309}
310
311fn try_direct_binary_dot_general(
312    session: &mut EagerSession<'_>,
313    inputs: &[&EagerTensor],
314    subscripts: &EinsumSubscripts,
315) -> Option<Result<EagerTensor>> {
316    if inputs.len() != 2 || subscripts.inputs.len() != 2 {
317        return None;
318    }
319
320    let lhs_labels = &subscripts.inputs[0];
321    let rhs_labels = &subscripts.inputs[1];
322    if lhs_labels.len() != inputs[0].shape().len() || rhs_labels.len() != inputs[1].shape().len() {
323        return None;
324    }
325
326    if let Some(plan) =
327        try_build_exact_output_binary_dot_plan(lhs_labels, rhs_labels, &subscripts.output)
328    {
329        let (lhs, rhs) = match plan.operand_order {
330            BinaryDotOperandOrder::Original => (inputs[0], inputs[1]),
331            BinaryDotOperandOrder::Swapped => (inputs[1], inputs[0]),
332        };
333        if !exact_dot_shapes(lhs.shape(), rhs.shape(), &plan.config) {
334            return None;
335        }
336        return Some(
337            session
338                .dot_general(lhs, rhs, plan.config)
339                .map_err(Error::Runtime),
340        );
341    }
342    None
343}
344
345fn requires_broadcast(inputs: &[&EagerTensor], subscripts: &EinsumSubscripts) -> bool {
346    let mut sizes = std::collections::HashMap::<u32, usize>::new();
347    for (tensor, labels) in inputs.iter().zip(&subscripts.inputs) {
348        for (&label, &size) in labels.iter().zip(tensor.shape()) {
349            if let Some(previous) = sizes.insert(label, size) {
350                if previous != size && (previous == 1 || size == 1) {
351                    return true;
352                }
353            }
354        }
355    }
356    false
357}
358
359fn exact_dot_shapes(
360    lhs_shape: &[usize],
361    rhs_shape: &[usize],
362    config: &tenferro_tensor::DotGeneralConfig,
363) -> bool {
364    config
365        .lhs_contracting_dims
366        .iter()
367        .zip(&config.rhs_contracting_dims)
368        .all(|(&lhs, &rhs)| lhs_shape[lhs] == rhs_shape[rhs])
369        && config
370            .lhs_batch_dims
371            .iter()
372            .zip(&config.rhs_batch_dims)
373            .all(|(&lhs, &rhs)| lhs_shape[lhs] == rhs_shape[rhs])
374}
375
376fn try_expand_eager_einsum(
377    session: &mut EagerSession<'_>,
378    inputs: &[&EagerTensor],
379    subscripts: &EinsumSubscripts,
380) -> Result<Option<EagerTensor>> {
381    if inputs.len() <= 1 {
382        return Ok(None);
383    }
384
385    let shapes: Vec<Vec<usize>> = inputs
386        .iter()
387        .map(|tensor| tensor.shape().to_vec())
388        .collect();
389    let shape_refs: Vec<&[usize]> = shapes.iter().map(Vec::as_slice).collect();
390    let subs = Subscripts::from(subscripts);
391    let plan_spec = EinsumPlanSpec::Auto(default_auto_options());
392
393    let program = cached_expanded_eager_program(
394        session,
395        subscripts,
396        &subs,
397        &plan_spec,
398        &shape_refs,
399        &shapes,
400    )?;
401    execute_eager_einsum_program_in_session(session, inputs, &program)
402}
403
404struct ExpandedEagerProgram {
405    compiled: CompiledProgram<StdTensorOp>,
406    input_slots: Vec<(usize, usize)>,
407}
408
409#[derive(Clone)]
410struct ExpandedEagerProgramCacheKeyData {
411    subscripts: EinsumSubscripts,
412    shapes: Vec<Vec<usize>>,
413    plan_spec: EinsumPlanSpec,
414}
415
416impl ExpandedEagerProgramCacheKeyData {
417    fn new(
418        subscripts: &EinsumSubscripts,
419        shapes: &[Vec<usize>],
420        plan_spec: &EinsumPlanSpec,
421    ) -> Self {
422        Self {
423            subscripts: subscripts.clone(),
424            shapes: shapes.to_vec(),
425            plan_spec: plan_spec.clone(),
426        }
427    }
428
429    fn matches_expanded_eager_program(
430        &self,
431        subscripts: &EinsumSubscripts,
432        shapes: &[Vec<usize>],
433        plan_spec: &EinsumPlanSpec,
434    ) -> bool {
435        self.subscripts == *subscripts
436            && self.shapes.as_slice() == shapes
437            && plan_specs_equal(&self.plan_spec, plan_spec)
438    }
439
440    fn retained_bytes(&self) -> usize {
441        saturating_sum([
442            crate::cache::einsum_subscripts_retained_bytes(&self.subscripts),
443            saturating_sum(self.shapes.iter().map(vec_retained_bytes)),
444            plan_spec_retained_bytes(&self.plan_spec),
445        ])
446    }
447}
448
449struct CachedExpandedEagerProgram {
450    key_data: ExpandedEagerProgramCacheKeyData,
451    program: Arc<ExpandedEagerProgram>,
452}
453
454fn cached_expanded_eager_program(
455    session: &mut EagerSession<'_>,
456    subscripts: &EinsumSubscripts,
457    subs: &Subscripts,
458    plan_spec: &EinsumPlanSpec,
459    shape_refs: &[&[usize]],
460    shapes: &[Vec<usize>],
461) -> Result<Arc<ExpandedEagerProgram>> {
462    session.with_extension_caches(|caches| {
463        let plan_hash = plan_spec_hash(plan_spec);
464        let key = expanded_eager_program_cache_key(subscripts, shapes, plan_hash);
465        if let Some(cached) = caches.get::<CachedExpandedEagerProgram>(&key) {
466            let key_data = &cached.key_data;
467            if key_data.matches_expanded_eager_program(subscripts, shapes, plan_spec) {
468                return Ok(Arc::clone(&cached.program));
469            }
470        }
471
472        let tree = resolve_plan_spec(plan_spec, subs, shape_refs)?;
473        let program = Arc::new(build_expanded_eager_program(&tree, shapes)?);
474        let key_data = ExpandedEagerProgramCacheKeyData::new(subscripts, shapes, plan_spec);
475        let retained_bytes = saturating_sum([
476            key_data.retained_bytes(),
477            expanded_eager_program_retained_bytes(&program),
478        ]);
479        caches.put(
480            key,
481            CachedExpandedEagerProgram {
482                key_data,
483                program: Arc::clone(&program),
484            },
485            retained_bytes,
486        );
487        Ok(program)
488    })?
489}
490
491fn expanded_eager_program_cache_key(
492    subscripts: &EinsumSubscripts,
493    shapes: &[Vec<usize>],
494    plan_hash: u64,
495) -> ExtensionCacheKey {
496    let mut hasher = DefaultHasher::new();
497    subscripts.hash(&mut hasher);
498    shapes.hash(&mut hasher);
499    plan_hash.hash(&mut hasher);
500    ExtensionCacheKey::new(
501        EINSUM_EXTENSION_FAMILY_ID,
502        EINSUM_EAGER_EXPANDED_PROGRAMS_CACHE,
503        hasher.finish(),
504    )
505}
506
507fn plan_spec_hash(plan_spec: &EinsumPlanSpec) -> u64 {
508    let mut hasher = DefaultHasher::new();
509    hash_einsum_plan_spec(plan_spec, &mut hasher);
510    hasher.finish()
511}
512
513fn plan_spec_retained_bytes(plan_spec: &EinsumPlanSpec) -> usize {
514    match plan_spec {
515        EinsumPlanSpec::Auto(options) => saturating_sum([
516            std::mem::size_of::<EinsumPlanSpec>(),
517            vec_retained_bytes(&options.betas),
518        ]),
519        EinsumPlanSpec::LeftToRight => std::mem::size_of::<EinsumPlanSpec>(),
520        EinsumPlanSpec::Path(path) | EinsumPlanSpec::FixedPairs(path) => saturating_sum([
521            std::mem::size_of::<EinsumPlanSpec>(),
522            vec_retained_bytes(path),
523        ]),
524    }
525}
526
527fn build_expanded_eager_program(
528    tree: &crate::ContractionTree,
529    shapes: &[Vec<usize>],
530) -> Result<ExpandedEagerProgram> {
531    let mut builder = GraphBuilder::<StdTensorOp>::new();
532    let mut input_vals = Vec::with_capacity(shapes.len());
533    for input_idx in 0..shapes.len() {
534        let local = builder.add_input(TensorInputKey::User {
535            id: input_idx as u64,
536        });
537        input_vals.push(ValueRef::Local(local));
538    }
539
540    let result_ref = build_einsum_graph(&mut builder, tree, &input_vals, shapes)?;
541    let ValueRef::Local(result_local) = result_ref else {
542        return Err(Error::Runtime(tenferro_runtime::Error::Internal(
543            "expanded eager einsum returned an external value".into(),
544        )));
545    };
546    builder.set_outputs(vec![result_local]);
547    let graph = Arc::new(builder.build());
548    let output_key = graph.values()[result_local].key.clone();
549    let view = resolve(vec![graph]);
550    let graph = materialize_merge(&view, &[output_key]);
551    let compiled = compile(&graph);
552    let input_slots = compiled
553        .input_slots
554        .iter()
555        .zip(graph.inputs.iter())
556        .map(|(&slot, key)| {
557            let ValueKey::Input(TensorInputKey::User { id }) = key else {
558                return Err(runtime_internal(format!(
559                    "expanded eager einsum saw unexpected input key: {key:?}"
560                )));
561            };
562            Ok((slot, *id as usize))
563        })
564        .collect::<Result<_>>()?;
565
566    Ok(ExpandedEagerProgram {
567        compiled,
568        input_slots,
569    })
570}
571
572fn execute_eager_einsum_program_in_session(
573    session: &mut EagerSession<'_>,
574    inputs: &[&EagerTensor],
575    program: &ExpandedEagerProgram,
576) -> Result<Option<EagerTensor>> {
577    let mut slots: Vec<Option<EagerTensor>> = vec![None; program.compiled.n_slots];
578    for &(slot, input_idx) in &program.input_slots {
579        let tensor = inputs.get(input_idx).ok_or_else(|| {
580            runtime_missing(format!(
581                "expanded eager einsum input {input_idx} is missing"
582            ))
583        })?;
584        slots[slot] = Some((*tensor).clone());
585    }
586
587    let mut instruction_idx = 0;
588    while instruction_idx < program.compiled.instructions.len() {
589        if let Some((output_slot, output)) = try_execute_eager_broadcast_multiply_pattern(
590            session,
591            &program.compiled.instructions,
592            instruction_idx,
593            &slots,
594            &program.compiled.output_slots,
595        )? {
596            slots[output_slot] = Some(output);
597            instruction_idx += 3;
598            continue;
599        }
600
601        let instr = &program.compiled.instructions[instruction_idx];
602        if instr.outputs.len() != 1 {
603            return Err(runtime_internal(format!(
604                "expanded eager einsum expected single-output op, got {} outputs",
605                instr.outputs.len()
606            )));
607        }
608        let input_refs: Vec<&EagerTensor> = instr
609            .inputs
610            .iter()
611            .map(|&slot| slot_tensor(&slots, slot))
612            .collect::<Result<_>>()?;
613        let output = session
614            .apply_standard_op(instr.operation.clone(), &input_refs)
615            .map_err(Error::Runtime)?;
616        slots[instr.outputs[0]] = Some(output);
617        instruction_idx += 1;
618    }
619
620    let [output_slot] = program.compiled.output_slots.as_slice() else {
621        return Err(runtime_internal(format!(
622            "expanded eager einsum expected one graph output, got {}",
623            program.compiled.output_slots.len()
624        )));
625    };
626    slots
627        .get_mut(*output_slot)
628        .and_then(Option::take)
629        .map(Some)
630        .ok_or_else(|| runtime_missing("expanded eager einsum output slot is missing"))
631}
632
633fn expanded_eager_program_retained_bytes(program: &ExpandedEagerProgram) -> usize {
634    saturating_sum([
635        size_of::<ExpandedEagerProgram>(),
636        vec_retained_bytes(&program.input_slots),
637        compiled_program_retained_bytes(&program.compiled),
638    ])
639}
640
641fn compiled_program_retained_bytes(program: &CompiledProgram<StdTensorOp>) -> usize {
642    saturating_sum([
643        size_of::<CompiledProgram<StdTensorOp>>(),
644        vec_retained_bytes(&program.instructions),
645        vec_retained_bytes(&program.input_slots),
646        vec_retained_bytes(&program.output_slots),
647        saturating_sum(program.instructions.iter().map(instruction_retained_bytes)),
648    ])
649}
650
651fn instruction_retained_bytes(instruction: &Instruction<StdTensorOp>) -> usize {
652    saturating_sum([
653        size_of::<Instruction<StdTensorOp>>(),
654        std_tensor_op_retained_bytes(&instruction.operation),
655        vec_retained_bytes(&instruction.inputs),
656        vec_retained_bytes(&instruction.outputs),
657    ])
658}
659
660fn std_tensor_op_retained_bytes(op: &StdTensorOp) -> usize {
661    match op {
662        // Inline axes are already counted in Instruction<StdTensorOp>.
663        StdTensorOp::DotGeneral { config } => saturating_sum(
664            [
665                &config.lhs_contracting_dims,
666                &config.rhs_contracting_dims,
667                &config.lhs_batch_dims,
668                &config.rhs_batch_dims,
669            ]
670            .into_iter()
671            .filter(|axes| axes.spilled())
672            .map(|axes| axes.capacity().saturating_mul(size_of::<usize>())),
673        ),
674        StdTensorOp::Transpose { perm } => vec_retained_bytes(perm),
675        StdTensorOp::Reshape { to_shape } => vec_retained_bytes(to_shape),
676        StdTensorOp::BroadcastInDim { shape, dims } => {
677            saturating_sum([vec_retained_bytes(shape), vec_retained_bytes(dims)])
678        }
679        StdTensorOp::Constant { bytes, .. } => vec_retained_bytes(bytes),
680        StdTensorOp::ReduceSum { axes }
681        | StdTensorOp::ReduceProd { axes }
682        | StdTensorOp::ReduceMax { axes }
683        | StdTensorOp::ReduceMin { axes }
684        | StdTensorOp::Reverse { axes } => vec_retained_bytes(axes),
685        StdTensorOp::DynamicSlice { slice_sizes } => vec_retained_bytes(slice_sizes),
686        StdTensorOp::GatherDynamicSliceSizes {
687            offset_dims,
688            collapsed_slice_dims,
689            start_index_map,
690            slice_sizes,
691            ..
692        } => saturating_sum([
693            vec_retained_bytes(offset_dims),
694            vec_retained_bytes(collapsed_slice_dims),
695            vec_retained_bytes(start_index_map),
696            vec_retained_bytes(slice_sizes),
697        ]),
698        _ => 0,
699    }
700}
701
702fn try_execute_eager_broadcast_multiply_pattern(
703    session: &mut EagerSession<'_>,
704    instructions: &[Instruction<StdTensorOp>],
705    instruction_idx: usize,
706    slots: &[Option<EagerTensor>],
707    output_slots: &[usize],
708) -> Result<Option<(usize, EagerTensor)>> {
709    if instruction_idx + 2 >= instructions.len() {
710        return Ok(None);
711    }
712    let lhs_bc = &instructions[instruction_idx];
713    let rhs_bc = &instructions[instruction_idx + 1];
714    let multiply = &instructions[instruction_idx + 2];
715
716    let StdTensorOp::BroadcastInDim {
717        shape: lhs_shape_exprs,
718        dims: lhs_dims,
719    } = &lhs_bc.operation
720    else {
721        return Ok(None);
722    };
723    let StdTensorOp::BroadcastInDim {
724        shape: rhs_shape_exprs,
725        dims: rhs_dims,
726    } = &rhs_bc.operation
727    else {
728        return Ok(None);
729    };
730    if !matches!(multiply.operation, StdTensorOp::Mul)
731        || lhs_bc.outputs.len() != 1
732        || rhs_bc.outputs.len() != 1
733        || multiply.outputs.len() != 1
734        || multiply.inputs.len() != 2
735        || lhs_bc.inputs.is_empty()
736        || rhs_bc.inputs.is_empty()
737        || multiply.inputs[0] != lhs_bc.outputs[0]
738        || multiply.inputs[1] != rhs_bc.outputs[0]
739    {
740        return Ok(None);
741    }
742
743    let lhs_bc_slot = lhs_bc.outputs[0];
744    let rhs_bc_slot = rhs_bc.outputs[0];
745    if output_slots.contains(&lhs_bc_slot)
746        || output_slots.contains(&rhs_bc_slot)
747        || instructions[instruction_idx + 3..]
748            .iter()
749            .any(|instr| instr.inputs.contains(&lhs_bc_slot) || instr.inputs.contains(&rhs_bc_slot))
750    {
751        return Ok(None);
752    }
753
754    let lhs = slot_tensor(slots, lhs_bc.inputs[0])?;
755    let rhs = slot_tensor(slots, rhs_bc.inputs[0])?;
756    let lhs_shape = eval_shape_exprs(slots, &lhs_bc.inputs, lhs_shape_exprs)?;
757    let rhs_shape = eval_shape_exprs(slots, &rhs_bc.inputs, rhs_shape_exprs)?;
758    let Some(output) = backend_broadcast_multiply_untracked(
759        session, lhs, &lhs_shape, lhs_dims, rhs, &rhs_shape, rhs_dims,
760    )?
761    else {
762        return Ok(None);
763    };
764
765    Ok(Some((multiply.outputs[0], output)))
766}
767
768#[allow(clippy::too_many_arguments)]
769fn backend_broadcast_multiply_untracked(
770    session: &mut EagerSession<'_>,
771    lhs: &EagerTensor,
772    lhs_shape: &[usize],
773    lhs_dims: &[usize],
774    rhs: &EagerTensor,
775    rhs_shape: &[usize],
776    rhs_dims: &[usize],
777) -> Result<Option<EagerTensor>> {
778    if !Arc::ptr_eq(lhs.runtime(), rhs.runtime()) {
779        return Err(tenferro_runtime::Error::ContextMismatch {
780            lhs: lhs.ctx_id(),
781            rhs: rhs.ctx_id(),
782        }
783        .into());
784    }
785    if lhs.tracks_grad() || rhs.tracks_grad() {
786        return Ok(None);
787    }
788
789    let runtime = lhs.runtime();
790    let value = session.backend_session().execute_broadcast_multiply_value(
791        lhs.tensor_read(),
792        lhs_shape,
793        lhs_dims,
794        rhs.tensor_read(),
795        rhs_shape,
796        rhs_dims,
797    )?;
798
799    Ok(value
800        .map(|value| adopt_untracked_eager_value(runtime.clone(), value))
801        .transpose()?)
802}
803
804fn eval_shape_exprs(
805    slots: &[Option<EagerTensor>],
806    input_slots: &[usize],
807    shape: &[DimExpr],
808) -> Result<Vec<usize>> {
809    let inputs = input_slots
810        .iter()
811        .map(|&slot| slot_tensor(slots, slot))
812        .collect::<Result<Vec<_>>>()?;
813    let input_shapes = inputs
814        .iter()
815        .map(|tensor| tensor.shape())
816        .collect::<Vec<_>>();
817    DimExpr::eval_all(shape, &input_shapes).map_err(|error| {
818        runtime_extension_error(
819            "einsum",
820            ErrorKind::Validation(ValidationKind::InvalidArgument),
821            error,
822        )
823    })
824}
825
826fn slot_tensor(slots: &[Option<EagerTensor>], slot: usize) -> Result<&EagerTensor> {
827    slots.get(slot).and_then(Option::as_ref).ok_or_else(|| {
828        Error::Runtime(tenferro_runtime::Error::MissingInput(format!(
829            "expanded eager einsum missing value for slot {slot}"
830        )))
831    })
832}
833
834fn infer_eager_output_shape(
835    subscripts: &EinsumSubscripts,
836    inputs: &[&EagerTensor],
837) -> Result<Vec<tenferro_runtime::SymDim>> {
838    if inputs.is_empty() {
839        return Err(Error::invalid_argument(
840            "einsum",
841            "inputs",
842            "einsum requires at least one input tensor",
843        ));
844    }
845    if subscripts.inputs.len() != inputs.len() {
846        return Err(Error::invalid_argument(
847            "einsum",
848            "inputs",
849            format!(
850                "einsum subscripts expect {} inputs, got {}",
851                subscripts.inputs.len(),
852                inputs.len()
853            ),
854        ));
855    }
856
857    let mut label_dims = std::collections::HashMap::new();
858    for (labels, tensor) in subscripts.inputs.iter().zip(inputs.iter()) {
859        let shape = tensor.shape();
860        if labels.len() != shape.len() {
861            return Err(Error::validation(
862                "einsum",
863                ValidationError::RankMismatch {
864                    expected: labels.len(),
865                    actual: shape.len(),
866                },
867            ));
868        }
869        for (&label, &dim) in labels.iter().zip(shape.iter()) {
870            if let Some(existing) = label_dims.get_mut(&label) {
871                if *existing != dim && *existing != 1 && dim != 1 {
872                    return Err(Error::validation(
873                        "einsum",
874                        ShapeMismatch::ExpectedActual {
875                            expected: tenferro_tensor::ShapeVec::from_vec(vec![*existing]),
876                            actual: tenferro_tensor::ShapeVec::from_vec(vec![dim]),
877                        }
878                        .into(),
879                    ));
880                }
881                if *existing == 1 {
882                    *existing = dim;
883                }
884            } else {
885                label_dims.insert(label, dim);
886            }
887        }
888    }
889
890    subscripts
891        .output
892        .iter()
893        .map(|label| {
894            label_dims
895                .get(label)
896                .copied()
897                .map(tenferro_runtime::SymDim::from)
898                .ok_or_else(|| {
899                    Error::invalid_argument(
900                        "einsum",
901                        "output",
902                        format!("einsum output label {label} is missing from input labels"),
903                    )
904                })
905        })
906        .collect()
907}
908
909fn runtime_extension_error<E>(op: &'static str, kind: ErrorKind, source: E) -> Error
910where
911    E: StdError + Send + Sync + 'static,
912{
913    Error::Runtime(tenferro_runtime::Error::extension(
914        op,
915        ErrorPhase::Execution,
916        EINSUM_EXTENSION_FAMILY_ID,
917        kind,
918        source,
919    ))
920}
921
922fn runtime_internal(message: impl Into<String>) -> Error {
923    Error::Runtime(tenferro_runtime::Error::Internal(message.into()))
924}
925
926fn runtime_missing(message: impl Into<String>) -> Error {
927    Error::Runtime(tenferro_runtime::Error::MissingInput(message.into()))
928}
929
930fn tensordot(
931    session: &mut EagerSession<'_>,
932    lhs: &EagerTensor,
933    rhs: &EagerTensor,
934    axes: TensorDotAxes<'_>,
935) -> Result<EagerTensor> {
936    let config = crate::tensordot::dot_general_config(axes, lhs.shape().len(), rhs.shape().len())?;
937    crate::tensordot::validate_concrete_contract_dims(lhs.shape(), rhs.shape(), &config)?;
938    session
939        .dot_general(lhs, rhs, config)
940        .map_err(Error::Runtime)
941}
942
943#[cfg(test)]
944mod tests;