Skip to main content

tenferro_bf16_proof/
reduction.rs

1//! Reductions over bfloat16 storage.
2//!
3//! #1785 asks for sum reduction with a specified accumulation rule and for tests that distinguish
4//! that rule from repeated rounding of the storage type. Two reductions do that here: one
5//! accumulates in `f32` and rounds once, which is the contract this crate promises, and one
6//! rounds at every step. The second is not offered as the contract; it exists so tests can measure
7//! the difference instead of assuming it.
8
9use crate::Bf16;
10
11/// Sum a slice, accumulating in `f32` and rounding once at the end.
12///
13/// # Examples
14///
15/// ```rust
16/// use tenferro_bf16_proof::{reduction::sum_in_f32_accumulation, Bf16};
17///
18/// let ones = vec![Bf16::from_f32(1.0); 300];
19/// // Three hundred ones are exactly representable after rounding, because the accumulation kept
20/// // the intermediate values in f32.
21/// assert_eq!(sum_in_f32_accumulation(&ones).to_f32(), 300.0);
22/// ```
23#[must_use]
24pub fn sum_in_f32_accumulation(values: &[Bf16]) -> Bf16 {
25    let total: f32 = values.iter().map(|value| value.to_f32()).sum();
26    Bf16::from_f32(total)
27}
28
29/// Sum a slice, rounding to bfloat16 after every step.
30///
31/// This is the weaker contract, and it is the one to avoid when the accumulation is promised in
32/// `f32`: on a slice of three hundred ones it stalls at `256.0`, because from there on bfloat16
33/// spacing above one is `2.0` and adding one rounds back down.
34///
35/// # Examples
36///
37/// ```rust
38/// use tenferro_bf16_proof::{reduction::sum_with_per_step_rounding, Bf16};
39///
40/// let ones = vec![Bf16::from_f32(1.0); 300];
41/// assert_eq!(sum_with_per_step_rounding(&ones).to_f32(), 256.0);
42/// ```
43#[must_use]
44pub fn sum_with_per_step_rounding(values: &[Bf16]) -> Bf16 {
45    values
46        .iter()
47        .fold(Bf16::zero(), |accumulator, value| accumulator + *value)
48}
49
50/// Product of a slice, accumulating in `f32` and rounding once.
51///
52/// # Examples
53///
54/// ```rust
55/// use tenferro_bf16_proof::{reduction::product_in_f32_accumulation, Bf16};
56///
57/// let values = [Bf16::from_f32(3.0), Bf16::from_f32(4.0)];
58/// assert_eq!(product_in_f32_accumulation(&values).to_f32(), 12.0);
59/// ```
60#[must_use]
61pub fn product_in_f32_accumulation(values: &[Bf16]) -> Bf16 {
62    let total: f32 = values.iter().map(|value| value.to_f32()).product();
63    Bf16::from_f32(total)
64}