Skip to main content

tenferro_ad/
extension.rs

1//! Eager AD support for out-of-tree extension primitives.
2
3use std::sync::Arc;
4
5use computegraph::GraphOperation;
6use tenferro_ops::std_tensor_op::StdTensorOp;
7use tenferro_runtime::{
8    Error, ErrorPhase, ExtensionModule, InputSignature, PrepareCapability, PrepareError,
9    PreparedOperationExecutorHandle, Result, Runtime, RuntimeConfigError,
10};
11use tenferro_tensor::{Tensor, TensorRead, TensorValue};
12
13use crate::eager::{
14    eager_capture_active, eager_grad_recording_enabled, record_eager_outputs, EagerRuntime,
15    EagerSession, EagerTensor,
16};
17
18pub use tenferro_runtime::extension::{
19    apply, ExtensionCacheKey, ExtensionCacheLimits, ExtensionCacheSelector, ExtensionCacheStore,
20    ExtensionExecutionContext, ExtensionFamilyId, ExtensionOp,
21};
22
23/// Closed backend kind selected by the eager runtime owner for an extension.
24///
25/// # Examples
26///
27/// ```rust
28/// use tenferro_ad::extension::{EagerExtensionBackendKind, EagerExtensionTarget};
29/// use tenferro_runtime::EngineId;
30///
31/// let target = EagerExtensionTarget {
32///     engine_id: EngineId::new("example.engine")?,
33///     backend_kind: EagerExtensionBackendKind::Cpu,
34/// };
35/// assert!(matches!(
36///     target.backend_kind,
37///     EagerExtensionBackendKind::Cpu
38/// ));
39/// assert_eq!(target.engine_id.as_str(), "example.engine");
40/// # Ok::<(), Box<dyn std::error::Error>>(())
41/// ```
42#[doc(hidden)]
43#[derive(Clone, Copy, Debug, Eq, PartialEq)]
44pub enum EagerExtensionBackendKind {
45    /// The eager runtime owns a CPU backend.
46    Cpu,
47    /// The eager runtime owns a CUDA backend.
48    #[cfg(feature = "cuda")]
49    Cuda,
50    /// The eager runtime owns a WebGPU backend.
51    #[cfg(feature = "webgpu")]
52    WebGpu,
53}
54
55/// Exact engine target selected by the eager runtime owner.
56///
57/// # Examples
58///
59/// ```rust
60/// use tenferro_ad::extension::{EagerExtensionBackendKind, EagerExtensionTarget};
61/// use tenferro_runtime::EngineId;
62///
63/// let target = EagerExtensionTarget {
64///     engine_id: EngineId::new("example.engine")?,
65///     backend_kind: EagerExtensionBackendKind::Cpu,
66/// };
67/// assert_eq!(target.backend_kind, EagerExtensionBackendKind::Cpu);
68/// # Ok::<(), Box<dyn std::error::Error>>(())
69/// ```
70#[doc(hidden)]
71#[derive(Clone, Debug, Eq, PartialEq)]
72pub struct EagerExtensionTarget {
73    /// Exact runtime engine selected for this eager context.
74    pub engine_id: tenferro_runtime::EngineId,
75    /// Closed backend kind selected for this eager context.
76    pub backend_kind: EagerExtensionBackendKind,
77}
78
79#[cfg(test)]
80mod tests;
81
82/// Validate recording eligibility before a consuming extension mutation.
83/// This does not grant write authority: callers must still consume the input
84/// through `EagerTensor::into_value` and obtain an exclusive tensor borrow.
85///
86/// # Errors
87/// Returns `Error::RuntimeState` for gradient-tracked values, legacy AD traces,
88/// or active capture, including capture nested inside no_grad. Saved-value
89/// ownership must additionally pass the structural into_value check.
90/// Input-signature validation, module-factory and installation errors are
91/// propagated unchanged.
92///
93/// # Examples
94/// ```
95/// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
96/// use tenferro_ad::extension::prepare_eager_in_place_input;
97/// let input = EagerTensor::requires_grad_in(
98///     Tensor::from_vec_col_major([1], vec![1.0_f64])?, EagerRuntime::new()?)?;
99/// let result = prepare_eager_in_place_input(&input, "example", |_| {
100///     panic!("tracked inputs must be rejected before module construction")
101/// });
102/// assert!(result.is_err());
103/// assert_eq!(input.shape(), &[1]);
104/// # Ok::<(), Box<dyn std::error::Error>>(())
105/// ```
106#[doc(hidden)]
107pub fn prepare_eager_in_place_input(
108    input: &EagerTensor,
109    family_id: &'static str,
110    module_factory: impl FnOnce(EagerExtensionTarget) -> Result<Arc<dyn ExtensionModule>>,
111) -> Result<()> {
112    if input.requires_grad || input.trace.is_some() || eager_capture_active() {
113        return Err(Error::runtime_state(
114            "eager in-place",
115            ErrorPhase::Execution,
116            "in-place execution requires an untracked value outside trace capture",
117        ));
118    }
119    let target = input.ctx.eager_extension_target()?;
120    validate_eager_extension_input_signature(&input.ctx, &target, &[input.tensor_read()])?;
121    let module = module_factory(target.clone())?;
122    input
123        .ctx
124        .ensure_extension_module_for_engine(module, family_id, &target.engine_id)?;
125    Ok(())
126}
127
128/// Adopt an untracked eager tensor value produced by this runtime's backend.
129///
130/// This is a low-level extension contract for eager composite operations that
131/// execute through a lifetime-bound backend session and receive a lazy
132/// [`TensorValue`] from the backend. The value must have been produced for the
133/// same eager runtime; this helper intentionally does not register gradient
134/// metadata and must not be used for tracked outputs.
135///
136/// # Examples
137///
138/// ```rust
139/// use tenferro_ad::extension::adopt_untracked_eager_value;
140/// use tenferro_ad::EagerRuntime;
141/// use tenferro_cpu::CpuBackend;
142/// use tenferro_tensor::{Tensor, TensorValue};
143///
144/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
145/// let value = TensorValue::from_tensor(
146///     Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(),
147/// );
148/// let eager = adopt_untracked_eager_value(ctx, value)?;
149/// assert_eq!(eager.shape(), &[1]);
150/// assert!(!eager.tracks_grad());
151/// # Ok::<(), tenferro_ad::Error>(())
152/// ```
153/// # Errors
154///
155/// Returns [`Error::RuntimeState`] when the value cannot be registered in the
156/// supplied runtime, including an invalid or incompatible retained descriptor.
157#[must_use = "the adopted eager tensor carries the runtime value"]
158pub fn adopt_untracked_eager_value(
159    ctx: Arc<EagerRuntime>,
160    value: TensorValue,
161) -> Result<EagerTensor> {
162    EagerTensor::new_untracked_value_result(ctx, value)
163}
164
165/// Apply an extension op to eager AD tensors.
166///
167/// # Examples
168///
169/// ```rust
170/// use tenferro_ad::extension::apply_eager;
171/// use tenferro_ad::{EagerRuntime, EagerTensor};
172/// use tenferro_cpu::CpuBackend;
173/// use tenferro_tensor::Tensor;
174///
175/// let ctx = EagerRuntime::with_cpu_backend(CpuBackend::new())?;
176/// let x = EagerTensor::from_tensor_in(
177///     Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(),
178///     ctx,
179/// ).unwrap();
180/// let _ = &x;
181/// let _apply = apply_eager;
182/// # Ok::<(), tenferro_ad::Error>(())
183/// ```
184/// # Errors
185///
186/// Returns `Error::Validation` with `InvalidArgument` when `inputs` is empty
187/// or its length differs from the extension's declared input count. Returns
188/// `Error::ContextMismatch` when tensors belong to different eager runtimes;
189/// backend, extension, and runtime-state failures retain their typed sources.
190pub fn apply_eager(op: Arc<dyn ExtensionOp>, inputs: &[&EagerTensor]) -> Result<Vec<EagerTensor>> {
191    let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
192    let std_op = StdTensorOp::Extension(Arc::clone(&op));
193    let input_reads: Vec<_> = inputs.iter().map(|tensor| tensor.tensor_read()).collect();
194    // Native immediate path: resolve the extension engine from the runtime
195    // snapshot, prepare the op, and execute through the prepared plan's
196    // scheduler-session executor when it is session-capable. This skips the
197    // SemanticProgram build/compile + run_compiled cost on every call.
198    if let Some(outputs) = try_prepared_eager_extension(&ctx, &std_op, &input_reads)? {
199        return finish_eager_extension_outputs(ctx, std_op, inputs, outputs, None);
200    }
201    let outputs = ctx.exec_extension_outputs_read(&op, &input_reads)?;
202    finish_eager_extension_outputs(ctx, std_op, inputs, outputs, None)
203}
204
205/// Execute a prepared eager extension on the caller's borrowed session.
206///
207/// The exact eager runtime, input signature, and selected extension engine are
208/// checked before dispatch. A legacy context-only executor requires a separate
209/// top-level call to [`apply_eager`] after this session is released; this
210/// function never re-enters the eager backend or silently opens that fallback.
211///
212/// # Errors
213///
214/// Returns a typed context mismatch, validation, unsupported capability, or
215/// backend/extension error without replacing its source.
216pub(crate) fn apply_eager_in_session(
217    session: &mut EagerSession<'_>,
218    op: Arc<dyn ExtensionOp>,
219    inputs: &[&EagerTensor],
220) -> Result<Vec<EagerTensor>> {
221    let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
222    if !Arc::ptr_eq(session.runtime(), &ctx) {
223        return Err(Error::ContextMismatch {
224            lhs: session.runtime().id(),
225            rhs: ctx.id(),
226        });
227    }
228    let std_op = StdTensorOp::Extension(Arc::clone(&op));
229    let input_reads: Vec<_> = inputs.iter().map(|tensor| tensor.tensor_read()).collect();
230    let target = ctx.eager_extension_target()?;
231    let executor = prepared_eager_extension_executor(&ctx, &target, &std_op, &input_reads)?
232        .ok_or_else(|| {
233            Error::unsupported(
234                "extension::apply_eager_in_session",
235                ErrorPhase::Execution,
236                "no session-capable prepared extension executor for this signature",
237            )
238        })?;
239    if !executor.supports_session() {
240        return Err(Error::unsupported(
241            "extension::apply_eager_in_session",
242            ErrorPhase::Execution,
243            "the native-context executor requires a separate top-level runtime region",
244        ));
245    }
246    let outputs = session.execute_prepared_extension(executor.as_ref(), &input_reads)?;
247    finish_eager_extension_outputs(ctx, std_op, inputs, outputs, Some(session))
248}
249
250/// Run one extension op through the snapshot-resolved native prepared path.
251///
252/// Returns `None` (so the caller falls back to the compiled-program path for
253/// the exact op and signature) when the eager runtime has no exact extension
254/// engine, the engine cannot prepare the op, or the prepared plan has no
255/// executor. Context-only executors use the native-context bridge in their own
256/// top-level region. AD recording remains with the caller.
257fn try_prepared_eager_extension(
258    ctx: &EagerRuntime,
259    op: &StdTensorOp,
260    input_reads: &[TensorRead<'_>],
261) -> Result<Option<Vec<Tensor>>> {
262    // The recording test backend has no extension engine; its existing
263    // top-level route still falls back to compiled execution.
264    let Ok(target) = ctx.eager_extension_target() else {
265        return Ok(None);
266    };
267    let Some(executor) = prepared_eager_extension_executor(ctx, &target, op, input_reads)? else {
268        return Ok(None);
269    };
270    if executor.supports_session() {
271        // native-session: scheduler-owned session executor.
272        let outputs = ctx.with_extension_execution_context(|extension_ctx| {
273            let (session, caches) = extension_ctx.parts_mut();
274            executor.execute_in_session(session, caches, input_reads)
275        })??;
276        Ok(Some(outputs))
277    } else {
278        // native-context: mandatory `execute` bridge over the erased backend
279        // context (for out-of-tree prepared ops without a session executor).
280        let outputs = ctx.with_extension_erased_context(|erased, caches| {
281            executor.execute(erased, caches, input_reads)
282        })??;
283        Ok(Some(outputs))
284    }
285}
286
287fn prepared_eager_extension_executor(
288    ctx: &EagerRuntime,
289    target: &EagerExtensionTarget,
290    op: &StdTensorOp,
291    input_reads: &[TensorRead<'_>],
292) -> Result<Option<PreparedOperationExecutorHandle>> {
293    let StdTensorOp::Extension(ext) = op else {
294        return Ok(None);
295    };
296    let signature = InputSignature::from_reads(input_reads).map_err(|source| {
297        Error::runtime_state_source("extension::apply_eager", ErrorPhase::Execution, source)
298    })?;
299    let PrepareCapability::Prepared(plan) =
300        ctx.runtime()
301            .prepare_extension_immediate(&target.engine_id, ext.as_ref(), &signature)?
302    else {
303        return Ok(None);
304    };
305    Ok(plan.executor().cloned())
306}
307
308/// Ensure an eager extension module is installed, then apply the op through
309/// the single [`apply_eager`] entry.
310///
311/// This thin wrapper is retained for module-owner eager call sites (linalg,
312/// einsum). Forward execution always routes through [`apply_eager`]'s native
313/// prepared path; this wrapper only owns the install-ensure step.
314///
315/// # Errors
316///
317/// Returns `Error::Validation` with `InvalidArgument` when `inputs` is empty
318/// or its length differs from the extension's declared input count. Returns
319/// `Error::ContextMismatch` when tensors belong to different eager runtimes;
320/// backend, extension, and runtime-state failures retain their typed sources.
321#[doc(hidden)]
322pub fn apply_eager_with_extension_session(
323    op: Arc<dyn ExtensionOp>,
324    inputs: &[&EagerTensor],
325    module: Arc<dyn ExtensionModule>,
326) -> Result<Vec<EagerTensor>> {
327    let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
328    ctx.install_extension_module(module)?;
329    apply_eager(op, inputs)
330}
331
332/// Apply an eager extension through the owner-selected engine and backend kind.
333///
334/// This narrow sibling-crate wrapper is used by FFT, whose module factory must
335/// follow the eager runtime's exact backend selection. Input, target, and
336/// ingress validation always run before `module_factory`; errors returned by
337/// the factory are propagated unchanged. The returned module is then passed to
338/// the owner-scoped ensure operation. Forward execution always routes through
339/// [`apply_eager`]'s native prepared path.
340///
341/// # Errors
342///
343/// Returns [`tenferro_runtime::Error::Validation`] with
344/// [`tenferro_tensor::ValidationError::InvalidArgument`] when `inputs` is empty
345/// or its length differs from the extension's declared input count. Returns
346/// [`tenferro_runtime::Error::ContextMismatch`] when tensors belong to
347/// different eager runtimes.
348///
349/// Returns [`tenferro_runtime::Error::RuntimeStateSource`] when the selected
350/// engine is missing through
351/// [`tenferro_runtime::RuntimeConfigError::MissingEngine`], an input has no
352/// ingress through [`tenferro_runtime::PrepareError::NoInputIngress`], or the
353/// cold/missing-registration ensure path rejects the module. Errors returned by
354/// `module_factory` are propagated unchanged. Session, cache, and output
355/// registration failures retain their typed runtime sources.
356#[doc(hidden)]
357pub fn apply_eager_with_targeted_extension_session(
358    op: Arc<dyn ExtensionOp>,
359    inputs: &[&EagerTensor],
360    module_factory: impl FnOnce(
361        EagerExtensionTarget,
362    ) -> tenferro_runtime::Result<Arc<dyn ExtensionModule>>,
363) -> Result<Vec<EagerTensor>> {
364    let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
365    let target = ctx.eager_extension_target()?;
366    let input_reads: Vec<_> = inputs.iter().map(|tensor| tensor.tensor_read()).collect();
367    validate_eager_extension_input_signature(&ctx, &target, &input_reads)?;
368    let module = module_factory(target.clone())?;
369    ctx.ensure_extension_module_for_engine(module, op.family_id(), &target.engine_id)?;
370    apply_eager(op, inputs)
371}
372
373/// Install the exact eager-owner extension and execute it on a borrowed session.
374/// Input and selected-engine ingress checks precede module construction; a
375/// context-only executor returns a typed unsupported error rather than nesting
376/// an erased-context entry inside the active session.
377///
378/// # Errors
379///
380/// Returns typed context, validation, module-installation, unsupported-executor,
381/// or backend errors from the corresponding boundary.
382#[doc(hidden)]
383pub fn apply_eager_with_targeted_extension_in_session(
384    session: &mut EagerSession<'_>,
385    op: Arc<dyn ExtensionOp>,
386    inputs: &[&EagerTensor],
387    module_factory: impl FnOnce(
388        EagerExtensionTarget,
389    ) -> tenferro_runtime::Result<Arc<dyn ExtensionModule>>,
390) -> Result<Vec<EagerTensor>> {
391    let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
392    if !Arc::ptr_eq(session.runtime(), &ctx) {
393        return Err(Error::ContextMismatch {
394            lhs: session.runtime().id(),
395            rhs: ctx.id(),
396        });
397    }
398    let target = ctx.eager_extension_target()?;
399    let input_reads: Vec<_> = inputs.iter().map(|tensor| tensor.tensor_read()).collect();
400    validate_eager_extension_input_signature(&ctx, &target, &input_reads)?;
401    let module = module_factory(target.clone())?;
402    ctx.ensure_extension_module_for_engine(module, op.family_id(), &target.engine_id)?;
403    apply_eager_in_session(session, op, inputs)
404}
405
406pub(crate) fn validate_eager_extension_target(
407    runtime: &Runtime,
408    target: &EagerExtensionTarget,
409) -> Result<()> {
410    let snapshot = runtime.snapshot().map_err(|source| {
411        Error::runtime_state_source(
412            "extension::apply_eager_with_extension_session",
413            ErrorPhase::Execution,
414            source,
415        )
416    })?;
417    if snapshot.engine(&target.engine_id).is_none() {
418        return Err(Error::runtime_state_source(
419            "extension::apply_eager_with_extension_session",
420            ErrorPhase::Execution,
421            RuntimeConfigError::MissingEngine {
422                engine_id: target.engine_id.clone(),
423            },
424        ));
425    }
426    Ok(())
427}
428
429fn validate_eager_extension_input_signature(
430    ctx: &EagerRuntime,
431    target: &EagerExtensionTarget,
432    input_reads: &[TensorRead<'_>],
433) -> Result<()> {
434    let signature = InputSignature::from_reads(input_reads).map_err(|source| {
435        Error::runtime_state_source(
436            "extension::apply_eager_with_extension_session",
437            ErrorPhase::Execution,
438            source,
439        )
440    })?;
441    let snapshot = ctx.runtime().snapshot().map_err(|source| {
442        Error::runtime_state_source(
443            "extension::apply_eager_with_extension_session",
444            ErrorPhase::Execution,
445            source,
446        )
447    })?;
448    let engine = snapshot.engine(&target.engine_id).ok_or_else(|| {
449        Error::runtime_state_source(
450            "extension::apply_eager_with_extension_session",
451            ErrorPhase::Execution,
452            RuntimeConfigError::MissingEngine {
453                engine_id: target.engine_id.clone(),
454            },
455        )
456    })?;
457    for (input_index, entry) in signature.entries().iter().enumerate() {
458        if !engine.accepts_input_signature(entry) {
459            return Err(Error::runtime_state_source(
460                "extension::apply_eager_with_extension_session",
461                ErrorPhase::Execution,
462                PrepareError::NoInputIngress {
463                    input_index,
464                    placement: entry.placement().clone(),
465                },
466            ));
467        }
468    }
469    Ok(())
470}
471
472fn validate_eager_extension_inputs(
473    op: &dyn ExtensionOp,
474    inputs: &[&EagerTensor],
475) -> Result<Arc<EagerRuntime>> {
476    let Some(first) = inputs.first() else {
477        return Err(Error::invalid_argument(
478            "extension::apply_eager",
479            ErrorPhase::Execution,
480            "inputs",
481            "at least one input tensor is required",
482        ));
483    };
484    if inputs.len() != op.input_count() {
485        return Err(Error::invalid_argument(
486            "extension::apply_eager",
487            ErrorPhase::Execution,
488            "inputs",
489            format!(
490                "op family {:?} expects {} inputs, got {}",
491                op.family_id(),
492                op.input_count(),
493                inputs.len()
494            ),
495        ));
496    }
497
498    let ctx = Arc::clone(&first.ctx);
499    for tensor in inputs.iter().skip(1) {
500        if !first.same_context(tensor) {
501            return Err(Error::ContextMismatch {
502                lhs: first.ctx_id(),
503                rhs: tensor.ctx_id(),
504            });
505        }
506    }
507    Ok(ctx)
508}
509
510fn finish_eager_extension_outputs(
511    ctx: Arc<EagerRuntime>,
512    op: StdTensorOp,
513    inputs: &[&EagerTensor],
514    outputs: Vec<Tensor>,
515    session: Option<&mut EagerSession<'_>>,
516) -> Result<Vec<EagerTensor>> {
517    if outputs.len() != op.output_count() {
518        return Err(Error::Internal(format!(
519            "expected {} eager outputs for {:?}, got {}",
520            op.output_count(),
521            op,
522            outputs.len()
523        )));
524    }
525
526    if !eager_grad_recording_enabled()
527        || (!eager_capture_active() && !inputs.iter().any(|input| input.requires_grad))
528    {
529        return outputs
530            .into_iter()
531            .map(|output| EagerTensor::new_untracked_result(Arc::clone(&ctx), output))
532            .collect();
533    }
534
535    let output_refs: Vec<&Tensor> = outputs.iter().collect();
536    let recorded = match session {
537        Some(session) => session.record_outputs(&op, &output_refs, inputs)?,
538        None => record_eager_outputs(&op, &output_refs, inputs)?,
539    };
540    if recorded.traces.len() != outputs.len() {
541        return Err(Error::Internal(format!(
542            "expected {} eager traces for {:?}, got {}",
543            outputs.len(),
544            op,
545            recorded.traces.len()
546        )));
547    }
548    let results = recorded
549        .traces
550        .into_iter()
551        .zip(recorded.semantic_traces)
552        .zip(outputs)
553        .map(|((trace, semantic_trace), output)| {
554            if trace.requires_grad {
555                EagerTensor::new_result_with_semantic_trace(
556                    Arc::clone(&ctx),
557                    trace.key,
558                    output,
559                    trace.requires_grad,
560                    trace.trace,
561                    semantic_trace,
562                )
563            } else {
564                EagerTensor::new_unregistered_result_with_semantic_trace(
565                    Arc::clone(&ctx),
566                    trace.key,
567                    output,
568                    trace.requires_grad,
569                    trace.trace,
570                    semantic_trace,
571                )
572            }
573        })
574        .collect::<Result<Vec<_>>>()?;
575    crate::eager::finish_residuals(&op, inputs, &results.iter().collect::<Vec<_>>())?;
576    Ok(results)
577}