1use std::cell::Cell;
4use std::sync::{Arc, Mutex, OnceLock};
5
6use tenferro::{CompiledGraph, GraphCompiler, Runtime, Tensor, TracedGraph};
7use tenferro_ad::{AdContext, EagerRuntime};
8use tenferro_cpu::{BufferPoolStats, CpuBackend, CpuContext};
9use tenferro_tensor::{BackendSession, BackendSessionHost};
10
11#[derive(Clone, Debug)]
16pub enum ExecutionContext {
17 Cpu(Arc<CpuExecutionContext>),
19 #[cfg(feature = "tenferro-cuda")]
21 Cuda(Arc<crate::cuda::CudaExecutionContext>),
22}
23
24impl ExecutionContext {
25 pub fn is_global_default_cpu(&self) -> bool {
32 #[cfg(feature = "global-defaults")]
33 {
34 match self {
35 ExecutionContext::Cpu(context) => {
36 let own = context.eager_runtime().map(|runtime| runtime.id());
37 let global = defaults::default_eager_ctx().map(|runtime| runtime.id());
38 matches!((own, global), (Ok(a), Ok(b)) if a == b)
39 }
40 #[cfg(feature = "tenferro-cuda")]
41 ExecutionContext::Cuda(_) => false,
42 }
43 }
44 #[cfg(not(feature = "global-defaults"))]
45 {
46 let _ = self;
47 false
48 }
49 }
50}
51
52#[derive(Debug, Clone, thiserror::Error)]
70pub enum CpuExecutionContextError {
71 #[error("failed to initialize {component}: {source}")]
73 Initialization {
74 component: &'static str,
76 #[source]
78 source: Arc<dyn std::error::Error + Send + Sync + 'static>,
79 },
80 #[error("CPU graph {operation} failed: {source}")]
82 Graph {
83 operation: &'static str,
85 #[source]
87 source: Arc<dyn std::error::Error + Send + Sync + 'static>,
88 },
89}
90
91const CANONICAL_SESSION_REENTRY_MESSAGE: &str = "recursive tensorbackend canonical session entry";
92
93thread_local! {
94 static CANONICAL_SESSION_ACTIVE: Cell<bool> = const { Cell::new(false) };
95}
96
97struct CanonicalSessionGuard {
98 previous: bool,
99}
100
101impl CanonicalSessionGuard {
102 fn assert_inactive() {
103 CANONICAL_SESSION_ACTIVE.with(|active| {
104 assert!(!active.get(), "{CANONICAL_SESSION_REENTRY_MESSAGE}");
105 });
106 }
107
108 fn enter() -> Self {
109 Self::assert_inactive();
110 CANONICAL_SESSION_ACTIVE.with(|active| Self {
111 previous: active.replace(true),
112 })
113 }
114}
115
116impl Drop for CanonicalSessionGuard {
117 fn drop(&mut self) {
118 CANONICAL_SESSION_ACTIVE.with(|active| active.set(self.previous));
119 }
120}
121
122fn run_canonical_session<R: Send>(
124 backend: &mut CpuBackend,
125 f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
126) -> R {
127 backend.with_backend_session(|session| {
128 let _guard = CanonicalSessionGuard::enter();
129 f(session)
130 })
131}
132
133impl CpuExecutionContextError {
134 fn initialization(
135 component: &'static str,
136 source: impl std::error::Error + Send + Sync + 'static,
137 ) -> Self {
138 Self::Initialization {
139 component,
140 source: Arc::new(source),
141 }
142 }
143
144 fn graph(
145 operation: &'static str,
146 source: impl std::error::Error + Send + Sync + 'static,
147 ) -> Self {
148 Self::Graph {
149 operation,
150 source: Arc::new(source),
151 }
152 }
153}
154
155struct GraphState {
156 compiler: GraphCompiler,
157 runtime: Runtime,
158 backend: CpuBackend,
159}
160
161pub struct CpuExecutionContext {
186 backend: Mutex<CpuBackend>,
187 inline: OnceLock<CpuBackend>,
188 graph: OnceLock<Result<Mutex<GraphState>, CpuExecutionContextError>>,
189 eager: OnceLock<Result<Arc<EagerRuntime>, CpuExecutionContextError>>,
190}
191
192impl std::fmt::Debug for CpuExecutionContext {
193 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
194 f.debug_struct("CpuExecutionContext")
195 .field("graph_initialized", &self.graph.get().is_some())
196 .field("eager_initialized", &self.eager.get().is_some())
197 .finish_non_exhaustive()
198 }
199}
200
201impl CpuExecutionContext {
202 pub fn from_backend(backend: CpuBackend) -> Self {
207 Self {
208 backend: Mutex::new(backend),
209 inline: OnceLock::new(),
210 graph: OnceLock::new(),
211 eager: OnceLock::new(),
212 }
213 }
214
215 pub fn with_backend<R>(&self, f: impl FnOnce(&mut CpuBackend) -> R) -> R {
220 let mut backend = match self.backend.lock() {
221 Ok(guard) => guard,
222 Err(poisoned) => poisoned.into_inner(),
223 };
224 f(&mut backend)
225 }
226
227 pub(crate) fn with_session<R: Send>(
228 &self,
229 f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
230 ) -> R {
231 CanonicalSessionGuard::assert_inactive();
232 if rayon::current_thread_index().is_some() {
239 let mut backend = self.inline_backend();
240 return run_canonical_session(&mut backend, f);
241 }
242 let mut backend = match self.backend.lock() {
243 Ok(guard) => guard,
244 Err(poisoned) => poisoned.into_inner(),
245 };
246 run_canonical_session(&mut backend, f)
247 }
248
249 fn inline_backend(&self) -> CpuBackend {
255 self.inline
256 .get_or_init(|| match CpuContext::with_threads(1) {
257 Ok(context) => CpuBackend::from_context(Arc::new(context)),
258 Err(_) => self.backend_clone(),
263 })
264 .clone()
265 }
266
267 fn backend_clone(&self) -> CpuBackend {
268 self.with_backend(|backend| backend.clone())
269 }
270
271 fn graph_state(&self) -> Result<&Mutex<GraphState>, CpuExecutionContextError> {
272 self.graph
273 .get_or_init(|| {
274 let backend = self.backend_clone();
275 build_graph_runtime(&backend).map(|runtime| {
276 Mutex::new(GraphState {
277 compiler: GraphCompiler::new(),
278 runtime,
279 backend,
280 })
281 })
282 })
283 .as_ref()
284 .map_err(Clone::clone)
285 }
286
287 fn with_graph_state<R>(
288 &self,
289 f: impl FnOnce(&mut GraphCompiler, &mut Runtime, &mut CpuBackend) -> R,
290 ) -> Result<R, CpuExecutionContextError> {
291 let mut graph = match self.graph_state()?.lock() {
292 Ok(guard) => guard,
293 Err(poisoned) => poisoned.into_inner(),
294 };
295 let GraphState {
296 compiler,
297 runtime,
298 backend,
299 } = &mut *graph;
300 Ok(f(compiler, runtime, backend))
301 }
302
303 pub fn compile_graph(
310 &self,
311 graph: &TracedGraph,
312 ) -> Result<CompiledGraph, CpuExecutionContextError> {
313 self.with_graph_state(|compiler, _, _| compiler.compile_traced_graph(graph))?
314 .map_err(|source| CpuExecutionContextError::graph("compilation", source))
315 }
316
317 pub fn run_graph(
327 &self,
328 graph: &CompiledGraph,
329 inputs: &[&Tensor],
330 ) -> Result<Vec<Tensor>, CpuExecutionContextError> {
331 self.with_graph_state(|_, runtime, _| runtime.run_compiled(graph, inputs))?
332 .map_err(|source| CpuExecutionContextError::graph("execution", source))
333 }
334
335 pub fn eager_runtime(&self) -> Result<Arc<EagerRuntime>, CpuExecutionContextError> {
345 self.eager
346 .get_or_init(|| build_eager_runtime(self.backend_clone()))
347 .as_ref()
348 .map(Arc::clone)
349 .map_err(Clone::clone)
350 }
351
352 pub fn graph_cache_stats(
359 &self,
360 ) -> Result<tenferro::RuntimeCacheStats, CpuExecutionContextError> {
361 self.with_graph_state(|_, runtime, _| runtime.cache_stats())?
362 .map_err(|source| CpuExecutionContextError::graph("cache statistics", source))
363 }
364
365 pub fn graph_buffer_pool_stats(&self) -> Result<BufferPoolStats, CpuExecutionContextError> {
372 self.with_graph_state(|_, _, backend| backend.buffer_pool_stats())?
373 .map_err(|source| CpuExecutionContextError::graph("buffer-pool statistics", source))
374 }
375
376 pub fn reset_graph_buffer_pool(&self) -> Result<(), CpuExecutionContextError> {
383 self.with_graph_state(|_, _, backend| backend.reset_buffer_pool())?
384 .map_err(|source| CpuExecutionContextError::graph("buffer-pool reset", source))
385 }
386
387 pub fn reset_graph_runtime(&self) -> Result<(), CpuExecutionContextError> {
394 self.with_graph_state(|compiler, runtime, backend| {
395 let replacement = build_graph_runtime(backend)?;
396 *compiler = GraphCompiler::new();
397 let old = std::mem::replace(runtime, replacement);
398 drop(old);
399 backend
400 .reset_buffer_pool()
401 .map_err(|source| CpuExecutionContextError::graph("buffer-pool reset", source))
402 })??;
403 Ok(())
404 }
405}
406
407fn build_graph_runtime(backend: &CpuBackend) -> Result<Runtime, CpuExecutionContextError> {
408 let mut builder = Runtime::builder();
409 builder
410 .register_engine(
411 tenferro_cpu::runtime_engine_registration(backend).map_err(|source| {
412 CpuExecutionContextError::initialization("graph CPU engine", source)
413 })?,
414 )
415 .map_err(|source| CpuExecutionContextError::initialization("graph CPU engine", source))?;
416 builder
417 .install_extension_module(
418 tenferro_einsum::extension_module::<CpuBackend>(
419 tenferro_cpu::runtime_engine_id().map_err(|source| {
420 CpuExecutionContextError::initialization("einsum extension", source)
421 })?,
422 )
423 .map_err(|source| {
424 CpuExecutionContextError::initialization("einsum extension", source)
425 })?,
426 )
427 .map_err(|source| CpuExecutionContextError::initialization("einsum extension", source))?;
428 builder
429 .build()
430 .map_err(|source| CpuExecutionContextError::initialization("graph runtime", source))
431}
432
433fn build_eager_runtime(backend: CpuBackend) -> Result<Arc<EagerRuntime>, CpuExecutionContextError> {
434 let ad_context = AdContext::builder()
435 .with_semantic_extension_rules(tenferro_linalg::semantic_ad_rules().map_err(|source| {
436 CpuExecutionContextError::initialization("linalg AD rules", source)
437 })?)
438 .map_err(|source| CpuExecutionContextError::initialization("linalg AD rules", source))?
439 .build()
440 .map_err(|source| CpuExecutionContextError::initialization("AD context", source))?;
441 let runtime = EagerRuntime::with_cpu_backend_and_ad_context(backend, &ad_context)
442 .map_err(|source| CpuExecutionContextError::initialization("eager runtime", source))?;
443 let engine_id = tenferro_cpu::runtime_engine_id()
448 .map_err(|source| CpuExecutionContextError::initialization("CPU runtime engine", source))?;
449 let einsum_module = tenferro_einsum::extension_module::<CpuBackend>(engine_id.clone())
450 .map_err(|source| CpuExecutionContextError::initialization("einsum extension", source))?;
451 runtime
452 .install_extension_module(einsum_module)
453 .map_err(|source| CpuExecutionContextError::initialization("einsum runtime", source))?;
454 let linalg_module = tenferro_linalg::extension_module::<CpuBackend>(engine_id)
455 .map_err(|source| CpuExecutionContextError::initialization("linalg extension", source))?;
456 runtime
457 .install_extension_module(linalg_module)
458 .map_err(|source| CpuExecutionContextError::initialization("linalg runtime", source))?;
459 Ok(runtime)
460}
461
462#[cfg(feature = "global-defaults")]
463mod defaults {
464 use super::*;
465 use tenferro_cpu::CpuContext;
466
467 static DEFAULT_CONTEXT: OnceLock<Arc<CpuExecutionContext>> = OnceLock::new();
468
469 #[cfg(test)]
470 thread_local! {
471 static FORCE_EAGER_CONTEXT_FAILURE: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
472 }
473
474 #[cfg(test)]
475 static DEFAULT_CONTEXT_HITS: std::sync::atomic::AtomicUsize =
476 std::sync::atomic::AtomicUsize::new(0);
477
478 fn default_context() -> &'static Arc<CpuExecutionContext> {
479 DEFAULT_CONTEXT.get_or_init(|| {
480 #[cfg(test)]
481 DEFAULT_CONTEXT_HITS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
482 Arc::new(CpuExecutionContext::from_backend(CpuBackend::from_context(
483 Arc::new(CpuContext::from_env()),
484 )))
485 })
486 }
487
488 #[derive(Debug, Clone, thiserror::Error)]
503 pub enum EagerContextError {
504 #[error("failed to register tenferro linalg AD rule: {source}")]
506 Registration {
507 #[source]
509 source: Arc<dyn std::error::Error + Send + Sync + 'static>,
510 },
511 }
512
513 pub fn with_default_backend<R>(f: impl FnOnce(&mut CpuBackend) -> R) -> R {
515 default_context().with_backend(f)
516 }
517
518 pub(crate) fn with_default_session<R: Send>(
519 f: impl FnOnce(&mut dyn BackendSession) -> R + Send,
520 ) -> R {
521 default_context().with_session(f)
522 }
523
524 pub(crate) fn with_default_graph_runtime<R>(
525 f: impl FnOnce(&mut GraphCompiler, &Runtime, &mut CpuBackend) -> R,
526 ) -> anyhow::Result<R> {
527 default_context()
528 .with_graph_state(|compiler, runtime, backend| f(compiler, runtime, backend))
529 .map_err(anyhow::Error::new)
530 }
531
532 pub(crate) fn default_engine_buffer_pool_stats() -> anyhow::Result<BufferPoolStats> {
533 default_context()
534 .graph_buffer_pool_stats()
535 .map_err(anyhow::Error::new)
536 }
537
538 pub(crate) fn reset_default_engine_buffer_pool() -> anyhow::Result<()> {
539 default_context()
540 .reset_graph_buffer_pool()
541 .map_err(anyhow::Error::new)
542 }
543
544 pub(crate) fn reset_default_engine() -> anyhow::Result<()> {
545 default_context()
546 .reset_graph_runtime()
547 .map_err(anyhow::Error::new)
548 }
549
550 pub fn default_eager_ctx() -> Result<Arc<EagerRuntime>, EagerContextError> {
568 #[cfg(test)]
569 if FORCE_EAGER_CONTEXT_FAILURE.with(std::cell::Cell::get) {
570 return Err(EagerContextError::Registration {
571 source: Arc::new(std::io::Error::other(
572 "forced default eager context registration failure",
573 )),
574 });
575 }
576 default_context()
577 .eager_runtime()
578 .map_err(|source| EagerContextError::Registration {
579 source: Arc::new(source),
580 })
581 }
582
583 pub fn default_cpu_execution_context() -> Arc<CpuExecutionContext> {
599 Arc::clone(default_context())
600 }
601
602 #[cfg(test)]
603 pub(crate) fn default_context_hits() -> usize {
604 DEFAULT_CONTEXT_HITS.load(std::sync::atomic::Ordering::Relaxed)
605 }
606
607 #[cfg(test)]
608 pub(crate) fn with_forced_eager_context_failure<T>(f: impl FnOnce() -> T) -> T {
609 let previous = FORCE_EAGER_CONTEXT_FAILURE.with(|failure| failure.replace(true));
610 let result = f();
611 FORCE_EAGER_CONTEXT_FAILURE.with(|failure| failure.set(previous));
612 result
613 }
614}
615
616#[cfg(all(test, feature = "global-defaults"))]
617pub(crate) use defaults::with_forced_eager_context_failure;
618#[cfg(feature = "global-defaults")]
619pub use defaults::{
620 default_cpu_execution_context, default_eager_ctx, with_default_backend, EagerContextError,
621};
622#[cfg(feature = "global-defaults")]
623pub(crate) use defaults::{
624 default_engine_buffer_pool_stats, reset_default_engine, reset_default_engine_buffer_pool,
625 with_default_graph_runtime, with_default_session,
626};
627
628#[cfg(test)]
629mod tests {
630 use std::num::NonZeroUsize;
631 use std::sync::mpsc;
632 use std::time::Duration;
633
634 use super::*;
635 use tenferro::program::{CoreSemanticOp, ProgramInputSpec};
636 use tenferro::{DType, TensorSessionOpsExt, TraceContext};
637 use tenferro_ad::EagerTensor;
638 use tenferro_cpu::{CpuContext, ExternalCpuDomain};
639 use tenferro_tensor::CpuDomainId;
640
641 fn context() -> CpuExecutionContext {
642 CpuExecutionContext::from_backend(CpuBackend::with_threads(1).unwrap())
643 }
644
645 #[test]
646 fn explicit_session_runs_a_concrete_operation() {
647 let context = context();
648 let lhs = Tensor::from_vec_col_major(vec![2, 2], vec![1.0_f64, 2.0, 3.0, 4.0]).unwrap();
649 let rhs = Tensor::from_vec_col_major(vec![2, 1], vec![5.0_f64, 6.0]).unwrap();
650 let result = context
651 .with_session(|session| lhs.matmul(&rhs, session))
652 .unwrap();
653
654 assert_eq!(result.as_slice::<f64>().unwrap(), &[23.0, 34.0]);
655 }
656
657 #[test]
658 fn recursive_session_entry_fails_before_lock_and_restores_guard() {
659 let context = context();
660 let panic = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
661 context.with_session(|_| context.with_session(|_| ()))
662 }))
663 .expect_err("recursive canonical session entry should panic");
664 let message = panic
665 .downcast_ref::<&str>()
666 .copied()
667 .or_else(|| panic.downcast_ref::<String>().map(String::as_str))
668 .expect("recursive entry panic should contain a string message");
669 assert_eq!(message, CANONICAL_SESSION_REENTRY_MESSAGE);
670
671 assert_eq!(context.with_session(|_| 7usize), 7);
672 }
673
674 #[test]
675 fn explicit_plain_graph_and_eager_paths_share_only_the_supplied_backend() {
676 let context = context();
677 assert!(format!("{context:?}").contains("graph_initialized: false"));
678 assert_eq!(context.with_backend(|backend| backend.num_threads()), 1);
679
680 let mut trace = TraceContext::new();
681 let input = trace
682 .input(ProgramInputSpec::new(DType::F64, [2_usize.into()]))
683 .unwrap();
684 let output = trace.add_op(CoreSemanticOp::Neg, &[input]).unwrap()[0];
685 let graph = trace.finish(&[output]).unwrap();
686 let compiled = context.compile_graph(&graph).unwrap();
687 let input = Tensor::from_vec_col_major(vec![2], vec![1.0_f64, -2.0]).unwrap();
688 let output = context.run_graph(&compiled, &[&input]).unwrap();
689 assert_eq!(output[0].as_slice::<f64>().unwrap(), &[-1.0, 2.0]);
690 context.run_graph(&compiled, &[&input]).unwrap();
691 let cached = context.graph_cache_stats().unwrap().prepared_plans;
692 assert!(cached.entries > 0);
693 assert!(cached.hits > 0);
694 context.reset_graph_runtime().unwrap();
695 assert_eq!(
696 context.graph_cache_stats().unwrap().prepared_plans.entries,
697 0
698 );
699
700 let eager = context.eager_runtime().unwrap();
701 assert!(Arc::ptr_eq(&eager, &context.eager_runtime().unwrap()));
702 }
703
704 #[test]
705 fn separate_eager_contexts_reject_cross_context_operations() {
706 let first = context().eager_runtime().unwrap();
707 let second = context().eager_runtime().unwrap();
708 let a = EagerTensor::from_tensor_in(
709 Tensor::from_vec_col_major(vec![1], vec![1.0_f64]).unwrap(),
710 first,
711 )
712 .unwrap();
713 let b = EagerTensor::from_tensor_in(
714 Tensor::from_vec_col_major(vec![1], vec![2.0_f64]).unwrap(),
715 second,
716 )
717 .unwrap();
718 assert!(matches!(
719 a.add(&b),
720 Err(tenferro_ad::Error::ContextMismatch { .. })
721 ));
722 }
723
724 #[test]
725 fn caller_managed_backend_remains_caller_owned_after_context_drop() {
726 let executor = Arc::new(CpuContext::with_threads(1).unwrap());
727 let id = CpuDomainId::new(7);
728 let domain =
729 ExternalCpuDomain::new_caller_managed(id, executor.clone(), NonZeroUsize::MIN).unwrap();
730 let backend = CpuBackend::from_external_managed_domains(id, [domain]).unwrap();
731 let context = CpuExecutionContext::from_backend(backend);
732 assert_eq!(context.with_backend(|backend| backend.num_threads()), 1);
733 drop(context);
734 assert_eq!(executor.num_threads(), 1);
735 }
736
737 #[test]
738 fn independent_contexts_do_not_share_a_backend_mutex() {
739 let first = Arc::new(context());
740 let second = Arc::new(context());
741 let (entered_tx, entered_rx) = mpsc::channel();
742 let (release_tx, release_rx) = mpsc::channel();
743 let release_rx = Arc::new(Mutex::new(release_rx));
744 let handles = [first, second].map(|context| {
745 let entered_tx = entered_tx.clone();
746 let release_rx = Arc::clone(&release_rx);
747 std::thread::spawn(move || {
748 context.with_backend(|_| {
749 entered_tx.send(()).unwrap();
750 release_rx.lock().unwrap().recv().unwrap();
751 });
752 })
753 });
754 entered_rx.recv_timeout(Duration::from_secs(2)).unwrap();
755 entered_rx.recv_timeout(Duration::from_secs(2)).unwrap();
756 release_tx.send(()).unwrap();
757 release_tx.send(()).unwrap();
758 for handle in handles {
759 handle.join().unwrap();
760 }
761 }
762
763 #[test]
764 fn session_from_a_rayon_worker_completes_without_waiting_on_the_context_pool() {
765 use rayon::prelude::*;
766
767 let context = Arc::new(CpuExecutionContext::from_backend(
773 CpuBackend::with_threads(2).unwrap(),
774 ));
775 for enclosing_workers in [1usize, 2] {
776 let pool = Arc::new(
777 rayon::ThreadPoolBuilder::new()
778 .num_threads(enclosing_workers)
779 .build()
780 .unwrap(),
781 );
782 let (sender, receiver) = mpsc::channel();
783 let context = Arc::clone(&context);
784 std::thread::spawn(move || {
785 let results = pool.install(|| {
786 (0..2usize)
787 .into_par_iter()
788 .map(|_| {
789 let lhs = Tensor::from_vec_col_major(
790 vec![2, 2],
791 vec![1.0_f64, 2.0, 3.0, 4.0],
792 )
793 .unwrap();
794 let rhs =
795 Tensor::from_vec_col_major(vec![2, 1], vec![5.0_f64, 6.0]).unwrap();
796 context
797 .with_session(|session| lhs.matmul(&rhs, session))
798 .unwrap()
799 .as_slice::<f64>()
800 .unwrap()
801 .to_vec()
802 })
803 .collect::<Vec<_>>()
804 });
805 let _ = sender.send(results);
806 });
807
808 let results = receiver.recv_timeout(Duration::from_secs(60)).expect(
809 "a session entered from a Rayon worker must finish instead of waiting on the context pool",
810 );
811 assert_eq!(results.len(), 2, "enclosing workers = {enclosing_workers}");
812 for values in results {
813 assert_eq!(values, vec![23.0, 34.0]);
814 }
815 }
816 }
817
818 #[cfg(feature = "global-defaults")]
819 #[test]
820 fn explicit_paths_do_not_initialize_the_default_context() {
821 let before = defaults::default_context_hits();
822 let context = context();
823 context.with_backend(|backend| assert_eq!(backend.num_threads(), 1));
824 context.eager_runtime().unwrap();
825 assert_eq!(defaults::default_context_hits(), before);
826 }
827}