tenferro_bf16_proof/conversion.rs
1//! Directed conversions between bfloat16 storage and `f32`.
2//!
3//! #1785 asks for explicit bf16/f32 conversions in both directions with rounding and range
4//! behaviour that is stated rather than implied. Widening is exact, because every bfloat16 value
5//! is an `f32`; narrowing rounds to the nearest bfloat16, with ties to even, which is what
6//! [`half::bf16::from_f32`] does. Neither direction is a promotion rule, and neither is implicit.
7
8use crate::Bf16;
9use tenferro_tensor::{DynRank, Host, TypedTensor};
10
11/// Widen every stored value to `f32`, which is exact.
12///
13/// # Errors
14///
15/// Returns a validation error carrying
16/// [`tenferro_tensor_core::ValidationError::ShapeDataLengthMismatch`] when the
17/// output shape cannot be built, which cannot happen for a shape that already
18/// exists.
19///
20/// # Examples
21///
22/// ```rust
23/// use tenferro_bf16_proof::{conversion::widen, Bf16};
24/// use tenferro_tensor::{DynRank, Host, TypedTensor};
25///
26/// let source = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![2], vec![Bf16::from_f32(1.0), Bf16::from_f32(2.5)])?;
27/// let widened = widen(&source)?;
28/// assert_eq!(widened.as_slice(), &[1.0_f32, 2.5]);
29/// # Ok::<(), tenferro_tensor::Error>(())
30/// ```
31pub fn widen(
32 source: &TypedTensor<Bf16, DynRank, Host>,
33) -> Result<TypedTensor<f32, DynRank, Host>, tenferro_tensor::Error> {
34 TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(
35 source.shape().to_vec(),
36 source
37 .as_slice()
38 .iter()
39 .map(|value| value.to_f32())
40 .collect(),
41 )
42}
43
44/// Round every value to the nearest bfloat16, with ties to even.
45///
46/// # Errors
47///
48/// Returns a validation error carrying
49/// [`tenferro_tensor_core::ValidationError::ShapeDataLengthMismatch`] when the
50/// output shape cannot be built, which cannot happen for a shape that already
51/// exists.
52///
53/// # Examples
54///
55/// ```rust
56/// use tenferro_bf16_proof::conversion::narrow;
57/// use tenferro_tensor::{DynRank, Host, TypedTensor};
58///
59/// // 1.00390625 rounds down to 1.0: bfloat16 keeps eight bits of significand.
60/// let source = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![1], vec![1.00390625_f32])?;
61/// assert_eq!(narrow(&source)?.as_slice()[0].to_f32(), 1.0);
62/// # Ok::<(), tenferro_tensor::Error>(())
63/// ```
64pub fn narrow(
65 source: &TypedTensor<f32, DynRank, Host>,
66) -> Result<TypedTensor<Bf16, DynRank, Host>, tenferro_tensor::Error> {
67 TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(
68 source.shape().to_vec(),
69 source
70 .as_slice()
71 .iter()
72 .map(|value| Bf16::from_f32(*value))
73 .collect(),
74 )
75}