Skip to main content

tenferro_einsum/
extension.rs

1use std::any::Any;
2use std::collections::hash_map::DefaultHasher;
3use std::collections::HashMap;
4#[cfg(feature = "autodiff")]
5use std::collections::HashSet;
6use std::hash::{Hash, Hasher};
7use std::sync::Arc;
8
9use computegraph::graph::GraphBuilder;
10use computegraph::types::ValueRef;
11#[cfg(feature = "autodiff")]
12use tenferro_ad::semantic_extension::{
13    AdValue, ResidualSpec, SemanticAdError, SemanticAdRuleRole, SemanticExtensionRegistryError,
14    SemanticExtensionRuleSet, SemanticLinearTransposeRequest, SemanticLinearTransposeRule,
15    SemanticLinearizeRequest, SemanticLinearizeResult, SemanticLinearizeRule,
16    SemanticPrimalVjpRequest, SemanticPrimalVjpRule,
17};
18use tenferro_extension_macros::define_extension_runtime;
19#[cfg(feature = "autodiff")]
20use tenferro_ops::dim_expr::DimExpr;
21use tenferro_ops::ext_op::{
22    ExtensionLoweringError, ExtensionLoweringResult, ExtensionOp, ExtensionStandardLowering,
23};
24use tenferro_ops::std_tensor_op::StdTensorOp;
25use tenferro_ops::sym_dim::SymDim;
26use tenferro_runtime::extension::{ExtensionCacheKey, ExtensionExecutionContext};
27#[cfg(feature = "autodiff")]
28use tenferro_runtime::program::{
29    CoreSemanticOp, ProgramValue, ProgramValueMetadata, SemanticProgramBuilder,
30};
31use tenferro_tensor::{BackendSession, DType, Error as TensorError, Tensor, TensorRead};
32
33use crate::builder::build_einsum_graph;
34use crate::cache::{
35    einsum_subscripts_retained_bytes, saturating_sum, vec_retained_bytes,
36    EINSUM_EXTENSION_FAMILY_ID, EINSUM_RUNTIME_PLANS_CACHE,
37};
38#[cfg(test)]
39use crate::optimize::default_auto_options;
40#[cfg(feature = "autodiff")]
41use crate::optimize::jax_path_to_v1_pairs;
42use crate::optimize::{hash_einsum_plan_spec, plan_specs_equal, resolve_plan_spec, EinsumPlanSpec};
43#[cfg(feature = "autodiff")]
44use crate::util::map_label_occurrences;
45use crate::{
46    ContractionTree, EinsumSubscripts, Error as EinsumError, Result as EinsumResult, Subscripts,
47};
48
49/// Standard einsum extension payload.
50///
51/// This mirrors the current `tenferro.einsum.v1` payload shape. Runtime-owned
52/// execution goes through [`EinsumRuntime`].
53#[derive(Clone)]
54pub(crate) struct EinsumExtensionOp {
55    subscripts: EinsumSubscripts,
56    plan_spec: EinsumPlanSpec,
57    output_shape_hint: Option<Vec<SymDim>>,
58    allow_broadcast: bool,
59}
60
61impl std::fmt::Debug for EinsumExtensionOp {
62    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63        f.debug_struct("EinsumExtensionOp")
64            .field("subscripts", &self.subscripts)
65            .field("plan_spec", &self.plan_spec)
66            .field("output_shape_hint", &self.output_shape_hint)
67            .field("allow_broadcast", &self.allow_broadcast)
68            .finish()
69    }
70}
71
72impl EinsumExtensionOp {
73    /// Create an einsum extension payload without a precomputed plan.
74    #[must_use]
75    #[cfg(test)]
76    pub(crate) fn new(subscripts: EinsumSubscripts) -> Self {
77        Self::with_plan_spec(subscripts, EinsumPlanSpec::Auto(default_auto_options()))
78    }
79
80    #[must_use]
81    pub(crate) fn with_plan_spec(subscripts: EinsumSubscripts, plan_spec: EinsumPlanSpec) -> Self {
82        Self {
83            subscripts,
84            plan_spec,
85            output_shape_hint: None,
86            allow_broadcast: false,
87        }
88    }
89
90    pub(crate) fn with_plan_spec_and_broadcast(
91        subscripts: EinsumSubscripts,
92        plan_spec: EinsumPlanSpec,
93        allow_broadcast: bool,
94    ) -> Self {
95        let mut op = Self::with_plan_spec(subscripts, plan_spec);
96        op.allow_broadcast = allow_broadcast;
97        op
98    }
99
100    /// Create an einsum extension payload with an explicit output shape hint.
101    #[must_use]
102    #[cfg(any(feature = "autodiff", test))]
103    pub(crate) fn with_output_shape_hint(
104        subscripts: EinsumSubscripts,
105        output_shape_hint: Vec<SymDim>,
106        plan_spec: EinsumPlanSpec,
107    ) -> Self {
108        let mut op = Self::with_plan_spec(subscripts, plan_spec);
109        op.output_shape_hint = Some(output_shape_hint);
110        op
111    }
112
113    #[must_use]
114    #[cfg(feature = "autodiff")]
115    pub(crate) fn with_output_shape_hint_and_broadcast(
116        subscripts: EinsumSubscripts,
117        output_shape_hint: Vec<SymDim>,
118        plan_spec: EinsumPlanSpec,
119        allow_broadcast: bool,
120    ) -> Self {
121        let mut op = Self::with_plan_spec_and_broadcast(subscripts, plan_spec, allow_broadcast);
122        op.output_shape_hint = Some(output_shape_hint);
123        op
124    }
125
126    /// Return the canonical subscripts.
127    #[must_use]
128    pub(crate) fn subscripts(&self) -> &EinsumSubscripts {
129        &self.subscripts
130    }
131
132    /// Return the shape-independent planning policy.
133    #[must_use]
134    pub(crate) fn plan_spec(&self) -> &EinsumPlanSpec {
135        &self.plan_spec
136    }
137
138    #[must_use]
139    #[cfg(feature = "autodiff")]
140    pub(crate) fn allow_broadcast(&self) -> bool {
141        self.allow_broadcast
142    }
143}
144
145impl ExtensionOp for EinsumExtensionOp {
146    fn family_id(&self) -> &'static str {
147        EINSUM_EXTENSION_FAMILY_ID
148    }
149
150    fn payload_hash(&self, hasher: &mut dyn Hasher) {
151        hasher.write_usize(self.subscripts.inputs.len());
152        for input in &self.subscripts.inputs {
153            hasher.write_usize(input.len());
154            for label in input {
155                hasher.write_u32(*label);
156            }
157        }
158        hasher.write_usize(self.subscripts.output.len());
159        for label in &self.subscripts.output {
160            hasher.write_u32(*label);
161        }
162        hash_einsum_plan_spec(self.plan_spec(), hasher);
163        hasher.write_u8(u8::from(self.allow_broadcast));
164        if let Some(shape) = &self.output_shape_hint {
165            hasher.write_usize(shape.len());
166            for dim in shape {
167                match dim.constant_value() {
168                    Some(value) => {
169                        hasher.write_u8(1);
170                        hasher.write_usize(value);
171                    }
172                    None => hasher.write_u8(0),
173                }
174            }
175        } else {
176            hasher.write_usize(usize::MAX);
177        }
178    }
179
180    fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
181        other.as_any().downcast_ref::<Self>().is_some_and(|that| {
182            self.subscripts == that.subscripts
183                && plan_specs_equal(self.plan_spec(), that.plan_spec())
184                && self.output_shape_hint == that.output_shape_hint
185                && self.allow_broadcast == that.allow_broadcast
186        })
187    }
188
189    fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
190        Arc::new(self.clone())
191    }
192
193    fn as_any(&self) -> &dyn Any {
194        self
195    }
196
197    fn input_count(&self) -> usize {
198        self.subscripts.inputs.len()
199    }
200
201    fn output_count(&self) -> usize {
202        1
203    }
204
205    fn semantic_effects(&self) -> tenferro_ops::ext_op::ExtensionEffectDeclaration<'_> {
206        tenferro_ops::ext_op::ExtensionEffectDeclaration::Declared(&[])
207    }
208
209    fn semantic_aliases(&self) -> tenferro_ops::ext_op::ExtensionAliasDeclaration<'_> {
210        tenferro_ops::ext_op::ExtensionAliasDeclaration::AllFresh
211    }
212
213    fn infer_output_meta(
214        &self,
215        ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
216    ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
217        let input_dtypes = (0..self.input_count())
218            .map(|input| ctx.input_dtype(input))
219            .collect::<Result<Vec<_>, _>>()?;
220        let input_shapes = (0..self.input_count())
221            .map(|input| ctx.input_shape(input).map(<[_]>::to_vec))
222            .collect::<Result<Vec<_>, _>>()?;
223
224        let mut label_dims: HashMap<u32, SymDim> = HashMap::new();
225        for (labels, shape) in self.subscripts.inputs.iter().zip(input_shapes.iter()) {
226            if labels.len() != shape.len() {
227                return Err(TensorError::rank_mismatch(
228                    "einsum",
229                    labels.len(),
230                    shape.len(),
231                ));
232            }
233            for (&label, dim) in labels.iter().zip(shape.iter()) {
234                if let Some(existing) = label_dims.get_mut(&label) {
235                    if !self.allow_broadcast {
236                        ctx.require_equal(existing.clone(), dim.clone())?;
237                        continue;
238                    }
239                    match (existing.constant_value(), dim.constant_value()) {
240                        (Some(lhs), Some(rhs)) if lhs == rhs || lhs == 1 || rhs == 1 => {
241                            *existing = SymDim::from(if lhs == 1 { rhs } else { lhs });
242                        }
243                        (Some(_lhs), Some(_rhs)) => {
244                            ctx.require_equal(existing.clone(), dim.clone())?;
245                        }
246                        (Some(1), None) => *existing = dim.clone(),
247                        (None, Some(1)) => {}
248                        _ => {}
249                    }
250                } else {
251                    label_dims.insert(label, dim.clone());
252                }
253            }
254        }
255
256        let output_shape = match &self.output_shape_hint {
257            Some(shape) if shape.iter().all(|dim| dim.constant_value().is_some()) => shape.clone(),
258            _ => self
259                .subscripts
260                .output
261                .iter()
262                .map(|label| label_dims.get(label).cloned())
263                .collect::<Option<Vec<_>>>()
264                .ok_or_else(|| {
265                    TensorError::invalid_argument(
266                        "einsum",
267                        "output labels",
268                        "must be present in input metadata",
269                    )
270                })?,
271        };
272        if output_shape.len() != self.subscripts.output.len() {
273            return Err(TensorError::rank_mismatch(
274                "einsum",
275                self.subscripts.output.len(),
276                output_shape.len(),
277            ));
278        }
279        if let Some(external) = input_dtypes
280            .iter()
281            .find(|dtype| matches!(dtype, DType::External(_)))
282        {
283            return Err(TensorError::unsupported_dtype(
284                "einsum",
285                *external,
286                "einsum takes preset scalars only; an externally defined scalar is not supported",
287            ));
288        }
289        Ok(vec![(
290            promote_dtypes(input_dtypes.iter().copied()),
291            output_shape,
292        )])
293    }
294
295    fn lower_to_standard_ops(
296        &self,
297        builder: &mut GraphBuilder<StdTensorOp>,
298        inputs: &[ValueRef<StdTensorOp>],
299        input_dtypes: &[DType],
300        input_shapes: &[&[SymDim]],
301    ) -> ExtensionLoweringResult {
302        if inputs.len() != self.input_count()
303            || input_dtypes.len() != self.input_count()
304            || input_shapes.len() != self.input_count()
305        {
306            return Err(ExtensionLoweringError::new(format!(
307                "einsum extension expects {} inputs, got values={}, dtypes={}, shapes={}",
308                self.input_count(),
309                inputs.len(),
310                input_dtypes.len(),
311                input_shapes.len()
312            )));
313        }
314
315        let Some(shapes) = concrete_sym_shape_slices(input_shapes) else {
316            return Ok(ExtensionStandardLowering::Unsupported);
317        };
318        let shape_refs: Vec<&[usize]> = shapes.iter().map(Vec::as_slice).collect();
319        let subs = Subscripts::from(&self.subscripts);
320        let tree = resolve_plan_spec(self.plan_spec(), &subs, &shape_refs).map_err(|source| {
321            ExtensionLoweringError::from_source_with_kind(source.kind(), source)
322        })?;
323        let output = build_einsum_graph(builder, &tree, inputs, &shapes).map_err(|source| {
324            ExtensionLoweringError::from_source_with_kind(source.kind(), source)
325        })?;
326        Ok(ExtensionStandardLowering::Lowered(vec![output]))
327    }
328}
329
330fn concrete_sym_shape_slices(input_shapes: &[&[SymDim]]) -> Option<Vec<Vec<usize>>> {
331    input_shapes
332        .iter()
333        .map(|shape| {
334            shape
335                .iter()
336                .map(SymDim::constant_value)
337                .collect::<Option<Vec<_>>>()
338        })
339        .collect()
340}
341
342/// Return the semantic-program einsum extension AD rules.
343#[cfg(feature = "autodiff")]
344///
345/// # Errors
346///
347/// Returns [`SemanticExtensionRegistryError::MalformedFamilyId`] for an
348/// invalid family identifier, or
349/// [`SemanticExtensionRegistryError::DuplicateRule`] for a duplicate role.
350pub fn semantic_ad_rules(
351) -> std::result::Result<SemanticExtensionRuleSet, SemanticExtensionRegistryError> {
352    SemanticExtensionRuleSet::new()
353        .with_linearize(Arc::new(EinsumAdRule))?
354        .with_linear_transpose(Arc::new(EinsumAdRule))?
355        .with_primal_vjp(Arc::new(EinsumAdRule))
356}
357
358#[derive(Debug)]
359#[cfg(feature = "autodiff")]
360struct EinsumAdRule;
361
362#[cfg(feature = "autodiff")]
363impl SemanticLinearizeRule for EinsumAdRule {
364    fn family_id(&self) -> &'static str {
365        EINSUM_EXTENSION_FAMILY_ID
366    }
367
368    fn linearize(
369        &self,
370        request: SemanticLinearizeRequest<'_>,
371        builder: &mut SemanticProgramBuilder,
372    ) -> std::result::Result<SemanticLinearizeResult, SemanticAdError> {
373        let op = semantic_einsum_payload(request.op(), SemanticAdRuleRole::Linearize)?;
374        if !request.active_outputs()[0] {
375            return Ok(SemanticLinearizeResult::new([AdValue::Absent], []));
376        }
377        let mut terms = Vec::new();
378        for (active_idx, tangent) in request.tangent_inputs().iter().copied().enumerate() {
379            let AdValue::Value(tangent) = tangent else {
380                continue;
381            };
382            let inputs: Vec<_> = request
383                .primal_inputs()
384                .iter()
385                .copied()
386                .enumerate()
387                .map(|(input_idx, primal)| {
388                    if input_idx == active_idx {
389                        tangent
390                    } else {
391                        primal
392                    }
393                })
394                .collect();
395            terms.push(builder.add_extension(Arc::new(op.clone()), &inputs)?[0]);
396        }
397        let tangent = semantic_sum_terms(builder, terms)?;
398        Ok(SemanticLinearizeResult::new([tangent], []))
399    }
400}
401
402#[cfg(feature = "autodiff")]
403impl SemanticLinearTransposeRule for EinsumAdRule {
404    fn family_id(&self) -> &'static str {
405        EINSUM_EXTENSION_FAMILY_ID
406    }
407
408    fn residual_mask(&self) -> ResidualSpec {
409        // The einsum VJP reads every non-active operand as a tensor operand
410        // (conjugate + contract with the cotangent); which operands are read
411        // depends on the active-input configuration, so all inputs are
412        // declared. The output is only used for its shape.
413        ResidualSpec::all_inputs()
414    }
415
416    fn linear_transpose(
417        &self,
418        request: SemanticLinearTransposeRequest<'_>,
419        builder: &mut SemanticProgramBuilder,
420    ) -> std::result::Result<Box<[AdValue]>, SemanticAdError> {
421        let primal_inputs = (0..request.primal_input_count())
422            .map(|index| request.primal_input_value(index))
423            .collect::<Result<Vec<_>, _>>()?;
424        let primal_output_metadata = request.primal_output_meta(0)?;
425        semantic_einsum_vjp(
426            request.op(),
427            &primal_inputs,
428            primal_output_metadata,
429            request.cotangent_outputs(),
430            request.active_inputs(),
431            request.residual_mask(),
432            builder,
433        )
434    }
435}
436
437#[cfg(feature = "autodiff")]
438impl SemanticPrimalVjpRule for EinsumAdRule {
439    fn family_id(&self) -> &'static str {
440        EINSUM_EXTENSION_FAMILY_ID
441    }
442
443    fn residual_mask(&self) -> ResidualSpec {
444        // Same operand set as the linear transpose: every non-active input is
445        // read as a tensor; the output only for its shape.
446        ResidualSpec::all_inputs()
447    }
448
449    fn primal_vjp(
450        &self,
451        request: SemanticPrimalVjpRequest<'_>,
452        builder: &mut SemanticProgramBuilder,
453    ) -> std::result::Result<Box<[AdValue]>, SemanticAdError> {
454        let primal_inputs = (0..request.primal_input_count())
455            .map(|index| request.primal_input_value(index))
456            .collect::<Result<Vec<_>, _>>()?;
457        let primal_output_metadata = request.primal_output_meta(0)?;
458        semantic_einsum_vjp(
459            request.op(),
460            &primal_inputs,
461            primal_output_metadata,
462            request.cotangent_outputs(),
463            request.active_inputs(),
464            request.residual_mask(),
465            builder,
466        )
467    }
468}
469
470#[cfg(feature = "autodiff")]
471fn semantic_einsum_vjp(
472    payload: &dyn ExtensionOp,
473    primal_inputs: &[ProgramValue],
474    primal_output_metadata: &ProgramValueMetadata,
475    cotangent_outputs: &[AdValue],
476    active_inputs: &[bool],
477    residual_mask: ResidualSpec,
478    builder: &mut SemanticProgramBuilder,
479) -> std::result::Result<Box<[AdValue]>, SemanticAdError> {
480    let op = semantic_einsum_payload(payload, SemanticAdRuleRole::LinearTranspose)?;
481    let input_count = op.subscripts.inputs.len();
482    let AdValue::Value(cotangent) = cotangent_outputs[0] else {
483        return Ok(vec![AdValue::Absent; input_count].into_boxed_slice());
484    };
485    let primal_input_shapes = primal_inputs
486        .iter()
487        .copied()
488        .map(|value| semantic_value_shape(builder, value))
489        .collect::<std::result::Result<Vec<_>, _>>()?;
490    let cotangent_shape = semantic_metadata_shape(primal_output_metadata)?;
491
492    let input_labels = &op.subscripts.inputs;
493    let output_labels = &op.subscripts.output;
494    let mut result = Vec::with_capacity(input_count);
495    for active_idx in 0..input_count {
496        if !active_inputs[active_idx] {
497            result.push(AdValue::Absent);
498            continue;
499        }
500        let mut available_labels: HashSet<u32> = output_labels.iter().copied().collect();
501        for (input_idx, labels) in input_labels.iter().enumerate() {
502            if input_idx != active_idx {
503                available_labels.extend(labels.iter().copied());
504            }
505        }
506        let vjp_output_labels: Vec<u32> = input_labels[active_idx]
507            .iter()
508            .copied()
509            .filter(|label| available_labels.contains(label))
510            .collect();
511        let mut vjp_input_labels = vec![output_labels.clone()];
512        let mut vjp_inputs = vec![cotangent];
513        let mut vjp_input_shapes = vec![cotangent_shape.clone()];
514        for input_idx in 0..input_count {
515            if input_idx == active_idx {
516                continue;
517            }
518            vjp_input_labels.push(input_labels[input_idx].clone());
519            vjp_input_shapes.push(primal_input_shapes[input_idx].clone());
520            debug_assert!(
521                residual_mask.declares_input(input_idx),
522                "einsum transpose read primal input {input_idx} as a tensor operand but the \
523                 residual mask does not declare it; declare it in the einsum rule's residual mask"
524            );
525            vjp_inputs.push(semantic_conjugate_if_complex(
526                builder,
527                primal_inputs[input_idx],
528            )?);
529        }
530        let vjp_op = semantic_vjp_einsum_op(
531            op,
532            active_idx,
533            EinsumSubscripts {
534                inputs: vjp_input_labels,
535                output: vjp_output_labels.clone(),
536            },
537            &vjp_input_shapes,
538        )?;
539        let mut input_cotangent = builder.add_extension(Arc::new(vjp_op), &vjp_inputs)?[0];
540        if vjp_output_labels != input_labels[active_idx] {
541            input_cotangent = semantic_broadcast_einsum_vjp(
542                builder,
543                input_cotangent,
544                &vjp_output_labels,
545                &input_labels[active_idx],
546                primal_input_shapes[active_idx].clone(),
547            )?;
548        }
549        result.push(AdValue::Value(input_cotangent));
550    }
551    Ok(result.into_boxed_slice())
552}
553
554#[cfg(feature = "autodiff")]
555fn semantic_vjp_einsum_op(
556    primal_op: &EinsumExtensionOp,
557    active_idx: usize,
558    subscripts: EinsumSubscripts,
559    input_shapes: &[Vec<DimExpr>],
560) -> std::result::Result<EinsumExtensionOp, SemanticAdError> {
561    let plan_spec =
562        vjp_plan_spec_for_active(primal_op.plan_spec(), primal_op.input_count(), active_idx)?;
563    let sym_shapes: Vec<Vec<SymDim>> = input_shapes
564        .iter()
565        .enumerate()
566        .map(|(input_idx, shape)| {
567            let tensor_id = u64::MAX - input_idx as u64;
568            shape
569                .iter()
570                .enumerate()
571                .map(|(axis, dim)| match dim {
572                    DimExpr::Const(value) => SymDim::from(*value),
573                    _ => SymDim::tensor_axis(tensor_id, axis),
574                })
575                .collect()
576        })
577        .collect();
578    if let Some(concrete_shapes) = concrete_sym_shapes(&sym_shapes) {
579        let shape_refs: Vec<&[usize]> = concrete_shapes.iter().map(Vec::as_slice).collect();
580        let raw_subscripts = Subscripts::from(&subscripts);
581        let _tree = resolve_plan_spec(&plan_spec, &raw_subscripts, &shape_refs)
582            .map_err(|source| semantic_einsum_unsupported(source.to_string()))?;
583    }
584    Ok(EinsumExtensionOp::with_plan_spec_and_broadcast(
585        subscripts,
586        plan_spec,
587        primal_op.allow_broadcast(),
588    ))
589}
590
591#[cfg(feature = "autodiff")]
592fn semantic_value_shape(
593    builder: &SemanticProgramBuilder,
594    value: ProgramValue,
595) -> std::result::Result<Vec<DimExpr>, SemanticAdError> {
596    semantic_metadata_shape(builder.value_metadata(value)?)
597}
598
599#[cfg(feature = "autodiff")]
600fn semantic_metadata_shape(
601    metadata: &ProgramValueMetadata,
602) -> std::result::Result<Vec<DimExpr>, SemanticAdError> {
603    metadata
604        .shape()
605        .iter()
606        .map(|extent| {
607            extent.bound_expr().cloned().ok_or_else(|| {
608                semantic_einsum_unsupported(
609                    "einsum semantic AD requires a symbolic expression for every extent",
610                )
611            })
612        })
613        .collect()
614}
615
616#[cfg(feature = "autodiff")]
617fn semantic_conjugate_if_complex(
618    builder: &mut SemanticProgramBuilder,
619    value: ProgramValue,
620) -> std::result::Result<ProgramValue, SemanticAdError> {
621    if matches!(
622        builder.value_metadata(value)?.dtype(),
623        DType::C32 | DType::C64
624    ) {
625        Ok(builder.add_op(CoreSemanticOp::Conj, &[value])?[0])
626    } else {
627        Ok(value)
628    }
629}
630
631#[cfg(feature = "autodiff")]
632fn semantic_broadcast_einsum_vjp(
633    builder: &mut SemanticProgramBuilder,
634    cotangent: ProgramValue,
635    cotangent_labels: &[u32],
636    input_labels: &[u32],
637    shape: Vec<DimExpr>,
638) -> std::result::Result<ProgramValue, SemanticAdError> {
639    let dims = map_label_occurrences(cotangent_labels, input_labels).ok_or_else(|| {
640        semantic_einsum_unsupported(format!(
641            "einsum VJP cannot remap labels {cotangent_labels:?} into {input_labels:?}"
642        ))
643    })?;
644    let broadcast =
645        builder.add_op(CoreSemanticOp::BroadcastInDim { shape, dims }, &[cotangent])?[0];
646    semantic_project_repeated_labels(builder, broadcast, input_labels)
647}
648
649#[cfg(feature = "autodiff")]
650fn semantic_project_repeated_labels(
651    builder: &mut SemanticProgramBuilder,
652    cotangent: ProgramValue,
653    labels: &[u32],
654) -> std::result::Result<ProgramValue, SemanticAdError> {
655    let mut result = cotangent;
656    let mut first_axis_by_label = HashMap::new();
657    for (axis_b, label) in labels.iter().copied().enumerate() {
658        let Some(&axis_a) = first_axis_by_label.get(&label) else {
659            first_axis_by_label.insert(label, axis_b);
660            continue;
661        };
662        let extracted =
663            builder.add_op(CoreSemanticOp::ExtractDiag { axis_a, axis_b }, &[result])?[0];
664        result = builder.add_op(CoreSemanticOp::EmbedDiag { axis_a, axis_b }, &[extracted])?[0];
665    }
666    Ok(result)
667}
668
669#[cfg(feature = "autodiff")]
670fn semantic_sum_terms(
671    builder: &mut SemanticProgramBuilder,
672    terms: Vec<ProgramValue>,
673) -> std::result::Result<AdValue, SemanticAdError> {
674    let mut terms = terms.into_iter();
675    let Some(mut sum) = terms.next() else {
676        return Ok(AdValue::Absent);
677    };
678    for term in terms {
679        sum = builder.add_op(CoreSemanticOp::Add, &[sum, term])?[0];
680    }
681    Ok(AdValue::Value(sum))
682}
683
684#[cfg(feature = "autodiff")]
685fn semantic_einsum_payload(
686    op: &dyn ExtensionOp,
687    role: SemanticAdRuleRole,
688) -> std::result::Result<&EinsumExtensionOp, SemanticAdError> {
689    op.as_any()
690        .downcast_ref::<EinsumExtensionOp>()
691        .ok_or_else(|| SemanticAdError::Unsupported {
692            family_id: EINSUM_EXTENSION_FAMILY_ID,
693            role,
694            message: "einsum semantic AD received an incompatible payload".into(),
695        })
696}
697
698#[cfg(feature = "autodiff")]
699fn semantic_einsum_unsupported(message: impl Into<String>) -> SemanticAdError {
700    SemanticAdError::Unsupported {
701        family_id: EINSUM_EXTENSION_FAMILY_ID,
702        role: SemanticAdRuleRole::LinearTranspose,
703        message: message.into(),
704    }
705}
706
707#[cfg(feature = "autodiff")]
708fn vjp_plan_spec_for_active(
709    primal_plan: &EinsumPlanSpec,
710    input_count: usize,
711    active_idx: usize,
712) -> std::result::Result<EinsumPlanSpec, SemanticAdError> {
713    if active_idx >= input_count {
714        return Err(semantic_einsum_unsupported(format!(
715            "einsum VJP active input {active_idx} is outside {input_count} inputs"
716        )));
717    }
718
719    match primal_plan {
720        EinsumPlanSpec::Auto(options) => Ok(EinsumPlanSpec::Auto(options.clone())),
721        EinsumPlanSpec::LeftToRight => Ok(EinsumPlanSpec::LeftToRight),
722        EinsumPlanSpec::Path(path) => {
723            let pairs = jax_path_to_v1_pairs(path, input_count).map_err(|err| {
724                semantic_einsum_unsupported(format!(
725                    "failed to inherit einsum Path plan for VJP active input {active_idx}: {err}"
726                ))
727            })?;
728            derive_vjp_fixed_pairs(&pairs, input_count, active_idx).map(EinsumPlanSpec::FixedPairs)
729        }
730        EinsumPlanSpec::FixedPairs(pairs) => {
731            derive_vjp_fixed_pairs(pairs, input_count, active_idx).map(EinsumPlanSpec::FixedPairs)
732        }
733    }
734}
735
736#[cfg(feature = "autodiff")]
737fn derive_vjp_fixed_pairs(
738    primal_pairs: &[(usize, usize)],
739    input_count: usize,
740    active_idx: usize,
741) -> std::result::Result<Vec<(usize, usize)>, SemanticAdError> {
742    if input_count == 0 {
743        return Err(semantic_einsum_unsupported(
744            "einsum VJP cannot derive a plan for zero primal inputs",
745        ));
746    }
747    if active_idx >= input_count {
748        return Err(semantic_einsum_unsupported(format!(
749            "einsum VJP active input {active_idx} is outside {input_count} inputs"
750        )));
751    }
752    let required_steps = input_count.saturating_sub(1);
753    if primal_pairs.len() != required_steps {
754        return Err(semantic_einsum_unsupported(format!(
755            "einsum VJP cannot inherit explicit plan for active input {active_idx}: \
756             expected {required_steps} primal steps for {input_count} inputs, got {}",
757            primal_pairs.len()
758        )));
759    }
760    if input_count == 1 {
761        return Ok(Vec::new());
762    }
763
764    let children = fixed_pair_children(primal_pairs, input_count, active_idx)?;
765    let mut primal_to_vjp = vec![None; input_count];
766    let mut next_vjp_input = 1;
767    for (input_idx, slot) in primal_to_vjp.iter_mut().enumerate() {
768        if input_idx != active_idx {
769            *slot = Some(next_vjp_input);
770            next_vjp_input += 1;
771        }
772    }
773
774    let root = input_count + primal_pairs.len() - 1;
775    let mut pairs = Vec::with_capacity(required_steps);
776    let final_id = emit_vjp_adjoint(
777        root,
778        0,
779        &children,
780        input_count,
781        active_idx,
782        &primal_to_vjp,
783        &mut pairs,
784    )?;
785    let expected_final = input_count + pairs.len() - 1;
786    if final_id != expected_final || pairs.len() != required_steps {
787        return Err(semantic_einsum_unsupported(format!(
788            "einsum VJP plan derivation for active input {active_idx} produced an invalid \
789             tree: final id {final_id}, expected {expected_final}, steps {}",
790            pairs.len()
791        )));
792    }
793    Ok(pairs)
794}
795
796#[cfg(feature = "autodiff")]
797fn fixed_pair_children(
798    pairs: &[(usize, usize)],
799    input_count: usize,
800    active_idx: usize,
801) -> std::result::Result<Vec<Option<(usize, usize)>>, SemanticAdError> {
802    let mut live = vec![false; input_count + pairs.len()];
803    for slot in live.iter_mut().take(input_count) {
804        *slot = true;
805    }
806    let mut children = vec![None; input_count + pairs.len()];
807
808    for (step_idx, &(left, right)) in pairs.iter().enumerate() {
809        let next_idx = input_count + step_idx;
810        if left == right {
811            return Err(invalid_vjp_plan_error(
812                active_idx,
813                format!("pair ({left}, {right}) references the same operand"),
814            ));
815        }
816        if left >= next_idx || right >= next_idx {
817            return Err(invalid_vjp_plan_error(
818                active_idx,
819                format!("pair ({left}, {right}) references a non-existent operand"),
820            ));
821        }
822        if !live[left] || !live[right] {
823            return Err(invalid_vjp_plan_error(
824                active_idx,
825                format!("pair ({left}, {right}) references an operand that is no longer live"),
826            ));
827        }
828
829        live[left] = false;
830        live[right] = false;
831        live[next_idx] = true;
832        children[next_idx] = Some((left, right));
833    }
834
835    let live_count = live.iter().filter(|&&is_live| is_live).count();
836    if live_count != 1 {
837        return Err(invalid_vjp_plan_error(
838            active_idx,
839            format!("explicit plan leaves {live_count} live operands"),
840        ));
841    }
842
843    Ok(children)
844}
845
846#[cfg(feature = "autodiff")]
847fn emit_vjp_adjoint(
848    node: usize,
849    cotangent_id: usize,
850    children: &[Option<(usize, usize)>],
851    input_count: usize,
852    active_idx: usize,
853    primal_to_vjp: &[Option<usize>],
854    pairs: &mut Vec<(usize, usize)>,
855) -> std::result::Result<usize, SemanticAdError> {
856    if node < input_count {
857        return if node == active_idx {
858            Ok(cotangent_id)
859        } else {
860            Err(invalid_vjp_plan_error(
861                active_idx,
862                format!("adjoint walk reached inactive leaf {node}"),
863            ))
864        };
865    }
866
867    let (left, right) = children.get(node).and_then(|child| *child).ok_or_else(|| {
868        invalid_vjp_plan_error(active_idx, format!("missing children for node {node}"))
869    })?;
870    let left_has_active = subtree_contains_active(left, children, input_count, active_idx)?;
871    let right_has_active = subtree_contains_active(right, children, input_count, active_idx)?;
872    match (left_has_active, right_has_active) {
873        (true, false) => {
874            let sibling_id = emit_vjp_subtree(
875                right,
876                children,
877                input_count,
878                active_idx,
879                primal_to_vjp,
880                pairs,
881            )?;
882            let next = push_vjp_pair(cotangent_id, sibling_id, input_count, pairs);
883            emit_vjp_adjoint(
884                left,
885                next,
886                children,
887                input_count,
888                active_idx,
889                primal_to_vjp,
890                pairs,
891            )
892        }
893        (false, true) => {
894            let sibling_id = emit_vjp_subtree(
895                left,
896                children,
897                input_count,
898                active_idx,
899                primal_to_vjp,
900                pairs,
901            )?;
902            let next = push_vjp_pair(cotangent_id, sibling_id, input_count, pairs);
903            emit_vjp_adjoint(
904                right,
905                next,
906                children,
907                input_count,
908                active_idx,
909                primal_to_vjp,
910                pairs,
911            )
912        }
913        (true, true) => Err(invalid_vjp_plan_error(
914            active_idx,
915            format!("both children of node {node} contain the active input"),
916        )),
917        (false, false) => Err(invalid_vjp_plan_error(
918            active_idx,
919            format!("neither child of node {node} contains the active input"),
920        )),
921    }
922}
923
924#[cfg(feature = "autodiff")]
925fn emit_vjp_subtree(
926    node: usize,
927    children: &[Option<(usize, usize)>],
928    input_count: usize,
929    active_idx: usize,
930    primal_to_vjp: &[Option<usize>],
931    pairs: &mut Vec<(usize, usize)>,
932) -> std::result::Result<usize, SemanticAdError> {
933    if node < input_count {
934        return primal_to_vjp[node].ok_or_else(|| {
935            invalid_vjp_plan_error(
936                active_idx,
937                format!("sibling subtree unexpectedly reached active leaf {node}"),
938            )
939        });
940    }
941
942    let (left, right) = children.get(node).and_then(|child| *child).ok_or_else(|| {
943        invalid_vjp_plan_error(active_idx, format!("missing children for node {node}"))
944    })?;
945    let left_id = emit_vjp_subtree(
946        left,
947        children,
948        input_count,
949        active_idx,
950        primal_to_vjp,
951        pairs,
952    )?;
953    let right_id = emit_vjp_subtree(
954        right,
955        children,
956        input_count,
957        active_idx,
958        primal_to_vjp,
959        pairs,
960    )?;
961    Ok(push_vjp_pair(left_id, right_id, input_count, pairs))
962}
963
964#[cfg(feature = "autodiff")]
965fn push_vjp_pair(
966    left: usize,
967    right: usize,
968    n_vjp_inputs: usize,
969    pairs: &mut Vec<(usize, usize)>,
970) -> usize {
971    pairs.push((left, right));
972    n_vjp_inputs + pairs.len() - 1
973}
974
975#[cfg(feature = "autodiff")]
976fn subtree_contains_active(
977    node: usize,
978    children: &[Option<(usize, usize)>],
979    input_count: usize,
980    active_idx: usize,
981) -> std::result::Result<bool, SemanticAdError> {
982    if node < input_count {
983        return Ok(node == active_idx);
984    }
985    let (left, right) = children.get(node).and_then(|child| *child).ok_or_else(|| {
986        invalid_vjp_plan_error(active_idx, format!("missing children for node {node}"))
987    })?;
988    Ok(
989        subtree_contains_active(left, children, input_count, active_idx)?
990            || subtree_contains_active(right, children, input_count, active_idx)?,
991    )
992}
993
994#[cfg(feature = "autodiff")]
995fn invalid_vjp_plan_error(active_idx: usize, reason: String) -> SemanticAdError {
996    semantic_einsum_unsupported(format!(
997        "einsum VJP cannot inherit explicit plan for active input {active_idx}: {reason}"
998    ))
999}
1000
1001#[cfg(feature = "autodiff")]
1002fn concrete_sym_shapes(shapes: &[Vec<SymDim>]) -> Option<Vec<Vec<usize>>> {
1003    shapes
1004        .iter()
1005        .map(|shape| shape.iter().map(SymDim::constant_value).collect())
1006        .collect()
1007}
1008
1009define_extension_runtime! {
1010    runtime = EinsumRuntime,
1011    family_id = EINSUM_EXTENSION_FAMILY_ID,
1012    op_type = EinsumExtensionOp,
1013    execute_in_session = execute_einsum_extension_reads_in_session,
1014    session_supported = einsum_session_supported,
1015}
1016
1017pub(crate) fn execute_einsum_extension_session_reads(
1018    op: &EinsumExtensionOp,
1019    inputs: &[TensorRead<'_>],
1020    ctx: &mut ExtensionExecutionContext<'_, dyn BackendSession + '_>,
1021) -> tenferro_tensor::Result<Vec<Tensor>> {
1022    if inputs.is_empty() {
1023        return Err(tenferro_tensor::Error::invalid_argument(
1024            "einsum_extension",
1025            "inputs",
1026            "einsum requires at least one input tensor",
1027        ));
1028    }
1029
1030    let shapes: Vec<Vec<usize>> = inputs.iter().map(|input| input.shape().to_vec()).collect();
1031    let shape_refs: Vec<&[usize]> = shapes.iter().map(Vec::as_slice).collect();
1032    let subs = Subscripts::from(op.subscripts());
1033    let tree = cached_runtime_tree(ctx, op.subscripts(), op.plan_spec(), &shapes, || {
1034        resolve_plan_spec(op.plan_spec(), &subs, &shape_refs)
1035    })?;
1036    let output = crate::eager::eager_einsum_exec_read(ctx.backend_mut(), inputs, &tree)?;
1037    Ok(vec![output])
1038}
1039
1040/// Adapter from the scheduler/`apply_eager` borrowed-session shape to the
1041/// eager session executor. Reuses the existing forward kernel; reimplementing
1042/// it here is explicitly out of scope for issue #1665.
1043fn execute_einsum_extension_reads_in_session(
1044    op: &EinsumExtensionOp,
1045    session: &mut dyn BackendSession,
1046    caches: &mut tenferro_runtime::ExtensionCacheStore,
1047    inputs: &[TensorRead<'_>],
1048) -> tenferro_tensor::Result<Vec<Tensor>> {
1049    let mut ctx = ExtensionExecutionContext::new(session, caches);
1050    execute_einsum_extension_session_reads(op, inputs, &mut ctx)
1051}
1052
1053fn einsum_session_supported<B: tenferro_tensor::TensorBackend + 'static>(
1054    _op: &EinsumExtensionOp,
1055) -> bool {
1056    // The session executor runs the same forward kernel on any backend session
1057    // below (the einsum execution only needs elementwise/dot session ops), so
1058    // CPU and CUDA sessions both qualify for `apply_eager`'s native path.
1059    let type_id = std::any::TypeId::of::<B>();
1060    type_id == std::any::TypeId::of::<tenferro_cpu::CpuBackend>() || {
1061        #[cfg(feature = "cuda")]
1062        {
1063            type_id == std::any::TypeId::of::<tenferro_gpu::cuda::CudaBackend>()
1064        }
1065        #[cfg(not(feature = "cuda"))]
1066        {
1067            false
1068        }
1069    }
1070}
1071
1072#[derive(Clone)]
1073struct RuntimeTreeCacheKeyData {
1074    subscripts: EinsumSubscripts,
1075    shapes: Vec<Vec<usize>>,
1076    plan_spec: EinsumPlanSpec,
1077}
1078
1079impl RuntimeTreeCacheKeyData {
1080    fn new(
1081        subscripts: &EinsumSubscripts,
1082        shapes: &[Vec<usize>],
1083        plan_spec: &EinsumPlanSpec,
1084    ) -> Self {
1085        Self {
1086            subscripts: subscripts.clone(),
1087            shapes: shapes.to_vec(),
1088            plan_spec: plan_spec.clone(),
1089        }
1090    }
1091
1092    fn matches_runtime_tree(
1093        &self,
1094        subscripts: &EinsumSubscripts,
1095        shapes: &[Vec<usize>],
1096        plan_spec: &EinsumPlanSpec,
1097    ) -> bool {
1098        self.subscripts == *subscripts
1099            && self.shapes.as_slice() == shapes
1100            && plan_specs_equal(&self.plan_spec, plan_spec)
1101    }
1102
1103    fn retained_bytes(&self) -> usize {
1104        saturating_sum([
1105            einsum_subscripts_retained_bytes(&self.subscripts),
1106            saturating_sum(self.shapes.iter().map(vec_retained_bytes)),
1107            plan_spec_retained_bytes(&self.plan_spec),
1108        ])
1109    }
1110}
1111
1112struct CachedRuntimeTree {
1113    key_data: RuntimeTreeCacheKeyData,
1114    tree: Arc<ContractionTree>,
1115}
1116
1117fn cached_runtime_tree<B: BackendSession + ?Sized>(
1118    ctx: &mut ExtensionExecutionContext<'_, B>,
1119    subscripts: &EinsumSubscripts,
1120    plan_spec: &EinsumPlanSpec,
1121    shapes: &[Vec<usize>],
1122    build: impl FnOnce() -> EinsumResult<ContractionTree>,
1123) -> tenferro_tensor::Result<Arc<ContractionTree>> {
1124    let plan_hash = plan_spec_hash(plan_spec);
1125    let key = ExtensionCacheKey::new(
1126        EINSUM_EXTENSION_FAMILY_ID,
1127        EINSUM_RUNTIME_PLANS_CACHE,
1128        runtime_tree_cache_discriminator(subscripts, shapes, plan_hash),
1129    );
1130    if let Some(cached) = ctx.caches_mut().get::<CachedRuntimeTree>(&key) {
1131        let key_data = &cached.key_data;
1132        if key_data.matches_runtime_tree(subscripts, shapes, plan_spec) {
1133            return Ok(Arc::clone(&cached.tree));
1134        }
1135    }
1136
1137    let tree = Arc::new(build().map_err(einsum_runtime_error)?);
1138    let key_data = RuntimeTreeCacheKeyData::new(subscripts, shapes, plan_spec);
1139    let retained_bytes = saturating_sum([
1140        key_data.retained_bytes(),
1141        tree.retained_bytes_for_cache_stats(),
1142    ]);
1143    ctx.caches_mut().put(
1144        key,
1145        CachedRuntimeTree {
1146            key_data,
1147            tree: Arc::clone(&tree),
1148        },
1149        retained_bytes,
1150    );
1151    Ok(tree)
1152}
1153
1154fn einsum_runtime_error(error: EinsumError) -> tenferro_tensor::Error {
1155    error.into_tensor_error("einsum_extension")
1156}
1157
1158fn runtime_tree_cache_discriminator(
1159    subscripts: &EinsumSubscripts,
1160    shapes: &[Vec<usize>],
1161    plan_hash: u64,
1162) -> u64 {
1163    let mut hasher = DefaultHasher::new();
1164    subscripts.hash(&mut hasher);
1165    shapes.hash(&mut hasher);
1166    plan_hash.hash(&mut hasher);
1167    hasher.finish()
1168}
1169
1170fn plan_spec_hash(plan_spec: &EinsumPlanSpec) -> u64 {
1171    let mut hasher = DefaultHasher::new();
1172    hash_einsum_plan_spec(plan_spec, &mut hasher);
1173    hasher.finish()
1174}
1175
1176fn plan_spec_retained_bytes(plan_spec: &EinsumPlanSpec) -> usize {
1177    match plan_spec {
1178        EinsumPlanSpec::Auto(options) => saturating_sum([
1179            std::mem::size_of::<EinsumPlanSpec>(),
1180            vec_retained_bytes(&options.betas),
1181        ]),
1182        EinsumPlanSpec::LeftToRight => std::mem::size_of::<EinsumPlanSpec>(),
1183        EinsumPlanSpec::Path(path) | EinsumPlanSpec::FixedPairs(path) => saturating_sum([
1184            std::mem::size_of::<EinsumPlanSpec>(),
1185            vec_retained_bytes(path),
1186        ]),
1187    }
1188}
1189
1190fn promote_dtypes(dtypes: impl IntoIterator<Item = DType>) -> DType {
1191    dtypes
1192        .into_iter()
1193        .reduce(tenferro_tensor::validate::promote_dtype)
1194        .unwrap_or(DType::F64)
1195}
1196
1197#[cfg(test)]
1198mod tests;