Skip to main content

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;