Files
tensor/internal/core/float16_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

979 lines
32 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"math/bits"
"strings"
"testing"
)
func mustFromFloat16s(t *testing.T, vals []float64, shape ...int) *Array {
t.Helper()
a, err := FromFloat16s(vals, shape...)
if err != nil {
t.Fatalf("FromFloat16s(%v, %v): %v", vals, shape, err)
}
return a
}
func mustFromHalves(t *testing.T, halves []uint16, shape ...int) *Array {
t.Helper()
a, err := HalvesFromArray(append([]uint16(nil), halves...), shape...)
if err != nil {
t.Fatalf("HalvesFromArray(%v, %v): %v", halves, shape, err)
}
return a
}
// TestHalfConversionTable pins the exact bit patterns of the narrowing
// under IEEE 754 round-to-nearest-even. Each expectation is derived from
// the standard: the half format has a 5-bit exponent (bias 15) and a
// 10-bit fraction, the largest finite half is 65504 (0x7BFF), the
// overflow boundary for RNE is |x| = 65520 (the midpoint between 65504
// and the next binade at 65536; the tie rounds away from 65504's odd
// mantissa), and the smallest subnormal is 2^-24 (0x0001).
func TestHalfConversionTable(t *testing.T) {
cases := []struct {
in float64
want uint16
why string
}{
{0, 0x0000, "zero"},
{math.Copysign(0, -1), 0x8000, "negative zero is preserved"},
{0.5, 0x3800, "exponent field 14, fraction 0"},
{-0.5, 0xB800, "sign preserved"},
{1, 0x3C00, "exponent field 15, fraction 0"},
{2, 0x4000, "exponent field 16"},
{-2, 0xC000, "sign preserved"},
{0.1, 0x2E66, "1.6 × 2^-4: fraction 0.6 × 1024 = 614.4 rounds down"},
{1 + math.Ldexp(1, -10), 0x3C01, "one ulp above 1"},
{1 + math.Ldexp(1, -11), 0x3C00, "exact tie between 1 and 1+2^-10: even mantissa 0 wins"},
{2048, 0x6800, "exponent field 23"},
{2049, 0x6800, "exact tie between 2048 and 2050: even mantissa 0 wins"},
{2049.9, 0x6801, "closer to 2050"},
{math.Ldexp(1, -14), 0x0400, "the smallest normal"},
{math.Ldexp(1, -15), 0x0200, "subnormal with fraction field 512"},
{math.Ldexp(1, -24), 0x0001, "the smallest subnormal"},
{math.Ldexp(1, -25), 0x0000, "halfway between 0 and 2^-24: ties to the even zero"},
{1.5 * math.Ldexp(1, -24), 0x0002, "halfway between 2^-24 and 2^-23: even mantissa 2 wins"},
{0.9999 * math.Ldexp(1, -14), 0x0400, "subnormal rounding crosses into the smallest normal"},
{65504, 0x7BFF, "the largest finite half"},
{65512, 0x7BFF, "below the overflow midpoint"},
{65519.999, 0x7BFF, "just below the midpoint still rounds to 65504"},
{65520, 0x7C00, "the midpoint itself: the tie leaves the odd mantissa 1023 for infinity"},
{65521, 0x7C00, "beyond the midpoint"},
{-65504, 0xFBFF, "sign preserved at the maximum"},
{-65520, 0xFC00, "negative overflow to -Inf"},
{math.Inf(1), 0x7C00, "+Inf"},
{math.Inf(-1), 0xFC00, "-Inf"},
{math.NaN(), 0x7E00, "every NaN narrows to the canonical quiet NaN"},
}
for _, c := range cases {
if got := HalfFromFloat64(c.in); got != c.want {
t.Errorf("HalfFromFloat64(%v) = 0x%04X, want 0x%04X (%s)", c.in, got, c.want, c.why)
}
}
}
// TestHalfWidening checks the exact widening: every half value is a
// float32 value too, so the classic half to float32 to float64 route is
// an independent oracle for HalfToFloat64 across all 65536 patterns,
// NaNs included.
func TestHalfWidening(t *testing.T) {
for b := range 65536 {
h := uint16(b)
sign := uint32(h&0x8000) << 16
e := uint32(h>>10) & 0x1F
fr := uint32(h & 0x3FF)
var b32 uint32
switch {
case e == 0:
if fr == 0 {
b32 = sign
break
}
// frac × 2^-24 renormalised: the leading bit at position k
// makes the float32 exponent k-24.
k := uint(bits.Len32(fr) - 1)
b32 = sign | uint32(127-24+k)<<23 | (fr-1<<k)<<(23-k)
case e == 0x1F:
b32 = sign | 0xFF<<23 | fr<<13
default:
b32 = sign | (e-15+127)<<23 | fr<<13
}
want := float64(math.Float32frombits(b32))
got := HalfToFloat64(h)
if want != got && !(math.IsNaN(want) && math.IsNaN(got)) {
t.Fatalf("HalfToFloat64(0x%04X) = %v, want %v", h, got, want)
}
}
}
// TestHalfRoundTrip pins the round-trip contract: widening then
// renarrowing is the identity for every finite and infinite half, and
// every NaN pattern collapses onto the sign-preserving canonical quiet
// NaN, the documented narrowing policy.
func TestHalfRoundTrip(t *testing.T) {
for b := range 65536 {
h := uint16(b)
w := HalfToFloat64(h)
got := HalfFromFloat64(w)
want := h
if math.IsNaN(w) {
want = (h & 0x8000) | 0x7E00
}
if got != want {
t.Fatalf("round trip 0x%04X: widen gives %v, renarrow gives 0x%04X, want 0x%04X", h, w, got, want)
}
}
// A float64 signalling NaN canonicalises the same way.
snan := math.Float64frombits(0x7FF0000000000001)
if got := HalfFromFloat64(snan); got != 0x7E00 {
t.Fatalf("HalfFromFloat64(sNaN) = 0x%04X, want 0x7E00", got)
}
}
func TestFloat16Constructors(t *testing.T) {
a := mustFromFloat16s(t, []float64{1, 0.5, -2}, 3)
if a.Dtype() != Float16 || a.Len() != 3 {
t.Fatalf("float16 array: %s len %d", a.Dtype(), a.Len())
}
// RawHalves holds the bit patterns; FloatAt widens exactly.
want := []uint16{0x3C00, 0x3800, 0xC000}
if got := a.RawHalves(); !slicesEqualU16(got, want) {
t.Fatalf("RawHalves: %v, want %v", got, want)
}
if a.FloatAt(0) != 1 || a.FloatAt(1) != 0.5 || a.FloatAt(2) != -2 {
t.Fatalf("FloatAt: %v %v %v", a.FloatAt(0), a.FloatAt(1), a.FloatAt(2))
}
// SetFloatAt narrows under the same contract.
a.SetFloatAt(0, 0.25)
if got := a.RawHalves()[0]; got != 0x3400 {
t.Fatalf("SetFloatAt(0.25) = 0x%04X, want 0x3400", got)
}
// Overflow to infinity, signed zero preserved, NaN canonicalised.
b := mustFromFloat16s(t, []float64{65520, math.Copysign(0, -1), math.NaN()}, 3)
want = []uint16{0x7C00, 0x8000, 0x7E00}
if got := b.RawHalves(); !slicesEqualU16(got, want) {
t.Fatalf("narrowing contract: %v, want %v", got, want)
}
// Bits-taking route: no conversion happens.
h := mustFromHalves(t, []uint16{0x0001, 0x7BFF, 0xFC00}, 3)
if got := h.RawHalves(); !slicesEqualU16(got, []uint16{0x0001, 0x7BFF, 0xFC00}) {
t.Fatalf("HalvesFromArray: %v", got)
}
if h.FloatAt(0) != math.Ldexp(1, -24) || h.FloatAt(1) != 65504 || !math.IsInf(h.FloatAt(2), -1) {
t.Fatal("HalvesFromArray widened wrong values")
}
// Fill family.
z, _ := Zeros(Float16, 2)
if z.Dtype() != Float16 || z.FloatAt(1) != 0 {
t.Fatalf("Zeros float16: %s", z)
}
o, _ := Ones(Float16, 2)
if got := o.RawHalves()[1]; got != 0x3C00 {
t.Fatalf("Ones float16: 0x%04X", got)
}
f, _ := FullF16(0.5, 2)
if got := f.RawHalves()[0]; got != 0x3800 {
t.Fatalf("FullF16(0.5): 0x%04X", got)
}
inf, _ := FullF16(1e300, 1)
if got := inf.RawHalves()[0]; got != 0x7C00 {
t.Fatalf("FullF16 overflow: 0x%04X", got)
}
// An unknown dtype is still refused by name.
if _, err := Zeros(Dtype(99), 1); err == nil || !strings.Contains(err.Error(), "unknown element type") {
t.Fatalf("Zeros(Dtype(99)): %v", err)
}
}
func slicesEqualU16(a, b []uint16) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
// TestFloat16PromotionMatrix walks the ladder: int < float16 < float32
// < float < complex, with max(a, b) deciding.
func TestFloat16PromotionMatrix(t *testing.T) {
h := mustFromFloat16s(t, []float64{1.5}, 1)
i := mustFromInts(t, []int64{2}, 1)
f32 := mustFromFloat32s(t, []float32{4}, 1)
f64 := mustFromFloats(t, []float64{8}, 1)
c := mustFromComplexes(t, []complex128{1}, 1)
pairs := []struct {
a, b *Array
want Dtype
val float64
}{
{i, h, Float16, 3.5},
{h, i, Float16, 3.5},
{h, f32, Float32, 5.5},
{f32, h, Float32, 5.5},
{h, f64, Float, 9.5},
{f64, h, Float, 9.5},
{h, c, Complex, 2.5},
{c, h, Complex, 2.5},
}
for _, p := range pairs {
got, err := Add(p.a, p.b)
if err != nil {
t.Fatalf("Add(%s, %s): %v", p.a.Dtype(), p.b.Dtype(), err)
}
if got.Dtype() != p.want {
t.Fatalf("Add(%s, %s) dtype = %s, want %s", p.a.Dtype(), p.b.Dtype(), got.Dtype(), p.want)
}
if got.Dtype() == Complex {
continue // checked through ComplexAt below
}
if v := got.FloatAt(0); v != p.val {
t.Fatalf("Add(%s, %s)[0] = %v, want %v", p.a.Dtype(), p.b.Dtype(), v, p.val)
}
}
// The complex pair checks through ComplexAt.
cg, err := Add(h, c)
if err != nil {
t.Fatal(err)
}
if got := cg.ComplexAt(0); got != complex(2.5, 0) {
t.Fatalf("Add(float16, complex) = %v", got)
}
// Half arithmetic narrows once: the sum runs in float64 over the
// exact widenings (half 0.1 = 0.0999755859375, half 0.2 =
// 0.199951171875) and the result's fraction 204.5 ties to the even
// mantissa 204.
x := mustFromFloat16s(t, []float64{0.1}, 1)
y := mustFromFloat16s(t, []float64{0.2}, 1)
s, err := Add(x, y)
if err != nil {
t.Fatal(err)
}
want := uint16(0x34CC)
if got := s.RawHalves()[0]; got != want {
t.Fatalf("0.1+0.2 in half = 0x%04X, want 0x%04X", got, want)
}
// Div is true division and stays half.
q, err := Div(i, h)
if err != nil {
t.Fatal(err)
}
if q.Dtype() != Float16 || q.FloatAt(0) != HalfToFloat64(HalfFromFloat64(2.0/1.5)) {
t.Fatalf("Div(int, float16) = %s %v", q.Dtype(), q.FloatAt(0))
}
// Scalar maps keep the half dtype.
if got := MulF(h, 2); got.Dtype() != Float16 || got.FloatAt(0) != 3 {
t.Fatalf("MulF on float16: %s %v", got.Dtype(), got.FloatAt(0))
}
if got := AddI(h, 1); got.Dtype() != Float16 || got.FloatAt(0) != 2.5 {
t.Fatalf("AddI on float16: %s %v", got.Dtype(), got.FloatAt(0))
}
// Comparisons read through the exact widening and answer bool.
eq, err := Eq(h, mustFromFloat16s(t, []float64{1.5}, 1))
if err != nil {
t.Fatal(err)
}
if !eq.bools[0] {
t.Fatal("Eq of equal halves must answer true")
}
lt, err := Lt(i, h)
if err != nil {
t.Fatal(err)
}
if lt.bools[0] {
t.Fatal("Lt(int, float16) answered wrong")
}
gt, err := Gt(i, h)
if err != nil {
t.Fatal(err)
}
if !gt.bools[0] {
t.Fatal("Gt(int, float16) answered wrong")
}
// Where picks through the promotion and writes half bits.
w, err := Where(mustFromInts(t, []int64{1}, 1), h, i)
if err != nil {
t.Fatal(err)
}
if w.Dtype() != Float16 || w.RawHalves()[0] != 0x3E00 {
t.Fatalf("Where dtype %s bits %04X", w.Dtype(), w.RawHalves()[0])
}
// Minimum and Maximum stay half.
m, err := Minimum(h, mustFromFloat16s(t, []float64{2}, 1))
if err != nil {
t.Fatal(err)
}
if m.Dtype() != Float16 || m.FloatAt(0) != 1.5 {
t.Fatalf("Minimum on halves: %s %v", m.Dtype(), m.FloatAt(0))
}
// Quo still refuses non-int operands loudly.
if _, err := Quo(h, h); err == nil || !strings.Contains(err.Error(), "needs int arrays") {
t.Fatalf("Quo on float16: %v", err)
}
}
// TestFloat16Reductions mirrors the float32 reduction semantics: the
// values here are exactly representable in half, so the answers are
// exact too.
func TestFloat16Reductions(t *testing.T) {
vals := []float64{1, 2, 0.5, -0.25, 4}
h := mustFromFloat16s(t, vals, 5)
f32 := mustFromFloat32s(t, []float32{1, 2, 0.5, -0.25, 4}, 5)
// Sum answers a float scalar for both, identically.
if got, want := Sum(h).Float(), Sum(f32).Float(); got != want {
t.Fatalf("Sum float16 %v != float32 %v", got, want)
}
if Sum(h).IsFloat() != true {
t.Fatal("Sum of float16 must answer a float scalar")
}
// Mean.
mean, err := Mean(h)
if err != nil {
t.Fatal(err)
}
if mean != 7.25/5 {
t.Fatalf("Mean float16 = %v", mean)
}
// Min, Max and NaN skipping.
if mn, _ := Min(h); mn.Float() != -0.25 {
t.Fatalf("Min = %v", mn.Float())
}
if mx, _ := Max(h); mx.Float() != 4 {
t.Fatalf("Max = %v", mx.Float())
}
nan := mustFromHalves(t, []uint16{0x7E00, 0x3C00, 0x4000}, 3)
if mx, _ := Max(nan); mx.Float() != 2 {
t.Fatalf("Max with NaN = %v, want 2", mx.Float())
}
// ArgMax and ArgMin.
if got, err := ArgMax(h); err != nil || got != 4 {
t.Fatalf("ArgMax = %d, %v", got, err)
}
if got, err := ArgMin(h); err != nil || got != 3 {
t.Fatalf("ArgMin = %d, %v", got, err)
}
// Dot.
d, err := Dot(h, h)
if err != nil {
t.Fatal(err)
}
if !d.IsFloat() || d.Float() != 1+4+0.25+0.0625+16 {
t.Fatalf("Dot = %v", d.Float())
}
// Sort and ArgSort, NaN to the end.
sm := mustFromHalves(t, []uint16{0x4000, 0x7E00, 0x3C00, 0x0000}, 4)
sorted, err := Sort(sm)
if err != nil {
t.Fatal(err)
}
if sorted.Dtype() != Float16 {
t.Fatalf("Sort dtype %s", sorted.Dtype())
}
wantHalves := []uint16{0x0000, 0x3C00, 0x4000, 0x7E00}
if got := sorted.RawHalves(); !slicesEqualU16(got, wantHalves) {
t.Fatalf("Sort halves = %v, want %v", got, wantHalves)
}
perm, err := ArgSort(sm)
if err != nil {
t.Fatal(err)
}
wantPerm := []int64{3, 2, 0, 1}
for i, w := range wantPerm {
if perm.ints[i] != w {
t.Fatalf("ArgSort[%d] = %d, want %d", i, perm.ints[i], w)
}
}
// Axis reductions: fold in float64, narrow once.
m2 := mustFromFloat16s(t, []float64{1, 2, 3, 4}, 2, 2)
sa, err := SumAxis(m2, 1)
if err != nil {
t.Fatal(err)
}
if sa.Dtype() != Float16 || sa.FloatAt(0) != 3 || sa.FloatAt(1) != 7 {
t.Fatalf("SumAxis = %s %v %v", sa.Dtype(), sa.FloatAt(0), sa.FloatAt(1))
}
// MinAxis reduces along dim 0: per-column minima of [[1, 2], [3, 4]].
ma, err := MinAxis(m2, 0)
if err != nil {
t.Fatal(err)
}
if ma.Dtype() != Float16 || ma.FloatAt(0) != 1 || ma.FloatAt(1) != 2 {
t.Fatalf("MinAxis = %v %v", ma.FloatAt(0), ma.FloatAt(1))
}
mea, err := MeanAxis(m2, 0)
if err != nil {
t.Fatal(err)
}
if mea.Dtype() != Float || mea.FloatAt(0) != 2 {
t.Fatalf("MeanAxis = %s %v", mea.Dtype(), mea.FloatAt(0))
}
// Scans and products keep the half dtype.
cs, err := CumSum(m2, 1)
if err != nil {
t.Fatal(err)
}
if cs.Dtype() != Float16 || cs.FloatAt(0) != 1 || cs.FloatAt(1) != 3 || cs.FloatAt(2) != 3 || cs.FloatAt(3) != 7 {
t.Fatalf("CumSum = %v", cs.RawHalves())
}
pr, err := Prod(m2, 0, false)
if err != nil {
t.Fatal(err)
}
if pr.Dtype() != Float16 || pr.FloatAt(0) != 3 || pr.FloatAt(1) != 8 {
t.Fatalf("Prod = %v %v", pr.FloatAt(0), pr.FloatAt(1))
}
// Norm is always float.
nrm, err := Norm(m2, 2, 1, false)
if err != nil {
t.Fatal(err)
}
if nrm.Dtype() != Float || math.Abs(nrm.FloatAt(0)-math.Sqrt(5)) > 1e-15 {
t.Fatalf("Norm = %s %v", nrm.Dtype(), nrm.FloatAt(0))
}
// TopK keeps half and skips NaN.
tv, ti, err := TopK(nan, 2, 0)
if err != nil {
t.Fatal(err)
}
if tv.Dtype() != Float16 || tv.FloatAt(0) != 2 || tv.FloatAt(1) != 1 {
t.Fatalf("TopK values = %v %v", tv.FloatAt(0), tv.FloatAt(1))
}
if ti.ints[0] != 2 || ti.ints[1] != 1 {
t.Fatalf("TopK indices = %v %v", ti.ints[0], ti.ints[1])
}
}
func TestFloat16MathAndShape(t *testing.T) {
h := mustFromFloat16s(t, []float64{4, 0.25}, 2)
// Element-wise math keeps the half dtype and narrows once.
sq, err := Sqrt(h)
if err != nil {
t.Fatal(err)
}
if sq.Dtype() != Float16 || sq.RawHalves()[0] != 0x4000 || sq.RawHalves()[1] != 0x3800 {
t.Fatalf("Sqrt = %v", sq.RawHalves())
}
ex, err := Exp(h)
if err != nil {
t.Fatal(err)
}
if ex.Dtype() != Float16 || ex.RawHalves()[0] != HalfFromFloat64(math.Exp(4)) {
t.Fatalf("Exp = 0x%04X", ex.RawHalves()[0])
}
fl, err := Floor(mustFromFloat16s(t, []float64{1.5}, 1))
if err != nil {
t.Fatal(err)
}
if fl.Dtype() != Float16 || fl.FloatAt(0) != 1 {
t.Fatalf("Floor = %s %v", fl.Dtype(), fl.FloatAt(0))
}
pi, err := PowI(h, 2)
if err != nil {
t.Fatal(err)
}
if pi.Dtype() != Float16 || pi.FloatAt(0) != 16 {
t.Fatalf("PowI = %s %v", pi.Dtype(), pi.FloatAt(0))
}
ab := Abs(mustFromFloat16s(t, []float64{-2}, 1))
if ab.Dtype() != Float16 || ab.FloatAt(0) != 2 {
t.Fatalf("Abs = %s %v", ab.Dtype(), ab.FloatAt(0))
}
// Clipping keeps half, in float64 order.
cl, err := ClipF(h, 0.5, 3)
if err != nil {
t.Fatal(err)
}
if cl.Dtype() != Float16 || cl.FloatAt(0) != 3 || cl.FloatAt(1) != 0.5 {
t.Fatalf("ClipF = %v %v", cl.FloatAt(0), cl.FloatAt(1))
}
cli, err := ClipI(h, 1, 3)
if err != nil {
t.Fatal(err)
}
if cli.Dtype() != Float16 || cli.FloatAt(0) != 3 || cli.FloatAt(1) != 1 {
t.Fatalf("ClipI = %v %v", cli.FloatAt(0), cli.FloatAt(1))
}
// Shape and layout operations are dtype agnostic.
r, err := Reshape(h, 1, 2)
if err != nil {
t.Fatal(err)
}
if r.Dtype() != Float16 || r.RawHalves()[1] != 0x3400 {
t.Fatalf("Reshape = %v", r.RawHalves())
}
if tr := Transpose(r); tr.NDim() != 2 || tr.RawHalves()[0] != 0x4400 || tr.RawHalves()[1] != 0x3400 {
t.Fatalf("Transpose = %v", tr.RawHalves())
}
row, err := Row(r, 0)
if err != nil {
t.Fatal(err)
}
if !slicesEqualU16(row.RawHalves(), []uint16{0x4400, 0x3400}) {
t.Fatalf("Row = %v", row.RawHalves())
}
col, err := Col(r, 1)
if err != nil {
t.Fatal(err)
}
if !slicesEqualU16(col.RawHalves(), []uint16{0x3400}) {
t.Fatalf("Col = %v", col.RawHalves())
}
id, err := Identity(Float16, 2)
if err != nil {
t.Fatal(err)
}
if id.FloatAt(0) != 1 || id.FloatAt(1) != 0 || id.FloatAt(3) != 1 {
t.Fatalf("Identity = %s", id)
}
zl := ZerosLike(h)
if zl.Dtype() != Float16 || zl.FloatAt(0) != 0 {
t.Fatal("ZerosLike must keep float16")
}
// Concat promotes through the ladder.
c, err := Concat(h, mustFromFloat32s(t, []float32{1}, 1), 0)
if err != nil {
t.Fatal(err)
}
if c.Dtype() != Float32 {
t.Fatalf("Concat(float16, float32) dtype = %s", c.Dtype())
}
// Diff widens to float, like every non-int, non-complex dtype.
d, err := Diff(h, 1, 0)
if err != nil {
t.Fatal(err)
}
if d.Dtype() != Float || d.FloatAt(0) != -3.75 {
t.Fatalf("Diff = %s %v", d.Dtype(), d.FloatAt(0))
}
}
// TestFloat16AstypeAndAccess walks the conversion and accessor routes.
func TestFloat16AstypeAndAccess(t *testing.T) {
h := mustFromFloat16s(t, []float64{1.5, -0.5, 65504}, 3)
// To Int truncates like float32 to Int does.
i, err := Astype(h, Int)
if err != nil {
t.Fatal(err)
}
if i.ints[0] != 1 || i.ints[1] != 0 || i.ints[2] != 65504 {
t.Fatalf("Astype to int = %v", i.RawInts())
}
// Up the ladder is exact.
f32, err := Astype(h, Float32)
if err != nil {
t.Fatal(err)
}
if f32.Dtype() != Float32 || f32.FloatAt(2) != 65504 {
t.Fatalf("Astype to float32 = %s", f32)
}
f64, err := Astype(h, Float)
if err != nil {
t.Fatal(err)
}
if f64.FloatAt(0) != 1.5 {
t.Fatalf("Astype to float = %v", f64.FloatAt(0))
}
// Down from float narrows under the RNE contract.
down, err := Astype(mustFromFloats(t, []float64{0.1, 1e10}, 2), Float16)
if err != nil {
t.Fatal(err)
}
if got := down.RawHalves(); !slicesEqualU16(got, []uint16{0x2E66, 0x7C00}) {
t.Fatalf("Astype float to float16 = %v", got)
}
// Same dtype copies, views included.
sl, err := Slice(h, 0, 1, 3)
if err != nil {
t.Fatal(err)
}
cp, err := Astype(sl, Float16)
if err != nil {
t.Fatal(err)
}
if !slicesEqualU16(cp.RawHalves(), []uint16{0xB800, 0x7BFF}) {
t.Fatalf("Astype same dtype on a view = %v", cp.RawHalves())
}
// Complex to float16 is a loud refusal, like complex to float32.
if _, err := Astype(mustFromComplexes(t, []complex128{1}, 1), Float16); err == nil ||
!strings.Contains(err.Error(), "cannot narrow complex to float16") {
t.Fatalf("Astype complex to float16: %v", err)
}
// WithFloat refuses a float16 array by name, as it refuses float32.
if _, err := WithFloat(h, 1, 0); err == nil || !strings.Contains(err.Error(), "float16") {
t.Fatalf("WithFloat on float16: %v", err)
}
// Scatter refuses a complex source into a half array (canStore).
dst := mustFromHalves(t, []uint16{0, 0}, 2)
_, err = Scatter(dst, 0, mustFromInts(t, []int64{0}, 1), mustFromComplexes(t, []complex128{1i}, 1))
if err == nil || !strings.Contains(err.Error(), "cannot store") {
t.Fatalf("Scatter complex into float16: %v", err)
}
// Pad's constant mode narrows the fill value.
p, err := Pad(h, []int{0, 1}, "constant", 0.5)
if err != nil {
t.Fatal(err)
}
if got := p.RawHalves(); !slicesEqualU16(got, []uint16{0x3E00, 0xB800, 0x7BFF, 0x3800}) {
t.Fatalf("Pad = %v", got)
}
}
// TestFloat16AxisAndConversions pins the remaining half paths: the
// strided axis folds, the axis arg-extremes, the sort-based TopK path,
// the scalar maps and the Interpolate2D half grid.
func TestFloat16AxisAndConversions(t *testing.T) {
m2 := mustFromFloat16s(t, []float64{1, 2, 3, 4}, 2, 2)
// SumAxis along dim 0 streams the strided half fold.
sa, err := SumAxis(m2, 0)
if err != nil {
t.Fatal(err)
}
if sa.Dtype() != Float16 || sa.FloatAt(0) != 4 || sa.FloatAt(1) != 6 {
t.Fatalf("SumAxis dim 0 = %v %v", sa.FloatAt(0), sa.FloatAt(1))
}
// MaxAxis with a NaN in the line: NaN never wins.
nan := mustFromHalves(t, []uint16{0x3C00, 0x7E00, 0x4400, 0x4200}, 2, 2)
ma, err := MaxAxis(nan, 0)
if err != nil {
t.Fatal(err)
}
if ma.Dtype() != Float16 || ma.FloatAt(0) != 4 || ma.FloatAt(1) != 3 {
t.Fatalf("MaxAxis with NaN = %v %v", ma.FloatAt(0), ma.FloatAt(1))
}
// The axis arg-extremes walk the half payload widened.
am, err := ArgMaxAxis(m2, 0)
if err != nil {
t.Fatal(err)
}
if !slicesEqualI64(am.RawInts(), []int64{1, 1}) {
t.Fatalf("ArgMaxAxis = %v", am.RawInts())
}
an, err := ArgMinAxis(m2, 0)
if err != nil {
t.Fatal(err)
}
if !slicesEqualI64(an.RawInts(), []int64{0, 0}) {
t.Fatalf("ArgMinAxis = %v", an.RawInts())
}
// k*8 > n switches TopK to its sort path; the NaN lands last. The
// output is (k, cols): row 0 holds each column's top value.
tv, ti, err := TopK(nan, 2, 0)
if err != nil {
t.Fatal(err)
}
if tv.Dtype() != Float16 || tv.FloatAt(0) != 4 || tv.FloatAt(1) != 3 ||
tv.FloatAt(2) != 1 || !math.IsNaN(tv.FloatAt(3)) {
t.Fatalf("TopK values = %s %v", tv.Dtype(), tv.RawHalves())
}
if !slicesEqualI64(ti.RawInts(), []int64{1, 1, 0, 0}) {
t.Fatalf("TopK indices = %v", ti.RawInts())
}
// Pow against an int exponent promotes to half and narrows once.
pw, err := Pow(m2, mustFromInts(t, []int64{2, 2, 2, 2}, 2, 2))
if err != nil {
t.Fatal(err)
}
if pw.Dtype() != Float16 || pw.FloatAt(3) != 16 {
t.Fatalf("Pow = %s %v", pw.Dtype(), pw.FloatAt(3))
}
// AddC forces complex; the half values widen exactly.
ac := AddC(mustFromFloat16s(t, []float64{1.5}, 1), 2i)
if ac.Dtype() != Complex || ac.ComplexAt(0) != complex(1.5, 2) {
t.Fatalf("AddC on float16 = %s %v", ac.Dtype(), ac.ComplexAt(0))
}
// Concat of int and half lands on half through setConverted.
ci, err := Concat(mustFromInts(t, []int64{1, 2}, 2), mustFromFloat16s(t, []float64{3, 4}, 2), 0)
if err != nil {
t.Fatal(err)
}
if ci.Dtype() != Float16 || ci.FloatAt(3) != 4 {
t.Fatalf("Concat(int, float16) = %s %v", ci.Dtype(), ci.FloatAt(3))
}
// Astype half to complex keeps the real route.
cx, err := Astype(mustFromFloat16s(t, []float64{1.5}, 1), Complex)
if err != nil {
t.Fatal(err)
}
if cx.Dtype() != Complex || cx.ComplexAt(0) != complex(1.5, 0) {
t.Fatalf("Astype to complex = %s %v", cx.Dtype(), cx.ComplexAt(0))
}
// The narrowing constructors refuse a shape their values cannot fill.
if _, err := FromFloat16s([]float64{1}, 2); err == nil {
t.Fatal("FromFloat16s shape mismatch must error")
}
if _, err := HalvesFromArray([]uint16{1}, 2); err == nil {
t.Fatal("HalvesFromArray shape mismatch must error")
}
// A half grid interpolates through the exact widening: a linear
// field is reproduced exactly.
grid := mustFromFloat16s(t, []float64{0, 1, 2, 0, 1, 2}, 2, 3)
xs := mustFromFloats(t, []float64{0.5}, 1)
ys := mustFromFloats(t, []float64{0.5}, 1)
ip, err := Interpolate2D(grid, xs, ys, 0, 0, 1, 1)
if err != nil {
t.Fatal(err)
}
if ip.FloatAt(0) != 0.5 {
t.Fatalf("Interpolate2D half grid = %v", ip.FloatAt(0))
}
}
func slicesEqualI64(a, b []int64) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
func TestFloat16StringAndEqual(t *testing.T) {
h := mustFromFloat16s(t, []float64{1, 0.5}, 2)
if got := h.String(); !strings.HasPrefix(got, "float16 (2) [") || !strings.Contains(got, "0.5") {
t.Fatalf("String = %q", got)
}
if Dtype(0).String() != "int" || Float16.String() != "float16" {
t.Fatal("dtype names")
}
// Equal compares values, so -0.0 and +0.0 compare equal and NaN
// never does, exactly as for the other float dtypes.
zero := mustFromHalves(t, []uint16{0x0000, 0x3C00}, 2)
nzero := mustFromHalves(t, []uint16{0x8000, 0x3C00}, 2)
if !Equal(zero, nzero) {
t.Fatal("±0.0 halves must compare equal")
}
nan1 := mustFromHalves(t, []uint16{0x7E00}, 1)
nan2 := mustFromHalves(t, []uint16{0x7E01}, 1)
if Equal(nan1, nan2) {
t.Fatal("NaN halves must never compare equal")
}
// Equal on the very same slice is the documented pointer shortcut
// and answers true before any NaN logic, for every dtype.
// The dtype is part of the identity.
f32 := mustFromFloat32s(t, []float32{1, 0.5}, 2)
if Equal(h, f32) {
t.Fatal("float16 must not equal float32")
}
}
// TestFloat16ViewsAndMisc exercises the raw-payload paths on strided
// views: every kernel must see exactly the view's own elements.
func TestFloat16ViewsAndMisc(t *testing.T) {
h := mustFromFloat16s(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3)
// An interior column slice copies; a leading slice views.
// An interior column slice copies: rows [1..2] of columns 1, 2.
col, err := Slice(h, 1, 1, 3)
if err != nil {
t.Fatal(err)
}
if !slicesEqualU16(col.RawHalves(), []uint16{0x4000, 0x4200, 0x4500, 0x4600}) {
t.Fatalf("column slice = %v", col.RawHalves())
}
if got := Sum(col).Float(); got != 16 {
t.Fatalf("Sum over the slice = %v", got)
}
// A leading-dimension slice is a view; ArgMax materialises it and
// still finds the right element.
view, err := Slice(mustFromFloat16s(t, []float64{1, 5, 2}, 3), 0, 0, 3)
if err != nil {
t.Fatal(err)
}
if got, err := ArgMax(view); err != nil || got != 1 {
t.Fatalf("ArgMax over the view = %d, %v", got, err)
}
s, err := Sort(col)
if err != nil {
t.Fatal(err)
}
if !slicesEqualU16(s.RawHalves(), []uint16{0x4000, 0x4200, 0x4500, 0x4600}) {
t.Fatalf("Sort over the slice = %v", s.RawHalves())
}
// Unique keeps half and dedupes by value.
u, err := Unique(mustFromHalves(t, []uint16{0x3C00, 0x3C00, 0x8000}, 3))
if err != nil {
t.Fatal(err)
}
if u.Dtype() != Float16 || u.Len() != 2 {
t.Fatalf("Unique = %s %d", u.Dtype(), u.Len())
}
// Gather gathers per coordinate along dim 0: index {1, 0, 1} picks
// row 1, row 0, row 1 at the three column positions.
g, err := Gather(h, 0, mustFromInts(t, []int64{1, 0, 1}, 1, 3))
if err != nil {
t.Fatal(err)
}
if !slicesEqualU16(g.RawHalves(), []uint16{0x4400, 0x4000, 0x4600}) {
t.Fatalf("Gather = %v", g.RawHalves())
}
tk, err := Take(h, mustFromInts(t, []int64{2, 0}, 2))
if err != nil {
t.Fatal(err)
}
if !slicesEqualU16(tk.RawHalves(), []uint16{0x4200, 0x3C00}) {
t.Fatalf("Take = %v", tk.RawHalves())
}
rev := Reverse(h)
if rev.FloatAt(0) != 6 || rev.FloatAt(5) != 1 {
t.Fatalf("Reverse = %s", rev)
}
// SparseFrom and Dense round-trip the bits.
sp, err := SparseFrom(mustFromHalves(t, []uint16{0x3C00, 0x0000, 0x4000}, 3))
if err != nil {
t.Fatal(err)
}
if sp.Values.Dtype() != Float16 || sp.NNZ() != 2 {
t.Fatalf("SparseFrom dtype %s nnz %d", sp.Values.Dtype(), sp.NNZ())
}
dn, err := sp.Dense()
if err != nil {
t.Fatal(err)
}
if dn.Dtype() != Float16 || !slicesEqualU16(dn.RawHalves(), []uint16{0x3C00, 0x0000, 0x4000}) {
t.Fatalf("Sparse Dense = %v", dn.RawHalves())
}
// SpMul promotes and narrows once.
sm, err := SpMul(sp, mustFromFloat16s(t, []float64{2, 0, 0.5}, 3))
if err != nil {
t.Fatal(err)
}
if sm.Dtype() != Float16 || sm.FloatAt(0) != 2 || sm.FloatAt(2) != 1 {
t.Fatalf("SpMul = %s %v %v", sm.Dtype(), sm.FloatAt(0), sm.FloatAt(2))
}
}
// TestFloat16Refusals pins the loud refusals: the matmul and einsum
// kernels are not offered for the half dtype, and conversion with
// Astype is the documented route.
func TestFloat16Refusals(t *testing.T) {
h2 := mustFromFloat16s(t, []float64{1, 2, 3, 4}, 2, 2)
f2 := mustFromFloats(t, []float64{1, 0, 0, 1}, 2, 2)
if _, err := MatMul2D(h2, f2); err == nil || !strings.Contains(err.Error(), "float16 operands are not supported") {
t.Fatalf("MatMul2D float16: %v", err)
}
if _, err := MatMul2D(f2, h2); err == nil || !strings.Contains(err.Error(), "float16 operands are not supported") {
t.Fatalf("MatMul2D float16 on the right: %v", err)
}
if _, err := Einsum("ij,jk->ik", h2, f2); err == nil || !strings.Contains(err.Error(), "float16 operands are not supported") {
t.Fatalf("Einsum float16: %v", err)
}
if _, err := Einsum("ij->", h2); err == nil || !strings.Contains(err.Error(), "float16 operands are not supported") {
t.Fatalf("Einsum reduce-all float16: %v", err)
}
sp := &SparseCOO{Indices: mustFromInts(t, []int64{0, 0}, 1, 2), Values: mustFromFloat16s(t, []float64{1}, 1), Shape: []int{2, 2}}
if _, err := SpMatMul(sp, f2); err == nil || !strings.Contains(err.Error(), "float16 is not supported") {
t.Fatalf("SpMatMul float16: %v", err)
}
// The float64 side is unchanged: float64 by float64 still computes.
got, err := MatMul2D(f2, f2)
if err != nil || got.Dtype() != Float || got.FloatAt(0) != 1 {
t.Fatalf("MatMul2D float64 baseline moved: %v %s", err, got)
}
// Kron promotes to half and narrows once.
k, err := Kron(h2, mustFromInts(t, []int64{1, 1, 1, 1}, 2, 2))
if err != nil {
t.Fatal(err)
}
if k.Dtype() != Float16 || k.FloatAt(0) != 1 || k.FloatAt(3) != 2 || k.FloatAt(15) != 4 {
t.Fatalf("Kron = %s %v %v %v", k.Dtype(), k.FloatAt(0), k.FloatAt(3), k.FloatAt(15))
}
// Elements refuses only the genuine narrowings.
h1 := mustFromFloat16s(t, []float64{1.5}, 1)
if _, err := h1.Elements[int64](); err == nil || !strings.Contains(err.Error(), "cannot narrow float16 to int64") {
t.Fatalf("Elements int64 on float16: %v", err)
}
if v, err := h1.Elements[float64](); err != nil || v[0] != 1.5 {
t.Fatalf("Elements float64 on float16: %v %v", v, err)
}
if v, err := h1.Elements[float32](); err != nil || v[0] != 1.5 {
t.Fatalf("Elements float32 on float16: %v %v", v, err)
}
}
// The ordering trap: half bits do not order as uint16 for negatives
// (−1.0 is 0xBC00, which bit order places after 2.0), so the kernels
// that order or clamp must widen to float64 first. Each pin below
// fails if the glue ever sorts or clamps the raw payload, and the
// overflow pin holds the compute-in-float64-narrow-once contract's
// sticky infinity.
func TestFloat16NegativeOrdering(t *testing.T) {
x, xerr := FromFloat16s([]float64{2, -1, 0.5}, 3)
if xerr != nil {
t.Fatal(xerr)
}
sorted, serr := Sort(x)
if serr != nil {
t.Fatal(serr)
}
want := []uint16{HalfFromFloat64(-1), HalfFromFloat64(0.5), HalfFromFloat64(2)}
for i, h := range want {
if got := sorted.RawHalves()[i]; got != h {
t.Fatalf("Sort ascending [%d] = %#04x, want %#04x", i, got, h)
}
}
order, oerr := ArgSort(x)
if oerr != nil {
t.Fatal(oerr)
}
if got0, got1, got2 := order.FloatAt(0), order.FloatAt(1), order.FloatAt(2); got0 != 1 || got1 != 2 || got2 != 0 {
t.Fatalf("ArgSort ascending = [%g %g %g], want [1 2 0]", got0, got1, got2)
}
fIn, ferr := FromFloat16s([]float64{-0.5, 0.25}, 2)
if ferr != nil {
t.Fatal(ferr)
}
clipped, cerr := ClipF(fIn, -0.25, 0.5)
if cerr != nil {
t.Fatal(cerr)
}
if got := clipped.RawHalves()[0]; got != HalfFromFloat64(-0.25) {
t.Fatalf("ClipF on a negative = %#04x, want %#04x", got, HalfFromFloat64(-0.25))
}
clipIn, clipErr := FromFloat16s([]float64{-5, 0.25}, 2)
if clipErr != nil {
t.Fatal(clipErr)
}
clamped, ierr := ClipI(clipIn, -1, 1)
if ierr != nil {
t.Fatal(ierr)
}
if math.Abs(clamped.FloatAt(0)+1) > 1e-12 {
t.Fatalf("ClipI on a negative past the wall = %g, want -1", clamped.FloatAt(0))
}
lhs, lerr := FromFloat16s([]float64{65504, 65504}, 2)
if lerr != nil {
t.Fatal(lerr)
}
rhs, rerr := FromFloat16s([]float64{2, 2}, 2)
if rerr != nil {
t.Fatal(rerr)
}
overflow, merr := Mul(lhs, rhs)
if merr != nil {
t.Fatal(merr)
}
if h := overflow.RawHalves()[0]; h != 0x7C00 {
t.Fatalf("65504*2 in half = %#04x, want the +Inf half 0x7C00", h)
}
back, serr2 := Sub(overflow, rhs)
if serr2 != nil {
t.Fatal(serr2)
}
if h := back.RawHalves()[0]; h != 0x7C00 {
t.Fatalf("+Inf minus 65504 in half = %#04x, want the sticky +Inf 0x7C00", h)
}
}