184 lines
6.4 KiB
Go
184 lines
6.4 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package core
|
||
|
||
import (
|
||
"math"
|
||
"math/bits"
|
||
)
|
||
|
||
// IEEE 754 binary16 (half precision) as a fifth dtype. The payload is a
|
||
// []uint16 holding the half bit patterns, exactly as the design note
|
||
// prescribes: a payload per dtype, and kernels that widen on read. The
|
||
// half dtype sits on the promotion ladder one step below float32 and
|
||
// behaves like it: the arithmetic runs in float64 (where every half
|
||
// value is exact) and narrows once per element.
|
||
//
|
||
// Conversion contract, for both directions:
|
||
// - widening (HalfToFloat64) is exact: every finite half value,
|
||
// subnormals included, is a float64 value; infinities widen to
|
||
// infinities; a half NaN widens to a float64 NaN carrying the same
|
||
// sign with its ten payload bits at the top of the mantissa.
|
||
// - narrowing (HalfFromFloat64) rounds to nearest, ties to even
|
||
// (IEEE 754 round-to-nearest-even), with gradual underflow through
|
||
// the half subnormals. |x| >= 65520 overflows to the signed
|
||
// infinity (65520 is the midpoint between the largest finite half
|
||
// 65504 and the next binade at 65536, and the tie rounds away from
|
||
// the odd mantissa of 65504); anything smaller rounds to a finite
|
||
// half. A float64 NaN, quiet or signalling, narrows to the
|
||
// canonical half NaN 0x7E00 with its sign preserved: the payload is
|
||
// deliberately not carried across, which keeps the conversion a
|
||
// pure function of the value and makes the constructor's NaN
|
||
// behaviour trivial to reason about. Signed zero is preserved.
|
||
//
|
||
// Constructors built on the narrowing (FromFloat16s, FullF16, SetFloatAt
|
||
// on a half array) inherit exactly this contract.
|
||
|
||
// Half bit patterns the package writes directly.
|
||
const (
|
||
// halfOne is the bit pattern of 1.0.
|
||
halfOne uint16 = 0x3C00
|
||
// halfNaN is the canonical quiet NaN.
|
||
halfNaN uint16 = 0x7E00
|
||
// halfInf is +Inf; OR-ing the sign bit turns it into -Inf.
|
||
halfInf uint16 = 0x7C00
|
||
)
|
||
|
||
// HalfToFloat64 widens the half bit pattern h to float64 exactly: no
|
||
// rounding, no range loss, NaN payloads preserved in the top mantissa
|
||
// bits.
|
||
func HalfToFloat64(h uint16) float64 {
|
||
sign := uint64(h&0x8000) << 48
|
||
exp := uint64(h>>10) & 0x1F
|
||
frac := uint64(h & 0x03FF)
|
||
switch {
|
||
case exp == 0x1F:
|
||
if frac == 0 {
|
||
return math.Float64frombits(sign | 0x7FF0_0000_0000_0000)
|
||
}
|
||
// NaN: the ten payload bits land above the float64 mantissa's
|
||
// low 42 bits, so the quiet bit travels with the payload.
|
||
return math.Float64frombits(sign | 0x7FF0_0000_0000_0000 | frac<<42)
|
||
case exp == 0:
|
||
if frac == 0 {
|
||
return math.Float64frombits(sign) // ±0
|
||
}
|
||
// Subnormal: the value is frac × 2^-24, renormalised into the
|
||
// float64 field. k is the zero-based position of frac's leading
|
||
// bit, so the value is 2^(k-24) × (1 + r/2^k) with r below.
|
||
k := uint(bits.Len64(frac) - 1)
|
||
mant := (frac - 1<<k) << (52 - k)
|
||
return math.Float64frombits(sign | uint64(1023-24+k)<<52 | mant)
|
||
default:
|
||
return math.Float64frombits(sign | (exp-15+1023)<<52 | frac<<42)
|
||
}
|
||
}
|
||
|
||
// HalfFromFloat64 narrows f to the half bit pattern with IEEE 754
|
||
// round-to-nearest-even: gradual underflow through the subnormals,
|
||
// overflow to the signed infinity at |f| >= 65520, NaN canonicalised to
|
||
// 0x7E00 (sign preserved), signed zero preserved.
|
||
func HalfFromFloat64(f float64) uint16 {
|
||
b := math.Float64bits(f)
|
||
sign := uint16(b>>63) << 15
|
||
exp := int64(b>>52) & 0x7FF
|
||
frac := b & (1<<52 - 1)
|
||
if exp == 0x7FF {
|
||
if frac != 0 {
|
||
return sign | halfNaN
|
||
}
|
||
return sign | halfInf
|
||
}
|
||
// The half exponent field the value would carry if it fit: float64
|
||
// bias 1023 against half bias 15.
|
||
e := exp - 1008
|
||
if e >= 0x1F {
|
||
// Above every finite half: ±Inf.
|
||
return sign | halfInf
|
||
}
|
||
if e >= 1 {
|
||
// A normal half: round the 52-bit mantissa to 10 bits, ties to
|
||
// even, then carry a mantissa overflow into the exponent (which
|
||
// may itself overflow into the infinity boundary).
|
||
hf := frac >> 42
|
||
rem := frac & (1<<42 - 1)
|
||
if rem > 1<<41 || (rem == 1<<41 && hf&1 == 1) {
|
||
hf++
|
||
}
|
||
if hf == 1<<10 {
|
||
hf = 0
|
||
e++
|
||
if e >= 0x1F {
|
||
return sign | halfInf
|
||
}
|
||
}
|
||
return sign | uint16(e)<<10 | uint16(hf)
|
||
}
|
||
// Subnormal (or zero): the half mantissa m counts 2^-24 units. The
|
||
// value is num × 2^(exp-1075) with num the mantissa plus its
|
||
// implicit leading bit, so m is num shifted right by 1051-exp, the
|
||
// guard and sticky bits deciding the tie. A shift of 53 is the
|
||
// smallest that can still round up: f = 2^-25 shifts by exactly 53
|
||
// and ties to the even zero. Anything further is below the tie and
|
||
// drops out; float64 subnormals (exp 0) never reach even that.
|
||
sh := uint(1051 - exp)
|
||
if sh > 54 {
|
||
return sign
|
||
}
|
||
num := 1<<52 | frac
|
||
m := num >> sh
|
||
rem := num & (1<<sh - 1)
|
||
halfPoint := uint64(1) << (sh - 1)
|
||
if rem > halfPoint || (rem == halfPoint && m&1 == 1) {
|
||
m++
|
||
}
|
||
if m == 1<<10 {
|
||
// The subnormal rounding crossed into the smallest normal.
|
||
return sign | 1<<10
|
||
}
|
||
return sign | uint16(m)
|
||
}
|
||
|
||
// FromFloat16s builds a float16 array of the given shape from vals,
|
||
// copying them and narrowing each value to the nearest half
|
||
// (round-to-nearest-even, overflow to ±Inf, NaN canonicalised to
|
||
// 0x7E00, signed zero preserved: the HalfFromFloat64 contract).
|
||
func FromFloat16s(vals []float64, shape ...int) (*Array, error) {
|
||
sh, err := shapeFor(shape, len(vals))
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
halves := make([]uint16, len(vals))
|
||
for i, v := range vals {
|
||
halves[i] = HalfFromFloat64(v)
|
||
}
|
||
return &Array{shape: sh, dt: Float16, halves: halves}, nil
|
||
}
|
||
|
||
// HalvesFromArray builds a float16 array that takes ownership of the
|
||
// raw half bit patterns in vals: no copy is made and no value is
|
||
// converted, so the caller must not touch the slice afterwards. The
|
||
// value count must fill the shape exactly. This is the bits-taking
|
||
// route for callers that already hold IEEE 754 binary16 data.
|
||
func HalvesFromArray(vals []uint16, shape ...int) (*Array, error) {
|
||
sh, err := shapeFor(shape, len(vals))
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
return &Array{shape: sh, dt: Float16, halves: vals}, nil
|
||
}
|
||
|
||
// RawHalves returns the float16 payload (the raw IEEE 754 binary16 bit
|
||
// patterns) with the RawFloats contract.
|
||
func (a *Array) RawHalves() []uint16 { return a.halves }
|
||
|
||
// halfAt returns element i widened from the half payload, the float16
|
||
// twin of floatAt.
|
||
func (a *Array) halfAt(i int) float64 {
|
||
if a.strides != nil {
|
||
i = a.physIndex(i)
|
||
}
|
||
return HalfToFloat64(a.halves[i])
|
||
}
|