Skip to main content

tenferro_tensor_core/
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 and gets the value enum, the tag,
5//! and the membership query from that single declaration. A crate that needs a
6//! different set declares it with [`define_scalar_set!`] in its own crate, and
7//! its values implement this trait without touching tenferro's set.
8
9/// Kind of arithmetic a set member belongs to.
10///
11/// # Examples
12///
13/// ```rust
14/// use tenferro_tensor_core::MemberKind;
15///
16/// assert_ne!(MemberKind::Integer, MemberKind::Float);
17/// ```
18#[derive(Clone, Copy, Debug, PartialEq, Eq)]
19#[non_exhaustive]
20pub enum MemberKind {
21    /// Boolean.
22    Boolean,
23    /// Signed integer.
24    Integer,
25    /// Real floating point.
26    Float,
27    /// Complex floating point.
28    Complex,
29    /// A scalar tenferro does not declare, so its facts are unknown.
30    External,
31}
32
33/// Promotion-relevant facts about one set member.
34///
35/// `level` orders members within a kind and `width` is the component width in
36/// bits. `level` is what orders two members of the same kind, so a set can rank
37/// an extended-precision member above a standard one without claiming a wider
38/// exponent range.
39///
40/// # Examples
41///
42/// ```rust
43/// use tenferro_tensor_core::{MemberKind, MemberSpec};
44///
45/// let widened = tenferro_tensor_core::promote_specs(
46///     MemberSpec::new(MemberKind::Float, 0, 32),
47///     MemberSpec::new(MemberKind::Float, 1, 64),
48/// );
49/// assert_eq!(widened, MemberSpec::new(MemberKind::Float, 1, 64));
50/// ```
51#[derive(Clone, Copy, Debug, PartialEq, Eq)]
52pub struct MemberSpec {
53    /// Arithmetic kind of the member.
54    pub kind: MemberKind,
55    /// Rank within the kind.
56    pub level: u32,
57    /// Component width in bits.
58    pub width: u32,
59}
60
61impl MemberSpec {
62    /// Build a member fact.
63    ///
64    /// # Examples
65    ///
66    /// ```rust
67    /// use tenferro_tensor_core::{MemberKind, MemberSpec};
68    ///
69    /// let spec = MemberSpec::new(MemberKind::Float, 1, 64);
70    /// assert_eq!(spec.level, 1);
71    /// ```
72    #[must_use]
73    pub const fn new(kind: MemberKind, level: u32, width: u32) -> Self {
74        Self { kind, level, width }
75    }
76}
77
78/// Combine two member facts under the ordinary numeric promotion rules.
79///
80/// A boolean yields to anything. Two members of one kind keep the higher level.
81/// An integer with a float or complex yields to that kind's widest member,
82/// because an integer is widened rather than mixed. A float with a complex takes
83/// the narrowest complex that can still hold both, which keeps `f32 + c32` in
84/// `c32` while `f64 + c32` becomes `c64`.
85///
86/// # Examples
87///
88/// ```rust
89/// use tenferro_tensor_core::{promote_specs, MemberKind, MemberSpec};
90///
91/// let f32_ = MemberSpec::new(MemberKind::Float, 0, 32);
92/// let c32 = MemberSpec::new(MemberKind::Complex, 0, 32);
93/// assert_eq!(promote_specs(f32_, c32), c32);
94/// ```
95#[must_use]
96pub const fn promote_specs(lhs: MemberSpec, rhs: MemberSpec) -> MemberSpec {
97    use MemberKind::{Boolean, Complex, External, Float, Integer};
98    match (lhs.kind, rhs.kind) {
99        (Boolean, _) => rhs,
100        (_, Boolean) => lhs,
101        (External, _) => lhs,
102        (_, External) => rhs,
103        (Integer, Float) | (Float, Integer) => MemberSpec::new(Float, u32::MAX, u32::MAX),
104        (Integer, Complex) | (Complex, Integer) => MemberSpec::new(Complex, u32::MAX, u32::MAX),
105        (Integer, Integer) | (Float, Float) | (Complex, Complex) => {
106            if lhs.level >= rhs.level {
107                lhs
108            } else {
109                rhs
110            }
111        }
112        (Float, Complex) | (Complex, Float) => {
113            let width = if lhs.width >= rhs.width {
114                lhs.width
115            } else {
116                rhs.width
117            };
118            MemberSpec::new(Complex, 0, width)
119        }
120    }
121}
122
123/// Promote two members of one set from its parallel tag and fact tables.
124///
125/// `tags` and `specs` are emitted together in the same order by
126/// [`define_scalar_tag!`], so a tag that is not listed in `tags` (an externally
127/// defined member) promotes to itself.
128///
129/// # Examples
130///
131/// ```rust
132/// use tenferro_tensor_core::{promote_in_set, DType};
133///
134/// assert_eq!(
135///     promote_in_set(DType::TAGS, DType::SPECS, DType::I32, DType::F32),
136///     DType::F64
137/// );
138/// ```
139#[must_use]
140pub fn promote_in_set<Tag: Copy + PartialEq>(
141    tags: &[Tag],
142    specs: &[MemberSpec],
143    lhs: Tag,
144    rhs: Tag,
145) -> Tag {
146    // INVARIANT: `define_scalar_tag!` emits `TAGS` and `SPECS` from one member
147    // list, so the two tables always hold the same entries in the same order.
148    let (Some(lhs_index), Some(rhs_index)) = (
149        tags.iter().position(|tag| *tag == lhs),
150        tags.iter().position(|tag| *tag == rhs),
151    ) else {
152        return lhs;
153    };
154    let (Some(lhs_spec), Some(rhs_spec)) = (specs.get(lhs_index), specs.get(rhs_index)) else {
155        return lhs;
156    };
157    let target = promote_specs(*lhs_spec, *rhs_spec);
158    for (index, spec) in specs.iter().enumerate() {
159        if *spec == target {
160            if let Some(tag) = tags.get(index) {
161                return *tag;
162            }
163        }
164    }
165    let mut chosen: Option<(Tag, MemberSpec)> = None;
166    for (index, spec) in specs.iter().enumerate() {
167        let Some(tag) = tags.get(index).copied() else {
168            break;
169        };
170        if spec.kind != target.kind {
171            continue;
172        }
173        let better = match chosen {
174            None => true,
175            Some((_, current)) => {
176                let current_wide_enough = current.width >= target.width;
177                let candidate_wide_enough = spec.width >= target.width;
178                match (current_wide_enough, candidate_wide_enough) {
179                    (false, true) => true,
180                    (true, false) => false,
181                    (true, true) => {
182                        spec.width < current.width
183                            || (spec.width == current.width && spec.level < current.level)
184                    }
185                    (false, false) => {
186                        spec.width > current.width
187                            || (spec.width == current.width && spec.level < current.level)
188                    }
189                }
190            }
191        };
192        if better {
193            chosen = Some((tag, *spec));
194        }
195    }
196    match chosen {
197        Some((tag, _)) => tag,
198        None => lhs,
199    }
200}
201
202/// Define a closed scalar set's tag type and its promotion facts.
203///
204/// The declaration lists each member once. The macro emits the tag enum and the
205/// member facts (`spec`, `TAGS`, `SPECS`), plus the `external` variant when the
206/// declaration names one. Use this form when the value enum lives in another
207/// crate: tenferro's own preset set declares its tag here and its value enum
208/// beside the tensor family that carries it.
209///
210/// # Examples
211///
212/// ```rust
213/// use tenferro_tensor_core::{define_scalar_tag, MemberKind};
214///
215/// define_scalar_tag! {
216///     /// Tag for a two-member set.
217///     pub enum PairTag {
218///         /// Double precision.
219///         F64 => f64 : Float 1 64,
220///         /// Single precision.
221///         F32 => f32 : Float 0 32,
222///     }
223/// }
224///
225/// assert_eq!(PairTag::F32.spec().kind, MemberKind::Float);
226/// assert_eq!(PairTag::TAGS, &[PairTag::F64, PairTag::F32]);
227/// ```
228#[macro_export]
229macro_rules! define_scalar_tag {
230    (
231        $(#[$tag_meta:meta])*
232        $tag_vis:vis enum $tag:ident {
233            $(
234                $(#[$variant_meta:meta])*
235                $variant:ident => $ty:ty : $kind:ident $level:literal $width:literal
236            ),+ $(,)?
237        }
238        $( external $ext_variant:ident($ext_ty:ty); )?
239    ) => {
240        $(#[$tag_meta])*
241        #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
242        $tag_vis enum $tag {
243            $(
244                $(#[$variant_meta])*
245                $variant,
246            )+
247            $( $ext_variant($ext_ty), )?
248        }
249
250        impl $tag {
251            /// Promotion facts of this member.
252            ///
253            /// # Examples
254            ///
255            /// ```rust
256            /// use tenferro_tensor_core::{MemberKind, DType};
257            ///
258            /// assert_eq!(DType::F64.spec().kind, MemberKind::Float);
259            /// ```
260            #[must_use]
261            pub const fn spec(self) -> $crate::MemberSpec {
262                match self {
263                    $(
264                        $tag::$variant => $crate::MemberSpec::new(
265                            $crate::MemberKind::$kind,
266                            $level,
267                            $width,
268                        ),
269                    )+
270                    $( $tag::$ext_variant(_) => {
271                        $crate::MemberSpec::new($crate::MemberKind::External, 0, 0)
272                    } )?
273                }
274            }
275
276            /// Every member tag, in declaration order.
277            pub const TAGS: &'static [Self] = &[
278                $(
279                    $tag::$variant,
280                )+
281            ];
282
283            /// Promotion facts of every member, in declaration order.
284            pub const SPECS: &'static [$crate::MemberSpec] = &[
285                $(
286                    $crate::MemberSpec::new($crate::MemberKind::$kind, $level, $width),
287                )+
288            ];
289        }
290    };
291}