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}