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#[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 #[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 #[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 #[must_use]
128 pub(crate) fn subscripts(&self) -> &EinsumSubscripts {
129 &self.subscripts
130 }
131
132 #[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#[cfg(feature = "autodiff")]
344pub 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 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 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
1040fn 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 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;