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])
|
|||
|
|
}
|