Skip to main content

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}