1use 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#[cfg_attr(docsrs, doc(cfg(feature = "autodiff")))]
71pub trait EagerSessionEinsumExt {
72 fn einsum(&mut self, inputs: &[&EagerTensor], subscripts: &str) -> Result<EagerTensor>;
94
95 fn einsum_notation(
117 &mut self,
118 inputs: &[&EagerTensor],
119 notation: &EinsumNotation,
120 ) -> Result<EagerTensor>;
121
122 fn einsum_subscripts(
145 &mut self,
146 inputs: &[&EagerTensor],
147 subscripts: &EinsumSubscripts,
148 ) -> Result<EagerTensor>;
149
150 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, ¬ation)
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(¬ation.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 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;