Files
tensor/internal/core/float16.go
T

184 lines
6.4 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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])
}