tenferro_tensor/scalar_set.rs
1//! Closed scalar sets.
2//!
3//! A scalar set is the closed list of scalar types one tensor value type can
4//! carry. tenferro declares its own set once (as
5//! [`DefaultScalars`](crate::DefaultScalars)) and gets the value enum, the tag,
6//! and the membership query from that single declaration. A crate that needs a
7//! different set declares it with [`define_scalar_set!`] in its own crate, and
8//! its values implement this trait without touching tenferro's set.
9//!
10//! The trait and its declaration macro live with the host container they build
11//! values from; the promotion facts themselves stay in `tenferro-tensor-core`.
12
13/// A closed set of scalar types carried by one tensor value type.
14///
15/// A set is represented by one value type: the default set's payload is opaque
16/// and a downstream set defines its own value enum, and [`ScalarSet::tag`]
17/// reports which member a value currently holds.
18///
19/// # Examples
20///
21/// ```rust
22/// use tenferro_tensor::{DefaultScalars, ScalarSet};
23///
24/// let value = DefaultScalars::from_vec_col_major(vec![1], vec![1.0_f64])?;
25/// assert_eq!(value.tag(), tenferro_tensor_core::DType::F64);
26/// # Ok::<(), tenferro_tensor::Error>(())
27/// ```
28pub trait ScalarSet: Clone + core::fmt::Debug + 'static {
29 /// Tag identifying one member of this set.
30 type Tag: Copy + Eq + core::fmt::Debug + 'static;
31
32 /// Tags of every member, in declaration order.
33 const TAGS: &'static [Self::Tag];
34
35 /// Tag of the member this value currently holds.
36 ///
37 /// # Examples
38 ///
39 /// ```rust
40 /// use tenferro_tensor::{DefaultScalars, DType, ScalarSet};
41 ///
42 /// let value = DefaultScalars::from_vec_col_major(vec![1], vec![1.0_f64])?;
43 /// assert_eq!(value.tag(), DType::F64);
44 /// # Ok::<(), tenferro_tensor::Error>(())
45 /// ```
46 fn tag(&self) -> Self::Tag;
47
48 /// Promote two members of this set to the member that represents both.
49 ///
50 /// # Examples
51 ///
52 /// ```rust
53 /// use tenferro_tensor::{DefaultScalars, DType, ScalarSet};
54 ///
55 /// assert_eq!(
56 /// <DefaultScalars as ScalarSet>::promote(DType::I32, DType::F32),
57 /// DType::F64
58 /// );
59 /// ```
60 fn promote(lhs: Self::Tag, rhs: Self::Tag) -> Self::Tag;
61}
62
63/// Define a closed scalar set: its tag type, its value enum, and its membership.
64///
65/// The declaration lists each member once. The macro emits the tag enum, the
66/// value enum whose variants hold a
67/// host [`TypedTensor`](crate::TypedTensor) (`TypedTensor<T, DynRank, Host>`)
68/// of the member type, and the
69/// [`ScalarSet`] implementation. A downstream crate invokes this in its own
70/// crate, so tenferro never needs to know the set.
71///
72/// # Examples
73///
74/// ```rust
75/// use tenferro_tensor::{define_scalar_set, DynRank, Host, ScalarSet, TypedTensor};
76///
77/// define_scalar_set! {
78/// /// Tag for a two-member set.
79/// pub enum PairTag {
80/// /// Double precision.
81/// F64 => f64 : Float 1 64,
82/// /// Single precision.
83/// F32 => f32 : Float 0 32,
84/// }
85/// /// Value enum for a two-member set.
86/// pub enum Pair;
87/// }
88///
89/// let value = Pair::F32(TypedTensor::<f32, DynRank, Host>::from_host_vec_col_major(
90/// vec![1],
91/// vec![1.0_f32],
92/// )?);
93/// assert_eq!(value.tag(), PairTag::F32);
94/// assert_eq!(<Pair as ScalarSet>::TAGS, &[PairTag::F64, PairTag::F32]);
95/// # Ok::<(), tenferro_tensor::Error>(())
96/// ```
97#[macro_export]
98macro_rules! define_scalar_set {
99 (
100 $(#[$tag_meta:meta])*
101 $tag_vis:vis enum $tag:ident {
102 $(
103 $(#[$variant_meta:meta])*
104 $variant:ident => $ty:ty : $kind:ident $level:literal $width:literal
105 ),+ $(,)?
106 }
107 $(#[$set_meta:meta])*
108 $set_vis:vis enum $set:ident;
109 $( external $ext_variant:ident($ext_ty:ty); )?
110 ) => {
111 ::tenferro_tensor_core::define_scalar_tag! {
112 $(#[$tag_meta])*
113 $tag_vis enum $tag {
114 $(
115 $(#[$variant_meta])*
116 $variant => $ty : $kind $level $width
117 ),+
118 }
119 $( external $ext_variant($ext_ty); )?
120 }
121
122 $(#[$set_meta])*
123 #[derive(Clone, Debug)]
124 $set_vis enum $set {
125 $(
126 $(#[$variant_meta])*
127 $variant($crate::TypedTensor<$ty, $crate::DynRank, $crate::Host>),
128 )+
129 }
130
131 impl $crate::ScalarSet for $set {
132 type Tag = $tag;
133
134 const TAGS: &'static [Self::Tag] = &[
135 $(
136 $tag::$variant,
137 )+
138 ];
139
140 fn tag(&self) -> Self::Tag {
141 match self {
142 $(
143 $set::$variant(_) => $tag::$variant,
144 )+
145 }
146 }
147
148 fn promote(lhs: Self::Tag, rhs: Self::Tag) -> Self::Tag {
149 $(
150 if matches!(lhs, $tag::$ext_variant(_)) {
151 return lhs;
152 }
153 if matches!(rhs, $tag::$ext_variant(_)) {
154 return rhs;
155 }
156 )?
157 ::tenferro_tensor_core::promote_in_set(<$tag>::TAGS, <$tag>::SPECS, lhs, rhs)
158 }
159 }
160 };
161}