Skip to main content

tenferro_ad/
eager.rs

1use std::borrow::Cow;
2use std::cell::{Cell, RefCell};
3use std::cmp::Reverse;
4use std::collections::HashMap;
5use std::env;
6use std::fmt;
7use std::marker::PhantomData;
8use std::mem::{size_of, size_of_val};
9use std::rc::Rc;
10#[cfg(test)]
11use std::sync::atomic::{AtomicUsize, Ordering};
12use std::sync::{Arc, Mutex, MutexGuard, OnceLock, Weak};
13use std::time::{Duration, Instant};
14
15use lru::LruCache;
16use num_complex::Complex64;
17
18use crate::extension::{
19    validate_eager_extension_target, EagerExtensionBackendKind, EagerExtensionTarget,
20};
21use crate::extension_cache::{ExtensionCacheLimits, ExtensionCacheSelector, ExtensionCacheStore};
22#[cfg(test)]
23use computegraph::graph::Graph;
24use computegraph::ValueKey;
25#[cfg(test)]
26use computegraph::ValueRef;
27use tenferro_cpu::{CpuBackend, CpuBackendError, CpuPlacement};
28#[cfg(feature = "cuda")]
29use tenferro_gpu::cuda::CudaBackend;
30#[cfg(feature = "webgpu")]
31use tenferro_gpu::webgpu::WebGpuBackend;
32#[cfg(test)]
33use tenferro_ops::input_key::TensorInputKey;
34use tenferro_ops::{std_tensor_op::StdTensorOp, SymDim, TensorMeta};
35use tenferro_runtime::ad_support::{
36    analyze_deferred_semantic_trace, compile_ad_source, ones_tensor, RetainedValue,
37};
38use tenferro_runtime::program::{ProgramValueMetadata, SemanticFingerprint, SemanticProgram};
39use tenferro_runtime::{
40    CompiledGraph, CoreCapabilityBundle, EngineId, ErrorPhase, ExecutionContextIdentity,
41    ExtensionModule, GraphCompiler, HardwareClassId, PreparedCompiledGraph, RegistrationIdentity,
42    Runtime, RuntimeConfigError, RuntimeConfigSnapshot, RuntimeEpoch, TracedTensor,
43};
44#[cfg(test)]
45use tenferro_tensor::TensorBackend;
46#[cfg(test)]
47use tenferro_tensor::TypedTensor;
48use tenferro_tensor::{
49    AllocationGroup, CacheStats, CompareDir, DType, DescriptorSlot, DotGeneralConfig, GatherConfig,
50    GroupError, IntoShapeVec, PadConfig, ScatterConfig, SliceConfig, Tensor, TensorRead,
51    TensorScalar, TensorValue, TensorView,
52};
53use tenferro_tensor::{BackendSession, BackendSessionHost};
54
55#[cfg(feature = "cuda")]
56use crate::eager_backend::cuda_runtime_engine_id;
57use crate::eager_backend::{
58    cpu_runtime_engine_id, cpu_runtime_hardware_class, eager_runtime_for_backend, EagerBackend,
59};
60#[cfg(test)]
61use crate::eager_exec::exec_standard_op_on_tensor_reads_in_session;
62use crate::eager_exec::{eager_input_promotion_plan, exec_extension_op_on_tensor_reads};
63use crate::error::{ContextId, Error, Result};
64use crate::metadata::tensor_meta_from_tensor;
65use crate::semantic_extension::SemanticExtensionRuleSet;
66use crate::traced::{derivative_trace_from_frozen_program, next_input_key};
67use crate::transform_cache::{AdTransformCache, AdTransformCacheLimits};
68
69use crate::AdContext;
70
71pub(crate) type GradSlot = Arc<Mutex<Option<Arc<AdValueRecord>>>>;
72pub(crate) type WeakGradSlot = Weak<Mutex<Option<Arc<AdValueRecord>>>>;
73
74mod composite;
75mod residuals;
76pub(crate) use residuals::{finish_residuals, EagerTrace};
77
78#[cfg(test)]
79pub(crate) static CPU_RUNTIME_SELECTION_REFRESHES: AtomicUsize = AtomicUsize::new(0);
80
81struct CpuRuntimeSelection {
82    snapshot: Arc<RuntimeConfigSnapshot>,
83    epoch: RuntimeEpoch,
84    engine_id: EngineId,
85    registration_identity: RegistrationIdentity,
86    capabilities: CoreCapabilityBundle,
87}
88
89#[derive(Debug, Default, Clone)]
90struct EagerOpProfileEntry {
91    calls: usize,
92    total_time: Duration,
93}
94
95thread_local! {
96    static EAGER_OP_PROFILE_STATE: RefCell<HashMap<&'static str, EagerOpProfileEntry>> =
97        RefCell::new(HashMap::new());
98    static EAGER_NO_GRAD_DEPTH: Cell<usize> = const { Cell::new(0) };
99    static EAGER_CAPTURE_DEPTH: Cell<usize> = const { Cell::new(0) };
100    /// Runtimes whose session callback is running on this thread.
101    static EAGER_ENTERED_RUNTIMES: RefCell<Vec<ContextId>> = const { RefCell::new(Vec::new()) };
102    #[cfg(test)]
103    static EAGER_OP_PROFILE_ENABLED_OVERRIDE: RefCell<Option<bool>> = const { RefCell::new(None) };
104    #[cfg(test)]
105    static EAGER_OP_PROFILE_PRINT_EVERY_OVERRIDE: RefCell<Option<Option<usize>>> = const { RefCell::new(None) };
106    #[cfg(test)]
107    static EAGER_SEMANTIC_VJP_ENABLED_OVERRIDE: RefCell<Option<bool>> = const { RefCell::new(None) };
108}
109
110#[cfg(test)]
111pub(crate) static EAGER_SEMANTIC_VJP_EXECUTIONS: AtomicUsize = AtomicUsize::new(0);
112
113pub(crate) fn eager_grad_recording_enabled() -> bool {
114    EAGER_NO_GRAD_DEPTH.with(|depth| depth.get() == 0)
115}
116
117pub(crate) fn eager_capture_active() -> bool {
118    EAGER_CAPTURE_DEPTH.with(|depth| depth.get() > 0)
119}
120
121/// The calling thread's `no_grad`/`capture_trace` depths, carried into a
122/// backend-session callback.
123///
124/// A CPU session may run its callback on an executor worker, where the calling
125/// thread's thread-local guards are invisible. The callback thread adds these
126/// depths for the callback's duration, so a guard held around a session entry
127/// governs the operations inside it; guards started inside the callback stay
128/// local to it. On the calling thread itself nothing changes.
129#[derive(Clone, Copy)]
130struct InheritedEagerModes {
131    thread: std::thread::ThreadId,
132    no_grad: usize,
133    capture: usize,
134}
135
136impl InheritedEagerModes {
137    fn capture() -> Self {
138        Self {
139            thread: std::thread::current().id(),
140            no_grad: EAGER_NO_GRAD_DEPTH.with(Cell::get),
141            capture: EAGER_CAPTURE_DEPTH.with(Cell::get),
142        }
143    }
144
145    /// Apply the captured depths on the current thread until the returned
146    /// scope drops, including on unwind.
147    fn enter(self) -> InheritedEagerModesScope {
148        let inherited = if std::thread::current().id() == self.thread {
149            Self {
150                no_grad: 0,
151                capture: 0,
152                ..self
153            }
154        } else {
155            self
156        };
157        EAGER_NO_GRAD_DEPTH.with(|depth| depth.set(depth.get() + inherited.no_grad));
158        EAGER_CAPTURE_DEPTH.with(|depth| depth.set(depth.get() + inherited.capture));
159        InheritedEagerModesScope {
160            inherited,
161            _not_send: PhantomData,
162        }
163    }
164}
165
166/// Marks one runtime's session callback as running on the current thread, so a
167/// nested entry into the same runtime from that callback is rejected before it
168/// waits on the runtime's own owner lock.
169struct EnteredRuntimeScope {
170    id: ContextId,
171    // Pops this thread's entry, so it must drop where it was created.
172    _not_send: PhantomData<Rc<()>>,
173}
174
175impl EnteredRuntimeScope {
176    fn enter(id: ContextId) -> Self {
177        EAGER_ENTERED_RUNTIMES.with(|entered| entered.borrow_mut().push(id));
178        Self {
179            id,
180            _not_send: PhantomData,
181        }
182    }
183
184    fn any_entered() -> bool {
185        EAGER_ENTERED_RUNTIMES.with(|entered| !entered.borrow().is_empty())
186    }
187}
188
189impl Drop for EnteredRuntimeScope {
190    fn drop(&mut self) {
191        EAGER_ENTERED_RUNTIMES.with(|entered| {
192            let mut entered = entered.borrow_mut();
193            if let Some(position) = entered.iter().rposition(|id| *id == self.id) {
194                entered.remove(position);
195            }
196        });
197    }
198}
199
200struct InheritedEagerModesScope {
201    inherited: InheritedEagerModes,
202    // Restores this thread's counters, so it must drop where it was created.
203    _not_send: PhantomData<Rc<()>>,
204}
205
206impl Drop for InheritedEagerModesScope {
207    fn drop(&mut self) {
208        let InheritedEagerModes {
209            no_grad, capture, ..
210        } = self.inherited;
211        EAGER_NO_GRAD_DEPTH.with(|depth| depth.set(depth.get().saturating_sub(no_grad)));
212        EAGER_CAPTURE_DEPTH.with(|depth| depth.set(depth.get().saturating_sub(capture)));
213    }
214}
215
216fn eager_semantic_vjp_enabled() -> bool {
217    #[cfg(test)]
218    if let Some(value) = EAGER_SEMANTIC_VJP_ENABLED_OVERRIDE.with(|state| *state.borrow()) {
219        return value;
220    }
221
222    // Semantic eager VJP/JVP on by default (Unification 7).
223    // Set TENFERRO_EAGER_SEMANTIC_VJP=0 to disable.
224    static ENABLED: OnceLock<bool> = OnceLock::new();
225    *ENABLED.get_or_init(|| env::var("TENFERRO_EAGER_SEMANTIC_VJP").map_or(true, |v| v != "0"))
226}
227
228/// Scope guard that temporarily disables eager operation recording.
229///
230/// Values computed while this guard is alive are concrete eager tensors, but
231/// they do not participate in reverse-mode gradient tracking.
232///
233/// # Examples
234///
235/// ```
236/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
237/// use tenferro_cpu::CpuBackend;
238///
239/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
240/// let x = EagerTensor::requires_grad_in(
241///     Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(),
242///     ctx.clone(),
243/// )?;
244/// let y = ctx.with_eager_session(|s| {
245///     let _guard = ctx.no_grad();
246///     s.mul(&x, &x)
247/// })?;
248/// assert!(!y.tracks_grad());
249/// # Ok::<(), tenferro_ad::Error>(())
250/// ```
251#[derive(Debug)]
252pub struct EagerNoGradGuard {
253    active: bool,
254    // Thread-local depth guard: must not be Send so it cannot be moved to and
255    // dropped on another thread (which would corrupt the creator's depth).
256    _not_send: PhantomData<Rc<()>>,
257}
258
259impl Drop for EagerNoGradGuard {
260    fn drop(&mut self) {
261        if !self.active {
262            return;
263        }
264        EAGER_NO_GRAD_DEPTH.with(|depth| {
265            depth.set(depth.get().saturating_sub(1));
266        });
267        self.active = false;
268    }
269}
270
271/// Scope guard that keeps semantic-trace recording active for untracked
272/// intermediates.
273///
274/// Under active-edge semantics (issue #1665 Def 1), an operation whose inputs
275/// are all untracked produces no autograd nodes and drops its semantic trace.
276/// Inside this guard, such operations still record their semantic trace, so a
277/// later functional JVP/VJP can differentiate with respect to an untracked or
278/// detached leaf. This replaces the pre-Def-1 implicit recording.
279///
280/// # Examples
281///
282/// ```
283/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
284/// use tenferro_cpu::CpuBackend;
285///
286/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
287/// let x = EagerTensor::from_tensor_in(
288///     Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap(),
289///     ctx.clone(),
290/// )?;
291/// let y = ctx.with_eager_session(|s| {
292///     let _capture = ctx.capture_trace();
293///     s.mul(&x, &x)
294/// })?;
295/// let seed = EagerTensor::from_tensor_in(
296///     Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 1.0]).unwrap(),
297///     ctx.clone(),
298/// )?;
299/// let dx = ctx.vjp(&y, &x, &seed)?;
300/// assert_eq!(dx.value()?.as_slice::<f64>().unwrap(), &[2.0, 4.0]);
301/// # Ok::<(), tenferro_ad::Error>(())
302/// ```
303#[derive(Debug)]
304pub struct EagerTraceCaptureGuard {
305    active: bool,
306    // Thread-local depth guard: must not be Send so it cannot be moved to and
307    // dropped on another thread (which would corrupt the creator's depth).
308    _not_send: PhantomData<Rc<()>>,
309}
310
311impl Drop for EagerTraceCaptureGuard {
312    fn drop(&mut self) {
313        if !self.active {
314            return;
315        }
316        EAGER_CAPTURE_DEPTH.with(|depth| {
317            depth.set(depth.get().saturating_sub(1));
318        });
319        self.active = false;
320    }
321}
322
323pub(crate) fn eager_op_profile_enabled() -> bool {
324    #[cfg(test)]
325    if let Some(value) = EAGER_OP_PROFILE_ENABLED_OVERRIDE.with(|state| *state.borrow()) {
326        return value;
327    }
328
329    static ENABLED: OnceLock<bool> = OnceLock::new();
330    *ENABLED.get_or_init(|| env::var("TENFERRO_PROFILE_EAGER_OP_AGG").is_ok())
331}
332
333pub(crate) fn eager_op_profile_start() -> Option<Instant> {
334    eager_op_profile_enabled().then(Instant::now)
335}
336
337pub(crate) fn record_eager_op_profile(section: &'static str, elapsed: Duration) {
338    if !eager_op_profile_enabled() {
339        return;
340    }
341    EAGER_OP_PROFILE_STATE.with(|state| {
342        let mut state = state.borrow_mut();
343        let entry = state.entry(section).or_default();
344        entry.calls += 1;
345        entry.total_time += elapsed;
346    });
347}
348
349pub(crate) fn profile_eager_op_section<T>(section: &'static str, f: impl FnOnce() -> T) -> T {
350    if !eager_op_profile_enabled() {
351        return f();
352    }
353    let started = Instant::now();
354    let result = f();
355    record_eager_op_profile(section, started.elapsed());
356    result
357}
358
359pub(crate) fn maybe_print_eager_op_profile() {
360    if !eager_op_profile_enabled() {
361        return;
362    }
363    let Some(print_every) = eager_op_profile_print_every() else {
364        return;
365    };
366    if print_every == 0 {
367        return;
368    }
369
370    let should_print = EAGER_OP_PROFILE_STATE.with(|state| {
371        state
372            .borrow()
373            .get("nary_op.total")
374            .is_some_and(|entry| entry.calls % print_every == 0)
375    });
376    if should_print {
377        print_and_reset_eager_op_profile();
378    }
379}
380
381fn eager_op_profile_print_every() -> Option<usize> {
382    #[cfg(test)]
383    if let Some(value) = EAGER_OP_PROFILE_PRINT_EVERY_OVERRIDE.with(|state| *state.borrow()) {
384        return value;
385    }
386
387    env::var("TENFERRO_PROFILE_EAGER_OP_PRINT_EVERY")
388        .ok()?
389        .parse()
390        .ok()
391}
392
393pub(crate) fn print_and_reset_eager_op_profile() {
394    EAGER_OP_PROFILE_STATE.with(|state| {
395        let mut entries: Vec<_> = state
396            .borrow()
397            .iter()
398            .map(|(section, entry)| (*section, entry.clone()))
399            .collect();
400        state.borrow_mut().clear();
401        entries.sort_by_key(|(_, entry)| Reverse(entry.total_time));
402
403        eprintln!("=== tenferro eager op profile ===");
404        for (section, entry) in entries {
405            let Some(per_call_us) = eager_op_profile_per_call_us(&entry) else {
406                continue;
407            };
408            eprintln!(
409                "{section}: calls={} total={:.6}ms per_call={:.3}us",
410                entry.calls,
411                entry.total_time.as_secs_f64() * 1.0e3,
412                per_call_us,
413            );
414        }
415    });
416}
417
418fn eager_op_profile_per_call_us(entry: &EagerOpProfileEntry) -> Option<f64> {
419    (entry.calls != 0).then(|| entry.total_time.as_secs_f64() * 1.0e6 / entry.calls as f64)
420}
421
422fn runtime_config_error(op: &'static str, source: RuntimeConfigError) -> Error {
423    Error::runtime_state_source(op, ErrorPhase::Execution, source)
424}
425
426fn runtime_state_source<E>(op: &'static str, source: E) -> Error
427where
428    E: std::error::Error + Send + Sync + 'static,
429{
430    Error::runtime_state_source(op, ErrorPhase::Execution, source)
431}
432
433fn cpu_runtime_bridge_unsupported(message: impl Into<String>) -> Error {
434    Error::unsupported(
435        "CpuPlacementBoundEager::refresh_runtime_selection",
436        ErrorPhase::Execution,
437        message,
438    )
439}
440
441fn select_cpu_runtime(runtime: &Runtime) -> Result<CpuRuntimeSelection> {
442    let snapshot = runtime
443        .snapshot()
444        .map_err(|source| runtime_state_source("EagerRuntime::runtime_snapshot", source))?;
445    let engine_id = cpu_runtime_engine_id()
446        .map_err(|source| runtime_config_error("EagerRuntime::cpu_runtime_engine_id", source))?;
447    let expected_hardware = cpu_runtime_hardware_class().map_err(|source| {
448        runtime_config_error("EagerRuntime::cpu_runtime_hardware_class", source)
449    })?;
450    let engine = snapshot
451        .engine(&engine_id)
452        .ok_or_else(|| cpu_runtime_bridge_unsupported("missing CPU runtime engine"))?;
453    validate_cpu_runtime_engine(
454        engine.context_identity(),
455        engine.hardware_class(),
456        engine.capabilities(),
457        &expected_hardware,
458    )?;
459    let epoch = snapshot.epoch();
460    let registration_identity = engine.registration_identity();
461    let capabilities = engine.capabilities().clone();
462    Ok(CpuRuntimeSelection {
463        snapshot,
464        epoch,
465        engine_id,
466        registration_identity,
467        capabilities,
468    })
469}
470
471fn validate_cpu_runtime_engine(
472    context_identity: ExecutionContextIdentity,
473    hardware_class: &HardwareClassId,
474    capabilities: &CoreCapabilityBundle,
475    expected_hardware: &HardwareClassId,
476) -> Result<()> {
477    if context_identity != ExecutionContextIdentity::of::<CpuBackend>() {
478        return Err(cpu_runtime_bridge_unsupported(
479            "CPU runtime context mismatch",
480        ));
481    }
482    if hardware_class != expected_hardware {
483        return Err(cpu_runtime_bridge_unsupported(
484            "CPU runtime hardware mismatch",
485        ));
486    }
487    if capabilities.elementwise().is_none() {
488        return Err(cpu_runtime_bridge_unsupported(
489            "missing CPU runtime capability: elementwise",
490        ));
491    }
492    if capabilities.reduction().is_none() {
493        return Err(cpu_runtime_bridge_unsupported(
494            "missing CPU runtime capability: reduction",
495        ));
496    }
497    if capabilities.indexing().is_none() {
498        return Err(cpu_runtime_bridge_unsupported(
499            "missing CPU runtime capability: indexing",
500        ));
501    }
502    if capabilities.dot_general().is_none() {
503        return Err(cpu_runtime_bridge_unsupported(
504            "missing CPU runtime capability: dot_general",
505        ));
506    }
507    if capabilities.layout().is_none() {
508        return Err(cpu_runtime_bridge_unsupported(
509            "missing CPU runtime capability: layout",
510        ));
511    }
512    Ok(())
513}
514
515/// Stats for caches owned by an [`EagerRuntime`].
516///
517/// `retained_bytes` fields are logical payload estimates, not process RSS.
518#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
519pub struct EagerRuntimeCacheStats {
520    /// Generic extension runtime caches.
521    pub extensions: CacheStats,
522    /// Eager AD transform memoization cache.
523    pub ad_transforms: CacheStats,
524    /// Prepared eager derivative program cache.
525    pub prepared_derivatives: CacheStats,
526}
527
528#[cfg(test)]
529pub(crate) struct EagerGraphExecution {
530    pub(crate) outputs: Vec<Tensor>,
531}
532
533/// A read-only value view retained by an eager tensor record.
534///
535/// The guard borrows the record's allocation group. It never owns a tensor and
536/// cannot be converted into a mutable view.
537///
538/// # Examples
539///
540/// ```
541/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
542/// use tenferro_cpu::CpuBackend;
543///
544/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
545/// let value = EagerTensor::from_tensor_in(
546///     Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?,
547///     ctx,
548/// )?;
549/// let view = value.value()?;
550/// assert_eq!(view.shape(), &[2]);
551/// # Ok::<(), tenferro_ad::Error>(())
552/// ```
553#[derive(Debug)]
554pub struct ValueGuard<'a> {
555    view: TensorView<'a>,
556}
557
558impl<'a> ValueGuard<'a> {
559    /// Return the scalar dtype of the retained value.
560    pub fn dtype(&self) -> DType {
561        self.view.dtype()
562    }
563
564    /// Return the logical shape of the retained value.
565    pub fn shape(&self) -> &[usize] {
566        self.view.shape()
567    }
568
569    /// Borrow the dtype-erased tensor view.
570    pub fn as_tensor_view(&self) -> &TensorView<'_> {
571        &self.view
572    }
573
574    /// Borrow compact host bytes through the tensor's explicit scalar type.
575    ///
576    /// Backend-resident values return the backend's typed host-access error;
577    /// this method does not download storage implicitly.
578    ///
579    /// # Errors
580    ///
581    /// Returns [`tenferro_tensor::ValidationError::DTypeMismatch`] when
582    /// `T` does not match the view dtype, [`tenferro_tensor::ValidationError::NonContiguousViewAsSlice`]
583    /// for a non-contiguous view, or [`tenferro_tensor::Error::HostAccess`]
584    /// when backend storage cannot be mapped as a host slice.
585    pub fn as_slice<T: TensorScalar>(&self) -> tenferro_tensor::Result<&'a [T]> {
586        self.view.as_slice()
587    }
588
589    fn duplicate_host_tensor(&self) -> tenferro_tensor::Result<Tensor> {
590        match &self.view {
591            TensorView::F32(view) => {
592                <f32 as TensorScalar>::into_tensor(view.shape().to_vec(), view.as_slice()?.to_vec())
593            }
594            TensorView::F64(view) => {
595                <f64 as TensorScalar>::into_tensor(view.shape().to_vec(), view.as_slice()?.to_vec())
596            }
597            TensorView::I32(view) => {
598                <i32 as TensorScalar>::into_tensor(view.shape().to_vec(), view.as_slice()?.to_vec())
599            }
600            TensorView::I64(view) => {
601                <i64 as TensorScalar>::into_tensor(view.shape().to_vec(), view.as_slice()?.to_vec())
602            }
603            TensorView::Bool(view) => <bool as TensorScalar>::into_tensor(
604                view.shape().to_vec(),
605                view.as_slice()?.to_vec(),
606            ),
607            TensorView::C32(view) => <num_complex::Complex32 as TensorScalar>::into_tensor(
608                view.shape().to_vec(),
609                view.as_slice()?.to_vec(),
610            ),
611            TensorView::C64(view) => <num_complex::Complex64 as TensorScalar>::into_tensor(
612                view.shape().to_vec(),
613                view.as_slice()?.to_vec(),
614            ),
615        }
616    }
617}
618
619/// Read-only retained gradient value.
620///
621/// # Examples
622///
623/// ```
624/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
625/// use tenferro_cpu::CpuBackend;
626///
627/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
628/// let x = EagerTensor::requires_grad_in(
629///     Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?,
630///     ctx,
631/// )?;
632/// let loss = x.runtime().with_eager_session(|s| {
633///     let squared = s.mul(&x, &x)?;
634///     s.reduce_sum(&squared, Some(&[0]))
635/// })?;
636/// let _gradients = loss.backward()?;
637/// let gradient = x.grad()?.expect("tracked leaf has a gradient");
638/// assert_eq!(gradient.shape(), &[2]);
639/// # Ok::<(), tenferro_ad::Error>(())
640/// ```
641#[derive(Clone, Debug)]
642pub struct GradientValue {
643    record: Arc<AdValueRecord>,
644    ctx: Arc<EagerRuntime>,
645}
646
647impl GradientValue {
648    /// Return the scalar dtype of the gradient.
649    pub fn dtype(&self) -> DType {
650        self.record.dtype()
651    }
652
653    /// Return the logical shape of the gradient.
654    pub fn shape(&self) -> &[usize] {
655        self.record.shape()
656    }
657
658    /// Borrow the gradient's value guard.
659    ///
660    /// # Errors
661    ///
662    /// Returns [`Error::RuntimeState`] when the retained gradient record is
663    /// unavailable or its allocation-group descriptor is invalid.
664    pub fn value(&self) -> Result<ValueGuard<'_>> {
665        self.record.value("GradientValue::value")
666    }
667
668    /// Borrow the gradient as a dtype-erased read target.
669    ///
670    /// # Errors
671    ///
672    /// Returns [`Error::RuntimeState`] when the retained gradient record or
673    /// its allocation-group descriptor is unavailable.
674    pub fn tensor_read(&self) -> Result<TensorRead<'_>> {
675        self.record.tensor_read("GradientValue::tensor_read")
676    }
677
678    /// Borrow a compact host slice without downloading backend storage.
679    ///
680    /// # Errors
681    ///
682    /// Returns [`Error::RuntimeState`] when the retained value is unavailable,
683    /// [`tenferro_tensor::ValidationError::DTypeMismatch`] when `T` does
684    /// not match the gradient dtype, or [`tenferro_tensor::Error::HostAccess`]
685    /// when backend storage cannot be mapped as a host slice.
686    pub fn as_slice<T: TensorScalar>(&self) -> tenferro_tensor::Result<&[T]> {
687        self.record
688            .value("GradientValue::as_slice")
689            .map_err(|error| {
690                tenferro_tensor::Error::runtime_state_source("GradientValue::as_slice", error)
691            })?
692            .as_slice()
693    }
694
695    /// Explicitly copy a host-resident gradient into a standalone tensor.
696    ///
697    /// # Errors
698    ///
699    /// Returns [`Error::RuntimeState`] when the retained value or execution
700    /// session is unavailable, or a typed backend/host-access error when the
701    /// gradient cannot be materialized as a contiguous tensor.
702    pub fn to_tensor(&self) -> Result<Tensor> {
703        let value = self
704            .record
705            .value("GradientValue::to_tensor")
706            .map_err(|error| {
707                Error::runtime_state_source(
708                    "GradientValue::to_tensor",
709                    ErrorPhase::Execution,
710                    error,
711                )
712            })?;
713        match value.duplicate_host_tensor() {
714            Ok(tensor) => Ok(tensor),
715            Err(_) => {
716                let read = self.record.tensor_read("GradientValue::to_tensor")?;
717                self.ctx
718                    .with_execution_session(|session| session.to_contiguous_read(read))?
719                    .map_err(Error::from)
720            }
721        }
722    }
723}
724
725/// Move-only accumulated gradient bundle backed by one allocation group.
726///
727/// # Examples
728///
729/// ```
730/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
731/// use tenferro_cpu::CpuBackend;
732///
733/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
734/// let x = EagerTensor::requires_grad_in(
735///     Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?,
736///     ctx,
737/// )?;
738/// let loss = x.runtime().with_eager_session(|s| {
739///     let squared = s.mul(&x, &x)?;
740///     s.reduce_sum(&squared, Some(&[0]))
741/// })?;
742/// let gradients = loss.backward()?;
743/// assert!(!gradients.is_empty());
744/// # Ok::<(), tenferro_ad::Error>(())
745/// ```
746#[derive(Debug)]
747pub struct Gradients {
748    group: AllocationGroup,
749    slots: HashMap<ValueKey<StdTensorOp>, DescriptorSlot>,
750}
751
752impl Gradients {
753    fn from_tensors(tensors: HashMap<ValueKey<StdTensorOp>, Tensor>) -> Result<Self> {
754        let (keys, values): (Vec<_>, Vec<_>) = tensors.into_iter().unzip();
755        let (group, bindings) = AllocationGroup::from_tensors(values).map_err(|error| {
756            Error::runtime_state_source("Gradients::from_tensors", ErrorPhase::Execution, error)
757        })?;
758        let slots = keys.into_iter().zip(bindings).collect();
759        Ok(Self { group, slots })
760    }
761
762    /// Return the number of retained gradient descriptors.
763    pub fn len(&self) -> usize {
764        self.slots.len()
765    }
766
767    /// Return whether no gradient was produced.
768    pub fn is_empty(&self) -> bool {
769        self.slots.is_empty()
770    }
771
772    /// Borrow one gradient view by its local value key.
773    pub fn grad(&self, key: &ValueKey<StdTensorOp>) -> Option<TensorView<'_>> {
774        let slot = self.slots.get(key).copied()?;
775        let mut reads = self.group.read_views(std::slice::from_ref(&slot)).ok()?;
776        match reads.pop()? {
777            TensorRead::View(view) => Some(view),
778            TensorRead::Tensor(_) => None,
779        }
780    }
781
782    /// Consume one gradient owner while leaving the bundle unchanged on failure.
783    ///
784    /// # Errors
785    ///
786    /// Returns [`tenferro_tensor::Error::RuntimeState`] when the descriptor is
787    /// invalid or its allocation is aliased. A missing key is reported as
788    /// `Ok(None)`.
789    pub fn take_grad(
790        &mut self,
791        key: &ValueKey<StdTensorOp>,
792    ) -> tenferro_tensor::Result<Option<Tensor>> {
793        let Some(&slot) = self.slots.get(key) else {
794            return Ok(None);
795        };
796        let tensor = self.group.take_tensor(slot).map_err(|error| {
797            tenferro_tensor::Error::runtime_state_source("Gradients::take_grad", error)
798        })?;
799        self.slots.remove(key);
800        Ok(Some(tensor))
801    }
802}
803
804/// Error returned when a value cannot be consumed without changing its owner.
805///
806/// # Examples
807///
808/// ```
809/// use tenferro_ad::{EagerRuntime, EagerTensor, IntoValueError, Tensor};
810/// use tenferro_cpu::CpuBackend;
811///
812/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
813/// let value = EagerTensor::from_tensor_in(
814///     Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?,
815///     ctx,
816/// )?;
817/// let _shared = value.clone();
818/// assert!(matches!(
819///     value.into_value(),
820///     Err(IntoValueError::NotUnique(_))
821/// ));
822/// # Ok::<(), tenferro_ad::Error>(())
823/// ```
824#[derive(Debug)]
825pub enum IntoValueError<H> {
826    /// Another eager handle, tape record, or checkpoint retains the value.
827    NotUnique(H),
828    /// Group extraction failed after the handle was uniquely acquired.
829    Extract { value: H, error: GroupError },
830}
831
832impl<H> std::fmt::Display for IntoValueError<H> {
833    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
834        match self {
835            Self::NotUnique(_) => formatter.write_str("eager value is retained by another handle"),
836            Self::Extract { error, .. } => {
837                write!(formatter, "eager value extraction failed: {error}")
838            }
839        }
840    }
841}
842
843impl<H: std::fmt::Debug + Send + Sync + 'static> std::error::Error for IntoValueError<H> {}
844
845/// What one direct retention container holds.
846// The pooled variant owns an allocation group inline; boxing it would add an
847// allocation to every retained value on the hot path.
848#[allow(clippy::large_enum_variant)]
849#[derive(Debug)]
850enum RetentionContainer {
851    /// Pooled storage owned through an allocation group and one descriptor slot.
852    Pooled {
853        group: AllocationGroup,
854        slot: DescriptorSlot,
855    },
856    /// A caller-owned value the runtime retains without taking pool ownership.
857    ///
858    /// The value returns to its owner when the record drops, and the runtime never
859    /// substitutes it for pooled storage. A caller-owned payload has no typed
860    /// descriptor view, so only its read and consume paths are available.
861    CallerOwned {
862        /// Boxed because a tensor value is much larger than the pooled variant.
863        tensor: Box<Tensor>,
864    },
865    /// An untracked result held directly.
866    ///
867    /// No AD group, residual or gradient will share it, so it needs no
868    /// allocation group: the tensor's own storage returns to its pool when the
869    /// record drops, and a unique handle hands the tensor back unchanged.
870    Owned { tensor: Tensor },
871}
872
873/// Read-only descriptor record used by eager handles and the AD registries.
874#[derive(Debug)]
875pub(crate) struct AdValueRecord {
876    container: Arc<RetentionContainer>,
877    dtype: DType,
878    shape: Box<[usize]>,
879}
880
881impl AdValueRecord {
882    fn from_group(
883        group: AllocationGroup,
884        slot: DescriptorSlot,
885        dtype: DType,
886        shape: Vec<usize>,
887    ) -> Arc<Self> {
888        Arc::new(Self {
889            container: Arc::new(RetentionContainer::Pooled { group, slot }),
890            dtype,
891            shape: shape.into_boxed_slice(),
892        })
893    }
894
895    fn from_tensor(tensor: Tensor, op: &'static str) -> Result<Arc<Self>> {
896        let dtype = tensor.dtype();
897        let shape = tensor.shape().to_vec();
898        if matches!(dtype, DType::External(_)) {
899            // A caller-owned payload is retained directly: it owns no pooled
900            // storage, so there is no group to build and nothing to return to a
901            // pool when the record drops.
902            return Ok(Arc::new(Self {
903                container: Arc::new(RetentionContainer::CallerOwned {
904                    tensor: Box::new(tensor),
905                }),
906                dtype,
907                shape: shape.into_boxed_slice(),
908            }));
909        }
910        let (group, bindings) = AllocationGroup::from_tensors(vec![tensor])
911            .map_err(|error| Error::runtime_state_source(op, ErrorPhase::Execution, error))?;
912        let slot = bindings.first().copied().ok_or_else(|| {
913            Error::runtime_state(op, ErrorPhase::Execution, "empty allocation-group binding")
914        })?;
915        Ok(Self::from_group(group, slot, dtype, shape))
916    }
917
918    /// Retain an untracked result without building an allocation group.
919    fn from_untracked_tensor(tensor: Tensor, op: &'static str) -> Result<Arc<Self>> {
920        if matches!(tensor.dtype(), DType::External(_)) {
921            return Self::from_tensor(tensor, op);
922        }
923        let dtype = tensor.dtype();
924        let shape = tensor.shape().to_vec().into_boxed_slice();
925        Ok(Arc::new(Self {
926            container: Arc::new(RetentionContainer::Owned { tensor }),
927            dtype,
928            shape,
929        }))
930    }
931
932    fn tensor_read(&self, op: &'static str) -> Result<TensorRead<'_>> {
933        match self.container.as_ref() {
934            RetentionContainer::Pooled { group, slot } => {
935                let mut reads = group
936                    .read_views(std::slice::from_ref(slot))
937                    .map_err(|error| {
938                        Error::runtime_state_source(op, ErrorPhase::Execution, error)
939                    })?;
940                reads.pop().ok_or_else(|| {
941                    Error::runtime_state(
942                        op,
943                        ErrorPhase::Execution,
944                        "empty allocation-group binding",
945                    )
946                })
947            }
948            RetentionContainer::CallerOwned { tensor } => Ok(TensorRead::from_tensor(tensor)),
949            RetentionContainer::Owned { tensor } => Ok(TensorRead::from_tensor(tensor)),
950        }
951    }
952
953    fn value(&self, op: &'static str) -> Result<ValueGuard<'_>> {
954        if let RetentionContainer::Owned { tensor } = self.container.as_ref() {
955            // A preset-dtype owned tensor always has a typed borrowed view.
956            return Ok(ValueGuard {
957                view: TensorRead::from_tensor(tensor).tensor_view(),
958            });
959        }
960        match self.tensor_read(op)? {
961            TensorRead::View(view) => Ok(ValueGuard { view }),
962            // A caller-owned payload has no typed descriptor view, so a path that
963            // needs one fails explicitly instead of borrowing the payload as bytes.
964            TensorRead::Tensor(_) => Err(Error::runtime_state(
965                op,
966                ErrorPhase::Execution,
967                "allocation-group value did not produce a borrowed descriptor view",
968            )),
969        }
970    }
971
972    fn dtype(&self) -> DType {
973        self.dtype
974    }
975
976    fn shape(&self) -> &[usize] {
977        &self.shape
978    }
979}
980
981/// Placement-selected CPU view of one [`EagerRuntime`].
982///
983/// The view snapshots the runtime's CPU coordinator/provider bundle and the
984/// immutable runtime registration metadata when [`EagerRuntime::on_cpu`] is
985/// called. It holds no resource permit while idle and enters one backend
986/// session only while [`Self::with_eager_session`] runs. The session exposes
987/// core [`BackendSession`] operations on concrete [`Tensor`] values. This
988/// bridge deliberately does not expose the eager runtime's linalg, FFT, einsum,
989/// or extension-runtime registries.
990///
991/// The value is intentionally not `Clone`: mutable use makes concurrent
992/// session ownership explicit without adding another backend mutex.
993///
994/// # Examples
995///
996/// ```rust
997/// use tenferro_ad::EagerRuntime;
998/// use tenferro_cpu::CpuPlacement;
999///
1000/// let runtime = EagerRuntime::new()?;
1001/// let cpu = runtime.on_cpu(CpuPlacement::Auto)?;
1002/// assert_eq!(cpu.runtime_id(), runtime.id());
1003/// # Ok::<(), tenferro_ad::Error>(())
1004/// ```
1005pub struct CpuPlacementBoundEager {
1006    runtime: Arc<EagerRuntime>,
1007    backend: CpuBackend,
1008    snapshot: Arc<RuntimeConfigSnapshot>,
1009    epoch: RuntimeEpoch,
1010    engine_id: EngineId,
1011    registration_identity: RegistrationIdentity,
1012    capabilities: CoreCapabilityBundle,
1013}
1014
1015impl fmt::Debug for CpuPlacementBoundEager {
1016    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1017        f.debug_struct("CpuPlacementBoundEager")
1018            .field("runtime_id", &self.runtime.id())
1019            .field("placement", &self.backend.placement())
1020            .field("runtime_epoch", &self.epoch)
1021            .field("engine_id", &self.engine_id)
1022            .field("registration_identity", &self.registration_identity)
1023            .finish_non_exhaustive()
1024    }
1025}
1026
1027impl CpuPlacementBoundEager {
1028    fn refresh_runtime_selection(&mut self) -> Result<()> {
1029        let current_epoch = self.runtime.runtime.epoch().map_err(|source| {
1030            runtime_state_source("CpuPlacementBoundEager::refresh_runtime_selection", source)
1031        })?;
1032        if current_epoch == self.epoch {
1033            return Ok(());
1034        }
1035
1036        #[cfg(test)]
1037        CPU_RUNTIME_SELECTION_REFRESHES.fetch_add(1, Ordering::SeqCst);
1038
1039        let selection = select_cpu_runtime(&self.runtime.runtime)?;
1040        self.snapshot = selection.snapshot;
1041        self.epoch = selection.epoch;
1042        self.engine_id = selection.engine_id;
1043        self.registration_identity = selection.registration_identity;
1044        self.capabilities = selection.capabilities;
1045        Ok(())
1046    }
1047
1048    /// Return the identity of the original eager runtime.
1049    ///
1050    /// # Examples
1051    ///
1052    /// ```rust
1053    /// use tenferro_ad::EagerRuntime;
1054    /// use tenferro_cpu::CpuPlacement;
1055    ///
1056    /// let runtime = EagerRuntime::new()?;
1057    /// let cpu = runtime.on_cpu(CpuPlacement::Auto)?;
1058    /// assert_eq!(cpu.runtime_id(), runtime.id());
1059    /// # Ok::<(), tenferro_ad::Error>(())
1060    /// ```
1061    pub fn runtime_id(&self) -> ContextId {
1062        self.runtime.id()
1063    }
1064
1065    /// Return the placement requested when this view was created.
1066    ///
1067    /// # Examples
1068    ///
1069    /// ```rust
1070    /// use tenferro_ad::EagerRuntime;
1071    /// use tenferro_cpu::CpuPlacement;
1072    ///
1073    /// let runtime = EagerRuntime::new()?;
1074    /// let cpu = runtime.on_cpu(CpuPlacement::Auto)?;
1075    /// assert_eq!(cpu.placement(), CpuPlacement::Auto);
1076    /// # Ok::<(), tenferro_ad::Error>(())
1077    /// ```
1078    pub fn placement(&self) -> CpuPlacement {
1079        self.backend.placement()
1080    }
1081
1082    /// Enter one CPU backend session and run core operations through it.
1083    ///
1084    /// One call creates one backend session. Tenferro-managed CPU executors
1085    /// enter once around the closure and core operations reuse that compatible
1086    /// execution scope. The closure may borrow stack data and need not be
1087    /// `'static`.
1088    ///
1089    /// This phase-2 bridge accepts only core [`BackendSession`] operations. It
1090    /// does not lock or dispatch the eager runtime's linalg, FFT, einsum, or
1091    /// extension registries.
1092    ///
1093    /// # Examples
1094    ///
1095    /// ```rust
1096    /// use tenferro_ad::{EagerRuntime, Error};
1097    /// use tenferro_cpu::CpuPlacement;
1098    /// use tenferro_tensor::{Tensor, TensorRead};
1099    ///
1100    /// let runtime = EagerRuntime::new()?;
1101    /// let mut cpu = runtime.on_cpu(CpuPlacement::Auto)?;
1102    /// let lhs = Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?;
1103    /// let rhs = Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?;
1104    /// let output = cpu.with_eager_session(|session| {
1105    ///     session
1106    ///         .add_read(TensorRead::from_tensor(&lhs), TensorRead::from_tensor(&rhs))
1107    ///         .map_err(Error::from)
1108    /// })?;
1109    /// assert_eq!(output.as_slice::<f64>().unwrap(), &[3.0]);
1110    /// # Ok::<(), Error>(())
1111    /// ```
1112    ///
1113    /// # Errors
1114    ///
1115    /// Returns the callback's error unchanged. Core backend operations may
1116    /// report validation, unsupported capability, backend, or runtime-state
1117    /// failures through that error. Returns `E::from(`[`Error::SessionEntry`]`)`
1118    /// without running the callback when the backend cannot admit the session,
1119    /// for example when it is called from inside another session on this
1120    /// thread ([`tenferro_tensor::SessionEntryError::Reentered`]), and
1121    /// `E::from` the runtime-selection error when the CPU placement cannot be
1122    /// refreshed. Use only the borrowed `session` for work inside the scope.
1123    pub fn with_eager_session<T: Send, E: From<Error> + Send>(
1124        &mut self,
1125        f: impl FnOnce(&mut dyn BackendSession) -> std::result::Result<T, E> + Send,
1126    ) -> std::result::Result<T, E> {
1127        self.refresh_runtime_selection().map_err(E::from)?;
1128        match self.backend.with_backend_session(f) {
1129            Ok(result) => result,
1130            Err(entry) => Err(E::from(Error::from(entry))),
1131        }
1132    }
1133}
1134
1135/// Shared eager execution context for tensors on a backend.
1136///
1137/// Reusing one context lets eager tensors share backend state, extension
1138/// runtime caches, and gradient storage across a computation.
1139///
1140/// # Examples
1141///
1142/// ```
1143/// use tenferro_cpu::CpuBackend;
1144/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1145///
1146/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1147/// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(), ctx.clone()).unwrap();
1148/// let y = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(), ctx.clone()).unwrap();
1149/// let z = ctx.with_eager_session(|session| session.add(&x, &y)).unwrap();
1150///
1151/// assert_eq!(z.value().unwrap().as_slice::<f64>().unwrap(), &[3.0]);
1152/// # Ok::<(), tenferro_ad::Error>(())
1153/// ```
1154pub struct EagerRuntime {
1155    id: ContextId,
1156    runtime: Runtime,
1157    // The backend and its exact runtime engine registration are selected
1158    // together during construction and remain paired for this runtime's
1159    // lifetime. The mutex only serializes mutable backend operations.
1160    backend: Mutex<EagerBackend>,
1161    // Fixed at construction with the backend/engine pair; extension dispatch
1162    // can inspect it while holding the borrowed backend session.
1163    extension_backend_kind: Option<EagerExtensionBackendKind>,
1164    extension_install_lock: Mutex<()>,
1165    pub(crate) extension_caches: Mutex<ExtensionCacheStore>,
1166    semantic_extension_rules: SemanticExtensionRuleSet,
1167    grad_slots: Mutex<HashMap<ValueKey<StdTensorOp>, WeakGradSlot>>,
1168    value_records: Mutex<HashMap<ValueKey<StdTensorOp>, Weak<EagerTensorRecord>>>,
1169    ad_transform_cache: Arc<AdTransformCache>,
1170    /// S2: prepared derivative programs keyed by semantic structure, wrt input,
1171    /// and concrete bound input metadata. Avoids re-running freeze+AD
1172    /// transform+compile_frozen on warm structure hits.
1173    prepared_derivative_cache: Mutex<PreparedDerivativeCache>,
1174}
1175
1176/// An eager runtime and its borrowed backend session for one execution boundary.
1177///
1178/// Obtain this only through [`EagerRuntime::with_eager_session`]. It rejects
1179/// tensors from another eager runtime even if both runtimes use the same backend
1180/// type, and it cannot escape the boundary closure.
1181///
1182/// # Examples
1183///
1184/// ```rust
1185/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1186/// use tenferro_cpu::CpuBackend;
1187///
1188/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1189/// let x = EagerTensor::from_tensor_in(
1190///     Tensor::from_vec_col_major(vec![1], vec![3.0_f64])?, ctx.clone(),
1191/// )?;
1192/// let y = ctx.with_eager_session(|session| session.neg(&x))?;
1193/// assert_eq!(y.value()?.as_slice::<f64>()?, &[-3.0]);
1194/// # Ok::<(), tenferro_ad::Error>(())
1195/// ```
1196///
1197/// The borrowed session cannot escape its execution boundary:
1198///
1199/// ```compile_fail
1200/// use tenferro_ad::EagerRuntime;
1201/// use tenferro_cpu::CpuBackend;
1202/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new()).unwrap();
1203/// let escaped = ctx.with_eager_session(|session| session).unwrap();
1204/// let _ = escaped;
1205/// ```
1206pub struct EagerSession<'a> {
1207    runtime: &'a Arc<EagerRuntime>,
1208    backend: &'a mut dyn BackendSession,
1209}
1210
1211impl fmt::Debug for EagerSession<'_> {
1212    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
1213        f.debug_struct("EagerSession")
1214            .field("runtime_id", &self.runtime.id())
1215            .finish_non_exhaustive()
1216    }
1217}
1218
1219impl EagerSession<'_> {
1220    /// Negate an eager tensor inside the caller's execution boundary.
1221    ///
1222    /// # Examples
1223    ///
1224    /// ```rust
1225    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1226    /// use tenferro_cpu::CpuBackend;
1227    ///
1228    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1229    /// let x = EagerTensor::from_tensor_in(
1230    ///     Tensor::from_vec_col_major(vec![1], vec![4.0_f64])?, ctx.clone(),
1231    /// )?;
1232    /// let y = ctx.with_eager_session(|session| session.neg(&x))?;
1233    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[-4.0]);
1234    /// # Ok::<(), tenferro_ad::Error>(())
1235    /// ```
1236    ///
1237    /// # Errors
1238    ///
1239    /// Returns [`Error::ContextMismatch`] for a tensor from another runtime,
1240    /// or a typed eager/backend error from the selected operation.
1241    pub fn neg(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1242        self.run_unary(input, StdTensorOp::Neg)
1243    }
1244
1245    /// Compute the elementwise exponential inside this eager session.
1246    ///
1247    /// # Examples
1248    ///
1249    /// ```rust
1250    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1251    /// use tenferro_cpu::CpuBackend;
1252    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1253    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![0.0_f64])?, ctx.clone())?;
1254    /// let y = ctx.with_eager_session(|session| session.exp(&x))?;
1255    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[1.0]);
1256    /// # Ok::<(), tenferro_ad::Error>(())
1257    /// ```
1258    ///
1259    /// # Errors
1260    ///
1261    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1262    /// unsupported/backend error for the input dtype.
1263    pub fn exp(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1264        self.run_unary(input, StdTensorOp::Exp)
1265    }
1266
1267    /// Compute the elementwise absolute value inside this eager session.
1268    ///
1269    /// # Examples
1270    ///
1271    /// ```rust
1272    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1273    /// use tenferro_cpu::CpuBackend;
1274    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1275    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![-2.0_f64])?, ctx.clone())?;
1276    /// let y = ctx.with_eager_session(|session| session.abs(&x))?;
1277    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[2.0]);
1278    /// # Ok::<(), tenferro_ad::Error>(())
1279    /// ```
1280    ///
1281    /// # Errors
1282    ///
1283    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1284    /// unsupported/backend error for the input dtype.
1285    pub fn abs(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1286        self.run_unary(input, StdTensorOp::Abs)
1287    }
1288
1289    /// Compute the elementwise conjugate inside this eager session.
1290    ///
1291    /// # Examples
1292    ///
1293    /// ```rust
1294    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1295    /// use tenferro_cpu::CpuBackend;
1296    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1297    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?, ctx.clone())?;
1298    /// let y = ctx.with_eager_session(|session| session.conj(&x))?;
1299    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[2.0]);
1300    /// # Ok::<(), tenferro_ad::Error>(())
1301    /// ```
1302    ///
1303    /// # Errors
1304    ///
1305    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1306    /// unsupported/backend error for the input dtype.
1307    pub fn conj(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1308        self.run_unary(input, StdTensorOp::Conj)
1309    }
1310
1311    /// Compute the elementwise sign on this borrowed session.
1312    ///
1313    /// # Examples
1314    /// ```rust
1315    /// use tenferro_ad::{EagerRuntime, Tensor};
1316    /// let ctx = EagerRuntime::new()?;
1317    /// let y = ctx.with_eager_session(|s| {
1318    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![-2.0_f64])?)?;
1319    ///     s.sign(&x)
1320    /// })?;
1321    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[-1.0]);
1322    /// # Ok::<(), tenferro_ad::Error>(())
1323    /// ```
1324    /// # Errors
1325    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1326    pub fn sign(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1327        self.run_unary(input, StdTensorOp::Sign)
1328    }
1329
1330    /// Compute the elementwise natural logarithm on this borrowed session.
1331    ///
1332    /// # Examples
1333    /// ```rust
1334    /// use tenferro_ad::{EagerRuntime, Tensor};
1335    /// let ctx = EagerRuntime::new()?;
1336    /// let y = ctx.with_eager_session(|s| {
1337    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?)?;
1338    ///     s.log(&x)
1339    /// })?;
1340    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.0]);
1341    /// # Ok::<(), tenferro_ad::Error>(())
1342    /// ```
1343    /// # Errors
1344    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1345    pub fn log(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1346        self.run_unary(input, StdTensorOp::Log)
1347    }
1348
1349    /// Compute the elementwise square root on this borrowed session.
1350    ///
1351    /// # Examples
1352    /// ```rust
1353    /// use tenferro_ad::{EagerRuntime, Tensor};
1354    /// let ctx = EagerRuntime::new()?;
1355    /// let y = ctx.with_eager_session(|s| {
1356    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![4.0_f64])?)?;
1357    ///     s.sqrt(&x)
1358    /// })?;
1359    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[2.0]);
1360    /// # Ok::<(), tenferro_ad::Error>(())
1361    /// ```
1362    /// # Errors
1363    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1364    pub fn sqrt(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1365        self.run_unary(input, StdTensorOp::Sqrt)
1366    }
1367
1368    /// Compute the elementwise reciprocal square root on this borrowed session.
1369    ///
1370    /// # Examples
1371    /// ```rust
1372    /// use tenferro_ad::{EagerRuntime, Tensor};
1373    /// let ctx = EagerRuntime::new()?;
1374    /// let y = ctx.with_eager_session(|s| {
1375    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![4.0_f64])?)?;
1376    ///     s.rsqrt(&x)
1377    /// })?;
1378    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.5]);
1379    /// # Ok::<(), tenferro_ad::Error>(())
1380    /// ```
1381    /// # Errors
1382    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1383    pub fn rsqrt(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1384        self.run_unary(input, StdTensorOp::Rsqrt)
1385    }
1386
1387    /// Compute the elementwise sine on this borrowed session.
1388    ///
1389    /// # Examples
1390    /// ```rust
1391    /// use tenferro_ad::{EagerRuntime, Tensor};
1392    /// let ctx = EagerRuntime::new()?;
1393    /// let y = ctx.with_eager_session(|s| {
1394    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![0.0_f64])?)?;
1395    ///     s.sin(&x)
1396    /// })?;
1397    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.0]);
1398    /// # Ok::<(), tenferro_ad::Error>(())
1399    /// ```
1400    /// # Errors
1401    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1402    pub fn sin(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1403        self.run_unary(input, StdTensorOp::Sin)
1404    }
1405
1406    /// Compute the elementwise cosine on this borrowed session.
1407    ///
1408    /// # Examples
1409    /// ```rust
1410    /// use tenferro_ad::{EagerRuntime, Tensor};
1411    /// let ctx = EagerRuntime::new()?;
1412    /// let y = ctx.with_eager_session(|s| {
1413    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![0.0_f64])?)?;
1414    ///     s.cos(&x)
1415    /// })?;
1416    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[1.0]);
1417    /// # Ok::<(), tenferro_ad::Error>(())
1418    /// ```
1419    /// # Errors
1420    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1421    pub fn cos(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1422        self.run_unary(input, StdTensorOp::Cos)
1423    }
1424
1425    /// Compute the elementwise hyperbolic tangent on this borrowed session.
1426    ///
1427    /// # Examples
1428    /// ```rust
1429    /// use tenferro_ad::{EagerRuntime, Tensor};
1430    /// let ctx = EagerRuntime::new()?;
1431    /// let y = ctx.with_eager_session(|s| {
1432    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![0.0_f64])?)?;
1433    ///     s.tanh(&x)
1434    /// })?;
1435    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.0]);
1436    /// # Ok::<(), tenferro_ad::Error>(())
1437    /// ```
1438    /// # Errors
1439    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1440    pub fn tanh(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1441        self.run_unary(input, StdTensorOp::Tanh)
1442    }
1443
1444    /// Compute `exp(x) - 1` elementwise on this borrowed session.
1445    ///
1446    /// # Examples
1447    /// ```rust
1448    /// use tenferro_ad::{EagerRuntime, Tensor};
1449    /// let ctx = EagerRuntime::new()?;
1450    /// let y = ctx.with_eager_session(|s| {
1451    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![0.0_f64])?)?;
1452    ///     s.expm1(&x)
1453    /// })?;
1454    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.0]);
1455    /// # Ok::<(), tenferro_ad::Error>(())
1456    /// ```
1457    /// # Errors
1458    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1459    pub fn expm1(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1460        self.run_unary(input, StdTensorOp::Expm1)
1461    }
1462
1463    /// Compute `log(1 + x)` elementwise on this borrowed session.
1464    ///
1465    /// # Examples
1466    /// ```rust
1467    /// use tenferro_ad::{EagerRuntime, Tensor};
1468    /// let ctx = EagerRuntime::new()?;
1469    /// let y = ctx.with_eager_session(|s| {
1470    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![0.0_f64])?)?;
1471    ///     s.log1p(&x)
1472    /// })?;
1473    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.0]);
1474    /// # Ok::<(), tenferro_ad::Error>(())
1475    /// ```
1476    /// # Errors
1477    /// Returns a typed foreign-runtime, unsupported-dtype, or backend error.
1478    pub fn log1p(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1479        self.run_unary(input, StdTensorOp::Log1p)
1480    }
1481
1482    /// Compute the error function `erf(x)` elementwise on this borrowed session.
1483    ///
1484    /// Defined for real `F32`/`F64` tensors; `erf(+-0) = +-0`,
1485    /// `erf(+-inf) = +-1`, and `NaN` stays `NaN`. The derivative is
1486    /// `2/sqrt(pi) * exp(-x^2)`.
1487    ///
1488    /// # Examples
1489    /// ```rust
1490    /// use tenferro_ad::{EagerRuntime, Tensor};
1491    /// let ctx = EagerRuntime::new()?;
1492    /// let y = ctx.with_eager_session(|s| {
1493    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![0.0_f64, 1.0])?)?;
1494    ///     s.erf(&x)
1495    /// })?;
1496    /// let y = y.value()?;
1497    /// let y = y.as_slice::<f64>()?;
1498    /// assert_eq!(y[0], 0.0);
1499    /// assert!((y[1] - 0.842_700_792_949_714_9).abs() < 1.0e-15);
1500    /// # Ok::<(), tenferro_ad::Error>(())
1501    /// ```
1502    /// # Errors
1503    /// Returns a typed foreign-runtime error, a typed unsupported-dtype error
1504    /// for complex, integer, or `Bool` input, or a backend error.
1505    pub fn erf(&mut self, input: &EagerTensor) -> Result<EagerTensor> {
1506        self.run_unary(input, StdTensorOp::Erf)
1507    }
1508
1509    /// Convert a tensor under the checked dtype-promotion lattice.
1510    /// Use [`Self::cast`] for intentional lossy projection.
1511    ///
1512    /// # Examples
1513    ///
1514    /// ```rust
1515    /// use tenferro_ad::{DType, EagerRuntime, Tensor};
1516    /// use tenferro_cpu::CpuBackend;
1517    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1518    /// let converted = ctx.with_eager_session(|session| {
1519    ///     let x = session.constant_from(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?)?;
1520    ///     session.convert(&x, DType::C64)
1521    /// })?;
1522    /// assert_eq!(converted.dtype(), DType::C64);
1523    /// # Ok::<(), tenferro_ad::Error>(())
1524    /// ```
1525    ///
1526    /// # Errors
1527    ///
1528    /// Returns [`Error::ContextMismatch`] for a foreign runtime or a typed
1529    /// unsupported dtype conversion/backend error.
1530    pub fn convert(&mut self, input: &EagerTensor, to: DType) -> Result<EagerTensor> {
1531        self.ensure_runtime(input)?;
1532        tenferro_tensor::validate::validate_convert_dtype(
1533            "EagerTensor::convert",
1534            input.dtype(),
1535            to,
1536        )
1537        .map_err(Error::TensorRuntime)?;
1538        self.cast(input, to)
1539    }
1540
1541    /// Cast a tensor to a dtype, permitting explicitly lossy projections.
1542    ///
1543    /// # Examples
1544    ///
1545    /// ```rust
1546    /// use tenferro_ad::{DType, EagerRuntime, Tensor};
1547    /// use tenferro_cpu::CpuBackend;
1548    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1549    /// let casted = ctx.with_eager_session(|session| {
1550    ///     let x = session.constant_from(Tensor::from_vec_col_major(vec![1], vec![2.8_f64])?)?;
1551    ///     session.cast(&x, DType::I32)
1552    /// })?;
1553    /// assert_eq!(casted.value()?.as_slice::<i32>()?, &[2]);
1554    /// # Ok::<(), tenferro_ad::Error>(())
1555    /// ```
1556    ///
1557    /// # Errors
1558    ///
1559    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1560    /// unsupported projection/backend error.
1561    pub fn cast(&mut self, input: &EagerTensor, to: DType) -> Result<EagerTensor> {
1562        self.run_unary(
1563            input,
1564            StdTensorOp::Convert {
1565                from: input.dtype(),
1566                to,
1567            },
1568        )
1569    }
1570
1571    /// Permute the axes of an eager tensor while preserving independent ownership.
1572    ///
1573    /// # Examples
1574    ///
1575    /// ```rust
1576    /// use tenferro_ad::{EagerRuntime, Tensor};
1577    /// use tenferro_cpu::CpuBackend;
1578    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1579    /// let copied = ctx.with_eager_session(|session| {
1580    ///     let x = session.constant_from(Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?)?;
1581    ///     let y = session.transpose(&x, &[1, 0])?;
1582    ///     session.duplicate_value(&y)
1583    /// })?;
1584    /// assert_eq!(copied.as_slice::<f64>()?, &[1.0, 3.0, 2.0, 4.0]);
1585    /// # Ok::<(), tenferro_ad::Error>(())
1586    /// ```
1587    ///
1588    /// # Errors
1589    ///
1590    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1591    /// axis/backend error for an invalid permutation or copy.
1592    pub fn transpose(&mut self, input: &EagerTensor, perm: &[usize]) -> Result<EagerTensor> {
1593        self.ensure_runtime(input)?;
1594        input.transpose_in_session(perm, self.backend)
1595    }
1596
1597    /// Reshape an eager tensor while retaining a separate result owner.
1598    ///
1599    /// # Examples
1600    ///
1601    /// ```rust
1602    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1603    /// use tenferro_cpu::CpuBackend;
1604    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1605    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?, ctx.clone())?;
1606    /// let y = ctx.with_eager_session(|session| session.reshape(&x, [1, 2]))?;
1607    /// assert_eq!(y.shape(), &[1, 2]);
1608    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[1.0, 2.0]);
1609    /// # Ok::<(), tenferro_ad::Error>(())
1610    /// ```
1611    ///
1612    /// # Errors
1613    ///
1614    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1615    /// validation/backend error when the target shape is incompatible.
1616    pub fn reshape(
1617        &mut self,
1618        input: &EagerTensor,
1619        shape: impl IntoShapeVec,
1620    ) -> Result<EagerTensor> {
1621        self.ensure_runtime(input)?;
1622        input.reshape_in_session(&shape.into_shape_vec(), self.backend)
1623    }
1624
1625    /// Slice an eager tensor with explicit start, limit, and stride per axis.
1626    ///
1627    /// # Examples
1628    ///
1629    /// ```rust
1630    /// use tenferro_ad::{EagerRuntime, SliceConfig, Tensor};
1631    /// use tenferro_cpu::CpuBackend;
1632    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1633    /// let y = ctx.with_eager_session(|session| {
1634    ///     let x = session.constant_from(Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0])?)?;
1635    ///     session.slice(&x, SliceConfig { starts: vec![1], limits: vec![3], strides: vec![1] })
1636    /// })?;
1637    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[2.0, 3.0]);
1638    /// # Ok::<(), tenferro_ad::Error>(())
1639    /// ```
1640    ///
1641    /// # Errors
1642    ///
1643    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1644    /// axis/stride/backend error for an invalid slice or copy.
1645    pub fn slice(&mut self, input: &EagerTensor, config: SliceConfig) -> Result<EagerTensor> {
1646        self.ensure_runtime(input)?;
1647        input.slice_in_session(config, self.backend)
1648    }
1649
1650    /// Broadcast an eager tensor into a larger shape on this session.
1651    ///
1652    /// # Examples
1653    ///
1654    /// ```rust
1655    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1656    /// use tenferro_cpu::CpuBackend;
1657    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1658    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?, ctx.clone())?;
1659    /// let copy = ctx.with_eager_session(|session| {
1660    ///     let y = session.broadcast_in_dim(&x, &[2, 2], &[0])?;
1661    ///     session.duplicate_value(&y)
1662    /// })?;
1663    /// assert_eq!(copy.as_slice::<f64>()?, &[1.0, 2.0, 1.0, 2.0]);
1664    /// # Ok::<(), tenferro_ad::Error>(())
1665    /// ```
1666    ///
1667    /// # Errors
1668    ///
1669    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1670    /// validation/backend error for an invalid broadcast mapping.
1671    pub fn broadcast_in_dim(
1672        &mut self,
1673        input: &EagerTensor,
1674        shape: &[usize],
1675        dims: &[usize],
1676    ) -> Result<EagerTensor> {
1677        self.ensure_runtime(input)?;
1678        input.broadcast_in_dim_in_session(shape, dims, self.backend)
1679    }
1680
1681    /// Keep the lower triangle of a matrix on this borrowed session.
1682    ///
1683    /// # Examples
1684    /// ```rust
1685    /// use tenferro_ad::{EagerRuntime, Tensor};
1686    /// let ctx = EagerRuntime::new()?;
1687    /// let lower = ctx.with_eager_session(|s| {
1688    ///     let matrix = s.constant_from(Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?)?;
1689    ///     s.tril(&matrix, 0)
1690    /// })?;
1691    /// assert_eq!(lower.value()?.as_slice::<f64>()?, &[1.0, 2.0, 0.0, 4.0]);
1692    /// # Ok::<(), tenferro_ad::Error>(())
1693    /// ```
1694    /// # Errors
1695    /// Returns a typed foreign-runtime, rank, unsupported-dtype, or backend error.
1696    pub fn tril(&mut self, input: &EagerTensor, k: i64) -> Result<EagerTensor> {
1697        self.run_unary(input, StdTensorOp::Tril { k })
1698    }
1699
1700    /// Keep the upper triangle of a matrix on this borrowed session.
1701    ///
1702    /// # Examples
1703    /// ```rust
1704    /// use tenferro_ad::{EagerRuntime, Tensor};
1705    /// let ctx = EagerRuntime::new()?;
1706    /// let upper = ctx.with_eager_session(|s| {
1707    ///     let matrix = s.constant_from(Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?)?;
1708    ///     s.triu(&matrix, 0)
1709    /// })?;
1710    /// assert_eq!(upper.value()?.as_slice::<f64>()?, &[1.0, 0.0, 3.0, 4.0]);
1711    /// # Ok::<(), tenferro_ad::Error>(())
1712    /// ```
1713    /// # Errors
1714    /// Returns a typed foreign-runtime, rank, unsupported-dtype, or backend error.
1715    pub fn triu(&mut self, input: &EagerTensor, k: i64) -> Result<EagerTensor> {
1716        self.run_unary(input, StdTensorOp::Triu { k })
1717    }
1718
1719    /// Pad an eager tensor with zeros on this borrowed session.
1720    ///
1721    /// # Examples
1722    /// ```rust
1723    /// use tenferro_ad::{EagerRuntime, PadConfig, Tensor};
1724    /// let ctx = EagerRuntime::new()?;
1725    /// let padded = ctx.with_eager_session(|s| {
1726    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?)?;
1727    ///     s.pad(&x, PadConfig {
1728    ///         edge_padding_low: vec![1],
1729    ///         edge_padding_high: vec![1],
1730    ///         interior_padding: vec![1],
1731    ///     })
1732    /// })?;
1733    /// assert_eq!(padded.value()?.as_slice::<f64>()?, &[0.0, 1.0, 0.0, 2.0, 0.0]);
1734    /// # Ok::<(), tenferro_ad::Error>(())
1735    /// ```
1736    /// # Errors
1737    /// Returns [`Error::ContextMismatch`] when an input belongs to another eager
1738    /// runtime, a validation error with
1739    /// `ValidationError::InvalidArgument` for a padding configuration whose
1740    /// length or extents do not match the input rank, or
1741    /// [`Error::TensorRuntime`] for a typed backend failure.
1742    pub fn pad(&mut self, input: &EagerTensor, config: PadConfig) -> Result<EagerTensor> {
1743        self.run_unary(input, StdTensorOp::Pad(config))
1744    }
1745
1746    /// Reverse the elements along selected axes on this borrowed session.
1747    ///
1748    /// # Examples
1749    /// ```rust
1750    /// use tenferro_ad::{EagerRuntime, Tensor};
1751    /// let ctx = EagerRuntime::new()?;
1752    /// let reversed = ctx.with_eager_session(|s| {
1753    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0])?)?;
1754    ///     s.reverse(&x, &[0])
1755    /// })?;
1756    /// assert_eq!(reversed.value()?.as_slice::<f64>()?, &[3.0, 2.0, 1.0]);
1757    /// # Ok::<(), tenferro_ad::Error>(())
1758    /// ```
1759    /// # Errors
1760    /// Returns a typed foreign-runtime, invalid-axis, or backend error.
1761    pub fn reverse(&mut self, input: &EagerTensor, axes: &[usize]) -> Result<EagerTensor> {
1762        self.ensure_runtime(input)?;
1763        crate::eager_ops::validate_eager_axes("EagerSession::reverse", input.shape().len(), axes)?;
1764        self.run_unary(
1765            input,
1766            StdTensorOp::Reverse {
1767                axes: axes.to_vec(),
1768            },
1769        )
1770    }
1771
1772    /// Slice an eager tensor using runtime start indices in this borrowed session.
1773    ///
1774    /// # Examples
1775    /// ```rust
1776    /// use tenferro_ad::{EagerRuntime, Tensor};
1777    /// let ctx = EagerRuntime::new()?;
1778    /// let selected = ctx.with_eager_session(|s| {
1779    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![4], vec![1.0_f64, 2.0, 3.0, 4.0])?)?;
1780    ///     let starts = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![1_i64])?)?;
1781    ///     s.dynamic_slice(&x, &starts, &[2])
1782    /// })?;
1783    /// assert_eq!(selected.value()?.as_slice::<f64>()?, &[2.0, 3.0]);
1784    /// # Ok::<(), tenferro_ad::Error>(())
1785    /// ```
1786    /// # Errors
1787    /// Returns a typed foreign-runtime, invalid-index or slice-shape, or backend error.
1788    pub fn dynamic_slice(
1789        &mut self,
1790        input: &EagerTensor,
1791        starts: &EagerTensor,
1792        sizes: &[usize],
1793    ) -> Result<EagerTensor> {
1794        self.ensure_runtime(input)?;
1795        self.ensure_runtime(starts)?;
1796        EagerTensor::nary_op_in_session(
1797            &[input, starts],
1798            StdTensorOp::DynamicSlice {
1799                slice_sizes: sizes.to_vec(),
1800            },
1801            self.backend,
1802        )
1803    }
1804
1805    /// Gather elements of an eager tensor in this borrowed session.
1806    ///
1807    /// # Examples
1808    /// ```rust
1809    /// use tenferro_ad::{EagerRuntime, GatherConfig, Tensor};
1810    /// let ctx = EagerRuntime::new()?;
1811    /// let result = ctx.with_eager_session(|s| {
1812    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![3], vec![10.0_f64, 20.0, 30.0])?)?;
1813    ///     let indices = s.constant_from(Tensor::from_vec_col_major(vec![2, 1], vec![2_i64, 0])?)?;
1814    ///     s.gather(&x, &indices, GatherConfig {
1815    ///         offset_dims: vec![], collapsed_slice_dims: vec![0],
1816    ///         start_index_map: vec![0], index_vector_dim: 1,
1817    ///         slice_sizes: vec![1],
1818    ///     })
1819    /// })?;
1820    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[30.0, 10.0]);
1821    /// # Ok::<(), tenferro_ad::Error>(())
1822    /// ```
1823    /// # Errors
1824    /// Returns a typed foreign-runtime, invalid-index/configuration, or backend error.
1825    pub fn gather(
1826        &mut self,
1827        input: &EagerTensor,
1828        indices: &EagerTensor,
1829        config: GatherConfig,
1830    ) -> Result<EagerTensor> {
1831        self.ensure_runtime(input)?;
1832        self.ensure_runtime(indices)?;
1833        EagerTensor::nary_op_in_session(
1834            &[input, indices],
1835            StdTensorOp::Gather(config),
1836            self.backend,
1837        )
1838    }
1839
1840    /// Concatenate eager tensors along one axis in this borrowed session.
1841    ///
1842    /// # Examples
1843    /// ```rust
1844    /// use tenferro_ad::{EagerRuntime, Tensor};
1845    /// let ctx = EagerRuntime::new()?;
1846    /// let result = ctx.with_eager_session(|s| {
1847    ///     let a = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?)?;
1848    ///     let b = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?)?;
1849    ///     s.concatenate(&[&a, &b], 0)
1850    /// })?;
1851    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[1.0, 2.0]);
1852    /// # Ok::<(), tenferro_ad::Error>(())
1853    /// ```
1854    /// # Errors
1855    /// Returns a typed empty-input, foreign-runtime, invalid-axis/shape, or backend error.
1856    pub fn concatenate(&mut self, inputs: &[&EagerTensor], axis: usize) -> Result<EagerTensor> {
1857        for input in inputs {
1858            self.ensure_runtime(input)?;
1859        }
1860        EagerTensor::nary_op_in_session(
1861            inputs,
1862            StdTensorOp::Concatenate {
1863                axis,
1864                input_count: inputs.len(),
1865            },
1866            self.backend,
1867        )
1868    }
1869
1870    /// Scatter updates into an eager tensor within this borrowed session.
1871    ///
1872    /// # Examples
1873    /// ```rust
1874    /// use tenferro_ad::{EagerRuntime, ScatterConfig, Tensor};
1875    /// let ctx = EagerRuntime::new()?;
1876    /// let result = ctx.with_eager_session(|s| {
1877    ///     let input = s.constant_from(Tensor::from_vec_col_major(vec![4], vec![0.0_f64; 4])?)?;
1878    ///     let indices = s.constant_from(Tensor::from_vec_col_major(vec![2, 1], vec![1_i64, 3])?)?;
1879    ///     let updates = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![5.0_f64, 7.0])?)?;
1880    ///     s.scatter(&input, &indices, &updates, ScatterConfig {
1881    ///         update_window_dims: vec![],
1882    ///         inserted_window_dims: vec![0],
1883    ///         scatter_dims_to_operand_dims: vec![0],
1884    ///         index_vector_dim: 1,
1885    ///     })
1886    /// })?;
1887    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[0.0, 5.0, 0.0, 7.0]);
1888    /// # Ok::<(), tenferro_ad::Error>(())
1889    /// ```
1890    /// # Errors
1891    /// Returns a typed foreign-runtime, invalid-index/configuration, or backend error.
1892    pub fn scatter(
1893        &mut self,
1894        input: &EagerTensor,
1895        indices: &EagerTensor,
1896        updates: &EagerTensor,
1897        config: ScatterConfig,
1898    ) -> Result<EagerTensor> {
1899        self.ensure_runtime(input)?;
1900        self.ensure_runtime(indices)?;
1901        self.ensure_runtime(updates)?;
1902        EagerTensor::nary_op_in_session(
1903            &[input, indices, updates],
1904            StdTensorOp::Scatter(config),
1905            self.backend,
1906        )
1907    }
1908
1909    /// Extract a diagonal along two axes in this borrowed session.
1910    ///
1911    /// # Examples
1912    /// ```rust
1913    /// use tenferro_ad::{EagerRuntime, Tensor};
1914    /// let ctx = EagerRuntime::new()?;
1915    /// let diagonal = ctx.with_eager_session(|s| {
1916    ///     let matrix = s.constant_from(Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0])?)?;
1917    ///     s.extract_diag(&matrix, 0, 1)
1918    /// })?;
1919    /// assert_eq!(diagonal.value()?.as_slice::<f64>()?, &[1.0, 4.0]);
1920    /// # Ok::<(), tenferro_ad::Error>(())
1921    /// ```
1922    /// # Errors
1923    /// Returns a typed foreign-runtime, invalid-axis, or backend error.
1924    pub fn extract_diag(
1925        &mut self,
1926        input: &EagerTensor,
1927        axis_a: usize,
1928        axis_b: usize,
1929    ) -> Result<EagerTensor> {
1930        self.ensure_runtime(input)?;
1931        self.run_unary(input, StdTensorOp::ExtractDiag { axis_a, axis_b })
1932    }
1933
1934    /// Embed the input along a diagonal in this borrowed session.
1935    ///
1936    /// # Examples
1937    /// ```rust
1938    /// use tenferro_ad::{EagerRuntime, Tensor};
1939    /// let ctx = EagerRuntime::new()?;
1940    /// let matrix = ctx.with_eager_session(|s| {
1941    ///     let diagonal = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?)?;
1942    ///     s.embed_diag(&diagonal, 0, 1)
1943    /// })?;
1944    /// assert_eq!(matrix.value()?.as_slice::<f64>()?, &[1.0, 0.0, 0.0, 2.0]);
1945    /// # Ok::<(), tenferro_ad::Error>(())
1946    /// ```
1947    /// # Errors
1948    /// Returns a typed foreign-runtime, invalid-axis, or backend error.
1949    pub fn embed_diag(
1950        &mut self,
1951        input: &EagerTensor,
1952        axis_a: usize,
1953        axis_b: usize,
1954    ) -> Result<EagerTensor> {
1955        self.ensure_runtime(input)?;
1956        self.run_unary(input, StdTensorOp::EmbedDiag { axis_a, axis_b })
1957    }
1958
1959    /// Reduce an eager tensor over selected axes within this borrowed session.
1960    /// `None` reduces all axes, while `Some(&[])` retains the input shape.
1961    ///
1962    /// # Examples
1963    ///
1964    /// ```rust
1965    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
1966    /// use tenferro_cpu::CpuBackend;
1967    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
1968    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?, ctx.clone())?;
1969    /// let sum = ctx.with_eager_session(|session| session.reduce_sum(&x, None))?;
1970    /// assert_eq!(sum.value()?.as_slice::<f64>()?, &[3.0]);
1971    /// # Ok::<(), tenferro_ad::Error>(())
1972    /// ```
1973    ///
1974    /// # Errors
1975    ///
1976    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
1977    /// validation/backend error for invalid axes or unsupported dtypes.
1978    pub fn reduce_sum(
1979        &mut self,
1980        input: &EagerTensor,
1981        axes: Option<&[usize]>,
1982    ) -> Result<EagerTensor> {
1983        self.ensure_runtime(input)?;
1984        input.reduce_sum_in_session(axes, self.backend)
1985    }
1986
1987    /// Sum elementwise squares over the selected axes in this borrowed session.
1988    /// Only `f32` and `f64` are supported. `None` reduces every axis, like the
1989    /// rest of the reduction family; `Some(&[])` squares each value.
1990    ///
1991    /// # Examples
1992    /// ```rust
1993    /// use tenferro_ad::{EagerRuntime, Tensor};
1994    /// let ctx = EagerRuntime::new()?;
1995    /// let (sum, all) = ctx.with_eager_session(|s| {
1996    ///     let input = s.constant_from(Tensor::from_vec_col_major([2], vec![3.0_f64, 4.0])?)?;
1997    ///     Ok::<_, tenferro_ad::Error>((
1998    ///         s.reduce_sum_squares(&input, Some(&[0]))?,
1999    ///         s.reduce_sum_squares(&input, None)?,
2000    ///     ))
2001    /// })?;
2002    /// assert_eq!(sum.value()?.as_slice::<f64>()?, &[25.0]);
2003    /// assert_eq!(all.value()?.as_slice::<f64>()?, &[25.0]);
2004    /// # Ok::<(), tenferro_ad::Error>(())
2005    /// ```
2006    /// # Errors
2007    /// Returns typed foreign-runtime, invalid-axis, unsupported-dtype, or backend errors.
2008    pub fn reduce_sum_squares(
2009        &mut self,
2010        input: &EagerTensor,
2011        axes: Option<&[usize]>,
2012    ) -> Result<EagerTensor> {
2013        self.ensure_runtime(input)?;
2014        let axes = axes.map_or_else(|| (0..input.shape().len()).collect(), <[usize]>::to_vec);
2015        crate::eager_ops::validate_eager_axes(
2016            "EagerSession::reduce_sum_squares",
2017            input.shape().len(),
2018            &axes,
2019        )?;
2020        self.run_unary(input, StdTensorOp::ReduceSumSquares { axes })
2021    }
2022
2023    /// Reduce the product of selected axes in this borrowed session.
2024    /// `None` reduces every axis.
2025    ///
2026    /// # Examples
2027    /// ```rust
2028    /// use tenferro_ad::{EagerRuntime, Tensor};
2029    /// let ctx = EagerRuntime::new()?;
2030    /// let result = ctx.with_eager_session(|s| {
2031    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?)?;
2032    ///     s.reduce_prod(&x, None)
2033    /// })?;
2034    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[6.0]);
2035    /// # Ok::<(), tenferro_ad::Error>(())
2036    /// ```
2037    /// # Errors
2038    /// Returns a typed foreign-runtime, invalid-axis, unsupported-dtype, or backend error.
2039    pub fn reduce_prod(
2040        &mut self,
2041        input: &EagerTensor,
2042        axes: Option<&[usize]>,
2043    ) -> Result<EagerTensor> {
2044        self.ensure_runtime(input)?;
2045        let axes = axes.map_or_else(|| (0..input.shape().len()).collect(), <[usize]>::to_vec);
2046        crate::eager_ops::validate_eager_axes(
2047            "EagerSession::reduce_prod",
2048            input.shape().len(),
2049            &axes,
2050        )?;
2051        self.run_unary(input, StdTensorOp::ReduceProd { axes })
2052    }
2053
2054    /// Reduce the maximum over selected axes in this borrowed session.
2055    /// `None` reduces every axis.
2056    ///
2057    /// # Examples
2058    /// ```rust
2059    /// use tenferro_ad::{EagerRuntime, Tensor};
2060    /// let ctx = EagerRuntime::new()?;
2061    /// let result = ctx.with_eager_session(|s| {
2062    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?)?;
2063    ///     s.reduce_max(&x, None)
2064    /// })?;
2065    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[3.0]);
2066    /// # Ok::<(), tenferro_ad::Error>(())
2067    /// ```
2068    /// # Errors
2069    /// Returns a typed foreign-runtime, invalid-axis, unsupported-dtype, or backend error.
2070    pub fn reduce_max(
2071        &mut self,
2072        input: &EagerTensor,
2073        axes: Option<&[usize]>,
2074    ) -> Result<EagerTensor> {
2075        self.ensure_runtime(input)?;
2076        let axes = axes.map_or_else(|| (0..input.shape().len()).collect(), <[usize]>::to_vec);
2077        crate::eager_ops::validate_eager_axes(
2078            "EagerSession::reduce_max",
2079            input.shape().len(),
2080            &axes,
2081        )?;
2082        self.run_unary(input, StdTensorOp::ReduceMax { axes })
2083    }
2084
2085    /// Reduce the minimum over selected axes in this borrowed session.
2086    /// `None` reduces every axis.
2087    ///
2088    /// # Examples
2089    /// ```rust
2090    /// use tenferro_ad::{EagerRuntime, Tensor};
2091    /// let ctx = EagerRuntime::new()?;
2092    /// let result = ctx.with_eager_session(|s| {
2093    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?)?;
2094    ///     s.reduce_min(&x, None)
2095    /// })?;
2096    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[2.0]);
2097    /// # Ok::<(), tenferro_ad::Error>(())
2098    /// ```
2099    /// # Errors
2100    /// Returns a typed foreign-runtime, invalid-axis, unsupported-dtype, or backend error.
2101    pub fn reduce_min(
2102        &mut self,
2103        input: &EagerTensor,
2104        axes: Option<&[usize]>,
2105    ) -> Result<EagerTensor> {
2106        self.ensure_runtime(input)?;
2107        let axes = axes.map_or_else(|| (0..input.shape().len()).collect(), <[usize]>::to_vec);
2108        crate::eager_ops::validate_eager_axes(
2109            "EagerSession::reduce_min",
2110            input.shape().len(),
2111            &axes,
2112        )?;
2113        self.run_unary(input, StdTensorOp::ReduceMin { axes })
2114    }
2115
2116    /// Duplicate an eager value into an independent tensor within the caller's
2117    /// borrowed session, preserving its dtype and placement.
2118    ///
2119    /// # Examples
2120    ///
2121    /// ```rust
2122    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
2123    /// use tenferro_cpu::CpuBackend;
2124    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2125    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?, ctx.clone())?;
2126    /// let copy = ctx.with_eager_session(|session| session.duplicate_value(&x))?;
2127    /// assert_eq!(copy.as_slice::<f64>()?, &[2.0]);
2128    /// # Ok::<(), tenferro_ad::Error>(())
2129    /// ```
2130    ///
2131    /// # Errors
2132    ///
2133    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
2134    /// runtime/backend error when the retained value cannot be duplicated.
2135    pub fn duplicate_value(&mut self, input: &EagerTensor) -> Result<Tensor> {
2136        self.ensure_runtime(input)?;
2137        input.duplicate_value_in_session(self.backend)
2138    }
2139
2140    /// Import an untracked leaf within this borrowed session.
2141    ///
2142    /// `tensor` must already be usable by this runtime's backend: a host
2143    /// tensor on a CPU runtime, or a tensor already on the device of a CUDA
2144    /// or WebGPU runtime. No host/device transfer happens here. To import host
2145    /// data into a device runtime, use [`Self::constant_from_host`], which
2146    /// uploads first; on a CPU runtime the two are equivalent.
2147    ///
2148    /// # Examples
2149    ///
2150    /// ```rust
2151    /// use tenferro_ad::{EagerRuntime, Tensor};
2152    /// use tenferro_cpu::CpuBackend;
2153    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2154    /// let c = ctx.with_eager_session(|session| {
2155    ///     session.constant_from(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?)
2156    /// })?;
2157    /// assert_eq!(c.value()?.as_slice::<f64>()?, &[2.0]);
2158    /// # Ok::<(), tenferro_ad::Error>(())
2159    /// ```
2160    ///
2161    /// # Errors
2162    ///
2163    /// Returns [`Error::TensorRuntime`] for a typed backend failure when the value cannot be
2164    /// registered in the session, or [`Error::RuntimeState`] when the runtime's
2165    /// value registry is unavailable.
2166    pub fn constant_from(&mut self, tensor: Tensor) -> Result<EagerTensor> {
2167        EagerTensor::new_leaf_in_session(Arc::clone(self.runtime), tensor, false, self.backend)
2168    }
2169
2170    /// Upload a host tensor and import it as an untracked leaf in this session.
2171    ///
2172    /// Unlike [`Self::constant_from`], this explicitly crosses the host/device
2173    /// boundary: the host `tensor` is uploaded to this runtime's backend
2174    /// (a host copy on a CPU runtime) and the uploaded value becomes the leaf.
2175    /// Use it whenever the source data lives on the host and the runtime may
2176    /// be a device runtime; use [`Self::constant_from`] for a tensor that is
2177    /// already resident on the backend.
2178    ///
2179    /// # Examples
2180    /// ```rust
2181    /// use tenferro_ad::{EagerRuntime, Tensor};
2182    /// let ctx = EagerRuntime::new()?;
2183    /// let c = ctx.with_eager_session(|s| {
2184    ///     s.constant_from_host(Tensor::from_vec_col_major([1], vec![2.0_f64])?)
2185    /// })?;
2186    /// assert_eq!(c.value()?.as_slice::<f64>()?, &[2.0]);
2187    /// # Ok::<(), tenferro_ad::Error>(())
2188    /// ```
2189    /// # Errors
2190    /// Returns [`Error::TensorRuntime`] for a typed backend failure, including a host-tensor
2191    /// upload failure, or [`Error::RuntimeState`] when the runtime's value
2192    /// registry is unavailable.
2193    pub fn constant_from_host(&mut self, tensor: Tensor) -> Result<EagerTensor> {
2194        let uploaded = self
2195            .backend
2196            .upload_host_tensor(TensorRead::from_tensor(&tensor))
2197            .map_err(Error::from)?;
2198        self.constant_from(uploaded)
2199    }
2200
2201    /// Import a trainable leaf within this borrowed session.
2202    ///
2203    /// # Examples
2204    ///
2205    /// ```rust
2206    /// use tenferro_ad::{EagerRuntime, Tensor};
2207    /// use tenferro_cpu::CpuBackend;
2208    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2209    /// let x = ctx.with_eager_session(|session| {
2210    ///     session.variable_from(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?)
2211    /// })?;
2212    /// assert!(x.tracks_grad());
2213    /// # Ok::<(), tenferro_ad::Error>(())
2214    /// ```
2215    ///
2216    /// # Errors
2217    ///
2218    /// Returns [`Error::TensorRuntime`] for a typed backend failure when the value cannot be
2219    /// registered, or [`Error::RuntimeState`] when the runtime's value or
2220    /// gradient registry is unavailable.
2221    pub fn variable_from(&mut self, tensor: Tensor) -> Result<EagerTensor> {
2222        EagerTensor::new_leaf_in_session(Arc::clone(self.runtime), tensor, true, self.backend)
2223    }
2224
2225    /// Add eager tensors with the same broadcast and AD rules as the eager
2226    /// operation surface, reusing this borrowed execution session.
2227    ///
2228    /// # Examples
2229    ///
2230    /// ```rust
2231    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
2232    /// use tenferro_cpu::CpuBackend;
2233    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2234    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?, ctx.clone())?;
2235    /// let scalar = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![], vec![3.0_f64])?, ctx.clone())?;
2236    /// let y = ctx.with_eager_session(|session| session.add(&x, &scalar))?;
2237    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[4.0, 5.0]);
2238    /// # Ok::<(), tenferro_ad::Error>(())
2239    /// ```
2240    ///
2241    /// # Errors
2242    ///
2243    /// Returns [`Error::ContextMismatch`] for a tensor from another runtime,
2244    /// or a typed broadcast/backend error for the operands.
2245    pub fn add(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2246        self.run_binary("add", lhs, rhs, StdTensorOp::Add)
2247    }
2248
2249    /// Subtract eager tensors within this borrowed session.
2250    ///
2251    /// # Examples
2252    ///
2253    /// ```rust
2254    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
2255    /// use tenferro_cpu::CpuBackend;
2256    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2257    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![3.0_f64])?, ctx.clone())?;
2258    /// let y = ctx.with_eager_session(|session| session.sub(&x, &x))?;
2259    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[0.0]);
2260    /// # Ok::<(), tenferro_ad::Error>(())
2261    /// ```
2262    ///
2263    /// # Errors
2264    ///
2265    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
2266    /// broadcast/backend error for the operands.
2267    pub fn sub(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2268        self.run_binary("sub", lhs, rhs, StdTensorOp::Sub)
2269    }
2270
2271    /// Multiply eager tensors within this borrowed session.
2272    ///
2273    /// # Examples
2274    ///
2275    /// ```rust
2276    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
2277    /// use tenferro_cpu::CpuBackend;
2278    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2279    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![3.0_f64])?, ctx.clone())?;
2280    /// let y = ctx.with_eager_session(|session| session.mul(&x, &x))?;
2281    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[9.0]);
2282    /// # Ok::<(), tenferro_ad::Error>(())
2283    /// ```
2284    ///
2285    /// # Errors
2286    ///
2287    /// Returns [`Error::ContextMismatch`] for a foreign runtime, or a typed
2288    /// broadcast/backend error for the operands.
2289    pub fn mul(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2290        self.run_binary("mul", lhs, rhs, StdTensorOp::Mul)
2291    }
2292
2293    /// Divide eager tensors elementwise with broadcast rules.
2294    ///
2295    /// # Examples
2296    /// ```rust
2297    /// use tenferro_ad::{EagerRuntime, Tensor};
2298    /// let ctx = EagerRuntime::new()?;
2299    /// let y = ctx.with_eager_session(|s| {
2300    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![6.0_f64])?)?;
2301    ///     let divisor = s.constant_from(Tensor::from_vec_col_major(vec![], vec![2.0_f64])?)?;
2302    ///     s.div(&x, &divisor)
2303    /// })?;
2304    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[3.0]);
2305    /// # Ok::<(), tenferro_ad::Error>(())
2306    /// ```
2307    /// # Errors
2308    /// Returns [`Error::ContextMismatch`] when an input belongs to another eager
2309    /// runtime, a validation error with
2310    /// `ValidationError::ShapeMismatch` when the operands cannot broadcast, or
2311    /// [`Error::TensorRuntime`] for a typed backend failure (including integer division by
2312    /// zero).
2313    pub fn div(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2314        self.run_binary("div", lhs, rhs, StdTensorOp::Div)
2315    }
2316
2317    /// Compute the elementwise remainder with broadcast rules.
2318    ///
2319    /// # Examples
2320    /// ```rust
2321    /// use tenferro_ad::{EagerRuntime, Tensor};
2322    /// let ctx = EagerRuntime::new()?;
2323    /// let y = ctx.with_eager_session(|s| {
2324    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![5.0_f64])?)?;
2325    ///     let divisor = s.constant_from(Tensor::from_vec_col_major(vec![], vec![2.0_f64])?)?;
2326    ///     s.rem(&x, &divisor)
2327    /// })?;
2328    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[1.0]);
2329    /// # Ok::<(), tenferro_ad::Error>(())
2330    /// ```
2331    /// # Errors
2332    /// Returns [`Error::ContextMismatch`] when an input belongs to another eager
2333    /// runtime, a validation error with
2334    /// `ValidationError::ShapeMismatch` when the operands cannot broadcast, or
2335    /// [`Error::TensorRuntime`] for a typed backend failure (including an integer remainder by
2336    /// zero).
2337    pub fn rem(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2338        self.run_binary("rem", lhs, rhs, StdTensorOp::Rem)
2339    }
2340
2341    /// Raise eager tensor elements to broadcast exponents.
2342    ///
2343    /// # Examples
2344    /// ```rust
2345    /// use tenferro_ad::{EagerRuntime, Tensor};
2346    /// let ctx = EagerRuntime::new()?;
2347    /// let y = ctx.with_eager_session(|s| {
2348    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?)?;
2349    ///     let exponent = s.constant_from(Tensor::from_vec_col_major(vec![], vec![3.0_f64])?)?;
2350    ///     s.pow(&x, &exponent)
2351    /// })?;
2352    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[8.0]);
2353    /// # Ok::<(), tenferro_ad::Error>(())
2354    /// ```
2355    /// # Errors
2356    /// Returns [`Error::ContextMismatch`] when an input belongs to another eager
2357    /// runtime, a validation error with
2358    /// `ValidationError::ShapeMismatch` when the operands cannot broadcast, or
2359    /// [`Error::TensorRuntime`] for a typed backend failure (including a negative integer
2360    /// exponent).
2361    pub fn pow(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2362        self.run_binary("pow", lhs, rhs, StdTensorOp::Pow)
2363    }
2364
2365    /// Compute the elementwise maximum under broadcast rules.
2366    ///
2367    /// # Examples
2368    /// ```rust
2369    /// use tenferro_ad::{EagerRuntime, Tensor};
2370    /// let ctx = EagerRuntime::new()?;
2371    /// let y = ctx.with_eager_session(|s| {
2372    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?)?;
2373    ///     let bound = s.constant_from(Tensor::from_vec_col_major(vec![], vec![3.0_f64])?)?;
2374    ///     s.maximum(&x, &bound)
2375    /// })?;
2376    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[3.0]);
2377    /// # Ok::<(), tenferro_ad::Error>(())
2378    /// ```
2379    /// # Errors
2380    /// Returns a typed foreign-runtime, broadcast, unsupported-dtype, or backend error.
2381    pub fn maximum(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2382        self.run_binary("maximum", lhs, rhs, StdTensorOp::Maximum)
2383    }
2384
2385    /// Compute the elementwise minimum under broadcast rules.
2386    ///
2387    /// # Examples
2388    /// ```rust
2389    /// use tenferro_ad::{EagerRuntime, Tensor};
2390    /// let ctx = EagerRuntime::new()?;
2391    /// let y = ctx.with_eager_session(|s| {
2392    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?)?;
2393    ///     let bound = s.constant_from(Tensor::from_vec_col_major(vec![], vec![3.0_f64])?)?;
2394    ///     s.minimum(&x, &bound)
2395    /// })?;
2396    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[2.0]);
2397    /// # Ok::<(), tenferro_ad::Error>(())
2398    /// ```
2399    /// # Errors
2400    /// Returns a typed foreign-runtime, broadcast, unsupported-dtype, or backend error.
2401    pub fn minimum(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2402        self.run_binary("minimum", lhs, rhs, StdTensorOp::Minimum)
2403    }
2404
2405    /// Compare eager tensors elementwise under broadcast rules.
2406    ///
2407    /// # Examples
2408    /// ```rust
2409    /// use tenferro_ad::{CompareDir, EagerRuntime, Tensor};
2410    /// let ctx = EagerRuntime::new()?;
2411    /// let y = ctx.with_eager_session(|s| {
2412    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?)?;
2413    ///     let bound = s.constant_from(Tensor::from_vec_col_major(vec![], vec![1.0_f64])?)?;
2414    ///     s.compare(&x, &bound, CompareDir::Gt)
2415    /// })?;
2416    /// assert_eq!(y.value()?.as_slice::<bool>()?, &[true]);
2417    /// # Ok::<(), tenferro_ad::Error>(())
2418    /// ```
2419    /// # Errors
2420    /// Returns a typed foreign-runtime, broadcast, unsupported-dtype, or backend error.
2421    pub fn compare(
2422        &mut self,
2423        lhs: &EagerTensor,
2424        rhs: &EagerTensor,
2425        dir: CompareDir,
2426    ) -> Result<EagerTensor> {
2427        self.run_binary("compare", lhs, rhs, StdTensorOp::Compare(dir))
2428    }
2429
2430    /// Select eager values elementwise using a broadcast boolean condition.
2431    ///
2432    /// # Examples
2433    /// ```rust
2434    /// use tenferro_ad::{EagerRuntime, Tensor};
2435    /// let ctx = EagerRuntime::new()?;
2436    /// let y = ctx.with_eager_session(|s| {
2437    ///     let condition = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![true, false])?)?;
2438    ///     let yes = s.constant_from(Tensor::from_vec_col_major(vec![], vec![10.0_f64])?)?;
2439    ///     let no = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?)?;
2440    ///     s.where_select(&condition, &yes, &no)
2441    /// })?;
2442    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[10.0, 2.0]);
2443    /// # Ok::<(), tenferro_ad::Error>(())
2444    /// ```
2445    /// # Errors
2446    /// Returns a typed foreign-runtime, broadcast, unsupported-dtype, or backend error.
2447    pub fn where_select(
2448        &mut self,
2449        condition: &EagerTensor,
2450        on_true: &EagerTensor,
2451        on_false: &EagerTensor,
2452    ) -> Result<EagerTensor> {
2453        self.run_ternary(
2454            "where_select",
2455            condition,
2456            on_true,
2457            on_false,
2458            StdTensorOp::Select,
2459        )
2460    }
2461
2462    /// Alias for [`Self::where_select`] with the same borrowed-session semantics.
2463    ///
2464    /// # Examples
2465    /// ```rust
2466    /// use tenferro_ad::{EagerRuntime, Tensor};
2467    /// let ctx = EagerRuntime::new()?;
2468    /// let y = ctx.with_eager_session(|s| {
2469    ///     let predicate = s.constant_from(Tensor::from_vec_col_major(vec![], vec![true])?)?;
2470    ///     let yes = s.constant_from(Tensor::from_vec_col_major(vec![], vec![3.0_f64])?)?;
2471    ///     let no = s.constant_from(Tensor::from_vec_col_major(vec![], vec![4.0_f64])?)?;
2472    ///     s.select(&predicate, &yes, &no)
2473    /// })?;
2474    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[3.0]);
2475    /// # Ok::<(), tenferro_ad::Error>(())
2476    /// ```
2477    /// # Errors
2478    /// Returns a typed foreign-runtime, broadcast, unsupported-dtype, or backend error.
2479    pub fn select(
2480        &mut self,
2481        condition: &EagerTensor,
2482        on_true: &EagerTensor,
2483        on_false: &EagerTensor,
2484    ) -> Result<EagerTensor> {
2485        self.where_select(condition, on_true, on_false)
2486    }
2487
2488    /// Clamp eager values elementwise between broadcast lower and upper bounds.
2489    ///
2490    /// # Examples
2491    /// ```rust
2492    /// use tenferro_ad::{EagerRuntime, Tensor};
2493    /// let ctx = EagerRuntime::new()?;
2494    /// let y = ctx.with_eager_session(|s| {
2495    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![-2.0_f64, 5.0])?)?;
2496    ///     let lo = s.constant_from(Tensor::from_vec_col_major(vec![], vec![-1.0_f64])?)?;
2497    ///     let hi = s.constant_from(Tensor::from_vec_col_major(vec![], vec![4.0_f64])?)?;
2498    ///     s.clamp(&x, &lo, &hi)
2499    /// })?;
2500    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[-1.0, 4.0]);
2501    /// # Ok::<(), tenferro_ad::Error>(())
2502    /// ```
2503    /// # Errors
2504    /// Returns a typed foreign-runtime, broadcast, unsupported-dtype, or backend error.
2505    pub fn clamp(
2506        &mut self,
2507        input: &EagerTensor,
2508        lower: &EagerTensor,
2509        upper: &EagerTensor,
2510    ) -> Result<EagerTensor> {
2511        self.run_ternary("clamp", input, lower, upper, StdTensorOp::Clamp)
2512    }
2513
2514    /// Contract eager tensors according to a dot-general dimension mapping.
2515    ///
2516    /// The output layout is `[lhs free..., rhs free..., batch...]`: batch axes
2517    /// come last (see [`DotGeneralConfig`](tenferro_runtime::DotGeneralConfig)).
2518    ///
2519    /// # Examples
2520    ///
2521    /// ```rust
2522    /// use tenferro_ad::{DotGeneralConfig, EagerRuntime, Tensor};
2523    /// use tenferro_cpu::CpuBackend;
2524    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2525    /// let result = ctx.with_eager_session(|session| {
2526    ///     let lhs = session.variable_from(Tensor::from_vec_col_major(vec![1, 2], vec![2.0_f64, 3.0])?)?;
2527    ///     let rhs = session.constant_from(Tensor::from_vec_col_major(vec![2, 1], vec![4.0_f64, 5.0])?)?;
2528    ///     session.dot_general(&lhs, &rhs, DotGeneralConfig {
2529    ///         lhs_contracting_dims: [1].as_slice().into(),
2530    ///         rhs_contracting_dims: [0].as_slice().into(),
2531    ///         lhs_batch_dims: [].as_slice().into(),
2532    ///         rhs_batch_dims: [].as_slice().into(),
2533    ///     })
2534    /// })?;
2535    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[23.0]);
2536    /// # Ok::<(), tenferro_ad::Error>(())
2537    /// ```
2538    ///
2539    /// # Errors
2540    ///
2541    /// Returns [`Error::ContextMismatch`] for a foreign eager runtime,
2542    /// a typed validation error for incompatible contraction dimensions,
2543    /// or the backend's typed execution error.
2544    pub fn dot_general(
2545        &mut self,
2546        lhs: &EagerTensor,
2547        rhs: &EagerTensor,
2548        config: DotGeneralConfig,
2549    ) -> Result<EagerTensor> {
2550        self.ensure_runtime(lhs)?;
2551        self.ensure_runtime(rhs)?;
2552        config
2553            .validate_dims_with_ranks(lhs.shape().len(), rhs.shape().len())
2554            .map_err(Error::TensorRuntime)?;
2555        EagerTensor::nary_op_in_session(
2556            &[lhs, rhs],
2557            StdTensorOp::DotGeneral { config },
2558            self.backend,
2559        )
2560    }
2561
2562    /// Scale an eager tensor by a real scalar in this borrowed session.
2563    /// Integer factors are rounded; finite zero maps to `false` for boolean inputs.
2564    ///
2565    /// # Examples
2566    /// ```rust
2567    /// use tenferro_ad::{EagerRuntime, Tensor};
2568    /// let ctx = EagerRuntime::new()?;
2569    /// let scaled = ctx.with_eager_session(|s| {
2570    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?)?;
2571    ///     s.scale_real(&x, 2.0)
2572    /// })?;
2573    /// assert_eq!(scaled.value()?.as_slice::<f64>()?, &[2.0, 4.0]);
2574    /// # Ok::<(), tenferro_ad::Error>(())
2575    /// ```
2576    /// # Errors
2577    /// Returns a typed foreign-runtime, invalid-factor/dtype, or backend error.
2578    pub fn scale_real(&mut self, input: &EagerTensor, factor: f64) -> Result<EagerTensor> {
2579        self.ensure_runtime(input)?;
2580        let scalar = tenferro_runtime::scale::real_scale_scalar(input.dtype(), factor)?;
2581        let scalar = self.constant_from(scalar)?;
2582        self.mul(input, &scalar)
2583    }
2584
2585    /// Scale a complex eager tensor by a complex scalar in this borrowed session.
2586    ///
2587    /// # Examples
2588    /// ```rust
2589    /// use num_complex::Complex64;
2590    /// use tenferro_ad::{EagerRuntime, Tensor};
2591    /// let ctx = EagerRuntime::new()?;
2592    /// let scaled = ctx.with_eager_session(|s| {
2593    ///     let x = s.constant_from(Tensor::from_vec_col_major(vec![1], vec![Complex64::new(1.0, 2.0)])?)?;
2594    ///     s.scale_complex(&x, Complex64::new(0.0, 1.0))
2595    /// })?;
2596    /// assert_eq!(scaled.value()?.as_slice::<Complex64>()?, &[Complex64::new(-2.0, 1.0)]);
2597    /// # Ok::<(), tenferro_ad::Error>(())
2598    /// ```
2599    /// # Errors
2600    /// Returns [`Error::ContextMismatch`] when an input belongs to another eager
2601    /// runtime, [`Error::TensorRuntime`] containing
2602    /// `ValidationError::InvalidArgument` when the input dtype is not complex,
2603    /// or [`Error::TensorRuntime`] for a typed backend failure.
2604    pub fn scale_complex(&mut self, input: &EagerTensor, factor: Complex64) -> Result<EagerTensor> {
2605        self.ensure_runtime(input)?;
2606        let scalar = tenferro_runtime::scale::complex_scale_scalar(input.dtype(), factor)?;
2607        let scalar = self.constant_from(scalar)?;
2608        self.mul(input, &scalar)
2609    }
2610
2611    /// Multiply two rank-2 eager tensors in this borrowed session.
2612    ///
2613    /// # Examples
2614    /// ```rust
2615    /// use tenferro_ad::{EagerRuntime, Tensor};
2616    /// let ctx = EagerRuntime::new()?;
2617    /// let result = ctx.with_eager_session(|s| {
2618    ///     let a = s.constant_from(Tensor::from_vec_col_major(vec![1, 1], vec![2.0_f64])?)?;
2619    ///     let b = s.constant_from(Tensor::from_vec_col_major(vec![1, 1], vec![3.0_f64])?)?;
2620    ///     s.matmul(&a, &b)
2621    /// })?;
2622    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[6.0]);
2623    /// # Ok::<(), tenferro_ad::Error>(())
2624    /// ```
2625    /// # Errors
2626    /// Returns [`Error::ContextMismatch`] when an input belongs to another eager
2627    /// runtime, a validation error with
2628    /// `ValidationError::RankMismatch` or `ValidationError::ShapeMismatch` when
2629    /// the operands are not rank-2 with matching inner dimensions, a dtype
2630    /// mismatch between the operands, or [`Error::TensorRuntime`] for a typed backend failure.
2631    pub fn matmul(&mut self, lhs: &EagerTensor, rhs: &EagerTensor) -> Result<EagerTensor> {
2632        self.ensure_runtime(lhs)?;
2633        self.ensure_runtime(rhs)?;
2634        let lhs_shape = lhs.shape();
2635        let rhs_shape = rhs.shape();
2636        if lhs_shape.len() != 2 {
2637            return Err(tenferro_tensor::Error::rank_mismatch("matmul", 2, lhs_shape.len()).into());
2638        }
2639        if rhs_shape.len() != 2 {
2640            return Err(tenferro_tensor::Error::rank_mismatch("matmul", 2, rhs_shape.len()).into());
2641        }
2642        if lhs_shape[1] != rhs_shape[0] {
2643            return Err(
2644                tenferro_tensor::Error::shape_mismatch("matmul", lhs_shape, rhs_shape).into(),
2645            );
2646        }
2647        self.dot_general(
2648            lhs,
2649            rhs,
2650            DotGeneralConfig {
2651                lhs_contracting_dims: [1].as_slice().into(),
2652                rhs_contracting_dims: [0].as_slice().into(),
2653                lhs_batch_dims: [].as_slice().into(),
2654                rhs_batch_dims: [].as_slice().into(),
2655            },
2656        )
2657    }
2658
2659    /// Contract eagerly with optional conjugation of either operand.
2660    /// Untracked operands use the backend's conjugating contraction directly;
2661    /// tracked operands record explicit conjugations for reverse-mode AD.
2662    ///
2663    /// # Examples
2664    ///
2665    /// ```rust
2666    /// use tenferro_ad::{DotGeneralConfig, EagerRuntime, Tensor};
2667    /// use tenferro_cpu::CpuBackend;
2668    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2669    /// let result = ctx.with_eager_session(|session| {
2670    ///     let lhs = session.constant_from(Tensor::from_vec_col_major(vec![1, 1], vec![2.0_f64])?)?;
2671    ///     let rhs = session.constant_from(Tensor::from_vec_col_major(vec![1, 1], vec![3.0_f64])?)?;
2672    ///     session.dot_general_with_conj(&lhs, &rhs, DotGeneralConfig {
2673    ///         lhs_contracting_dims: [1].as_slice().into(),
2674    ///         rhs_contracting_dims: [0].as_slice().into(),
2675    ///         lhs_batch_dims: [].as_slice().into(),
2676    ///         rhs_batch_dims: [].as_slice().into(),
2677    ///     }, true, false)
2678    /// })?;
2679    /// assert_eq!(result.value()?.as_slice::<f64>()?, &[6.0]);
2680    /// # Ok::<(), tenferro_ad::Error>(())
2681    /// ```
2682    ///
2683    /// # Errors
2684    ///
2685    /// Returns [`Error::ContextMismatch`] for a foreign eager runtime,
2686    /// a typed validation error for invalid dimensions, or a backend error.
2687    pub fn dot_general_with_conj(
2688        &mut self,
2689        lhs: &EagerTensor,
2690        rhs: &EagerTensor,
2691        config: DotGeneralConfig,
2692        lhs_conj: bool,
2693        rhs_conj: bool,
2694    ) -> Result<EagerTensor> {
2695        self.ensure_runtime(lhs)?;
2696        self.ensure_runtime(rhs)?;
2697        config
2698            .validate_dims_with_ranks(lhs.shape().len(), rhs.shape().len())
2699            .map_err(Error::TensorRuntime)?;
2700        if !lhs.requires_grad && !rhs.requires_grad {
2701            let output = crate::eager_exec::exec_dot_general_with_conj_on_tensor_reads_in_session(
2702                lhs.tensor_read(),
2703                rhs.tensor_read(),
2704                &config,
2705                lhs_conj,
2706                rhs_conj,
2707                self.backend,
2708            )?;
2709            return EagerTensor::new_untracked_result(Arc::clone(self.runtime), output);
2710        }
2711        let lhs = if lhs_conj {
2712            self.conj(lhs)?
2713        } else {
2714            lhs.clone()
2715        };
2716        let rhs = if rhs_conj {
2717            self.conj(rhs)?
2718        } else {
2719            rhs.clone()
2720        };
2721        self.dot_general(&lhs, &rhs, config)
2722    }
2723
2724    fn run_binary(
2725        &mut self,
2726        name: &'static str,
2727        lhs: &EagerTensor,
2728        rhs: &EagerTensor,
2729        op: StdTensorOp,
2730    ) -> Result<EagerTensor> {
2731        self.ensure_runtime(lhs)?;
2732        self.ensure_runtime(rhs)?;
2733        let (lhs, rhs) = crate::eager_ops::broadcast_binary_in_session(name, lhs, rhs, self)?;
2734        EagerTensor::nary_op_in_session(&[&lhs, &rhs], op, self.backend)
2735    }
2736
2737    fn run_ternary(
2738        &mut self,
2739        name: &'static str,
2740        first: &EagerTensor,
2741        second: &EagerTensor,
2742        third: &EagerTensor,
2743        op: StdTensorOp,
2744    ) -> Result<EagerTensor> {
2745        self.ensure_runtime(first)?;
2746        self.ensure_runtime(second)?;
2747        self.ensure_runtime(third)?;
2748        let (first, second, third) =
2749            crate::eager_ops::broadcast_ternary_in_session(name, first, second, third, self)?;
2750        EagerTensor::nary_op_in_session(&[&first, &second, &third], op, self.backend)
2751    }
2752
2753    /// Apply one standard tensor op in this borrowed session and record it
2754    /// for AD when needed.
2755    ///
2756    /// Extension crates use this when an extension-level eager operation
2757    /// expands into ordinary `StdTensorOp` nodes instead of a custom extension
2758    /// primitive: all of them run in this one backend session.
2759    ///
2760    /// # Examples
2761    ///
2762    /// ```rust
2763    /// use tenferro_ad::{EagerRuntime, Tensor};
2764    /// use tenferro_cpu::CpuBackend;
2765    /// use tenferro_ops::std_tensor_op::StdTensorOp;
2766    ///
2767    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2768    /// let y = ctx.with_eager_session(|s| {
2769    ///     let x = s.variable_from(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?)?;
2770    ///     let negated = s.apply_standard_op(StdTensorOp::Neg, &[&x])?;
2771    ///     s.apply_standard_op(StdTensorOp::Mul, &[&negated, &x])
2772    /// })?;
2773    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[-1.0, -4.0]);
2774    /// assert!(y.tracks_grad());
2775    /// # Ok::<(), tenferro_ad::Error>(())
2776    /// ```
2777    ///
2778    /// # Errors
2779    ///
2780    /// Returns [`Error::TensorRuntime`] containing
2781    /// [`tenferro_tensor::ValidationError::InvalidArgument`] for an extension
2782    /// op, [`Error::ContextMismatch`] for a tensor from another runtime, a
2783    /// typed input-count error, or the backend's typed execution error.
2784    pub fn apply_standard_op(
2785        &mut self,
2786        op: StdTensorOp,
2787        inputs: &[&EagerTensor],
2788    ) -> Result<EagerTensor> {
2789        if matches!(op, StdTensorOp::Extension(_)) {
2790            return Err(Error::invalid_argument(
2791                "EagerSession::apply_standard_op",
2792                ErrorPhase::Execution,
2793                "op",
2794                "Extension ops must be passed to apply_eager",
2795            ));
2796        }
2797        for input in inputs {
2798            self.ensure_runtime(input)?;
2799        }
2800        EagerTensor::nary_op_in_session(inputs, op, self.backend)
2801    }
2802
2803    /// Borrow the backend session this eager session runs on.
2804    ///
2805    /// Extension crates use it to run their backend kernels on untracked
2806    /// values inside the same execution region instead of entering a second
2807    /// session, which would be rejected as reentry. It grants the same access
2808    /// as [`EagerRuntime::with_execution_session`].
2809    ///
2810    /// # Examples
2811    ///
2812    /// ```rust
2813    /// use tenferro_ad::{EagerRuntime, Tensor};
2814    /// use tenferro_cpu::CpuBackend;
2815    /// use tenferro_tensor::TensorRead;
2816    ///
2817    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
2818    /// let x = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, -2.0])?;
2819    /// let copy = ctx.with_eager_session(|s| {
2820    ///     s.backend_session()
2821    ///         .to_contiguous_read(TensorRead::from_tensor(&x))
2822    ///         .map_err(tenferro_ad::Error::from)
2823    /// })?;
2824    /// assert_eq!(copy.as_slice::<f64>()?, &[1.0, -2.0]);
2825    /// # Ok::<(), Box<dyn std::error::Error>>(())
2826    /// ```
2827    pub fn backend_session(&mut self) -> &mut dyn BackendSession {
2828        &mut *self.backend
2829    }
2830
2831    fn run_unary(&mut self, input: &EagerTensor, op: StdTensorOp) -> Result<EagerTensor> {
2832        self.ensure_runtime(input)?;
2833        EagerTensor::nary_op_in_session(&[input], op, self.backend)
2834    }
2835
2836    pub(crate) fn record_outputs(
2837        &mut self,
2838        op: &StdTensorOp,
2839        outputs: &[&Tensor],
2840        inputs: &[&EagerTensor],
2841    ) -> Result<RecordedEagerOutputs> {
2842        record_eager_outputs_in_session(op, outputs, inputs, self.backend)
2843    }
2844
2845    pub(crate) fn ensure_runtime(&self, input: &EagerTensor) -> Result<()> {
2846        if !Arc::ptr_eq(self.runtime, &input.ctx) {
2847            return Err(Error::ContextMismatch {
2848                lhs: self.runtime.id(),
2849                rhs: input.ctx_id(),
2850            });
2851        }
2852        Ok(())
2853    }
2854
2855    pub(crate) fn runtime(&self) -> &Arc<EagerRuntime> {
2856        self.runtime
2857    }
2858
2859    /// Run `f` on this runtime's extension cache store from inside the
2860    /// session.
2861    ///
2862    /// Operation families use this for their own prepared-plan caches without
2863    /// reopening the runtime: the eager owner is already locked, and the cache
2864    /// lock is taken second, as in every extension execution region.
2865    ///
2866    /// # Examples
2867    ///
2868    /// ```rust
2869    /// use tenferro_ad::EagerRuntime;
2870    ///
2871    /// let ctx = EagerRuntime::new()?;
2872    /// let entries = ctx.with_eager_session(|session| {
2873    ///     session.with_extension_caches(|caches| caches.len())
2874    /// })?;
2875    /// assert_eq!(entries, 0);
2876    /// # Ok::<(), tenferro_ad::Error>(())
2877    /// ```
2878    ///
2879    /// # Errors
2880    ///
2881    /// Returns a runtime-state error when the extension cache lock is
2882    /// poisoned.
2883    pub fn with_extension_caches<R>(
2884        &mut self,
2885        f: impl FnOnce(&mut tenferro_runtime::ExtensionCacheStore) -> R,
2886    ) -> Result<R> {
2887        let mut caches = self.runtime.lock_extension_caches()?;
2888        Ok(f(&mut caches))
2889    }
2890
2891    pub(crate) fn execute_prepared_extension(
2892        &mut self,
2893        executor: &dyn tenferro_runtime::PreparedOperationExecutor,
2894        inputs: &[TensorRead<'_>],
2895    ) -> Result<Vec<Tensor>> {
2896        // The eager owner is already locked; acquire the extension-cache lock
2897        // second, as in the top-level extension execution region.
2898        let mut caches = self.runtime.lock_extension_caches()?;
2899        executor.execute_in_session(self.backend, &mut caches, inputs)
2900    }
2901}
2902
2903impl fmt::Debug for EagerRuntime {
2904    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
2905        let mut debug = f.debug_struct("EagerRuntime");
2906        debug.field("id", &self.id);
2907        debug.field("runtime_id", &self.runtime.id());
2908        debug.field("runtime_epoch", &self.runtime.epoch().ok());
2909        match self.backend.try_lock() {
2910            Ok(backend) => {
2911                debug.field("backend", &*backend);
2912            }
2913            Err(_) => {
2914                debug.field("backend", &"<locked>");
2915            }
2916        }
2917        match self.extension_caches.try_lock() {
2918            Ok(caches) => {
2919                debug.field(
2920                    "extension_cache_stats",
2921                    &caches.stats(ExtensionCacheSelector::All),
2922                );
2923            }
2924            Err(_) => {
2925                debug.field("extension_cache_stats", &"<locked>");
2926            }
2927        }
2928        match self.extension_install_lock.try_lock() {
2929            Ok(_) => {
2930                debug.field("extension_install_lock", &"<unlocked>");
2931            }
2932            Err(_) => {
2933                debug.field("extension_install_lock", &"<locked>");
2934            }
2935        }
2936        debug.field("semantic_extension_rules", &self.semantic_extension_rules);
2937        match self.grad_slots.try_lock() {
2938            Ok(slots) => {
2939                debug.field("grad_slots_len", &slots.len());
2940            }
2941            Err(_) => {
2942                debug.field("grad_slots_len", &"<locked>");
2943            }
2944        }
2945        match self.value_records.try_lock() {
2946            Ok(records) => {
2947                debug.field("value_records_len", &records.len());
2948            }
2949            Err(_) => {
2950                debug.field("value_records_len", &"<locked>");
2951            }
2952        }
2953        match self.ad_transform_cache.stats() {
2954            Ok(stats) => {
2955                debug.field("ad_transform_cache_stats", &stats);
2956            }
2957            Err(err) => {
2958                debug.field("ad_transform_cache_stats", &format_args!("{err}"));
2959            }
2960        }
2961        match self.prepared_derivative_cache.try_lock() {
2962            Ok(cache) => {
2963                debug.field("prepared_derivative_cache_stats", &cache.stats());
2964            }
2965            Err(_) => {
2966                debug.field("prepared_derivative_cache_stats", &"<locked>");
2967            }
2968        }
2969        debug.finish_non_exhaustive()
2970    }
2971}
2972
2973impl EagerRuntime {
2974    pub(crate) fn lock_backend(&self) -> Result<MutexGuard<'_, EagerBackend>> {
2975        // A thread inside a session must not wait on an owner lock: a callback
2976        // of this runtime holds it (the wait never returns), or another thread
2977        // may hold it while waiting for the permit this thread holds (#1946
2978        // F1). Report the reentry before blocking instead.
2979        if EnteredRuntimeScope::any_entered()
2980            || tenferro_tensor::has_active_backend_session()
2981            || tenferro_cpu::current_cpu_execution() == tenferro_cpu::CpuThreadExecution::Active
2982        {
2983            return Err(tenferro_tensor::SessionEntryError::Reentered {
2984                backend: "EagerRuntime",
2985            }
2986            .into());
2987        }
2988        let poisoned =
2989            || Error::runtime_state("eager_backend", ErrorPhase::Execution, "lock poisoned");
2990        // A shared execution scope already holds the CPU permit, so waiting for
2991        // the owner could deadlock the same way. Take it only when it is free.
2992        if tenferro_cpu::current_cpu_execution() == tenferro_cpu::CpuThreadExecution::SharedScope {
2993            return match self.backend.try_lock() {
2994                Ok(backend) => Ok(backend),
2995                Err(std::sync::TryLockError::Poisoned(_)) => Err(poisoned()),
2996                Err(std::sync::TryLockError::WouldBlock) => {
2997                    Err(tenferro_tensor::SessionEntryError::Contended {
2998                        backend: "EagerRuntime",
2999                        message: "the runtime is in use by another thread while this thread's \
3000                                  CPU execution scope holds the permit; waiting could deadlock"
3001                            .to_owned(),
3002                    }
3003                    .into())
3004                }
3005            };
3006        }
3007        // Independent top-level callers wait for the owner and are served in turn.
3008        self.backend.lock().map_err(|_| poisoned())
3009    }
3010
3011    fn lock_extension_caches(&self) -> Result<MutexGuard<'_, ExtensionCacheStore>> {
3012        self.extension_caches.lock().map_err(|_| {
3013            Error::runtime_state(
3014                "eager_extension_caches",
3015                ErrorPhase::Execution,
3016                "lock poisoned",
3017            )
3018        })
3019    }
3020
3021    fn lock_extension_install(&self) -> Result<MutexGuard<'_, ()>> {
3022        self.extension_install_lock.lock().map_err(|_| {
3023            Error::runtime_state(
3024                "eager_extension_install",
3025                ErrorPhase::Execution,
3026                "lock poisoned",
3027            )
3028        })
3029    }
3030
3031    fn lock_prepared_derivative_cache(&self) -> Result<MutexGuard<'_, PreparedDerivativeCache>> {
3032        self.prepared_derivative_cache.lock().map_err(|_| {
3033            Error::runtime_state(
3034                "prepared_derivative_cache",
3035                ErrorPhase::Execution,
3036                "lock poisoned",
3037            )
3038        })
3039    }
3040
3041    fn lock_grad_slots(
3042        &self,
3043    ) -> Result<MutexGuard<'_, HashMap<ValueKey<StdTensorOp>, WeakGradSlot>>> {
3044        self.grad_slots.lock().map_err(|_| {
3045            Error::runtime_state(
3046                "eager_gradient_slots",
3047                ErrorPhase::Execution,
3048                "lock poisoned",
3049            )
3050        })
3051    }
3052
3053    fn lock_value_records(
3054        &self,
3055    ) -> Result<MutexGuard<'_, HashMap<ValueKey<StdTensorOp>, Weak<EagerTensorRecord>>>> {
3056        self.value_records.lock().map_err(|_| {
3057            Error::runtime_state(
3058                "eager_value_registry",
3059                ErrorPhase::Execution,
3060                "lock poisoned",
3061            )
3062        })
3063    }
3064
3065    fn from_backend(backend: EagerBackend) -> Result<Self> {
3066        Self::from_backend_with_rules_and_cache(
3067            backend,
3068            SemanticExtensionRuleSet::default(),
3069            Arc::new(AdTransformCache::new()),
3070        )
3071    }
3072
3073    fn from_backend_with_rules_and_cache(
3074        backend: EagerBackend,
3075        semantic_extension_rules: SemanticExtensionRuleSet,
3076        ad_transform_cache: Arc<AdTransformCache>,
3077    ) -> Result<Self> {
3078        let runtime = eager_runtime_for_backend(&backend)
3079            .map_err(|source| runtime_config_error("EagerRuntime::from_backend", source))?;
3080        let extension_backend_kind = match &backend {
3081            EagerBackend::Cpu(_) => Some(EagerExtensionBackendKind::Cpu),
3082            #[cfg(test)]
3083            EagerBackend::Recording(_) => None,
3084            #[cfg(feature = "cuda")]
3085            EagerBackend::Cuda(_) => Some(EagerExtensionBackendKind::Cuda),
3086            #[cfg(feature = "webgpu")]
3087            EagerBackend::WebGpu(_) => Some(EagerExtensionBackendKind::WebGpu),
3088        };
3089        Ok(Self {
3090            id: ContextId::fresh(),
3091            runtime,
3092            backend: Mutex::new(backend),
3093            extension_backend_kind,
3094            extension_install_lock: Mutex::new(()),
3095            extension_caches: Mutex::new(ExtensionCacheStore::new()),
3096            semantic_extension_rules,
3097            grad_slots: Mutex::new(HashMap::new()),
3098            value_records: Mutex::new(HashMap::new()),
3099            ad_transform_cache,
3100            prepared_derivative_cache: Mutex::new(PreparedDerivativeCache::default()),
3101        })
3102    }
3103
3104    /// Create a shared CPU eager execution context.
3105    ///
3106    /// # Examples
3107    ///
3108    /// ```
3109    /// use tenferro_ad::EagerRuntime;
3110    ///
3111    /// let ctx = EagerRuntime::new()?;
3112    /// assert_eq!(std::sync::Arc::strong_count(&ctx), 1);
3113    /// # Ok::<(), tenferro_ad::Error>(())
3114    /// ```
3115    ///
3116    /// # Errors
3117    ///
3118    /// Returns [`Error::RuntimeStateSource`] when provider runtime
3119    /// registration cannot be configured, preserving the underlying
3120    /// [`RuntimeConfigError`] as the typed error source.
3121    pub fn new() -> Result<Arc<Self>> {
3122        Self::with_cpu_backend(CpuBackend::new())
3123    }
3124
3125    /// Create a shared eager execution context from a configured CPU backend.
3126    ///
3127    /// # Examples
3128    ///
3129    /// ```
3130    /// use tenferro_cpu::CpuBackend;
3131    /// use tenferro_ad::{EagerRuntime};
3132    ///
3133    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::with_threads(1)?)?;
3134    /// assert_eq!(std::sync::Arc::strong_count(&ctx), 1);
3135    /// # Ok::<(), Box<dyn std::error::Error>>(())
3136    /// ```
3137    ///
3138    /// # Errors
3139    ///
3140    /// Returns [`Error::RuntimeStateSource`] when provider runtime
3141    /// registration cannot be configured, preserving the underlying
3142    /// [`RuntimeConfigError`] as the typed error source.
3143    pub fn with_cpu_backend(backend: CpuBackend) -> Result<Arc<Self>> {
3144        Ok(Arc::new(Self::from_backend(EagerBackend::cpu(backend))?))
3145    }
3146
3147    /// Snapshot a placement-selected CPU handle from this eager runtime.
3148    ///
3149    /// The eager backend lock is held only long enough to verify the backend
3150    /// kind and clone its CPU coordinator/provider snapshot. Placement
3151    /// resolution happens after that guard is dropped. The returned value does
3152    /// not hold a resource permit or a second runtime/backend mutex while idle.
3153    ///
3154    /// # Examples
3155    ///
3156    /// ```rust
3157    /// use tenferro_ad::EagerRuntime;
3158    /// use tenferro_cpu::CpuPlacement;
3159    ///
3160    /// let runtime = EagerRuntime::new()?;
3161    /// let cpu = runtime.on_cpu(CpuPlacement::Auto)?;
3162    /// assert_eq!(cpu.runtime_id(), runtime.id());
3163    /// # Ok::<(), tenferro_ad::Error>(())
3164    /// ```
3165    ///
3166    /// # Errors
3167    ///
3168    /// Returns [`Error::RuntimeState`] if the eager backend lock is poisoned,
3169    /// [`Error::Unsupported`] if the runtime is not CPU-backed, or a typed
3170    /// tensor runtime error retaining [`tenferro_cpu::CpuPlacementError`] when
3171    /// the requested placement cannot be resolved.
3172    pub fn on_cpu(self: &Arc<Self>, placement: CpuPlacement) -> Result<CpuPlacementBoundEager> {
3173        let backend = {
3174            let backend = self.lock_backend()?;
3175            backend.cpu_snapshot().ok_or_else(|| {
3176                Error::unsupported(
3177                    "EagerRuntime::on_cpu",
3178                    ErrorPhase::Execution,
3179                    "the eager runtime is not CPU-backed",
3180                )
3181            })?
3182        };
3183        let selection = select_cpu_runtime(&self.runtime)?;
3184        let backend = backend.for_placement(placement).map_err(|source| {
3185            let error: tenferro_tensor::Error = CpuBackendError::Placement {
3186                op: "EagerRuntime::on_cpu",
3187                source,
3188            }
3189            .into();
3190            Error::from(error)
3191        })?;
3192        Ok(CpuPlacementBoundEager {
3193            runtime: Arc::clone(self),
3194            backend,
3195            snapshot: selection.snapshot,
3196            epoch: selection.epoch,
3197            engine_id: selection.engine_id,
3198            registration_identity: selection.registration_identity,
3199            capabilities: selection.capabilities,
3200        })
3201    }
3202
3203    /// Create a shared CPU eager context with explicit AD extension rules.
3204    ///
3205    /// # Examples
3206    ///
3207    /// ```rust
3208    /// use tenferro_cpu::CpuBackend;
3209    /// use tenferro_ad::{AdContext, EagerRuntime};
3210    ///
3211    /// let ad = AdContext::builder().build().unwrap();
3212    /// let ctx = EagerRuntime::with_cpu_backend_and_ad_context(CpuBackend::new(), &ad)?;
3213    /// assert_eq!(std::sync::Arc::strong_count(&ctx), 1);
3214    /// # Ok::<(), tenferro_ad::Error>(())
3215    /// ```
3216    ///
3217    /// # Errors
3218    ///
3219    /// Returns [`Error::RuntimeStateSource`] when provider runtime
3220    /// registration cannot be configured, preserving the underlying
3221    /// [`RuntimeConfigError`] as the typed error source.
3222    pub fn with_cpu_backend_and_ad_context(
3223        backend: CpuBackend,
3224        ad: &AdContext,
3225    ) -> Result<Arc<Self>> {
3226        Ok(Arc::new(Self::from_backend_with_rules_and_cache(
3227            EagerBackend::cpu(backend),
3228            ad.semantic_extension_rules().clone(),
3229            ad.ad_transform_cache(),
3230        )?))
3231    }
3232
3233    /// Create a shared eager execution context from a configured CUDA backend.
3234    ///
3235    /// # Examples
3236    ///
3237    /// ```
3238    /// use tenferro_gpu::cuda::CudaBackend;
3239    /// use tenferro_ad::EagerRuntime;
3240    ///
3241    /// let _ctor: fn(CudaBackend) -> tenferro_ad::Result<std::sync::Arc<EagerRuntime>> =
3242    ///     EagerRuntime::with_cuda_backend;
3243    /// ```
3244    #[cfg(feature = "cuda")]
3245    ///
3246    /// # Errors
3247    ///
3248    /// Returns [`Error::RuntimeStateSource`] when provider runtime
3249    /// registration cannot be configured, preserving the underlying
3250    /// [`RuntimeConfigError`] as the typed error source.
3251    pub fn with_cuda_backend(backend: CudaBackend) -> Result<Arc<Self>> {
3252        Ok(Arc::new(Self::from_backend(EagerBackend::cuda(backend))?))
3253    }
3254
3255    /// Create a shared CUDA eager context with explicit AD extension rules.
3256    ///
3257    /// # Examples
3258    ///
3259    /// ```rust
3260    /// use tenferro_ad::{AdContext, EagerRuntime};
3261    /// use tenferro_gpu::cuda::CudaBackend;
3262    ///
3263    /// let _ctor: fn(CudaBackend, &AdContext) -> tenferro_ad::Result<std::sync::Arc<EagerRuntime>> =
3264    ///     EagerRuntime::with_cuda_backend_and_ad_context;
3265    /// ```
3266    #[cfg(feature = "cuda")]
3267    ///
3268    /// # Errors
3269    ///
3270    /// Returns [`Error::RuntimeStateSource`] when provider runtime
3271    /// registration cannot be configured, preserving the underlying
3272    /// [`RuntimeConfigError`] as the typed error source.
3273    pub fn with_cuda_backend_and_ad_context(
3274        backend: CudaBackend,
3275        ad: &AdContext,
3276    ) -> Result<Arc<Self>> {
3277        Ok(Arc::new(Self::from_backend_with_rules_and_cache(
3278            EagerBackend::cuda(backend),
3279            ad.semantic_extension_rules().clone(),
3280            ad.ad_transform_cache(),
3281        )?))
3282    }
3283
3284    /// Create a shared eager execution context from a configured WebGPU backend.
3285    ///
3286    /// # Examples
3287    ///
3288    /// ```
3289    /// use tenferro_ad::EagerRuntime;
3290    /// use tenferro_gpu::webgpu::WebGpuBackend;
3291    ///
3292    /// let _ctor: fn(WebGpuBackend) -> tenferro_ad::Result<std::sync::Arc<EagerRuntime>> =
3293    ///     EagerRuntime::with_webgpu_backend;
3294    /// ```
3295    #[cfg(feature = "webgpu")]
3296    ///
3297    /// # Errors
3298    ///
3299    /// Returns [`Error::RuntimeStateSource`] when provider runtime
3300    /// registration cannot be configured, preserving the underlying
3301    /// [`RuntimeConfigError`] as the typed error source.
3302    pub fn with_webgpu_backend(backend: WebGpuBackend) -> Result<Arc<Self>> {
3303        Ok(Arc::new(Self::from_backend(EagerBackend::webgpu(backend))?))
3304    }
3305
3306    /// Create a shared WebGPU eager context with explicit AD extension rules.
3307    ///
3308    /// # Examples
3309    ///
3310    /// ```rust
3311    /// use tenferro_ad::{AdContext, EagerRuntime};
3312    /// use tenferro_gpu::webgpu::WebGpuBackend;
3313    ///
3314    /// let _ctor: fn(WebGpuBackend, &AdContext) -> tenferro_ad::Result<std::sync::Arc<EagerRuntime>> =
3315    ///     EagerRuntime::with_webgpu_backend_and_ad_context;
3316    /// ```
3317    #[cfg(feature = "webgpu")]
3318    ///
3319    /// # Errors
3320    ///
3321    /// Returns [`Error::RuntimeStateSource`] when provider runtime
3322    /// registration cannot be configured, preserving the underlying
3323    /// [`RuntimeConfigError`] as the typed error source.
3324    pub fn with_webgpu_backend_and_ad_context(
3325        backend: WebGpuBackend,
3326        ad: &AdContext,
3327    ) -> Result<Arc<Self>> {
3328        Ok(Arc::new(Self::from_backend_with_rules_and_cache(
3329            EagerBackend::webgpu(backend),
3330            ad.semantic_extension_rules().clone(),
3331            ad.ad_transform_cache(),
3332        )?))
3333    }
3334
3335    /// Return an opaque identifier for this context.
3336    ///
3337    /// # Examples
3338    ///
3339    /// ```
3340    /// use tenferro_cpu::CpuBackend;
3341    /// use tenferro_ad::{EagerRuntime};
3342    ///
3343    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3344    /// assert_ne!(ctx.id(), EagerRuntime::with_cpu_backend(CpuBackend::new())?.id());
3345    /// # Ok::<(), tenferro_ad::Error>(())
3346    /// ```
3347    pub fn id(&self) -> ContextId {
3348        self.id
3349    }
3350
3351    /// Disable eager operation recording on the current thread until the guard is dropped.
3352    ///
3353    /// This is useful for optimizer updates, metric calculations, and other
3354    /// eager computations that should not become part of the AD tape.
3355    ///
3356    /// # Examples
3357    ///
3358    /// ```
3359    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
3360    /// use tenferro_cpu::CpuBackend;
3361    ///
3362    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3363    /// let x = EagerTensor::requires_grad_in(
3364    ///     Tensor::from_vec_col_major(vec![1], vec![3.0_f64]).unwrap(),
3365    ///     ctx.clone(),
3366    /// )?;
3367    /// let y = ctx.with_eager_session(|s| {
3368    ///     let _guard = ctx.no_grad();
3369    ///     s.mul(&x, &x)
3370    /// })?;
3371    /// assert!(!y.tracks_grad());
3372    /// # Ok::<(), tenferro_ad::Error>(())
3373    /// ```
3374    pub fn no_grad(&self) -> EagerNoGradGuard {
3375        EAGER_NO_GRAD_DEPTH.with(|depth| {
3376            depth.set(depth.get().saturating_add(1));
3377        });
3378        EagerNoGradGuard {
3379            active: true,
3380            _not_send: PhantomData,
3381        }
3382    }
3383
3384    /// Keep semantic-trace recording active for untracked intermediates.
3385    ///
3386    /// See [`EagerTraceCaptureGuard`] for the full contract and an example.
3387    pub fn capture_trace(&self) -> EagerTraceCaptureGuard {
3388        EAGER_CAPTURE_DEPTH.with(|depth| {
3389            depth.set(depth.get().saturating_add(1));
3390        });
3391        EagerTraceCaptureGuard {
3392            active: true,
3393            _not_send: PhantomData,
3394        }
3395    }
3396
3397    /// Install or replace one extension module on this eager context's runtime.
3398    ///
3399    /// Eager extension wrappers call this as an idempotent "ensure installed"
3400    /// step. When the exact module instance (same module ID and allocation) is
3401    /// already installed, this is a read-only no-op that returns the current
3402    /// runtime epoch without acquiring the install lock or reconfiguring. The
3403    /// cold or replacement paths keep the transactional install-or-replace
3404    /// behavior, serialized so parallel first-use of the same extension family
3405    /// cannot publish over another thread's base snapshot.
3406    ///
3407    /// # Errors
3408    ///
3409    /// Returns [`tenferro_runtime::Error::RuntimeState`] when runtime
3410    /// reconfiguration fails or the extension module transaction is invalid.
3411    pub fn install_extension_module(
3412        &self,
3413        module: Arc<dyn ExtensionModule>,
3414    ) -> Result<RuntimeEpoch> {
3415        let snapshot = self.runtime.snapshot().map_err(|source| {
3416            runtime_state_source("EagerRuntime::install_extension_module", source)
3417        })?;
3418        if snapshot.has_extension_module_identical(&module) {
3419            return Ok(snapshot.epoch());
3420        }
3421        let _install_guard = self.lock_extension_install()?;
3422        self.runtime
3423            .reconfigure(|edit| {
3424                edit.replace_extension_module(module)?;
3425                Ok(())
3426            })
3427            .map_err(|source| {
3428                runtime_state_source("EagerRuntime::install_extension_module", source)
3429            })
3430    }
3431
3432    pub(crate) fn ensure_extension_module_for_engine(
3433        &self,
3434        module: Arc<dyn ExtensionModule>,
3435        family_id: &'static str,
3436        engine_id: &EngineId,
3437    ) -> Result<RuntimeEpoch> {
3438        let snapshot = self.runtime.snapshot().map_err(|source| {
3439            runtime_state_source("EagerRuntime::ensure_extension_module_for_engine", source)
3440        })?;
3441        if snapshot.has_extension_module_engine(module.module_id(), family_id, engine_id) {
3442            return Ok(snapshot.epoch());
3443        }
3444        let _install_guard = self.lock_extension_install()?;
3445        self.runtime
3446            .reconfigure(|edit| {
3447                edit.ensure_extension_module_for_engine(module, family_id, engine_id)?;
3448                Ok(())
3449            })
3450            .map_err(|source| {
3451                runtime_state_source("EagerRuntime::ensure_extension_module_for_engine", source)
3452            })
3453    }
3454
3455    pub(crate) fn runtime(&self) -> &Runtime {
3456        &self.runtime
3457    }
3458
3459    pub(crate) fn eager_extension_target(&self) -> Result<EagerExtensionTarget> {
3460        let backend_kind = self.extension_backend_kind.ok_or_else(|| {
3461            Error::unsupported(
3462                "EagerRuntime::eager_extension_target",
3463                ErrorPhase::Execution,
3464                "the recording backend has no registered eager extension engine",
3465            )
3466        })?;
3467        let engine_id = match backend_kind {
3468            EagerExtensionBackendKind::Cpu => cpu_runtime_engine_id(),
3469            #[cfg(feature = "cuda")]
3470            EagerExtensionBackendKind::Cuda => cuda_runtime_engine_id(),
3471            #[cfg(feature = "webgpu")]
3472            EagerExtensionBackendKind::WebGpu => tenferro_gpu::webgpu::webgpu_runtime_engine_id(),
3473        }
3474        .map_err(|source| runtime_config_error("EagerRuntime::eager_extension_target", source))?;
3475        let target = EagerExtensionTarget {
3476            engine_id,
3477            backend_kind,
3478        };
3479        validate_eager_extension_target(&self.runtime, &target)?;
3480        Ok(target)
3481    }
3482
3483    /// Clear generic extension runtime cache entries.
3484    ///
3485    /// # Examples
3486    ///
3487    /// ```
3488    /// use tenferro_cpu::CpuBackend;
3489    /// use tenferro_ad::{EagerRuntime};
3490    ///
3491    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3492    /// ctx.clear_extension_caches()?;
3493    /// assert_eq!(ctx.cache_stats()?.extensions.entries, 0);
3494    /// # Ok::<(), tenferro_ad::Error>(())
3495    /// ```
3496    ///
3497    /// # Errors
3498    ///
3499    /// Returns [`tenferro_runtime::Error::RuntimeState`] when the extension
3500    /// cache lock is poisoned.
3501    pub fn clear_extension_caches(&self) -> Result<()> {
3502        self.lock_extension_caches()?.clear();
3503        Ok(())
3504    }
3505
3506    /// Clear every cache owned by this eager context.
3507    ///
3508    /// # Examples
3509    ///
3510    /// ```
3511    /// use tenferro_cpu::CpuBackend;
3512    /// use tenferro_ad::{EagerRuntime};
3513    ///
3514    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3515    /// ctx.clear_caches()?;
3516    /// assert_eq!(ctx.cache_stats()?.extensions.entries, 0);
3517    /// assert_eq!(ctx.cache_stats()?.ad_transforms.entries, 0);
3518    /// assert_eq!(ctx.cache_stats()?.prepared_derivatives.entries, 0);
3519    /// # Ok::<(), tenferro_ad::Error>(())
3520    /// ```
3521    ///
3522    /// # Errors
3523    ///
3524    /// Returns [`tenferro_runtime::Error::RuntimeState`] when either the
3525    /// extension cache or AD-transform cache is poisoned.
3526    pub fn clear_caches(&self) -> Result<()> {
3527        self.clear_extension_caches()?;
3528        self.clear_ad_transform_caches()?;
3529        self.clear_prepared_derivative_cache()?;
3530        Ok(())
3531    }
3532
3533    /// Clear prepared derivative program cache entries.
3534    ///
3535    /// # Examples
3536    ///
3537    /// ```rust
3538    /// use tenferro_ad::EagerRuntime;
3539    /// use tenferro_cpu::CpuBackend;
3540    ///
3541    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3542    /// ctx.clear_prepared_derivative_cache()?;
3543    /// assert_eq!(ctx.cache_stats()?.prepared_derivatives.entries, 0);
3544    /// # Ok::<(), tenferro_ad::Error>(())
3545    /// ```
3546    ///
3547    /// # Errors
3548    ///
3549    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the prepared
3550    /// derivative cache lock is poisoned.
3551    pub fn clear_prepared_derivative_cache(&self) -> Result<()> {
3552        self.lock_prepared_derivative_cache()?.clear();
3553        Ok(())
3554    }
3555
3556    /// Return eager runtime cache-entry and retained-byte stats.
3557    ///
3558    /// # Examples
3559    ///
3560    /// ```
3561    /// use tenferro_cpu::CpuBackend;
3562    /// use tenferro_ad::{EagerRuntime};
3563    ///
3564    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3565    /// let stats = ctx.cache_stats()?;
3566    /// assert_eq!(stats.extensions.entries, 0);
3567    /// assert_eq!(stats.ad_transforms.entries, 0);
3568    /// assert_eq!(stats.prepared_derivatives.entries, 0);
3569    /// # Ok::<(), tenferro_ad::Error>(())
3570    /// ```
3571    ///
3572    /// # Errors
3573    ///
3574    /// Returns [`tenferro_runtime::Error::RuntimeState`] when a cache or
3575    /// AD-transform cache lock is poisoned.
3576    pub fn cache_stats(&self) -> Result<EagerRuntimeCacheStats> {
3577        Ok(EagerRuntimeCacheStats {
3578            extensions: self
3579                .lock_extension_caches()?
3580                .stats(ExtensionCacheSelector::All),
3581            ad_transforms: self.ad_transform_cache.stats()?,
3582            prepared_derivatives: self.lock_prepared_derivative_cache()?.stats(),
3583        })
3584    }
3585
3586    /// Return the AD transform cache retention limits.
3587    ///
3588    /// # Examples
3589    ///
3590    /// ```
3591    /// use tenferro_ad::EagerRuntime;
3592    /// use tenferro_cpu::CpuBackend;
3593    ///
3594    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3595    /// assert!(ctx.ad_transform_cache_limits()?.max_entries().get() > 0);
3596    /// # Ok::<(), tenferro_ad::Error>(())
3597    /// ```
3598    ///
3599    /// # Errors
3600    ///
3601    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the AD-transform
3602    /// cache lock is poisoned.
3603    pub fn ad_transform_cache_limits(&self) -> Result<AdTransformCacheLimits> {
3604        self.ad_transform_cache.limits()
3605    }
3606
3607    /// Replace AD transform cache retention limits.
3608    ///
3609    /// # Examples
3610    ///
3611    /// ```
3612    /// use std::num::NonZeroUsize;
3613    /// use tenferro_ad::{AdTransformCacheLimits, EagerRuntime};
3614    /// use tenferro_cpu::CpuBackend;
3615    ///
3616    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3617    /// let limits = AdTransformCacheLimits::new(NonZeroUsize::new(1).unwrap());
3618    /// ctx.set_ad_transform_cache_limits(limits)?;
3619    /// assert_eq!(ctx.ad_transform_cache_limits()?, limits);
3620    /// # Ok::<(), tenferro_ad::Error>(())
3621    /// ```
3622    ///
3623    /// # Errors
3624    ///
3625    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the AD-transform
3626    /// cache lock is poisoned while updating limits.
3627    pub fn set_ad_transform_cache_limits(&self, limits: AdTransformCacheLimits) -> Result<()> {
3628        self.ad_transform_cache.set_limits(limits)
3629    }
3630
3631    /// Clear AD transform cache entries visible through this eager runtime.
3632    ///
3633    /// # Examples
3634    ///
3635    /// ```
3636    /// use tenferro_ad::EagerRuntime;
3637    /// use tenferro_cpu::CpuBackend;
3638    ///
3639    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3640    /// ctx.clear_ad_transform_caches()?;
3641    /// assert_eq!(ctx.cache_stats()?.ad_transforms.entries, 0);
3642    /// # Ok::<(), tenferro_ad::Error>(())
3643    /// ```
3644    ///
3645    /// # Errors
3646    ///
3647    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the AD-transform
3648    /// cache lock is poisoned while clearing entries.
3649    pub fn clear_ad_transform_caches(&self) -> Result<()> {
3650        self.ad_transform_cache.clear()
3651    }
3652
3653    /// Return prepared derivative cache retention limits.
3654    ///
3655    /// # Examples
3656    ///
3657    /// ```rust
3658    /// use tenferro_ad::EagerRuntime;
3659    /// use tenferro_cpu::CpuBackend;
3660    ///
3661    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3662    /// assert!(ctx.prepared_derivative_cache_limits()?.max_entries().get() > 0);
3663    /// # Ok::<(), tenferro_ad::Error>(())
3664    /// ```
3665    ///
3666    /// # Errors
3667    ///
3668    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the prepared
3669    /// derivative cache lock is poisoned.
3670    pub fn prepared_derivative_cache_limits(&self) -> Result<AdTransformCacheLimits> {
3671        Ok(self.lock_prepared_derivative_cache()?.limits())
3672    }
3673
3674    /// Replace prepared derivative cache retention limits.
3675    ///
3676    /// # Examples
3677    ///
3678    /// ```rust
3679    /// use std::num::NonZeroUsize;
3680    /// use tenferro_ad::{AdTransformCacheLimits, EagerRuntime};
3681    /// use tenferro_cpu::CpuBackend;
3682    ///
3683    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3684    /// let limits = AdTransformCacheLimits::new(NonZeroUsize::new(1).unwrap());
3685    /// ctx.set_prepared_derivative_cache_limits(limits)?;
3686    /// assert_eq!(ctx.prepared_derivative_cache_limits()?, limits);
3687    /// # Ok::<(), tenferro_ad::Error>(())
3688    /// ```
3689    ///
3690    /// # Errors
3691    ///
3692    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the prepared
3693    /// derivative cache lock is poisoned.
3694    pub fn set_prepared_derivative_cache_limits(
3695        &self,
3696        limits: AdTransformCacheLimits,
3697    ) -> Result<()> {
3698        self.lock_prepared_derivative_cache()?.set_limits(limits);
3699        Ok(())
3700    }
3701
3702    /// Return the extension cache retention limits.
3703    ///
3704    /// # Errors
3705    ///
3706    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the extension
3707    /// cache lock is poisoned.
3708    pub fn extension_cache_limits(&self) -> Result<ExtensionCacheLimits> {
3709        Ok(self.lock_extension_caches()?.limits())
3710    }
3711
3712    /// Replace extension cache retention limits.
3713    ///
3714    /// # Errors
3715    ///
3716    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the extension
3717    /// cache lock is poisoned.
3718    pub fn set_extension_cache_limits(&self, limits: ExtensionCacheLimits) -> Result<()> {
3719        self.lock_extension_caches()?.set_limits(limits);
3720        Ok(())
3721    }
3722
3723    /// Enter one backend execution session and run provider-neutral operations.
3724    ///
3725    /// The callback receives only a lifetime-bound, non-owning backend session.
3726    /// The backend and its engine registration are fixed when the eager runtime
3727    /// is constructed. Extension modules are installed separately and remain
3728    /// available to later extension operations.
3729    ///
3730    /// # Examples
3731    ///
3732    /// ```
3733    /// use tenferro_ad::EagerRuntime;
3734    /// use tenferro_cpu::CpuBackend;
3735    /// use tenferro_tensor::{Tensor, TensorRead};
3736    ///
3737    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3738    /// let lhs = Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?;
3739    /// let rhs = Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?;
3740    /// let output = ctx.with_execution_session(|session| {
3741    ///     session.add_read(TensorRead::from_tensor(&lhs), TensorRead::from_tensor(&rhs))
3742    /// })??;
3743    /// assert_eq!(output.as_slice::<f64>()?, &[3.0]);
3744    /// # Ok::<(), tenferro_ad::Error>(())
3745    /// ```
3746    ///
3747    /// # Errors
3748    ///
3749    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the eager backend
3750    /// lock is poisoned, and [`tenferro_runtime::Error::SessionEntry`] without
3751    /// running the callback when the backend cannot admit the session.
3752    /// Backend operations retain their typed tensor/backend errors inside the
3753    /// callback result.
3754    pub fn with_execution_session<R: Send>(
3755        &self,
3756        f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
3757    ) -> Result<R> {
3758        // Lock order: the eager backend owner lock is taken before admission,
3759        // and admission never waits on this lock while holding a permit.
3760        let mut backend = self.lock_backend()?;
3761        let modes = InheritedEagerModes::capture();
3762        let id = self.id;
3763        Ok(backend.with_backend_session(move |session| {
3764            let _modes = modes.enter();
3765            let _entered = EnteredRuntimeScope::enter(id);
3766            f(session)
3767        })?)
3768    }
3769
3770    /// Enter a runtime-bound eager session for one or more eager operations.
3771    ///
3772    /// Only tensors owned by this runtime may execute on the borrowed session;
3773    /// the backend lock and the CPU execution permit remain live for the callback.
3774    /// The CPU backend may run the callback on a worker thread; the calling
3775    /// thread's [`Self::no_grad`] and [`Self::capture_trace`] guards still govern
3776    /// it, because the callback inherits their state for its duration. Guards
3777    /// started inside the callback end with their own scope and never reach the
3778    /// calling thread.
3779    ///
3780    /// # Examples
3781    ///
3782    /// ```rust
3783    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
3784    /// use tenferro_cpu::CpuBackend;
3785    ///
3786    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3787    /// let x = EagerTensor::from_tensor_in(
3788    ///     Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?, ctx.clone(),
3789    /// )?;
3790    /// let y = ctx.with_eager_session(|session| session.neg(&x))?;
3791    /// assert_eq!(y.value()?.as_slice::<f64>()?, &[-2.0]);
3792    /// # Ok::<(), tenferro_ad::Error>(())
3793    /// ```
3794    ///
3795    /// The callback's error type only needs `From<tenferro_ad::Error>`, so a
3796    /// downstream error type flows through with a single `?`:
3797    ///
3798    /// ```rust
3799    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
3800    /// use tenferro_cpu::CpuBackend;
3801    ///
3802    /// #[derive(Debug)]
3803    /// enum AppError {
3804    ///     Tenferro(tenferro_ad::Error),
3805    ///     Negative,
3806    /// }
3807    /// impl From<tenferro_ad::Error> for AppError {
3808    ///     fn from(error: tenferro_ad::Error) -> Self {
3809    ///         AppError::Tenferro(error)
3810    ///     }
3811    /// }
3812    ///
3813    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new()).unwrap();
3814    /// let x = EagerTensor::from_tensor_in(
3815    ///     Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(),
3816    ///     ctx.clone(),
3817    /// )
3818    /// .unwrap();
3819    /// let squared = ctx.with_eager_session(|session| -> Result<EagerTensor, AppError> {
3820    ///     let y = session.mul(&x, &x)?;
3821    ///     let value = y.value()?;
3822    ///     if value.as_slice::<f64>().map_err(tenferro_ad::Error::from)?[0] < 0.0 {
3823    ///         return Err(AppError::Negative);
3824    ///     }
3825    ///     Ok(y)
3826    /// });
3827    /// assert_eq!(squared.unwrap().value().unwrap().as_slice::<f64>().unwrap(), &[4.0]);
3828    /// ```
3829    ///
3830    /// # Errors
3831    ///
3832    /// Returns the callback's error unchanged. Before the callback runs,
3833    /// returns `E::from(`[`Error::RuntimeState`]`)` if the backend lock is
3834    /// poisoned, or `E::from(`[`tenferro_runtime::Error::SessionEntry`]`)`
3835    /// when the backend cannot admit the session (for example same-thread
3836    /// reentry). `T` and `E` must be `Send` while the CPU session may run the
3837    /// callback on a pool thread.
3838    pub fn with_eager_session<T: Send, E: From<Error> + Send>(
3839        self: &Arc<Self>,
3840        f: impl FnOnce(&mut EagerSession<'_>) -> std::result::Result<T, E> + Send,
3841    ) -> std::result::Result<T, E> {
3842        match self.with_execution_session(|backend| {
3843            f(&mut EagerSession {
3844                runtime: self,
3845                backend,
3846            })
3847        }) {
3848            Ok(result) => result,
3849            Err(entry) => Err(E::from(entry)),
3850        }
3851    }
3852
3853    /// Materialize a host-placement read without entering a backend session.
3854    ///
3855    /// Returns `None` when this runtime's backend has no session-free host
3856    /// materialization path, in which case the caller must enter a session.
3857    ///
3858    /// # Errors
3859    ///
3860    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the eager backend
3861    /// lock is poisoned, or the backend's typed materialization error.
3862    pub(crate) fn to_contiguous_host_read(&self, input: &TensorRead<'_>) -> Result<Option<Tensor>> {
3863        let backend = self.lock_backend()?;
3864        match backend.to_contiguous_host_read(input) {
3865            Some(materialized) => Ok(Some(materialized.map_err(Error::from)?)),
3866            None => Ok(None),
3867        }
3868    }
3869
3870    // Lock ordering: the eager backend owner is locked first; the
3871    // extension-cache lock is acquired only after it and remains held through
3872    // the borrowed session callback.
3873    /// Run an extension-owned eager operation with a borrowed backend session
3874    /// and the eager runtime's extension cache store.
3875    ///
3876    /// The eager backend owner is locked before the extension-cache lock is
3877    /// acquired. The callback receives an
3878    /// [`tenferro_runtime::ExtensionExecutionContext`] so cache access and
3879    /// backend execution share one lifetime-bound context without exposing the
3880    /// owning eager backend. The backend and its engine registration remain
3881    /// fixed for the eager runtime's lifetime.
3882    ///
3883    /// # Examples
3884    ///
3885    /// ```
3886    /// use tenferro_ad::EagerRuntime;
3887    /// use tenferro_cpu::CpuBackend;
3888    /// use tenferro_tensor::{Tensor, TensorRead};
3889    ///
3890    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3891    /// let lhs = Tensor::from_vec_col_major(vec![1], vec![1.0_f64])?;
3892    /// let rhs = Tensor::from_vec_col_major(vec![1], vec![2.0_f64])?;
3893    /// let output = ctx.with_extension_execution_context(|extension_ctx| {
3894    ///     extension_ctx
3895    ///         .backend_mut()
3896    ///         .add_read(TensorRead::from_tensor(&lhs), TensorRead::from_tensor(&rhs))
3897    /// })??;
3898    /// assert_eq!(output.as_slice::<f64>()?, &[3.0]);
3899    /// # Ok::<(), tenferro_ad::Error>(())
3900    /// ```
3901    ///
3902    /// # Errors
3903    ///
3904    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the eager backend
3905    /// or extension-cache lock is poisoned. Errors returned by the callback
3906    /// remain in its result value.
3907    pub fn with_extension_execution_context<R: Send>(
3908        &self,
3909        f: impl FnOnce(
3910                &mut tenferro_runtime::ExtensionExecutionContext<'_, dyn BackendSession + '_>,
3911            ) -> R
3912            + Send,
3913    ) -> Result<R> {
3914        let mut backend = self.lock_backend()?;
3915        let mut extension_cache_guard = self.lock_extension_caches()?;
3916        let extension_caches: &mut ExtensionCacheStore = &mut extension_cache_guard;
3917        let modes = InheritedEagerModes::capture();
3918        let id = self.id;
3919        Ok(backend.with_backend_session(move |session| {
3920            let _modes = modes.enter();
3921            let _entered = EnteredRuntimeScope::enter(id);
3922            let mut extension_ctx =
3923                tenferro_runtime::ExtensionExecutionContext::new(session, extension_caches);
3924            f(&mut extension_ctx)
3925        })?)
3926    }
3927
3928    /// Run a prepared extension executor through the runtime-owned erased
3929    /// backend context (the native-context path).
3930    ///
3931    /// This is the sibling of [`Self::with_extension_execution_context`] for
3932    /// prepared operations whose executor does not support the scheduler-owned
3933    /// session but implements the mandatory `execute` bridge. The concrete
3934    /// backend is exposed as an erased context whose type identity matches the
3935    /// executor's binding.
3936    pub(crate) fn with_extension_erased_context<R: Send>(
3937        &self,
3938        f: impl FnOnce(&mut tenferro_runtime::ErasedExecutionContext<'_>, &mut ExtensionCacheStore) -> R
3939            + Send,
3940    ) -> Result<R> {
3941        let mut backend = self.lock_backend()?;
3942        let mut extension_cache_guard = self.lock_extension_caches()?;
3943        let extension_caches: &mut ExtensionCacheStore = &mut extension_cache_guard;
3944        let mut erased = backend.erased_context();
3945        Ok(f(&mut erased, extension_caches))
3946    }
3947
3948    /// Block the current thread until backend work submitted by this eager runtime completes.
3949    ///
3950    /// CPU runtimes return immediately. CUDA and WebGPU runtimes synchronize
3951    /// their current backend work queue.
3952    ///
3953    /// # Examples
3954    ///
3955    /// ```
3956    /// use tenferro_cpu::CpuBackend;
3957    /// use tenferro_ad::EagerRuntime;
3958    ///
3959    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
3960    /// ctx.synchronize().unwrap();
3961    /// # Ok::<(), tenferro_ad::Error>(())
3962    /// ```
3963    ///
3964    /// # Errors
3965    ///
3966    /// Returns [`tenferro_runtime::Error::RuntimeState`] if the backend lock is
3967    /// poisoned, or a typed tensor backend error if synchronization fails.
3968    pub fn synchronize(&self) -> Result<()> {
3969        self.lock_backend()?.synchronize().map_err(Error::from)
3970    }
3971
3972    /// Owner-context extension fallback used by `extension::apply_eager` when
3973    /// the extension has no prepared session executor.
3974    pub(crate) fn exec_extension_outputs_read(
3975        &self,
3976        op: &Arc<dyn tenferro_ops::ext_op::ExtensionOp>,
3977        inputs: &[TensorRead<'_>],
3978    ) -> Result<Vec<Tensor>> {
3979        // Lock ordering: the backend lock is held for the input session; the
3980        // runtime's extension cache locks are acquired only after it.
3981        let mut backend =
3982            profile_eager_op_section("exec_extension_outputs_read.lock_backend", || {
3983                self.lock_backend()
3984            })?;
3985        profile_eager_op_section("exec_extension_outputs_read.exec_op", || {
3986            exec_extension_op_on_tensor_reads(op, inputs, &mut *backend, &self.runtime)
3987        })
3988    }
3989
3990    #[cfg(test)]
3991    pub(crate) fn exec_standard_graph_outputs(
3992        &self,
3993        graph: &Graph<StdTensorOp>,
3994        initial_data: HashMap<ValueKey<StdTensorOp>, Tensor>,
3995    ) -> Result<EagerGraphExecution> {
3996        let mut backend =
3997            profile_eager_op_section("exec_graph.lock_backend", || self.lock_backend())?;
3998        let mut all_values = initial_data;
3999
4000        profile_eager_op_section("exec_graph.with_backend_session", || {
4001            backend.with_backend_session(|exec| -> Result<()> {
4002                for op_node in graph.operations() {
4003                    let outputs = {
4004                        let input_values = op_node
4005                            .inputs
4006                            .iter()
4007                            .map(|input| {
4008                                let key = match input {
4009                                    ValueRef::Local(local_id) => &graph.values()[*local_id].key,
4010                                    ValueRef::External(key) => key,
4011                                };
4012                                all_values.get(key).ok_or_else(|| {
4013                                    Error::Internal(format!(
4014                                        "standard graph eager execution missing value for {key:?}"
4015                                    ))
4016                                })
4017                            })
4018                            .collect::<Result<Vec<_>>>()?;
4019                        let input_reads = input_values
4020                            .iter()
4021                            .map(|value| TensorRead::from_tensor(value))
4022                            .collect::<Vec<_>>();
4023                        exec_standard_op_on_tensor_reads_in_session(
4024                            &op_node.operation,
4025                            &input_reads,
4026                            exec,
4027                        )?
4028                    };
4029
4030                    if outputs.len() != op_node.outputs.len() {
4031                        return Err(Error::Internal(format!(
4032                            "standard graph eager execution expected {} outputs for {:?}, got {}",
4033                            op_node.outputs.len(),
4034                            op_node.operation,
4035                            outputs.len()
4036                        )));
4037                    }
4038
4039                    for (output_id, output) in op_node.outputs.iter().zip(outputs) {
4040                        let key = graph.values()[*output_id].key.clone();
4041                        all_values.insert(key, output);
4042                    }
4043                }
4044                Ok(())
4045            })?
4046        })?;
4047
4048        let outputs = graph
4049            .outputs()
4050            .iter()
4051            .map(|&output_id| {
4052                let key = &graph.values()[output_id].key;
4053                all_values
4054                    .get(key)
4055                    .ok_or_else(|| {
4056                        Error::Internal(format!(
4057                            "standard graph eager execution missing graph output {key:?}"
4058                        ))
4059                    })?
4060                    .duplicate()
4061                    .map_err(Error::from)
4062            })
4063            .collect::<Result<Vec<_>>>()?;
4064
4065        Ok(EagerGraphExecution { outputs })
4066    }
4067
4068    pub(crate) fn try_register_grad_slot(
4069        &self,
4070        key: &ValueKey<StdTensorOp>,
4071        slot: &GradSlot,
4072    ) -> Result<()> {
4073        insert_pruning_dead(
4074            &mut *self.lock_grad_slots()?,
4075            key.clone(),
4076            Arc::downgrade(slot),
4077        );
4078        Ok(())
4079    }
4080
4081    pub(crate) fn try_register_value_record(
4082        &self,
4083        key: &ValueKey<StdTensorOp>,
4084        record: &Arc<EagerTensorRecord>,
4085    ) -> Result<()> {
4086        insert_pruning_dead(
4087            &mut *self.lock_value_records()?,
4088            key.clone(),
4089            Arc::downgrade(record),
4090        );
4091        Ok(())
4092    }
4093
4094    pub(crate) fn value_record(
4095        &self,
4096        key: &ValueKey<StdTensorOp>,
4097    ) -> Result<Option<Arc<EagerTensorRecord>>> {
4098        let mut records = self.lock_value_records()?;
4099        let Some(record) = records.get(key).cloned() else {
4100            return Ok(None);
4101        };
4102        match record.upgrade() {
4103            Some(record) => Ok(Some(record)),
4104            None => {
4105                records.remove(key);
4106                Ok(None)
4107            }
4108        }
4109    }
4110
4111    /// Clear all live gradient slots tracked by this context.
4112    ///
4113    /// This resets the stored gradients to `None` without unregistering the
4114    /// tensors, so future `backward()` calls can accumulate again.
4115    ///
4116    /// # Examples
4117    ///
4118    /// ```
4119    /// use tenferro_cpu::CpuBackend;
4120    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4121    ///
4122    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4123    /// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0]).unwrap(), ctx.clone()).unwrap();
4124    /// let y = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![3], vec![4.0_f64, 5.0, 6.0]).unwrap(), ctx.clone()).unwrap();
4125    /// let loss = ctx.with_eager_session(|s| {
4126    ///     let product = s.mul(&x, &y)?;
4127    ///     s.reduce_sum(&product, Some(&[0]))
4128    /// })?;
4129    /// let _ = loss.backward().unwrap();
4130    ///
4131    /// ctx.clear_grads()?;
4132    ///
4133    /// assert!(x.grad()?.is_none());
4134    /// assert!(y.grad()?.is_none());
4135    /// # Ok::<(), tenferro_ad::Error>(())
4136    /// ```
4137    ///
4138    /// # Errors
4139    ///
4140    /// Returns [`tenferro_runtime::Error::RuntimeState`] if a gradient-slot
4141    /// lock is poisoned while clearing live gradients.
4142    pub fn clear_grads(&self) -> Result<()> {
4143        let live_slots = {
4144            let mut live_slots = Vec::new();
4145            self.lock_grad_slots()?.retain(|_, slot| {
4146                if let Some(slot) = slot.upgrade() {
4147                    live_slots.push(slot);
4148                    true
4149                } else {
4150                    false
4151                }
4152            });
4153            live_slots
4154        };
4155
4156        let mut poisoned_slot = false;
4157        for slot in live_slots {
4158            match slot.lock() {
4159                Ok(mut current) => {
4160                    *current = None;
4161                }
4162                Err(_) => {
4163                    poisoned_slot = true;
4164                }
4165            }
4166        }
4167        if poisoned_slot {
4168            return Err(Error::runtime_state(
4169                "eager_gradient_slot",
4170                ErrorPhase::Execution,
4171                "lock poisoned",
4172            ));
4173        }
4174        Ok(())
4175    }
4176
4177    /// Import a concrete tensor into this context as an untracked constant.
4178    ///
4179    /// The returned tensor does not participate in gradient tracking.
4180    /// Use this for fixed masks, quadrature weights, physical constants,
4181    /// and other data that should not receive gradients. Like
4182    /// [`EagerSession::constant_from`], it performs no host/device transfer;
4183    /// use [`EagerSession::constant_from_host`] to upload host data into a
4184    /// device runtime.
4185    ///
4186    /// # Examples
4187    ///
4188    /// ```
4189    /// use tenferro_cpu::CpuBackend;
4190    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4191    ///
4192    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4193    /// let c = ctx.constant_from(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap())?;
4194    /// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0]).unwrap(), ctx.clone())?;
4195    /// let z = ctx.with_eager_session(|s| s.add(&x, &c))?;
4196    ///
4197    /// assert_eq!(z.value()?.as_slice::<f64>().unwrap(), &[4.0, 6.0]);
4198    /// # Ok::<(), tenferro_ad::Error>(())
4199    /// ```
4200    ///
4201    /// # Errors
4202    ///
4203    /// Returns [`tenferro_runtime::Error::RuntimeState`] when metadata cannot
4204    /// be registered or the backend lock is poisoned.
4205    pub fn constant_from(self: &Arc<Self>, tensor: Tensor) -> Result<EagerTensor> {
4206        EagerTensor::new_leaf(Arc::clone(self), tensor, false)
4207    }
4208
4209    /// Import a concrete tensor into this context as a trainable variable.
4210    ///
4211    /// The returned tensor participates in gradient tracking; its gradient
4212    /// slot is registered in this context.
4213    ///
4214    /// # Examples
4215    ///
4216    /// ```
4217    /// use tenferro_cpu::CpuBackend;
4218    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4219    ///
4220    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4221    /// let p = ctx.variable_from(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap())?;
4222    /// let loss = ctx.with_eager_session(|s| {
4223    ///     let y = s.exp(&p)?;
4224    ///     s.reduce_sum(&y, Some(&[0]))
4225    /// })?;
4226    /// let _ = loss.backward().unwrap();
4227    ///
4228    /// let grad = p.grad().unwrap().unwrap();
4229    /// assert_eq!(grad.shape(), &[2]);
4230    /// # Ok::<(), tenferro_ad::Error>(())
4231    /// ```
4232    ///
4233    /// # Errors
4234    ///
4235    /// Returns [`tenferro_runtime::Error::RuntimeState`] when gradient metadata
4236    /// or the eager backend state cannot be registered.
4237    pub fn variable_from(self: &Arc<Self>, tensor: Tensor) -> Result<EagerTensor> {
4238        EagerTensor::new_leaf(Arc::clone(self), tensor, true)
4239    }
4240
4241    /// Gradient of a scalar eager output with respect to an eager tensor.
4242    ///
4243    /// Functional eager gradients return ordinary eager tensors and do not
4244    /// write into `grad()` slots. The returned tensor keeps a trace when the
4245    /// derivative computation depends on tracked eager values.
4246    ///
4247    /// # Examples
4248    ///
4249    /// ```
4250    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4251    /// use tenferro_cpu::CpuBackend;
4252    ///
4253    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4254    /// let x = EagerTensor::requires_grad_in(
4255    ///     Tensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap(),
4256    ///     ctx.clone(),
4257    /// )?;
4258    /// let loss = ctx.with_eager_session(|s| s.mul(&x, &x))?;
4259    /// let dx = ctx.grad(&loss, &x)?;
4260    /// assert_eq!(dx.value()?.as_slice::<f64>().unwrap(), &[6.0]);
4261    /// # Ok::<(), tenferro_ad::Error>(())
4262    /// ```
4263    ///
4264    /// # Errors
4265    ///
4266    /// Returns [`tenferro_runtime::Error::NonScalarGrad`] for a non-scalar
4267    /// output, [`Error::ContextMismatch`] for tensors from another runtime,
4268    /// [`Error::UnsupportedAdRule`] when an AD rule is unavailable, or a typed
4269    /// validation/backend error from eager execution. An inactive `wrt` returns
4270    /// [`Error::Validation`] with `argument: "wrt"`; use
4271    /// [`grad_optional`](Self::grad_optional) to observe that state.
4272    pub fn grad(self: &Arc<Self>, output: &EagerTensor, wrt: &EagerTensor) -> Result<EagerTensor> {
4273        self.grad_optional(output, wrt)?
4274            .ok_or_else(|| crate::traced::inactive_wrt_error("grad", &wrt.key))
4275    }
4276
4277    /// Gradient that returns `None` when `wrt` is inactive.
4278    ///
4279    /// # Examples
4280    ///
4281    /// ```
4282    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4283    /// use tenferro_cpu::CpuBackend;
4284    ///
4285    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4286    /// let x = EagerTensor::requires_grad_in(
4287    ///     Tensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap(),
4288    ///     ctx.clone(),
4289    /// )?;
4290    /// let y = EagerTensor::requires_grad_in(
4291    ///     Tensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap(),
4292    ///     ctx.clone(),
4293    /// )?;
4294    /// let loss = ctx.with_eager_session(|s| s.mul(&y, &y))?;
4295    /// assert!(ctx.grad_optional(&loss, &x)?.is_none());
4296    /// # Ok::<(), tenferro_ad::Error>(())
4297    /// ```
4298    ///
4299    /// # Errors
4300    ///
4301    /// Returns [`tenferro_runtime::Error::NonScalarGrad`] for a non-scalar
4302    /// output, [`Error::ContextMismatch`] for a foreign runtime, or a typed
4303    /// validation/backend/runtime-state error from eager execution.
4304    pub fn grad_optional(
4305        self: &Arc<Self>,
4306        output: &EagerTensor,
4307        wrt: &EagerTensor,
4308    ) -> Result<Option<EagerTensor>> {
4309        if !output.shape().is_empty() {
4310            return Err(Error::NonScalarGrad {
4311                shape: output.shape().to_vec(),
4312            });
4313        }
4314
4315        let value = output.to_tensor()?;
4316        let seed = self.with_execution_session(|session| one_like_tensor(&value, session))??;
4317        let seed = EagerTensor::new_result(Arc::clone(self), eager_val_key(), seed, false, None)?;
4318        self.vjp_optional(output, wrt, &seed)
4319    }
4320
4321    /// Reverse-mode vector-Jacobian product for eager tensors.
4322    ///
4323    /// # Examples
4324    ///
4325    /// ```
4326    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4327    /// use tenferro_cpu::CpuBackend;
4328    ///
4329    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4330    /// let x = EagerTensor::requires_grad_in(
4331    ///     Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0]).unwrap(),
4332    ///     ctx.clone(),
4333    /// )?;
4334    /// let y = ctx.with_eager_session(|s| s.mul(&x, &x))?;
4335    /// let seed = EagerTensor::from_tensor_in(
4336    ///     Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 1.0]).unwrap(),
4337    ///     ctx.clone(),
4338    /// )?;
4339    /// let dx = ctx.vjp(&y, &x, &seed)?;
4340    /// assert_eq!(dx.value()?.as_slice::<f64>().unwrap(), &[4.0, 6.0]);
4341    /// # Ok::<(), tenferro_ad::Error>(())
4342    /// ```
4343    ///
4344    /// # Errors
4345    ///
4346    /// Returns [`Error::ContextMismatch`] for tensors from different eager
4347    /// runtimes, [`Error::Validation`] when the cotangent shape or dtype does
4348    /// not match the output, [`Error::UnsupportedAdRule`] when a rule is not
4349    /// registered, or a typed backend/runtime-state error. An inactive `wrt`
4350    /// returns [`Error::Validation`] with `argument: "wrt"`; use
4351    /// [`vjp_optional`](Self::vjp_optional) to observe that state.
4352    pub fn vjp(
4353        self: &Arc<Self>,
4354        output: &EagerTensor,
4355        wrt: &EagerTensor,
4356        cotangent: &EagerTensor,
4357    ) -> Result<EagerTensor> {
4358        self.vjp_optional(output, wrt, cotangent)?
4359            .ok_or_else(|| crate::traced::inactive_wrt_error("vjp", &wrt.key))
4360    }
4361
4362    /// Reverse-mode vector-Jacobian product that returns `None` for inactive inputs.
4363    ///
4364    /// # Examples
4365    ///
4366    /// ```
4367    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4368    /// use tenferro_cpu::CpuBackend;
4369    ///
4370    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4371    /// let x = EagerTensor::requires_grad_in(
4372    ///     Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(),
4373    ///     ctx.clone(),
4374    /// )?;
4375    /// let y = EagerTensor::requires_grad_in(
4376    ///     Tensor::from_vec_col_major(vec![1], vec![4.0_f64]).unwrap(),
4377    ///     ctx.clone(),
4378    /// )?;
4379    /// let seed = EagerTensor::from_tensor_in(
4380    ///     Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(),
4381    ///     ctx.clone(),
4382    /// )?;
4383    /// let loss = ctx.with_eager_session(|s| s.mul(&y, &y))?;
4384    /// assert!(ctx.vjp_optional(&loss, &x, &seed)?.is_none());
4385    /// # Ok::<(), tenferro_ad::Error>(())
4386    /// ```
4387    ///
4388    /// # Errors
4389    ///
4390    /// Returns [`Error::ContextMismatch`] for tensors from different eager
4391    /// runtimes, [`Error::Validation`] when the cotangent shape or dtype does
4392    /// not match the output, [`Error::UnsupportedAdRule`] when a rule is not
4393    /// registered, or a typed backend/runtime-state error.
4394    pub fn vjp_optional(
4395        self: &Arc<Self>,
4396        output: &EagerTensor,
4397        wrt: &EagerTensor,
4398        cotangent: &EagerTensor,
4399    ) -> Result<Option<EagerTensor>> {
4400        validate_same_runtime(self, output, "vjp output")?;
4401        validate_same_runtime(self, wrt, "vjp wrt")?;
4402        validate_same_runtime(self, cotangent, "vjp cotangent")?;
4403        validate_seed_tensor("vjp", output, cotangent)?;
4404        Ok(semantic_eager_vjp_many(self, output, &[wrt], cotangent)?
4405            .pop()
4406            .flatten())
4407    }
4408
4409    /// Forward-mode Jacobian-vector product for eager tensors.
4410    ///
4411    /// # Examples
4412    ///
4413    /// ```
4414    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4415    /// use tenferro_cpu::CpuBackend;
4416    ///
4417    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4418    /// let x = EagerTensor::requires_grad_in(
4419    ///     Tensor::from_vec_col_major(vec![1], vec![3.0_f64]).unwrap(),
4420    ///     ctx.clone(),
4421    /// )?;
4422    /// let tangent = EagerTensor::from_tensor_in(
4423    ///     Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(),
4424    ///     ctx.clone(),
4425    /// )?;
4426    /// let y = ctx.with_eager_session(|s| s.mul(&x, &x))?;
4427    /// let dy = ctx.jvp(&y, &x, &tangent)?;
4428    /// assert_eq!(dy.value()?.as_slice::<f64>().unwrap(), &[6.0]);
4429    /// # Ok::<(), tenferro_ad::Error>(())
4430    /// ```
4431    ///
4432    /// # Errors
4433    ///
4434    /// Returns [`Error::ContextMismatch`] for tensors from different eager
4435    /// runtimes, [`Error::Validation`] when the tangent shape or dtype does not
4436    /// match `wrt`, [`Error::UnsupportedAdRule`] when a rule is unavailable, or
4437    /// a typed backend/runtime-state error. An inactive `wrt` returns
4438    /// [`Error::Validation`] with `argument: "wrt"`; use
4439    /// [`jvp_optional`](Self::jvp_optional) to observe that state.
4440    pub fn jvp(
4441        self: &Arc<Self>,
4442        output: &EagerTensor,
4443        wrt: &EagerTensor,
4444        tangent: &EagerTensor,
4445    ) -> Result<EagerTensor> {
4446        self.jvp_optional(output, wrt, tangent)?
4447            .ok_or_else(|| crate::traced::inactive_wrt_error("jvp", &wrt.key))
4448    }
4449
4450    /// Forward-mode Jacobian-vector product that returns `None` for inactive outputs.
4451    ///
4452    /// # Examples
4453    ///
4454    /// ```
4455    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
4456    /// use tenferro_cpu::CpuBackend;
4457    ///
4458    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
4459    /// let x = EagerTensor::requires_grad_in(
4460    ///     Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(),
4461    ///     ctx.clone(),
4462    /// )?;
4463    /// let y = EagerTensor::requires_grad_in(
4464    ///     Tensor::from_vec_col_major(vec![1], vec![4.0_f64]).unwrap(),
4465    ///     ctx.clone(),
4466    /// )?;
4467    /// let tangent = EagerTensor::from_tensor_in(
4468    ///     Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(),
4469    ///     ctx.clone(),
4470    /// )?;
4471    /// let loss = ctx.with_eager_session(|s| s.mul(&y, &y))?;
4472    /// assert!(ctx.jvp_optional(&loss, &x, &tangent)?.is_none());
4473    /// # Ok::<(), tenferro_ad::Error>(())
4474    /// ```
4475    ///
4476    /// # Errors
4477    ///
4478    /// Returns [`Error::ContextMismatch`] for tensors from different eager
4479    /// runtimes, [`Error::Validation`] when the tangent shape or dtype does not
4480    /// match `wrt`, [`Error::UnsupportedAdRule`] when a rule is unavailable, or
4481    /// a typed backend/runtime-state error.
4482    pub fn jvp_optional(
4483        self: &Arc<Self>,
4484        output: &EagerTensor,
4485        wrt: &EagerTensor,
4486        tangent: &EagerTensor,
4487    ) -> Result<Option<EagerTensor>> {
4488        validate_same_runtime(self, output, "jvp output")?;
4489        validate_same_runtime(self, wrt, "jvp wrt")?;
4490        validate_same_runtime(self, tangent, "jvp tangent")?;
4491        validate_seed_tensor("jvp", wrt, tangent)?;
4492        // Unification 7: semantic path is the only JVP path.
4493        match semantic_eager_jvp_optional(self, output, wrt, tangent)? {
4494            Some(result) => Ok(result),
4495            None => Ok(None),
4496        }
4497    }
4498
4499    fn store_grads(
4500        &self,
4501        cotangents: &HashMap<ValueKey<StdTensorOp>, Tensor>,
4502        session: &mut dyn BackendSession,
4503    ) -> Result<()> {
4504        let mut updates = Vec::new();
4505
4506        {
4507            let mut slots = self.lock_grad_slots()?;
4508            slots.retain(|key, slot| {
4509                let Some(slot) = slot.upgrade() else {
4510                    return false;
4511                };
4512
4513                if let Some(incoming) = cotangents.get(key) {
4514                    updates.push((slot, incoming));
4515                }
4516
4517                true
4518            });
4519        }
4520
4521        for (slot, incoming) in updates {
4522            let mut current = slot.lock().map_err(|_| {
4523                Error::runtime_state(
4524                    "eager_gradient_slot",
4525                    ErrorPhase::Execution,
4526                    "lock poisoned",
4527                )
4528            })?;
4529            let next = match current.as_ref() {
4530                Some(existing) => {
4531                    let existing_read = existing.tensor_read("EagerRuntime::store_grads")?;
4532                    let incoming_read = TensorRead::from_tensor(incoming);
4533                    let tensor = session
4534                        .add_read(existing_read, incoming_read)
4535                        .map_err(Error::from)?;
4536                    AdValueRecord::from_tensor(tensor, "EagerRuntime::store_grads")?
4537                }
4538                None => {
4539                    let duplicate = session
4540                        .to_contiguous_read(TensorRead::from_tensor(incoming))
4541                        .map_err(Error::from)?;
4542                    AdValueRecord::from_tensor(duplicate, "EagerRuntime::store_grads")?
4543                }
4544            };
4545            *current = Some(next);
4546        }
4547
4548        Ok(())
4549    }
4550}
4551
4552#[derive(Clone, Debug, PartialEq, Eq, Hash)]
4553struct PreparedDerivativeCacheKey {
4554    semantic_fingerprint: SemanticFingerprint,
4555    runtime_epoch: RuntimeEpoch,
4556    active_inputs: Box<[bool]>,
4557    input_metadata: Box<[ProgramValueMetadata]>,
4558}
4559
4560/// Cached prepared derivative: program + index metadata.
4561#[derive(Debug)]
4562struct PreparedDerivative {
4563    program: Arc<CompiledGraph>,
4564    execution_program: Arc<CompiledGraph>,
4565    saved_input_indices: Vec<usize>,
4566    prepared: Arc<PreparedCompiledGraph>,
4567    seed_input_index: usize,
4568    derivative_output_indices: Box<[Option<usize>]>,
4569}
4570
4571#[derive(Debug)]
4572struct PreparedDerivativeCache {
4573    limits: AdTransformCacheLimits,
4574    entries: LruCache<PreparedDerivativeCacheKey, PreparedDerivativeCacheEntry>,
4575    stats: CacheStats,
4576}
4577
4578impl PreparedDerivativeCache {
4579    fn limits(&self) -> AdTransformCacheLimits {
4580        self.limits
4581    }
4582
4583    fn set_limits(&mut self, limits: AdTransformCacheLimits) {
4584        self.limits = limits;
4585        self.evict_to_limits();
4586    }
4587
4588    fn clear(&mut self) {
4589        let clears = self.stats.clears.saturating_add(1);
4590        self.entries.clear();
4591        self.stats = CacheStats {
4592            clears,
4593            ..CacheStats::empty()
4594        };
4595    }
4596
4597    fn stats(&self) -> CacheStats {
4598        self.stats
4599    }
4600
4601    fn get(&mut self, key: &PreparedDerivativeCacheKey) -> Option<Arc<PreparedDerivative>> {
4602        match self.entries.get(key) {
4603            Some(entry) => {
4604                self.stats.hits = self.stats.hits.saturating_add(1);
4605                Some(Arc::clone(&entry.value))
4606            }
4607            None => {
4608                self.stats.misses = self.stats.misses.saturating_add(1);
4609                None
4610            }
4611        }
4612    }
4613
4614    fn insert(&mut self, key: PreparedDerivativeCacheKey, value: Arc<PreparedDerivative>) {
4615        let retained_bytes = prepared_derivative_cache_entry_retained_bytes(&key, value.as_ref());
4616        let entry = PreparedDerivativeCacheEntry {
4617            value,
4618            retained_bytes,
4619        };
4620        self.stats.retained_bytes = self.stats.retained_bytes.saturating_add(retained_bytes);
4621        if let Some((_old_key, old_entry)) = self.entries.push(key, entry) {
4622            self.stats.retained_bytes = self
4623                .stats
4624                .retained_bytes
4625                .saturating_sub(old_entry.retained_bytes);
4626        }
4627        self.stats.entries = self.entries.len();
4628        self.evict_to_limits();
4629    }
4630
4631    fn evict_to_limits(&mut self) {
4632        while self.entries.len() > self.limits.max_entries().get()
4633            || self
4634                .limits
4635                .max_retained_bytes()
4636                .is_some_and(|limit| self.stats.retained_bytes > limit.get())
4637        {
4638            let Some((_key, entry)) = self.entries.pop_lru() else {
4639                break;
4640            };
4641            self.stats.retained_bytes = self
4642                .stats
4643                .retained_bytes
4644                .saturating_sub(entry.retained_bytes);
4645            self.stats.evictions = self.stats.evictions.saturating_add(1);
4646        }
4647        self.stats.entries = self.entries.len();
4648    }
4649}
4650
4651impl Default for PreparedDerivativeCache {
4652    fn default() -> Self {
4653        Self {
4654            limits: AdTransformCacheLimits::default(),
4655            entries: LruCache::unbounded(),
4656            stats: CacheStats::empty(),
4657        }
4658    }
4659}
4660
4661#[derive(Debug)]
4662struct PreparedDerivativeCacheEntry {
4663    value: Arc<PreparedDerivative>,
4664    retained_bytes: usize,
4665}
4666
4667fn prepared_derivative_cache_entry_retained_bytes(
4668    key: &PreparedDerivativeCacheKey,
4669    value: &PreparedDerivative,
4670) -> usize {
4671    size_of::<PreparedDerivativeCacheKey>()
4672        .saturating_add(size_of_val(key.active_inputs.as_ref()))
4673        .saturating_add(size_of_val(value.derivative_output_indices.as_ref()))
4674        .saturating_add(size_of_val(value.saved_input_indices.as_slice()))
4675        .saturating_add(compiled_graph_retained_bytes(
4676            value.execution_program.as_ref(),
4677        ))
4678        .saturating_add(
4679            key.input_metadata
4680                .len()
4681                .saturating_mul(size_of::<ProgramValueMetadata>()),
4682        )
4683        .saturating_add(size_of::<PreparedDerivative>())
4684        .saturating_add(compiled_graph_retained_bytes(value.program.as_ref()))
4685        .saturating_add(prepared_compiled_graph_retained_bytes(
4686            value.prepared.as_ref(),
4687            value.program.as_ref(),
4688        ))
4689}
4690
4691fn prepared_compiled_graph_retained_bytes(
4692    prepared: &PreparedCompiledGraph,
4693    derivative_program: &CompiledGraph,
4694) -> usize {
4695    size_of_val(prepared).saturating_add(compiled_graph_retained_bytes(derivative_program))
4696}
4697
4698fn compiled_graph_retained_bytes(program: &CompiledGraph) -> usize {
4699    size_of::<CompiledGraph>()
4700        .saturating_add(size_of_val(program.input_keys()))
4701        .saturating_add(program.bindings().len().saturating_mul(size_of::<usize>()))
4702        .saturating_add(semantic_program_retained_bytes(program.program()))
4703}
4704
4705fn semantic_program_retained_bytes(program: &SemanticProgram) -> usize {
4706    size_of::<SemanticProgram>()
4707        .saturating_add(size_of_val(program.inputs()))
4708        .saturating_add(size_of_val(program.outputs()))
4709        .saturating_add(
4710            program
4711                .operations()
4712                .len()
4713                .saturating_mul(size_of::<usize>()),
4714        )
4715        .saturating_add(
4716            program
4717                .shape_guards()
4718                .len()
4719                .saturating_mul(size_of::<usize>()),
4720        )
4721}
4722
4723fn semantic_eager_vjp_many(
4724    ctx: &Arc<EagerRuntime>,
4725    output: &EagerTensor,
4726    wrts: &[&EagerTensor],
4727    cotangent: &EagerTensor,
4728) -> Result<Vec<Option<EagerTensor>>> {
4729    if !eager_semantic_vjp_enabled() || wrts.is_empty() {
4730        return Ok(vec![None; wrts.len()]);
4731    }
4732    let Some(raw_output_trace) = output.semantic_trace.as_ref() else {
4733        return Ok(vec![None; wrts.len()]);
4734    };
4735    if !wrts.iter().any(|wrt| {
4736        wrt.semantic_trace
4737            .as_ref()
4738            .and_then(TracedTensor::input_key)
4739            .is_some_and(|key| raw_output_trace.has_attached_input_key(&key))
4740    }) {
4741        return Ok(vec![None; wrts.len()]);
4742    }
4743
4744    // First AD request on this output: run the deferred graph analysis over
4745    // the whole raw carrier chain once (metadata registration + constraint
4746    // scopes), so `compile_ad_source` sees the same analyzed graph the eager
4747    // forward used to append.
4748    let output_trace = analyze_deferred_semantic_trace(raw_output_trace)?;
4749    let saved = output
4750        .trace
4751        .as_ref()
4752        .map(EagerTrace::collect)
4753        .unwrap_or_default();
4754    let mut source_outputs = vec![&output_trace];
4755    source_outputs.extend(saved.iter().map(|value| &value.trace));
4756
4757    // Residual roots keep the semantic producers available to differentiation;
4758    // only the execution-only derivative replaces their numerical values.
4759    let mut compiler = GraphCompiler::new();
4760    let source =
4761        tenferro_runtime::ad_support::compile_ad_source_many(&mut compiler, &source_outputs)?;
4762    if source.output_count() != source_outputs.len()
4763        || source.input_keys().len() != source.input_count()
4764        || source.bindings().len() != source.input_count()
4765    {
4766        return Ok(vec![None; wrts.len()]);
4767    }
4768    let wrt_input_indices = wrts
4769        .iter()
4770        .map(|wrt| {
4771            let key = wrt.semantic_trace.as_ref()?.input_key()?;
4772            source.input_key_index(&key)
4773        })
4774        .collect::<Vec<_>>();
4775    let mut active_inputs = vec![false; source.input_count()];
4776    for &index in wrt_input_indices.iter().flatten() {
4777        active_inputs[index] = true;
4778    }
4779    if !active_inputs.iter().any(|&active| active) {
4780        return Ok(vec![None; wrts.len()]);
4781    }
4782
4783    // One transform and execution for the complete active set shares primal
4784    // work and cotangents between leaves instead of replaying them per target.
4785    // S2: check prepared-derivative cache before AD transform + compile_frozen.
4786    let cache_key = PreparedDerivativeCacheKey {
4787        semantic_fingerprint: source.program().semantic_fingerprint(),
4788        runtime_epoch: ctx.runtime.epoch().map_err(|source| {
4789            Error::runtime_state_source("semantic_eager_vjp", ErrorPhase::Execution, source)
4790        })?,
4791        active_inputs: active_inputs.clone().into_boxed_slice(),
4792        input_metadata: source.frozen_program().input_metadata_with_bound_shapes(),
4793    };
4794    let prepared = { ctx.lock_prepared_derivative_cache()?.get(&cache_key) };
4795    let (
4796        seed_input_index,
4797        derivative_output_indices,
4798        derivative_program,
4799        execution_program,
4800        saved_input_indices,
4801        prepared_runtime,
4802    ) = if let Some(prepared) = prepared {
4803        (
4804            prepared.seed_input_index,
4805            prepared.derivative_output_indices.clone(),
4806            Arc::clone(&prepared.program),
4807            Arc::clone(&prepared.execution_program),
4808            prepared.saved_input_indices.clone(),
4809            Some(Arc::clone(&prepared.prepared)),
4810        )
4811    } else {
4812        let mut active_outputs = vec![false; source.output_count()];
4813        active_outputs[0] = true;
4814        let ad = AdContext::with_rules_and_transform_cache(
4815            ctx.semantic_extension_rules.clone(),
4816            Arc::clone(&ctx.ad_transform_cache),
4817        );
4818        let derivative = ad
4819            .vjp_program(source.frozen_program(), &active_inputs, &active_outputs)
4820            .map_err(|source| {
4821                Error::runtime_state_source("semantic_eager_vjp", ErrorPhase::GraphBuild, source)
4822            })?;
4823        let seed_input_index = derivative
4824            .derivative_input_indices()
4825            .first()
4826            .copied()
4827            .flatten();
4828        let Some(seed_input_index) = seed_input_index else {
4829            return Ok(vec![None; wrts.len()]);
4830        };
4831        let derivative_output_indices = derivative
4832            .derivative_output_indices()
4833            .to_vec()
4834            .into_boxed_slice();
4835        let program = Arc::new(compiler.compile_frozen_program(derivative.frozen())?);
4836        let (execution_program, saved_input_indices) = if saved.is_empty() {
4837            (Arc::clone(&program), Vec::new())
4838        } else {
4839            let (execution, indices) = crate::semantic_transform::semantic_vjp_with_saved_outputs(
4840                source.frozen_program(),
4841                &active_inputs,
4842                &active_outputs,
4843                &ctx.semantic_extension_rules,
4844                &(1..source.output_count()).collect::<Vec<_>>(),
4845            )
4846            .map_err(|source| {
4847                Error::runtime_state_source("semantic_eager_vjp", ErrorPhase::GraphBuild, source)
4848            })?;
4849            (
4850                Arc::new(compiler.compile_frozen_program(execution.frozen())?),
4851                indices,
4852            )
4853        };
4854        (
4855            seed_input_index,
4856            derivative_output_indices,
4857            program,
4858            execution_program,
4859            saved_input_indices,
4860            None,
4861        )
4862    };
4863
4864    let cotangent_tensor = Arc::new(RetainedValue::from_tensor(cotangent.to_tensor()?));
4865    let input_count = execution_program.input_count();
4866    let mut owned_inputs: Vec<Option<Tensor>> = (0..input_count).map(|_| None).collect();
4867    // Residuals, primal bindings and the seed are staged in one backend session
4868    // rather than one entry per value.
4869    ctx.with_execution_session(|session| -> Result<()> {
4870        for (value, &index) in saved.iter().zip(&saved_input_indices) {
4871            let read = value.value.tensor_read("eager residual")?;
4872            let Some(slot) = owned_inputs.get_mut(index) else {
4873                return Err(Error::Internal(format!(
4874                    "semantic eager VJP residual index {index} is outside {input_count} inputs"
4875                )));
4876            };
4877            *slot = Some(session.to_contiguous_read(read)?);
4878        }
4879        for (source_input_index, (_, tensor)) in source.bindings().iter().enumerate() {
4880            let Some(slot) = owned_inputs.get_mut(source_input_index) else {
4881                return Err(Error::Internal(format!(
4882                    "semantic eager VJP derivative program has no primal input slot {source_input_index}"
4883                )));
4884            };
4885            *slot = Some(copy_value_in_session(session, tensor)?);
4886        }
4887        let Some(slot) = owned_inputs.get_mut(seed_input_index) else {
4888            return Err(Error::Internal(format!(
4889                "semantic eager VJP seed input index {seed_input_index} is outside {input_count} inputs"
4890            )));
4891        };
4892        *slot = Some(copy_value_in_session(session, cotangent_tensor.as_ref())?);
4893        Ok(())
4894    })??;
4895    let input_refs = owned_inputs
4896        .iter()
4897        .enumerate()
4898        .map(|(index, tensor)| {
4899            tensor.as_ref().ok_or_else(|| {
4900                Error::Internal(format!(
4901                    "semantic eager VJP derivative input {index} was not populated"
4902                ))
4903            })
4904        })
4905        .collect::<Result<Vec<_>>>()?;
4906    let prepared_runtime = if let Some(prepared_runtime) = prepared_runtime {
4907        prepared_runtime
4908    } else {
4909        let prepared_runtime = Arc::new(
4910            ctx.runtime
4911                .prepare_compiled(&execution_program, &input_refs)?,
4912        );
4913        let entry = Arc::new(PreparedDerivative {
4914            program: Arc::clone(&derivative_program),
4915            execution_program: Arc::clone(&execution_program),
4916            saved_input_indices: saved_input_indices.clone(),
4917            prepared: Arc::clone(&prepared_runtime),
4918            seed_input_index,
4919            derivative_output_indices: derivative_output_indices.clone(),
4920        });
4921        ctx.lock_prepared_derivative_cache()?
4922            .insert(cache_key, entry);
4923        prepared_runtime
4924    };
4925    let mut outputs = ctx
4926        .runtime
4927        .run_prepared(&prepared_runtime, &input_refs)?
4928        .into_iter()
4929        .map(Some)
4930        .collect::<Vec<_>>();
4931    let cotangent_trace =
4932        TracedTensor::from_shared_tensor_value_symbolic_shape(Arc::clone(&cotangent_tensor))?;
4933
4934    #[cfg(test)]
4935    EAGER_SEMANTIC_VJP_EXECUTIONS.fetch_add(1, Ordering::Relaxed);
4936
4937    wrts.iter()
4938        .zip(wrt_input_indices)
4939        .map(|(wrt, input_index)| {
4940            let Some(derivative_output_index) = input_index
4941                .and_then(|index| derivative_output_indices.get(index).copied().flatten())
4942            else {
4943                return Ok(None);
4944            };
4945            let result = outputs
4946                .get_mut(derivative_output_index)
4947                .and_then(Option::take)
4948                .ok_or_else(|| {
4949                    Error::Internal(format!(
4950                "semantic eager VJP derivative output {derivative_output_index} unavailable"
4951            ))
4952                })?;
4953            let wrt_trace = wrt.semantic_trace.as_ref().ok_or_else(|| {
4954                Error::Internal("active eager VJP input has no semantic trace".into())
4955            })?;
4956            let semantic_trace = derivative_trace_from_frozen_program(
4957                &source,
4958                derivative_program.frozen_program(),
4959                derivative_output_index,
4960                &[(seed_input_index, Arc::clone(&cotangent_tensor))],
4961                &[&output_trace, wrt_trace, &cotangent_trace],
4962                None,
4963                "semantic_eager_vjp",
4964            )?;
4965            Ok(Some(EagerTensor::new_result_with_semantic_trace(
4966                Arc::clone(ctx),
4967                eager_val_key(),
4968                result,
4969                true,
4970                output.trace.clone(),
4971                Some(semantic_trace),
4972            )?))
4973        })
4974        .collect()
4975}
4976
4977fn semantic_eager_jvp_optional(
4978    ctx: &Arc<EagerRuntime>,
4979    output: &EagerTensor,
4980    wrt: &EagerTensor,
4981    tangent: &EagerTensor,
4982) -> Result<Option<Option<EagerTensor>>> {
4983    if !eager_semantic_vjp_enabled() {
4984        return Ok(None);
4985    }
4986    let (Some(raw_output_trace), Some(wrt_trace)) =
4987        (output.semantic_trace.as_ref(), wrt.semantic_trace.as_ref())
4988    else {
4989        return Ok(None);
4990    };
4991    let Some(wrt_key) = wrt_trace.input_key() else {
4992        return Ok(None);
4993    };
4994    if !raw_output_trace.has_attached_input_key(&wrt_key) {
4995        return Ok(None);
4996    }
4997
4998    // First AD request on this output: run the deferred graph analysis once
4999    // over the whole raw carrier chain before compiling.
5000    let output_trace = analyze_deferred_semantic_trace(raw_output_trace)?;
5001
5002    let mut compiler = GraphCompiler::new();
5003    let source = compile_ad_source(&mut compiler, &output_trace)?;
5004    if source.output_count() != 1
5005        || source.input_keys().len() != source.input_count()
5006        || source.bindings().len() != source.input_count()
5007    {
5008        return Ok(None);
5009    }
5010    let Some(wrt_input_index) = source.input_key_index(&wrt_key) else {
5011        return Ok(None);
5012    };
5013
5014    let mut active_inputs = vec![false; source.input_count()];
5015    if let Some(active) = active_inputs.get_mut(wrt_input_index) {
5016        *active = true;
5017    } else {
5018        return Ok(None);
5019    }
5020    let ad = AdContext::with_rules_and_transform_cache(
5021        ctx.semantic_extension_rules.clone(),
5022        Arc::clone(&ctx.ad_transform_cache),
5023    );
5024    let derivative = ad
5025        .jvp_program(source.frozen_program(), &active_inputs)
5026        .map_err(|source| {
5027            Error::runtime_state_source("semantic_eager_jvp", ErrorPhase::GraphBuild, source)
5028        })?;
5029    // derivative_input_indices maps source input → derivative seed input.
5030    let Some(seed_input_index) = derivative
5031        .derivative_input_indices()
5032        .get(wrt_input_index)
5033        .copied()
5034        .flatten()
5035    else {
5036        return Ok(Some(None));
5037    };
5038    // derivative_output_indices maps source output → derivative output.
5039    // There is always exactly one source output (guarded above).
5040    let Some(derivative_output_index) = derivative
5041        .derivative_output_indices()
5042        .first()
5043        .copied()
5044        .flatten()
5045    else {
5046        return Ok(Some(None));
5047    };
5048
5049    let derivative_program = compiler.compile_frozen_program(derivative.frozen())?;
5050    let tangent_tensor = Arc::new(RetainedValue::from_tensor(tangent.to_tensor()?));
5051    let input_count = derivative_program.input_count();
5052    let mut owned_inputs: Vec<Option<Tensor>> = (0..input_count).map(|_| None).collect();
5053    // Primal bindings and the seed are staged in one backend session.
5054    ctx.with_execution_session(|session| -> Result<()> {
5055        for (source_input_index, (_, tensor)) in source.bindings().iter().enumerate() {
5056            let Some(slot) = owned_inputs.get_mut(source_input_index) else {
5057                return Err(Error::Internal(format!(
5058                    "semantic eager JVP derivative program has no primal input slot {source_input_index}"
5059                )));
5060            };
5061            *slot = Some(copy_value_in_session(session, tensor)?);
5062        }
5063        let Some(slot) = owned_inputs.get_mut(seed_input_index) else {
5064            return Err(Error::Internal(format!(
5065                "semantic eager JVP seed input index {seed_input_index} is outside {input_count} inputs"
5066            )));
5067        };
5068        *slot = Some(copy_value_in_session(session, tangent_tensor.as_ref())?);
5069        Ok(())
5070    })??;
5071    let input_refs = owned_inputs
5072        .iter()
5073        .enumerate()
5074        .map(|(index, tensor)| {
5075            tensor.as_ref().ok_or_else(|| {
5076                Error::Internal(format!(
5077                    "semantic eager JVP derivative input {index} was not populated"
5078                ))
5079            })
5080        })
5081        .collect::<Result<Vec<_>>>()?;
5082    let outputs = ctx.runtime.run_compiled(&derivative_program, &input_refs)?;
5083    let output_count = outputs.len();
5084    let Some(result) = outputs.into_iter().nth(derivative_output_index) else {
5085        return Err(Error::Internal(format!(
5086            "semantic eager JVP derivative output index {derivative_output_index} is outside {} outputs",
5087            output_count
5088        )));
5089    };
5090    let tangent_trace =
5091        TracedTensor::from_shared_tensor_value_symbolic_shape(Arc::clone(&tangent_tensor))?;
5092    let semantic_trace = derivative_trace_from_frozen_program(
5093        &source,
5094        derivative.frozen(),
5095        derivative_output_index,
5096        &[(seed_input_index, Arc::clone(&tangent_tensor))],
5097        &[&output_trace, wrt_trace, &tangent_trace],
5098        None,
5099        "semantic_eager_jvp",
5100    )?;
5101
5102    Ok(Some(Some(EagerTensor::new_result_with_semantic_trace(
5103        Arc::clone(ctx),
5104        eager_val_key(),
5105        result,
5106        true,
5107        None,
5108        Some(semantic_trace),
5109    )?)))
5110}
5111
5112fn validate_same_runtime(
5113    runtime: &Arc<EagerRuntime>,
5114    tensor: &EagerTensor,
5115    role: &'static str,
5116) -> Result<()> {
5117    if tensor.ctx_id() != runtime.id() {
5118        return Err(Error::ContextMismatch {
5119            lhs: runtime.id(),
5120            rhs: tensor.ctx_id(),
5121        });
5122    }
5123    let _ = role;
5124    Ok(())
5125}
5126
5127fn copy_value_in_session(
5128    session: &mut dyn BackendSession,
5129    value: &RetainedValue,
5130) -> Result<Tensor> {
5131    let read = value.tensor_read().map_err(|error| {
5132        Error::runtime_state_source("copy_value_for_runtime", ErrorPhase::Execution, error)
5133    })?;
5134    session.to_contiguous_read(read).map_err(Error::from)
5135}
5136
5137fn validate_seed_tensor(op: &'static str, primal: &EagerTensor, seed: &EagerTensor) -> Result<()> {
5138    if primal.dtype() != seed.dtype() {
5139        return Err(
5140            tenferro_tensor::Error::dtype_mismatch(op, primal.dtype(), seed.dtype()).into(),
5141        );
5142    }
5143    if primal.shape() != seed.shape() {
5144        return Err(
5145            tenferro_tensor::Error::shape_mismatch(op, primal.shape(), seed.shape()).into(),
5146        );
5147    }
5148    Ok(())
5149}
5150
5151/// Eager tensor with reverse-mode autodiff over concrete tensor values.
5152///
5153/// This executes each primitive immediately and records a lightweight reverse
5154/// DAG for `backward()`. Gradients accumulate across repeated `backward()`
5155/// calls until they are cleared explicitly.
5156///
5157/// # Examples
5158///
5159/// ```
5160/// use tenferro_cpu::CpuBackend;
5161/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5162///
5163/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5164/// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0]).unwrap(), ctx)?;
5165/// for _ in 0..2 {
5166///     let loss = x.runtime().with_eager_session(|s| {
5167///         let squared = s.mul(&x, &x)?;
5168///         s.reduce_sum(&squared, Some(&[0]))
5169///     })?;
5170///     loss.backward()?;
5171/// }
5172///
5173/// assert_eq!(x.grad()?.unwrap().as_slice::<f64>().unwrap(), &[4.0, 8.0, 12.0]);
5174/// x.clear_grad()?;
5175///
5176/// assert!(x.grad().unwrap().is_none());
5177/// # Ok::<(), tenferro_ad::Error>(())
5178/// ```
5179#[derive(Clone)]
5180pub struct EagerTensor {
5181    pub(crate) key: ValueKey<StdTensorOp>,
5182    pub(crate) trace: Option<EagerTrace>,
5183    pub(crate) semantic_trace: Option<TracedTensor>,
5184    pub(crate) requires_grad: bool,
5185    grad_slot: GradSlot,
5186    pub(crate) ctx: Arc<EagerRuntime>,
5187    _record: Arc<EagerTensorRecord>,
5188}
5189
5190pub(crate) struct EagerTensorRecord {
5191    value: Arc<AdValueRecord>,
5192    key: ValueKey<StdTensorOp>,
5193    trace: Option<EagerTrace>,
5194    semantic_trace: Option<TracedTensor>,
5195    requires_grad: bool,
5196    grad_slot: GradSlot,
5197    ctx: Arc<EagerRuntime>,
5198}
5199
5200struct EagerTensorParts {
5201    ctx: Arc<EagerRuntime>,
5202    key: ValueKey<StdTensorOp>,
5203    requires_grad: bool,
5204    trace: Option<EagerTrace>,
5205    semantic_trace: Option<TracedTensor>,
5206    value: Arc<AdValueRecord>,
5207    register_value: bool,
5208}
5209
5210impl fmt::Debug for EagerTensor {
5211    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
5212        f.debug_struct("EagerTensor")
5213            .field("dtype", &self.dtype())
5214            .field("shape", &self.shape())
5215            .field("key", &self.key)
5216            .field("requires_grad", &self.requires_grad)
5217            .field("has_trace", &self.trace.is_some())
5218            .field("has_semantic_trace", &self.semantic_trace.is_some())
5219            .field("ctx_id", &self.ctx_id())
5220            .finish_non_exhaustive()
5221    }
5222}
5223
5224impl EagerTensor {
5225    /// Create an untracked eager tensor inside an existing eager context.
5226    ///
5227    /// # Examples
5228    ///
5229    /// ```
5230    /// use tenferro_cpu::CpuBackend;
5231    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5232    ///
5233    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5234    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap(), ctx)?;
5235    ///
5236    /// assert_eq!(x.value()?.as_slice::<f64>().unwrap(), &[1.0, 2.0]);
5237    /// # Ok::<(), tenferro_ad::Error>(())
5238    /// ```
5239    ///
5240    /// # Errors
5241    ///
5242    /// Returns [`tenferro_runtime::Error::RuntimeState`] when metadata cannot
5243    /// be registered in the target context, or a typed tensor/backend error
5244    /// while materializing the source value.
5245    pub fn from_tensor_in(tensor: Tensor, ctx: Arc<EagerRuntime>) -> Result<Self> {
5246        Self::new_leaf(ctx, tensor, false)
5247    }
5248
5249    /// Create an untracked eager tensor from compact column-major data inside
5250    /// an existing eager runtime.
5251    ///
5252    /// # Errors
5253    ///
5254    /// Returns [`Error::TensorRuntime`] with
5255    /// [`tenferro_tensor::ValidationError::ShapeMismatch`] when the shape and
5256    /// data length disagree, or with
5257    /// [`tenferro_tensor::ValidationError::IntegerOverflow`] when shape
5258    /// arithmetic overflows. Returns [`Error::RuntimeState`] when eager
5259    /// metadata cannot be registered.
5260    pub fn from_vec_col_major_in<T: TensorScalar>(
5261        shape: impl IntoShapeVec,
5262        data: Vec<T>,
5263        ctx: Arc<EagerRuntime>,
5264    ) -> Result<Self> {
5265        Self::from_tensor_in(Tensor::from_vec_col_major(shape, data)?, ctx)
5266    }
5267
5268    /// Create a tracked eager leaf inside an existing eager context.
5269    ///
5270    /// # Examples
5271    ///
5272    /// ```
5273    /// use tenferro_cpu::CpuBackend;
5274    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5275    ///
5276    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5277    /// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap(), ctx)?;
5278    ///
5279    /// assert!(x.grad().unwrap().is_none());
5280    /// # Ok::<(), tenferro_ad::Error>(())
5281    /// ```
5282    ///
5283    /// # Errors
5284    ///
5285    /// Returns [`tenferro_runtime::Error::RuntimeState`] when gradient metadata
5286    /// cannot be registered in the target context, or a typed tensor/backend
5287    /// error while creating the leaf.
5288    pub fn requires_grad_in(tensor: Tensor, ctx: Arc<EagerRuntime>) -> Result<Self> {
5289        Self::new_leaf(ctx, tensor, true)
5290    }
5291
5292    pub(crate) fn new_leaf(
5293        ctx: Arc<EagerRuntime>,
5294        tensor: Tensor,
5295        requires_grad: bool,
5296    ) -> Result<Self> {
5297        Self::new_leaf_with_session(ctx, tensor, requires_grad, None)
5298    }
5299
5300    pub(crate) fn new_leaf_in_session(
5301        ctx: Arc<EagerRuntime>,
5302        tensor: Tensor,
5303        requires_grad: bool,
5304        session: &mut dyn BackendSession,
5305    ) -> Result<Self> {
5306        Self::new_leaf_with_session(ctx, tensor, requires_grad, Some(session))
5307    }
5308
5309    fn new_leaf_with_session(
5310        ctx: Arc<EagerRuntime>,
5311        tensor: Tensor,
5312        requires_grad: bool,
5313        session: Option<&mut dyn BackendSession>,
5314    ) -> Result<Self> {
5315        let key = eager_val_key();
5316        // A host-placement tensor needs no backend session: the CPU backend only
5317        // copies the host buffer for it, so entering a session would add
5318        // admission, provider exclusion, and session construction without doing
5319        // any provider work (#1704). Views and backend-family reads keep the
5320        // session path, and device runtimes keep it for every read.
5321        let read = TensorRead::from_tensor(&tensor);
5322        let semantic_tensor = match session {
5323            Some(session) => session.to_contiguous_read(read).map_err(Error::from)?,
5324            None => match ctx.to_contiguous_host_read(&read)? {
5325                Some(materialized) => materialized,
5326                None => ctx
5327                    .with_execution_session(|session| session.to_contiguous_read(read))?
5328                    .map_err(Error::from)?,
5329            },
5330        };
5331        let semantic_value = Arc::new(RetainedValue::from_tensor(semantic_tensor));
5332        let semantic_trace = TracedTensor::from_shared_tensor_value_symbolic_shape(semantic_value)?;
5333        // Deferred materialization: the per-op/leaf global-metadata registry
5334        // write for the eager tensor key was unreadable after the semantic
5335        // trace became the sole AD carrier, so it is dropped.
5336        // ponytail: leaf input-key metadata is still registered by
5337        // `from_shared_tensor_value_symbolic_shape`; the eager-key entry was
5338        // vestigial and removed. Add back only if something reads it.
5339        let value = AdValueRecord::from_tensor(tensor, "EagerTensor::new_leaf")?;
5340        Self::from_parts(EagerTensorParts {
5341            ctx,
5342            key,
5343            requires_grad,
5344            trace: None,
5345            semantic_trace: Some(semantic_trace),
5346            value,
5347            register_value: true,
5348        })
5349    }
5350
5351    pub(crate) fn new_result(
5352        ctx: Arc<EagerRuntime>,
5353        key: ValueKey<StdTensorOp>,
5354        tensor: Tensor,
5355        requires_grad: bool,
5356        trace: Option<EagerTrace>,
5357    ) -> Result<Self> {
5358        Self::new_result_with_semantic_trace(ctx, key, tensor, requires_grad, trace, None)
5359    }
5360
5361    pub(crate) fn new_result_with_semantic_trace(
5362        ctx: Arc<EagerRuntime>,
5363        key: ValueKey<StdTensorOp>,
5364        tensor: Tensor,
5365        requires_grad: bool,
5366        trace: Option<EagerTrace>,
5367        semantic_trace: Option<TracedTensor>,
5368    ) -> Result<Self> {
5369        let value = AdValueRecord::from_tensor(tensor, "EagerTensor::new_result")?;
5370        Self::from_parts(EagerTensorParts {
5371            ctx,
5372            key,
5373            requires_grad,
5374            trace,
5375            semantic_trace,
5376            value,
5377            register_value: true,
5378        })
5379    }
5380
5381    pub(crate) fn new_unregistered_result_with_semantic_trace(
5382        ctx: Arc<EagerRuntime>,
5383        key: ValueKey<StdTensorOp>,
5384        tensor: Tensor,
5385        requires_grad: bool,
5386        trace: Option<EagerTrace>,
5387        semantic_trace: Option<TracedTensor>,
5388    ) -> Result<Self> {
5389        let value = AdValueRecord::from_tensor(tensor, "EagerTensor::new_unregistered_result")?;
5390        Self::from_parts(EagerTensorParts {
5391            ctx,
5392            key,
5393            requires_grad,
5394            trace,
5395            semantic_trace,
5396            value,
5397            register_value: false,
5398        })
5399    }
5400
5401    pub(crate) fn new_result_value(
5402        ctx: Arc<EagerRuntime>,
5403        key: ValueKey<StdTensorOp>,
5404        value: TensorValue,
5405        requires_grad: bool,
5406        trace: Option<EagerTrace>,
5407        semantic_trace: Option<TracedTensor>,
5408    ) -> Result<Self> {
5409        let (group, slot, dtype, shape) = value.try_into_group_parts().map_err(|_| {
5410            Error::runtime_state(
5411                "EagerTensor::new_result_value",
5412                ErrorPhase::Execution,
5413                "a TensorValue could not be transferred into its allocation group",
5414            )
5415        })?;
5416        let value = AdValueRecord::from_group(group, slot, dtype, shape);
5417        Self::from_parts(EagerTensorParts {
5418            ctx,
5419            key,
5420            requires_grad,
5421            trace,
5422            semantic_trace,
5423            value,
5424            register_value: true,
5425        })
5426    }
5427
5428    fn from_parts(parts: EagerTensorParts) -> Result<Self> {
5429        let EagerTensorParts {
5430            ctx,
5431            key,
5432            requires_grad,
5433            trace,
5434            semantic_trace,
5435            value,
5436            register_value,
5437        } = parts;
5438        let grad_slot = Arc::new(Mutex::new(None));
5439        if requires_grad {
5440            ctx.try_register_grad_slot(&key, &grad_slot)?;
5441        }
5442        let record = Arc::new(EagerTensorRecord {
5443            value: Arc::clone(&value),
5444            key: key.clone(),
5445            trace: trace.clone(),
5446            semantic_trace: semantic_trace.clone(),
5447            requires_grad,
5448            grad_slot: Arc::clone(&grad_slot),
5449            ctx: Arc::clone(&ctx),
5450        });
5451        if register_value {
5452            ctx.try_register_value_record(&key, &record)?;
5453        }
5454
5455        Ok(Self {
5456            key,
5457            trace,
5458            semantic_trace,
5459            requires_grad,
5460            grad_slot,
5461            ctx,
5462            _record: record,
5463        })
5464    }
5465
5466    pub(crate) fn new_untracked_result(ctx: Arc<EagerRuntime>, tensor: Tensor) -> Result<Self> {
5467        let value =
5468            AdValueRecord::from_untracked_tensor(tensor, "EagerTensor::new_untracked_result")?;
5469        Ok(Self::new_untracked_value_record(ctx, value, None))
5470    }
5471
5472    pub(crate) fn new_untracked_value_result(
5473        ctx: Arc<EagerRuntime>,
5474        value: TensorValue,
5475    ) -> Result<Self> {
5476        Self::new_untracked_value_result_with_semantic_trace(ctx, value, None)
5477    }
5478
5479    pub(crate) fn new_untracked_value_result_with_semantic_trace(
5480        ctx: Arc<EagerRuntime>,
5481        value: TensorValue,
5482        semantic_trace: Option<TracedTensor>,
5483    ) -> Result<Self> {
5484        let (group, slot, dtype, shape) = value.try_into_group_parts().map_err(|_| {
5485            Error::runtime_state(
5486                "EagerTensor::new_untracked_value_result",
5487                ErrorPhase::Execution,
5488                "a TensorValue could not be transferred into its allocation group",
5489            )
5490        })?;
5491        let value = AdValueRecord::from_group(group, slot, dtype, shape);
5492        Ok(Self::new_untracked_value_record(ctx, value, semantic_trace))
5493    }
5494
5495    fn new_untracked_value_record(
5496        ctx: Arc<EagerRuntime>,
5497        value: Arc<AdValueRecord>,
5498        semantic_trace: Option<TracedTensor>,
5499    ) -> Self {
5500        let key = eager_val_key();
5501        let grad_slot = Arc::new(Mutex::new(None));
5502        let record = Arc::new(EagerTensorRecord {
5503            value,
5504            key: key.clone(),
5505            trace: None,
5506            semantic_trace: semantic_trace.clone(),
5507            requires_grad: false,
5508            grad_slot: Arc::clone(&grad_slot),
5509            ctx: Arc::clone(&ctx),
5510        });
5511        Self {
5512            key,
5513            trace: None,
5514            semantic_trace,
5515            requires_grad: false,
5516            grad_slot,
5517            ctx,
5518            _record: record,
5519        }
5520    }
5521
5522    pub(crate) fn from_record(record: Arc<EagerTensorRecord>) -> Self {
5523        Self {
5524            key: record.key.clone(),
5525            trace: record.trace.clone(),
5526            semantic_trace: record.semantic_trace.clone(),
5527            requires_grad: record.requires_grad,
5528            grad_slot: Arc::clone(&record.grad_slot),
5529            ctx: Arc::clone(&record.ctx),
5530            _record: record,
5531        }
5532    }
5533
5534    /// Detach this tensor from the reverse graph.
5535    ///
5536    /// The returned tensor keeps the concrete value but no longer contributes
5537    /// gradients to the original graph.
5538    ///
5539    /// # Examples
5540    ///
5541    /// ```
5542    /// use tenferro_cpu::CpuBackend;
5543    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5544    ///
5545    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5546    /// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap(), ctx)?;
5547    /// let y = x.detach();
5548    ///
5549    /// assert_eq!(y.value()?.as_slice::<f64>().unwrap(), &[1.0, 2.0]);
5550    /// assert!(y.grad().unwrap().is_none());
5551    /// # Ok::<(), tenferro_ad::Error>(())
5552    /// ```
5553    pub fn detach(&self) -> Self {
5554        let semantic_trace = self
5555            .duplicate_value()
5556            .ok()
5557            .and_then(|tensor| TracedTensor::from_tensor_symbolic_shape(tensor).ok());
5558        Self::new_untracked_value_record(
5559            self.ctx.clone(),
5560            Arc::clone(&self._record.value),
5561            semantic_trace,
5562        )
5563    }
5564
5565    /// Detach this tensor from its graph and re-register it in a different
5566    /// context as an untracked leaf.
5567    ///
5568    /// # Examples
5569    ///
5570    /// ```
5571    /// use tenferro_cpu::CpuBackend;
5572    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5573    ///
5574    /// let ctx_a = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5575    /// let ctx_b = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5576    /// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap(), ctx_a)?;
5577    /// let d = x.detach_into(&ctx_b)?;
5578    ///
5579    /// assert!(!d.tracks_grad());
5580    /// assert_eq!(d.ctx_id(), ctx_b.id());
5581    /// # Ok::<(), tenferro_ad::Error>(())
5582    /// ```
5583    ///
5584    /// # Errors
5585    ///
5586    /// Returns [`Error::RuntimeState`] if the source cannot be materialized or
5587    /// the target context cannot register its metadata.
5588    pub fn detach_into(&self, ctx: &Arc<EagerRuntime>) -> Result<Self> {
5589        Self::from_tensor_in(self.to_tensor()?, Arc::clone(ctx))
5590    }
5591
5592    /// Borrow the retained value without creating an owner or copy.
5593    ///
5594    /// # Errors
5595    ///
5596    /// Returns [`Error::RuntimeState`] when the retained allocation-group
5597    /// descriptor is unavailable or invalid.
5598    pub fn value(&self) -> Result<ValueGuard<'_>> {
5599        self._record.value.value("EagerTensor::value")
5600    }
5601
5602    /// Explicitly duplicate this value into a fresh standalone allocation.
5603    ///
5604    /// # Errors
5605    ///
5606    /// Returns [`Error::RuntimeState`] when the retained value or execution
5607    /// session is unavailable, or a typed host/backend error when the value
5608    /// cannot be materialized as a contiguous tensor.
5609    ///
5610    /// # Examples
5611    ///
5612    /// ```
5613    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5614    /// use tenferro_cpu::CpuBackend;
5615    ///
5616    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5617    /// let value = EagerTensor::from_tensor_in(
5618    ///     Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0])?,
5619    ///     ctx,
5620    /// )?;
5621    /// let duplicate = value.duplicate_value()?;
5622    /// assert_eq!(duplicate.as_slice::<f64>()?, &[1.0, 2.0]);
5623    /// # Ok::<(), tenferro_ad::Error>(())
5624    /// ```
5625    pub fn duplicate_value(&self) -> Result<Tensor> {
5626        if let Some(tensor) = self.duplicate_host_value() {
5627            return Ok(tensor);
5628        }
5629        let read = self
5630            ._record
5631            .value
5632            .tensor_read("EagerTensor::duplicate_value")?;
5633        self.ctx
5634            .with_execution_session(|session| session.to_contiguous_read(read))?
5635            .map_err(Error::from)
5636    }
5637
5638    fn duplicate_host_value(&self) -> Option<Tensor> {
5639        // A pooled value duplicates through its descriptor view; a caller-owned
5640        // payload has no such view and duplicates through its own read path, which
5641        // copies the value while keeping its element type.
5642        self.value().ok()?.duplicate_host_tensor().ok()
5643    }
5644
5645    pub(crate) fn duplicate_value_in_session(
5646        &self,
5647        session: &mut dyn BackendSession,
5648    ) -> Result<Tensor> {
5649        if let Some(tensor) = self.duplicate_host_value() {
5650            return Ok(tensor);
5651        }
5652        let read = self
5653            ._record
5654            .value
5655            .tensor_read("EagerTensor::duplicate_value")?;
5656        session.to_contiguous_read(read).map_err(Error::from)
5657    }
5658
5659    // INVARIANT: the error variants return the unchanged eager handle so a
5660    // caller can retry ownership extraction without an implicit copy.
5661    #[allow(clippy::result_large_err)]
5662    /// Consume this handle and structurally extract its retained allocation.
5663    ///
5664    /// A shared handle is returned unchanged as [`IntoValueError::NotUnique`].
5665    /// Group extraction failures return the unchanged handle and typed group
5666    /// error; no copy or fallback materialization is attempted.
5667    ///
5668    /// # Errors
5669    ///
5670    /// Returns [`IntoValueError::NotUnique`] when another handle retains the
5671    /// value, or [`IntoValueError::Extract`] when structural group extraction
5672    /// fails because the allocation is aliased or its descriptor is invalid.
5673    ///
5674    /// # Examples
5675    ///
5676    /// ```
5677    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5678    /// use tenferro_cpu::CpuBackend;
5679    ///
5680    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5681    /// let value = EagerTensor::from_tensor_in(
5682    ///     Tensor::from_vec_col_major(vec![1], vec![3.0_f64])?,
5683    ///     ctx,
5684    /// )?;
5685    /// let owner = value
5686    ///     .into_value()
5687    ///     .expect("a uniquely owned value should be extractable");
5688    /// assert_eq!(owner.as_slice::<f64>()?, &[3.0]);
5689    /// # Ok::<(), tenferro_ad::Error>(())
5690    /// ```
5691    pub fn into_value(self) -> std::result::Result<Tensor, IntoValueError<Self>> {
5692        if Arc::strong_count(&self._record) != 1 {
5693            return Err(IntoValueError::NotUnique(self));
5694        }
5695        let Self { _record, .. } = self;
5696        let record = match Arc::try_unwrap(_record) {
5697            Ok(record) => record,
5698            Err(record) => return Err(IntoValueError::NotUnique(Self::from_record(record))),
5699        };
5700        let EagerTensorRecord {
5701            value,
5702            key,
5703            trace,
5704            semantic_trace,
5705            requires_grad,
5706            grad_slot,
5707            ctx,
5708        } = record;
5709        let value = match Arc::try_unwrap(value) {
5710            Ok(value) => value,
5711            Err(value) => {
5712                let record = Arc::new(EagerTensorRecord {
5713                    value,
5714                    key,
5715                    trace,
5716                    semantic_trace,
5717                    requires_grad,
5718                    grad_slot,
5719                    ctx,
5720                });
5721                return Err(IntoValueError::NotUnique(Self::from_record(record)));
5722            }
5723        };
5724        let AdValueRecord {
5725            container,
5726            dtype,
5727            shape,
5728        } = value;
5729        let container = match Arc::try_unwrap(container) {
5730            Ok(container) => container,
5731            Err(container) => {
5732                let record = Arc::new(EagerTensorRecord {
5733                    value: Arc::new(AdValueRecord {
5734                        container,
5735                        dtype,
5736                        shape,
5737                    }),
5738                    key,
5739                    trace,
5740                    semantic_trace,
5741                    requires_grad,
5742                    grad_slot,
5743                    ctx,
5744                });
5745                return Err(IntoValueError::NotUnique(Self::from_record(record)));
5746            }
5747        };
5748        let (group, slot) = match container {
5749            // A caller-owned payload is handed back to its owner unchanged.
5750            RetentionContainer::CallerOwned { tensor } => return Ok(*tensor),
5751            RetentionContainer::Owned { tensor } => return Ok(tensor),
5752            RetentionContainer::Pooled { group, slot } => (group, slot),
5753        };
5754        match group.into_tensor(slot) {
5755            Ok(tensor) => Ok(tensor),
5756            Err((group, error)) => {
5757                let record = Arc::new(EagerTensorRecord {
5758                    value: Arc::new(AdValueRecord {
5759                        container: Arc::new(RetentionContainer::Pooled { group, slot }),
5760                        dtype,
5761                        shape,
5762                    }),
5763                    key,
5764                    trace,
5765                    semantic_trace,
5766                    requires_grad,
5767                    grad_slot,
5768                    ctx,
5769                });
5770                Err(IntoValueError::Extract {
5771                    value: Self::from_record(record),
5772                    error,
5773                })
5774            }
5775        }
5776    }
5777
5778    /// Return this tensor's scalar dtype without materializing through
5779    /// [`value`](Self::value).
5780    pub fn dtype(&self) -> DType {
5781        self._record.value.dtype()
5782    }
5783
5784    /// Return this tensor's logical shape without materializing through
5785    /// [`value`](Self::value).
5786    pub fn shape(&self) -> &[usize] {
5787        self._record.value.shape()
5788    }
5789
5790    /// Borrow this tensor value as a [`TensorRead`].
5791    ///
5792    /// This is the preferred borrowed input boundary for executor calls. It
5793    /// preserves the option to replace eager storage with non-contiguous views
5794    /// without forcing callers through [`value`](Self::value).
5795    ///
5796    /// # Panics
5797    ///
5798    /// Panics if a validated eager value record becomes unavailable, which
5799    /// indicates an internal invariant violation.
5800    pub fn tensor_read(&self) -> TensorRead<'_> {
5801        self._record
5802            .value
5803            .tensor_read("EagerTensor::tensor_read")
5804            .expect("validated eager value record")
5805    }
5806
5807    /// Materialize this eager tensor as an owned [`Tensor`].
5808    ///
5809    /// This is the owned materialization boundary for callers that need a
5810    /// standalone compact tensor. The operation is fallible because eager
5811    /// values may be backed by lazy or backend-resident storage.
5812    ///
5813    /// # Errors
5814    ///
5815    /// Returns [`Error::RuntimeState`] if backend state is unavailable, or a
5816    /// typed tensor backend error when contiguous materialization fails.
5817    pub fn to_tensor(&self) -> Result<Tensor> {
5818        self.duplicate_value()
5819    }
5820
5821    /// Return the accumulated gradient currently stored for this tensor.
5822    ///
5823    /// The stored gradient accumulates across repeated `backward()` calls
5824    /// until it is cleared explicitly.
5825    ///
5826    /// For complex scalar losses, stored gradients use tenferro's
5827    /// Hermitian-adjoint cotangent convention. See
5828    /// <https://tensor4all.org/tenferro-rs/guides/complex-ad.html>.
5829    ///
5830    /// # Examples
5831    ///
5832    /// ```
5833    /// use tenferro_cpu::CpuBackend;
5834    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5835    ///
5836    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5837    /// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap(), ctx.clone()).unwrap();
5838    /// let loss = ctx.with_eager_session(|s| {
5839    ///     let y = s.exp(&x)?;
5840    ///     s.reduce_sum(&y, Some(&[0]))
5841    /// })?;
5842    /// let _cotangents = loss.backward().unwrap();
5843    ///
5844    /// let grad = x.grad()?.unwrap();
5845    /// assert_eq!(grad.shape(), &[2]);
5846    /// # Ok::<(), tenferro_ad::Error>(())
5847    /// ```
5848    ///
5849    /// # Errors
5850    ///
5851    /// Returns [`Error::RuntimeState`] if the gradient slot is poisoned or no
5852    /// longer available.
5853    pub fn grad(&self) -> Result<Option<GradientValue>> {
5854        self.grad_slot
5855            .lock()
5856            .map_err(|_| {
5857                Error::runtime_state(
5858                    "eager_gradient_slot",
5859                    ErrorPhase::Execution,
5860                    "lock poisoned",
5861                )
5862            })
5863            .map(|slot| {
5864                slot.as_ref().map(|record| GradientValue {
5865                    record: Arc::clone(record),
5866                    ctx: Arc::clone(&self.ctx),
5867                })
5868            })
5869    }
5870
5871    /// Clear the accumulated gradient stored for this tensor.
5872    ///
5873    /// This only affects this tensor's gradient slot. Other tensors in the
5874    /// same context retain their gradients until they are cleared explicitly or
5875    /// overwritten by later accumulation.
5876    ///
5877    /// # Examples
5878    ///
5879    /// ```
5880    /// use tenferro_cpu::CpuBackend;
5881    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5882    ///
5883    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5884    /// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0]).unwrap(), ctx.clone()).unwrap();
5885    /// let y = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![3], vec![4.0_f64, 5.0, 6.0]).unwrap(), ctx).unwrap();
5886    /// let loss = x.runtime().with_eager_session(|s| {
5887    ///     let product = s.mul(&x, &y)?;
5888    ///     s.reduce_sum(&product, Some(&[0]))
5889    /// })?;
5890    /// let _ = loss.backward().unwrap();
5891    ///
5892    /// x.clear_grad()?;
5893    ///
5894    /// assert!(x.grad()?.is_none());
5895    /// assert!(y.grad()?.is_some());
5896    /// # Ok::<(), tenferro_ad::Error>(())
5897    /// ```
5898    ///
5899    /// # Errors
5900    ///
5901    /// Returns [`Error::RuntimeState`] if the gradient slot lock is poisoned.
5902    pub fn clear_grad(&self) -> Result<()> {
5903        *self.grad_slot.lock().map_err(|_| {
5904            Error::runtime_state(
5905                "eager_gradient_slot",
5906                ErrorPhase::Execution,
5907                "lock poisoned",
5908            )
5909        })? = None;
5910        Ok(())
5911    }
5912
5913    /// Report whether this tensor participates in gradient tracking.
5914    ///
5915    /// Tracked tensors keep a gradient slot in their eager context; untracked
5916    /// tensors and detached tensors do not.
5917    ///
5918    /// # Examples
5919    ///
5920    /// ```
5921    /// use tenferro_cpu::CpuBackend;
5922    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5923    ///
5924    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5925    /// let plain = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap(), ctx.clone()).unwrap();
5926    /// let tracked = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![2], vec![3.0_f64, 4.0]).unwrap(), ctx.clone()).unwrap();
5927    /// let detached = tracked.detach();
5928    ///
5929    /// assert!(!plain.tracks_grad());
5930    /// assert!(tracked.tracks_grad());
5931    /// assert!(!detached.tracks_grad());
5932    /// # Ok::<(), tenferro_ad::Error>(())
5933    /// ```
5934    pub fn tracks_grad(&self) -> bool {
5935        self.requires_grad
5936    }
5937
5938    #[cfg(test)]
5939    fn debug_trace_saved_value_count(&self) -> Option<usize> {
5940        None
5941    }
5942
5943    /// Return the opaque identifier of the context this tensor belongs to.
5944    ///
5945    /// # Examples
5946    ///
5947    /// ```
5948    /// use tenferro_cpu::CpuBackend;
5949    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5950    ///
5951    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5952    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(), ctx.clone()).unwrap();
5953    ///
5954    /// assert_eq!(x.ctx_id(), ctx.id());
5955    /// # Ok::<(), tenferro_ad::Error>(())
5956    /// ```
5957    pub fn ctx_id(&self) -> ContextId {
5958        self.ctx.id()
5959    }
5960
5961    /// Borrow the eager runtime context that owns this tensor.
5962    pub fn runtime(&self) -> &Arc<EagerRuntime> {
5963        &self.ctx
5964    }
5965
5966    /// Check whether two tensors belong to the same eager context.
5967    ///
5968    /// # Examples
5969    ///
5970    /// ```
5971    /// use tenferro_cpu::CpuBackend;
5972    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
5973    ///
5974    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
5975    /// let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(), ctx.clone()).unwrap();
5976    /// let y = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(), ctx).unwrap();
5977    ///
5978    /// assert!(x.same_context(&y));
5979    /// # Ok::<(), tenferro_ad::Error>(())
5980    /// ```
5981    pub fn same_context(&self, other: &Self) -> bool {
5982        self.ctx_id() == other.ctx_id()
5983    }
5984
5985    #[cfg(test)]
5986    pub(crate) fn standard_graph_op(
5987        inputs: &[&Self],
5988        build_graph: impl FnOnce(&[TensorInputKey]) -> Result<Arc<Graph<StdTensorOp>>>,
5989    ) -> Result<Vec<Self>> {
5990        let Some(first) = inputs.first() else {
5991            return Err(Error::Internal(
5992                "standard eager graph op requires at least one input tensor".to_string(),
5993            ));
5994        };
5995        let ctx = Arc::clone(&first.ctx);
5996        for tensor in inputs.iter().skip(1) {
5997            if !first.same_context(tensor) {
5998                return Err(Error::ContextMismatch {
5999                    lhs: first.ctx_id(),
6000                    rhs: tensor.ctx_id(),
6001                });
6002            }
6003        }
6004
6005        let graph_input_keys = (0..inputs.len())
6006            .map(|_| next_input_key())
6007            .collect::<Vec<_>>();
6008        let graph = build_graph(&graph_input_keys)?;
6009        let initial_data = graph_input_keys
6010            .iter()
6011            .zip(inputs.iter())
6012            .map(|(key, tensor)| Ok((ValueKey::Input(key.clone()), tensor.to_tensor()?)))
6013            .collect::<Result<HashMap<_, _>>>()?;
6014        let execution = ctx.exec_standard_graph_outputs(graph.as_ref(), initial_data)?;
6015        if execution.outputs.len() != graph.outputs().len() {
6016            return Err(Error::Internal(format!(
6017                "standard eager graph op expected {} graph outputs, got {}",
6018                graph.outputs().len(),
6019                execution.outputs.len()
6020            )));
6021        }
6022
6023        if !eager_grad_recording_enabled() || !inputs.iter().any(|input| input.requires_grad) {
6024            return execution
6025                .outputs
6026                .into_iter()
6027                .map(|output| {
6028                    Self::new_unregistered_result_with_semantic_trace(
6029                        Arc::clone(&ctx),
6030                        eager_val_key(),
6031                        output,
6032                        false,
6033                        None,
6034                        None,
6035                    )
6036                })
6037                .collect();
6038        }
6039
6040        let recorded = record_eager_graph_outputs(
6041            graph.as_ref(),
6042            &graph_input_keys,
6043            &execution.outputs,
6044            inputs,
6045        )?;
6046        if recorded.traces.len() != execution.outputs.len() {
6047            return Err(Error::Internal(format!(
6048                "standard eager graph op expected {} eager traces, got {}",
6049                execution.outputs.len(),
6050                recorded.traces.len()
6051            )));
6052        }
6053
6054        recorded
6055            .traces
6056            .into_iter()
6057            .zip(recorded.semantic_traces)
6058            .zip(execution.outputs)
6059            .map(|((trace, semantic_trace), output)| {
6060                Self::new_result_with_semantic_trace(
6061                    Arc::clone(&ctx),
6062                    trace.key,
6063                    output,
6064                    trace.requires_grad,
6065                    trace.trace,
6066                    semantic_trace,
6067                )
6068            })
6069            .collect()
6070    }
6071
6072    /// Run reverse-mode AD from this scalar output.
6073    ///
6074    /// Returns the full cotangent map produced by the reverse pass and also
6075    /// accumulates into `grad()` for tracked eager tensors reachable from this
6076    /// output.
6077    ///
6078    /// For complex scalar outputs, cotangents use tenferro's Hermitian
6079    /// real-inner-product convention. See
6080    /// <https://tensor4all.org/tenferro-rs/guides/complex-ad.html>.
6081    ///
6082    /// # Examples
6083    ///
6084    /// ```
6085    /// use tenferro_cpu::CpuBackend;
6086    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
6087    ///
6088    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
6089    /// let x = EagerTensor::requires_grad_in(Tensor::from_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0]).unwrap(), ctx).unwrap();
6090    /// for _ in 0..2 {
6091    ///     let loss = x.runtime().with_eager_session(|s| {
6092    ///         let doubled = s.add(&x, &x)?;
6093    ///         s.reduce_sum(&doubled, Some(&[0]))
6094    ///     })?;
6095    ///     loss.backward()?;
6096    /// }
6097    ///
6098    /// assert_eq!(x.grad().unwrap().unwrap().as_slice::<f64>().unwrap(), &[4.0, 4.0, 4.0]);
6099    /// # Ok::<(), tenferro_ad::Error>(())
6100    /// ```
6101    ///
6102    /// # Errors
6103    ///
6104    /// Returns [`Error::NonScalarGrad`] when this output is not scalar,
6105    /// [`Error::UnsupportedAdRule`] when a graph operation lacks a reverse rule,
6106    /// or a typed validation/backend/runtime-state error during the reverse pass.
6107    pub fn backward(&self) -> Result<Gradients> {
6108        if !self.shape().is_empty() {
6109            return Err(Error::NonScalarGrad {
6110                shape: self.shape().to_vec(),
6111            });
6112        }
6113
6114        let value = self.to_tensor()?;
6115        let seed = self
6116            .ctx
6117            .with_execution_session(|session| one_like_tensor(&value, session))??;
6118        self.backward_from_seed(seed)
6119    }
6120
6121    /// Run reverse-mode AD from this output with an explicit cotangent seed.
6122    ///
6123    /// This is the stateful eager VJP sugar: it returns the cotangent map and
6124    /// accumulates reachable tracked leaves into their `grad()` slots. Use
6125    /// [`EagerRuntime::vjp`] when the VJP result should be returned as a
6126    /// composable eager tensor without touching grad slots.
6127    ///
6128    /// # Examples
6129    ///
6130    /// ```
6131    /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
6132    /// use tenferro_cpu::CpuBackend;
6133    ///
6134    /// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
6135    /// let x = EagerTensor::requires_grad_in(
6136    ///     Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0]).unwrap(),
6137    ///     ctx.clone(),
6138    /// )?;
6139    /// let seed = EagerTensor::from_tensor_in(
6140    ///     Tensor::from_vec_col_major(vec![2], vec![1.0_f64, 2.0]).unwrap(),
6141    ///     ctx,
6142    /// )?;
6143    /// let y = x.runtime().with_eager_session(|s| s.mul(&x, &x))?;
6144    /// y.backward_with(&seed)?;
6145    /// assert_eq!(x.grad()?.unwrap().as_slice::<f64>().unwrap(), &[4.0, 12.0]);
6146    /// # Ok::<(), tenferro_ad::Error>(())
6147    /// ```
6148    ///
6149    /// # Errors
6150    ///
6151    /// Returns [`Error::ContextMismatch`] when `cotangent` belongs to another
6152    /// eager runtime, [`Error::Validation`] when its shape or dtype is not a
6153    /// valid seed, [`Error::UnsupportedAdRule`] for an unavailable reverse
6154    /// rule, or a typed backend/runtime-state error during execution.
6155    pub fn backward_with(&self, cotangent: &EagerTensor) -> Result<Gradients> {
6156        if !self.same_context(cotangent) {
6157            return Err(Error::ContextMismatch {
6158                lhs: self.ctx_id(),
6159                rhs: cotangent.ctx_id(),
6160            });
6161        }
6162        validate_seed_tensor("backward", self, cotangent)?;
6163        self.backward_from_seed(cotangent.to_tensor()?)
6164    }
6165
6166    fn backward_from_seed(&self, seed: Tensor) -> Result<Gradients> {
6167        let cotangent =
6168            EagerTensor::new_result(Arc::clone(&self.ctx), eager_val_key(), seed, false, None)?;
6169        let candidate_keys = {
6170            let mut slots = self.ctx.lock_grad_slots()?;
6171            let mut keys = Vec::new();
6172            slots.retain(|key, slot| {
6173                if slot.upgrade().is_some() {
6174                    keys.push(key.clone());
6175                    true
6176                } else {
6177                    false
6178                }
6179            });
6180            keys
6181        };
6182
6183        let mut targets = Vec::new();
6184        for key in candidate_keys {
6185            let Some(record) = self.ctx.value_record(&key)? else {
6186                continue;
6187            };
6188            if !record.requires_grad {
6189                continue;
6190            }
6191            targets.push((key, EagerTensor::from_record(record)));
6192        }
6193        let wrts = targets.iter().map(|(_, wrt)| wrt).collect::<Vec<_>>();
6194        let gradients = semantic_eager_vjp_many(&self.ctx, self, &wrts, &cotangent)?;
6195        // Shared-handle duplication and gradient storage share one session.
6196        let cotangents = self.ctx.with_execution_session(|session| {
6197            let mut cotangents = HashMap::new();
6198            for ((key, _), grad) in targets.into_iter().zip(gradients) {
6199                let Some(grad) = grad else {
6200                    continue;
6201                };
6202                let tensor = match grad.into_value() {
6203                    Ok(tensor) => tensor,
6204                    Err(IntoValueError::NotUnique(handle)) => {
6205                        handle.duplicate_value_in_session(session)?
6206                    }
6207                    // A gradient whose retained layout is a view (for example a
6208                    // transpose) cannot be extracted as an owned tensor; copy
6209                    // it to a compact tensor instead.
6210                    Err(IntoValueError::Extract { value, .. }) => {
6211                        value.duplicate_value_in_session(session)?
6212                    }
6213                };
6214                cotangents.insert(key, tensor);
6215            }
6216            self.ctx.store_grads(&cotangents, session)?;
6217            Ok::<_, Error>(cotangents)
6218        })??;
6219        Gradients::from_tensors(cotangents)
6220    }
6221}
6222
6223/// Insert a weak registry entry, first dropping dead entries when the insert
6224/// would otherwise grow the table.
6225///
6226/// Dead entries are otherwise removed only when their own key is looked up, so
6227/// a long-running runtime would keep one entry per value it ever created.
6228/// Pruning at the growth point keeps the cost amortized O(1) per insert: a
6229/// sweep runs at most once per capacity doubling.
6230fn insert_pruning_dead<K: std::hash::Hash + Eq, V>(
6231    map: &mut HashMap<K, Weak<V>>,
6232    key: K,
6233    value: Weak<V>,
6234) {
6235    if map.len() == map.capacity() {
6236        map.retain(|_, entry| entry.strong_count() > 0);
6237    }
6238    map.insert(key, value);
6239}
6240
6241pub(crate) fn eager_val_key() -> ValueKey<StdTensorOp> {
6242    ValueKey::Input(next_input_key())
6243}
6244
6245pub(crate) struct RecordedEagerTrace {
6246    pub(crate) key: ValueKey<StdTensorOp>,
6247    pub(crate) trace: Option<EagerTrace>,
6248    pub(crate) requires_grad: bool,
6249}
6250
6251pub(crate) struct RecordedEagerOutputs {
6252    pub(crate) traces: Vec<RecordedEagerTrace>,
6253    pub(crate) semantic_traces: Vec<Option<TracedTensor>>,
6254}
6255
6256pub(crate) fn record_eager_outputs(
6257    op: &StdTensorOp,
6258    outputs: &[&Tensor],
6259    inputs: &[&EagerTensor],
6260) -> Result<RecordedEagerOutputs> {
6261    let output_metadata = outputs
6262        .iter()
6263        .map(|output| tensor_meta_from_tensor(output))
6264        .collect::<Vec<_>>();
6265    record_eager_outputs_inner(op, output_metadata, inputs, None)
6266}
6267
6268pub(crate) fn record_eager_outputs_in_session(
6269    op: &StdTensorOp,
6270    outputs: &[&Tensor],
6271    inputs: &[&EagerTensor],
6272    session: &mut dyn BackendSession,
6273) -> Result<RecordedEagerOutputs> {
6274    let metadata = outputs
6275        .iter()
6276        .map(|output| tensor_meta_from_tensor(output))
6277        .collect();
6278    record_eager_outputs_inner(op, metadata, inputs, Some(session))
6279}
6280
6281pub(crate) fn record_eager_value_outputs_in_session(
6282    op: &StdTensorOp,
6283    outputs: &[&TensorValue],
6284    inputs: &[&EagerTensor],
6285    session: &mut dyn BackendSession,
6286) -> Result<RecordedEagerOutputs> {
6287    let metadata = outputs
6288        .iter()
6289        .map(|output| tensor_meta_from_value(output))
6290        .collect();
6291    record_eager_outputs_inner(op, metadata, inputs, Some(session))
6292}
6293
6294fn record_eager_outputs_inner(
6295    op: &StdTensorOp,
6296    output_metadata: Vec<TensorMeta>,
6297    inputs: &[&EagerTensor],
6298    session: Option<&mut dyn BackendSession>,
6299) -> Result<RecordedEagerOutputs> {
6300    let semantic_traces = record_semantic_eager_outputs(op, &output_metadata, inputs, session)?;
6301    record_eager_outputs_from_metadata(output_metadata, semantic_traces, inputs)
6302}
6303
6304fn record_semantic_eager_outputs(
6305    op: &StdTensorOp,
6306    output_metadata: &[TensorMeta],
6307    inputs: &[&EagerTensor],
6308    mut session: Option<&mut dyn BackendSession>,
6309) -> Result<Vec<Option<TracedTensor>>> {
6310    // Materialize a constant semantic leaf for any untracked input that lost
6311    // its implicit semantic trace on the active-edge fast path. This keeps
6312    // "untracked constant feeds tracked AD" working (PyTorch-style: untracked
6313    // = constant leaf, no gradient flows to it) without re-recording every
6314    // untracked op at creation time.
6315    let mut owned_constants = Vec::<TracedTensor>::new();
6316    for input in inputs {
6317        if input.semantic_trace.is_none() {
6318            let value = match session.as_deref_mut() {
6319                Some(session) => input.duplicate_value_in_session(session)?,
6320                None => input.to_tensor()?,
6321            };
6322            owned_constants.push(TracedTensor::from_tensor_symbolic_shape(value)?);
6323        }
6324    }
6325    let mut constants = owned_constants.iter();
6326    let mut semantic_inputs: Vec<&TracedTensor> = inputs
6327        .iter()
6328        .map(|input| {
6329            input
6330                .semantic_trace
6331                .as_ref()
6332                .unwrap_or_else(|| constants.next().expect("materialized constant"))
6333        })
6334        .collect();
6335    let promotion_plan =
6336        eager_input_promotion_plan(op, inputs.len(), |index| inputs[index].dtype());
6337    // Mirror eager execution in the deferred carrier only. The concrete
6338    // tensors have already been promoted at the execution boundary, so these
6339    // casts add semantic graph nodes without an eager copy or backend kernel.
6340    let promoted_semantic_inputs = if semantic_inputs.iter().enumerate().any(|(index, semantic)| {
6341        semantic.dtype != promotion_plan.target_dtype(index, semantic.dtype)
6342    }) {
6343        Some(
6344            semantic_inputs
6345                .iter()
6346                .enumerate()
6347                .map(|(index, &semantic)| {
6348                    let target = promotion_plan.target_dtype(index, semantic.dtype);
6349                    if semantic.dtype == target {
6350                        Ok(Cow::Borrowed(semantic))
6351                    } else {
6352                        semantic.cast(target).map(Cow::Owned)
6353                    }
6354                })
6355                .collect::<Result<Vec<Cow<'_, TracedTensor>>>>()?,
6356        )
6357    } else {
6358        None
6359    };
6360    if let Some(promoted_semantic_inputs) = &promoted_semantic_inputs {
6361        semantic_inputs = promoted_semantic_inputs.iter().map(Cow::as_ref).collect();
6362    }
6363    let exact_semantic_inputs = if matches!(op, StdTensorOp::Concatenate { .. }) {
6364        Some(
6365            semantic_inputs
6366                .iter()
6367                .zip(inputs)
6368                .map(|(&semantic, input)| {
6369                    if semantic.is_concrete_shape() {
6370                        Ok(semantic.clone())
6371                    } else {
6372                        semantic.reshape(input.shape())
6373                    }
6374                })
6375                .collect::<Result<Vec<_>>>()?,
6376        )
6377    } else {
6378        None
6379    };
6380    if let Some(exact_semantic_inputs) = &exact_semantic_inputs {
6381        semantic_inputs = exact_semantic_inputs.iter().collect();
6382    }
6383    // Deferred materialization (issue #1665 steps 6-7): append only a raw
6384    // carrier. The runtime helper retains metadata scopes introduced by the
6385    // promotion/exactification helpers without analyzing this operation.
6386    let outputs = tenferro_runtime::extension::append_raw_eager_outputs(
6387        op.clone(),
6388        &semantic_inputs,
6389        output_metadata,
6390    )?;
6391    Ok(outputs.into_iter().map(Some).collect())
6392}
6393
6394#[cfg(test)]
6395fn record_eager_graph_outputs(
6396    graph: &Graph<StdTensorOp>,
6397    graph_input_keys: &[TensorInputKey],
6398    outputs: &[Tensor],
6399    inputs: &[&EagerTensor],
6400) -> Result<RecordedEagerOutputs> {
6401    let semantic_traces = record_semantic_eager_graph_outputs(graph, graph_input_keys, inputs)?;
6402    let output_metadata = outputs.iter().map(tensor_meta_from_tensor);
6403    record_eager_outputs_from_metadata(output_metadata, semantic_traces, inputs)
6404}
6405
6406#[cfg(test)]
6407fn record_semantic_eager_graph_outputs(
6408    graph: &Graph<StdTensorOp>,
6409    graph_input_keys: &[TensorInputKey],
6410    inputs: &[&EagerTensor],
6411) -> Result<Vec<Option<TracedTensor>>> {
6412    let Some(semantic_inputs) = inputs
6413        .iter()
6414        .map(|input| input.semantic_trace.as_ref())
6415        .collect::<Option<Vec<_>>>()
6416    else {
6417        return Ok(vec![None; graph.outputs().len()]);
6418    };
6419    if graph_input_keys.len() != semantic_inputs.len() {
6420        return Err(Error::Internal(format!(
6421            "semantic graph recording expected {} input keys, got {}",
6422            semantic_inputs.len(),
6423            graph_input_keys.len()
6424        )));
6425    }
6426
6427    let mut values = HashMap::new();
6428    for (key, tensor) in graph_input_keys.iter().zip(semantic_inputs) {
6429        values.insert(ValueKey::Input(key.clone()), tensor.clone());
6430    }
6431
6432    for op_node in graph.operations() {
6433        let input_values = op_node
6434            .inputs
6435            .iter()
6436            .map(|input| {
6437                let key = match input {
6438                    ValueRef::Local(local_id) => &graph.values()[*local_id].key,
6439                    ValueRef::External(key) => key,
6440                };
6441                values.get(key).cloned().ok_or_else(|| {
6442                    Error::Internal(format!(
6443                        "semantic graph recording missing value for {key:?}"
6444                    ))
6445                })
6446            })
6447            .collect::<Result<Vec<_>>>()?;
6448        let input_refs = input_values.iter().collect::<Vec<_>>();
6449        let semantic_outputs = match &op_node.operation {
6450            StdTensorOp::Extension(ext) => {
6451                tenferro_runtime::extension::apply(Arc::clone(ext), &input_refs)?
6452            }
6453            op => tenferro_runtime::extension::apply_standard_op(op.clone(), &input_refs)?,
6454        };
6455        if semantic_outputs.len() != op_node.outputs.len() {
6456            return Err(Error::Internal(format!(
6457                "semantic graph recording expected {} outputs for {:?}, got {}",
6458                op_node.outputs.len(),
6459                op_node.operation,
6460                semantic_outputs.len()
6461            )));
6462        }
6463        for (output_id, output) in op_node.outputs.iter().copied().zip(semantic_outputs) {
6464            values.insert(graph.values()[output_id].key.clone(), output);
6465        }
6466    }
6467
6468    graph
6469        .outputs()
6470        .iter()
6471        .map(|&output_id| {
6472            let key = &graph.values()[output_id].key;
6473            values.get(key).cloned().map(Some).ok_or_else(|| {
6474                Error::Internal(format!(
6475                    "semantic graph recording missing output for {key:?}"
6476                ))
6477            })
6478        })
6479        .collect()
6480}
6481
6482fn record_eager_outputs_from_metadata(
6483    output_metadata: impl IntoIterator<Item = TensorMeta>,
6484    semantic_traces: Vec<Option<TracedTensor>>,
6485    inputs: &[&EagerTensor],
6486) -> Result<RecordedEagerOutputs> {
6487    let output_metadata = output_metadata.into_iter().collect::<Vec<_>>();
6488    if semantic_traces.len() != output_metadata.len() {
6489        return Err(Error::Internal(format!(
6490            "eager recording expected {} semantic traces, got {}",
6491            output_metadata.len(),
6492            semantic_traces.len()
6493        )));
6494    }
6495    let requires_grad =
6496        eager_grad_recording_enabled() && inputs.iter().any(|input| input.requires_grad);
6497    let trace_count = output_metadata.len();
6498    let residual_trace = EagerTrace::new(inputs);
6499    let traces = (0..trace_count)
6500        .map(|_| RecordedEagerTrace {
6501            key: eager_val_key(),
6502            trace: Some(residual_trace.clone()),
6503            requires_grad,
6504        })
6505        .collect();
6506
6507    Ok(RecordedEagerOutputs {
6508        traces,
6509        semantic_traces,
6510    })
6511}
6512
6513fn tensor_meta_from_value(value: &TensorValue) -> TensorMeta {
6514    TensorMeta::exact(
6515        value.dtype(),
6516        value.shape().iter().copied().map(SymDim::from).collect(),
6517    )
6518}
6519
6520#[cfg(test)]
6521pub(crate) fn zero_like_tensor<B: TensorBackend>(
6522    input: &Tensor,
6523    backend: &mut B,
6524) -> Result<Tensor> {
6525    let host = match input.dtype() {
6526        // A caller-owned payload has no zero-like runtime tensor.
6527        DType::External(type_id) => {
6528            return Err(Error::unsupported(
6529                "zero_like_tensor",
6530                ErrorPhase::GraphBuild,
6531                format!(
6532                    "an externally defined payload ({:?}) has no zero-like runtime tensor",
6533                    DType::External(type_id)
6534                ),
6535            ));
6536        }
6537        DType::F32 => Tensor::from_typed::<f32>(TypedTensor::zeros(input.shape().to_vec())?),
6538        DType::F64 => Tensor::from_typed::<f64>(TypedTensor::zeros(input.shape().to_vec())?),
6539        DType::I32 => Tensor::from_typed::<i32>(TypedTensor::zeros(input.shape().to_vec())?),
6540        DType::I64 => Tensor::from_typed::<i64>(TypedTensor::zeros(input.shape().to_vec())?),
6541        DType::Bool => Tensor::from_typed::<bool>(TypedTensor::from_vec_col_major(
6542            input.shape().to_vec(),
6543            vec![false; input.shape().iter().product()],
6544        )?),
6545        DType::C32 => Tensor::from_typed::<tenferro_tensor::Complex32>(TypedTensor::zeros(
6546            input.shape().to_vec(),
6547        )?),
6548        DType::C64 => Tensor::from_typed::<tenferro_tensor::Complex64>(TypedTensor::zeros(
6549            input.shape().to_vec(),
6550        )?),
6551    };
6552    backend
6553        .upload_host_tensor(TensorRead::from_tensor(&host))
6554        .map_err(Error::from)
6555}
6556
6557pub(crate) fn one_like_tensor(input: &Tensor, session: &mut dyn BackendSession) -> Result<Tensor> {
6558    let host = ones_tensor(input.dtype(), input.shape().to_vec())?;
6559    session
6560        .upload_host_tensor(TensorRead::from_tensor(&host))
6561        .map_err(Error::from)
6562}
6563
6564#[cfg(test)]
6565mod tests;