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;