tenferro_ops/ad/context.rs
1//! AD context for guard-based shape resolution and value metadata queries.
2//!
3//! During AD graph construction, linalg rules such as SVD, QR, and LU need
4//! concrete matrix dimensions to choose between structurally different
5//! subgraphs. `ShapeGuardContext` records those dimension comparisons as guards
6//! so cached AD graphs can later be invalidated when the observed shape
7//! relationship changes.
8
9use std::cmp::Ordering;
10use std::collections::HashMap;
11use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
12#[cfg(feature = "autodiff")]
13use std::sync::Arc;
14use std::sync::{Mutex, OnceLock};
15
16use computegraph::graph::Graph;
17use computegraph::types::{ValueKey, ValueRef};
18use tenferro_tensor::DType;
19
20#[cfg(feature = "autodiff")]
21use crate::ad::{ADRuleError, ADRuleKind, ExtensionAdDispatcher};
22use crate::dim_expr::{DimExpr, DimExprEvalError};
23use crate::shape_extent::ShapeExtent;
24use crate::std_tensor_op::StdTensorOp;
25use crate::sym_dim::SymDim;
26
27type MetadataMap = HashMap<ValueKey<StdTensorOp>, TensorMeta>;
28
29type GlobalMetadataMap = HashMap<ValueKey<StdTensorOp>, GlobalMetadataEntry>;
30
31#[derive(Clone, Debug)]
32struct GlobalMetadataEntry {
33 stack: Vec<GlobalMetadataRegistration>,
34}
35
36#[derive(Clone, Debug)]
37struct GlobalMetadataRegistration {
38 token: u64,
39 meta: TensorMeta,
40}
41
42#[derive(Clone, Debug)]
43struct ScopedGlobalMetadataRegistration {
44 key: ValueKey<StdTensorOp>,
45 token: u64,
46}
47
48/// Error returned when the process-global AD metadata registry is unavailable.
49#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
50pub enum MetadataRegistryError {
51 /// A previous panic poisoned the global metadata mutex.
52 #[error("AD global metadata registry lock poisoned")]
53 LockPoisoned,
54}
55
56/// Error returned when shape-guard metadata cannot be resolved.
57///
58/// # Examples
59///
60/// ```
61/// use tenferro_ops::ShapeGuardError;
62///
63/// let error = ShapeGuardError::LocalWithoutAttachedGraph { local_id: 0 };
64/// assert!(error.to_string().contains("attached graph"));
65/// ```
66#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
67pub enum ShapeGuardError {
68 /// A local graph value was queried before a graph was attached.
69 #[error("cannot resolve local value {local_id} without an attached graph")]
70 LocalWithoutAttachedGraph {
71 /// Graph-local value id.
72 local_id: usize,
73 },
74 /// A local graph value id is outside the attached graph's value table.
75 #[error("local value {local_id} is out of bounds for the attached graph")]
76 LocalOutOfBounds {
77 /// Graph-local value id.
78 local_id: usize,
79 },
80 /// No metadata was registered for the resolved value key.
81 #[error("missing TensorMeta for {key:?}")]
82 MissingMetadata {
83 /// Resolved value key.
84 key: ValueKey<StdTensorOp>,
85 },
86 /// Metadata exists, but at least one axis is only bounded or unknown.
87 #[error("TensorMeta for {key:?} does not have an exact shape; query extents instead")]
88 NonExactShape {
89 /// Resolved value key.
90 key: ValueKey<StdTensorOp>,
91 },
92}
93
94/// Result type used by shape-guard metadata queries.
95///
96/// The error side preserves a [`ShapeGuardFailure`] wrapper so an AD callback
97/// can retain the original [`ShapeGuardError`] even when a foreign callback
98/// protocol accepts only a rendered message.
99///
100/// # Examples
101///
102/// ```
103/// use tenferro_ops::ShapeGuardResult;
104///
105/// let result: ShapeGuardResult<()> = Ok(());
106/// assert!(result.is_ok());
107/// ```
108pub type ShapeGuardResult<T> = Result<T, ShapeGuardFailure>;
109
110#[cfg(feature = "autodiff")]
111impl From<ShapeGuardFailure> for ADRuleError {
112 fn from(err: ShapeGuardFailure) -> Self {
113 err.record_for_ad_boundary();
114 ADRuleError::invalid_input(
115 "tenferro.shape_guard",
116 ADRuleKind::Jvp,
117 err.typed_source().to_string(),
118 )
119 }
120}
121
122/// Error returned by a shape-guard metadata query.
123///
124/// The public [`ShapeGuardError`] remains the typed source. The private side
125/// channel is shared with the owning [`ShapeGuardContext`] so an external
126/// message-only AD callback can report the same typed source at the runtime
127/// boundary without changing the callback protocol.
128///
129/// # Examples
130///
131/// ```
132/// use computegraph::types::{ValueKey, ValueRef};
133/// use tenferro_ops::input_key::TensorInputKey;
134/// use tenferro_ops::std_tensor_op::StdTensorOp;
135/// use tenferro_ops::{ShapeGuardContext, ShapeGuardError};
136///
137/// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 8 });
138/// let value = ValueRef::External(key);
139/// let mut ctx = ShapeGuardContext::default();
140/// let failure = ctx.shape_of(&value).unwrap_err();
141/// assert!(matches!(
142/// failure.typed_source(),
143/// ShapeGuardError::MissingMetadata { .. }
144/// ));
145/// ```
146#[derive(Clone, Debug)]
147pub struct ShapeGuardFailure {
148 source: ShapeGuardError,
149 #[cfg(feature = "autodiff")]
150 deferred: Arc<Mutex<Option<ShapeGuardError>>>,
151}
152
153impl ShapeGuardFailure {
154 #[cfg(feature = "autodiff")]
155 fn new(source: ShapeGuardError, deferred: Arc<Mutex<Option<ShapeGuardError>>>) -> Self {
156 Self { source, deferred }
157 }
158
159 #[cfg(not(feature = "autodiff"))]
160 fn new(source: ShapeGuardError) -> Self {
161 Self { source }
162 }
163
164 /// Return the original typed shape-guard failure.
165 pub fn typed_source(&self) -> &ShapeGuardError {
166 &self.source
167 }
168
169 /// Consume this boundary error and return its original typed failure.
170 pub fn into_typed_source(self) -> ShapeGuardError {
171 self.source
172 }
173
174 #[cfg(feature = "autodiff")]
175 fn record_for_ad_boundary(&self) {
176 if let Ok(mut deferred) = self.deferred.lock() {
177 if deferred.is_none() {
178 *deferred = Some(self.source.clone());
179 }
180 }
181 }
182}
183
184impl std::fmt::Display for ShapeGuardFailure {
185 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
186 self.source.fmt(f)
187 }
188}
189
190impl std::error::Error for ShapeGuardFailure {
191 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
192 Some(&self.source)
193 }
194}
195
196impl PartialEq for ShapeGuardFailure {
197 fn eq(&self, other: &Self) -> bool {
198 self.source == other.source
199 }
200}
201
202impl Eq for ShapeGuardFailure {}
203
204/// Global metadata registry.
205///
206/// Stored as a tokenized stack per value key: duplicate scoped registrations
207/// shadow older metadata while they are live, and dropping scopes in any order
208/// removes only the matching token. `ShapeGuardContext::metadata_of` reaches into
209/// the registry lazily via [`lookup_global_metadata`] and caches the result into
210/// the context's local map.
211///
212/// Earlier designs either cloned the whole map up-front into each AD
213/// `ShapeGuardContext` or kept the map in an `Arc` and cloned on every write.
214/// Both variants were quadratic across the monotonically growing registry and
215/// dominated oracle_replay runtime.
216static GLOBAL_METADATA: OnceLock<Mutex<GlobalMetadataMap>> = OnceLock::new();
217static NEXT_GLOBAL_METADATA_TOKEN: AtomicU64 = AtomicU64::new(0);
218
219fn global_metadata_registry() -> &'static Mutex<GlobalMetadataMap> {
220 GLOBAL_METADATA.get_or_init(|| Mutex::new(HashMap::new()))
221}
222
223/// Lifetime token for graph-scoped global metadata.
224///
225/// Dropping the last frontend owner of a traced graph drops this scope and
226/// releases the metadata keys that were registered for that graph graph.
227#[doc(hidden)]
228#[derive(Debug)]
229pub struct GlobalMetadataScope {
230 registrations: Vec<ScopedGlobalMetadataRegistration>,
231}
232
233impl Drop for GlobalMetadataScope {
234 fn drop(&mut self) {
235 release_scoped_global_metadata(&self.registrations);
236 }
237}
238
239/// Per-value tensor metadata used by AD rules.
240///
241/// Shape information is stored as per-axis [`ShapeExtent`] values. Callers must
242/// explicitly choose whether they need an exact shape or only a known bound.
243///
244/// # Examples
245///
246/// ```
247/// use tenferro_ops::{SymDim, TensorMeta};
248/// use tenferro_tensor::DType;
249///
250/// let meta = TensorMeta::exact(DType::F64, vec![SymDim::from(2usize), SymDim::from(3usize)]);
251/// assert_eq!(meta.rank(), 2);
252/// ```
253#[derive(Clone, Debug, PartialEq, Eq)]
254pub struct TensorMeta {
255 /// Element dtype of the tensor value.
256 pub dtype: DType,
257 /// Per-axis shape guarantees.
258 pub extents: Vec<ShapeExtent<SymDim>>,
259 /// Canonical identity of an externally defined scalar, when the value carries one.
260 ///
261 /// A semantic program's identity must be reproducible across processes, while an
262 /// externally defined scalar's tag is a process-local `TypeId`, so a traced value
263 /// whose dtype is external declares the stable name here.
264 pub scalar_identity: Option<&'static str>,
265}
266
267impl TensorMeta {
268 /// Construct metadata whose every axis is exact.
269 ///
270 /// # Examples
271 ///
272 /// ```
273 /// use tenferro_ops::{SymDim, TensorMeta};
274 /// use tenferro_tensor::DType;
275 ///
276 /// let meta = TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]);
277 /// assert_eq!(meta.exact_shape(), Some(vec![SymDim::from(4usize)]));
278 /// ```
279 pub fn exact(dtype: DType, shape: Vec<SymDim>) -> Self {
280 let extents = shape.iter().cloned().map(ShapeExtent::exact).collect();
281 Self {
282 dtype,
283 extents,
284 scalar_identity: None,
285 }
286 }
287
288 /// Construct metadata from per-axis extents.
289 ///
290 /// # Examples
291 ///
292 /// ```
293 /// use tenferro_ops::{ShapeExtent, SymDim, TensorMeta};
294 /// use tenferro_tensor::DType;
295 ///
296 /// let meta = TensorMeta::with_extents(
297 /// DType::F64,
298 /// vec![ShapeExtent::upper_bound(SymDim::from(8usize))],
299 /// );
300 /// assert_eq!(meta.exact_shape(), None);
301 /// ```
302 pub fn with_extents(dtype: DType, extents: Vec<ShapeExtent<SymDim>>) -> Self {
303 Self {
304 dtype,
305 extents,
306 scalar_identity: None,
307 }
308 }
309
310 /// Declare the canonical identity of an externally defined scalar.
311 ///
312 /// # Examples
313 ///
314 /// ```rust
315 /// use tenferro_ops::{SymDim, TensorMeta};
316 /// use tenferro_tensor::DType;
317 ///
318 /// let dtype = DType::External(std::any::TypeId::of::<f64>());
319 /// let meta = TensorMeta::exact(dtype, vec![SymDim::from(2usize)])
320 /// .with_scalar_identity("example.scalar.v1");
321 /// assert_eq!(meta.scalar_identity(), Some("example.scalar.v1"));
322 /// ```
323 #[must_use]
324 pub fn with_scalar_identity(mut self, identity: &'static str) -> Self {
325 self.scalar_identity = Some(identity);
326 self
327 }
328
329 /// Return the declared identity of an externally defined scalar, if any.
330 ///
331 /// # Examples
332 ///
333 /// ```rust
334 /// use tenferro_ops::{SymDim, TensorMeta};
335 /// use tenferro_tensor::DType;
336 ///
337 /// let dtype = DType::External(std::any::TypeId::of::<f64>());
338 /// assert_eq!(TensorMeta::exact(dtype, vec![SymDim::from(1usize)]).scalar_identity(), None);
339 /// assert_eq!(
340 /// TensorMeta::exact(dtype, vec![SymDim::from(1usize)])
341 /// .with_scalar_identity("example.scalar.v1")
342 /// .scalar_identity(),
343 /// Some("example.scalar.v1")
344 /// );
345 /// ```
346 #[must_use]
347 pub const fn scalar_identity(&self) -> Option<&'static str> {
348 self.scalar_identity
349 }
350
351 /// Return the tensor rank known by this metadata record.
352 pub fn rank(&self) -> usize {
353 self.extents.len()
354 }
355
356 /// Return the per-axis shape guarantees.
357 ///
358 /// # Examples
359 ///
360 /// ```
361 /// use tenferro_ops::{SymDim, TensorMeta};
362 /// use tenferro_tensor::DType;
363 ///
364 /// let meta = TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]);
365 /// assert_eq!(meta.extents().len(), 1);
366 /// ```
367 pub fn extents(&self) -> &[ShapeExtent<SymDim>] {
368 &self.extents
369 }
370
371 /// Return the shape only when every axis is exact.
372 ///
373 /// # Examples
374 ///
375 /// ```
376 /// use tenferro_ops::{ShapeExtent, SymDim, TensorMeta};
377 /// use tenferro_tensor::DType;
378 ///
379 /// let meta = TensorMeta::with_extents(
380 /// DType::F64,
381 /// vec![ShapeExtent::upper_bound(SymDim::from(8usize))],
382 /// );
383 /// assert_eq!(meta.exact_shape(), None);
384 /// ```
385 pub fn exact_shape(&self) -> Option<Vec<SymDim>> {
386 self.extents
387 .iter()
388 .map(|extent| extent.as_exact().cloned())
389 .collect()
390 }
391
392 /// Return one known bound per axis when every axis has a bound.
393 ///
394 /// This is intentionally separate from [`TensorMeta::exact_shape`]: a bound
395 /// is not proof of the runtime size.
396 pub fn bound_shape(&self) -> Option<Vec<SymDim>> {
397 self.extents
398 .iter()
399 .map(|extent| extent.bound_expr().cloned())
400 .collect()
401 }
402}
403
404/// A recorded dimension comparison made during AD graph construction.
405///
406/// # Examples
407///
408/// ```
409/// use std::cmp::Ordering;
410/// use tenferro_ops::ShapeGuard;
411///
412/// let guard = ShapeGuard {
413/// dim_a: 5,
414/// dim_b: 3,
415/// ordering: Ordering::Greater,
416/// };
417///
418/// assert_eq!(guard.ordering, Ordering::Greater);
419/// ```
420#[derive(Clone, Debug, PartialEq, Eq)]
421pub struct ShapeGuard {
422 /// First dimension value, such as `m`.
423 pub dim_a: usize,
424 /// Second dimension value, such as `n`.
425 pub dim_b: usize,
426 /// The observed ordering `dim_a.cmp(&dim_b)`.
427 pub ordering: Ordering,
428}
429
430/// AD context providing dimension resolution, guard recording, and value metadata.
431///
432/// # Examples
433///
434/// ```
435/// use tenferro_ops::ShapeGuardContext;
436///
437/// let ctx = ShapeGuardContext::default();
438/// assert!(ctx.guards().is_empty());
439/// ```
440#[derive(Clone, Debug, Default)]
441pub struct ShapeGuardContext {
442 guards: Vec<ShapeGuard>,
443 metadata: MetadataMap,
444 shape_sources: HashMap<u64, ValueKey<StdTensorOp>>,
445 use_global_registry: bool,
446 local_keys: Option<Vec<ValueKey<StdTensorOp>>>,
447 #[cfg(feature = "autodiff")]
448 deferred_shape_error: Arc<Mutex<Option<ShapeGuardError>>>,
449 #[cfg(feature = "autodiff")]
450 extension_ad_dispatcher: Option<Arc<dyn ExtensionAdDispatcher>>,
451 #[cfg(feature = "autodiff")]
452 active_value_keys: Option<std::sync::Arc<std::collections::HashSet<ValueKey<StdTensorOp>>>>,
453 #[cfg(feature = "autodiff")]
454 transpose_primal_outputs: Option<Vec<ValueKey<StdTensorOp>>>,
455 #[cfg(feature = "autodiff")]
456 transpose_primal_outputs_used: bool,
457}
458
459impl ShapeGuardContext {
460 /// Create a context backed by the global metadata registry.
461 ///
462 /// Instead of cloning the entire global registry up-front (which used
463 /// to be O(N) per AD pass and quadratic across oracle_replay), the
464 /// context keeps a flag and lazily fetches entries from the shared
465 /// [`lookup_global_metadata`] on first miss, caching into its local
466 /// `metadata` map for subsequent reads within the same pass.
467 ///
468 /// # Examples
469 ///
470 /// ```
471 /// let ctx = tenferro_ops::ShapeGuardContext::with_global_metadata();
472 /// assert!(ctx.guards().is_empty());
473 /// ```
474 pub fn with_global_metadata() -> Self {
475 Self {
476 use_global_registry: true,
477 ..Self::default()
478 }
479 }
480
481 #[doc(hidden)]
482 /// Keep global-registry lookup enabled after a pass boundary.
483 ///
484 /// This is intentionally a no-op for cached entries: global metadata is
485 /// already read lazily on cache misses, and clearing the local cache would
486 /// also discard metadata inserted directly into this context.
487 pub fn refresh_global_metadata(&mut self) {
488 self.use_global_registry = true;
489 }
490
491 #[doc(hidden)]
492 pub fn insert_shape_source(&mut self, tensor_id: u64, key: ValueKey<StdTensorOp>) {
493 self.shape_sources.entry(tensor_id).or_insert(key);
494 }
495
496 #[doc(hidden)]
497 pub fn shape_source(&self, tensor_id: u64) -> Option<&ValueKey<StdTensorOp>> {
498 self.shape_sources.get(&tensor_id)
499 }
500
501 #[doc(hidden)]
502 #[cfg(feature = "autodiff")]
503 pub fn with_extension_ad_dispatcher(
504 mut self,
505 dispatcher: Arc<dyn ExtensionAdDispatcher>,
506 ) -> Self {
507 self.extension_ad_dispatcher = Some(dispatcher);
508 self
509 }
510
511 #[doc(hidden)]
512 #[cfg(feature = "autodiff")]
513 pub(crate) fn extension_ad_dispatcher(&self) -> Option<Arc<dyn ExtensionAdDispatcher>> {
514 self.extension_ad_dispatcher.as_ref().map(Arc::clone)
515 }
516
517 #[cfg(feature = "autodiff")]
518 pub fn with_linearize_active_values(
519 mut self,
520 keys: std::sync::Arc<std::collections::HashSet<ValueKey<StdTensorOp>>>,
521 ) -> Self {
522 self.active_value_keys = Some(keys);
523 self
524 }
525
526 /// Whether a primal value lies on a path from the current linearize targets.
527 ///
528 /// When no active set was attached, every value is treated as active so
529 /// existing callers keep the conservative full JVP graphs.
530 #[cfg(feature = "autodiff")]
531 pub fn is_value_active_in_linearize(&self, key: &ValueKey<StdTensorOp>) -> bool {
532 self.active_value_keys
533 .as_ref()
534 .is_none_or(|set| set.contains(key))
535 }
536
537 /// Primal output keys for the operation currently being transposed.
538 ///
539 /// Primary-mode extension transpose rules such as `Eigh` use these to reuse
540 /// forward eigenvectors instead of recomputing a decomposition.
541 #[cfg(feature = "autodiff")]
542 pub fn set_transpose_primal_outputs(&mut self, keys: Option<Vec<ValueKey<StdTensorOp>>>) {
543 self.transpose_primal_outputs = keys;
544 self.transpose_primal_outputs_used = false;
545 }
546
547 /// Return the current primal outputs and mark them as consumed by this rule.
548 #[cfg(feature = "autodiff")]
549 pub fn transpose_primal_outputs(&mut self) -> Option<&[ValueKey<StdTensorOp>]> {
550 if self.transpose_primal_outputs.is_some() {
551 self.transpose_primal_outputs_used = true;
552 }
553 self.transpose_primal_outputs.as_deref()
554 }
555
556 #[cfg(feature = "autodiff")]
557 pub fn transpose_primal_outputs_were_used(&self) -> bool {
558 self.transpose_primal_outputs_used
559 }
560
561 /// Returns the guards recorded so far.
562 ///
563 /// # Examples
564 ///
565 /// ```
566 /// use tenferro_ops::ShapeGuardContext;
567 ///
568 /// let ctx = ShapeGuardContext::default();
569 /// assert_eq!(ctx.guards(), &[]);
570 /// ```
571 pub fn guards(&self) -> &[ShapeGuard] {
572 &self.guards
573 }
574
575 /// Clears all recorded guards.
576 ///
577 /// # Examples
578 ///
579 /// ```
580 /// use tenferro_ops::ShapeGuardContext;
581 ///
582 /// let mut ctx = ShapeGuardContext::default();
583 /// ctx.clear_guards();
584 /// assert!(ctx.guards().is_empty());
585 /// ```
586 pub fn clear_guards(&mut self) {
587 self.guards.clear();
588 }
589
590 /// Take the first typed shape-guard failure recorded while crossing an AD
591 /// callback boundary.
592 ///
593 /// AD callbacks expose only a message-bearing error. AD
594 /// frontends call this after the callback returns and attach the typed
595 /// value to their public runtime error.
596 #[doc(hidden)]
597 #[cfg(feature = "autodiff")]
598 pub fn take_deferred_shape_error(&mut self) -> Option<ShapeGuardError> {
599 self.deferred_shape_error
600 .lock()
601 .ok()
602 .and_then(|mut error| error.take())
603 }
604
605 /// Return the shape metadata for a value reference.
606 ///
607 /// # Examples
608 ///
609 /// ```
610 /// use computegraph::types::{ValueKey, ValueRef};
611 /// use tenferro_ops::input_key::TensorInputKey;
612 /// use tenferro_ops::std_tensor_op::StdTensorOp;
613 /// use tenferro_ops::{ShapeGuardContext, SymDim, TensorMeta};
614 /// use tenferro_tensor::DType;
615 ///
616 /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
617 /// let value = ValueRef::External(key.clone());
618 /// let mut ctx = ShapeGuardContext::default();
619 /// ctx.insert_metadata(key, TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]));
620 ///
621 /// let shape = ctx.shape_of(&value).unwrap();
622 /// assert_eq!(shape, &[SymDim::from(4usize)]);
623 /// ```
624 ///
625 /// # Errors
626 ///
627 /// Returns [`ShapeGuardError`] when the value cannot be resolved, metadata
628 /// is missing, or the metadata does not describe an exact shape.
629 pub fn shape_of(&mut self, val: &ValueRef<StdTensorOp>) -> ShapeGuardResult<Vec<SymDim>> {
630 let key = self.resolve_key(val)?.clone();
631 self.ensure_metadata_loaded(&key);
632 let meta = self.metadata.get(&key).ok_or_else(|| {
633 self.shape_guard_failure(ShapeGuardError::MissingMetadata { key: key.clone() })
634 })?;
635 meta.exact_shape()
636 .ok_or_else(|| self.shape_guard_failure(ShapeGuardError::NonExactShape { key }))
637 }
638
639 /// Return the rank for a value reference without requiring exact extents.
640 ///
641 /// Use this when an AD rule only needs axis count or needs to build
642 /// runtime-shape references. Calling [`ShapeGuardContext::shape_of`] in those
643 /// cases would reject valid values such as `DynamicTruncate` outputs whose
644 /// runtime extent is known only as an upper bound.
645 ///
646 /// # Examples
647 ///
648 /// ```
649 /// use computegraph::types::{ValueKey, ValueRef};
650 /// use tenferro_ops::input_key::TensorInputKey;
651 /// use tenferro_ops::std_tensor_op::StdTensorOp;
652 /// use tenferro_ops::{ShapeExtent, ShapeGuardContext, SymDim, TensorMeta};
653 /// use tenferro_tensor::DType;
654 ///
655 /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
656 /// let value = ValueRef::External(key.clone());
657 /// let mut ctx = ShapeGuardContext::default();
658 /// ctx.insert_metadata(
659 /// key,
660 /// TensorMeta::with_extents(DType::F64, vec![ShapeExtent::upper_bound(SymDim::from(8usize))]),
661 /// );
662 ///
663 /// assert_eq!(ctx.rank_of(&value).unwrap(), 1);
664 /// ```
665 ///
666 /// # Errors
667 ///
668 /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
669 /// metadata is unavailable.
670 pub fn rank_of(&mut self, val: &ValueRef<StdTensorOp>) -> ShapeGuardResult<usize> {
671 self.metadata_of(val).map(TensorMeta::rank)
672 }
673
674 /// Return per-axis shape guarantees for a value reference.
675 ///
676 /// # Examples
677 ///
678 /// ```
679 /// use computegraph::types::{ValueKey, ValueRef};
680 /// use tenferro_ops::input_key::TensorInputKey;
681 /// use tenferro_ops::std_tensor_op::StdTensorOp;
682 /// use tenferro_ops::{ShapeExtent, ShapeGuardContext, SymDim, TensorMeta};
683 /// use tenferro_tensor::DType;
684 ///
685 /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
686 /// let value = ValueRef::External(key.clone());
687 /// let mut ctx = ShapeGuardContext::default();
688 /// ctx.insert_metadata(
689 /// key,
690 /// TensorMeta::with_extents(DType::F64, vec![ShapeExtent::upper_bound(SymDim::from(8usize))]),
691 /// );
692 ///
693 /// let extents = ctx.extents_of(&value).unwrap();
694 /// assert_eq!(extents[0], ShapeExtent::upper_bound(SymDim::from(8usize)));
695 /// ```
696 ///
697 /// # Errors
698 ///
699 /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
700 /// metadata is unavailable.
701 pub fn extents_of(
702 &mut self,
703 val: &ValueRef<StdTensorOp>,
704 ) -> ShapeGuardResult<&[ShapeExtent<SymDim>]> {
705 self.metadata_of(val).map(TensorMeta::extents)
706 }
707
708 /// Return the exact shape for a value reference, if all axes are exact.
709 ///
710 /// # Examples
711 ///
712 /// ```
713 /// use computegraph::types::{ValueKey, ValueRef};
714 /// use tenferro_ops::input_key::TensorInputKey;
715 /// use tenferro_ops::std_tensor_op::StdTensorOp;
716 /// use tenferro_ops::{ShapeExtent, ShapeGuardContext, SymDim, TensorMeta};
717 /// use tenferro_tensor::DType;
718 ///
719 /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
720 /// let value = ValueRef::External(key.clone());
721 /// let mut ctx = ShapeGuardContext::default();
722 /// ctx.insert_metadata(
723 /// key,
724 /// TensorMeta::with_extents(DType::F64, vec![ShapeExtent::upper_bound(SymDim::from(8usize))]),
725 /// );
726 ///
727 /// let maybe_shape = ctx.exact_shape_of(&value).unwrap();
728 /// assert_eq!(maybe_shape, None);
729 /// ```
730 ///
731 /// # Errors
732 ///
733 /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
734 /// metadata is unavailable.
735 pub fn exact_shape_of(
736 &mut self,
737 val: &ValueRef<StdTensorOp>,
738 ) -> ShapeGuardResult<Option<Vec<SymDim>>> {
739 self.metadata_of(val).map(TensorMeta::exact_shape)
740 }
741
742 #[doc(hidden)]
743 pub fn shape_if_available(&mut self, val: &ValueRef<StdTensorOp>) -> Option<Vec<SymDim>> {
744 self.metadata_if_available(val)
745 .and_then(TensorMeta::exact_shape)
746 }
747
748 /// Return the dtype metadata for a value reference.
749 ///
750 /// # Examples
751 ///
752 /// ```
753 /// use computegraph::types::{ValueKey, ValueRef};
754 /// use tenferro_ops::input_key::TensorInputKey;
755 /// use tenferro_ops::std_tensor_op::StdTensorOp;
756 /// use tenferro_ops::{ShapeGuardContext, SymDim, TensorMeta};
757 /// use tenferro_tensor::DType;
758 ///
759 /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
760 /// let value = ValueRef::External(key.clone());
761 /// let mut ctx = ShapeGuardContext::default();
762 /// ctx.insert_metadata(key, TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]));
763 ///
764 /// let dtype = ctx.dtype_of(&value).unwrap();
765 /// assert_eq!(dtype, DType::F64);
766 /// ```
767 ///
768 /// # Errors
769 ///
770 /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
771 /// metadata is unavailable.
772 pub fn dtype_of(&mut self, val: &ValueRef<StdTensorOp>) -> ShapeGuardResult<DType> {
773 self.metadata_of(val).map(|meta| meta.dtype)
774 }
775
776 /// Return the complete metadata record for a value reference.
777 ///
778 /// # Examples
779 ///
780 /// ```
781 /// use computegraph::types::{ValueKey, ValueRef};
782 /// use tenferro_ops::input_key::TensorInputKey;
783 /// use tenferro_ops::std_tensor_op::StdTensorOp;
784 /// use tenferro_ops::{ShapeGuardContext, SymDim, TensorMeta};
785 /// use tenferro_tensor::DType;
786 ///
787 /// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 1 });
788 /// let value = ValueRef::External(key.clone());
789 /// let mut ctx = ShapeGuardContext::default();
790 /// ctx.insert_metadata(key, TensorMeta::exact(DType::F64, vec![SymDim::from(4usize)]));
791 ///
792 /// let meta = ctx.metadata_of(&value).unwrap();
793 /// assert_eq!(meta.dtype, DType::F64);
794 /// ```
795 ///
796 /// # Errors
797 ///
798 /// Returns [`ShapeGuardError`] when the value cannot be resolved or its
799 /// metadata is unavailable.
800 pub fn metadata_of(&mut self, val: &ValueRef<StdTensorOp>) -> ShapeGuardResult<&TensorMeta> {
801 let key = self.resolve_key(val)?.clone();
802 self.ensure_metadata_loaded(&key);
803 self.metadata
804 .get(&key)
805 .ok_or_else(|| self.shape_guard_failure(ShapeGuardError::MissingMetadata { key }))
806 }
807
808 #[doc(hidden)]
809 pub fn metadata_if_available(&mut self, val: &ValueRef<StdTensorOp>) -> Option<&TensorMeta> {
810 let key = self.resolve_key_if_available(val)?.clone();
811 self.ensure_metadata_loaded(&key);
812 self.metadata.get(&key)
813 }
814
815 #[doc(hidden)]
816 pub fn attach_graph(&mut self, graph: &Graph<StdTensorOp>) {
817 self.local_keys = Some(graph.values().iter().map(|node| node.key.clone()).collect());
818 }
819
820 #[doc(hidden)]
821 pub fn insert_metadata(&mut self, key: ValueKey<StdTensorOp>, meta: TensorMeta) {
822 self.metadata.insert(key, meta);
823 }
824
825 #[doc(hidden)]
826 pub fn extend_metadata<I>(&mut self, entries: I)
827 where
828 I: IntoIterator<Item = (ValueKey<StdTensorOp>, TensorMeta)>,
829 {
830 self.metadata.extend(entries);
831 }
832
833 fn resolve_key_if_available<'a>(
834 &'a self,
835 val: &'a ValueRef<StdTensorOp>,
836 ) -> Option<&'a ValueKey<StdTensorOp>> {
837 match val {
838 ValueRef::External(key) => Some(key),
839 ValueRef::Local(local_id) => self
840 .local_keys
841 .as_ref()
842 .and_then(|keys| keys.get(*local_id)),
843 }
844 }
845
846 fn resolve_key<'a>(
847 &'a self,
848 val: &'a ValueRef<StdTensorOp>,
849 ) -> ShapeGuardResult<&'a ValueKey<StdTensorOp>> {
850 match val {
851 ValueRef::External(key) => Ok(key),
852 ValueRef::Local(local_id) if self.local_keys.is_none() => Err(self
853 .shape_guard_failure(ShapeGuardError::LocalWithoutAttachedGraph {
854 local_id: *local_id,
855 })),
856 ValueRef::Local(local_id) => self
857 .local_keys
858 .as_ref()
859 .and_then(|keys| keys.get(*local_id))
860 .ok_or_else(|| {
861 self.shape_guard_failure(ShapeGuardError::LocalOutOfBounds {
862 local_id: *local_id,
863 })
864 }),
865 }
866 }
867
868 #[cfg(feature = "autodiff")]
869 fn shape_guard_failure(&self, source: ShapeGuardError) -> ShapeGuardFailure {
870 ShapeGuardFailure::new(source, Arc::clone(&self.deferred_shape_error))
871 }
872
873 #[cfg(not(feature = "autodiff"))]
874 fn shape_guard_failure(&self, source: ShapeGuardError) -> ShapeGuardFailure {
875 ShapeGuardFailure::new(source)
876 }
877
878 fn ensure_metadata_loaded(&mut self, key: &ValueKey<StdTensorOp>) {
879 if !self.metadata.contains_key(key) && self.use_global_registry {
880 if let Ok(Some(meta)) = lookup_global_metadata(key) {
881 self.metadata.insert(key.clone(), meta);
882 }
883 }
884 }
885}
886
887/// Look up a single metadata entry from the global registry.
888///
889/// Locks the registry briefly for a single `HashMap::get` + clone.
890///
891/// # Examples
892///
893/// ```
894/// use computegraph::types::ValueKey;
895/// use tenferro_ops::ad::context::lookup_global_metadata;
896/// use tenferro_ops::input_key::TensorInputKey;
897/// use tenferro_ops::std_tensor_op::StdTensorOp;
898///
899/// let key = ValueKey::<StdTensorOp>::Input(TensorInputKey::User { id: 99 });
900/// let meta = lookup_global_metadata(&key).unwrap();
901/// assert!(meta.is_none());
902/// ```
903///
904/// # Errors
905///
906/// Returns [`MetadataRegistryError::LockPoisoned`] when the global metadata
907/// registry lock is poisoned.
908pub fn lookup_global_metadata(
909 key: &ValueKey<StdTensorOp>,
910) -> Result<Option<TensorMeta>, MetadataRegistryError> {
911 let guard = global_metadata_registry()
912 .lock()
913 .map_err(|_| MetadataRegistryError::LockPoisoned)?;
914 Ok(guard
915 .get(key)
916 .and_then(|entry| entry.stack.last())
917 .map(|registration| registration.meta.clone()))
918}
919
920#[doc(hidden)]
921///
922/// # Errors
923///
924/// Returns [`MetadataRegistryError::LockPoisoned`] when the global metadata
925/// registry lock is poisoned.
926pub fn register_scoped_global_metadata_batch<I>(
927 entries: I,
928) -> Result<GlobalMetadataScope, MetadataRegistryError>
929where
930 I: IntoIterator<Item = (ValueKey<StdTensorOp>, TensorMeta)>,
931{
932 let mut guard = global_metadata_registry()
933 .lock()
934 .map_err(|_| MetadataRegistryError::LockPoisoned)?;
935 let mut registrations = Vec::new();
936 for (key, meta) in entries {
937 let token = NEXT_GLOBAL_METADATA_TOKEN.fetch_add(1, AtomicOrdering::Relaxed);
938 let entry = guard
939 .entry(key.clone())
940 .or_insert_with(|| GlobalMetadataEntry { stack: Vec::new() });
941 entry.stack.push(GlobalMetadataRegistration { token, meta });
942 registrations.push(ScopedGlobalMetadataRegistration { key, token });
943 }
944 Ok(GlobalMetadataScope { registrations })
945}
946
947fn release_scoped_global_metadata(registrations: &[ScopedGlobalMetadataRegistration]) {
948 let Ok(mut guard) = global_metadata_registry().lock() else {
949 // Drop cannot return an error. Failing closed here avoids reading or
950 // mutating data from a poisoned registry at the cost of leaking entries
951 // until process exit.
952 return;
953 };
954 for registration in registrations {
955 let should_remove = if let Some(entry) = guard.get_mut(®istration.key) {
956 if let Some(position) = entry
957 .stack
958 .iter()
959 .rposition(|candidate| candidate.token == registration.token)
960 {
961 entry.stack.remove(position);
962 }
963 entry.stack.is_empty()
964 } else {
965 false
966 };
967 if should_remove {
968 guard.remove(®istration.key);
969 }
970 }
971}
972
973/// Resolve a [`DimExpr`] to a concrete `usize`.
974#[doc(hidden)]
975pub fn resolve_dim(dim: &DimExpr) -> Result<usize, DimExprEvalError> {
976 dim.eval(&[])
977}
978
979/// Resolve matrix dimensions and record their ordering as a guard.
980#[doc(hidden)]
981pub fn resolve_and_guard(
982 m: &DimExpr,
983 n: &DimExpr,
984 ctx: &mut ShapeGuardContext,
985) -> Result<(usize, usize), DimExprEvalError> {
986 let m_size = resolve_dim(m)?;
987 let n_size = resolve_dim(n)?;
988 ctx.guards.push(ShapeGuard {
989 dim_a: m_size,
990 dim_b: n_size,
991 ordering: m_size.cmp(&n_size),
992 });
993 Ok((m_size, n_size))
994}