use std::sync::Arc;
use computegraph::GraphOperation;
use tenferro_ops::std_tensor_op::StdTensorOp;
use tenferro_runtime::ad_support::push_metadata_scope;
use tenferro_runtime::{
Error, ErrorPhase, ExtensionModule, InputSignature, PrepareError, Result, Runtime,
RuntimeConfigError,
};
use tenferro_tensor::{BackendSession, Tensor, TensorRead, TensorValue};
use crate::eager::{eager_grad_recording_enabled, record_eager_outputs, EagerRuntime, EagerTensor};
pub use tenferro_runtime::extension::{
apply, ExtensionCacheKey, ExtensionCacheLimits, ExtensionCacheSelector, ExtensionCacheStore,
ExtensionExecutionContext, ExtensionFamilyId, ExtensionOp,
};
#[doc(hidden)]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum EagerExtensionBackendKind {
Cpu,
#[cfg(feature = "cuda")]
Cuda,
#[cfg(feature = "webgpu")]
WebGpu,
}
#[doc(hidden)]
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct EagerExtensionTarget {
pub engine_id: tenferro_runtime::EngineId,
pub backend_kind: EagerExtensionBackendKind,
}
#[cfg(test)]
mod tests;
#[must_use = "the adopted eager tensor carries the runtime value"]
pub fn adopt_untracked_eager_value(
ctx: Arc<EagerRuntime>,
value: TensorValue,
) -> Result<EagerTensor> {
EagerTensor::new_untracked_value_result(ctx, value)
}
pub fn apply_eager(op: Arc<dyn ExtensionOp>, inputs: &[&EagerTensor]) -> Result<Vec<EagerTensor>> {
let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
let std_op = StdTensorOp::Extension(op);
let input_reads: Vec<_> = inputs.iter().map(|tensor| tensor.tensor_read()).collect();
let outputs = ctx.exec_outputs_read(&std_op, &input_reads)?;
finish_eager_extension_outputs(ctx, std_op, inputs, outputs)
}
pub fn apply_eager_with_extension_session(
op: Arc<dyn ExtensionOp>,
inputs: &[&EagerTensor],
module: Arc<dyn ExtensionModule>,
execute: impl FnOnce(
&dyn ExtensionOp,
&[TensorRead<'_>],
&mut ExtensionExecutionContext<'_, dyn BackendSession + '_>,
) -> tenferro_tensor::Result<Vec<Tensor>>
+ Send,
) -> Result<Vec<EagerTensor>> {
let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
ctx.install_extension_module(module)?;
let input_reads: Vec<_> = inputs.iter().map(|tensor| tensor.tensor_read()).collect();
let outputs = ctx.with_extension_execution_context(|extension_ctx| {
execute(op.as_ref(), &input_reads, extension_ctx)
})??;
finish_eager_extension_outputs(ctx, StdTensorOp::Extension(op), inputs, outputs)
}
#[doc(hidden)]
pub fn apply_eager_with_targeted_extension_session(
op: Arc<dyn ExtensionOp>,
inputs: &[&EagerTensor],
module_factory: impl FnOnce(
EagerExtensionTarget,
) -> tenferro_runtime::Result<Arc<dyn ExtensionModule>>,
execute: impl FnOnce(
&dyn ExtensionOp,
&[TensorRead<'_>],
&mut ExtensionExecutionContext<'_, dyn BackendSession + '_>,
) -> tenferro_tensor::Result<Vec<Tensor>>
+ Send,
) -> Result<Vec<EagerTensor>> {
let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
let target = ctx.eager_extension_target()?;
let input_reads: Vec<_> = inputs.iter().map(|tensor| tensor.tensor_read()).collect();
validate_eager_extension_input_signature(&ctx, &target, &input_reads)?;
let module = module_factory(target.clone())?;
ctx.ensure_extension_module_for_engine(module, op.family_id(), &target.engine_id)?;
let outputs = ctx.with_extension_execution_context(|extension_ctx| {
execute(op.as_ref(), &input_reads, extension_ctx)
})??;
finish_eager_extension_outputs(ctx, StdTensorOp::Extension(op), inputs, outputs)
}
pub(crate) fn validate_eager_extension_target(
runtime: &Runtime,
target: &EagerExtensionTarget,
) -> Result<()> {
let snapshot = runtime.snapshot().map_err(|source| {
Error::runtime_state_source(
"extension::apply_eager_with_extension_session",
ErrorPhase::Execution,
source,
)
})?;
if snapshot.engine(&target.engine_id).is_none() {
return Err(Error::runtime_state_source(
"extension::apply_eager_with_extension_session",
ErrorPhase::Execution,
RuntimeConfigError::MissingEngine {
engine_id: target.engine_id.clone(),
},
));
}
Ok(())
}
fn validate_eager_extension_input_signature(
ctx: &EagerRuntime,
target: &EagerExtensionTarget,
input_reads: &[TensorRead<'_>],
) -> Result<()> {
let signature = InputSignature::from_reads(input_reads).map_err(|source| {
Error::runtime_state_source(
"extension::apply_eager_with_extension_session",
ErrorPhase::Execution,
source,
)
})?;
let snapshot = ctx.runtime().snapshot().map_err(|source| {
Error::runtime_state_source(
"extension::apply_eager_with_extension_session",
ErrorPhase::Execution,
source,
)
})?;
let engine = snapshot.engine(&target.engine_id).ok_or_else(|| {
Error::runtime_state_source(
"extension::apply_eager_with_extension_session",
ErrorPhase::Execution,
RuntimeConfigError::MissingEngine {
engine_id: target.engine_id.clone(),
},
)
})?;
for (input_index, entry) in signature.entries().iter().enumerate() {
if !engine.accepts_input_signature(entry) {
return Err(Error::runtime_state_source(
"extension::apply_eager_with_extension_session",
ErrorPhase::Execution,
PrepareError::NoInputIngress {
input_index,
placement: entry.placement().clone(),
},
));
}
}
Ok(())
}
fn validate_eager_extension_inputs(
op: &dyn ExtensionOp,
inputs: &[&EagerTensor],
) -> Result<Arc<EagerRuntime>> {
let Some(first) = inputs.first() else {
return Err(Error::invalid_argument(
"extension::apply_eager",
ErrorPhase::Execution,
"inputs",
"at least one input tensor is required",
));
};
if inputs.len() != op.input_count() {
return Err(Error::invalid_argument(
"extension::apply_eager",
ErrorPhase::Execution,
"inputs",
format!(
"op family {:?} expects {} inputs, got {}",
op.family_id(),
op.input_count(),
inputs.len()
),
));
}
let ctx = Arc::clone(&first.ctx);
for tensor in inputs.iter().skip(1) {
if !first.same_context(tensor) {
return Err(Error::ContextMismatch {
lhs: first.ctx_id(),
rhs: tensor.ctx_id(),
});
}
}
Ok(ctx)
}
fn finish_eager_extension_outputs(
ctx: Arc<EagerRuntime>,
op: StdTensorOp,
inputs: &[&EagerTensor],
outputs: Vec<Tensor>,
) -> Result<Vec<EagerTensor>> {
if outputs.len() != op.output_count() {
return Err(Error::Internal(format!(
"expected {} eager outputs for {:?}, got {}",
op.output_count(),
op,
outputs.len()
)));
}
if !eager_grad_recording_enabled() {
return outputs
.into_iter()
.map(|output| EagerTensor::new_untracked_result(Arc::clone(&ctx), output))
.collect();
}
let output_refs: Vec<&Tensor> = outputs.iter().collect();
let recorded = record_eager_outputs(&op, &output_refs, inputs)?;
if recorded.traces.len() != outputs.len() {
return Err(Error::Internal(format!(
"expected {} eager traces for {:?}, got {}",
outputs.len(),
op,
recorded.traces.len()
)));
}
let mut metadata_scopes = vec![Arc::clone(&recorded.metadata_scope)];
for input in inputs {
for scope in &input.metadata_scopes {
push_metadata_scope(&mut metadata_scopes, Arc::clone(scope));
}
}
recorded
.traces
.into_iter()
.zip(recorded.semantic_traces)
.zip(outputs)
.map(|((trace, semantic_trace), output)| {
if trace.requires_grad {
EagerTensor::new_result_with_semantic_trace(
Arc::clone(&ctx),
trace.key,
output,
trace.requires_grad,
trace.trace,
semantic_trace,
metadata_scopes.clone(),
)
} else {
EagerTensor::new_unregistered_result_with_semantic_trace(
Arc::clone(&ctx),
trace.key,
output,
trace.requires_grad,
trace.trace,
semantic_trace,
metadata_scopes.clone(),
)
}
})
.collect()
}
pub fn apply_standard_op(op: StdTensorOp, inputs: &[&EagerTensor]) -> Result<EagerTensor> {
if matches!(op, StdTensorOp::Extension(_)) {
return Err(Error::invalid_argument(
"extension::apply_standard_op",
ErrorPhase::Execution,
"op",
"Extension ops must be passed to apply_eager",
));
}
EagerTensor::nary_op(inputs, op)
}