use std::sync::Arc;
use computegraph::GraphOperation;
use tenferro_ops::std_tensor_op::StdTensorOp;
use tenferro_runtime::{
Error, ErrorPhase, ExtensionModule, InputSignature, PrepareCapability, PrepareError, Result,
Runtime, RuntimeConfigError,
};
use tenferro_tensor::{Tensor, TensorRead, TensorValue};
use crate::eager::{
eager_capture_active, 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();
if let Some(outputs) = try_prepared_eager_extension(&ctx, &std_op, &input_reads)? {
return finish_eager_extension_outputs(ctx, std_op, inputs, outputs);
}
let outputs = ctx.exec_outputs_read(&std_op, &input_reads)?;
finish_eager_extension_outputs(ctx, std_op, inputs, outputs)
}
fn try_prepared_eager_extension(
ctx: &EagerRuntime,
op: &StdTensorOp,
input_reads: &[TensorRead<'_>],
) -> Result<Option<Vec<Tensor>>> {
let StdTensorOp::Extension(ext) = op else {
return Ok(None);
};
let Ok(target) = ctx.eager_extension_target() else {
return Ok(None);
};
let signature = InputSignature::from_reads(input_reads).map_err(|source| {
Error::runtime_state_source("extension::apply_eager", ErrorPhase::Execution, source)
})?;
let PrepareCapability::Prepared(plan) =
ctx.runtime()
.prepare_extension_immediate(&target.engine_id, ext.as_ref(), &signature)?
else {
return Ok(None);
};
let Some(executor) = plan.executor() else {
return Ok(None);
};
let executor = Arc::clone(executor);
if executor.supports_session() {
let outputs = ctx.with_extension_execution_context(|extension_ctx| {
let (session, caches) = extension_ctx.parts_mut();
executor.execute_in_session(session, caches, input_reads)
})??;
Ok(Some(outputs))
} else {
let outputs = ctx.with_extension_erased_context(|erased, caches| {
executor.execute(erased, caches, input_reads)
})??;
Ok(Some(outputs))
}
}
#[doc(hidden)]
pub fn apply_eager_with_extension_session(
op: Arc<dyn ExtensionOp>,
inputs: &[&EagerTensor],
module: Arc<dyn ExtensionModule>,
) -> Result<Vec<EagerTensor>> {
let ctx = validate_eager_extension_inputs(op.as_ref(), inputs)?;
ctx.install_extension_module(module)?;
apply_eager(op, inputs)
}
#[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>>,
) -> 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)?;
apply_eager(op, inputs)
}
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()
|| (!eager_capture_active() && !inputs.iter().any(|input| input.requires_grad))
{
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()
)));
}
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,
)
} else {
EagerTensor::new_unregistered_result_with_semantic_trace(
Arc::clone(&ctx),
trace.key,
output,
trace.requires_grad,
trace.trace,
semantic_trace,
)
}
})
.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)
}