tenferro_runtime/traced.rs
1use std::collections::HashMap;
2use std::error::Error as StdError;
3use std::fmt;
4use std::sync::atomic::{AtomicU64, Ordering};
5use std::sync::Arc;
6
7use computegraph::graph::{Graph, GraphBuilder};
8use computegraph::types::{OperationRole, ValueKey, ValueRef};
9use computegraph::LocalValueId;
10use num_complex::{Complex32, Complex64};
11use tenferro_ops::ad::context::GlobalMetadataScope;
12use tenferro_ops::broadcast::{
13 broadcast_error_to_validation, broadcast_in_dim_extent_error, broadcast_input_plan,
14 broadcast_shape, broadcast_shapes, BroadcastError,
15};
16use tenferro_ops::dim_expr::DimExpr;
17use tenferro_ops::input_key::TensorInputKey;
18use tenferro_ops::std_tensor_op::StdTensorOp;
19use tenferro_ops::TensorMeta;
20use tenferro_tensor::{
21 CompareDir, DType, DotGeneralConfig, Error as TensorError, GatherConfig, IntoShapeVec,
22 PadConfig, ScatterConfig, ShapeMismatch, SliceConfig, Tensor, TensorScalar, TensorValue,
23 ValidationError,
24};
25
26use super::error::{Error, ErrorPhase, Result};
27use super::sym_dim::SymDim;
28use crate::checkpoint::{CheckpointNode, RetainedInputMap, RetainedValue};
29use crate::metadata::{
30 concrete_tensor_meta, register_scoped_graph_metadata, register_scoped_value_metadata,
31 symbolic_input_meta, tensor_meta, MetadataScopeChain,
32};
33use crate::scalar_semantics::{bool_from_real_for_op, round_real_to_i32_for_op, round_real_to_i64};
34use crate::shape_constraint::ConstraintScopeChain;
35
36static NEXT_INPUT_ID: AtomicU64 = AtomicU64::new(0);
37static NEXT_TRACED_ID: AtomicU64 = AtomicU64::new(0);
38
39pub type TracedTensorId = u64;
40
41pub(crate) fn next_input_key() -> TensorInputKey {
42 TensorInputKey::User {
43 id: NEXT_INPUT_ID.fetch_add(1, Ordering::Relaxed),
44 }
45}
46
47pub(crate) fn next_traced_id() -> TracedTensorId {
48 NEXT_TRACED_ID.fetch_add(1, Ordering::Relaxed)
49}
50
51type TracedInputMap = RetainedInputMap;
52
53#[derive(Clone)]
54pub struct TracedTensor {
55 pub id: TracedTensorId,
56 pub rank: usize,
57 pub dtype: DType,
58 pub(crate) graph: Arc<Graph<StdTensorOp>>,
59 pub val: LocalValueId,
60 pub(crate) data: Option<Arc<RetainedValue>>,
61 pub(crate) shape_hint: Option<Vec<SymDim>>,
62 pub(crate) inputs_map: Arc<TracedInputMap>,
63 /// Retained construction-time metadata per bound leaf input key.
64 ///
65 /// Populated by data-attached leaf constructors with the exact
66 /// `TensorMeta` they register (symbolic `tensor_axis` extents for
67 /// [`TracedTensor::from_shared_tensor_value_symbolic_shape`]), so the
68 /// deferred eager-AD analysis can seed symbolic leaf metadata without
69 /// deriving concrete extents from the bound values.
70 pub(crate) leaf_metas: Arc<HashMap<TensorInputKey, TensorMeta>>,
71 pub(crate) extra_roots: Vec<Arc<Graph<StdTensorOp>>>,
72 pub(crate) checkpoint_chain: Option<Arc<CheckpointNode>>,
73 pub(crate) metadata_scopes: MetadataScopeChain,
74 pub(crate) constraint_scopes: ConstraintScopeChain,
75}
76
77impl fmt::Debug for TracedTensor {
78 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
79 f.debug_struct("TracedTensor")
80 .field("id", &self.id)
81 .field("rank", &self.rank)
82 .field("dtype", &self.dtype)
83 .field("val", &self.val)
84 .field("shape_hint", &self.shape_hint)
85 .field("has_data", &self.data.is_some())
86 .finish_non_exhaustive()
87 }
88}
89
90pub(crate) fn merge_traced_inputs_map<'a>(
91 inputs: impl IntoIterator<Item = &'a TracedTensor>,
92) -> Arc<TracedInputMap> {
93 let maps: Vec<_> = inputs
94 .into_iter()
95 .map(|input| &input.inputs_map)
96 .filter(|map| !map.is_empty())
97 .collect();
98 match maps.as_slice() {
99 [] => return Arc::new(HashMap::new()),
100 [single] => return Arc::clone(*single),
101 _ => {}
102 }
103
104 for &candidate in &maps {
105 if input_map_matches_ordered_merge(candidate.as_ref(), &maps) {
106 return Arc::clone(candidate);
107 }
108 }
109
110 let mut merged = (**maps[0]).clone();
111 for map in maps.iter().skip(1) {
112 merged.extend(
113 map.iter()
114 .map(|(key, tensor)| (key.clone(), tensor.clone())),
115 );
116 }
117 Arc::new(merged)
118}
119
120fn input_map_matches_ordered_merge(
121 candidate: &TracedInputMap,
122 maps: &[&Arc<TracedInputMap>],
123) -> bool {
124 merged_map_matches_ordered(candidate, maps, Arc::ptr_eq)
125}
126
127/// Merge the retained leaf-metadata maps of several traced tensors into one.
128///
129/// Sibling of [`merge_traced_inputs_map`]: the merged map covers exactly the
130/// same leaf keys as the merged bindings map, mapping each to the
131/// construction-time `TensorMeta` its leaf registered (symbolic for symbolic
132/// leaves).
133pub(crate) fn merge_traced_leaf_metas<'a>(
134 inputs: impl IntoIterator<Item = &'a TracedTensor>,
135) -> Arc<HashMap<TensorInputKey, TensorMeta>> {
136 let maps: Vec<_> = inputs
137 .into_iter()
138 .map(|input| &input.leaf_metas)
139 .filter(|map| !map.is_empty())
140 .collect();
141 match maps.as_slice() {
142 [] => return Arc::new(HashMap::new()),
143 [single] => return Arc::clone(*single),
144 _ => {}
145 }
146
147 for &candidate in &maps {
148 if merged_map_matches_ordered(candidate.as_ref(), &maps, |a, b| a == b) {
149 return Arc::clone(candidate);
150 }
151 }
152
153 let mut merged = (**maps[0]).clone();
154 for map in maps.iter().skip(1) {
155 merged.extend(map.iter().map(|(key, meta)| (key.clone(), meta.clone())));
156 }
157 Arc::new(merged)
158}
159
160fn merged_map_matches_ordered<V>(
161 candidate: &HashMap<TensorInputKey, V>,
162 maps: &[&Arc<HashMap<TensorInputKey, V>>],
163 matches: impl Fn(&V, &V) -> bool,
164) -> bool {
165 for map in maps {
166 for key in map.keys() {
167 let Some(final_value) = maps.iter().rev().find_map(|source| source.get(key)) else {
168 return false;
169 };
170 let Some(candidate_value) = candidate.get(key) else {
171 return false;
172 };
173 if !matches(candidate_value, final_value) {
174 return false;
175 }
176 }
177 }
178 true
179}
180
181pub(crate) fn try_concrete_shape(tensor: &TracedTensor) -> Option<Vec<usize>> {
182 tensor
183 .shape_hint
184 .as_ref()?
185 .iter()
186 .map(SymDim::constant_value)
187 .collect()
188}
189
190fn graph_validation(op: &'static str, source: impl Into<ValidationError>) -> Error {
191 Error::validation(op, ErrorPhase::GraphBuild, source.into())
192}
193
194fn graph_invalid_argument(
195 op: &'static str,
196 argument: &'static str,
197 message: impl Into<String>,
198) -> Error {
199 graph_validation(
200 op,
201 ValidationError::InvalidArgument {
202 argument,
203 message: message.into(),
204 },
205 )
206}
207
208fn graph_broadcast_error(op: &'static str, error: BroadcastError) -> Error {
209 graph_validation(op, broadcast_error_to_validation(error))
210}
211
212fn graph_tensor_error(op: &'static str, error: TensorError) -> Error {
213 match error {
214 TensorError::Validation { source, .. } => graph_validation(op, source),
215 other => Error::TensorRuntime(other),
216 }
217}
218
219fn graph_error_with_context(op: &'static str, error: Error) -> Error {
220 match error {
221 Error::Validation { source, .. } => graph_validation(op, source),
222 other => other,
223 }
224}
225
226pub(crate) fn concrete_shape(tensor: &TracedTensor) -> Result<Vec<usize>> {
227 tensor
228 .shape_hint
229 .as_ref()
230 .ok_or_else(|| {
231 graph_invalid_argument(
232 "TracedTensor::concrete_shape",
233 "shape",
234 format!("missing shape hint for traced tensor {}", tensor.id),
235 )
236 })?
237 .iter()
238 .map(|dim| {
239 dim.constant_value().ok_or_else(|| {
240 graph_invalid_argument(
241 "TracedTensor::concrete_shape",
242 "shape",
243 format!("symbolic dimension in shape hint for tensor {}", tensor.id),
244 )
245 })
246 })
247 .collect()
248}
249
250/// Broadcast a traced tensor to `target_shape`.
251///
252/// Expanding singleton axes are first reshaped away so the existing
253/// `BroadcastInDim` transpose rule reduces them correctly during VJP.
254pub(crate) fn broadcast_to(tensor: &TracedTensor, target_shape: &[usize]) -> Result<TracedTensor> {
255 let tensor_shape = concrete_shape(tensor)?;
256 if tensor_shape == target_shape {
257 return Ok(tensor.clone());
258 }
259
260 let plan = broadcast_input_plan(&tensor_shape, target_shape)
261 .map_err(|err| graph_broadcast_error("broadcast_to", err))?;
262
263 let source = if plan.source_shape == tensor_shape {
264 tensor.clone()
265 } else {
266 tensor.reshape(&plan.source_shape)?
267 };
268 source.broadcast_in_dim(target_shape, &plan.dims)
269}
270
271/// Broadcast two tensors to a common shape.
272pub(crate) fn broadcast_binary(
273 a: &TracedTensor,
274 b: &TracedTensor,
275) -> Result<(TracedTensor, TracedTensor)> {
276 if a.shape_hint == b.shape_hint && a.rank == b.rank {
277 return Ok((a.clone(), b.clone()));
278 }
279 if (try_concrete_shape(a).is_none() || try_concrete_shape(b).is_none()) && a.rank == b.rank {
280 return Ok((a.clone(), b.clone()));
281 }
282 let a_shape = concrete_shape(a)?;
283 let b_shape = concrete_shape(b)?;
284 let target = broadcast_shape(&a_shape, &b_shape)
285 .map_err(|err| graph_broadcast_error("broadcast_binary", err))?;
286 Ok((broadcast_to(a, &target)?, broadcast_to(b, &target)?))
287}
288
289pub(crate) fn broadcast_ternary(
290 a: &TracedTensor,
291 b: &TracedTensor,
292 c: &TracedTensor,
293) -> Result<(TracedTensor, TracedTensor, TracedTensor)> {
294 let a_shape = concrete_shape(a)?;
295 let b_shape = concrete_shape(b)?;
296 let c_shape = concrete_shape(c)?;
297 let target = broadcast_shapes([a_shape.as_slice(), b_shape.as_slice(), c_shape.as_slice()])
298 .map_err(|err| graph_broadcast_error("broadcast_ternary", err))?;
299 Ok((
300 broadcast_to(a, &target)?,
301 broadcast_to(b, &target)?,
302 broadcast_to(c, &target)?,
303 ))
304}
305
306fn scale_with_constant(input: &TracedTensor, op: StdTensorOp) -> Result<TracedTensor> {
307 let scalar = apply_nullary(op, 0, input.dtype, Some(vec![]))?;
308 apply_binary(
309 StdTensorOp::Mul,
310 input,
311 &scalar,
312 input.rank,
313 input.shape_hint.clone(),
314 )
315}
316
317fn try_inferred_output_dtype(
318 op: &StdTensorOp,
319 inputs: &[DType],
320 context: &'static str,
321) -> Result<DType> {
322 crate::shape_infer::infer_output_dtype_at(op, inputs, ErrorPhase::GraphBuild)
323 .map_err(|err| graph_error_with_context(context, err))
324}
325
326fn checked_shape_product_for_graph_build(
327 shape: &[usize],
328 context: &'static str,
329 _role: &'static str,
330) -> Result<usize> {
331 shape.iter().copied().try_fold(1usize, |acc, dim| {
332 acc.checked_mul(dim)
333 .ok_or_else(|| graph_validation(context, ValidationError::IntegerOverflow))
334 })
335}
336
337fn validate_concrete_reshape_shape(input: &TracedTensor, shape: &[usize]) -> Result<()> {
338 let to = checked_shape_product_for_graph_build(shape, "TracedTensor::reshape", "target")?;
339 let Some(input_shape) = try_concrete_shape(input) else {
340 return Ok(());
341 };
342 let from =
343 checked_shape_product_for_graph_build(&input_shape, "TracedTensor::reshape", "input")?;
344 if from != to {
345 return Err(graph_validation(
346 "TracedTensor::reshape",
347 ShapeMismatch::ReshapeElementCount { from, to },
348 ));
349 }
350 Ok(())
351}
352
353fn traced_input_shape_exprs(input_idx: usize, tensor: &TracedTensor) -> Vec<DimExpr> {
354 match tensor.shape_hint.as_ref() {
355 Some(shape) => shape
356 .iter()
357 .enumerate()
358 .map(|(axis, dim)| {
359 dim.constant_value()
360 .map_or(DimExpr::InputDim { input_idx, axis }, DimExpr::Const)
361 })
362 .collect(),
363 None => (0..tensor.rank)
364 .map(|axis| DimExpr::InputDim { input_idx, axis })
365 .collect(),
366 }
367}
368
369fn traced_input_sym_shape(tensor: &TracedTensor) -> Vec<SymDim> {
370 tensor.shape_hint.clone().unwrap_or_else(|| {
371 (0..tensor.rank)
372 .map(|axis| SymDim::tensor_axis(tensor.id, axis))
373 .collect()
374 })
375}
376
377pub(crate) fn infer_traced_single_output_shape(
378 op_name: &'static str,
379 op: &StdTensorOp,
380 inputs: &[&TracedTensor],
381) -> Result<(usize, Option<Vec<SymDim>>)> {
382 let input_shape_exprs: Vec<Vec<DimExpr>> = inputs
383 .iter()
384 .enumerate()
385 .map(|(input_idx, tensor)| traced_input_shape_exprs(input_idx, tensor))
386 .collect();
387 let input_shape_refs: Vec<&[DimExpr]> = input_shape_exprs.iter().map(Vec::as_slice).collect();
388 let output_shapes = crate::shape_infer::infer_output_shapes(op, &input_shape_refs)
389 .map_err(|err| graph_error_with_context(op_name, err))?;
390 let output_shape = output_shapes.first().ok_or_else(|| {
391 Error::Internal(format!("{op_name}: shape inference returned no outputs"))
392 })?;
393 if output_shapes.len() != 1 {
394 return Err(Error::Internal(format!(
395 "{op_name}: expected single-output shape inference, got {} outputs",
396 output_shapes.len()
397 )));
398 }
399
400 let input_sym_shapes: Vec<Vec<SymDim>> = inputs
401 .iter()
402 .map(|tensor| traced_input_sym_shape(tensor))
403 .collect();
404 let input_sym_refs: Vec<&[SymDim]> = input_sym_shapes.iter().map(Vec::as_slice).collect();
405 let out_shape_hint = output_shape
406 .iter()
407 .map(|dim| SymDim::from_dim_expr(dim, &input_sym_refs))
408 .collect();
409 Ok((output_shape.len(), Some(out_shape_hint)))
410}
411
412pub(crate) fn register_metadata_or_runtime_state<E>(
413 result: std::result::Result<GlobalMetadataScope, E>,
414) -> Result<GlobalMetadataScope>
415where
416 E: StdError + Send + Sync + 'static,
417{
418 result.map_err(|err| Error::runtime_state_source("metadata", ErrorPhase::Compile, err))
419}
420
421fn reduction_output_meta(
422 tensor: &TracedTensor,
423 axes: &[usize],
424 op: &'static str,
425) -> Result<(usize, Option<Vec<SymDim>>)> {
426 let mut seen = vec![false; tensor.rank];
427 for &axis in axes {
428 if axis >= tensor.rank {
429 return Err(graph_validation(
430 op,
431 ValidationError::AxisOutOfBounds {
432 axis,
433 rank: tensor.rank,
434 },
435 ));
436 }
437 if seen[axis] {
438 return Err(graph_validation(
439 op,
440 ValidationError::DuplicateAxis {
441 axis,
442 role: "reduction",
443 },
444 ));
445 }
446 seen[axis] = true;
447 }
448
449 let out_shape_hint = tensor.shape_hint.as_ref().map(|shape| {
450 (0..shape.len())
451 .filter(|d| !axes.contains(d))
452 .map(|d| shape[d].clone())
453 .collect()
454 });
455 Ok((tensor.rank - axes.len(), out_shape_hint))
456}
457
458fn validate_traced_axis(tensor: &TracedTensor, axis: usize, op: &'static str) -> Result<()> {
459 if axis >= tensor.rank {
460 return Err(graph_validation(
461 op,
462 ValidationError::AxisOutOfBounds {
463 axis,
464 rank: tensor.rank,
465 },
466 ));
467 }
468 Ok(())
469}
470
471fn validate_traced_axes(rank: usize, axes: &[usize], op: &'static str) -> Result<()> {
472 let mut seen = vec![false; rank];
473 for &axis in axes {
474 if axis >= rank {
475 return Err(graph_validation(
476 op,
477 ValidationError::AxisOutOfBounds { axis, rank },
478 ));
479 }
480 if seen[axis] {
481 return Err(graph_validation(
482 op,
483 ValidationError::DuplicateAxis { axis, role: "axis" },
484 ));
485 }
486 seen[axis] = true;
487 }
488 Ok(())
489}
490
491fn validate_traced_insert_axis(rank: usize, axis: usize, op: &'static str) -> Result<()> {
492 if axis > rank {
493 return Err(graph_invalid_argument(
494 op,
495 "axis",
496 format!("axis {axis} out of bounds for rank {rank} insertion"),
497 ));
498 }
499 Ok(())
500}
501
502fn validate_traced_perm(rank: usize, perm: &[usize], op: &'static str) -> Result<()> {
503 if perm.len() != rank {
504 return Err(graph_validation(
505 op,
506 ValidationError::InvalidPermutationLength {
507 expected: rank,
508 actual: perm.len(),
509 },
510 ));
511 }
512 let mut seen = vec![false; rank];
513 for &axis in perm {
514 if axis >= rank {
515 return Err(graph_validation(
516 op,
517 ValidationError::AxisOutOfBounds { axis, rank },
518 ));
519 }
520 if seen[axis] {
521 return Err(graph_validation(
522 op,
523 ValidationError::DuplicateAxis {
524 axis,
525 role: "permutation",
526 },
527 ));
528 }
529 seen[axis] = true;
530 }
531 Ok(())
532}
533
534fn validate_broadcast_in_dim_args(
535 input: &TracedTensor,
536 output_shape: &[SymDim],
537 dims: &[usize],
538 op: &'static str,
539) -> Result<()> {
540 if dims.len() != input.rank {
541 return Err(graph_validation(
542 op,
543 ValidationError::RankMismatch {
544 expected: input.rank,
545 actual: dims.len(),
546 },
547 ));
548 }
549
550 let concrete_input_shape: Option<Vec<usize>> = input
551 .shape_hint
552 .as_ref()
553 .and_then(|shape| shape.iter().map(SymDim::constant_value).collect());
554 let concrete_output_shape = output_shape
555 .iter()
556 .map(SymDim::constant_value)
557 .collect::<Option<Vec<_>>>();
558 if let (Some(input_shape), Some(output_shape)) = (
559 concrete_input_shape.as_deref(),
560 concrete_output_shape.as_deref(),
561 ) && let Some(error) = broadcast_in_dim_extent_error(input_shape, output_shape, dims)
562 {
563 return Err(graph_broadcast_error(op, error));
564 }
565
566 let mut seen = vec![false; output_shape.len()];
567 for &dim in dims {
568 if dim >= output_shape.len() {
569 return Err(graph_validation(
570 op,
571 ValidationError::AxisOutOfBounds {
572 axis: dim,
573 rank: output_shape.len(),
574 },
575 ));
576 }
577 if seen[dim] {
578 return Err(graph_validation(
579 op,
580 ValidationError::DuplicateAxis {
581 axis: dim,
582 role: "broadcast",
583 },
584 ));
585 }
586 seen[dim] = true;
587 }
588
589 Ok(())
590}
591
592impl std::ops::Add for &TracedTensor {
593 type Output = Result<TracedTensor>;
594
595 fn add(self, rhs: &TracedTensor) -> Result<TracedTensor> {
596 TracedTensor::add(self, rhs)
597 }
598}
599
600impl std::ops::Sub for &TracedTensor {
601 type Output = Result<TracedTensor>;
602
603 fn sub(self, rhs: &TracedTensor) -> Result<TracedTensor> {
604 TracedTensor::sub(self, rhs)
605 }
606}
607
608impl std::ops::Mul for &TracedTensor {
609 type Output = Result<TracedTensor>;
610
611 fn mul(self, rhs: &TracedTensor) -> Result<TracedTensor> {
612 TracedTensor::mul(self, rhs)
613 }
614}
615
616impl std::ops::Mul<f64> for &TracedTensor {
617 type Output = Result<TracedTensor>;
618
619 fn mul(self, rhs: f64) -> Result<TracedTensor> {
620 self.scale_real(rhs)
621 }
622}
623
624impl std::ops::Mul<&TracedTensor> for f64 {
625 type Output = Result<TracedTensor>;
626
627 fn mul(self, rhs: &TracedTensor) -> Result<TracedTensor> {
628 rhs.scale_real(self)
629 }
630}
631
632impl std::ops::Neg for &TracedTensor {
633 type Output = Result<TracedTensor>;
634
635 fn neg(self) -> Self::Output {
636 TracedTensor::neg(self)
637 }
638}
639
640impl std::ops::Div for &TracedTensor {
641 type Output = Result<TracedTensor>;
642
643 fn div(self, rhs: &TracedTensor) -> Result<TracedTensor> {
644 TracedTensor::div(self, rhs)
645 }
646}
647
648impl std::ops::Rem for &TracedTensor {
649 type Output = Result<TracedTensor>;
650
651 fn rem(self, rhs: &TracedTensor) -> Result<TracedTensor> {
652 TracedTensor::rem(self, rhs)
653 }
654}
655
656impl TracedTensor {
657 /// Return the graph that owns this traced tensor's current value.
658 ///
659 /// # Examples
660 ///
661 /// ```
662 /// use tenferro_runtime::TracedTensor;
663 ///
664 /// let x = TracedTensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap();
665 /// let _graph = x.graph();
666 /// ```
667 pub fn graph(&self) -> &Arc<Graph<StdTensorOp>> {
668 &self.graph
669 }
670
671 /// Return the concrete tensor data attached to this traced value, if any.
672 ///
673 /// Placeholder tensors created with `input_concrete_shape` or
674 /// `input_symbolic_shape` have no attached data until execution bindings
675 /// provide it.
676 ///
677 /// # Examples
678 ///
679 /// ```
680 /// use tenferro_runtime::{DType, TracedTensor};
681 ///
682 /// let concrete = TracedTensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap();
683 /// assert!(concrete.attached_value().is_some());
684 ///
685 /// let placeholder = TracedTensor::input_symbolic_shape(DType::F64, 1).unwrap();
686 /// assert!(placeholder.attached_value().is_none());
687 /// ```
688 pub fn attached_value(&self) -> Option<&Arc<RetainedValue>> {
689 self.data.as_ref()
690 }
691
692 /// Build a [`TracedTensor`] leaf from a concrete [`Tensor`], keeping its
693 /// shape as a concrete `shape_hint`.
694 ///
695 /// This is the common constructor when you have concrete tensor data that
696 /// you want to use both for graph building and for evaluation. The
697 /// resulting tensor is treated as a concrete-shape leaf by downstream
698 /// passes (binary einsum decomposition, build-time reshape folding, etc.).
699 ///
700 /// # Examples
701 ///
702 /// ```
703 /// use tenferro_runtime::{Tensor, TracedTensor};
704 ///
705 /// let a = TracedTensor::from_tensor_concrete_shape(
706 /// Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap(),
707 /// )
708 /// .unwrap();
709 /// assert_eq!(a.rank, 2);
710 /// assert!(a.is_concrete_shape());
711 /// ```
712 ///
713 /// # Errors
714 ///
715 /// Returns [`Error::RuntimeStateSource`] when graph metadata registration
716 /// cannot retain the concrete tensor's shape or dtype.
717 pub fn from_tensor_concrete_shape(tensor: Tensor) -> Result<Self> {
718 Self::from_tensor_concrete_shape_with_identity(tensor, None)
719 }
720
721 /// Build a [`TracedTensor`] leaf from a concrete tensor that declares the
722 /// canonical identity of its externally defined scalar.
723 ///
724 /// # Examples
725 ///
726 /// ```rust
727 /// use tenferro_runtime::TracedTensor;
728 /// use tenferro_tensor::DType;
729 ///
730 /// let dtype = DType::External(std::any::TypeId::of::<f64>());
731 /// let leaf = TracedTensor::input_concrete_shape_declaring_scalar(dtype, &[2], "example.scalar.v1")?;
732 /// assert!(leaf.is_concrete_shape());
733 /// # Ok::<(), tenferro_runtime::Error>(())
734 /// ```
735 ///
736 /// # Errors
737 ///
738 /// Returns [`Error::RuntimeStateSource`] when graph metadata registration
739 /// cannot retain the concrete tensor's shape or dtype.
740 pub fn from_tensor_concrete_shape_declaring_scalar(
741 tensor: Tensor,
742 identity: &'static str,
743 ) -> Result<Self> {
744 Self::from_tensor_concrete_shape_with_identity(tensor, Some(identity))
745 }
746
747 fn from_tensor_concrete_shape_with_identity(
748 tensor: Tensor,
749 identity: Option<&'static str>,
750 ) -> Result<Self> {
751 let shape = tensor.shape().to_vec();
752 let rank = shape.len();
753 let dtype = tensor.dtype();
754 let key = next_input_key();
755 let id = next_traced_id();
756 let data = Arc::new(RetainedValue::from_tensor(tensor));
757 let meta = concrete_tensor_meta(dtype, &shape);
758 let meta = match identity {
759 Some(identity) => meta.with_scalar_identity(identity),
760 None => meta,
761 };
762
763 let mut builder = GraphBuilder::new();
764 let val = builder.add_input(key.clone());
765 builder.set_outputs(vec![val]);
766 let graph = Arc::new(builder.build());
767 let metadata_scope = register_metadata_or_runtime_state(register_scoped_value_metadata(
768 graph.values()[val].key.clone(),
769 meta.clone(),
770 ))?;
771
772 let mut map = HashMap::new();
773 map.insert(key.clone(), Arc::clone(&data));
774 let mut leaf_metas = HashMap::new();
775 leaf_metas.insert(key, meta);
776
777 Ok(Self {
778 id,
779 rank,
780 dtype,
781 graph,
782 val,
783 data: Some(data),
784 shape_hint: Some(shape.into_iter().map(SymDim::from).collect()),
785 inputs_map: Arc::new(map),
786 leaf_metas: Arc::new(leaf_metas),
787 extra_roots: Vec::new(),
788 checkpoint_chain: None,
789 metadata_scopes: MetadataScopeChain::from_scope(metadata_scope),
790 constraint_scopes: ConstraintScopeChain::empty(),
791 })
792 }
793
794 /// Build a [`TracedTensor`] leaf from a concrete [`Tensor`] but advertise
795 /// a symbolic shape during graph construction.
796 ///
797 /// The tensor data is still attached (so plain `eval` works without
798 /// bindings), but graph passes see the leaf as shape-symbolic. This is
799 /// useful for building a single traced program that should not bake in
800 /// shape-specific optimizations.
801 ///
802 /// # Examples
803 ///
804 /// ```
805 /// use tenferro_runtime::{Tensor, TracedTensor};
806 ///
807 /// let t = TracedTensor::from_tensor_symbolic_shape(
808 /// Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap(),
809 /// )
810 /// .unwrap();
811 /// assert_eq!(t.rank, 2);
812 /// assert!(!t.is_concrete_shape());
813 /// ```
814 ///
815 /// # Errors
816 ///
817 /// Returns [`Error::RuntimeStateSource`] when symbolic graph metadata
818 /// registration is unavailable or its registry state is poisoned.
819 pub fn from_tensor_symbolic_shape(tensor: Tensor) -> Result<Self> {
820 Self::from_tensor_value_symbolic_shape(TensorValue::from_tensor(tensor))
821 }
822
823 /// Build a data-attached symbolic-shape traced leaf from shared tensor data.
824 ///
825 /// This is an internal crate-boundary helper for eager AD recording. It
826 /// preserves the same tensor allocation used by the eager value while graph
827 /// passes see symbolic input extents.
828 ///
829 /// # Examples
830 ///
831 /// ```
832 /// use std::sync::Arc;
833 /// use tenferro_runtime::{Tensor, TracedTensor};
834 ///
835 /// let tensor = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
836 /// let traced = TracedTensor::from_tensor_symbolic_shape(tensor)?;
837 /// assert_eq!(traced.rank, 1);
838 /// assert!(!traced.is_concrete_shape());
839 /// # Ok::<(), tenferro_runtime::Error>(())
840 /// ```
841 ///
842 /// # Errors
843 ///
844 /// Returns [`Error::RuntimeStateSource`] when symbolic graph metadata
845 /// registration is unavailable or its registry state is poisoned.
846 #[doc(hidden)]
847 pub fn from_tensor_value_symbolic_shape(value: TensorValue) -> Result<Self> {
848 let data = Arc::new(RetainedValue::from_tensor_value(value)?);
849 Self::from_shared_tensor_value_symbolic_shape(data)
850 }
851
852 #[doc(hidden)]
853 pub fn from_shared_tensor_value_symbolic_shape(data: Arc<RetainedValue>) -> Result<Self> {
854 let rank = data.shape().len();
855 let dtype = data.dtype();
856 let key = next_input_key();
857 let id = next_traced_id();
858 let meta = symbolic_input_meta(dtype, id, rank);
859
860 let mut builder = GraphBuilder::new();
861 let val = builder.add_input(key.clone());
862 builder.set_outputs(vec![val]);
863 let graph = Arc::new(builder.build());
864 let metadata_scope = register_metadata_or_runtime_state(register_scoped_value_metadata(
865 graph.values()[val].key.clone(),
866 meta.clone(),
867 ))?;
868
869 let mut map = HashMap::new();
870 map.insert(key.clone(), Arc::clone(&data));
871 let mut leaf_metas = HashMap::new();
872 leaf_metas.insert(key, meta);
873
874 Ok(Self {
875 id,
876 rank,
877 dtype,
878 graph,
879 val,
880 data: Some(data),
881 shape_hint: None,
882 inputs_map: Arc::new(map),
883 leaf_metas: Arc::new(leaf_metas),
884 extra_roots: Vec::new(),
885 checkpoint_chain: None,
886 metadata_scopes: MetadataScopeChain::from_scope(metadata_scope),
887 constraint_scopes: ConstraintScopeChain::empty(),
888 })
889 }
890
891 /// Build a data-less placeholder leaf with a fixed (concrete) shape.
892 ///
893 /// Must be passed as an input to [`crate::Runtime::run_compiled`] before evaluation.
894 /// Use this when you know the exact shape of the input but want to build
895 /// the graph once and feed different concrete tensors at execution time.
896 ///
897 /// # Examples
898 ///
899 /// ```
900 /// use tenferro_tensor::DType;
901 /// use tenferro_runtime::TracedTensor;
902 ///
903 /// let x = TracedTensor::input_concrete_shape(DType::F64, &[2, 3]).unwrap();
904 /// assert_eq!(x.rank, 2);
905 /// assert!(x.is_concrete_shape());
906 /// ```
907 ///
908 /// # Errors
909 ///
910 /// Returns [`Error::RuntimeStateSource`] when graph metadata registration
911 /// fails or the registry state is poisoned. `dtype` and `shape` are
912 /// metadata values and are not revalidated by this constructor.
913 pub fn input_concrete_shape(dtype: DType, shape: &[usize]) -> Result<Self> {
914 Self::input_concrete_shape_with_identity(dtype, shape, None)
915 }
916
917 /// Build a data-less placeholder leaf of a concrete shape that declares the
918 /// canonical identity of its externally defined scalar.
919 ///
920 /// A semantic program's identity must be reproducible across processes, while an
921 /// externally defined scalar's tag is a process-local `TypeId`, so an input whose
922 /// dtype is external must declare the stable name its contribution owns.
923 ///
924 /// # Examples
925 ///
926 /// ```rust
927 /// use tenferro_runtime::TracedTensor;
928 /// use tenferro_tensor::DType;
929 ///
930 /// let dtype = DType::External(std::any::TypeId::of::<f64>());
931 /// let x = TracedTensor::input_concrete_shape_declaring_scalar(dtype, &[2], "example.scalar.v1")?;
932 /// assert_eq!(x.rank, 1);
933 /// assert!(x.is_concrete_shape());
934 /// # Ok::<(), tenferro_runtime::Error>(())
935 /// ```
936 ///
937 /// # Errors
938 ///
939 /// Returns [`Error::RuntimeStateSource`] when graph metadata registration fails
940 /// or the registry state is poisoned.
941 pub fn input_concrete_shape_declaring_scalar(
942 dtype: DType,
943 shape: &[usize],
944 identity: &'static str,
945 ) -> Result<Self> {
946 Self::input_concrete_shape_with_identity(dtype, shape, Some(identity))
947 }
948
949 fn input_concrete_shape_with_identity(
950 dtype: DType,
951 shape: &[usize],
952 identity: Option<&'static str>,
953 ) -> Result<Self> {
954 let shape = shape.to_vec();
955 let rank = shape.len();
956 let key = next_input_key();
957 let id = next_traced_id();
958
959 let mut builder = GraphBuilder::new();
960 let val = builder.add_input(key.clone());
961 builder.set_outputs(vec![val]);
962 let graph = Arc::new(builder.build());
963 let meta = concrete_tensor_meta(dtype, &shape);
964 let meta = match identity {
965 Some(identity) => meta.with_scalar_identity(identity),
966 None => meta,
967 };
968 let metadata_scope = register_metadata_or_runtime_state(register_scoped_value_metadata(
969 graph.values()[val].key.clone(),
970 meta,
971 ))?;
972
973 Ok(Self {
974 id,
975 rank,
976 dtype,
977 graph,
978 val,
979 data: None,
980 shape_hint: Some(shape.into_iter().map(SymDim::from).collect()),
981 inputs_map: Arc::new(HashMap::new()),
982 leaf_metas: Arc::new(HashMap::new()),
983 extra_roots: Vec::new(),
984 checkpoint_chain: None,
985 metadata_scopes: MetadataScopeChain::from_scope(metadata_scope),
986 constraint_scopes: ConstraintScopeChain::empty(),
987 })
988 }
989
990 /// Build a data-less placeholder leaf with the given rank but fully
991 /// symbolic shape (every dim is a distinct `SymDim::TensorAxis`).
992 ///
993 /// Must be passed as an input to [`crate::Runtime::run_compiled`] before
994 /// evaluation. Use this to build shape-agnostic graphs.
995 ///
996 /// # Examples
997 ///
998 /// ```
999 /// use tenferro_tensor::DType;
1000 /// use tenferro_runtime::TracedTensor;
1001 ///
1002 /// let x = TracedTensor::input_symbolic_shape(DType::F64, 2).unwrap();
1003 /// assert_eq!(x.rank, 2);
1004 /// assert!(!x.is_concrete_shape());
1005 /// ```
1006 ///
1007 /// # Errors
1008 ///
1009 /// Returns [`Error::RuntimeStateSource`] when graph metadata registration
1010 /// fails or the registry state is poisoned. `rank` is recorded as the
1011 /// symbolic placeholder rank and is not otherwise rejected here.
1012 pub fn input_symbolic_shape(dtype: DType, rank: usize) -> Result<Self> {
1013 let key = next_input_key();
1014 let id = next_traced_id();
1015
1016 let mut builder = GraphBuilder::new();
1017 let val = builder.add_input(key.clone());
1018 builder.set_outputs(vec![val]);
1019 let graph = Arc::new(builder.build());
1020 let metadata_scope = register_metadata_or_runtime_state(register_scoped_value_metadata(
1021 graph.values()[val].key.clone(),
1022 symbolic_input_meta(dtype, id, rank),
1023 ))?;
1024
1025 Ok(Self {
1026 id,
1027 rank,
1028 dtype,
1029 graph,
1030 val,
1031 data: None,
1032 shape_hint: None,
1033 inputs_map: Arc::new(HashMap::new()),
1034 leaf_metas: Arc::new(HashMap::new()),
1035 extra_roots: Vec::new(),
1036 checkpoint_chain: None,
1037 metadata_scopes: MetadataScopeChain::from_scope(metadata_scope),
1038 constraint_scopes: ConstraintScopeChain::empty(),
1039 })
1040 }
1041
1042 /// Build a concrete-shape [`TracedTensor`] leaf from column-major typed
1043 /// `Vec<T>` data.
1044 ///
1045 /// The data must already be in tenferro's physical column-major order.
1046 ///
1047 /// # Examples
1048 ///
1049 /// ```
1050 /// use tenferro_runtime::TracedTensor;
1051 ///
1052 /// let a = TracedTensor::from_vec_col_major(
1053 /// vec![2, 3],
1054 /// vec![1.0_f64, 4.0, 2.0, 5.0, 3.0, 6.0],
1055 /// )?;
1056 /// assert_eq!(a.rank, 2);
1057 /// # Ok::<(), tenferro_runtime::Error>(())
1058 /// ```
1059 ///
1060 /// # Errors
1061 ///
1062 /// Returns [`Error::TensorRuntime`] containing
1063 /// `ValidationError::ShapeDataLengthMismatch` when the shape product does
1064 /// not equal `data.len()`, or `ValidationError::IntegerOverflow` when the
1065 /// shape product cannot be represented by `usize`.
1066 pub fn from_vec_col_major<T: TensorScalar>(
1067 shape: impl IntoShapeVec,
1068 data: Vec<T>,
1069 ) -> Result<Self> {
1070 Self::from_tensor_concrete_shape(Tensor::from_vec_col_major(shape, data)?)
1071 }
1072
1073 /// Return the tensor element dtype recorded for this traced value.
1074 pub fn dtype(&self) -> DType {
1075 self.dtype
1076 }
1077
1078 /// Returns `true` iff every dim of this tensor's `shape_hint` is a
1079 /// constant `SymDim` (i.e. the shape is fully known at graph-build time).
1080 ///
1081 /// # Examples
1082 ///
1083 /// ```
1084 /// use tenferro_tensor::DType;
1085 /// use tenferro_runtime::TracedTensor;
1086 ///
1087 /// let a = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
1088 /// let b = TracedTensor::input_symbolic_shape(DType::F64, 2).unwrap();
1089 /// assert!(a.is_concrete_shape());
1090 /// assert!(!b.is_concrete_shape());
1091 /// ```
1092 pub fn is_concrete_shape(&self) -> bool {
1093 try_concrete_shape(self).is_some()
1094 }
1095
1096 /// Return the fully-concrete shape of this tensor, if every dim of
1097 /// its shape-hint is a constant `SymDim`. Returns `None` if any
1098 /// dimension is symbolic.
1099 ///
1100 /// This is the counterpart to [`Self::is_concrete_shape`] for callers
1101 /// that need to *use* the concrete shape (e.g. external composition
1102 /// wrappers building `broadcast_in_dim` payloads from known shapes).
1103 ///
1104 /// # Examples
1105 ///
1106 /// ```
1107 /// use tenferro_tensor::DType;
1108 /// use tenferro_runtime::TracedTensor;
1109 ///
1110 /// let a = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
1111 /// assert_eq!(a.try_concrete_shape(), Some(vec![2, 3]));
1112 ///
1113 /// let b = TracedTensor::input_symbolic_shape(DType::F64, 2).unwrap();
1114 /// assert!(b.try_concrete_shape().is_none());
1115 /// ```
1116 pub fn try_concrete_shape(&self) -> Option<Vec<usize>> {
1117 try_concrete_shape(self)
1118 }
1119
1120 /// Return the concrete tensor shape.
1121 ///
1122 /// Returns an error when a shape hint is missing or any dimension is
1123 /// symbolic. Composite traced ops that require concrete sizes should
1124 /// propagate this error instead of panicking.
1125 ///
1126 /// # Errors
1127 ///
1128 /// Returns [`Error::Validation`] with `InvalidArgument` when this tensor
1129 /// has no shape hint or any dimension is symbolic.
1130 pub fn concrete_shape(&self) -> Result<Vec<usize>> {
1131 concrete_shape(self)
1132 }
1133
1134 /// If this `TracedTensor` is a leaf (single-node input graph),
1135 /// return its input key. Computed tensors return `None`.
1136 pub fn input_key(&self) -> Option<TensorInputKey> {
1137 match &self.graph.values()[self.val].key {
1138 ValueKey::Input(key) => Some(key.clone()),
1139 _ => None,
1140 }
1141 }
1142
1143 /// Return whether this traced graph carries default data for `key`.
1144 #[doc(hidden)]
1145 pub fn has_attached_input_key(&self, key: &TensorInputKey) -> bool {
1146 self.inputs_map.contains_key(key)
1147 }
1148
1149 /// Elementwise addition with NumPy-style broadcasting.
1150 ///
1151 /// Prefer using the `+` operator when it reads naturally.
1152 ///
1153 /// A longer expression such as `a + b + c` does not compose because the
1154 /// first `+` returns `Result<TracedTensor, Error>`, so the second `+`
1155 /// would receive a result rather than a tensor. Use `?` at each step or
1156 /// the explicit fallible method chain shown below when the operation
1157 /// sequence is more important than notation:
1158 ///
1159 /// # Examples
1160 ///
1161 /// ```rust
1162 /// # use tenferro_runtime::{Error, TracedTensor};
1163 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1164 /// # let z = TracedTensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0]).unwrap();
1165 /// let y = x.add(&z);
1166 /// let y2 = &x + &z;
1167 /// # fn add_three(
1168 /// # a: &TracedTensor,
1169 /// # b: &TracedTensor,
1170 /// # c: &TracedTensor,
1171 /// # ) -> Result<TracedTensor, Error> {
1172 /// let ab = (a + b)?;
1173 /// let sum = (&ab + c)?;
1174 /// let method_chain = a.add(b)?.add(c)?;
1175 /// let _ = method_chain;
1176 /// # Ok(sum)
1177 /// # }
1178 /// ```
1179 ///
1180 /// Tenferro prioritizes robust error handling over the conciseness of
1181 /// chained operator notation; the explicit fallible methods are the
1182 /// canonical form for longer sequences.
1183 ///
1184 /// # Errors
1185 ///
1186 /// Returns [`Error::Validation`] with `ShapeMismatch` when operand shapes
1187 /// cannot be broadcast, or [`Error::RuntimeStateSource`] when graph
1188 /// metadata registration fails.
1189 ///
1190 /// # Deferred errors
1191 ///
1192 /// If symbolic dimensions prevent shape comparison during graph
1193 /// construction, the same `ShapeMismatch` can be reported during
1194 /// compilation or execution, with the corresponding [`ErrorPhase`].
1195 pub fn add(&self, other: &TracedTensor) -> Result<TracedTensor> {
1196 let (lhs, rhs) = broadcast_binary(self, other)?;
1197 apply_binary(
1198 StdTensorOp::Add,
1199 &lhs,
1200 &rhs,
1201 lhs.rank,
1202 lhs.shape_hint.clone(),
1203 )
1204 }
1205
1206 /// Elementwise subtraction with NumPy-style broadcasting.
1207 ///
1208 /// Prefer using the `-` operator when it reads naturally.
1209 ///
1210 /// # Errors
1211 ///
1212 /// Returns [`Error::Validation`] with `ShapeMismatch` when operand shapes
1213 /// cannot be broadcast, or [`Error::RuntimeStateSource`] when graph
1214 /// metadata registration fails.
1215 ///
1216 /// # Deferred errors
1217 ///
1218 /// If symbolic dimensions prevent shape comparison during graph
1219 /// construction, the same `ShapeMismatch` can be reported during
1220 /// compilation or execution, with the corresponding [`ErrorPhase`].
1221 pub fn sub(&self, other: &TracedTensor) -> Result<TracedTensor> {
1222 let (lhs, rhs) = broadcast_binary(self, other)?;
1223 apply_binary(
1224 StdTensorOp::Sub,
1225 &lhs,
1226 &rhs,
1227 lhs.rank,
1228 lhs.shape_hint.clone(),
1229 )
1230 }
1231
1232 /// Elementwise multiplication with NumPy-style broadcasting.
1233 ///
1234 /// Prefer using the `*` operator when it reads naturally.
1235 ///
1236 /// # Examples
1237 ///
1238 /// ```rust
1239 /// # use tenferro_runtime::TracedTensor;
1240 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1241 /// # let z = TracedTensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0]).unwrap();
1242 /// let y = x.mul(&z);
1243 /// let y2 = &x * &z;
1244 /// ```
1245 ///
1246 /// # Errors
1247 ///
1248 /// Returns [`Error::Validation`] with `ShapeMismatch` when operand shapes
1249 /// cannot be broadcast, or [`Error::RuntimeStateSource`] when graph
1250 /// metadata registration fails.
1251 ///
1252 /// # Deferred errors
1253 ///
1254 /// If symbolic ranks prevent shape comparison during graph construction,
1255 /// the same `ShapeMismatch` can be reported during compilation or
1256 /// execution, with the corresponding [`ErrorPhase`].
1257 pub fn mul(&self, other: &TracedTensor) -> Result<TracedTensor> {
1258 let (lhs, rhs) = broadcast_binary(self, other)?;
1259 apply_binary(
1260 StdTensorOp::Mul,
1261 &lhs,
1262 &rhs,
1263 lhs.rank,
1264 lhs.shape_hint.clone(),
1265 )
1266 }
1267
1268 /// Elementwise division with NumPy-style broadcasting.
1269 ///
1270 /// Prefer using the `/` operator when it reads naturally.
1271 ///
1272 /// # Examples
1273 ///
1274 /// ```rust
1275 /// # use tenferro_runtime::TracedTensor;
1276 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1277 /// # let z = TracedTensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0]).unwrap();
1278 /// let y = x.div(&z);
1279 /// let y2 = &x / &z;
1280 /// ```
1281 ///
1282 /// # Errors
1283 ///
1284 /// Returns [`Error::Validation`] with `ShapeMismatch` when operand shapes
1285 /// cannot be broadcast, or [`Error::RuntimeStateSource`] when graph
1286 /// metadata registration fails.
1287 ///
1288 /// # Deferred errors
1289 ///
1290 /// If symbolic ranks prevent shape comparison during graph construction,
1291 /// the same `ShapeMismatch` can be reported during compilation or
1292 /// execution, with the corresponding [`ErrorPhase`]. For integer inputs,
1293 /// a zero divisor is reported during execution as
1294 /// [`Error::TensorRuntime`] containing a
1295 /// [`tenferro_tensor::Error::Extension`] classified as
1296 /// `tenferro_tensor::ErrorKind::NumericalFailure` and retaining the typed
1297 /// backend source; floating-point and complex zero divisors follow their
1298 /// numeric semantics instead.
1299 pub fn div(&self, other: &TracedTensor) -> Result<TracedTensor> {
1300 let (lhs, rhs) = broadcast_binary(self, other)?;
1301 apply_binary(
1302 StdTensorOp::Div,
1303 &lhs,
1304 &rhs,
1305 lhs.rank,
1306 lhs.shape_hint.clone(),
1307 )
1308 }
1309
1310 /// Elementwise remainder with NumPy-style broadcasting.
1311 ///
1312 /// Prefer using the `%` operator when it reads naturally.
1313 ///
1314 /// # Errors
1315 ///
1316 /// Returns [`Error::Validation`] with `ShapeMismatch` when operand shapes
1317 /// cannot be broadcast, [`Error::Unsupported`] at
1318 /// [`ErrorPhase::GraphBuild`] when either operand has a complex dtype, or
1319 /// [`Error::RuntimeStateSource`] when graph metadata registration fails.
1320 ///
1321 /// # Deferred errors
1322 ///
1323 /// If symbolic ranks prevent shape comparison during graph construction,
1324 /// the same `ShapeMismatch` can be reported during compilation or
1325 /// execution, with the corresponding [`ErrorPhase`]. For integer inputs,
1326 /// a zero divisor is reported during execution as
1327 /// [`Error::TensorRuntime`] containing a
1328 /// [`tenferro_tensor::Error::Extension`] classified as
1329 /// `tenferro_tensor::ErrorKind::NumericalFailure` and retaining the typed
1330 /// backend source; floating-point zero divisors follow their numeric
1331 /// semantics.
1332 pub fn rem(&self, other: &TracedTensor) -> Result<TracedTensor> {
1333 let (lhs, rhs) = broadcast_binary(self, other)?;
1334 apply_binary(
1335 StdTensorOp::Rem,
1336 &lhs,
1337 &rhs,
1338 lhs.rank,
1339 lhs.shape_hint.clone(),
1340 )
1341 }
1342
1343 /// Elementwise comparison with NumPy-style broadcasting.
1344 ///
1345 /// # Errors
1346 ///
1347 /// Returns [`Error::Validation`] with `ShapeMismatch` when the concrete
1348 /// operands cannot be broadcast, [`Error::Unsupported`] when ordered
1349 /// comparison rejects a complex dtype, or [`Error::RuntimeStateSource`]
1350 /// when result metadata cannot be registered.
1351 ///
1352 /// # Deferred errors
1353 ///
1354 /// With same-rank symbolic operands, shape compatibility is retained as a
1355 /// graph constraint. A concrete mismatch is reported later as
1356 /// [`Error::TensorRuntime`] containing a typed validation source, with the
1357 /// failure phase identifying compilation or execution.
1358 pub fn compare(&self, other: &TracedTensor, dir: CompareDir) -> Result<TracedTensor> {
1359 apply_broadcast_binary_op(StdTensorOp::Compare(dir), self, other)
1360 }
1361
1362 /// Elementwise maximum with NumPy-style broadcasting.
1363 ///
1364 /// # Errors
1365 ///
1366 /// Returns [`Error::Validation`] with `ShapeMismatch` when the concrete
1367 /// operands cannot be broadcast, [`Error::Unsupported`] when ordered
1368 /// maximum rejects a complex dtype, or [`Error::RuntimeStateSource`] when
1369 /// result metadata cannot be registered.
1370 ///
1371 /// # Deferred errors
1372 ///
1373 /// With same-rank symbolic operands, the broadcast constraint may fail at
1374 /// compile or execution and is returned as [`Error::TensorRuntime`] with
1375 /// its typed validation source.
1376 pub fn maximum(&self, other: &TracedTensor) -> Result<TracedTensor> {
1377 apply_broadcast_binary_op(StdTensorOp::Maximum, self, other)
1378 }
1379
1380 /// Elementwise minimum with NumPy-style broadcasting.
1381 ///
1382 /// # Errors
1383 ///
1384 /// Returns [`Error::Validation`] with `ShapeMismatch` when the concrete
1385 /// operands cannot be broadcast, [`Error::Unsupported`] when ordered
1386 /// minimum rejects a complex dtype, or [`Error::RuntimeStateSource`] when
1387 /// result metadata cannot be registered.
1388 ///
1389 /// # Deferred errors
1390 ///
1391 /// With same-rank symbolic operands, the broadcast constraint may fail at
1392 /// compile or execution and is returned as [`Error::TensorRuntime`] with
1393 /// its typed validation source.
1394 pub fn minimum(&self, other: &TracedTensor) -> Result<TracedTensor> {
1395 apply_broadcast_binary_op(StdTensorOp::Minimum, self, other)
1396 }
1397
1398 /// Select values from `on_true` or `on_false` using `condition`.
1399 ///
1400 /// # Errors
1401 ///
1402 /// Returns [`Error::Validation`] with `InvalidArgument` when an operand
1403 /// lacks concrete shape metadata, or `ShapeMismatch` when the concrete
1404 /// condition and branches cannot share a broadcast shape. Dtype promotion
1405 /// failures are returned as [`Error::TensorRuntime`] with the typed
1406 /// `UnsupportedDTypeConversion` source; metadata failures retain
1407 /// [`Error::RuntimeStateSource`].
1408 pub fn where_select(
1409 condition: &TracedTensor,
1410 on_true: &TracedTensor,
1411 on_false: &TracedTensor,
1412 ) -> Result<TracedTensor> {
1413 apply_broadcast_ternary_op(StdTensorOp::Select, condition, on_true, on_false)
1414 }
1415
1416 /// Alias for [`Self::where_select`].
1417 ///
1418 /// # Errors
1419 ///
1420 /// Returns the same concrete failures as [`Self::where_select`]:
1421 /// [`Error::Validation`] with `InvalidArgument`/`ShapeMismatch` for shape
1422 /// metadata or broadcasting, [`Error::TensorRuntime`] with
1423 /// `UnsupportedDTypeConversion` for failed promotion, and
1424 /// [`Error::RuntimeStateSource`] for metadata registration.
1425 pub fn select(
1426 condition: &TracedTensor,
1427 on_true: &TracedTensor,
1428 on_false: &TracedTensor,
1429 ) -> Result<TracedTensor> {
1430 Self::where_select(condition, on_true, on_false)
1431 }
1432
1433 /// Clamp values elementwise between lower and upper bounds.
1434 ///
1435 /// # Errors
1436 ///
1437 /// Returns [`Error::Validation`] with `InvalidArgument` when an operand
1438 /// lacks concrete shape metadata, `ShapeMismatch` when bounds cannot be
1439 /// broadcast with the input, [`Error::Unsupported`] for an ordered
1440 /// complex dtype, or [`Error::RuntimeStateSource`] when metadata cannot be
1441 /// registered.
1442 pub fn clamp(&self, lower: &TracedTensor, upper: &TracedTensor) -> Result<TracedTensor> {
1443 apply_broadcast_ternary_op(StdTensorOp::Clamp, self, lower, upper)
1444 }
1445
1446 fn apply_same_shape_unary(&self, op: StdTensorOp) -> Result<TracedTensor> {
1447 apply_unary(op, self, self.rank, self.shape_hint.clone())
1448 }
1449
1450 /// Elementwise negation.
1451 ///
1452 /// Prefer using the unary `-` operator when it reads naturally.
1453 ///
1454 /// # Examples
1455 ///
1456 /// ```rust
1457 /// # use tenferro_runtime::TracedTensor;
1458 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1459 /// let y = x.neg().unwrap();
1460 /// let y2 = (-&x).unwrap();
1461 /// ```
1462 ///
1463 /// # Errors
1464 ///
1465 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1466 /// is unavailable or poisoned while recording the unary result.
1467 pub fn neg(&self) -> Result<TracedTensor> {
1468 self.apply_same_shape_unary(StdTensorOp::Neg)
1469 }
1470
1471 /// Elementwise complex conjugate.
1472 ///
1473 /// # Examples
1474 ///
1475 /// ```rust
1476 /// # use num_complex::Complex64;
1477 /// # use tenferro_runtime::TracedTensor;
1478 /// # let x = TracedTensor::from_vec_col_major(
1479 /// # vec![2],
1480 /// # vec![Complex64::new(1.0, 2.0), Complex64::new(3.0, 4.0)],
1481 /// # )
1482 /// # .unwrap();
1483 /// let y = x.conj().unwrap();
1484 /// ```
1485 ///
1486 /// # Errors
1487 ///
1488 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1489 /// is unavailable or poisoned while recording the unary result.
1490 pub fn conj(&self) -> Result<TracedTensor> {
1491 self.apply_same_shape_unary(StdTensorOp::Conj)
1492 }
1493
1494 /// Elementwise absolute value.
1495 ///
1496 /// Complex inputs return real magnitudes (`C32 -> F32`, `C64 -> F64`).
1497 ///
1498 /// # Examples
1499 ///
1500 /// ```rust
1501 /// # use tenferro_runtime::TracedTensor;
1502 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![-1.0_f64, 2.0]).unwrap();
1503 /// let y = x.abs().unwrap();
1504 /// ```
1505 ///
1506 /// # Errors
1507 ///
1508 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1509 /// is unavailable or poisoned while recording the unary result.
1510 pub fn abs(&self) -> Result<TracedTensor> {
1511 self.apply_same_shape_unary(StdTensorOp::Abs)
1512 }
1513
1514 /// Elementwise sign.
1515 ///
1516 /// # Examples
1517 ///
1518 /// ```rust
1519 /// # use tenferro_runtime::TracedTensor;
1520 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![-1.0_f64, 2.0]).unwrap();
1521 /// let y = x.sign().unwrap();
1522 /// ```
1523 ///
1524 /// # Errors
1525 ///
1526 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1527 /// is unavailable or poisoned while recording the unary result.
1528 pub fn sign(&self) -> Result<TracedTensor> {
1529 self.apply_same_shape_unary(StdTensorOp::Sign)
1530 }
1531
1532 /// Scale by a real scalar: `y = factor * x`.
1533 ///
1534 /// # Examples
1535 ///
1536 /// ```rust
1537 /// # use tenferro_runtime::TracedTensor;
1538 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1539 /// let y = x.scale_real(2.0)?;
1540 /// # Ok::<(), tenferro_runtime::Error>(())
1541 /// ```
1542 ///
1543 /// # Errors
1544 ///
1545 /// Returns [`Error::Validation`] with `InvalidArgument` when an integer or
1546 /// boolean factor is non-finite or out of range for the input dtype, or
1547 /// [`Error::RuntimeStateSource`] when output metadata registration fails.
1548 pub fn scale_real(&self, factor: f64) -> Result<TracedTensor> {
1549 let op = match self.dtype {
1550 DType::F64 => StdTensorOp::constant(factor),
1551 DType::F32 => StdTensorOp::constant(factor as f32),
1552 DType::I32 => StdTensorOp::constant(round_real_to_i32_for_op("scale_real", factor)?),
1553 DType::I64 => StdTensorOp::constant(round_real_to_i64(factor)?),
1554 DType::Bool => StdTensorOp::constant(bool_from_real_for_op("scale_real", factor)?),
1555 DType::C64 => StdTensorOp::constant(Complex64::new(factor, 0.0)),
1556 DType::C32 => StdTensorOp::constant(Complex32::new(factor as f32, 0.0)),
1557 // An externally defined scalar has no traced constant, so the closure
1558 // rejects it instead of guessing one.
1559 DType::External(_) => {
1560 return Err(graph_invalid_argument(
1561 "scale_real",
1562 "dtype",
1563 format!("requires a preset tensor dtype, got {:?}", self.dtype),
1564 ));
1565 }
1566 };
1567 scale_with_constant(self, op)
1568 }
1569
1570 /// Scale by a complex scalar: `y = factor * x`.
1571 ///
1572 /// Only complex tensors support complex scaling. For a real scalar factor
1573 /// that should preserve the input dtype, prefer [`scale_real`](Self::scale_real).
1574 ///
1575 /// # Examples
1576 ///
1577 /// ```rust
1578 /// use num_complex::Complex64;
1579 /// # use tenferro_runtime::TracedTensor;
1580 /// # let x = TracedTensor::from_vec_col_major(
1581 /// # vec![2],
1582 /// # vec![Complex64::new(1.0, 0.0), Complex64::new(2.0, 0.0)],
1583 /// # )
1584 /// # .unwrap();
1585 /// let y = x.scale_complex(Complex64::new(0.0, 1.0)).unwrap(); // multiply by i
1586 /// ```
1587 ///
1588 /// # Errors
1589 ///
1590 /// Returns [`Error::Validation`] with `InvalidArgument` when a complex
1591 /// factor is applied to a non-complex dtype, or
1592 /// [`Error::RuntimeStateSource`] when output metadata registration fails.
1593 pub fn scale_complex(&self, factor: Complex64) -> Result<TracedTensor> {
1594 match self.dtype {
1595 DType::C64 => scale_with_constant(self, StdTensorOp::constant(factor)),
1596 DType::C32 => scale_with_constant(
1597 self,
1598 StdTensorOp::constant(Complex32::new(factor.re as f32, factor.im as f32)),
1599 ),
1600 DType::F32
1601 | DType::F64
1602 | DType::I32
1603 | DType::I64
1604 | DType::Bool
1605 | DType::External(_) => Err(graph_invalid_argument(
1606 "scale_complex",
1607 "dtype",
1608 format!("requires complex tensor dtype, got {:?}", self.dtype),
1609 )),
1610 }
1611 }
1612
1613 /// Elementwise exponential.
1614 ///
1615 /// # Examples
1616 ///
1617 /// ```rust
1618 /// # use tenferro_runtime::TracedTensor;
1619 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1620 /// let y = x.exp().unwrap();
1621 /// ```
1622 ///
1623 /// # Errors
1624 ///
1625 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1626 /// is unavailable or poisoned while recording the unary result.
1627 pub fn exp(&self) -> Result<TracedTensor> {
1628 self.apply_same_shape_unary(StdTensorOp::Exp)
1629 }
1630
1631 /// Elementwise natural logarithm.
1632 ///
1633 /// # Examples
1634 ///
1635 /// ```rust
1636 /// # use tenferro_runtime::TracedTensor;
1637 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1638 /// let y = x.log().unwrap();
1639 /// ```
1640 ///
1641 /// # Errors
1642 ///
1643 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1644 /// is unavailable or poisoned while recording the unary result.
1645 pub fn log(&self) -> Result<TracedTensor> {
1646 self.apply_same_shape_unary(StdTensorOp::Log)
1647 }
1648
1649 /// Elementwise sine.
1650 ///
1651 /// # Examples
1652 ///
1653 /// ```rust
1654 /// # use tenferro_runtime::TracedTensor;
1655 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1656 /// let y = x.sin().unwrap();
1657 /// ```
1658 ///
1659 /// # Errors
1660 ///
1661 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1662 /// is unavailable or poisoned while recording the unary result.
1663 pub fn sin(&self) -> Result<TracedTensor> {
1664 self.apply_same_shape_unary(StdTensorOp::Sin)
1665 }
1666
1667 /// Elementwise cosine.
1668 ///
1669 /// # Examples
1670 ///
1671 /// ```rust
1672 /// # use tenferro_runtime::TracedTensor;
1673 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1674 /// let y = x.cos().unwrap();
1675 /// ```
1676 ///
1677 /// # Errors
1678 ///
1679 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1680 /// is unavailable or poisoned while recording the unary result.
1681 pub fn cos(&self) -> Result<TracedTensor> {
1682 self.apply_same_shape_unary(StdTensorOp::Cos)
1683 }
1684
1685 /// Elementwise hyperbolic tangent.
1686 ///
1687 /// # Examples
1688 ///
1689 /// ```rust
1690 /// # use tenferro_runtime::TracedTensor;
1691 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1692 /// let y = x.tanh().unwrap();
1693 /// ```
1694 ///
1695 /// # Errors
1696 ///
1697 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1698 /// is unavailable or poisoned while recording the unary result.
1699 pub fn tanh(&self) -> Result<TracedTensor> {
1700 self.apply_same_shape_unary(StdTensorOp::Tanh)
1701 }
1702
1703 /// Elementwise square root.
1704 ///
1705 /// # Examples
1706 ///
1707 /// ```rust
1708 /// # use tenferro_runtime::TracedTensor;
1709 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 4.0]).unwrap();
1710 /// let y = x.sqrt().unwrap();
1711 /// ```
1712 ///
1713 /// # Errors
1714 ///
1715 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1716 /// is unavailable or poisoned while recording the unary result.
1717 pub fn sqrt(&self) -> Result<TracedTensor> {
1718 self.apply_same_shape_unary(StdTensorOp::Sqrt)
1719 }
1720
1721 /// Elementwise reciprocal square root.
1722 ///
1723 /// # Examples
1724 ///
1725 /// ```rust
1726 /// # use tenferro_runtime::TracedTensor;
1727 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 4.0]).unwrap();
1728 /// let y = x.rsqrt().unwrap();
1729 /// ```
1730 ///
1731 /// # Errors
1732 ///
1733 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1734 /// is unavailable or poisoned while recording the unary result.
1735 pub fn rsqrt(&self) -> Result<TracedTensor> {
1736 self.apply_same_shape_unary(StdTensorOp::Rsqrt)
1737 }
1738
1739 /// Elementwise power with NumPy-style broadcasting.
1740 ///
1741 /// # Examples
1742 ///
1743 /// ```rust
1744 /// # use tenferro_runtime::TracedTensor;
1745 /// # let base = TracedTensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0]).unwrap();
1746 /// # let exp = TracedTensor::from_vec_col_major(vec![2], vec![3.0_f64, 2.0]).unwrap();
1747 /// let y = base.pow(&exp);
1748 /// ```
1749 ///
1750 /// # Errors
1751 ///
1752 /// Returns [`Error::Validation`] with `ShapeMismatch` when the concrete
1753 /// operands cannot be broadcast, or [`Error::RuntimeStateSource`] when
1754 /// result metadata cannot be registered.
1755 ///
1756 /// # Deferred errors
1757 ///
1758 /// A symbolic broadcast mismatch or integer negative exponent is
1759 /// discovered at compile or execution and is returned as
1760 /// [`Error::TensorRuntime`] with a typed `ShapeMismatch` or
1761 /// `NegativeIntegerExponent` numerical source and the corresponding
1762 /// [`ErrorPhase`].
1763 pub fn pow(&self, other: &TracedTensor) -> Result<TracedTensor> {
1764 let (lhs, rhs) = broadcast_binary(self, other)?;
1765 apply_binary(
1766 StdTensorOp::Pow,
1767 &lhs,
1768 &rhs,
1769 lhs.rank,
1770 lhs.shape_hint.clone(),
1771 )
1772 }
1773
1774 /// Elementwise `exp(x) - 1`.
1775 ///
1776 /// # Examples
1777 ///
1778 /// ```rust
1779 /// # use tenferro_runtime::TracedTensor;
1780 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1781 /// let y = x.expm1().unwrap();
1782 /// ```
1783 ///
1784 /// # Errors
1785 ///
1786 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1787 /// is unavailable or poisoned while recording the unary result.
1788 pub fn expm1(&self) -> Result<TracedTensor> {
1789 self.apply_same_shape_unary(StdTensorOp::Expm1)
1790 }
1791
1792 /// Elementwise `log(1 + x)`.
1793 ///
1794 /// # Examples
1795 ///
1796 /// ```rust
1797 /// # use tenferro_runtime::TracedTensor;
1798 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1799 /// let y = x.log1p().unwrap();
1800 /// ```
1801 ///
1802 /// # Errors
1803 ///
1804 /// Returns [`Error::RuntimeStateSource`] when the graph metadata registry
1805 /// is unavailable or poisoned while recording the unary result.
1806 pub fn log1p(&self) -> Result<TracedTensor> {
1807 self.apply_same_shape_unary(StdTensorOp::Log1p)
1808 }
1809
1810 /// Elementwise error function `erf(x)`, for real `F32`/`F64` tensors.
1811 ///
1812 /// `erf(+-0) = +-0`, `erf(+-inf) = +-1`, and `NaN` stays `NaN`. The
1813 /// derivative is `2/sqrt(pi) * exp(-x^2)`.
1814 ///
1815 /// # Examples
1816 ///
1817 /// ```rust
1818 /// # use tenferro_runtime::TracedTensor;
1819 /// let x = TracedTensor::from_vec_col_major(vec![2], vec![0.0_f64, 1.0])?;
1820 /// let y = x.erf()?;
1821 /// assert_eq!(y.dtype(), tenferro_runtime::DType::F64);
1822 /// # Ok::<(), tenferro_runtime::Error>(())
1823 /// ```
1824 ///
1825 /// # Errors
1826 ///
1827 /// Returns [`Error::Unsupported`] (phase `GraphBuild`) for complex,
1828 /// integer, or `Bool` input, or [`Error::RuntimeStateSource`] when the
1829 /// graph metadata registry is unavailable or poisoned while recording the
1830 /// unary result.
1831 pub fn erf(&self) -> Result<TracedTensor> {
1832 if !matches!(self.dtype, DType::F32 | DType::F64) {
1833 return Err(Error::unsupported(
1834 "TracedTensor::erf",
1835 ErrorPhase::GraphBuild,
1836 format!(
1837 "erf is defined for real F32/F64 tensors, got {:?}",
1838 self.dtype
1839 ),
1840 ));
1841 }
1842 self.apply_same_shape_unary(StdTensorOp::Erf)
1843 }
1844
1845 /// Convert the tensor to a different dtype using checked conversion.
1846 ///
1847 /// Use [`cast`](Self::cast) when a lossy dtype projection is intended.
1848 ///
1849 /// # Examples
1850 ///
1851 /// ```rust
1852 /// use tenferro_runtime::DType;
1853 /// # use tenferro_runtime::TracedTensor;
1854 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
1855 ///
1856 /// let y = x.convert(DType::C64)?;
1857 /// # Ok::<(), tenferro_runtime::Error>(())
1858 /// ```
1859 ///
1860 /// # Errors
1861 ///
1862 /// Returns [`tenferro_tensor::Error::UnsupportedDTypeConversion`] when the
1863 /// requested pair is outside tenferro's checked dtype-promotion lattice,
1864 /// or [`Error::Validation`] when graph metadata rejects the conversion.
1865 /// Use [`cast`](Self::cast) for explicit lossy dtype projection.
1866 pub fn convert(&self, to: DType) -> Result<TracedTensor> {
1867 tenferro_tensor::validate::validate_convert_dtype("TracedTensor::convert", self.dtype, to)?;
1868 self.cast(to)
1869 }
1870
1871 /// Cast the tensor to a different dtype using explicit dtype projection.
1872 ///
1873 /// `cast` may truncate, narrow precision, project complex values to their
1874 /// real component, or use boolean truthiness where the backend supports the
1875 /// requested projection.
1876 ///
1877 /// # Examples
1878 ///
1879 /// ```rust
1880 /// use tenferro_runtime::DType;
1881 /// # use tenferro_runtime::TracedTensor;
1882 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.2_f64, -2.8]).unwrap();
1883 ///
1884 /// let y = x.cast(DType::I32).unwrap();
1885 /// ```
1886 ///
1887 /// # Errors
1888 ///
1889 /// Returns [`Error::TensorRuntime`] containing
1890 /// `UnsupportedDTypeConversion` when the requested input-to-target
1891 /// projection is not supported, or [`Error::RuntimeStateSource`] when
1892 /// converted-output metadata cannot be registered.
1893 pub fn cast(&self, to: DType) -> Result<TracedTensor> {
1894 if self.dtype == to {
1895 return Ok(self.clone());
1896 }
1897
1898 apply_unary_with_dtype(
1899 StdTensorOp::Convert {
1900 from: self.dtype,
1901 to,
1902 },
1903 self,
1904 self.rank,
1905 self.shape_hint.clone(),
1906 to,
1907 )
1908 }
1909
1910 /// Generalized tensor contraction.
1911 ///
1912 /// The output layout is `[lhs free..., rhs free..., batch...]`: batch axes
1913 /// come last (see [`DotGeneralConfig`]).
1914 ///
1915 /// # Examples
1916 ///
1917 /// ```rust
1918 /// # use tenferro_runtime::{DotGeneralConfig, TracedTensor};
1919 /// # let a = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
1920 /// # let b = TracedTensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
1921 /// # let config = DotGeneralConfig {
1922 /// # lhs_contracting_dims: [1].as_slice().into(),
1923 /// # rhs_contracting_dims: [0].as_slice().into(),
1924 /// # lhs_batch_dims: [].as_slice().into(),
1925 /// # rhs_batch_dims: [].as_slice().into(),
1926 /// # };
1927 /// let y = a.dot_general(&b, config)?;
1928 /// # Ok::<(), tenferro_runtime::Error>(())
1929 /// ```
1930 ///
1931 /// # Errors
1932 ///
1933 /// Returns [`Error::Validation`] with `RankMismatch`, `AxisOutOfBounds`,
1934 /// `DuplicateAxis`, or `AxisRoleConflict` when dimension numbers are
1935 /// invalid for the operand ranks, and [`Error::RuntimeStateSource`] when
1936 /// output metadata cannot be registered.
1937 ///
1938 /// # Deferred errors
1939 ///
1940 /// Contracting or batch dimensions whose sizes are symbolic are checked
1941 /// when concrete inputs reach compilation or execution. A mismatch is
1942 /// returned as [`Error::TensorRuntime`] with a typed `ShapeMismatch`
1943 /// source and its corresponding [`ErrorPhase`].
1944 pub fn dot_general(
1945 &self,
1946 other: &TracedTensor,
1947 config: DotGeneralConfig,
1948 ) -> Result<TracedTensor> {
1949 config
1950 .validate_dims_with_ranks(self.rank, other.rank)
1951 .map_err(|err| graph_tensor_error("dot_general", err))?;
1952 let lhs_free: Vec<usize> = (0..self.rank)
1953 .filter(|d| {
1954 !config.lhs_contracting_dims.contains(d) && !config.lhs_batch_dims.contains(d)
1955 })
1956 .collect();
1957 let rhs_free: Vec<usize> = (0..other.rank)
1958 .filter(|d| {
1959 !config.rhs_contracting_dims.contains(d) && !config.rhs_batch_dims.contains(d)
1960 })
1961 .collect();
1962 let out_rank = config.lhs_batch_dims.len() + lhs_free.len() + rhs_free.len();
1963 let out_shape_hint = match (&self.shape_hint, &other.shape_hint) {
1964 (Some(lhs_shape), Some(rhs_shape)) => {
1965 let mut out_shape = Vec::with_capacity(out_rank);
1966 for &d in &lhs_free {
1967 out_shape.push(lhs_shape[d].clone());
1968 }
1969 for &d in &rhs_free {
1970 out_shape.push(rhs_shape[d].clone());
1971 }
1972 for &d in &config.lhs_batch_dims {
1973 out_shape.push(lhs_shape[d].clone());
1974 }
1975 Some(out_shape)
1976 }
1977 _ => None,
1978 };
1979
1980 apply_binary(
1981 StdTensorOp::DotGeneral { config },
1982 self,
1983 other,
1984 out_rank,
1985 out_shape_hint,
1986 )
1987 }
1988
1989 /// Matrix multiplication for rank-2 tensors.
1990 ///
1991 /// # Errors
1992 ///
1993 /// Returns [`Error::Validation`] with `RankMismatch` when either operand is
1994 /// not rank 2, `ShapeMismatch::ContractedDimensions` when known matrix
1995 /// dimensions differ, or [`Error::RuntimeStateSource`] when output
1996 /// metadata cannot be registered.
1997 ///
1998 /// # Deferred errors
1999 ///
2000 /// If either contracted dimension is symbolic, the mismatch is discovered
2001 /// at compilation or execution and returned as [`Error::TensorRuntime`]
2002 /// with its typed `ShapeMismatch` source.
2003 pub fn matmul(&self, other: &TracedTensor) -> Result<TracedTensor> {
2004 if self.rank != 2 {
2005 return Err(graph_validation(
2006 "TracedTensor::matmul",
2007 ValidationError::RankMismatch {
2008 expected: 2,
2009 actual: self.rank,
2010 },
2011 ));
2012 }
2013 if other.rank != 2 {
2014 return Err(graph_validation(
2015 "TracedTensor::matmul",
2016 ValidationError::RankMismatch {
2017 expected: 2,
2018 actual: other.rank,
2019 },
2020 ));
2021 }
2022 if let (Some(lhs_shape), Some(rhs_shape)) = (&self.shape_hint, &other.shape_hint)
2023 && let (Some(lhs_cols), Some(rhs_rows)) =
2024 (lhs_shape[1].constant_value(), rhs_shape[0].constant_value())
2025 && lhs_cols != rhs_rows
2026 {
2027 return Err(graph_validation(
2028 "TracedTensor::matmul",
2029 ShapeMismatch::ContractedDimensions {
2030 lhs_axis: 1,
2031 lhs_size: lhs_cols,
2032 rhs_axis: 0,
2033 rhs_size: rhs_rows,
2034 },
2035 ));
2036 }
2037 self.dot_general(
2038 other,
2039 DotGeneralConfig {
2040 lhs_contracting_dims: [1].as_slice().into(),
2041 rhs_contracting_dims: [0].as_slice().into(),
2042 lhs_batch_dims: [].as_slice().into(),
2043 rhs_batch_dims: [].as_slice().into(),
2044 },
2045 )
2046 }
2047
2048 /// Sum over the given axes.
2049 ///
2050 /// # Examples
2051 ///
2052 /// ```rust
2053 /// # use tenferro_runtime::TracedTensor;
2054 /// # let x = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64; 4]).unwrap();
2055 /// let total = x.reduce_sum(None)?;
2056 /// let rows = x.reduce_sum(Some(&[1]))?;
2057 /// let identity = x.reduce_sum(Some(&[]))?;
2058 /// assert_eq!(total.rank, 0);
2059 /// assert_eq!(rows.rank, 1);
2060 /// assert_eq!(identity.rank, 2);
2061 /// # Ok::<(), tenferro_runtime::Error>(())
2062 /// ```
2063 ///
2064 /// # Errors
2065 ///
2066 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when an axis is
2067 /// outside the input rank or `DuplicateAxis` when `axes` repeats an axis,
2068 /// or [`Error::RuntimeStateSource`] when output metadata cannot be
2069 /// registered.
2070 pub fn reduce_sum(&self, axes: Option<&[usize]>) -> Result<TracedTensor> {
2071 let axes = axes.map_or_else(|| (0..self.rank).collect(), <[usize]>::to_vec);
2072 let (out_rank, out_shape_hint) =
2073 reduction_output_meta(self, &axes, "TracedTensor::reduce_sum")?;
2074 apply_unary(
2075 StdTensorOp::ReduceSum { axes },
2076 self,
2077 out_rank,
2078 out_shape_hint,
2079 )
2080 }
2081
2082 /// Sum elementwise squares over the requested axes.
2083 ///
2084 /// Each value is squared in its input dtype before reduction. The initial
2085 /// supported dtypes are `f32` and `f64`; other dtypes return a typed
2086 /// unsupported error during execution. `None` reduces every axis, like
2087 /// the rest of the reduction family; `Some(&[])` returns the elementwise
2088 /// square without reducing rank.
2089 ///
2090 /// This operation is useful when the squared sum is needed directly. Use
2091 /// the linalg norm APIs when a square root or complex magnitude semantics
2092 /// are required.
2093 ///
2094 /// # Errors
2095 ///
2096 /// Returns a typed validation error for invalid axes or a typed
2097 /// runtime-state error while registering output metadata.
2098 ///
2099 /// # Deferred errors
2100 ///
2101 /// Unsupported dtypes and backend execution failures are reported when the
2102 /// compiled graph is executed.
2103 /// # Examples
2104 ///
2105 /// ```rust
2106 /// # use tenferro_runtime::TracedTensor;
2107 /// # let x = TracedTensor::from_vec_col_major(
2108 /// # vec![2, 2],
2109 /// # vec![1.0_f64, 2.0, 3.0, 4.0],
2110 /// # )?;
2111 /// let squares = x.reduce_sum_squares(Some(&[1]))?;
2112 /// assert_eq!(squares.rank, 1);
2113 /// let total = x.reduce_sum_squares(None)?;
2114 /// assert_eq!(total.rank, 0);
2115 /// # Ok::<(), tenferro_runtime::Error>(())
2116 /// ```
2117 pub fn reduce_sum_squares(&self, axes: Option<&[usize]>) -> Result<TracedTensor> {
2118 let axes = axes.map_or_else(|| (0..self.rank).collect(), <[usize]>::to_vec);
2119 let (out_rank, out_shape_hint) =
2120 reduction_output_meta(self, &axes, "TracedTensor::reduce_sum_squares")?;
2121 apply_unary(
2122 StdTensorOp::ReduceSumSquares { axes },
2123 self,
2124 out_rank,
2125 out_shape_hint,
2126 )
2127 }
2128
2129 /// Reduce by taking the maximum along the given axes.
2130 ///
2131 /// Used by tropical (max-plus) compositions: a max-plus reduction over
2132 /// an axis is `ReduceMax` on that axis.
2133 ///
2134 /// # Examples
2135 ///
2136 /// ```rust
2137 /// # use tenferro_runtime::TracedTensor;
2138 /// # let x = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64; 4]).unwrap();
2139 /// let y = x.reduce_max(Some(&[0]))?;
2140 /// # Ok::<(), tenferro_runtime::Error>(())
2141 /// ```
2142 ///
2143 /// # Errors
2144 ///
2145 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when an axis is
2146 /// outside the input rank or `DuplicateAxis` when `axes` repeats an axis,
2147 /// [`Error::Unsupported`] when a non-empty maximum reduction receives a
2148 /// complex dtype, or [`Error::RuntimeStateSource`] when output metadata
2149 /// cannot be registered.
2150 pub fn reduce_max(&self, axes: Option<&[usize]>) -> Result<TracedTensor> {
2151 let axes = axes.map_or_else(|| (0..self.rank).collect(), <[usize]>::to_vec);
2152 let (out_rank, out_shape_hint) =
2153 reduction_output_meta(self, &axes, "TracedTensor::reduce_max")?;
2154 try_apply_unary(
2155 StdTensorOp::ReduceMax { axes },
2156 self,
2157 out_rank,
2158 out_shape_hint,
2159 "TracedTensor::reduce_max",
2160 )
2161 }
2162
2163 /// Reduce by taking the minimum along the given axes.
2164 ///
2165 /// Used by tropical (min-plus) compositions: a min-plus reduction over
2166 /// an axis is `ReduceMin` on that axis.
2167 ///
2168 /// # Examples
2169 ///
2170 /// ```rust
2171 /// # use tenferro_runtime::TracedTensor;
2172 /// # let x = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64; 4]).unwrap();
2173 /// let y = x.reduce_min(Some(&[0]))?;
2174 /// # Ok::<(), tenferro_runtime::Error>(())
2175 /// ```
2176 ///
2177 /// # Errors
2178 ///
2179 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when an axis is
2180 /// outside the input rank or `DuplicateAxis` when `axes` repeats an axis,
2181 /// [`Error::Unsupported`] when a non-empty minimum reduction receives a
2182 /// complex dtype, or [`Error::RuntimeStateSource`] when output metadata
2183 /// cannot be registered.
2184 pub fn reduce_min(&self, axes: Option<&[usize]>) -> Result<TracedTensor> {
2185 let axes = axes.map_or_else(|| (0..self.rank).collect(), <[usize]>::to_vec);
2186 let (out_rank, out_shape_hint) =
2187 reduction_output_meta(self, &axes, "TracedTensor::reduce_min")?;
2188 try_apply_unary(
2189 StdTensorOp::ReduceMin { axes },
2190 self,
2191 out_rank,
2192 out_shape_hint,
2193 "TracedTensor::reduce_min",
2194 )
2195 }
2196
2197 /// Reduce by taking the product along the given axes.
2198 ///
2199 /// # Examples
2200 ///
2201 /// ```rust
2202 /// # use tenferro_runtime::TracedTensor;
2203 /// # let x = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64; 4]).unwrap();
2204 /// let y = x.reduce_prod(Some(&[0]))?;
2205 /// # Ok::<(), tenferro_runtime::Error>(())
2206 /// ```
2207 ///
2208 /// # Errors
2209 ///
2210 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when an axis is
2211 /// outside the input rank or `DuplicateAxis` when `axes` repeats an axis,
2212 /// or [`Error::RuntimeStateSource`] when output metadata cannot be
2213 /// registered.
2214 pub fn reduce_prod(&self, axes: Option<&[usize]>) -> Result<TracedTensor> {
2215 let axes = axes.map_or_else(|| (0..self.rank).collect(), <[usize]>::to_vec);
2216 let (out_rank, out_shape_hint) =
2217 reduction_output_meta(self, &axes, "TracedTensor::reduce_prod")?;
2218 apply_unary(
2219 StdTensorOp::ReduceProd { axes },
2220 self,
2221 out_rank,
2222 out_shape_hint,
2223 )
2224 }
2225
2226 /// Reshape without changing element order.
2227 ///
2228 /// # Examples
2229 ///
2230 /// ```rust
2231 /// # use tenferro_runtime::TracedTensor;
2232 /// # let x = TracedTensor::from_vec_col_major(vec![4], vec![1.0_f64; 4]).unwrap();
2233 /// for y in [x.reshape([2, 2])?, x.reshape(vec![2, 2])?, x.reshape(&[2, 2][..])?] {
2234 /// assert_eq!(y.try_concrete_shape(), Some(vec![2, 2]));
2235 /// }
2236 /// assert!(x.reshape([3]).is_err());
2237 /// # Ok::<(), tenferro_runtime::Error>(())
2238 /// ```
2239 ///
2240 /// # Errors
2241 ///
2242 /// Returns [`Error::Validation`] with `ShapeMismatch::ReshapeElementCount`
2243 /// when a concrete input has a different element count, or
2244 /// `IntegerOverflow` when the target shape product overflows `usize`.
2245 pub fn reshape(&self, shape: impl IntoShapeVec) -> Result<TracedTensor> {
2246 let shape = shape.into_shape_vec();
2247 validate_concrete_reshape_shape(self, &shape)?;
2248 apply_unary_with_dtype(
2249 StdTensorOp::Reshape {
2250 to_shape: DimExpr::from_concrete(&shape),
2251 },
2252 self,
2253 shape.len(),
2254 Some(shape.iter().copied().map(SymDim::from).collect()),
2255 self.dtype,
2256 )
2257 }
2258
2259 /// Return a symbolic expression for the size of one axis, suitable as
2260 /// an `InputDim`-style reference when composing with
2261 /// [`TracedTensor::reshape_sym`].
2262 ///
2263 /// Semantics: if this tensor's `shape_hint` has a symbolic
2264 /// (non-constant) entry for `axis`, that entry is returned
2265 /// verbatim. Otherwise — including when `shape_hint[axis]` is a
2266 /// concrete `SymDim::Concrete(n)` — a
2267 /// `SymDim::tensor_axis(self.id, axis)` reference is returned so the
2268 /// resulting graph remains shape-polymorphic if the same graph is
2269 /// later evaluated against a differently-shaped binding.
2270 ///
2271 /// For a canonical "what is the size of this axis?" query that
2272 /// reports the concrete size when it is known, prefer
2273 /// [`Self::axis_sym_dim`].
2274 ///
2275 /// # Examples
2276 ///
2277 /// ```rust
2278 /// # use tenferro_runtime::TracedTensor;
2279 /// # let x = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
2280 /// let rows = x.sym_size(0)?;
2281 /// let cols = x.sym_size(1)?;
2282 /// let y = x.reshape_sym(&[rows * cols]).unwrap();
2283 /// # Ok::<(), tenferro_runtime::Error>(())
2284 /// ```
2285 ///
2286 /// # Errors
2287 ///
2288 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when `axis` is
2289 /// outside this tensor's rank.
2290 pub fn sym_size(&self, axis: usize) -> Result<SymDim> {
2291 validate_traced_axis(self, axis, "TracedTensor::sym_size")?;
2292 Ok(self
2293 .shape_hint
2294 .as_ref()
2295 .and_then(|shape| shape.get(axis))
2296 .filter(|dim| dim.constant_value().is_none())
2297 .cloned()
2298 .unwrap_or_else(|| SymDim::tensor_axis(self.id, axis)))
2299 }
2300
2301 /// Return the canonical `SymDim` for `axis` — the concrete
2302 /// `SymDim::Concrete(n)` when the size is known, otherwise a symbolic
2303 /// expression identifying this tensor's axis.
2304 ///
2305 /// Unlike [`Self::sym_size`], this method does **not** rewrite
2306 /// concrete axes into `TensorAxis` references. It is the accessor
2307 /// external composition wrappers should use when building mixed
2308 /// concrete/symbolic target shapes for operations like
2309 /// [`Self::broadcast_in_dim_sym`].
2310 ///
2311 /// # Examples
2312 ///
2313 /// ```
2314 /// use tenferro_tensor::DType;
2315 /// use tenferro_runtime::TracedTensor;
2316 ///
2317 /// let a = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
2318 /// // Concrete axis: reports the constant size.
2319 /// assert_eq!(a.axis_sym_dim(0).unwrap().constant_value(), Some(2));
2320 ///
2321 /// let b = TracedTensor::input_symbolic_shape(DType::F64, 2).unwrap();
2322 /// // Fully symbolic leaf: reports a TensorAxis reference.
2323 /// assert!(b.axis_sym_dim(0).unwrap().constant_value().is_none());
2324 /// ```
2325 ///
2326 /// # Errors
2327 ///
2328 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when `axis` is
2329 /// outside this tensor's rank.
2330 pub fn axis_sym_dim(&self, axis: usize) -> Result<SymDim> {
2331 validate_traced_axis(self, axis, "TracedTensor::axis_sym_dim")?;
2332 match self.shape_hint.as_ref().and_then(|shape| shape.get(axis)) {
2333 Some(dim) => Ok(dim.clone()),
2334 None => Ok(SymDim::tensor_axis(self.id, axis)),
2335 }
2336 }
2337
2338 /// Return the full symbolic shape of this tensor when a `shape_hint`
2339 /// is present.
2340 ///
2341 /// Returns `None` for fully-symbolic placeholders produced via
2342 /// [`Self::input_symbolic_shape`] (where `shape_hint` is intentionally
2343 /// absent). For those, build the shape axis-by-axis via
2344 /// [`Self::axis_sym_dim`].
2345 ///
2346 /// # Examples
2347 ///
2348 /// ```
2349 /// use tenferro_tensor::DType;
2350 /// use tenferro_runtime::TracedTensor;
2351 ///
2352 /// let a = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
2353 /// assert!(a.sym_shape().is_some());
2354 /// assert_eq!(a.sym_shape().unwrap().len(), 2);
2355 ///
2356 /// let b = TracedTensor::input_symbolic_shape(DType::F64, 2).unwrap();
2357 /// assert!(b.sym_shape().is_none());
2358 /// ```
2359 pub fn sym_shape(&self) -> Option<&[SymDim]> {
2360 self.shape_hint.as_deref()
2361 }
2362
2363 /// Reshape using symbolic dimensions derived from traced tensor axes.
2364 ///
2365 /// # Examples
2366 ///
2367 /// ```rust
2368 /// # use tenferro_runtime::TracedTensor;
2369 /// # let x = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
2370 /// let rows = x.sym_size(0)?;
2371 /// let cols = x.sym_size(1)?;
2372 /// let y = x.reshape_sym(&[rows * cols]).unwrap();
2373 /// # Ok::<(), tenferro_runtime::Error>(())
2374 /// ```
2375 ///
2376 /// # Errors
2377 ///
2378 /// Returns [`Error::SymbolicShapeConversion`] when a supplied symbolic
2379 /// dimension cannot be mapped to this graph, or [`Error::RuntimeStateSource`]
2380 /// when result metadata cannot be registered.
2381 ///
2382 /// # Deferred errors
2383 ///
2384 /// Element-count compatibility for symbolic dimensions is checked when
2385 /// concrete inputs reach compilation or execution. A mismatch is returned
2386 /// as [`Error::TensorRuntime`] with a typed `ShapeMismatch` source.
2387 pub fn reshape_sym(&self, shape: &[SymDim]) -> Result<TracedTensor> {
2388 let tensor_map = [(self.id, 0usize)];
2389 let to_shape = shape
2390 .iter()
2391 .map(|dim| {
2392 dim.to_dim_expr(&tensor_map)
2393 .map_err(|source| Error::SymbolicShapeConversion {
2394 op: "TracedTensor::reshape_sym",
2395 phase: ErrorPhase::GraphBuild,
2396 source,
2397 })
2398 })
2399 .collect::<Result<Vec<_>>>()?;
2400 let out_shape_hint = Some(shape.to_vec());
2401 apply_unary(
2402 StdTensorOp::Reshape { to_shape },
2403 self,
2404 shape.len(),
2405 out_shape_hint,
2406 )
2407 }
2408
2409 /// Broadcast into a larger shape with explicit dimension placement.
2410 ///
2411 /// # Examples
2412 ///
2413 /// ```rust
2414 /// # use tenferro_runtime::TracedTensor;
2415 /// # let x = TracedTensor::from_vec_col_major(vec![3], vec![1.0_f64; 3]).unwrap();
2416 /// let y = x.broadcast_in_dim(&[2, 3], &[1])?;
2417 /// # Ok::<(), tenferro_runtime::Error>(())
2418 /// ```
2419 ///
2420 /// # Errors
2421 ///
2422 /// Returns [`Error::Validation`] with `RankMismatch` when `dims` does not
2423 /// have one entry per input axis, `AxisOutOfBounds` or `DuplicateAxis` for
2424 /// an invalid output mapping, or `InvalidArgument` when known dimensions
2425 /// cannot broadcast. [`Error::RuntimeStateSource`] reports failure to
2426 /// register the result metadata.
2427 pub fn broadcast_in_dim(&self, shape: &[usize], dims: &[usize]) -> Result<TracedTensor> {
2428 let out_shape_hint: Vec<SymDim> = shape.iter().copied().map(SymDim::from).collect();
2429 validate_broadcast_in_dim_args(
2430 self,
2431 &out_shape_hint,
2432 dims,
2433 "TracedTensor::broadcast_in_dim",
2434 )?;
2435 apply_unary(
2436 StdTensorOp::BroadcastInDim {
2437 shape: DimExpr::from_concrete(shape),
2438 dims: dims.to_vec(),
2439 },
2440 self,
2441 shape.len(),
2442 Some(out_shape_hint),
2443 )
2444 }
2445
2446 /// Broadcast into a symbolic target shape with explicit dimension
2447 /// placement.
2448 ///
2449 /// Unlike [`Self::broadcast_in_dim`], each axis of `shape` is a
2450 /// [`SymDim`], so the target shape can mix concrete sizes (via
2451 /// `SymDim::from(n)`) with symbolic references to this tensor's axes
2452 /// (via [`Self::axis_sym_dim`]) or to axes of other traced tensors.
2453 ///
2454 /// When `shape` contains a `SymDim` that references a traced tensor
2455 /// other than `self`, the referenced tensor(s) must be supplied in
2456 /// `shape_refs`. They are wired into the built op as auxiliary
2457 /// shape-reference inputs — the op does not read their data, only
2458 /// their runtime shape. `shape_refs` must be listed in the same order
2459 /// in which their tensor IDs first appear when walking `shape` after
2460 /// any references to `self`. Usually the simplest correct thing is to
2461 /// pass each unique non-self reference tensor once.
2462 ///
2463 /// # Examples
2464 ///
2465 /// ```
2466 /// use tenferro_runtime::TracedTensor;
2467 ///
2468 /// let a = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
2469 /// let b = TracedTensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
2470 /// let m = a.axis_sym_dim(0)?;
2471 /// let k = a.axis_sym_dim(1)?;
2472 /// let n = b.axis_sym_dim(1)?;
2473 /// // Broadcast `a[m, k]` to `[m, k, n]`, placing `a`'s axes at 0, 1
2474 /// // and taking `n` from `b` as an auxiliary shape reference.
2475 /// let a_b = a.broadcast_in_dim_sym(&[m, k, n], &[0, 1], &[&b])?;
2476 /// assert_eq!(a_b.rank, 3);
2477 /// # Ok::<(), tenferro_runtime::Error>(())
2478 /// ```
2479 ///
2480 /// # Errors
2481 ///
2482 /// Returns [`Error::Validation`] with `RankMismatch`, `AxisOutOfBounds`,
2483 /// `DuplicateAxis`, or `InvalidArgument` when the output mapping or shape
2484 /// references are invalid, [`Error::SymbolicShapeConversion`] for an
2485 /// unmappable symbolic dimension, or [`Error::RuntimeStateSource`] when
2486 /// metadata cannot be registered.
2487 ///
2488 /// # Deferred errors
2489 ///
2490 /// If a symbolic output dimension is smaller than a non-unit input axis,
2491 /// the concrete broadcast check is deferred to compilation or execution
2492 /// and is returned as [`Error::TensorRuntime`] with a typed validation
2493 /// source.
2494 pub fn broadcast_in_dim_sym(
2495 &self,
2496 shape: &[SymDim],
2497 dims: &[usize],
2498 shape_refs: &[&TracedTensor],
2499 ) -> Result<TracedTensor> {
2500 validate_broadcast_in_dim_args(self, shape, dims, "TracedTensor::broadcast_in_dim_sym")?;
2501
2502 // Build a dedup'd list of shape-reference tensors (first occurrence
2503 // wins) and index them starting at 1 — the primary input `self`
2504 // is at 0.
2505 let mut dedup_refs: Vec<&TracedTensor> = Vec::with_capacity(shape_refs.len());
2506 let mut tensor_map: Vec<(u64, usize)> = vec![(self.id, 0)];
2507 for &t in shape_refs {
2508 if !tensor_map.iter().any(|(id, _)| *id == t.id) {
2509 let idx = tensor_map.len();
2510 tensor_map.push((t.id, idx));
2511 dedup_refs.push(t);
2512 }
2513 }
2514
2515 let to_shape: Vec<DimExpr> = shape
2516 .iter()
2517 .map(|dim| {
2518 dim.to_dim_expr(&tensor_map)
2519 .map_err(|source| Error::SymbolicShapeConversion {
2520 op: "broadcast_in_dim_sym",
2521 phase: ErrorPhase::GraphBuild,
2522 source,
2523 })
2524 })
2525 .collect::<Result<Vec<_>>>()?;
2526
2527 // Trim auxiliary shape-reference inputs down to those actually
2528 // used by the generated `DimExpr`s. If the target shape resolved
2529 // to all constants (the concrete-shape case) the op is a plain
2530 // unary broadcast with no extra parents. Otherwise the op needs
2531 // a contiguous prefix of shape-ref inputs covering every
2532 // referenced `input_idx`.
2533 let max_used_idx = DimExpr::max_input_idx_all(&to_shape).unwrap_or(0);
2534 let used_refs: Vec<&TracedTensor> = dedup_refs.into_iter().take(max_used_idx).collect();
2535
2536 let out_shape_hint = Some(shape.to_vec());
2537 apply_unary_with_shape_refs(
2538 StdTensorOp::BroadcastInDim {
2539 shape: to_shape,
2540 dims: dims.to_vec(),
2541 },
2542 self,
2543 &used_refs,
2544 shape.len(),
2545 out_shape_hint,
2546 )
2547 }
2548
2549 /// Slice with explicit start, limit, and stride per axis.
2550 ///
2551 /// # Errors
2552 ///
2553 /// Returns [`Error::Validation`] with `RankMismatch` when the start/limit/
2554 /// stride vectors do not match the input rank, `InvalidSliceStep` when a
2555 /// stride is zero, `InvalidSliceBounds` when a limit precedes its start,
2556 /// or [`Error::RuntimeStateSource`] when output metadata cannot be
2557 /// registered.
2558 pub fn slice(&self, config: SliceConfig) -> Result<TracedTensor> {
2559 let op = StdTensorOp::Slice(config);
2560 let (out_rank, out_shape_hint) =
2561 infer_traced_single_output_shape("TracedTensor::slice", &op, &[self])?;
2562 apply_unary(op, self, out_rank, out_shape_hint)
2563 }
2564
2565 /// Pad with zeros using StableHLO-style edge and interior padding.
2566 ///
2567 /// # Errors
2568 ///
2569 /// Returns [`Error::Validation`] with `RankMismatch` when padding vectors
2570 /// do not match the input rank, `InvalidArgument` for negative interior
2571 /// padding, or `IntegerOverflow` when the padded extent exceeds `usize`.
2572 /// [`Error::RuntimeStateSource`] is returned when output metadata cannot be
2573 /// registered.
2574 pub fn pad(&self, config: PadConfig) -> Result<TracedTensor> {
2575 let op = StdTensorOp::Pad(config);
2576 let (out_rank, out_shape_hint) =
2577 infer_traced_single_output_shape("TracedTensor::pad", &op, &[self])?;
2578 apply_unary(op, self, out_rank, out_shape_hint)
2579 }
2580
2581 /// Reverse the order of elements along the requested axes.
2582 ///
2583 /// # Errors
2584 ///
2585 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when an axis is
2586 /// outside the input rank or `DuplicateAxis` when `axes` repeats one, or
2587 /// [`Error::RuntimeStateSource`] when result metadata cannot be
2588 /// registered.
2589 pub fn reverse(&self, axes: &[usize]) -> Result<TracedTensor> {
2590 validate_traced_axes(self.rank, axes, "TracedTensor::reverse")?;
2591 apply_unary(
2592 StdTensorOp::Reverse {
2593 axes: axes.to_vec(),
2594 },
2595 self,
2596 self.rank,
2597 self.shape_hint.clone(),
2598 )
2599 }
2600
2601 /// Gather slices from `self` using integer start indices.
2602 ///
2603 /// # Errors
2604 ///
2605 /// Returns [`Error::Validation`] with `RankMismatch`, `AxisOutOfBounds`,
2606 /// `DuplicateAxis`, or `ShapeMismatch` when indices or the gather
2607 /// configuration is incompatible with the input, and
2608 /// [`Error::RuntimeStateSource`] when output metadata cannot be
2609 /// registered.
2610 ///
2611 /// # Deferred errors
2612 ///
2613 /// Runtime index values are checked after binding. An out-of-range index
2614 /// is returned as [`Error::TensorRuntime`] with the backend's typed
2615 /// validation source and [`ErrorPhase::Execution`].
2616 pub fn gather(&self, indices: &TracedTensor, config: GatherConfig) -> Result<TracedTensor> {
2617 let op = StdTensorOp::Gather(config);
2618 let (out_rank, out_shape_hint) =
2619 infer_traced_single_output_shape("TracedTensor::gather", &op, &[self, indices])?;
2620 apply_binary_preserve_input_dtypes(op, self, indices, out_rank, out_shape_hint, self.dtype)
2621 }
2622
2623 /// Scatter updates into `self` using StableHLO scatter semantics.
2624 ///
2625 /// # Errors
2626 ///
2627 /// Returns [`Error::Validation`] with `RankMismatch`, `AxisOutOfBounds`,
2628 /// `DuplicateAxis`, or `ShapeMismatch` when indices, updates, or the
2629 /// scatter configuration is incompatible, [`Error::TensorRuntime`] with
2630 /// `UnsupportedDTypeConversion` when dtype promotion cannot be
2631 /// represented, or [`Error::RuntimeStateSource`] when output metadata
2632 /// cannot be registered.
2633 ///
2634 /// # Deferred errors
2635 ///
2636 /// Runtime index/update values are checked after binding. An invalid
2637 /// index or update shape is returned as [`Error::TensorRuntime`] with its
2638 /// typed validation source and [`ErrorPhase::Execution`].
2639 pub fn scatter(
2640 &self,
2641 indices: &TracedTensor,
2642 updates: &TracedTensor,
2643 config: ScatterConfig,
2644 ) -> Result<TracedTensor> {
2645 let op = StdTensorOp::Scatter(config);
2646 let (out_rank, out_shape_hint) = infer_traced_single_output_shape(
2647 "TracedTensor::scatter",
2648 &op,
2649 &[self, indices, updates],
2650 )?;
2651 let out_dtype = crate::shape_infer::promote_dtype(self.dtype, updates.dtype);
2652 let operand = if self.dtype != out_dtype {
2653 self.cast(out_dtype)?
2654 } else {
2655 self.clone()
2656 };
2657 let updates = if updates.dtype != out_dtype {
2658 updates.cast(out_dtype)?
2659 } else {
2660 updates.clone()
2661 };
2662 apply_ternary_with_output_dtype(
2663 op,
2664 &operand,
2665 indices,
2666 &updates,
2667 out_rank,
2668 out_shape_hint,
2669 out_dtype,
2670 )
2671 }
2672
2673 /// Slice using runtime start indices.
2674 ///
2675 /// # Errors
2676 ///
2677 /// Returns [`Error::Validation`] with `RankMismatch`, `AxisOutOfBounds`,
2678 /// or `InvalidArgument` when `starts` or `sizes` has an incompatible rank
2679 /// or extent, and [`Error::RuntimeStateSource`] when output metadata cannot
2680 /// be registered.
2681 ///
2682 /// # Deferred errors
2683 ///
2684 /// Runtime start values are checked after binding. An out-of-range start
2685 /// is returned as [`Error::TensorRuntime`] with the backend's typed
2686 /// validation source and [`ErrorPhase::Execution`].
2687 pub fn dynamic_slice(&self, starts: &TracedTensor, sizes: &[usize]) -> Result<TracedTensor> {
2688 let op = StdTensorOp::DynamicSlice {
2689 slice_sizes: sizes.to_vec(),
2690 };
2691 let (out_rank, out_shape_hint) =
2692 infer_traced_single_output_shape("TracedTensor::dynamic_slice", &op, &[self, starts])?;
2693 apply_binary_preserve_input_dtypes(op, self, starts, out_rank, out_shape_hint, self.dtype)
2694 }
2695
2696 /// Keep the lower triangle and zero the rest.
2697 ///
2698 /// # Examples
2699 ///
2700 /// ```rust
2701 /// # use tenferro_runtime::TracedTensor;
2702 /// let matrix = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64; 4])?;
2703 /// let lower = matrix.tril(0)?;
2704 /// assert_eq!(lower.rank, 2);
2705 /// # Ok::<(), tenferro_runtime::Error>(())
2706 /// ```
2707 ///
2708 /// # Errors
2709 ///
2710 /// Returns [`Error::RuntimeStateSource`] when traced output metadata
2711 /// registration is unavailable or inconsistent with the graph.
2712 pub fn tril(&self, k: i64) -> Result<TracedTensor> {
2713 apply_unary(
2714 StdTensorOp::Tril { k },
2715 self,
2716 self.rank,
2717 self.shape_hint.clone(),
2718 )
2719 }
2720
2721 /// Keep the upper triangle and zero the rest.
2722 ///
2723 /// # Examples
2724 ///
2725 /// ```rust
2726 /// # use tenferro_runtime::TracedTensor;
2727 /// let matrix = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64; 4])?;
2728 /// let upper = matrix.triu(0)?;
2729 /// assert_eq!(upper.rank, 2);
2730 /// # Ok::<(), tenferro_runtime::Error>(())
2731 /// ```
2732 ///
2733 /// # Errors
2734 ///
2735 /// Returns [`Error::RuntimeStateSource`] when traced output metadata
2736 /// registration is unavailable or inconsistent with the graph.
2737 pub fn triu(&self, k: i64) -> Result<TracedTensor> {
2738 apply_unary(
2739 StdTensorOp::Triu { k },
2740 self,
2741 self.rank,
2742 self.shape_hint.clone(),
2743 )
2744 }
2745
2746 /// Permute tensor axes.
2747 ///
2748 /// # Examples
2749 ///
2750 /// ```rust
2751 /// # use tenferro_runtime::TracedTensor;
2752 /// # let x = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
2753 /// let y = x.transpose(&[1, 0])?;
2754 /// # Ok::<(), tenferro_runtime::Error>(())
2755 /// ```
2756 ///
2757 /// # Errors
2758 ///
2759 /// Returns [`Error::Validation`] with `InvalidPermutationLength`,
2760 /// `AxisOutOfBounds`, or `DuplicateAxis` when `perm` is not a valid
2761 /// permutation of the tensor axes, or [`Error::RuntimeStateSource`] when
2762 /// output metadata registration fails.
2763 pub fn transpose(&self, perm: &[usize]) -> Result<TracedTensor> {
2764 validate_traced_perm(self.rank, perm, "TracedTensor::transpose")?;
2765 let out_shape_hint = self
2766 .shape_hint
2767 .as_ref()
2768 .map(|shape| perm.iter().map(|&p| shape[p].clone()).collect());
2769 apply_unary(
2770 StdTensorOp::Transpose {
2771 perm: perm.to_vec(),
2772 },
2773 self,
2774 self.rank,
2775 out_shape_hint,
2776 )
2777 }
2778
2779 /// Extract the diagonal along two axes.
2780 ///
2781 /// # Examples
2782 ///
2783 /// ```rust
2784 /// # use tenferro_runtime::TracedTensor;
2785 /// # let x = TracedTensor::from_vec_col_major(vec![2, 2], vec![1.0_f64; 4]).unwrap();
2786 /// let y = x.extract_diag(0, 1)?;
2787 /// # Ok::<(), tenferro_runtime::Error>(())
2788 /// ```
2789 ///
2790 /// # Errors
2791 ///
2792 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when either axis
2793 /// is outside the input rank or `InvalidArgument` when `axis_a == axis_b`.
2794 pub fn extract_diag(&self, axis_a: usize, axis_b: usize) -> Result<TracedTensor> {
2795 validate_traced_axis(self, axis_a, "TracedTensor::extract_diag")?;
2796 validate_traced_axis(self, axis_b, "TracedTensor::extract_diag")?;
2797 if axis_a == axis_b {
2798 return Err(graph_invalid_argument(
2799 "TracedTensor::extract_diag",
2800 "axes",
2801 "diagonal axes must be distinct",
2802 ));
2803 }
2804 let op = StdTensorOp::ExtractDiag { axis_a, axis_b };
2805 let (out_rank, out_shape_hint) =
2806 infer_traced_single_output_shape("TracedTensor::extract_diag", &op, &[self])?;
2807 apply_unary(op, self, out_rank, out_shape_hint)
2808 }
2809
2810 /// Embed a vector or lower-rank tensor along a diagonal.
2811 ///
2812 /// # Examples
2813 ///
2814 /// ```rust
2815 /// # use tenferro_runtime::TracedTensor;
2816 /// # let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64; 2]).unwrap();
2817 /// let y = x.embed_diag(0, 1)?;
2818 /// # Ok::<(), tenferro_runtime::Error>(())
2819 /// ```
2820 ///
2821 /// # Errors
2822 ///
2823 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when `axis_a` is
2824 /// outside the input rank or `InvalidArgument` when `axis_b` is not a
2825 /// valid insertion axis.
2826 pub fn embed_diag(&self, axis_a: usize, axis_b: usize) -> Result<TracedTensor> {
2827 validate_traced_axis(self, axis_a, "TracedTensor::embed_diag")?;
2828 validate_traced_insert_axis(self.rank, axis_b, "TracedTensor::embed_diag")?;
2829 let out_shape_hint = self.shape_hint.as_ref().map(|shape| {
2830 let mut out_shape = shape.clone();
2831 out_shape.insert(axis_b, shape[axis_a].clone());
2832 out_shape
2833 });
2834 apply_unary(
2835 StdTensorOp::EmbedDiag { axis_a, axis_b },
2836 self,
2837 self.rank + 1,
2838 out_shape_hint,
2839 )
2840 }
2841
2842 /// Return the runtime size of one axis as a scalar `f64` tensor.
2843 ///
2844 /// The result is metadata-derived and therefore has no gradient.
2845 ///
2846 /// # Examples
2847 ///
2848 /// ```
2849 /// use tenferro_cpu::CpuBackend;
2850 /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
2851 ///
2852 /// let x = TracedTensor::from_vec_col_major(vec![2, 3], vec![1.0_f64, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap();
2853 /// let cols = x.shape_of(1)?;
2854 /// let mut compiler = GraphCompiler::new();
2855 /// let program = compiler.compile(&cols).unwrap();
2856 /// let backend = CpuBackend::new();
2857 /// let mut builder = Runtime::builder();
2858 /// builder
2859 /// .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
2860 /// .unwrap();
2861 /// let runtime = builder.build().unwrap();
2862 /// let outputs = runtime.run_compiled(&program, &[]).unwrap();
2863 /// let out = &outputs[0];
2864 /// assert_eq!(out.shape(), &[] as &[usize]);
2865 /// # Ok::<(), tenferro_runtime::Error>(())
2866 /// ```
2867 ///
2868 /// # Errors
2869 ///
2870 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when `axis` is
2871 /// outside the input rank, or [`Error::RuntimeStateSource`] when scalar
2872 /// output metadata cannot be registered.
2873 pub fn shape_of(&self, axis: usize) -> Result<TracedTensor> {
2874 validate_traced_axis(self, axis, "TracedTensor::shape_of")?;
2875 apply_unary_with_dtype(
2876 StdTensorOp::ShapeOf { axis },
2877 self,
2878 0,
2879 Some(vec![]),
2880 DType::F64,
2881 )
2882 }
2883
2884 /// Truncate this tensor along `axis` to the first `size` elements.
2885 ///
2886 /// `size` is read at runtime from a scalar traced tensor. Values are
2887 /// rounded to the nearest integer, clamped to `[0, self.shape[axis]]`,
2888 /// and the output keeps the same element dtype as the input.
2889 ///
2890 /// # Examples
2891 ///
2892 /// ```
2893 /// use tenferro_cpu::CpuBackend;
2894 /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
2895 ///
2896 /// let x = TracedTensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap();
2897 /// let size = TracedTensor::from_vec_col_major(vec![], vec![2.0_f64]).unwrap();
2898 /// let y = x.dynamic_truncate(&size, 0)?;
2899 /// let mut compiler = GraphCompiler::new();
2900 /// let program = compiler.compile(&y).unwrap();
2901 /// let backend = CpuBackend::new();
2902 /// let mut builder = Runtime::builder();
2903 /// builder
2904 /// .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
2905 /// .unwrap();
2906 /// let runtime = builder.build().unwrap();
2907 /// let outputs = runtime.run_compiled(&program, &[]).unwrap();
2908 /// let out = &outputs[0];
2909 /// assert_eq!(out.shape(), &[2]);
2910 /// # Ok::<(), tenferro_runtime::Error>(())
2911 /// ```
2912 ///
2913 /// # Errors
2914 ///
2915 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when `axis` is
2916 /// outside the input rank or `RankMismatch` when `size` is not scalar.
2917 ///
2918 /// # Deferred errors
2919 ///
2920 /// At execution, non-`f32`/`f64`/`i64` size dtypes return
2921 /// [`Error::TensorRuntime`] with `Unsupported`, non-finite size values
2922 /// return a typed `InvalidArgument`, and an empty scalar buffer returns a
2923 /// typed runtime-state source.
2924 pub fn dynamic_truncate(&self, size: &TracedTensor, axis: usize) -> Result<TracedTensor> {
2925 validate_traced_axis(self, axis, "TracedTensor::dynamic_truncate")?;
2926 if size.rank != 0 {
2927 return Err(graph_validation(
2928 "TracedTensor::dynamic_truncate",
2929 ValidationError::RankMismatch {
2930 expected: 0,
2931 actual: size.rank,
2932 },
2933 ));
2934 }
2935 apply_binary_preserve_input_dtypes(
2936 StdTensorOp::DynamicTruncate { axis },
2937 self,
2938 size,
2939 self.rank,
2940 None,
2941 self.dtype,
2942 )
2943 }
2944
2945 /// Pad this tensor with zeros along `axis` to match `reference.shape[axis]`.
2946 ///
2947 /// If `reference` is smaller along that axis, this is a no-op.
2948 ///
2949 /// # Examples
2950 ///
2951 /// ```
2952 /// use tenferro_cpu::CpuBackend;
2953 /// use tenferro_runtime::{GraphCompiler, Runtime, TracedTensor};
2954 ///
2955 /// let x = TracedTensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap();
2956 /// let reference = TracedTensor::from_vec_col_major(vec![4], vec![0.0_f64, 0.0, 0.0, 0.0]).unwrap();
2957 /// let y = x.pad_to_match(&reference, 0)?;
2958 /// let mut compiler = GraphCompiler::new();
2959 /// let program = compiler.compile(&y).unwrap();
2960 /// let backend = CpuBackend::new();
2961 /// let mut builder = Runtime::builder();
2962 /// builder
2963 /// .register_engine(tenferro_cpu::runtime_engine_registration(&backend).unwrap())
2964 /// .unwrap();
2965 /// let runtime = builder.build().unwrap();
2966 /// let outputs = runtime.run_compiled(&program, &[]).unwrap();
2967 /// let out = &outputs[0];
2968 /// assert_eq!(out.shape(), &[4]);
2969 /// # Ok::<(), tenferro_runtime::Error>(())
2970 /// ```
2971 ///
2972 /// # Errors
2973 ///
2974 /// Returns [`Error::Validation`] with `AxisOutOfBounds` when `axis` is
2975 /// outside either tensor's rank, or [`Error::RuntimeStateSource`] when
2976 /// output metadata cannot be registered.
2977 pub fn pad_to_match(&self, reference: &TracedTensor, axis: usize) -> Result<TracedTensor> {
2978 validate_traced_axis(self, axis, "TracedTensor::pad_to_match")?;
2979 validate_traced_axis(reference, axis, "TracedTensor::pad_to_match")?;
2980 let op = StdTensorOp::PadToMatch { axis };
2981 let (out_rank, out_shape_hint) = infer_traced_single_output_shape(
2982 "TracedTensor::pad_to_match",
2983 &op,
2984 &[self, reference],
2985 )?;
2986 apply_binary_preserve_input_dtypes(
2987 op,
2988 self,
2989 reference,
2990 out_rank,
2991 out_shape_hint,
2992 self.dtype,
2993 )
2994 }
2995}
2996
2997pub(crate) fn apply_unary(
2998 op: StdTensorOp,
2999 input: &TracedTensor,
3000 out_rank: usize,
3001 out_shape_hint: Option<Vec<SymDim>>,
3002) -> Result<TracedTensor> {
3003 let out_dtype = try_inferred_output_dtype(&op, &[input.dtype], "apply_unary")?;
3004 apply_unary_with_dtype(op, input, out_rank, out_shape_hint, out_dtype)
3005}
3006
3007fn try_apply_unary(
3008 op: StdTensorOp,
3009 input: &TracedTensor,
3010 out_rank: usize,
3011 out_shape_hint: Option<Vec<SymDim>>,
3012 context: &'static str,
3013) -> Result<TracedTensor> {
3014 let out_dtype = try_inferred_output_dtype(&op, &[input.dtype], context)?;
3015 apply_unary_with_dtype(op, input, out_rank, out_shape_hint, out_dtype)
3016}
3017
3018pub(crate) fn apply_unary_with_dtype(
3019 op: StdTensorOp,
3020 input: &TracedTensor,
3021 out_rank: usize,
3022 out_shape_hint: Option<Vec<SymDim>>,
3023 out_dtype: DType,
3024) -> Result<TracedTensor> {
3025 let mut builder = GraphBuilder::new();
3026 builder.add_parent(input.graph.clone());
3027 let input_ref = ValueRef::External(input.graph.values()[input.val].key.clone());
3028 let outputs = builder.add_operation(op, vec![input_ref], OperationRole::Primary);
3029 builder.set_outputs(outputs.clone());
3030 let graph = Arc::new(builder.build());
3031 let metadata_scope =
3032 register_single_output_metadata(graph.as_ref(), outputs[0], out_dtype, &out_shape_hint)?;
3033
3034 Ok(TracedTensor {
3035 id: next_traced_id(),
3036 rank: out_rank,
3037 dtype: out_dtype,
3038 graph,
3039 val: outputs[0],
3040 data: None,
3041 shape_hint: out_shape_hint,
3042 inputs_map: input.inputs_map.clone(),
3043 leaf_metas: input.leaf_metas.clone(),
3044 extra_roots: input.extra_roots.clone(),
3045 checkpoint_chain: input.checkpoint_chain.clone(),
3046 metadata_scopes: MetadataScopeChain::with_new(metadata_scope, [&input.metadata_scopes]),
3047 constraint_scopes: input.constraint_scopes.clone(),
3048 })
3049}
3050
3051/// Apply a unary-primary op that additionally references one or more
3052/// tensors for shape resolution only.
3053///
3054/// The primary `input` becomes op input 0; each tensor in `shape_refs`
3055/// becomes op input 1, 2, … in order. Used by
3056/// [`TracedTensor::broadcast_in_dim_sym`] when the target shape
3057/// references axes of tensors other than the primary input; the op
3058/// reads only their runtime shape, not their data.
3059pub(crate) fn apply_unary_with_shape_refs(
3060 op: StdTensorOp,
3061 input: &TracedTensor,
3062 shape_refs: &[&TracedTensor],
3063 out_rank: usize,
3064 out_shape_hint: Option<Vec<SymDim>>,
3065) -> Result<TracedTensor> {
3066 let mut builder = GraphBuilder::new();
3067 builder.add_parent(input.graph.clone());
3068 for t in shape_refs {
3069 builder.add_parent(t.graph.clone());
3070 }
3071 let mut op_inputs: Vec<ValueRef<StdTensorOp>> = Vec::with_capacity(1 + shape_refs.len());
3072 op_inputs.push(ValueRef::External(
3073 input.graph.values()[input.val].key.clone(),
3074 ));
3075 for t in shape_refs {
3076 op_inputs.push(ValueRef::External(t.graph.values()[t.val].key.clone()));
3077 }
3078 let outputs = builder.add_operation(op, op_inputs, OperationRole::Primary);
3079 builder.set_outputs(outputs.clone());
3080 let graph = Arc::new(builder.build());
3081 let metadata_scope =
3082 register_single_output_metadata(graph.as_ref(), outputs[0], input.dtype, &out_shape_hint)?;
3083
3084 let inputs_map =
3085 merge_traced_inputs_map(std::iter::once(input).chain(shape_refs.iter().copied()));
3086 let leaf_metas =
3087 merge_traced_leaf_metas(std::iter::once(input).chain(shape_refs.iter().copied()));
3088
3089 let mut extra_roots = input.extra_roots.clone();
3090 for t in shape_refs {
3091 extra_roots.extend(t.extra_roots.iter().cloned());
3092 }
3093
3094 let mut checkpoint_chain = input.checkpoint_chain.clone();
3095 for t in shape_refs {
3096 checkpoint_chain =
3097 CheckpointNode::merge_chains(checkpoint_chain, t.checkpoint_chain.clone());
3098 }
3099
3100 Ok(TracedTensor {
3101 id: next_traced_id(),
3102 rank: out_rank,
3103 dtype: input.dtype,
3104 graph,
3105 val: outputs[0],
3106 data: None,
3107 shape_hint: out_shape_hint,
3108 inputs_map,
3109 leaf_metas,
3110 extra_roots,
3111 checkpoint_chain,
3112 metadata_scopes: MetadataScopeChain::with_new(
3113 metadata_scope,
3114 std::iter::once(&input.metadata_scopes)
3115 .chain(shape_refs.iter().map(|tensor| &tensor.metadata_scopes)),
3116 ),
3117 constraint_scopes: ConstraintScopeChain::merge(
3118 std::iter::once(&input.constraint_scopes)
3119 .chain(shape_refs.iter().map(|tensor| &tensor.constraint_scopes)),
3120 ),
3121 })
3122}
3123
3124pub(crate) fn apply_nullary(
3125 op: StdTensorOp,
3126 rank: usize,
3127 dtype: DType,
3128 shape_hint: Option<Vec<SymDim>>,
3129) -> Result<TracedTensor> {
3130 let mut builder = GraphBuilder::new();
3131 let outputs = builder.add_operation(op, vec![], OperationRole::Primary);
3132 builder.set_outputs(outputs.clone());
3133 let graph = Arc::new(builder.build());
3134 let metadata_scope =
3135 register_single_output_metadata(graph.as_ref(), outputs[0], dtype, &shape_hint)?;
3136
3137 Ok(TracedTensor {
3138 id: next_traced_id(),
3139 rank,
3140 dtype,
3141 graph,
3142 val: outputs[0],
3143 data: None,
3144 shape_hint,
3145 inputs_map: Arc::new(HashMap::new()),
3146 leaf_metas: Arc::new(HashMap::new()),
3147 extra_roots: Vec::new(),
3148 checkpoint_chain: None,
3149 metadata_scopes: MetadataScopeChain::from_scope(metadata_scope),
3150 constraint_scopes: ConstraintScopeChain::empty(),
3151 })
3152}
3153
3154pub(crate) fn apply_binary(
3155 op: StdTensorOp,
3156 lhs: &TracedTensor,
3157 rhs: &TracedTensor,
3158 out_rank: usize,
3159 out_shape_hint: Option<Vec<SymDim>>,
3160) -> Result<TracedTensor> {
3161 let input_dtype = crate::shape_infer::promote_dtype_for_binary_op(&op, lhs.dtype, rhs.dtype);
3162 let out_dtype = try_inferred_output_dtype(&op, &[lhs.dtype, rhs.dtype], "apply_binary")?;
3163
3164 // Insert Convert ops when an input dtype differs from the primitive input dtype.
3165 let lhs = if lhs.dtype != input_dtype {
3166 lhs.cast(input_dtype)?
3167 } else {
3168 lhs.clone()
3169 };
3170 let rhs = if rhs.dtype != input_dtype {
3171 rhs.cast(input_dtype)?
3172 } else {
3173 rhs.clone()
3174 };
3175
3176 apply_binary_with_output_dtype(op, &lhs, &rhs, out_rank, out_shape_hint, out_dtype)
3177}
3178
3179fn try_apply_binary(
3180 op: StdTensorOp,
3181 lhs: &TracedTensor,
3182 rhs: &TracedTensor,
3183 out_rank: usize,
3184 out_shape_hint: Option<Vec<SymDim>>,
3185 context: &'static str,
3186) -> Result<TracedTensor> {
3187 let input_dtype = crate::shape_infer::promote_dtype_for_binary_op(&op, lhs.dtype, rhs.dtype);
3188 let out_dtype = try_inferred_output_dtype(&op, &[lhs.dtype, rhs.dtype], context)?;
3189
3190 let lhs = if lhs.dtype != input_dtype {
3191 lhs.cast(input_dtype)?
3192 } else {
3193 lhs.clone()
3194 };
3195 let rhs = if rhs.dtype != input_dtype {
3196 rhs.cast(input_dtype)?
3197 } else {
3198 rhs.clone()
3199 };
3200
3201 apply_binary_with_output_dtype(op, &lhs, &rhs, out_rank, out_shape_hint, out_dtype)
3202}
3203
3204pub(crate) fn apply_binary_preserve_input_dtypes(
3205 op: StdTensorOp,
3206 lhs: &TracedTensor,
3207 rhs: &TracedTensor,
3208 out_rank: usize,
3209 out_shape_hint: Option<Vec<SymDim>>,
3210 out_dtype: DType,
3211) -> Result<TracedTensor> {
3212 apply_binary_with_output_dtype(op, lhs, rhs, out_rank, out_shape_hint, out_dtype)
3213}
3214
3215pub(crate) fn apply_broadcast_binary_op(
3216 op: StdTensorOp,
3217 lhs: &TracedTensor,
3218 rhs: &TracedTensor,
3219) -> Result<TracedTensor> {
3220 let (lhs, rhs) = broadcast_binary(lhs, rhs)?;
3221 try_apply_binary(
3222 op,
3223 &lhs,
3224 &rhs,
3225 lhs.rank,
3226 lhs.shape_hint.clone(),
3227 "broadcast_binary",
3228 )
3229}
3230
3231pub(crate) fn apply_broadcast_ternary_op(
3232 op: StdTensorOp,
3233 first: &TracedTensor,
3234 second: &TracedTensor,
3235 third: &TracedTensor,
3236) -> Result<TracedTensor> {
3237 let (first, second, third) = broadcast_ternary(first, second, third)?;
3238 try_apply_ternary(
3239 op,
3240 &first,
3241 &second,
3242 &third,
3243 first.rank,
3244 first.shape_hint.clone(),
3245 "broadcast_ternary",
3246 )
3247}
3248
3249fn try_apply_ternary(
3250 op: StdTensorOp,
3251 first: &TracedTensor,
3252 second: &TracedTensor,
3253 third: &TracedTensor,
3254 out_rank: usize,
3255 out_shape_hint: Option<Vec<SymDim>>,
3256 context: &'static str,
3257) -> Result<TracedTensor> {
3258 let out_dtype =
3259 try_inferred_output_dtype(&op, &[first.dtype, second.dtype, third.dtype], context)?;
3260 let (first, second, third) = match op {
3261 StdTensorOp::Select => {
3262 let value_dtype = crate::shape_infer::promote_dtype(second.dtype, third.dtype);
3263 let second = if second.dtype != value_dtype {
3264 second.cast(value_dtype)?
3265 } else {
3266 second.clone()
3267 };
3268 let third = if third.dtype != value_dtype {
3269 third.cast(value_dtype)?
3270 } else {
3271 third.clone()
3272 };
3273 (first.clone(), second, third)
3274 }
3275 _ => {
3276 let input_dtype =
3277 crate::shape_infer::promote_dtypes([first.dtype, second.dtype, third.dtype]);
3278 let first = if first.dtype != input_dtype {
3279 first.cast(input_dtype)?
3280 } else {
3281 first.clone()
3282 };
3283 let second = if second.dtype != input_dtype {
3284 second.cast(input_dtype)?
3285 } else {
3286 second.clone()
3287 };
3288 let third = if third.dtype != input_dtype {
3289 third.cast(input_dtype)?
3290 } else {
3291 third.clone()
3292 };
3293 (first, second, third)
3294 }
3295 };
3296 apply_ternary_with_output_dtype(
3297 op,
3298 &first,
3299 &second,
3300 &third,
3301 out_rank,
3302 out_shape_hint,
3303 out_dtype,
3304 )
3305}
3306
3307fn apply_binary_with_output_dtype(
3308 op: StdTensorOp,
3309 lhs: &TracedTensor,
3310 rhs: &TracedTensor,
3311 out_rank: usize,
3312 out_shape_hint: Option<Vec<SymDim>>,
3313 out_dtype: DType,
3314) -> Result<TracedTensor> {
3315 let lhs_ref = ValueRef::External(lhs.graph.values()[lhs.val].key.clone());
3316 let rhs_ref = ValueRef::External(rhs.graph.values()[rhs.val].key.clone());
3317
3318 let mut builder = GraphBuilder::new();
3319 builder.add_parent(lhs.graph.clone());
3320 builder.add_parent(rhs.graph.clone());
3321 let outputs = builder.add_operation(op, vec![lhs_ref, rhs_ref], OperationRole::Primary);
3322 builder.set_outputs(outputs.clone());
3323 let graph = Arc::new(builder.build());
3324 let metadata_scope =
3325 register_single_output_metadata(graph.as_ref(), outputs[0], out_dtype, &out_shape_hint)?;
3326
3327 let mut extra_roots = lhs.extra_roots.clone();
3328 extra_roots.extend(rhs.extra_roots.iter().cloned());
3329
3330 Ok(TracedTensor {
3331 id: next_traced_id(),
3332 rank: out_rank,
3333 dtype: out_dtype,
3334 graph,
3335 val: outputs[0],
3336 data: None,
3337 shape_hint: out_shape_hint,
3338 inputs_map: merge_traced_inputs_map([lhs, rhs]),
3339 leaf_metas: merge_traced_leaf_metas([lhs, rhs]),
3340 extra_roots,
3341 checkpoint_chain: CheckpointNode::merge_chains(
3342 lhs.checkpoint_chain.clone(),
3343 rhs.checkpoint_chain.clone(),
3344 ),
3345 metadata_scopes: MetadataScopeChain::with_new(
3346 metadata_scope,
3347 [&lhs.metadata_scopes, &rhs.metadata_scopes],
3348 ),
3349 constraint_scopes: ConstraintScopeChain::merge([
3350 &lhs.constraint_scopes,
3351 &rhs.constraint_scopes,
3352 ]),
3353 })
3354}
3355
3356fn apply_ternary_with_output_dtype(
3357 op: StdTensorOp,
3358 first: &TracedTensor,
3359 second: &TracedTensor,
3360 third: &TracedTensor,
3361 out_rank: usize,
3362 out_shape_hint: Option<Vec<SymDim>>,
3363 out_dtype: DType,
3364) -> Result<TracedTensor> {
3365 let first_ref = ValueRef::External(first.graph.values()[first.val].key.clone());
3366 let second_ref = ValueRef::External(second.graph.values()[second.val].key.clone());
3367 let third_ref = ValueRef::External(third.graph.values()[third.val].key.clone());
3368
3369 let mut builder = GraphBuilder::new();
3370 builder.add_parent(first.graph.clone());
3371 builder.add_parent(second.graph.clone());
3372 builder.add_parent(third.graph.clone());
3373 let outputs = builder.add_operation(
3374 op,
3375 vec![first_ref, second_ref, third_ref],
3376 OperationRole::Primary,
3377 );
3378 builder.set_outputs(outputs.clone());
3379 let graph = Arc::new(builder.build());
3380 let metadata_scope =
3381 register_single_output_metadata(graph.as_ref(), outputs[0], out_dtype, &out_shape_hint)?;
3382
3383 let mut extra_roots = first.extra_roots.clone();
3384 extra_roots.extend(second.extra_roots.iter().cloned());
3385 extra_roots.extend(third.extra_roots.iter().cloned());
3386
3387 let checkpoint_chain = CheckpointNode::merge_chains(
3388 CheckpointNode::merge_chains(
3389 first.checkpoint_chain.clone(),
3390 second.checkpoint_chain.clone(),
3391 ),
3392 third.checkpoint_chain.clone(),
3393 );
3394
3395 Ok(TracedTensor {
3396 id: next_traced_id(),
3397 rank: out_rank,
3398 dtype: out_dtype,
3399 graph,
3400 val: outputs[0],
3401 data: None,
3402 shape_hint: out_shape_hint,
3403 inputs_map: merge_traced_inputs_map([first, second, third]),
3404 leaf_metas: merge_traced_leaf_metas([first, second, third]),
3405 extra_roots,
3406 checkpoint_chain,
3407 metadata_scopes: MetadataScopeChain::with_new(
3408 metadata_scope,
3409 [
3410 &first.metadata_scopes,
3411 &second.metadata_scopes,
3412 &third.metadata_scopes,
3413 ],
3414 ),
3415 constraint_scopes: ConstraintScopeChain::merge([
3416 &first.constraint_scopes,
3417 &second.constraint_scopes,
3418 &third.constraint_scopes,
3419 ]),
3420 })
3421}
3422
3423fn register_single_output_metadata(
3424 graph: &Graph<StdTensorOp>,
3425 output: LocalValueId,
3426 dtype: DType,
3427 shape_hint: &Option<Vec<SymDim>>,
3428) -> Result<GlobalMetadataScope> {
3429 if let Some(shape) = shape_hint {
3430 // Fresh graph output keys are generated in this builder, so metadata
3431 // registration failure would indicate a global metadata invariant bug.
3432 register_metadata_or_runtime_state(register_scoped_value_metadata(
3433 graph.values()[output].key.clone(),
3434 tensor_meta(dtype, shape.clone()),
3435 ))
3436 } else {
3437 // Fresh graph output keys are generated in this builder, so metadata
3438 // registration failure would indicate a global metadata invariant bug.
3439 register_metadata_or_runtime_state(register_scoped_graph_metadata(
3440 graph,
3441 std::iter::empty(),
3442 ))
3443 }
3444}
3445
3446impl TracedTensor {
3447 pub(crate) fn resolve_roots(&self) -> Vec<Arc<Graph<StdTensorOp>>> {
3448 let mut roots = Vec::with_capacity(1 + self.extra_roots.len());
3449 roots.push(self.graph.clone());
3450 roots.extend(self.extra_roots.iter().cloned());
3451 roots
3452 }
3453}
3454
3455mod composite_ops;
3456
3457#[cfg(test)]
3458mod tests;