Files
tensor/internal/core/float16_test.go
T

979 lines
32 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"
"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)
}
}