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}