Skip to main content

tenferro_bf16_proof/
lib.rs

1//! Standard bfloat16 carried through the external-scalar boundary.
2//!
3//! #1785 asks for `half::bf16` as a standard scalar representation, with construction, views,
4//! materialization, explicit bf16/f32 conversions, basic forward arithmetic, and sum reduction
5//! through the same public boundary the other external scalars use. This crate supplies exactly
6//! that: the representation *is* [`half::bf16`], and the thin wrapper exists only because a
7//! foreign type cannot implement tenferro's local scalar traits, which #1785 accepts for an
8//! external contribution.
9//!
10//! # Declared arithmetic, accumulation, and rounding
11//!
12//! The behaviour is specified rather than implied, which is what #1785 asks for:
13//!
14//! - **Storage** is `half::bf16`, so a stored value is the nearest bfloat16 to what was written.
15//! - **A single operation** is computed in `f32` and rounded back to bfloat16 once, so the result
16//!   carries the rounding error of the operation and nothing else.
17//! - **A reduction** accumulates in `f32` and rounds once at the end, so the accumulation does not
18//!   quantize at every step. [`reduction::sum_in_f32_accumulation`] states this, and
19//!   `tests/reduction_precision.rs` distinguishes that contract from repeated bfloat16 rounding,
20//!   which is what #1785 requires of a promise like this one.
21//!
22//! Repeated bfloat16 rounding is a different and weaker contract, so the tests measure the
23//! difference rather than asserting only that the sum is close.
24
25#![deny(missing_docs)]
26
27pub mod conversion;
28pub mod einsum;
29pub mod reduction;
30
31/// The standard bfloat16 representation this contribution stores.
32pub use half::bf16;
33
34/// A bfloat16 scalar in tenferro's scalar contract.
35///
36/// The representation is [`bf16`], and the wrapper carries no extra state.
37///
38/// # Examples
39///
40/// ```rust
41/// use tenferro_bf16_proof::Bf16;
42///
43/// let value = Bf16::from_f32(1.5);
44/// assert_eq!(value.to_f32(), 1.5);
45/// assert_eq!(value.narrow(), half::bf16::from_f32(1.5));
46/// ```
47#[derive(Clone, Copy, Debug, Default, PartialEq)]
48pub struct Bf16(pub bf16);
49
50impl Bf16 {
51    /// Wrap a bfloat16 value.
52    ///
53    /// # Examples
54    ///
55    /// ```rust
56    /// use tenferro_bf16_proof::Bf16;
57    ///
58    /// assert_eq!(Bf16::of(half::bf16::from_f32(2.0)).to_f32(), 2.0);
59    /// ```
60    #[must_use]
61    pub const fn of(value: bf16) -> Self {
62        Self(value)
63    }
64
65    /// The stored bfloat16 value.
66    ///
67    /// # Examples
68    ///
69    /// ```rust
70    /// use tenferro_bf16_proof::Bf16;
71    ///
72    /// assert_eq!(Bf16::from_f32(1.0).narrow(), half::bf16::from_f32(1.0));
73    /// ```
74    #[must_use]
75    pub const fn narrow(self) -> bf16 {
76        self.0
77    }
78
79    /// Round an `f64` to bfloat16, which is the only way a stored value is produced.
80    ///
81    /// # Examples
82    ///
83    /// ```rust
84    /// use tenferro_bf16_proof::Bf16;
85    ///
86    /// // The spacing of bfloat16 on [1, 2) is 2^-8, so 1.001 is nearer to 1.0 than to the next
87    /// // representable value above it.
88    /// let stored = Bf16::from_f64(1.001);
89    /// assert_eq!(stored.to_f32(), 1.0);
90    /// ```
91    #[must_use]
92    pub fn from_f64(value: f64) -> Self {
93        Self(bf16::from_f64(value))
94    }
95
96    /// Round an `f32` to bfloat16.
97    ///
98    /// # Examples
99    ///
100    /// ```rust
101    /// use tenferro_bf16_proof::Bf16;
102    ///
103    /// assert_eq!(Bf16::from_f32(0.5).to_f32(), 0.5);
104    /// ```
105    #[must_use]
106    pub fn from_f32(value: f32) -> Self {
107        Self(bf16::from_f32(value))
108    }
109
110    /// Widen the stored value to `f32`, which is exact.
111    ///
112    /// # Examples
113    ///
114    /// ```rust
115    /// use tenferro_bf16_proof::Bf16;
116    ///
117    /// let stored = Bf16::from_f64(1.0 + 2f64.powi(-10));
118    /// assert_eq!(stored.to_f32(), 1.0);
119    /// ```
120    #[must_use]
121    pub fn to_f32(self) -> f32 {
122        self.0.to_f32()
123    }
124
125    /// Widen the stored value to `f64`, which is exact.
126    ///
127    /// # Examples
128    ///
129    /// ```rust
130    /// use tenferro_bf16_proof::Bf16;
131    ///
132    /// assert_eq!(Bf16::from_f32(2.5).to_f64(), 2.5);
133    /// ```
134    #[must_use]
135    pub fn to_f64(self) -> f64 {
136        f64::from(self.0.to_f32())
137    }
138
139    /// Additive identity.
140    ///
141    /// # Examples
142    ///
143    /// ```rust
144    /// use tenferro_bf16_proof::Bf16;
145    ///
146    /// assert_eq!(Bf16::zero().to_f32(), 0.0);
147    /// ```
148    #[must_use]
149    pub const fn zero() -> Self {
150        Self(bf16::from_bits(0x0000))
151    }
152
153    /// Multiplicative identity.
154    ///
155    /// # Examples
156    ///
157    /// ```rust
158    /// use tenferro_bf16_proof::Bf16;
159    ///
160    /// assert_eq!(Bf16::one().to_f32(), 1.0);
161    /// ```
162    #[must_use]
163    pub fn one() -> Self {
164        Self::from_f32(1.0)
165    }
166}
167
168impl std::ops::Add for Bf16 {
169    type Output = Self;
170
171    /// Add in `f32` and round the sum once.
172    ///
173    /// # Examples
174    ///
175    /// ```rust
176    /// use tenferro_bf16_proof::Bf16;
177    ///
178    /// assert_eq!((Bf16::from_f32(1.0) + Bf16::from_f32(2.0)).to_f32(), 3.0);
179    /// ```
180    fn add(self, rhs: Self) -> Self {
181        Self::from_f32(self.to_f32() + rhs.to_f32())
182    }
183}
184
185impl std::ops::Sub for Bf16 {
186    type Output = Self;
187
188    /// Subtract in `f32` and round the difference once.
189    ///
190    /// # Examples
191    ///
192    /// ```rust
193    /// use tenferro_bf16_proof::Bf16;
194    ///
195    /// assert_eq!((Bf16::from_f32(3.0) - Bf16::from_f32(2.0)).to_f32(), 1.0);
196    /// ```
197    fn sub(self, rhs: Self) -> Self {
198        Self::from_f32(self.to_f32() - rhs.to_f32())
199    }
200}
201
202impl std::ops::Mul for Bf16 {
203    type Output = Self;
204
205    /// Multiply in `f32` and round the product once.
206    ///
207    /// # Examples
208    ///
209    /// ```rust
210    /// use tenferro_bf16_proof::Bf16;
211    ///
212    /// assert_eq!((Bf16::from_f32(3.0) * Bf16::from_f32(4.0)).to_f32(), 12.0);
213    /// ```
214    fn mul(self, rhs: Self) -> Self {
215        Self::from_f32(self.to_f32() * rhs.to_f32())
216    }
217}
218
219impl std::ops::Neg for Bf16 {
220    type Output = Self;
221
222    /// Negate, which flips the sign bit and needs no rounding.
223    ///
224    /// # Examples
225    ///
226    /// ```rust
227    /// use tenferro_bf16_proof::Bf16;
228    ///
229    /// assert_eq!((-Bf16::from_f32(2.0)).to_f32(), -2.0);
230    /// ```
231    fn neg(self) -> Self {
232        Self(-self.0)
233    }
234}
235
236impl tenferro_tensor_core::Scalar for Bf16 {
237    const DOMAIN: tenferro_tensor_core::ScalarDomain = tenferro_tensor_core::ScalarDomain::Field;
238}
239
240impl tenferro_tensor_core::ScalarArithmetic for Bf16 {
241    fn scalar_zero() -> Self {
242        Self::zero()
243    }
244
245    fn scalar_one() -> Self {
246        Self::one()
247    }
248
249    fn scalar_add(self, rhs: Self) -> Self {
250        self + rhs
251    }
252
253    fn scalar_sub(self, rhs: Self) -> Self {
254        self - rhs
255    }
256
257    fn scalar_mul(self, rhs: Self) -> Self {
258        self * rhs
259    }
260}
261
262/// bfloat16 addition as an operation type.
263///
264/// The operation is a type rather than a closure, so tenferro's kernels are instantiated once per
265/// element type and operation instead of once per call site.
266///
267/// # Examples
268///
269/// ```rust
270/// use tenferro_bf16_proof::{Bf16, Bf16Add};
271/// use tenferro_cpu::BinaryScalarOp;
272///
273/// let sum = <Bf16Add as BinaryScalarOp<Bf16>>::apply(Bf16::from_f32(1.0), Bf16::from_f32(2.0));
274/// assert_eq!(sum.to_f32(), 3.0);
275/// ```
276pub struct Bf16Add;
277
278impl tenferro_cpu::BinaryScalarOp<Bf16> for Bf16Add {
279    fn apply(lhs: Bf16, rhs: Bf16) -> Bf16 {
280        lhs + rhs
281    }
282}
283
284/// bfloat16 subtraction as an operation type.
285///
286/// # Examples
287///
288/// ```rust
289/// use tenferro_bf16_proof::{Bf16, Bf16Sub};
290/// use tenferro_cpu::BinaryScalarOp;
291///
292/// let difference =
293///     <Bf16Sub as BinaryScalarOp<Bf16>>::apply(Bf16::from_f32(3.0), Bf16::from_f32(2.0));
294/// assert_eq!(difference.to_f32(), 1.0);
295/// ```
296pub struct Bf16Sub;
297
298impl tenferro_cpu::BinaryScalarOp<Bf16> for Bf16Sub {
299    fn apply(lhs: Bf16, rhs: Bf16) -> Bf16 {
300        lhs - rhs
301    }
302}
303
304/// bfloat16 multiplication as an operation type.
305///
306/// # Examples
307///
308/// ```rust
309/// use tenferro_bf16_proof::{Bf16, Bf16Mul};
310/// use tenferro_cpu::BinaryScalarOp;
311///
312/// let product =
313///     <Bf16Mul as BinaryScalarOp<Bf16>>::apply(Bf16::from_f32(3.0), Bf16::from_f32(4.0));
314/// assert_eq!(product.to_f32(), 12.0);
315/// ```
316pub struct Bf16Mul;
317
318impl tenferro_cpu::BinaryScalarOp<Bf16> for Bf16Mul {
319    fn apply(lhs: Bf16, rhs: Bf16) -> Bf16 {
320        lhs * rhs
321    }
322}
323
324tenferro_tensor::define_scalar_set! {
325    /// Tag for the set that pairs the standard narrow type with `f32`.
326    ///
327    /// # Examples
328    ///
329    /// ```rust
330    /// use tenferro_bf16_proof::Bf16Tag;
331    ///
332    /// assert_ne!(Bf16Tag::F32, Bf16Tag::Bf16);
333    /// ```
334    pub enum Bf16Tag {
335        /// Standard single precision.
336        F32 => f32 : Float 0 32,
337        /// Bfloat16, ranked below single precision.
338        Bf16 => Bf16 : Float 1 16,
339    }
340    /// Value enum for the set that carries bfloat16 beside `f32`.
341    ///
342    /// # Examples
343    ///
344    /// ```rust
345    /// use tenferro_bf16_proof::{Bf16, Bf16Set, Bf16Tag};
346    /// use tenferro_tensor::{DynRank, Host, ScalarSet, TypedTensor};
347    ///
348    /// let value = Bf16Set::Bf16(TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![1], vec![Bf16::from_f32(2.0)])?);
349    /// assert_eq!(value.tag(), Bf16Tag::Bf16);
350    /// # Ok::<(), tenferro_tensor::Error>(())
351    /// ```
352    pub enum Bf16Set;
353}