1use num_complex::{Complex32, Complex64};
7use std::collections::{HashMap, HashSet};
8use std::error::Error as StdError;
9use std::fmt;
10use std::panic::{catch_unwind, AssertUnwindSafe};
11use std::sync::{Arc, Condvar, Mutex, MutexGuard};
12use std::thread::{self, ThreadId};
13
14use smallvec::SmallVec;
15use tenferro_ops::shape_extent::ShapeExtent;
16use tenferro_tensor::{
17 AllocationGroup, DType, DescriptorSlot, GroupError, MemoryKind, Tensor, TensorBackend,
18 TensorRead, TensorValue, TensorView,
19};
20
21use crate::error::ErrorPhase;
22use crate::exec::{ExecInstruction, ExecProgram, ExecSlot, ExtensionExecutionDispatch};
23use crate::extension_cache::{ExtensionCacheSelector, ExtensionCacheStore};
24use crate::graph::CompiledGraph;
25use crate::runtime::schedule::{
26 EventDependency, ExecutionLocation, ScheduledGraph, ScheduledNode, ScheduledNodeKind,
27 ScheduledTransfer, UnsupportedScheduledNodeError,
28};
29use crate::runtime::{
30 CacheOwnerError, CacheStats, EventDomainError, EventDomainOperation, EventDomainRun,
31 EventToken, InputSignature, PrepareError, PrepareOptions, PreparedOperationPlan, Runtime,
32 RuntimeCacheOwner, SubmissionError, TransferError, TransferProviderContractError,
33 TransferRequest,
34};
35use crate::{Error, Result};
36
37type RuntimeInputReads<'a> = SmallVec<[TensorRead<'a>; 8]>;
38type RuntimeInputShapes<'a> = SmallVec<[&'a [usize]; 8]>;
39type RuntimeShapeScratch = SmallVec<[usize; 8]>;
40
41#[derive(Clone, Copy, Debug)]
42pub(super) enum RuntimeOutputMode {
43 Tensor,
44 Value,
45}
46
47#[derive(Clone)]
53pub struct PreparedCompiledGraph {
54 runtime_id: super::RuntimeId,
55 epoch: super::RuntimeEpoch,
56 program: CompiledGraph,
57 prepared: Arc<super::preparation::PreparedProgram>,
58}
59
60impl PreparedCompiledGraph {
61 #[doc(hidden)]
64 #[must_use]
65 pub fn elementwise_region_summary(&self) -> (usize, usize) {
66 super::region::region_summary(self.prepared.root().regions())
67 }
68
69 #[doc(hidden)]
72 #[must_use]
73 pub fn elementwise_region_execution_counts(&self) -> (usize, usize) {
74 let counters = self.prepared.root().region_counters();
75 (counters.fused(), counters.fallbacks())
76 }
77
78 #[doc(hidden)]
83 #[must_use]
84 pub fn execution_submission_count(&self) -> usize {
85 self.prepared.root().region_counters().submissions()
86 }
87
88 #[doc(hidden)]
94 #[must_use]
95 pub fn execution_command_count(&self) -> usize {
96 let root = self.prepared.root();
97 super::region::command_count(root.schedule(), root.regions())
98 }
99}
100
101pub struct ExecutionHandle {
103 submission: Arc<InFlightSubmission>,
104}
105
106pub struct ExecutionInputs {
118 group: AllocationGroup,
119 bindings: Box<[DescriptorSlot]>,
120}
121
122impl ExecutionInputs {
123 pub fn new(tensors: Vec<Tensor>) -> Result<Self> {
130 let (group, bindings) = AllocationGroup::from_tensors(tensors).map_err(|error| {
131 Error::runtime_state(
132 "ExecutionInputs::new",
133 ErrorPhase::Execution,
134 error.to_string(),
135 )
136 })?;
137 Ok(Self { group, bindings })
138 }
139
140 pub(crate) fn as_reads(&self) -> Result<Vec<TensorRead<'_>>> {
141 self.group.read_views(&self.bindings).map_err(|error| {
142 Error::runtime_state(
143 "ExecutionInputs::as_reads",
144 ErrorPhase::Execution,
145 error.to_string(),
146 )
147 })
148 }
149}
150
151impl fmt::Debug for ExecutionInputs {
152 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
153 formatter
154 .debug_struct("ExecutionInputs")
155 .field("len", &self.bindings.len())
156 .finish()
157 }
158}
159
160pub struct ScopedReadInputs<'env> {
162 bindings: Box<[ScopedReadBinding<'env>]>,
163}
164
165impl<'env> ScopedReadInputs<'env> {
166 pub fn new(bindings: Vec<TensorView<'env>>) -> Self {
168 Self {
169 bindings: bindings
170 .into_iter()
171 .map(|tensor| ScopedReadBinding { tensor })
172 .collect(),
173 }
174 }
175
176 fn as_reads(&self) -> Vec<TensorRead<'env>> {
177 self.bindings
178 .iter()
179 .map(|binding| TensorRead::from_view(binding.tensor.clone()))
180 .collect()
181 }
182
183 fn has_non_host_provider(&self) -> bool {
184 self.bindings.iter().any(|binding| {
185 !matches!(
186 binding.tensor.placement().memory_kind,
187 MemoryKind::PinnedHost | MemoryKind::UnpinnedHost
188 )
189 })
190 }
191
192 pub fn len(&self) -> usize {
194 self.bindings.len()
195 }
196
197 pub fn is_empty(&self) -> bool {
199 self.bindings.is_empty()
200 }
201}
202
203pub struct ScopedReadBinding<'env> {
205 tensor: TensorView<'env>,
206}
207
208impl<'env> ScopedReadBinding<'env> {
209 pub fn new(tensor: TensorView<'env>) -> Self {
211 Self { tensor }
212 }
213
214 pub fn tensor(&self) -> &TensorView<'env> {
216 &self.tensor
217 }
218}
219
220impl fmt::Debug for ScopedReadInputs<'_> {
221 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
222 formatter
223 .debug_struct("ScopedReadInputs")
224 .field("len", &self.bindings.len())
225 .finish()
226 }
227}
228
229#[allow(clippy::large_enum_variant)]
233#[derive(Debug)]
234pub enum ScopedExecutionOutcome<'env> {
235 Completed(ScopedExecutionBundle<'env>),
237 RetiredFailed {
239 error: Error,
240 inputs: ScopedReadInputs<'env>,
241 },
242}
243
244#[derive(Clone, Debug, PartialEq, Eq)]
246pub struct OutputMetadata {
247 pub dtype: DType,
249 pub shape: Box<[usize]>,
251}
252
253pub struct ScopedExecutionBundle<'env> {
255 owned: AllocationGroup,
256 outputs: Box<[ScopedOutput<'env>]>,
257}
258
259impl fmt::Debug for ScopedExecutionBundle<'_> {
260 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
261 formatter
262 .debug_struct("ScopedExecutionBundle")
263 .field("len", &self.outputs.len())
264 .finish()
265 }
266}
267
268impl<'env> ScopedExecutionBundle<'env> {
269 pub fn output(&self, index: usize) -> std::result::Result<OutputRef<'_>, OutputAccessError> {
277 let output = self
278 .outputs
279 .get(index)
280 .ok_or(OutputAccessError::InvalidOutput { index })?;
281 match output {
282 ScopedOutput::Borrowed(view) => Ok(OutputRef::Tensor(view.clone())),
283 ScopedOutput::Owned(slot) => output_ref_from_group(&self.owned, *slot),
284 ScopedOutput::Metadata(metadata) => Ok(OutputRef::Metadata(metadata)),
285 }
286 }
287
288 #[allow(clippy::result_large_err)]
298 pub fn into_owned_output(
299 self,
300 index: usize,
301 ) -> std::result::Result<Tensor, (Self, ScopedOutputExtractError)> {
302 let slot = match self.outputs.get(index) {
303 Some(ScopedOutput::Owned(slot)) => *slot,
304 Some(ScopedOutput::Borrowed(_)) => {
305 return Err((self, ScopedOutputExtractError::BorrowedOutput))
306 }
307 Some(ScopedOutput::Metadata(_)) => {
308 return Err((self, ScopedOutputExtractError::MetadataOutput))
309 }
310 None => return Err((self, ScopedOutputExtractError::InvalidOutput { index })),
311 };
312 let Self { owned, outputs } = self;
313 match owned.into_tensor(slot) {
314 Ok(tensor) => Ok(tensor),
315 Err((owned, error)) => Err((
316 Self { owned, outputs },
317 ScopedOutputExtractError::Output(OutputExtractError::Group(error)),
318 )),
319 }
320 }
321}
322
323#[allow(clippy::large_enum_variant)]
327#[derive(Debug)]
328pub enum ScopedOutput<'env> {
329 Borrowed(TensorView<'env>),
331 Owned(DescriptorSlot),
333 Metadata(OutputMetadata),
335}
336
337#[derive(Debug, thiserror::Error)]
339pub enum ScopedOutputExtractError {
340 #[error("scoped output index {index} is outside the output set")]
341 InvalidOutput { index: usize },
342 #[error("scoped output is borrowed from the caller")]
343 BorrowedOutput,
344 #[error("scoped output contains metadata only")]
345 MetadataOutput,
346 #[error("scoped owned output cannot be extracted: {0}")]
347 Output(#[from] OutputExtractError),
348}
349
350#[derive(Debug)]
352pub struct ScopedSubmitRejected<'env> {
353 source: Box<Error>,
354 inputs: ScopedReadInputs<'env>,
355}
356
357impl<'env> ScopedSubmitRejected<'env> {
358 pub fn into_parts(self) -> (Error, ScopedReadInputs<'env>) {
360 (*self.source, self.inputs)
361 }
362}
363
364impl fmt::Display for ScopedSubmitRejected<'_> {
365 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
366 write!(
367 formatter,
368 "scoped submission rejected before admission: {}",
369 self.source
370 )
371 }
372}
373
374impl StdError for ScopedSubmitRejected<'_> {
375 fn source(&self) -> Option<&(dyn StdError + 'static)> {
376 Some(self.source.as_ref())
377 }
378}
379
380#[derive(Debug)]
382pub struct ExecutionBundle {
383 group: AllocationGroup,
384 outputs: Box<[DescriptorSlot]>,
385}
386
387#[allow(clippy::large_enum_variant)]
391#[derive(Debug)]
392pub enum OutputRef<'a> {
393 Tensor(TensorView<'a>),
394 Metadata(&'a OutputMetadata),
395}
396
397#[derive(Debug, thiserror::Error)]
398pub enum OutputAccessError {
399 #[error("execution output index {index} is outside the output set")]
400 InvalidOutput { index: usize },
401 #[error("execution output group is invalid: {0}")]
402 Group(#[from] GroupError),
403}
404
405#[derive(Debug, thiserror::Error)]
406pub enum OutputExtractError {
407 #[error("execution output index {index} is outside the output set")]
408 InvalidOutput { index: usize },
409 #[error("execution output cannot be extracted: {0}")]
410 Group(#[from] GroupError),
411}
412
413impl ExecutionBundle {
414 fn from_inputs_and_outputs(inputs: ExecutionInputs, outputs: Vec<Tensor>) -> Result<Self> {
415 let mut group = inputs.group;
416 let mut output_slots = Vec::with_capacity(outputs.len());
417 for output in outputs {
418 let slot = group.append_tensor(output).map_err(|error| {
419 Error::runtime_state(
420 "ExecutionBundle::from_inputs_and_outputs",
421 ErrorPhase::Execution,
422 error.to_string(),
423 )
424 })?;
425 output_slots.push(slot);
426 }
427 Ok(Self {
428 group,
429 outputs: output_slots.into_boxed_slice(),
430 })
431 }
432
433 #[cfg(test)]
434 pub(super) fn from_outputs(outputs: Vec<Tensor>) -> Result<Self> {
435 let (group, bindings) = AllocationGroup::from_tensors(outputs).map_err(|error| {
436 Error::runtime_state(
437 "ExecutionBundle::from_outputs",
438 ErrorPhase::Execution,
439 error.to_string(),
440 )
441 })?;
442 Ok(Self {
443 group,
444 outputs: bindings,
445 })
446 }
447
448 pub fn output(&self, index: usize) -> std::result::Result<OutputRef<'_>, OutputAccessError> {
454 let slot = *self
455 .outputs
456 .get(index)
457 .ok_or(OutputAccessError::InvalidOutput { index })?;
458 let mut views = self.group.read_views(std::slice::from_ref(&slot))?;
459 let view = views
460 .pop()
461 .ok_or(OutputAccessError::InvalidOutput { index })?;
462 match view {
463 TensorRead::View(view) => Ok(OutputRef::Tensor(view)),
464 TensorRead::Tensor(_) => Err(OutputAccessError::InvalidOutput { index }),
465 }
466 }
467
468 #[allow(clippy::result_large_err)]
477 pub fn into_output(
478 self,
479 index: usize,
480 ) -> std::result::Result<Tensor, (Self, OutputExtractError)> {
481 let slot = match self.outputs.get(index).copied() {
482 Some(slot) => slot,
483 None => return Err((self, OutputExtractError::InvalidOutput { index })),
484 };
485 let ExecutionBundle { group, outputs } = self;
486 match group.into_tensor(slot) {
487 Ok(tensor) => Ok(tensor),
488 Err((group, error)) => Err((Self { group, outputs }, OutputExtractError::Group(error))),
489 }
490 }
491}
492
493#[derive(Debug)]
495pub enum ExecutionOutcome {
496 Completed(Box<ExecutionBundle>),
498 RetiredFailed {
501 error: Error,
502 inputs: Box<ExecutionInputs>,
503 },
504 CompletionUnproven {
507 error: Error,
508 diagnostic_keys: Box<[String]>,
509 },
510}
511
512#[derive(Debug)]
528pub enum SubmitError {
529 PreAdmission {
532 source: Box<Error>,
533 inputs: Box<ExecutionInputs>,
534 },
535}
536
537impl SubmitError {
538 pub fn into_pre_admission(self) -> Option<(Error, ExecutionInputs)> {
553 match self {
554 Self::PreAdmission { source, inputs } => Some((*source, *inputs)),
555 }
556 }
557}
558
559impl fmt::Display for SubmitError {
560 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
561 match self {
562 Self::PreAdmission { source, .. } => {
563 write!(formatter, "submission rejected before admission: {source}")
564 }
565 }
566 }
567}
568
569impl StdError for SubmitError {
570 fn source(&self) -> Option<&(dyn StdError + 'static)> {
571 match self {
572 Self::PreAdmission { source, .. } => {
573 source.source().or(Some(source.as_ref()))
577 }
578 }
579 }
580}
581
582impl From<SubmitError> for Error {
583 fn from(error: SubmitError) -> Self {
584 match error {
585 SubmitError::PreAdmission { source, .. } => *source,
586 }
587 }
588}
589
590impl ExecutionHandle {
591 pub fn wait(self) -> Result<ExecutionOutcome> {
599 self.submission.wait()
600 }
601}
602
603impl fmt::Debug for ExecutionHandle {
604 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
605 formatter
606 .debug_struct("ExecutionHandle")
607 .field("pending", &self.submission.is_pending())
608 .finish_non_exhaustive()
609 }
610}
611
612pub(super) struct InFlightSubmission {
613 work: Mutex<Option<InFlightWork>>,
614 completion: Mutex<Option<Result<ExecutionOutcome>>>,
615 completed: Condvar,
616}
617
618struct AdmittedExecution {
619 prepared: PreparedCompiledGraph,
620 inputs: ExecutionInputs,
621}
622
623enum InFlightWork {
624 Admitted(Box<AdmittedExecution>),
625 #[cfg(test)]
626 Test {
627 inputs: Box<ExecutionInputs>,
628 work: Box<dyn FnOnce() -> Result<Vec<Tensor>> + Send>,
629 },
630}
631
632impl InFlightSubmission {
633 fn new(prepared: PreparedCompiledGraph, inputs: ExecutionInputs) -> Self {
634 Self {
635 work: Mutex::new(Some(InFlightWork::Admitted(Box::new(AdmittedExecution {
636 prepared,
637 inputs,
638 })))),
639 completion: Mutex::new(None),
640 completed: Condvar::new(),
641 }
642 }
643
644 #[cfg(test)]
645 pub(super) fn for_test(work: impl FnOnce() -> Result<Vec<Tensor>> + Send + 'static) -> Self {
646 Self {
647 work: Mutex::new(Some(InFlightWork::Test {
648 inputs: Box::new(ExecutionInputs::new(Vec::new()).expect("empty test inputs")),
649 work: Box::new(work),
650 })),
651 completion: Mutex::new(None),
652 completed: Condvar::new(),
653 }
654 }
655
656 fn into_unstarted_inputs(self) -> ExecutionInputs {
657 let work = match self.work.into_inner() {
658 Ok(work) => work,
659 Err(poisoned) => poisoned.into_inner(),
660 };
661 match work {
662 Some(InFlightWork::Admitted(admitted)) => admitted.inputs,
663 #[cfg(test)]
664 Some(InFlightWork::Test { inputs, .. }) => *inputs,
665 None => unreachable!("unstarted submission must still contain its owner"),
666 }
667 }
668
669 pub(super) fn run(&self) {
670 let work = match self.work.lock() {
671 Ok(mut work) => work.take(),
672 Err(poisoned) => poisoned.into_inner().take(),
673 };
674 let result = match work {
675 Some(InFlightWork::Admitted(admitted)) => run_admitted_work(admitted),
676 #[cfg(test)]
677 Some(InFlightWork::Test { work, .. }) => {
678 let result = catch_unwind(AssertUnwindSafe(work));
679 match result {
680 Ok(Ok(outputs)) => ExecutionBundle::from_outputs(outputs)
681 .map(|bundle| ExecutionOutcome::Completed(Box::new(bundle))),
682 Ok(Err(error)) => Err(error),
683 Err(payload) => Ok(ExecutionOutcome::CompletionUnproven {
684 error: Error::runtime_state(
685 "ExecutionHandle::wait",
686 ErrorPhase::Execution,
687 panic_payload_message(payload),
688 ),
689 diagnostic_keys: Box::from(["execution.retirement-unproven".to_owned()]),
690 }),
691 }
692 }
693 None => Err(Error::runtime_state(
694 "Runtime::submit",
695 ErrorPhase::Execution,
696 "in-flight submission work was already consumed",
697 )),
698 };
699 match self.completion.lock() {
700 Ok(mut completion) => {
701 *completion = Some(result);
702 }
703 Err(poisoned) => {
704 *poisoned.into_inner() = Some(Err(Error::runtime_state(
705 "ExecutionHandle::wait",
706 ErrorPhase::Execution,
707 "in-flight completion lock poisoned",
708 )));
709 }
710 }
711 self.completed.notify_all();
712 }
713
714 fn wait(&self) -> Result<ExecutionOutcome> {
715 let mut completion = self.completion.lock().map_err(|_| {
716 Error::runtime_state(
717 "ExecutionHandle::wait",
718 ErrorPhase::Execution,
719 "in-flight completion lock poisoned",
720 )
721 })?;
722 loop {
723 if let Some(result) = completion.take() {
724 return result;
725 }
726 completion = self.completed.wait(completion).map_err(|_| {
727 Error::runtime_state(
728 "ExecutionHandle::wait",
729 ErrorPhase::Execution,
730 "in-flight completion lock poisoned while waiting",
731 )
732 })?;
733 }
734 }
735
736 fn is_pending(&self) -> bool {
737 self.completion
738 .lock()
739 .map_or(true, |completion| completion.is_none())
740 }
741}
742
743fn run_admitted_work(admitted: Box<AdmittedExecution>) -> Result<ExecutionOutcome> {
744 let execution = catch_unwind(AssertUnwindSafe(|| {
745 let input_reads = admitted.inputs.as_reads()?;
746 execute_admitted(&admitted.prepared, &input_reads)
747 }));
748 match execution {
749 Ok(Ok(outputs)) => {
750 let AdmittedExecution { inputs, .. } = *admitted;
751 match ExecutionBundle::from_inputs_and_outputs(inputs, outputs) {
752 Ok(bundle) => Ok(ExecutionOutcome::Completed(Box::new(bundle))),
753 Err(error) => Ok(ExecutionOutcome::CompletionUnproven {
754 error,
755 diagnostic_keys: Box::from(["execution.bundle-build".to_owned()]),
756 }),
757 }
758 }
759 Ok(Err(error)) => {
760 let AdmittedExecution { inputs, .. } = *admitted;
761 Ok(ExecutionOutcome::RetiredFailed {
762 error,
763 inputs: Box::new(inputs),
764 })
765 }
766 Err(payload) => {
767 let error = Error::runtime_state(
768 "ExecutionHandle::wait",
769 ErrorPhase::Execution,
770 panic_payload_message(payload),
771 );
772 Box::leak(admitted);
776 Ok(ExecutionOutcome::CompletionUnproven {
777 error,
778 diagnostic_keys: Box::from(["execution.retirement-unproven".to_owned()]),
779 })
780 }
781 }
782}
783
784pub(super) trait SubmissionSpawner {
785 fn spawn(&self, submission: Arc<InFlightSubmission>) -> std::io::Result<()>;
786}
787
788pub(super) fn spawn_in_flight(
789 submission: Arc<InFlightSubmission>,
790 spawner: &dyn SubmissionSpawner,
791) -> std::result::Result<ExecutionHandle, Box<(ExecutionInputs, Error)>> {
792 if let Err(source) = spawner.spawn(Arc::clone(&submission)) {
793 let submission = match Arc::try_unwrap(submission) {
797 Ok(submission) => submission,
798 Err(_) => unreachable!("failed spawner must not retain submission"),
799 };
800 let inputs = submission.into_unstarted_inputs();
801 let error = Error::runtime_state_source(
802 "Runtime::submit",
803 ErrorPhase::Execution,
804 SubmissionError::WorkerSpawn { source },
805 );
806 return Err(Box::new((inputs, error)));
807 }
808 Ok(ExecutionHandle { submission })
809}
810
811pub(super) struct OsThreadSpawner;
812
813impl SubmissionSpawner for OsThreadSpawner {
814 fn spawn(&self, submission: Arc<InFlightSubmission>) -> std::io::Result<()> {
815 thread::Builder::new()
816 .name("tenferro-runtime-submit".to_string())
817 .spawn(move || submission.run())
818 .map(drop)
819 }
820}
821
822impl fmt::Debug for PreparedCompiledGraph {
823 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
824 formatter
825 .debug_struct("PreparedCompiledGraph")
826 .field("runtime_id", &self.runtime_id)
827 .field("epoch", &self.epoch)
828 .field("program", &self.program)
829 .finish_non_exhaustive()
830 }
831}
832
833fn panic_payload_message(payload: Box<dyn std::any::Any + Send + 'static>) -> String {
834 if let Some(message) = payload.downcast_ref::<&str>() {
835 format!("submitted execution panicked: {message}")
836 } else if let Some(message) = payload.downcast_ref::<String>() {
837 format!("submitted execution panicked: {message}")
838 } else {
839 "submitted execution panicked".to_string()
840 }
841}
842
843#[allow(
844 dead_code,
845 reason = "Phase 5 runtime execution task adds erased dispatch methods"
846)]
847pub(super) trait ErasedTensorBackendExecutor: fmt::Debug + Send + Sync {
848 fn backend_type_name(&self) -> &'static str;
849 fn extension_cache_stats(&self) -> std::result::Result<CacheStats, CacheOwnerError>;
850 fn clear_extension_caches(&self) -> std::result::Result<(), CacheOwnerError>;
851 fn execute(
852 &self,
853 program: &ExecProgram,
854 operations: &[PreparedOperationPlan],
855 inputs: Vec<Tensor>,
856 ) -> Result<Vec<Tensor>>;
857 fn execute_tensor_refs(
858 &self,
859 program: &ExecProgram,
860 operations: &[PreparedOperationPlan],
861 inputs: &[&Tensor],
862 ) -> Result<Vec<Tensor>>;
863 fn execute_values(
864 &self,
865 program: &ExecProgram,
866 operations: &[PreparedOperationPlan],
867 inputs: Vec<Tensor>,
868 ) -> Result<Vec<TensorValue>>;
869 fn execute_value_refs(
870 &self,
871 program: &ExecProgram,
872 operations: &[PreparedOperationPlan],
873 inputs: &[&Tensor],
874 ) -> Result<Vec<TensorValue>>;
875 fn execute_slot_instruction<'input>(
876 &self,
877 instruction_index: usize,
878 instruction: &ExecInstruction,
879 operations: &[PreparedOperationPlan],
880 slots: &mut [Option<ExecSlot<'input>>],
881 output_mode: RuntimeOutputMode,
882 terminal_slots: &[bool],
883 ) -> Result<()>;
884 fn execute_elementwise_fusion_slots<'input>(
891 &self,
892 input_slots: &[usize],
893 instruction_count: usize,
894 plan: &tenferro_tensor::backend::ElementwiseFusionPlan,
895 slots: &mut [Option<ExecSlot<'input>>],
896 output_slots: &[usize],
897 ) -> Result<bool>;
898 fn materialize_slot<'input>(&self, slot: ExecSlot<'input>) -> Result<Tensor>;
899 fn materialize_slot_value<'input>(&self, slot: ExecSlot<'input>) -> Result<TensorValue>;
900}
901
902pub(super) fn erased_tensor_backend_executor<B>(backend: B) -> Arc<dyn ErasedTensorBackendExecutor>
903where
904 B: TensorBackend + Send + Sync + 'static,
905{
906 Arc::new(TensorBackendExecutor::<B>::new(backend))
907}
908
909pub(super) fn extension_cache_owner(
910 executor: Arc<dyn ErasedTensorBackendExecutor>,
911) -> Arc<dyn RuntimeCacheOwner> {
912 Arc::new(TensorBackendExtensionCacheOwner { executor })
913}
914
915#[derive(Debug)]
916struct TensorBackendExtensionCacheOwner {
917 executor: Arc<dyn ErasedTensorBackendExecutor>,
918}
919
920impl RuntimeCacheOwner for TensorBackendExtensionCacheOwner {
921 fn cache_stats(&self) -> std::result::Result<CacheStats, CacheOwnerError> {
922 self.executor.extension_cache_stats()
923 }
924
925 fn clear_caches(&self) -> std::result::Result<(), CacheOwnerError> {
926 self.executor.clear_extension_caches()
927 }
928}
929
930#[allow(
931 dead_code,
932 reason = "Phase 5 runtime execution task consumes backend execution state"
933)]
934struct TensorBackendExecutorState<B: TensorBackend + 'static> {
935 backend: B,
936 backend_cache: B::RuntimeCache,
937 extension_caches: ExtensionCacheStore,
938 slot_workspace: Vec<Option<ExecSlot<'static>>>,
939 borrowed_slot_workspace_capacity: usize,
940}
941
942struct TensorBackendExecutor<B: TensorBackend + 'static> {
943 state: Mutex<TensorBackendExecutorSlot<B>>,
944 available: Condvar,
945}
946
947struct TensorBackendExecutorSlot<B: TensorBackend + 'static> {
948 state: Option<TensorBackendExecutorState<B>>,
949 active_thread: Option<ThreadId>,
950 execution_poisoned: bool,
951}
952
953impl<B> TensorBackendExecutor<B>
954where
955 B: TensorBackend + Send + Sync + 'static,
956{
957 fn new(backend: B) -> Self {
958 Self {
959 state: Mutex::new(TensorBackendExecutorSlot {
960 state: Some(TensorBackendExecutorState {
961 backend,
962 backend_cache: B::RuntimeCache::default(),
963 extension_caches: ExtensionCacheStore::new(),
964 slot_workspace: Vec::new(),
965 borrowed_slot_workspace_capacity: 0,
966 }),
967 active_thread: None,
968 execution_poisoned: false,
969 }),
970 available: Condvar::new(),
971 }
972 }
973
974 fn lease_state(&self, caller: &'static str) -> Result<TensorBackendExecutorLease<'_, B>> {
975 let current = thread::current().id();
976 let mut slot = self.lock_slot(caller)?;
977 loop {
978 if slot.execution_poisoned {
979 return Err(Error::runtime_state(
980 caller,
981 ErrorPhase::Execution,
982 "tensor backend executor state poisoned by panic during prior execution",
983 ));
984 }
985 if let Some(state) = slot.state.take() {
986 slot.active_thread = Some(current);
987 return Ok(TensorBackendExecutorLease {
988 executor: self,
989 state: Some(state),
990 });
991 }
992 if slot.active_thread == Some(current) {
993 return Err(Error::runtime_state(
994 caller,
995 ErrorPhase::Execution,
996 "reentrant tensor backend executor call would deadlock",
997 ));
998 }
999 slot = self.available.wait(slot).map_err(|_| {
1000 Error::runtime_state(
1001 caller,
1002 ErrorPhase::Execution,
1003 "tensor backend executor state lock poisoned while waiting",
1004 )
1005 })?;
1006 }
1007 }
1008
1009 fn lock_slot(
1010 &self,
1011 caller: &'static str,
1012 ) -> Result<MutexGuard<'_, TensorBackendExecutorSlot<B>>> {
1013 self.state.lock().map_err(|_| {
1014 Error::runtime_state(
1015 caller,
1016 ErrorPhase::Execution,
1017 "tensor backend executor state lock poisoned",
1018 )
1019 })
1020 }
1021}
1022
1023struct TensorBackendExecutorLease<'a, B: TensorBackend + 'static> {
1024 executor: &'a TensorBackendExecutor<B>,
1025 state: Option<TensorBackendExecutorState<B>>,
1026}
1027
1028impl<B: TensorBackend + 'static> TensorBackendExecutorLease<'_, B> {
1029 fn state_mut(&mut self) -> &mut TensorBackendExecutorState<B> {
1030 self.state
1031 .as_mut()
1032 .expect("executor state lease always owns state before drop")
1033 }
1034}
1035
1036impl<B: TensorBackend + 'static> Drop for TensorBackendExecutorLease<'_, B> {
1037 fn drop(&mut self) {
1038 let Some(state) = self.state.take() else {
1039 return;
1040 };
1041 let panicking = thread::panicking();
1042 let mut wake_all_waiters = panicking;
1043 match self.executor.state.lock() {
1044 Ok(mut slot) => {
1045 slot.state = Some(state);
1046 slot.active_thread = None;
1047 slot.execution_poisoned |= panicking;
1048 }
1049 Err(poisoned_lock) => {
1050 let mut slot = poisoned_lock.into_inner();
1051 slot.state = Some(state);
1052 slot.active_thread = None;
1053 slot.execution_poisoned = true;
1054 wake_all_waiters = true;
1055 }
1056 }
1057 if wake_all_waiters {
1058 self.executor.available.notify_all();
1059 } else {
1060 self.executor.available.notify_one();
1061 }
1062 }
1063}
1064
1065impl<B: TensorBackend + 'static> fmt::Debug for TensorBackendExecutor<B> {
1066 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
1067 formatter
1068 .debug_struct("TensorBackendExecutor")
1069 .field("backend_type", &std::any::type_name::<B>())
1070 .field("state_lock_poisoned", &self.state.is_poisoned())
1071 .finish_non_exhaustive()
1072 }
1073}
1074
1075impl<B> ErasedTensorBackendExecutor for TensorBackendExecutor<B>
1076where
1077 B: TensorBackend + Send + Sync + 'static,
1078{
1079 fn backend_type_name(&self) -> &'static str {
1080 std::any::type_name::<B>()
1081 }
1082
1083 fn extension_cache_stats(&self) -> std::result::Result<CacheStats, CacheOwnerError> {
1084 let lease = self
1085 .lease_state("Runtime::extension_cache_stats")
1086 .map_err(cache_owner_error)?;
1087 let state = lease
1088 .state
1089 .as_ref()
1090 .expect("executor state lease always owns state before drop");
1091 Ok(cache_stats_from_tensor_stats(
1092 state.extension_caches.stats(ExtensionCacheSelector::All),
1093 ))
1094 }
1095
1096 fn clear_extension_caches(&self) -> std::result::Result<(), CacheOwnerError> {
1097 let mut lease = self
1098 .lease_state("Runtime::clear_extension_caches")
1099 .map_err(cache_owner_error)?;
1100 lease.state_mut().extension_caches.clear();
1101 Ok(())
1102 }
1103
1104 fn execute(
1105 &self,
1106 program: &ExecProgram,
1107 operations: &[PreparedOperationPlan],
1108 inputs: Vec<Tensor>,
1109 ) -> Result<Vec<Tensor>> {
1110 validate_exec_input_count(program, inputs.len())?;
1111 let mut lease = self.lease_state("Runtime::run_compiled")?;
1112 let TensorBackendExecutorState {
1113 backend,
1114 backend_cache,
1115 extension_caches,
1116 slot_workspace,
1117 borrowed_slot_workspace_capacity: _,
1118 } = lease.state_mut();
1119 let mut extension_dispatch = ExtensionExecutionDispatch {
1120 operations,
1121 caches: extension_caches,
1122 };
1123 crate::segment::eval_exec_segmented_with_cache_and_workspace(
1124 backend,
1125 program,
1126 inputs,
1127 slot_workspace,
1128 backend_cache,
1129 Some(&mut extension_dispatch),
1130 )
1131 }
1132
1133 fn execute_tensor_refs(
1134 &self,
1135 program: &ExecProgram,
1136 operations: &[PreparedOperationPlan],
1137 inputs: &[&Tensor],
1138 ) -> Result<Vec<Tensor>> {
1139 validate_exec_input_count(program, inputs.len())?;
1140 let inputs = inputs
1141 .iter()
1142 .map(|tensor| ExecSlot::Read(TensorRead::from_tensor(tensor)))
1143 .collect();
1144 let mut lease = self.lease_state("Runtime::run_compiled")?;
1145 let TensorBackendExecutorState {
1146 backend,
1147 backend_cache,
1148 extension_caches,
1149 borrowed_slot_workspace_capacity,
1150 ..
1151 } = lease.state_mut();
1152 let mut extension_dispatch = ExtensionExecutionDispatch {
1153 operations,
1154 caches: extension_caches,
1155 };
1156 let mut slot_workspace = Vec::with_capacity(*borrowed_slot_workspace_capacity);
1157 let result = crate::segment::eval_exec_segmented_slots_with_cache_and_workspace(
1158 backend,
1159 program,
1160 inputs,
1161 &mut slot_workspace,
1162 backend_cache,
1163 Some(&mut extension_dispatch),
1164 );
1165 *borrowed_slot_workspace_capacity = slot_workspace.capacity();
1166 result
1167 }
1168
1169 fn execute_values(
1170 &self,
1171 program: &ExecProgram,
1172 operations: &[PreparedOperationPlan],
1173 inputs: Vec<Tensor>,
1174 ) -> Result<Vec<TensorValue>> {
1175 validate_exec_input_count(program, inputs.len())?;
1176 let inputs = inputs.into_iter().map(ExecSlot::Owned).collect();
1177 let mut lease = self.lease_state("Runtime::run_compiled_values")?;
1178 let TensorBackendExecutorState {
1179 backend,
1180 backend_cache,
1181 extension_caches,
1182 slot_workspace,
1183 borrowed_slot_workspace_capacity: _,
1184 } = lease.state_mut();
1185 let mut extension_dispatch = ExtensionExecutionDispatch {
1186 operations,
1187 caches: extension_caches,
1188 };
1189 crate::segment::eval_exec_segmented_slot_values_with_cache_and_workspace(
1190 backend,
1191 program,
1192 inputs,
1193 slot_workspace,
1194 backend_cache,
1195 Some(&mut extension_dispatch),
1196 )
1197 }
1198
1199 fn execute_value_refs(
1200 &self,
1201 program: &ExecProgram,
1202 operations: &[PreparedOperationPlan],
1203 inputs: &[&Tensor],
1204 ) -> Result<Vec<TensorValue>> {
1205 validate_exec_input_count(program, inputs.len())?;
1206 let inputs = inputs
1207 .iter()
1208 .map(|tensor| ExecSlot::Read(TensorRead::from_tensor(tensor)))
1209 .collect();
1210 let mut lease = self.lease_state("Runtime::run_compiled_values")?;
1211 let TensorBackendExecutorState {
1212 backend,
1213 backend_cache,
1214 extension_caches,
1215 borrowed_slot_workspace_capacity,
1216 ..
1217 } = lease.state_mut();
1218 let mut extension_dispatch = ExtensionExecutionDispatch {
1219 operations,
1220 caches: extension_caches,
1221 };
1222 let mut slot_workspace = Vec::with_capacity(*borrowed_slot_workspace_capacity);
1223 let result = crate::segment::eval_exec_segmented_slot_values_with_cache_and_workspace(
1224 backend,
1225 program,
1226 inputs,
1227 &mut slot_workspace,
1228 backend_cache,
1229 Some(&mut extension_dispatch),
1230 );
1231 *borrowed_slot_workspace_capacity = slot_workspace.capacity();
1232 result
1233 }
1234
1235 fn execute_slot_instruction<'input>(
1236 &self,
1237 instruction_index: usize,
1238 instruction: &ExecInstruction,
1239 operations: &[PreparedOperationPlan],
1240 slots: &mut [Option<ExecSlot<'input>>],
1241 output_mode: RuntimeOutputMode,
1242 terminal_slots: &[bool],
1243 ) -> Result<()> {
1244 let mut lease = self.lease_state("Runtime::run_compiled scheduled instruction")?;
1245 let TensorBackendExecutorState {
1246 backend,
1247 backend_cache,
1248 extension_caches,
1249 ..
1250 } = lease.state_mut();
1251 let mut extension_dispatch = ExtensionExecutionDispatch {
1252 operations,
1253 caches: extension_caches,
1254 };
1255
1256 let value_mode = matches!(output_mode, RuntimeOutputMode::Value);
1261 if crate::exec::is_host_instruction(instruction) {
1262 backend.with_backend_session(|exec| -> crate::Result<()> {
1263 if value_mode
1264 && crate::exec::try_execute_terminal_value_instruction(
1265 exec,
1266 slots,
1267 instruction,
1268 terminal_slots,
1269 )?
1270 {
1271 crate::exec::reclaim_last_use_inputs_exec(slots, instruction, exec);
1272 return Ok(());
1273 }
1274 crate::exec::execute_host_instruction_exec(exec, slots, instruction)?;
1275 crate::exec::reclaim_last_use_inputs_exec(slots, instruction, exec);
1276 Ok(())
1277 })??;
1278 } else if crate::exec::is_ffi_instruction(instruction) {
1279 if crate::exec::needs_owner_extension_fallback(instruction, Some(&extension_dispatch)) {
1280 if !(value_mode
1284 && crate::exec::instruction_may_be_terminal_value(instruction, terminal_slots)
1285 && backend.with_backend_session(|exec| {
1286 crate::exec::try_execute_terminal_value_instruction(
1287 exec,
1288 slots,
1289 instruction,
1290 terminal_slots,
1291 )
1292 })??)
1293 {
1294 crate::exec::execute_owner_extension_fallback(
1295 backend,
1296 slots,
1297 instruction,
1298 Some(&mut extension_dispatch),
1299 )?;
1300 }
1301 crate::exec::reclaim_last_use_inputs_via_session(backend, slots, instruction);
1302 } else {
1303 backend.with_backend_session_cached(
1304 backend_cache,
1305 |exec| -> crate::Result<()> {
1306 if value_mode
1307 && crate::exec::try_execute_terminal_value_instruction(
1308 exec,
1309 slots,
1310 instruction,
1311 terminal_slots,
1312 )?
1313 {
1314 crate::exec::reclaim_last_use_inputs_exec(slots, instruction, exec);
1315 return Ok(());
1316 }
1317 crate::exec::execute_ffi_instruction_exec(
1318 exec,
1319 slots,
1320 instruction,
1321 Some(instruction_index),
1322 Some(&mut extension_dispatch),
1323 )?;
1324 crate::exec::reclaim_last_use_inputs_exec(slots, instruction, exec);
1325 Ok(())
1326 },
1327 )??;
1328 }
1329 } else if value_mode
1330 && crate::exec::instruction_may_be_terminal_value(instruction, terminal_slots)
1331 && backend.with_backend_session(|exec| {
1332 crate::exec::try_execute_terminal_value_instruction(
1333 exec,
1334 slots,
1335 instruction,
1336 terminal_slots,
1337 )
1338 })??
1339 {
1340 crate::exec::reclaim_last_use_inputs_via_session(backend, slots, instruction);
1342 } else {
1343 backend.with_backend_session(|exec| -> crate::Result<()> {
1344 let result = crate::exec::execute_backend_op(exec, slots, instruction)?;
1345 slots[instruction.output_slots[0]] = Some(ExecSlot::Owned(result));
1346 crate::exec::reclaim_last_use_inputs_exec(slots, instruction, exec);
1347 Ok(())
1348 })??;
1349 }
1350 Ok(())
1351 }
1352
1353 fn execute_elementwise_fusion_slots<'input>(
1354 &self,
1355 input_slots: &[usize],
1356 instruction_count: usize,
1357 plan: &tenferro_tensor::backend::ElementwiseFusionPlan,
1358 slots: &mut [Option<ExecSlot<'input>>],
1359 output_slots: &[usize],
1360 ) -> Result<bool> {
1361 let mut views: Vec<usize> = Vec::new();
1366 for &slot in input_slots {
1367 let value = slots
1368 .get(slot)
1369 .and_then(Option::as_ref)
1370 .ok_or_else(|| crate::Error::from(tenferro_tensor::Error::MissingValue { slot }))?;
1371 if value.as_tensor("elementwise_region").is_err() {
1372 views.push(slot);
1373 }
1374 }
1375 if !views.is_empty() {
1376 if instruction_count < super::region::VIEW_COPY_MIN_INSTRUCTIONS {
1377 return Ok(false);
1378 }
1379 let mut lease = self.lease_state("Runtime::run_prepared elementwise region copy")?;
1380 let backend = &mut lease.state_mut().backend;
1381 for &slot in &views {
1382 let tensor = {
1383 let read = slots[slot].as_ref().map(ExecSlot::as_read).ok_or_else(|| {
1384 crate::Error::from(tenferro_tensor::Error::MissingValue { slot })
1385 })?;
1386 backend.with_backend_session(|exec| exec.to_contiguous_read(read))??
1387 };
1388 slots[slot] = Some(ExecSlot::Owned(tensor));
1389 }
1390 }
1391
1392 let mut inputs: Vec<&Tensor> = Vec::with_capacity(input_slots.len());
1393 for &slot in input_slots {
1394 let value = slots
1395 .get(slot)
1396 .and_then(Option::as_ref)
1397 .ok_or_else(|| crate::Error::from(tenferro_tensor::Error::MissingValue { slot }))?;
1398 inputs.push(value.as_tensor("elementwise_region")?);
1399 }
1400
1401 if !crate::segment::plan_matches_input_dtypes(plan, &inputs) {
1402 return Ok(false);
1403 }
1404
1405 let mut lease = self.lease_state("Runtime::run_prepared elementwise region")?;
1406 let backend = &mut lease.state_mut().backend;
1407 let outputs = backend
1408 .with_backend_session(|exec| exec.execute_elementwise_fusion(&inputs, plan))??;
1409 let Some(outputs) = outputs else {
1410 return Ok(false);
1411 };
1412 if outputs.len() != output_slots.len() {
1413 return Err(crate::Error::Internal(format!(
1414 "fused elementwise region produced {} outputs for {} slots",
1415 outputs.len(),
1416 output_slots.len()
1417 )));
1418 }
1419 for (&slot, tensor) in output_slots.iter().zip(outputs) {
1420 slots[slot] = Some(ExecSlot::Owned(tensor));
1421 }
1422 Ok(true)
1423 }
1424
1425 fn materialize_slot<'input>(&self, slot: ExecSlot<'input>) -> Result<Tensor> {
1426 let mut lease = self.lease_state("Runtime::run_compiled collect outputs")?;
1427 let backend = &mut lease.state_mut().backend;
1428 backend.with_backend_session(|exec| slot.into_tensor(exec))?
1429 }
1430
1431 fn materialize_slot_value<'input>(&self, slot: ExecSlot<'input>) -> Result<TensorValue> {
1432 let mut lease = self.lease_state("Runtime::run_compiled_values collect outputs")?;
1433 let backend = &mut lease.state_mut().backend;
1434 backend.with_backend_session(|exec| slot.into_value(exec))?
1435 }
1436}
1437
1438fn cache_owner_error(error: Error) -> CacheOwnerError {
1439 CacheOwnerError::new(Arc::new(error))
1440}
1441
1442fn cache_stats_from_tensor_stats(stats: tenferro_tensor::CacheStats) -> CacheStats {
1443 CacheStats {
1444 entries: stats.entries,
1445 retained_bytes: stats.retained_bytes,
1446 hits: stats.hits,
1447 misses: stats.misses,
1448 evictions: stats.evictions,
1449 clears: stats.clears,
1450 }
1451}
1452
1453pub(super) fn run_compiled(
1454 runtime: &Runtime,
1455 program: &CompiledGraph,
1456 inputs: &[&Tensor],
1457) -> Result<Vec<Tensor>> {
1458 let inputs = resolve_input_refs(program, inputs)?;
1459 let signature = input_signature_reads(&inputs)?;
1460 let prepared = prepare(runtime, program, &signature)?;
1461 validate_prepared_epoch(runtime, prepared.root().epoch(), "Runtime::run_compiled")?;
1462 execute_scheduled_reads(
1463 prepared.root().staging(),
1464 prepared.root().schedule(),
1465 prepared.root().regions(),
1466 prepared.root().region_counters(),
1467 prepared.operations(),
1468 &inputs,
1469 )
1470}
1471
1472pub(super) fn execute_scoped_read_only<'env>(
1477 runtime: &Runtime,
1478 program: &CompiledGraph,
1479 inputs: ScopedReadInputs<'env>,
1480) -> std::result::Result<ScopedExecutionOutcome<'env>, ScopedSubmitRejected<'env>> {
1481 if inputs.has_non_host_provider() {
1482 return Err(ScopedSubmitRejected {
1483 source: Box::new(Error::unsupported(
1484 "Runtime::execute_scoped_read_only",
1485 ErrorPhase::Execution,
1486 "scoped borrowed execution requires host/CPU storage",
1487 )),
1488 inputs,
1489 });
1490 }
1491
1492 let reads = inputs.as_reads();
1493 let prepared = match prepare_compiled_reads(runtime, program, &reads) {
1494 Ok(prepared) => prepared,
1495 Err(source) => {
1496 return Err(ScopedSubmitRejected {
1497 source: Box::new(source),
1498 inputs,
1499 })
1500 }
1501 };
1502 if !schedule_supports_scoped_execution(prepared.prepared.root().schedule()) {
1503 return Err(ScopedSubmitRejected {
1504 source: Box::new(Error::unsupported(
1505 "Runtime::execute_scoped_read_only",
1506 ErrorPhase::Execution,
1507 "the selected provider does not prove retirement before scoped return",
1508 )),
1509 inputs,
1510 });
1511 }
1512
1513 match execute_scoped_admitted(
1514 prepared.prepared.root().staging(),
1515 prepared.prepared.root().schedule(),
1516 prepared.prepared.operations(),
1517 &reads,
1518 ) {
1519 Ok(bundle) => Ok(ScopedExecutionOutcome::Completed(bundle)),
1520 Err(error) => Ok(ScopedExecutionOutcome::RetiredFailed { error, inputs }),
1521 }
1522}
1523
1524fn execute_scoped_admitted<'env>(
1525 program: &ExecProgram,
1526 schedule: &ScheduledGraph,
1527 operations: &[PreparedOperationPlan],
1528 inputs: &[TensorRead<'env>],
1529) -> Result<ScopedExecutionBundle<'env>> {
1530 validate_exec_input_count(program, inputs.len())?;
1531 crate::exec::validate_exec_program(program, "scoped tensor executor")?;
1532 let input_slots = inputs.iter().cloned().map(ExecSlot::Read).collect();
1533 let mut located = execute_scheduled_slots(
1534 program,
1535 schedule,
1536 &[],
1537 &super::region::RegionExecutionCounters::default(),
1538 operations,
1539 input_slots,
1540 RuntimeOutputMode::Tensor,
1541 )?;
1542 collect_scoped_outputs(program, &mut located)
1543}
1544
1545pub(crate) fn collect_scoped_outputs<'env>(
1546 program: &ExecProgram,
1547 outputs: &mut [Option<LocatedExecSlot<'env>>],
1548) -> Result<ScopedExecutionBundle<'env>> {
1549 let mut owned = AllocationGroup::from_tensors(Vec::new())
1550 .map(|(group, _)| group)
1551 .map_err(|error| {
1552 Error::runtime_state(
1553 "Runtime::collect_scoped_outputs",
1554 ErrorPhase::Execution,
1555 error.to_string(),
1556 )
1557 })?;
1558 let mut scoped = Vec::with_capacity(program.output_slots.len());
1559 for &slot in &program.output_slots {
1560 let located = outputs
1561 .get_mut(slot)
1562 .and_then(Option::take)
1563 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
1564 let output = match located.value {
1565 ExecSlot::Read(read) => ScopedOutput::Borrowed(read.tensor_view()),
1566 ExecSlot::Owned(tensor) => {
1567 let slot = owned.append_tensor(tensor).map_err(|error| {
1568 Error::runtime_state(
1569 "Runtime::collect_scoped_outputs",
1570 ErrorPhase::Execution,
1571 error.to_string(),
1572 )
1573 })?;
1574 ScopedOutput::Owned(slot)
1575 }
1576 ExecSlot::Value(value) => {
1577 let (group, slot, dtype, shape) = value.try_into_group_parts().map_err(|_| {
1578 Error::runtime_state(
1579 "Runtime::collect_scoped_outputs",
1580 ErrorPhase::Execution,
1581 "scoped value output was retained by another owner",
1582 )
1583 })?;
1584 let slot = owned.append_group(group, slot).map_err(|error| {
1585 Error::runtime_state(
1586 "Runtime::collect_scoped_outputs",
1587 ErrorPhase::Execution,
1588 error.to_string(),
1589 )
1590 })?;
1591 let _ = (dtype, shape);
1592 ScopedOutput::Owned(slot)
1593 }
1594 };
1595 scoped.push(output);
1596 }
1597 outputs.iter_mut().for_each(|output| *output = None);
1598 Ok(ScopedExecutionBundle {
1599 owned,
1600 outputs: scoped.into_boxed_slice(),
1601 })
1602}
1603
1604fn output_ref_from_group<'a>(
1605 group: &'a AllocationGroup,
1606 slot: DescriptorSlot,
1607) -> std::result::Result<OutputRef<'a>, OutputAccessError> {
1608 let mut views = group.read_views(std::slice::from_ref(&slot))?;
1609 match views.pop() {
1610 Some(TensorRead::View(view)) => Ok(OutputRef::Tensor(view)),
1611 Some(TensorRead::Tensor(_)) | None => Err(OutputAccessError::InvalidOutput {
1612 index: slot.index(),
1613 }),
1614 }
1615}
1616
1617fn schedule_supports_scoped_execution(schedule: &ScheduledGraph) -> bool {
1618 let supports = |location: &ExecutionLocation| {
1619 location
1620 .witness()
1621 .event_domain_driver()
1622 .supports_scoped_read_only()
1623 };
1624 supports(schedule.root_location())
1625 && schedule.input_locations().iter().all(supports)
1626 && schedule.operation_locations().iter().all(supports)
1627 && schedule.nodes().iter().all(|node| {
1628 node.event_domain_witness()
1629 .event_domain_driver()
1630 .supports_scoped_read_only()
1631 })
1632}
1633
1634pub(super) fn prepare_compiled(
1635 runtime: &Runtime,
1636 program: &CompiledGraph,
1637 inputs: &[&Tensor],
1638) -> Result<PreparedCompiledGraph> {
1639 let inputs = resolve_input_refs(program, inputs)?;
1640 let signature = input_signature_reads(&inputs)?;
1641 let prepared = prepare(runtime, program, &signature)?;
1642 validate_prepared_epoch(
1643 runtime,
1644 prepared.root().epoch(),
1645 "Runtime::prepare_compiled",
1646 )?;
1647 Ok(PreparedCompiledGraph {
1648 runtime_id: runtime.id(),
1649 epoch: prepared.root().epoch(),
1650 program: program.clone(),
1651 prepared,
1652 })
1653}
1654
1655pub(super) fn submit(
1656 runtime: &Runtime,
1657 program: &CompiledGraph,
1658 inputs: ExecutionInputs,
1659) -> std::result::Result<ExecutionHandle, SubmitError> {
1660 submit_with_spawner(runtime, program, inputs, &OsThreadSpawner)
1661}
1662
1663pub(super) fn submit_with_spawner(
1664 runtime: &Runtime,
1665 program: &CompiledGraph,
1666 inputs: ExecutionInputs,
1667 spawner: &dyn SubmissionSpawner,
1668) -> std::result::Result<ExecutionHandle, SubmitError> {
1669 let prepared = match prepare_submission(runtime, program, &inputs) {
1670 Ok(prepared) => prepared,
1671 Err(source) => {
1672 return Err(SubmitError::PreAdmission {
1673 source: Box::new(source),
1674 inputs: Box::new(inputs),
1675 })
1676 }
1677 };
1678 let submission = Arc::new(InFlightSubmission::new(prepared, inputs));
1679 match spawn_in_flight(submission, spawner) {
1680 Ok(handle) => Ok(handle),
1681 Err(failure) => {
1682 let (inputs, source) = *failure;
1683 Err(SubmitError::PreAdmission {
1684 source: Box::new(source),
1685 inputs: Box::new(inputs),
1686 })
1687 }
1688 }
1689}
1690
1691pub(super) fn run_prepared(
1692 runtime: &Runtime,
1693 prepared: &PreparedCompiledGraph,
1694 inputs: &[&Tensor],
1695) -> Result<Vec<Tensor>> {
1696 validate_prepared_runtime(runtime, prepared, "Runtime::run_prepared")?;
1697 let inputs = resolve_input_refs(&prepared.program, inputs)?;
1698 execute_scheduled_reads(
1699 prepared.prepared.root().staging(),
1700 prepared.prepared.root().schedule(),
1701 prepared.prepared.root().regions(),
1702 prepared.prepared.root().region_counters(),
1703 prepared.prepared.operations(),
1704 &inputs,
1705 )
1706}
1707
1708fn prepare_submission(
1709 runtime: &Runtime,
1710 program: &CompiledGraph,
1711 inputs: &ExecutionInputs,
1712) -> Result<PreparedCompiledGraph> {
1713 let input_reads = inputs.as_reads()?;
1714 prepare_compiled_reads(runtime, program, &input_reads)
1715}
1716
1717fn prepare_compiled_reads(
1718 runtime: &Runtime,
1719 program: &CompiledGraph,
1720 inputs: &[TensorRead<'_>],
1721) -> Result<PreparedCompiledGraph> {
1722 validate_ordered_input_metadata_reads(program, inputs)?;
1723 let signature = input_signature_reads(inputs)?;
1724 let prepared = prepare(runtime, program, &signature)?;
1725 validate_prepared_epoch(runtime, prepared.root().epoch(), "Runtime::submit")?;
1726 Ok(PreparedCompiledGraph {
1727 runtime_id: runtime.id(),
1728 epoch: prepared.root().epoch(),
1729 program: program.clone(),
1730 prepared,
1731 })
1732}
1733
1734fn execute_admitted(
1735 prepared: &PreparedCompiledGraph,
1736 inputs: &[TensorRead<'_>],
1737) -> Result<Vec<Tensor>> {
1738 execute_scheduled_reads(
1739 prepared.prepared.root().staging(),
1740 prepared.prepared.root().schedule(),
1741 prepared.prepared.root().regions(),
1742 prepared.prepared.root().region_counters(),
1743 prepared.prepared.operations(),
1744 inputs,
1745 )
1746}
1747
1748pub(super) fn run_compiled_values(
1749 runtime: &Runtime,
1750 program: &CompiledGraph,
1751 inputs: &[&Tensor],
1752) -> Result<Vec<TensorValue>> {
1753 let inputs = resolve_input_refs(program, inputs)?;
1754 let signature = input_signature_reads(&inputs)?;
1755 let prepared = prepare(runtime, program, &signature)?;
1756 validate_prepared_epoch(
1757 runtime,
1758 prepared.root().epoch(),
1759 "Runtime::run_compiled_values",
1760 )?;
1761 execute_scheduled_value_reads(
1762 prepared.root().staging(),
1763 prepared.root().schedule(),
1764 prepared.operations(),
1765 &inputs,
1766 )
1767}
1768
1769fn validate_prepared_runtime(
1770 runtime: &Runtime,
1771 prepared: &PreparedCompiledGraph,
1772 caller: &'static str,
1773) -> Result<()> {
1774 if runtime.id() != prepared.runtime_id {
1775 return Err(Error::runtime_state(
1776 caller,
1777 ErrorPhase::Execution,
1778 "prepared compiled graph belongs to a different runtime",
1779 ));
1780 }
1781 let epoch = runtime
1782 .epoch()
1783 .map_err(|source| Error::runtime_state_source(caller, ErrorPhase::Execution, source))?;
1784 if epoch != prepared.epoch {
1785 return Err(Error::runtime_state(
1786 caller,
1787 ErrorPhase::Execution,
1788 format!(
1789 "prepared epoch {:?} does not match current epoch {:?}",
1790 prepared.epoch, epoch
1791 ),
1792 ));
1793 }
1794 Ok(())
1795}
1796
1797fn prepare(
1798 runtime: &Runtime,
1799 program: &CompiledGraph,
1800 signature: &InputSignature,
1801) -> Result<Arc<super::preparation::PreparedProgram>> {
1802 runtime
1803 .prepare_compiled_for(program, signature, &PrepareOptions::new())
1804 .map_err(prepare_error)
1805}
1806
1807fn validate_prepared_epoch(
1808 runtime: &Runtime,
1809 prepared_epoch: super::RuntimeEpoch,
1810 caller: &'static str,
1811) -> Result<()> {
1812 let epoch = runtime
1813 .epoch()
1814 .map_err(|source| Error::runtime_state_source(caller, ErrorPhase::Execution, source))?;
1815 if epoch != prepared_epoch {
1816 return Err(Error::runtime_state(
1817 caller,
1818 ErrorPhase::Execution,
1819 format!(
1820 "prepared epoch {:?} does not match current epoch {:?}",
1821 prepared_epoch, epoch
1822 ),
1823 ));
1824 }
1825 Ok(())
1826}
1827
1828fn execute_scheduled_reads(
1829 program: &ExecProgram,
1830 schedule: &ScheduledGraph,
1831 regions: &[super::region::ElementwiseRegion],
1832 region_counters: &super::region::RegionExecutionCounters,
1833 operations: &[PreparedOperationPlan],
1834 inputs: &[TensorRead<'_>],
1835) -> Result<Vec<Tensor>> {
1836 validate_exec_input_count(program, inputs.len())?;
1837 crate::exec::validate_exec_program(program, "scheduled tensor executor")?;
1838 let inputs = inputs.iter().cloned().map(ExecSlot::Read).collect();
1839 execute_scheduled_slots(
1840 program,
1841 schedule,
1842 regions,
1843 region_counters,
1844 operations,
1845 inputs,
1846 RuntimeOutputMode::Tensor,
1847 )
1848 .and_then(|mut slots| {
1849 collect_tensor_outputs_with(program, &mut slots, |location, slot| {
1850 location.witness().executor().materialize_slot(slot)
1851 })
1852 })
1853}
1854
1855fn execute_scheduled_value_reads(
1856 program: &ExecProgram,
1857 schedule: &ScheduledGraph,
1858 operations: &[PreparedOperationPlan],
1859 inputs: &[TensorRead<'_>],
1860) -> Result<Vec<TensorValue>> {
1861 validate_exec_input_count(program, inputs.len())?;
1862 crate::exec::validate_exec_program(program, "scheduled value executor")?;
1863 let inputs = inputs.iter().cloned().map(ExecSlot::Read).collect();
1864 execute_scheduled_slots(
1865 program,
1866 schedule,
1867 &[],
1868 &super::region::RegionExecutionCounters::default(),
1869 operations,
1870 inputs,
1871 RuntimeOutputMode::Value,
1872 )
1873 .and_then(|mut slots| {
1874 collect_value_outputs_with(program, &mut slots, |location, slot| {
1875 location.witness().executor().materialize_slot_value(slot)
1876 })
1877 })
1878}
1879
1880fn execute_scheduled_slots<'input>(
1881 program: &ExecProgram,
1882 schedule: &ScheduledGraph,
1883 regions: &[super::region::ElementwiseRegion],
1884 region_counters: &super::region::RegionExecutionCounters,
1885 operations: &[PreparedOperationPlan],
1886 inputs: Vec<ExecSlot<'input>>,
1887 output_mode: RuntimeOutputMode,
1888) -> Result<Vec<Option<LocatedExecSlot<'input>>>> {
1889 let terminal_slots = if matches!(output_mode, RuntimeOutputMode::Value) {
1890 crate::exec::terminal_output_slots(program)
1891 } else {
1892 Vec::new()
1893 };
1894 let mut staged = Vec::new();
1895 let mut located = (0..program.n_slots)
1896 .map(|_| Vec::new())
1897 .collect::<Vec<Vec<LocatedExecSlot<'input>>>>();
1898 let regions_applicable = matches!(output_mode, RuntimeOutputMode::Tensor);
1906 let mut region_starts: HashMap<usize, usize> = HashMap::new();
1907 let mut region_covered: HashSet<usize> = HashSet::new();
1908 if regions_applicable {
1909 for (region_index, region) in regions.iter().enumerate() {
1910 if let Some(&first) = region.node_indices.first() {
1911 region_starts.insert(first, region_index);
1912 }
1913 region_covered.extend(region.node_indices.iter().copied());
1914 }
1915 }
1916
1917 let mut event_domains = ScheduledEventDomains::new(schedule)?;
1918 let result = (|| {
1919 crate::exec::initialize_exec_slots_in(program, inputs, &mut staged)?;
1920 if schedule.input_locations().len() != program.input_slots.len() {
1921 return Err(Error::runtime_state(
1922 "Runtime::run_compiled",
1923 ErrorPhase::Execution,
1924 format!(
1925 "prepared schedule has {} input locations for {} inputs",
1926 schedule.input_locations().len(),
1927 program.input_slots.len()
1928 ),
1929 ));
1930 }
1931 for (&slot, location) in program.input_slots.iter().zip(schedule.input_locations()) {
1932 let value = staged
1933 .get_mut(slot)
1934 .and_then(Option::take)
1935 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
1936 validate_runtime_input_ingress(location, &value.as_read(), slot)?;
1937 located[slot].push(LocatedExecSlot {
1938 location: location.clone(),
1939 value,
1940 });
1941 }
1942 for (node_index, node) in schedule.nodes().iter().enumerate() {
1943 match node {
1944 ScheduledNode::Operation(operation_node) => {
1945 if let Some(®ion_index) = region_starts.get(&node_index) {
1946 let region = ®ions[region_index];
1947 let first_instruction = program
1948 .instructions
1949 .get(region.instruction_range.start)
1950 .ok_or_else(|| {
1951 Error::runtime_state(
1952 "Runtime::run_prepared",
1953 ErrorPhase::Execution,
1954 format!(
1955 "elementwise region references instruction {}, but the execution program has {} instructions",
1956 region.instruction_range.start,
1957 program.instructions.len()
1958 ),
1959 )
1960 })?;
1961 let operation = instruction_execution(schedule, first_instruction)?;
1962 if operation.location() != operation_node.location() {
1963 return Err(Error::runtime_state(
1964 "Runtime::run_prepared",
1965 ErrorPhase::Execution,
1966 "elementwise region location does not match its first prepared operation"
1967 .to_string(),
1968 ));
1969 }
1970 let location = operation.location().clone();
1971 let mut launch = || {
1972 for &slot in ®ion.input_slots {
1973 stage_slot_input(slot, &location, &mut located, &mut staged)?;
1974 }
1975 let applied = operation.executor().execute_elementwise_fusion_slots(
1976 ®ion.input_slots,
1977 region.instructions.len(),
1978 ®ion.plan,
1979 &mut staged,
1980 ®ion.output_slots,
1981 )?;
1982 if applied {
1983 region_counters.record_fused();
1984 validate_region_outputs(region, &location, &staged)?;
1985 return retain_region_results(
1986 region,
1987 &location,
1988 &mut located,
1989 &mut staged,
1990 );
1991 }
1992 region_counters.record_fallback();
1995 for (offset, instruction) in region.instructions.iter().enumerate() {
1996 let instruction_index = region.instruction_range.start + offset;
1997 let member = instruction_execution(schedule, instruction)?;
1998 if member.location() != operation_node.location() {
1999 return Err(Error::runtime_state(
2000 "Runtime::run_prepared",
2001 ErrorPhase::Execution,
2002 "elementwise region member location does not match the region location"
2003 .to_string(),
2004 ));
2005 }
2006 stage_instruction_inputs(
2007 instruction,
2008 member.location(),
2009 &mut located,
2010 &mut staged,
2011 )?;
2012 member.executor().execute_slot_instruction(
2013 instruction_index,
2014 instruction,
2015 operations,
2016 &mut staged,
2017 output_mode,
2018 &terminal_slots,
2019 )?;
2020 validate_instruction_outputs(
2021 instruction_index,
2022 instruction,
2023 member.location(),
2024 &staged,
2025 )?;
2026 retain_instruction_results(
2027 instruction,
2028 member.location(),
2029 &mut located,
2030 &mut staged,
2031 )?;
2032 }
2033 Ok(())
2034 };
2035 event_domains.enqueue(node_index, node, &mut launch)?;
2036 region_counters.record_submission();
2037 let completion = EventDependency::from_completion(node.completion());
2041 if let Some(token) = event_domains.completions.get(&completion).cloned() {
2042 for &covered in ®ion.node_indices {
2043 if covered == node_index {
2044 continue;
2045 }
2046 let covered_node = schedule.nodes().get(covered).ok_or_else(|| {
2047 Error::runtime_state(
2048 "Runtime::run_prepared",
2049 ErrorPhase::Execution,
2050 format!(
2051 "elementwise region references schedule node {covered}, but the schedule has {} nodes",
2052 schedule.nodes().len()
2053 ),
2054 )
2055 })?;
2056 event_domains.completions.insert(
2057 EventDependency::from_completion(covered_node.completion()),
2058 Arc::clone(&token),
2059 );
2060 }
2061 }
2062 continue;
2063 }
2064 if region_covered.contains(&node_index) {
2065 continue;
2067 }
2068 let mut launch = || {
2069 let instruction_index = operation_node.instruction_index();
2070 let instruction =
2071 program.instructions.get(instruction_index).ok_or_else(|| {
2072 Error::runtime_state(
2073 "Runtime::run_compiled",
2074 ErrorPhase::Execution,
2075 format!(
2076 "scheduled operation references instruction \
2077 {instruction_index}, but the execution program has {} \
2078 instructions",
2079 program.instructions.len()
2080 ),
2081 )
2082 })?;
2083 let operation = instruction_execution(schedule, instruction)?;
2084 if operation.location() != operation_node.location() {
2085 return Err(Error::runtime_state(
2086 "Runtime::run_compiled",
2087 ErrorPhase::Execution,
2088 format!(
2089 "scheduled instruction {instruction_index} location does not \
2090 match its prepared executor"
2091 ),
2092 ));
2093 }
2094 stage_instruction_inputs(
2095 instruction,
2096 operation.location(),
2097 &mut located,
2098 &mut staged,
2099 )?;
2100 operation.executor().execute_slot_instruction(
2101 instruction_index,
2102 instruction,
2103 operations,
2104 &mut staged,
2105 output_mode,
2106 &terminal_slots,
2107 )?;
2108 validate_instruction_outputs(
2109 instruction_index,
2110 instruction,
2111 operation.location(),
2112 &staged,
2113 )?;
2114 retain_instruction_results(
2115 instruction,
2116 operation.location(),
2117 &mut located,
2118 &mut staged,
2119 )
2120 };
2121 event_domains.enqueue(node_index, node, &mut launch)?;
2122 region_counters.record_submission();
2123 }
2124 ScheduledNode::Transfer(transfer) => {
2125 let mut launch = || execute_scheduled_transfer(transfer, &mut located);
2126 event_domains.enqueue(node_index, node, &mut launch)?;
2127 region_counters.record_submission();
2128 }
2129 ScheduledNode::Collective(_) => {
2130 return Err(Error::runtime_state_source(
2131 "Runtime::run_compiled",
2132 ErrorPhase::Execution,
2133 UnsupportedScheduledNodeError {
2134 node_index,
2135 node_kind: ScheduledNodeKind::Collective,
2136 },
2137 ));
2138 }
2139 ScheduledNode::Barrier(_) => {
2140 let mut launch = || Ok(());
2141 event_domains.enqueue(node_index, node, &mut launch)?;
2142 region_counters.record_submission();
2143 }
2144 }
2145 }
2146 Ok(())
2147 })();
2148 let drain = event_domains.drain();
2149 match (result, drain) {
2150 (Ok(()), Ok(())) => collect_located_outputs(program, &mut located),
2151 (Err(error), Ok(())) | (Ok(()), Err(error)) => {
2152 staged.clear();
2153 located.clear();
2154 Err(error)
2155 }
2156 (Err(primary), Err(cleanup)) => {
2157 staged.clear();
2158 located.clear();
2159 Err(scheduled_execution_cleanup_error(primary, cleanup))
2160 }
2161 }
2162}
2163
2164#[derive(Debug)]
2165pub(crate) struct ScheduledEventDomains {
2166 runs: Vec<RuntimeOwnedEventDomainRun>,
2167 completions: HashMap<EventDependency, Arc<dyn EventToken>>,
2168}
2169
2170#[derive(Debug, thiserror::Error)]
2171#[error(
2172 "scheduled node {node_index} depends on completion {dependency:?}, but no completion token was recorded"
2173)]
2174pub(crate) struct MissingScheduledDependencyCompletionError {
2175 pub(crate) dependency: EventDependency,
2176 pub(crate) node_index: usize,
2177}
2178
2179#[derive(Debug, thiserror::Error)]
2180#[error("scheduled transfer destination {destination:?} already contains value slot {value_slot}")]
2181pub(crate) struct DuplicateTransferDestinationError {
2182 pub(crate) value_slot: usize,
2183 pub(crate) destination: ExecutionLocation,
2184}
2185
2186impl ScheduledEventDomains {
2187 pub(crate) fn new(schedule: &ScheduledGraph) -> Result<Self> {
2188 schedule.preflight().map_err(|source| {
2189 Error::runtime_state_source("Runtime::run_compiled", ErrorPhase::Execution, source)
2190 })?;
2191 let mut drivers = Vec::new();
2192 let mut seen_domains = HashSet::new();
2193 for node in schedule.nodes() {
2194 let domain = node.completion().domain();
2195 if !seen_domains.insert(domain) {
2196 continue;
2197 }
2198 let witness = node.event_domain_witness();
2199 debug_assert_eq!(witness.event_domain_id(), domain);
2200 drivers.push((domain, witness.event_domain_driver().clone()));
2201 }
2202 let mut runs = Vec::with_capacity(drivers.len());
2203 for (domain, driver) in drivers {
2204 let run = RuntimeOwnedEventDomainRun::new(domain, driver.begin_run(domain)?);
2205 let actual = run.domain(EventDomainOperation::BeginRun)?;
2206 if actual != domain {
2207 return Err(event_domain_error(EventDomainError::RunDomainMismatch {
2208 operation: EventDomainOperation::BeginRun,
2209 node_index: None,
2210 expected: domain,
2211 actual,
2212 }));
2213 }
2214 runs.push(run);
2215 }
2216 Ok(Self {
2217 runs,
2218 completions: HashMap::new(),
2219 })
2220 }
2221
2222 #[cfg(test)]
2223 pub(crate) fn for_test(
2224 drivers: Vec<(super::EventDomainId, Arc<dyn super::EventDomainDriver>)>,
2225 ) -> Result<Self> {
2226 let mut runs = Vec::with_capacity(drivers.len());
2227 for (domain, driver) in drivers {
2228 let run = RuntimeOwnedEventDomainRun::new(domain, driver.begin_run(domain)?);
2229 let actual = run.domain(EventDomainOperation::BeginRun)?;
2230 if actual != domain {
2231 return Err(event_domain_error(EventDomainError::RunDomainMismatch {
2232 operation: EventDomainOperation::BeginRun,
2233 node_index: None,
2234 expected: domain,
2235 actual,
2236 }));
2237 }
2238 runs.push(run);
2239 }
2240 Ok(Self {
2241 runs,
2242 completions: HashMap::new(),
2243 })
2244 }
2245
2246 pub(crate) fn enqueue(
2247 &mut self,
2248 node_index: usize,
2249 node: &ScheduledNode,
2250 launch: &mut dyn FnMut() -> Result<()>,
2251 ) -> Result<()> {
2252 let completion = node.completion();
2253 let destination = completion.domain();
2254 let run_index = self.run_index(completion.domain())?;
2255 let actual_preflight_domain = self.runs[run_index].domain(EventDomainOperation::Enqueue)?;
2256 if actual_preflight_domain != destination {
2257 return Err(event_domain_error(EventDomainError::RunDomainMismatch {
2258 operation: EventDomainOperation::Enqueue,
2259 node_index: Some(node_index),
2260 expected: destination,
2261 actual: actual_preflight_domain,
2262 }));
2263 }
2264 let dependencies = self.classify_dependencies(node_index, node, destination)?;
2265 let actual_run_domain = self.runs[run_index].domain(EventDomainOperation::Enqueue)?;
2266 if actual_run_domain != destination {
2267 return Err(event_domain_error(EventDomainError::RunDomainMismatch {
2268 operation: EventDomainOperation::Enqueue,
2269 node_index: Some(node_index),
2270 expected: destination,
2271 actual: actual_run_domain,
2272 }));
2273 }
2274 let completion_event = self.runs[run_index].enqueue(&dependencies, launch)?;
2275 let actual = completion_event.origin();
2276 if actual != completion.domain() {
2277 return Err(event_domain_error(
2278 EventDomainError::CompletionTokenDomainMismatch {
2279 operation: EventDomainOperation::ValidateCompletion,
2280 node_index: Some(node_index),
2281 expected: completion.domain(),
2282 actual,
2283 },
2284 ));
2285 }
2286 self.completions.insert(
2287 EventDependency::from_completion(completion),
2288 completion_event,
2289 );
2290 Ok(())
2291 }
2292
2293 fn classify_dependencies(
2294 &self,
2295 node_index: usize,
2296 node: &ScheduledNode,
2297 destination: super::EventDomainId,
2298 ) -> Result<SmallVec<[Arc<dyn EventToken>; 4]>> {
2299 let mut admitted = SmallVec::with_capacity(node.dependencies().len());
2300 for dependency in node.dependencies() {
2301 let dependency_completion =
2302 self.completions.get(dependency).cloned().ok_or_else(|| {
2303 Error::runtime_state_source(
2304 "Runtime::run_compiled",
2305 ErrorPhase::Execution,
2306 MissingScheduledDependencyCompletionError {
2307 dependency: *dependency,
2308 node_index,
2309 },
2310 )
2311 })?;
2312 let actual = dependency_completion.origin();
2313 if actual != dependency.domain() {
2314 return Err(event_domain_error(
2315 EventDomainError::DependencyDomainMismatch {
2316 operation: match node {
2317 ScheduledNode::Transfer(_) => EventDomainOperation::TransferBridge,
2318 ScheduledNode::Operation(_)
2319 | ScheduledNode::Collective(_)
2320 | ScheduledNode::Barrier(_) => EventDomainOperation::Enqueue,
2321 },
2322 node_index: Some(node_index),
2323 expected: dependency.domain(),
2324 actual,
2325 },
2326 ));
2327 }
2328 match node {
2329 ScheduledNode::Transfer(transfer) => {
2330 let source = transfer.source_event_domain();
2331 if actual == destination {
2332 admitted.push(dependency_completion);
2333 } else if actual == source {
2334 dependency_completion.wait().map_err(|source_error| {
2335 event_domain_error(EventDomainError::DependencyWaitFailed {
2336 operation: EventDomainOperation::TransferBridge,
2337 node_index: Some(node_index),
2338 expected: destination,
2339 actual,
2340 source: Box::new(source_error),
2341 })
2342 })?;
2343 } else {
2344 return Err(event_domain_error(
2345 EventDomainError::DependencyDomainMismatch {
2346 operation: EventDomainOperation::TransferBridge,
2347 node_index: Some(node_index),
2348 expected: source,
2349 actual,
2350 },
2351 ));
2352 }
2353 }
2354 ScheduledNode::Operation(_)
2355 | ScheduledNode::Collective(_)
2356 | ScheduledNode::Barrier(_) => {
2357 if actual != destination {
2358 return Err(event_domain_error(
2359 EventDomainError::DependencyDomainMismatch {
2360 operation: EventDomainOperation::Enqueue,
2361 node_index: Some(node_index),
2362 expected: destination,
2363 actual,
2364 },
2365 ));
2366 }
2367 admitted.push(dependency_completion);
2368 }
2369 }
2370 }
2371 Ok(admitted)
2372 }
2373
2374 fn run_index(&self, domain: super::EventDomainId) -> Result<usize> {
2375 if let Some(index) = self
2376 .runs
2377 .iter()
2378 .position(|run| run.requested_domain() == domain)
2379 {
2380 return Ok(index);
2381 }
2382 Err(missing_event_domain_driver(domain))
2383 }
2384
2385 pub(crate) fn drain(&mut self) -> Result<()> {
2386 let mut failures = Vec::new();
2387 for run in &mut self.runs {
2388 if let Err(error) = run.drain() {
2389 failures.push(error);
2390 }
2391 }
2392 let mut failures = failures.into_iter();
2393 let Some(mut error) = failures.next() else {
2394 return Ok(());
2395 };
2396 for failure in failures {
2397 error = Error::with_suppressed(error, failure);
2398 }
2399 Err(error)
2400 }
2401}
2402
2403#[derive(Debug)]
2404enum RuntimeOwnedEventDomainRunState {
2405 Pending(Box<dyn EventDomainRun>),
2406 Retired,
2407 Failed,
2408}
2409
2410#[derive(Clone, Copy, Debug, Eq, PartialEq)]
2411enum EventDomainRunTerminalState {
2412 Retired,
2413 Failed,
2414}
2415
2416impl fmt::Display for EventDomainRunTerminalState {
2417 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
2418 match self {
2419 Self::Retired => formatter.write_str("retired"),
2420 Self::Failed => formatter.write_str("failed"),
2421 }
2422 }
2423}
2424
2425#[derive(Debug)]
2426struct RuntimeOwnedEventDomainRun {
2427 requested_domain: super::EventDomainId,
2428 state: RuntimeOwnedEventDomainRunState,
2429}
2430
2431impl RuntimeOwnedEventDomainRun {
2432 fn new(requested_domain: super::EventDomainId, inner: Box<dyn EventDomainRun>) -> Self {
2433 Self {
2434 requested_domain,
2435 state: RuntimeOwnedEventDomainRunState::Pending(inner),
2436 }
2437 }
2438
2439 fn requested_domain(&self) -> super::EventDomainId {
2440 self.requested_domain
2441 }
2442
2443 fn domain(&self, operation: EventDomainOperation) -> Result<super::EventDomainId> {
2444 match &self.state {
2445 RuntimeOwnedEventDomainRunState::Pending(run) => Ok(run.domain()),
2446 RuntimeOwnedEventDomainRunState::Retired => Err(event_domain_run_state_error(
2447 operation,
2448 self.requested_domain,
2449 EventDomainRunTerminalState::Retired,
2450 )),
2451 RuntimeOwnedEventDomainRunState::Failed => Err(event_domain_run_state_error(
2452 operation,
2453 self.requested_domain,
2454 EventDomainRunTerminalState::Failed,
2455 )),
2456 }
2457 }
2458
2459 fn enqueue(
2460 &mut self,
2461 dependencies: &[Arc<dyn EventToken>],
2462 launch: &mut dyn FnMut() -> Result<()>,
2463 ) -> Result<Arc<dyn EventToken>> {
2464 match &mut self.state {
2465 RuntimeOwnedEventDomainRunState::Pending(run) => run.enqueue(dependencies, launch),
2466 RuntimeOwnedEventDomainRunState::Retired => Err(event_domain_run_state_error(
2467 EventDomainOperation::Enqueue,
2468 self.requested_domain,
2469 EventDomainRunTerminalState::Retired,
2470 )),
2471 RuntimeOwnedEventDomainRunState::Failed => Err(event_domain_run_state_error(
2472 EventDomainOperation::Enqueue,
2473 self.requested_domain,
2474 EventDomainRunTerminalState::Failed,
2475 )),
2476 }
2477 }
2478
2479 fn drain(&mut self) -> Result<()> {
2480 let run = match std::mem::replace(&mut self.state, RuntimeOwnedEventDomainRunState::Failed)
2481 {
2482 RuntimeOwnedEventDomainRunState::Pending(run) => run,
2483 RuntimeOwnedEventDomainRunState::Retired => {
2484 self.state = RuntimeOwnedEventDomainRunState::Retired;
2485 return Err(event_domain_run_state_error(
2486 EventDomainOperation::Drain,
2487 self.requested_domain,
2488 EventDomainRunTerminalState::Retired,
2489 ));
2490 }
2491 RuntimeOwnedEventDomainRunState::Failed => {
2492 self.state = RuntimeOwnedEventDomainRunState::Failed;
2493 return Err(event_domain_run_state_error(
2494 EventDomainOperation::Drain,
2495 self.requested_domain,
2496 EventDomainRunTerminalState::Failed,
2497 ));
2498 }
2499 };
2500 let domain = self.requested_domain;
2501 let mut run = run;
2502 let drain_result = catch_unwind(AssertUnwindSafe(|| run.drain()));
2503 drop_event_domain_run(run);
2504 match drain_result {
2505 Ok(Ok(())) => {
2506 self.state = RuntimeOwnedEventDomainRunState::Retired;
2507 Ok(())
2508 }
2509 Ok(Err(error)) => {
2510 self.state = RuntimeOwnedEventDomainRunState::Failed;
2511 Err(error)
2512 }
2513 Err(payload) => {
2514 self.state = RuntimeOwnedEventDomainRunState::Failed;
2515 Err(event_domain_error(EventDomainError::DrainPanicked {
2516 operation: EventDomainOperation::Drain,
2517 domain,
2518 message: safe_event_domain_panic_message(payload),
2519 }))
2520 }
2521 }
2522 }
2523}
2524
2525#[derive(Debug, thiserror::Error)]
2528#[error("{operation} used event-domain run {domain:?} after it reached terminal state {state}")]
2529pub(crate) struct EventDomainRunLifecycleError {
2530 operation: EventDomainOperation,
2531 domain: super::EventDomainId,
2532 state: EventDomainRunTerminalState,
2533}
2534
2535fn event_domain_run_state_error(
2536 operation: EventDomainOperation,
2537 domain: super::EventDomainId,
2538 state: EventDomainRunTerminalState,
2539) -> Error {
2540 Error::runtime_state_source(
2541 "Runtime::run_compiled",
2542 ErrorPhase::Execution,
2543 EventDomainRunLifecycleError {
2544 operation,
2545 domain,
2546 state,
2547 },
2548 )
2549}
2550
2551fn drop_event_domain_run(run: Box<dyn EventDomainRun>) {
2552 if catch_unwind(AssertUnwindSafe(|| drop(run))).is_err() {
2553 }
2555}
2556
2557impl Drop for RuntimeOwnedEventDomainRun {
2558 fn drop(&mut self) {
2559 let state = std::mem::replace(&mut self.state, RuntimeOwnedEventDomainRunState::Failed);
2560 if let RuntimeOwnedEventDomainRunState::Pending(run) = state {
2561 drop_event_domain_run(run);
2562 }
2563 }
2564}
2565
2566fn missing_event_domain_driver(domain: super::EventDomainId) -> Error {
2567 Error::from(EventDomainError::MissingDriver { domain })
2568}
2569
2570fn event_domain_error(source: EventDomainError) -> Error {
2571 Error::from(source)
2572}
2573
2574fn safe_event_domain_panic_message(payload: Box<dyn std::any::Any + Send + 'static>) -> String {
2575 match payload.downcast::<&'static str>() {
2576 Ok(message) => (*message).to_owned(),
2577 Err(payload) => match payload.downcast::<String>() {
2578 Ok(message) => *message,
2579 Err(_) => "non-string panic payload".to_owned(),
2580 },
2581 }
2582}
2583
2584fn scheduled_execution_cleanup_error(primary: Error, cleanup: Error) -> Error {
2585 Error::with_suppressed(primary, cleanup)
2586}
2587
2588fn instruction_execution<'a>(
2589 schedule: &'a ScheduledGraph,
2590 instruction: &ExecInstruction,
2591) -> Result<InstructionExecution<'a>> {
2592 let Some(operation_index) = instruction.semantic_operation_index else {
2593 let location = schedule.root_location();
2594 return Ok(InstructionExecution {
2595 witness: location.witness(),
2596 location,
2597 });
2598 };
2599 let location = schedule
2600 .operation_locations()
2601 .get(operation_index)
2602 .ok_or_else(|| {
2603 Error::runtime_state(
2604 "Runtime::run_compiled",
2605 ErrorPhase::Execution,
2606 format!(
2607 "instruction references semantic operation {operation_index}, but prepared schedule has {} operations",
2608 schedule.operation_locations().len()
2609 ),
2610 )
2611 })?;
2612 Ok(InstructionExecution {
2613 witness: location.witness(),
2614 location,
2615 })
2616}
2617
2618struct InstructionExecution<'a> {
2619 witness: &'a super::snapshot::ExecutableEngineSnapshot,
2620 location: &'a ExecutionLocation,
2621}
2622
2623impl InstructionExecution<'_> {
2624 fn executor(&self) -> &Arc<dyn ErasedTensorBackendExecutor> {
2625 self.witness.executor()
2626 }
2627
2628 fn location(&self) -> &ExecutionLocation {
2629 self.location
2630 }
2631}
2632
2633pub(crate) struct LocatedExecSlot<'input> {
2634 pub(crate) location: ExecutionLocation,
2635 pub(crate) value: ExecSlot<'input>,
2636}
2637
2638fn stage_instruction_inputs<'input>(
2639 instruction: &ExecInstruction,
2640 location: &ExecutionLocation,
2641 located: &mut [Vec<LocatedExecSlot<'input>>],
2642 staged: &mut [Option<ExecSlot<'input>>],
2643) -> Result<()> {
2644 for &slot in &instruction.input_slots {
2645 stage_slot_input(slot, location, located, staged)?;
2646 }
2647 Ok(())
2648}
2649
2650fn stage_slot_input<'input>(
2652 slot: usize,
2653 location: &ExecutionLocation,
2654 located: &mut [Vec<LocatedExecSlot<'input>>],
2655 staged: &mut [Option<ExecSlot<'input>>],
2656) -> Result<()> {
2657 if staged
2658 .get(slot)
2659 .ok_or(tenferro_tensor::Error::MissingValue { slot })?
2660 .is_some()
2661 {
2662 return Ok(());
2663 }
2664 let values = located
2665 .get_mut(slot)
2666 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
2667 let value_index = values
2668 .iter()
2669 .position(|value| &value.location == location)
2670 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
2671 staged[slot] = Some(values.swap_remove(value_index).value);
2672 Ok(())
2673}
2674
2675fn validate_region_outputs(
2678 region: &super::region::ElementwiseRegion,
2679 location: &ExecutionLocation,
2680 staged: &[Option<ExecSlot<'_>>],
2681) -> Result<()> {
2682 for &output_slot in ®ion.output_slots {
2683 let output = staged
2684 .get(output_slot)
2685 .and_then(Option::as_ref)
2686 .ok_or(tenferro_tensor::Error::MissingValue { slot: output_slot })?;
2687 let output = output.as_read();
2688 if !location
2689 .witness()
2690 .owns_resident_tensor(&output, location.storage_class())
2691 {
2692 return Err(Error::runtime_state_source(
2693 "Runtime::run_prepared",
2694 ErrorPhase::Execution,
2695 super::EngineExecutionContractError::OutputResidencyMismatch {
2696 instruction_index: region.instruction_range.start,
2697 output_slot,
2698 engine_id: location.engine_id().clone(),
2699 storage_class: location.storage_class().clone(),
2700 backend_family: output.backend_family(),
2701 allocation_domain: output.allocation_domain(),
2702 },
2703 ));
2704 }
2705 }
2706 Ok(())
2707}
2708
2709fn retain_region_results<'input>(
2712 region: &super::region::ElementwiseRegion,
2713 location: &ExecutionLocation,
2714 located: &mut [Vec<LocatedExecSlot<'input>>],
2715 staged: &mut [Option<ExecSlot<'input>>],
2716) -> Result<()> {
2717 for (&slot, &last_use) in region.input_slots.iter().zip(region.input_last_use.iter()) {
2718 if last_use {
2719 located[slot].clear();
2720 staged[slot].take();
2721 }
2722 }
2723 for &slot in ®ion.output_slots {
2724 let value = staged
2725 .get_mut(slot)
2726 .and_then(Option::take)
2727 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
2728 let values = located
2729 .get_mut(slot)
2730 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
2731 values.clear();
2732 values.push(LocatedExecSlot {
2733 location: location.clone(),
2734 value,
2735 });
2736 }
2737 Ok(())
2738}
2739
2740fn validate_instruction_outputs(
2741 instruction_index: usize,
2742 instruction: &ExecInstruction,
2743 location: &ExecutionLocation,
2744 staged: &[Option<ExecSlot<'_>>],
2745) -> Result<()> {
2746 for &output_slot in &instruction.output_slots {
2747 let output = staged
2748 .get(output_slot)
2749 .and_then(Option::as_ref)
2750 .ok_or(tenferro_tensor::Error::MissingValue { slot: output_slot })?;
2751 let output = output.as_read();
2752 if !location
2753 .witness()
2754 .owns_resident_tensor(&output, location.storage_class())
2755 {
2756 return Err(Error::runtime_state_source(
2757 "Runtime::run_compiled",
2758 ErrorPhase::Execution,
2759 super::EngineExecutionContractError::OutputResidencyMismatch {
2760 instruction_index,
2761 output_slot,
2762 engine_id: location.engine_id().clone(),
2763 storage_class: location.storage_class().clone(),
2764 backend_family: output.backend_family(),
2765 allocation_domain: output.allocation_domain(),
2766 },
2767 ));
2768 }
2769 }
2770 Ok(())
2771}
2772
2773pub(crate) fn retain_instruction_results<'input>(
2774 instruction: &ExecInstruction,
2775 location: &ExecutionLocation,
2776 located: &mut [Vec<LocatedExecSlot<'input>>],
2777 staged: &mut [Option<ExecSlot<'input>>],
2778) -> Result<()> {
2779 for &slot in &instruction.input_slots {
2780 let is_output = instruction.output_slots.contains(&slot);
2781 let is_last_use = instruction
2782 .input_slots
2783 .iter()
2784 .enumerate()
2785 .any(|(index, &candidate)| {
2786 candidate == slot && instruction.last_use.get(index).copied().unwrap_or(false)
2787 });
2788 if is_last_use {
2789 located[slot].clear();
2790 if !is_output {
2791 staged[slot].take();
2792 }
2793 } else if !is_output && let Some(value) = staged[slot].take() {
2794 located[slot].push(LocatedExecSlot {
2795 location: location.clone(),
2796 value,
2797 });
2798 }
2799 }
2800
2801 for &slot in &instruction.output_slots {
2802 let value = staged
2803 .get_mut(slot)
2804 .and_then(Option::take)
2805 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
2806 let values = located
2807 .get_mut(slot)
2808 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
2809 values.clear();
2810 values.push(LocatedExecSlot {
2811 location: location.clone(),
2812 value,
2813 });
2814 }
2815 Ok(())
2816}
2817
2818fn validate_runtime_input_ingress(
2819 location: &ExecutionLocation,
2820 input: &TensorRead<'_>,
2821 slot: usize,
2822) -> Result<()> {
2823 let accepted = location
2824 .witness()
2825 .accepts_runtime_input(input, location.storage_class());
2826 if accepted {
2827 return Ok(());
2828 }
2829 Err(Error::runtime_state_source(
2830 "Runtime::run_compiled",
2831 ErrorPhase::Execution,
2832 super::InputIngressContractError::ResidencyMismatch {
2833 input_slot: slot,
2834 ingress_engine_id: location.engine_id().clone(),
2835 ingress_storage_class: location.storage_class().clone(),
2836 placement: input.placement().clone(),
2837 backend_family: input.backend_family(),
2838 allocation_domain: input.allocation_domain(),
2839 },
2840 ))
2841}
2842
2843fn execute_scheduled_transfer<'input>(
2844 transfer: &ScheduledTransfer,
2845 located: &mut [Vec<LocatedExecSlot<'input>>],
2846) -> Result<()> {
2847 let source = transfer.source_location();
2848 let destination = transfer.destination_location();
2849 let provider = transfer.provider();
2850 let values =
2851 located
2852 .get_mut(transfer.value_slot())
2853 .ok_or(tenferro_tensor::Error::MissingValue {
2854 slot: transfer.value_slot(),
2855 })?;
2856 if values.iter().any(|value| &value.location == destination) {
2857 return Err(Error::runtime_state_source(
2858 "Runtime::run_compiled",
2859 ErrorPhase::Execution,
2860 DuplicateTransferDestinationError {
2861 value_slot: transfer.value_slot(),
2862 destination: destination.clone(),
2863 },
2864 ));
2865 }
2866 let transferred = {
2867 let source_value = values
2868 .iter()
2869 .find(|value| &value.location == source)
2870 .ok_or(tenferro_tensor::Error::MissingValue {
2871 slot: transfer.value_slot(),
2872 })?;
2873 let source_read = source_value.value.as_read();
2874 let expected_dtype = source_read.dtype();
2875 let expected_shape = source_read.shape().to_vec();
2876 let transferred =
2877 provider.transfer_blocking(TransferRequest::new(source, destination, source_read))?;
2878 validate_transfer_output(destination, expected_dtype, &expected_shape, &transferred)?;
2879 transferred
2880 };
2881 values.push(LocatedExecSlot {
2882 location: destination.clone(),
2883 value: ExecSlot::Owned(transferred),
2884 });
2885 Ok(())
2886}
2887
2888fn validate_transfer_output(
2889 destination: &ExecutionLocation,
2890 expected_dtype: tenferro_tensor::DType,
2891 expected_shape: &[usize],
2892 output: &Tensor,
2893) -> Result<()> {
2894 let expected_elements = checked_transfer_element_count(expected_shape).map_err(|source| {
2895 Error::runtime_state_source(
2896 "Runtime::run_compiled",
2897 ErrorPhase::Execution,
2898 TransferError::ProviderContract { source },
2899 )
2900 })?;
2901 let contract_error = if output.dtype() != expected_dtype {
2902 Some(TransferProviderContractError::DTypeMismatch {
2903 expected: expected_dtype,
2904 actual: output.dtype(),
2905 })
2906 } else if output.shape() != expected_shape {
2907 Some(TransferProviderContractError::ShapeMismatch {
2908 expected: expected_shape.to_vec(),
2909 actual: output.shape().to_vec(),
2910 })
2911 } else if tensor_buffer_len(output) != expected_elements {
2912 Some(TransferProviderContractError::InvalidBufferLength {
2913 expected: expected_elements,
2914 actual: tensor_buffer_len(output),
2915 })
2916 } else if !destination
2917 .witness()
2918 .accepts_input_placement(output.placement(), destination.storage_class())
2919 {
2920 Some(
2921 TransferProviderContractError::DestinationPlacementMismatch {
2922 destination_engine_id: destination.engine_id().clone(),
2923 destination_storage_class: destination.storage_class().clone(),
2924 actual: output.placement().clone(),
2925 },
2926 )
2927 } else if !destination.witness().owns_resident_tensor(
2928 &TensorRead::from_tensor(output),
2929 destination.storage_class(),
2930 ) {
2931 Some(
2932 TransferProviderContractError::DestinationResidencyMismatch {
2933 destination_engine_id: destination.engine_id().clone(),
2934 destination_storage_class: destination.storage_class().clone(),
2935 actual_backend_family: TensorRead::from_tensor(output).backend_family(),
2936 actual_allocation_domain: TensorRead::from_tensor(output).allocation_domain(),
2937 },
2938 )
2939 } else {
2940 None
2941 };
2942 match contract_error {
2943 None => Ok(()),
2944 Some(source) => Err(Error::runtime_state_source(
2945 "Runtime::run_compiled",
2946 ErrorPhase::Execution,
2947 TransferError::ProviderContract { source },
2948 )),
2949 }
2950}
2951
2952fn checked_transfer_element_count(
2953 shape: &[usize],
2954) -> std::result::Result<usize, TransferProviderContractError> {
2955 tenferro_tensor::validate::checked_shape_product(
2956 "Runtime::run_compiled",
2957 "transfer source shape",
2958 shape,
2959 )
2960 .map_err(|source| TransferProviderContractError::LogicalElementCount { source })
2961}
2962
2963fn tensor_buffer_len(tensor: &Tensor) -> usize {
2964 match tensor.dtype() {
2965 DType::F32 => tensor.as_typed::<f32>().map_or(0, |t| t.buffer().len()),
2966 DType::F64 => tensor.as_typed::<f64>().map_or(0, |t| t.buffer().len()),
2967 DType::I32 => tensor.as_typed::<i32>().map_or(0, |t| t.buffer().len()),
2968 DType::I64 => tensor.as_typed::<i64>().map_or(0, |t| t.buffer().len()),
2969 DType::Bool => tensor.as_typed::<bool>().map_or(0, |t| t.buffer().len()),
2970 DType::C32 => tensor
2971 .as_typed::<Complex32>()
2972 .map_or(0, |t| t.buffer().len()),
2973 DType::C64 => tensor
2974 .as_typed::<Complex64>()
2975 .map_or(0, |t| t.buffer().len()),
2976 DType::External(_) => 0,
2980 }
2981}
2982
2983#[cfg(test)]
2984mod transfer_validation_tests {
2985 use std::error::Error as _;
2986
2987 use super::checked_transfer_element_count;
2988 use crate::TransferProviderContractError;
2989
2990 #[test]
2991 fn transfer_element_count_overflow_is_typed_and_preserves_source() {
2992 let error = checked_transfer_element_count(&[usize::MAX, 2]).unwrap_err();
2993
2994 assert!(matches!(
2995 error,
2996 TransferProviderContractError::LogicalElementCount { .. }
2997 ));
2998 assert!(error.source().is_some());
2999 }
3000}
3001
3002fn collect_located_outputs<'input>(
3003 program: &ExecProgram,
3004 located: &mut [Vec<LocatedExecSlot<'input>>],
3005) -> Result<Vec<Option<LocatedExecSlot<'input>>>> {
3006 let mut outputs = (0..program.n_slots).map(|_| None).collect::<Vec<_>>();
3007 for &slot in &program.output_slots {
3008 if outputs[slot].is_some() {
3009 continue;
3010 }
3011 let values = located
3012 .get_mut(slot)
3013 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
3014 let value = values
3015 .pop()
3016 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
3017 values.clear();
3018 outputs[slot] = Some(value);
3019 }
3020 located.iter_mut().for_each(Vec::clear);
3021 Ok(outputs)
3022}
3023
3024pub(super) fn collect_tensor_outputs_with<'input>(
3025 program: &ExecProgram,
3026 outputs: &mut [Option<LocatedExecSlot<'input>>],
3027 mut materialize: impl FnMut(&ExecutionLocation, ExecSlot<'input>) -> Result<Tensor>,
3028) -> Result<Vec<Tensor>> {
3029 program
3030 .output_slots
3031 .iter()
3032 .map(|&slot| {
3033 let located = outputs
3034 .get_mut(slot)
3035 .and_then(Option::take)
3036 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
3037 materialize(&located.location, located.value)
3038 })
3039 .collect()
3040}
3041
3042fn collect_value_outputs_with<'input>(
3043 program: &ExecProgram,
3044 outputs: &mut [Option<LocatedExecSlot<'input>>],
3045 mut materialize: impl FnMut(&ExecutionLocation, ExecSlot<'input>) -> Result<TensorValue>,
3046) -> Result<Vec<TensorValue>> {
3047 program
3048 .output_slots
3049 .iter()
3050 .map(|&slot| {
3051 let located = outputs
3052 .get_mut(slot)
3053 .and_then(Option::take)
3054 .ok_or(tenferro_tensor::Error::MissingValue { slot })?;
3055 materialize(&located.location, located.value)
3056 })
3057 .collect()
3058}
3059
3060fn input_signature_reads(inputs: &[TensorRead<'_>]) -> Result<InputSignature> {
3061 InputSignature::from_reads(inputs).map_err(|source| prepare_error(Arc::new(source)))
3062}
3063
3064fn resolve_input_refs<'a>(
3065 program: &'a CompiledGraph,
3066 inputs: &'a [&'a Tensor],
3067) -> Result<RuntimeInputReads<'a>> {
3068 let expected = program.program().inputs().len();
3069 if inputs.len() > expected {
3070 return Err(Error::GraphInputCountMismatch {
3071 expected,
3072 actual: inputs.len(),
3073 });
3074 }
3075 let resolved = if inputs.is_empty() {
3076 semantic_default_inputs(program)?
3077 } else if inputs.len() == expected {
3078 inputs
3079 .iter()
3080 .map(|tensor| TensorRead::from_tensor(tensor))
3081 .collect()
3082 } else {
3083 let mut explicit = inputs.iter();
3084 let mut resolved = RuntimeInputReads::new();
3085 for value in program.program().inputs() {
3086 if let Some(retained) = program.bindings().tensor_ref_for_input(*value) {
3087 resolved.push(retained.tensor_read()?);
3088 } else {
3089 let tensor = explicit.next().ok_or_else(|| {
3090 Error::invalid_argument(
3091 "Runtime::resolve_input_refs",
3092 ErrorPhase::Execution,
3093 "inputs",
3094 format!(
3095 "expected {expected} inputs, or one explicit input for each unbound placeholder"
3096 ),
3097 )
3098 })?;
3099 resolved.push(TensorRead::from_tensor(tensor));
3100 }
3101 }
3102 if explicit.next().is_some() {
3103 return Err(Error::invalid_argument(
3104 "Runtime::resolve_input_refs",
3105 ErrorPhase::Execution,
3106 "inputs",
3107 format!("expected {expected} inputs, received {}", inputs.len()),
3108 ));
3109 }
3110 resolved
3111 };
3112 validate_ordered_input_metadata_reads(program, &resolved)?;
3113 Ok(resolved)
3114}
3115
3116fn semantic_default_inputs(program: &CompiledGraph) -> Result<RuntimeInputReads<'_>> {
3117 program
3118 .program()
3119 .inputs()
3120 .iter()
3121 .enumerate()
3122 .map(|(input_index, value)| {
3123 let retained = program
3124 .bindings()
3125 .tensor_ref_for_input(*value)
3126 .ok_or_else(|| Error::UnboundPlaceholder {
3127 input_key: format!("semantic input {input_index}"),
3128 })?;
3129 retained.tensor_read()
3130 })
3131 .collect()
3132}
3133
3134fn validate_ordered_input_metadata_reads(
3135 program: &CompiledGraph,
3136 inputs: &[TensorRead<'_>],
3137) -> Result<()> {
3138 let actuals = inputs
3139 .iter()
3140 .map(|input| (input.dtype(), input.shape().to_vec()))
3141 .collect::<Vec<_>>();
3142 validate_ordered_input_metadata_values(program, &actuals, "Runtime::submit")
3143}
3144
3145fn validate_ordered_input_metadata_values(
3146 program: &CompiledGraph,
3147 actuals: &[(tenferro_tensor::DType, Vec<usize>)],
3148 caller: &'static str,
3149) -> Result<()> {
3150 let expected = program.input_count();
3151 if actuals.len() != expected {
3152 return Err(Error::GraphInputCountMismatch {
3153 expected,
3154 actual: actuals.len(),
3155 });
3156 }
3157 let input_shapes: RuntimeInputShapes<'_> =
3158 actuals.iter().map(|(_, shape)| shape.as_slice()).collect();
3159 for (input_value, (actual_dtype, actual_shape)) in
3160 program.program().inputs().iter().zip(actuals)
3161 {
3162 let metadata = program
3163 .program()
3164 .value_metadata(*input_value)
3165 .map_err(|source| Error::runtime_state_source(caller, ErrorPhase::Execution, source))?;
3166 if metadata.dtype() != *actual_dtype {
3167 return Err(Error::PlaceholderDtypeMismatch {
3168 expected: metadata.dtype(),
3169 actual: *actual_dtype,
3170 });
3171 }
3172 if metadata.shape().len() != actual_shape.len() {
3173 return Err(Error::PlaceholderRankMismatch {
3174 expected: metadata.shape().len(),
3175 actual: actual_shape.len(),
3176 });
3177 }
3178 let mut expected_shape: RuntimeShapeScratch = actual_shape.iter().copied().collect();
3179 let mut exact_mismatch = false;
3180 for (axis, (extent, actual_size)) in metadata.shape().iter().zip(actual_shape).enumerate() {
3181 match extent {
3182 ShapeExtent::Exact(expression) => {
3183 let expected = expression.eval(&input_shapes).map_err(|source| {
3184 Error::runtime_state_source(caller, ErrorPhase::Execution, source)
3185 })?;
3186 expected_shape[axis] = expected;
3187 exact_mismatch |= expected != *actual_size;
3188 }
3189 ShapeExtent::UpperBound(expression) => {
3190 let bound = expression.eval(&input_shapes).map_err(|source| {
3191 Error::runtime_state_source(caller, ErrorPhase::Execution, source)
3192 })?;
3193 if *actual_size > bound {
3194 return Err(Error::PlaceholderShapeBoundExceeded {
3195 axis,
3196 bound,
3197 actual: *actual_size,
3198 });
3199 }
3200 }
3201 ShapeExtent::Unknown => {}
3202 }
3203 }
3204 if exact_mismatch {
3205 return Err(Error::PlaceholderShapeMismatch {
3206 expected: expected_shape.into_vec(),
3207 actual: actual_shape.clone(),
3208 });
3209 }
3210 }
3211 Ok(())
3212}
3213
3214fn validate_exec_input_count(program: &ExecProgram, actual: usize) -> Result<()> {
3215 let expected = program.input_slots.len();
3216 if actual != expected {
3217 return Err(Error::runtime_state(
3218 "Runtime::run_compiled",
3219 ErrorPhase::Execution,
3220 format!("expected {expected} inputs for execution program, got {actual}"),
3221 ));
3222 }
3223 Ok(())
3224}
3225
3226fn prepare_error(source: Arc<PrepareError>) -> Error {
3227 Error::runtime_state_source(
3228 "Runtime::run_compiled",
3229 ErrorPhase::Execution,
3230 SharedPrepareError(source),
3231 )
3232}
3233
3234#[derive(Clone, Debug)]
3235struct SharedPrepareError(Arc<PrepareError>);
3236
3237impl fmt::Display for SharedPrepareError {
3238 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
3239 write!(formatter, "{}", self.0)
3240 }
3241}
3242
3243impl StdError for SharedPrepareError {
3244 fn source(&self) -> Option<&(dyn StdError + 'static)> {
3245 Some(self.0.as_ref())
3246 }
3247}
3248
3249#[cfg(test)]
3250mod tests {
3251 use std::any::Any;
3252 use std::error::Error as StdError;
3253 use std::hash::Hasher;
3254 use std::num::NonZeroU64;
3255 use std::sync::atomic::{AtomicBool, Ordering};
3256 use std::sync::{Arc, TryLockError, Weak};
3257
3258 use tenferro_cpu::CpuBackend;
3259 use tenferro_ops::dim_expr::DimExpr;
3260 use tenferro_ops::ext_op::ExtensionOp;
3261 use tenferro_ops::SymDim;
3262 use tenferro_tensor::{
3263 BackendSession, BackendSessionHost, DType, Tensor, TensorBackend, TensorRead, TypedTensor,
3264 };
3265
3266 use crate::exec::{ExecInstruction, ExecOp, ExecProgram, ExecSlot};
3267 use crate::runtime::{
3268 CoreCapabilityBundle, EngineId, ErasedExecutionContext, EventDomainDriver, EventDomainId,
3269 ExecutableEngineContract, ExecutionContextIdentity, HardwareClassId, InputIngressContract,
3270 InputSignature, PreparedOperation, PreparedOperationBinding, PreparedOperationExecutor,
3271 PreparedOperationPlan, ProviderDeviceIdentity, ProviderId, RegistrationIdentity,
3272 RuntimeCacheOwner, RuntimeEpoch, RuntimeId, SpecializationProjection,
3273 SpecializationRequirements, StorageClass,
3274 };
3275 use crate::{Error, ErrorPhase, ExtensionCacheStore, Result};
3276
3277 use super::{
3278 DuplicateTransferDestinationError, ErasedTensorBackendExecutor, LocatedExecSlot,
3279 TensorBackendExecutor,
3280 };
3281 use crate::runtime::schedule::{
3282 EventCompletion, EventSlotId, ExecutionLocation, ScheduledTransfer,
3283 };
3284
3285 const LOCK_PROBE_FAMILY: &str = "runtime.lock-probe.v1";
3286 const REENTRANT_PROBE_FAMILY: &str = "runtime.reentrant-probe.v1";
3287
3288 #[derive(Clone, Debug)]
3289 struct LockProbeOp;
3290
3291 impl ExtensionOp for LockProbeOp {
3292 fn family_id(&self) -> &'static str {
3293 LOCK_PROBE_FAMILY
3294 }
3295
3296 fn payload_hash(&self, _hasher: &mut dyn Hasher) {}
3297
3298 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
3299 other.as_any().downcast_ref::<Self>().is_some()
3300 }
3301
3302 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
3303 Arc::new(self.clone())
3304 }
3305
3306 fn as_any(&self) -> &dyn Any {
3307 self
3308 }
3309
3310 fn input_count(&self) -> usize {
3311 1
3312 }
3313
3314 fn output_count(&self) -> usize {
3315 1
3316 }
3317
3318 fn infer_output_meta(
3319 &self,
3320 ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
3321 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
3322 Ok(vec![(ctx.input_dtype(0)?, ctx.input_shape(0)?.to_vec())])
3323 }
3324 }
3325
3326 #[derive(Clone, Debug)]
3327 struct ReentrantProbeOp;
3328
3329 impl ExtensionOp for ReentrantProbeOp {
3330 fn family_id(&self) -> &'static str {
3331 REENTRANT_PROBE_FAMILY
3332 }
3333
3334 fn payload_hash(&self, _hasher: &mut dyn Hasher) {}
3335
3336 fn payload_eq(&self, other: &dyn ExtensionOp) -> bool {
3337 other.as_any().downcast_ref::<Self>().is_some()
3338 }
3339
3340 fn clone_arc(&self) -> Arc<dyn ExtensionOp> {
3341 Arc::new(self.clone())
3342 }
3343
3344 fn as_any(&self) -> &dyn Any {
3345 self
3346 }
3347
3348 fn input_count(&self) -> usize {
3349 1
3350 }
3351
3352 fn output_count(&self) -> usize {
3353 1
3354 }
3355
3356 fn infer_output_meta(
3357 &self,
3358 ctx: &mut tenferro_ops::ExtensionShapeContext<'_>,
3359 ) -> tenferro_tensor::Result<Vec<(DType, Vec<SymDim>)>> {
3360 Ok(vec![(ctx.input_dtype(0)?, ctx.input_shape(0)?.to_vec())])
3361 }
3362 }
3363
3364 #[derive(Debug)]
3365 struct LockProbePreparedOperation {
3366 binding: PreparedOperationBinding,
3367 specialization: SpecializationProjection,
3368 executor: Weak<TensorBackendExecutor<CpuBackend>>,
3369 observed_unlocked_state: Arc<AtomicBool>,
3370 }
3371
3372 impl PreparedOperation for LockProbePreparedOperation {
3373 fn binding(&self) -> &PreparedOperationBinding {
3374 &self.binding
3375 }
3376
3377 fn specialization(&self) -> &SpecializationProjection {
3378 &self.specialization
3379 }
3380
3381 fn retained_bytes(&self) -> usize {
3382 0
3383 }
3384 }
3385
3386 impl PreparedOperationExecutor for LockProbePreparedOperation {
3387 fn execute(
3388 &self,
3389 context: &mut ErasedExecutionContext<'_>,
3390 _extension_caches: &mut ExtensionCacheStore,
3391 inputs: &[TensorRead<'_>],
3392 ) -> Result<Vec<Tensor>> {
3393 let executor = self.executor.upgrade().expect("executor still alive");
3394 let unlocked = match executor.state.try_lock() {
3395 Ok(_guard) => true,
3396 Err(TryLockError::WouldBlock) => false,
3397 Err(TryLockError::Poisoned(_)) => false,
3398 };
3399 self.observed_unlocked_state
3400 .store(unlocked, Ordering::SeqCst);
3401 let backend = context
3402 .downcast_mut::<CpuBackend>(self.binding.context_identity())
3403 .map_err(|source| {
3404 Error::runtime_state_source("lock_probe", ErrorPhase::Execution, source)
3405 })?;
3406 Ok(vec![backend.with_backend_session(|exec| {
3407 exec.to_contiguous_read(inputs[0].clone())
3408 })??])
3409 }
3410 }
3411
3412 #[derive(Debug)]
3413 struct ReentrantProbePreparedOperation {
3414 binding: PreparedOperationBinding,
3415 specialization: SpecializationProjection,
3416 executor: Weak<TensorBackendExecutor<CpuBackend>>,
3417 observed_reentrant_error: Arc<AtomicBool>,
3418 }
3419
3420 #[derive(Debug)]
3421 struct SessionProbePreparedOperation {
3422 binding: PreparedOperationBinding,
3423 specialization: SpecializationProjection,
3424 observed_session: Arc<AtomicBool>,
3425 }
3426
3427 impl PreparedOperation for SessionProbePreparedOperation {
3428 fn binding(&self) -> &PreparedOperationBinding {
3429 &self.binding
3430 }
3431
3432 fn specialization(&self) -> &SpecializationProjection {
3433 &self.specialization
3434 }
3435
3436 fn retained_bytes(&self) -> usize {
3437 0
3438 }
3439 }
3440
3441 impl PreparedOperationExecutor for SessionProbePreparedOperation {
3442 fn execute(
3443 &self,
3444 _context: &mut ErasedExecutionContext<'_>,
3445 _extension_caches: &mut ExtensionCacheStore,
3446 _inputs: &[TensorRead<'_>],
3447 ) -> Result<Vec<Tensor>> {
3448 Err(Error::unsupported(
3449 "session_probe",
3450 ErrorPhase::Execution,
3451 "session probe must use the scheduler-owned session path",
3452 ))
3453 }
3454
3455 fn supports_session(&self) -> bool {
3456 true
3457 }
3458
3459 fn execute_in_session(
3460 &self,
3461 session: &mut dyn BackendSession,
3462 _extension_caches: &mut ExtensionCacheStore,
3463 inputs: &[TensorRead<'_>],
3464 ) -> Result<Vec<Tensor>> {
3465 self.observed_session.store(true, Ordering::SeqCst);
3466 Ok(vec![session.to_contiguous_read(inputs[0].clone())?])
3467 }
3468 }
3469
3470 impl ReentrantProbePreparedOperation {
3471 fn probe_reentrant_call(&self, input: Tensor) {
3472 let executor = self.executor.upgrade().expect("executor still alive");
3473 let nested = ErasedTensorBackendExecutor::execute(
3474 executor.as_ref(),
3475 &passthrough_program(),
3476 &[],
3477 vec![input],
3478 );
3479 let observed = nested.is_err_and(|error| {
3480 error
3481 .to_string()
3482 .contains("reentrant tensor backend executor call would deadlock")
3483 });
3484 self.observed_reentrant_error
3485 .store(observed, Ordering::SeqCst);
3486 }
3487 }
3488
3489 impl PreparedOperation for ReentrantProbePreparedOperation {
3490 fn binding(&self) -> &PreparedOperationBinding {
3491 &self.binding
3492 }
3493
3494 fn specialization(&self) -> &SpecializationProjection {
3495 &self.specialization
3496 }
3497
3498 fn retained_bytes(&self) -> usize {
3499 0
3500 }
3501 }
3502
3503 impl PreparedOperationExecutor for ReentrantProbePreparedOperation {
3504 fn execute(
3505 &self,
3506 context: &mut ErasedExecutionContext<'_>,
3507 _extension_caches: &mut ExtensionCacheStore,
3508 inputs: &[TensorRead<'_>],
3509 ) -> Result<Vec<Tensor>> {
3510 let backend = context
3511 .downcast_mut::<CpuBackend>(self.binding.context_identity())
3512 .map_err(|source| {
3513 Error::runtime_state_source("reentrant_probe", ErrorPhase::Execution, source)
3514 })?;
3515 let materialized = backend
3516 .with_backend_session(|exec| exec.to_contiguous_read(inputs[0].clone()))??;
3517 self.probe_reentrant_call(materialized.duplicate()?);
3518 Ok(vec![materialized])
3519 }
3520 }
3521
3522 fn lock_probe_program() -> ExecProgram {
3523 ExecProgram {
3524 instructions: vec![ExecInstruction {
3525 op: ExecOp::Extension(Arc::new(LockProbeOp)),
3526 semantic_operation_index: Some(0),
3527 input_slots: vec![0],
3528 output_slots: vec![1],
3529 dtype: DType::F64,
3530 output_shapes: vec![vec![DimExpr::InputDim {
3531 input_idx: 0,
3532 axis: 0,
3533 }]]
3534 .into(),
3535 output_extents: vec![vec![]].into(),
3536 last_use: vec![false],
3537 }],
3538 input_slots: vec![0],
3539 output_slots: vec![1],
3540 n_slots: 2,
3541 shape_guards: vec![],
3542 }
3543 }
3544
3545 fn partial_session_probe_program() -> ExecProgram {
3546 let output_shape = || {
3547 vec![DimExpr::InputDim {
3548 input_idx: 0,
3549 axis: 0,
3550 }]
3551 };
3552 ExecProgram {
3553 instructions: vec![
3554 ExecInstruction {
3555 op: ExecOp::Extension(Arc::new(LockProbeOp)),
3556 semantic_operation_index: Some(0),
3557 input_slots: vec![0],
3558 output_slots: vec![1],
3559 dtype: DType::F64,
3560 output_shapes: vec![output_shape()].into(),
3561 output_extents: vec![vec![]].into(),
3562 last_use: vec![false],
3563 },
3564 ExecInstruction {
3565 op: ExecOp::Extension(Arc::new(LockProbeOp)),
3566 semantic_operation_index: Some(1),
3567 input_slots: vec![1],
3568 output_slots: vec![2],
3569 dtype: DType::F64,
3570 output_shapes: vec![output_shape()].into(),
3571 output_extents: vec![vec![]].into(),
3572 last_use: vec![false],
3573 },
3574 ],
3575 input_slots: vec![0],
3576 output_slots: vec![2],
3577 n_slots: 3,
3578 shape_guards: vec![],
3579 }
3580 }
3581
3582 fn reentrant_probe_program() -> ExecProgram {
3583 ExecProgram {
3584 instructions: vec![ExecInstruction {
3585 op: ExecOp::Extension(Arc::new(ReentrantProbeOp)),
3586 semantic_operation_index: Some(0),
3587 input_slots: vec![0],
3588 output_slots: vec![1],
3589 dtype: DType::F64,
3590 output_shapes: vec![vec![DimExpr::InputDim {
3591 input_idx: 0,
3592 axis: 0,
3593 }]]
3594 .into(),
3595 output_extents: vec![vec![]].into(),
3596 last_use: vec![false],
3597 }],
3598 input_slots: vec![0],
3599 output_slots: vec![1],
3600 n_slots: 2,
3601 shape_guards: vec![],
3602 }
3603 }
3604
3605 fn passthrough_program() -> ExecProgram {
3606 ExecProgram {
3607 instructions: vec![],
3608 input_slots: vec![0],
3609 output_slots: vec![0],
3610 n_slots: 1,
3611 shape_guards: vec![],
3612 }
3613 }
3614
3615 fn f64_zeros(shape: Vec<usize>) -> Tensor {
3616 Tensor::from_typed::<f64>(TypedTensor::zeros(shape).unwrap())
3617 }
3618
3619 fn nz(value: u64) -> NonZeroU64 {
3620 NonZeroU64::new(value).unwrap_or(NonZeroU64::MIN)
3621 }
3622
3623 fn probe_binding() -> PreparedOperationBinding {
3624 PreparedOperationBinding::new(
3625 RuntimeId::from_nonzero(nz(1)),
3626 RuntimeEpoch::from_nonzero(nz(2)),
3627 EngineId::new("tenferro.cpu").unwrap(),
3628 RegistrationIdentity::new(nz(3), nz(4)),
3629 ExecutionContextIdentity::of::<CpuBackend>(),
3630 HardwareClassId::new("tenferro.cpu.host").unwrap(),
3631 )
3632 }
3633
3634 fn probe_specialization() -> SpecializationProjection {
3635 SpecializationRequirements::polymorphic(0)
3636 .project(&InputSignature::new(Vec::new()))
3637 .unwrap()
3638 }
3639
3640 #[test]
3641 fn tensor_backend_executor_releases_state_lock_during_extension_execution() {
3642 let executor = Arc::new(TensorBackendExecutor::<CpuBackend>::new(CpuBackend::new()));
3643 let observed_unlocked_state = Arc::new(AtomicBool::new(false));
3644 let prepared = Arc::new(LockProbePreparedOperation {
3645 binding: probe_binding(),
3646 specialization: probe_specialization(),
3647 executor: Arc::downgrade(&executor),
3648 observed_unlocked_state: Arc::clone(&observed_unlocked_state),
3649 });
3650 let operations = vec![PreparedOperationPlan::executable(
3651 prepared.clone(),
3652 prepared,
3653 )];
3654
3655 let output = ErasedTensorBackendExecutor::execute(
3656 executor.as_ref(),
3657 &lock_probe_program(),
3658 &operations,
3659 vec![f64_zeros(vec![2])],
3660 )
3661 .expect("extension executes");
3662
3663 assert_eq!(output[0].shape(), &[2]);
3664 assert!(
3665 observed_unlocked_state.load(Ordering::SeqCst),
3666 "executor state lock must not be held while extension runtime callbacks execute"
3667 );
3668 }
3669
3670 #[test]
3671 fn scheduled_execution_cleanup_error_preserves_primary_error() {
3672 let primary = Error::runtime_state(
3673 "primary-execution",
3674 ErrorPhase::Execution,
3675 "primary execution failure",
3676 );
3677 let cleanup = Error::runtime_state(
3678 "event-domain-cleanup",
3679 ErrorPhase::Execution,
3680 "cleanup failure",
3681 );
3682 let primary_display = primary.to_string();
3683
3684 let combined = super::scheduled_execution_cleanup_error(primary, cleanup);
3685 let primary_error = combined.primary().expect("primary execution error");
3686 let cleanup_error = combined.suppressed().expect("cleanup error");
3687 assert_eq!(primary_error.to_string(), primary_display);
3688 assert_eq!(
3689 cleanup_error.to_string(),
3690 "event-domain-cleanup (Execution): runtime state failure: cleanup failure"
3691 );
3692 assert_eq!(
3693 StdError::source(&combined)
3694 .expect("primary error in the standard source chain")
3695 .to_string(),
3696 primary_display
3697 );
3698 }
3699
3700 #[test]
3701 fn duplicate_scheduled_transfer_destination_reports_typed_fields() {
3702 let domain = EventDomainId::runtime_created_for_test(
3703 RuntimeId::from_nonzero(nz(1)),
3704 RuntimeEpoch::from_nonzero(nz(1)),
3705 RegistrationIdentity::new(nz(1), nz(1)),
3706 );
3707 let source = ExecutionLocation::new(
3708 EngineId::new("tenferro-test.transfer-source").expect("source engine"),
3709 ProviderDeviceIdentity::new(
3710 ProviderId::new("tenferro-test.transfer").expect("provider id"),
3711 "source",
3712 )
3713 .expect("source provider target"),
3714 domain,
3715 StorageClass::new("tenferro-test.transfer-source").expect("source storage"),
3716 );
3717 let destination = ExecutionLocation::new(
3718 EngineId::new("tenferro-test.transfer-destination").expect("destination engine"),
3719 ProviderDeviceIdentity::new(
3720 ProviderId::new("tenferro-test.transfer").expect("provider id"),
3721 "destination",
3722 )
3723 .expect("destination provider target"),
3724 domain,
3725 StorageClass::new("tenferro-test.transfer-destination").expect("destination storage"),
3726 );
3727 let transfer = ScheduledTransfer::new(
3728 3,
3729 source,
3730 destination.clone(),
3731 [],
3732 EventCompletion::new(domain, EventSlotId::new(0), 0),
3733 );
3734 let mut located = vec![
3735 Vec::new(),
3736 Vec::new(),
3737 Vec::new(),
3738 vec![LocatedExecSlot {
3739 location: destination.clone(),
3740 value: ExecSlot::Owned(f64_zeros(vec![1])),
3741 }],
3742 ];
3743
3744 let error = super::execute_scheduled_transfer(&transfer, &mut located)
3745 .expect_err("duplicate transfer destination");
3746 let Error::RuntimeStateSource { source, .. } = error else {
3747 panic!("duplicate transfer destination must retain a typed source");
3748 };
3749 let duplicate = source
3750 .downcast_ref::<DuplicateTransferDestinationError>()
3751 .expect("typed duplicate transfer destination source");
3752 assert_eq!(duplicate.value_slot, 3);
3753 assert_eq!(duplicate.destination, destination);
3754 }
3755
3756 #[test]
3757 fn tensor_backend_executor_reentrant_call_returns_error_instead_of_deadlocking() {
3758 let executor = Arc::new(TensorBackendExecutor::<CpuBackend>::new(CpuBackend::new()));
3759 let observed_reentrant_error = Arc::new(AtomicBool::new(false));
3760 let prepared = Arc::new(ReentrantProbePreparedOperation {
3761 binding: probe_binding(),
3762 specialization: probe_specialization(),
3763 executor: Arc::downgrade(&executor),
3764 observed_reentrant_error: Arc::clone(&observed_reentrant_error),
3765 });
3766 let operations = vec![PreparedOperationPlan::executable(
3767 prepared.clone(),
3768 prepared,
3769 )];
3770
3771 let output = ErasedTensorBackendExecutor::execute(
3772 executor.as_ref(),
3773 &reentrant_probe_program(),
3774 &operations,
3775 vec![f64_zeros(vec![2])],
3776 )
3777 .expect("outer extension executes");
3778
3779 assert_eq!(output[0].shape(), &[2]);
3780 assert!(
3781 observed_reentrant_error.load(Ordering::SeqCst),
3782 "same-thread reentrant executor call must fail immediately instead of deadlocking"
3783 );
3784 }
3785
3786 #[test]
3787 fn tensor_backend_executor_dispatches_session_capable_extension_in_one_session_path() {
3788 let executor = TensorBackendExecutor::<CpuBackend>::new(CpuBackend::new());
3789 let observed_session = Arc::new(AtomicBool::new(false));
3790 let prepared = Arc::new(SessionProbePreparedOperation {
3791 binding: probe_binding(),
3792 specialization: probe_specialization(),
3793 observed_session: Arc::clone(&observed_session),
3794 });
3795 let operations = vec![PreparedOperationPlan::executable(
3796 prepared.clone(),
3797 prepared,
3798 )];
3799
3800 let output = ErasedTensorBackendExecutor::execute(
3801 &executor,
3802 &lock_probe_program(),
3803 &operations,
3804 vec![f64_zeros(vec![2])],
3805 )
3806 .expect("session-capable extension executes");
3807
3808 assert_eq!(output[0].shape(), &[2]);
3809 assert!(observed_session.load(Ordering::SeqCst));
3810 }
3811
3812 #[test]
3813 fn tensor_backend_executor_batches_session_capable_region_after_boundary() {
3814 let executor = Arc::new(TensorBackendExecutor::<CpuBackend>::new(CpuBackend::new()));
3815 let observed_unlocked_state = Arc::new(AtomicBool::new(false));
3816 let observed_session = Arc::new(AtomicBool::new(false));
3817 let ordinary = Arc::new(LockProbePreparedOperation {
3818 binding: probe_binding(),
3819 specialization: probe_specialization(),
3820 executor: Arc::downgrade(&executor),
3821 observed_unlocked_state: Arc::clone(&observed_unlocked_state),
3822 });
3823 let session = Arc::new(SessionProbePreparedOperation {
3824 binding: probe_binding(),
3825 specialization: probe_specialization(),
3826 observed_session: Arc::clone(&observed_session),
3827 });
3828 let operations = vec![
3829 PreparedOperationPlan::executable(ordinary.clone(), ordinary),
3830 PreparedOperationPlan::executable(session.clone(), session),
3831 ];
3832
3833 let output = ErasedTensorBackendExecutor::execute(
3834 executor.as_ref(),
3835 &partial_session_probe_program(),
3836 &operations,
3837 vec![f64_zeros(vec![2])],
3838 )
3839 .expect("mixed session regions execute");
3840
3841 assert_eq!(output[0].shape(), &[2]);
3842 assert!(observed_unlocked_state.load(Ordering::SeqCst));
3843 assert!(observed_session.load(Ordering::SeqCst));
3844
3845 let values = ErasedTensorBackendExecutor::execute_values(
3846 executor.as_ref(),
3847 &partial_session_probe_program(),
3848 &operations,
3849 vec![f64_zeros(vec![2])],
3850 )
3851 .expect("mixed session regions execute in value mode");
3852 assert_eq!(values[0].shape(), &[2]);
3853 }
3854
3855 #[test]
3856 fn tensor_backend_executor_bridge_does_not_require_clone_source_contract() {
3857 type ContractConstructor<B> = fn(
3858 ProviderDeviceIdentity,
3859 CoreCapabilityBundle,
3860 B,
3861 Arc<dyn EventDomainDriver>,
3862 InputIngressContract,
3863 Option<Arc<dyn RuntimeCacheOwner>>,
3864 ) -> ExecutableEngineContract;
3865
3866 fn factory_accepts_backend_without_clone_bound<B>()
3867 where
3868 B: TensorBackend + Send + Sync + 'static,
3869 {
3870 let _factory: fn(B) -> Arc<dyn ErasedTensorBackendExecutor> =
3871 super::erased_tensor_backend_executor::<B>;
3872 }
3873
3874 fn contract_accepts_backend_without_clone_bound<B>()
3875 where
3876 B: TensorBackend + Send + Sync + 'static,
3877 {
3878 let _constructor: ContractConstructor<B> = ExecutableEngineContract::new::<B>;
3879 }
3880
3881 factory_accepts_backend_without_clone_bound::<CpuBackend>();
3882 contract_accepts_backend_without_clone_bound::<CpuBackend>();
3883 }
3884}