979 lines
32 KiB
Go
979 lines
32 KiB
Go
// 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)
|
||
}
|
||
}
|