1use 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#[doc(hidden)]
43#[derive(Clone, Copy, Debug, Eq, PartialEq)]
44pub enum EagerExtensionBackendKind {
45 Cpu,
47 #[cfg(feature = "cuda")]
49 Cuda,
50 #[cfg(feature = "webgpu")]
52 WebGpu,
53}
54
55#[doc(hidden)]
71#[derive(Clone, Debug, Eq, PartialEq)]
72pub struct EagerExtensionTarget {
73 pub engine_id: tenferro_runtime::EngineId,
75 pub backend_kind: EagerExtensionBackendKind,
77}
78
79#[cfg(test)]
80mod tests;
81
82#[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#[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
165pub 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 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
205pub(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
250fn try_prepared_eager_extension(
258 ctx: &EagerRuntime,
259 op: &StdTensorOp,
260 input_reads: &[TensorRead<'_>],
261) -> Result<Option<Vec<Tensor>>> {
262 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 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 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#[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#[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#[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}