feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,334 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package core
|
||||
|
||||
import (
|
||||
"math"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func mustFromFloat32s(t *testing.T, vals []float32, shape ...int) *Array {
|
||||
t.Helper()
|
||||
a, err := FromFloat32s(vals, shape...)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat32s(%v, %v): %v", vals, shape, err)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
func TestFloat32Constructors(t *testing.T) {
|
||||
a := mustFromFloat32s(t, []float32{1.5, -2.5}, 2)
|
||||
if a.Dtype() != Float32 || a.Len() != 2 {
|
||||
t.Fatalf("float32 array: %s len %d", a.Dtype(), a.Len())
|
||||
}
|
||||
// floatAt reads exactly.
|
||||
if v := a.FloatAt(0); v != 1.5 {
|
||||
t.Fatalf("floatAt: %v", v)
|
||||
}
|
||||
|
||||
z, _ := Zeros(Float32, 3)
|
||||
if z.Dtype() != Float32 || z.FloatAt(1) != 0 {
|
||||
t.Fatalf("Zeros float32: %s", z)
|
||||
}
|
||||
o, _ := Ones(Float32, 2)
|
||||
if o.FloatAt(0) != 1 {
|
||||
t.Fatalf("Ones float32: %s", o)
|
||||
}
|
||||
f, _ := FullF32s(0.5, 2)
|
||||
if f.FloatAt(1) != 0.5 {
|
||||
t.Fatalf("FullF32s: %s", f)
|
||||
}
|
||||
// The dtype is part of the identity: float never equals float32.
|
||||
wide := mustFromFloats(t, []float64{1.5, -2.5}, 2)
|
||||
if Equal(wide, mustFromFloat32s(t, []float32{1.5, -2.5}, 2)) {
|
||||
t.Fatal("float and float32 must not compare equal")
|
||||
}
|
||||
if got := a.String(); !strings.Contains(got, "float32") {
|
||||
t.Fatalf("String: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFloat32PromotionMatrix(t *testing.T) {
|
||||
f32 := mustFromFloat32s(t, []float32{1.5}, 1)
|
||||
i := mustFromInts(t, []int64{1}, 1)
|
||||
f64 := mustFromFloats(t, []float64{1.5}, 1)
|
||||
c := mustFromComplexes(t, []complex128{1}, 1)
|
||||
|
||||
cases := []struct {
|
||||
b *Array
|
||||
want Dtype
|
||||
}{
|
||||
{f32, Float32},
|
||||
{i, Float32},
|
||||
{f64, Float},
|
||||
{c, Complex},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
sum, err := Add(f32, tc.b)
|
||||
if err != nil {
|
||||
t.Fatalf("Add(%s): %v", tc.b.Dtype(), err)
|
||||
}
|
||||
if sum.Dtype() != tc.want {
|
||||
t.Fatalf("float32 + %s = %s, want %s", tc.b.Dtype(), sum.Dtype(), tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFloat32Arithmetic(t *testing.T) {
|
||||
a := mustFromFloat32s(t, []float32{0.1, 2}, 2)
|
||||
b := mustFromFloat32s(t, []float32{0.2, 4}, 2)
|
||||
|
||||
sum, err := Add(a, b)
|
||||
if err != nil || sum.Dtype() != Float32 {
|
||||
t.Fatalf("Add: %s %v", sum, err)
|
||||
}
|
||||
// 0.1+0.2 computed in float64 and rounded once, the float32 result.
|
||||
if v := sum.FloatAt(0); v != 0.30000001192092896 {
|
||||
t.Fatalf("Add value: %v", v)
|
||||
}
|
||||
|
||||
prod, _ := Mul(a, b)
|
||||
if v := prod.FloatAt(1); v != 8 {
|
||||
t.Fatalf("Mul value: %v", v)
|
||||
}
|
||||
|
||||
// Division keeps float32.
|
||||
q, _ := Div(b, a)
|
||||
if q.Dtype() != Float32 {
|
||||
t.Fatalf("Div dtype: %s", q.Dtype())
|
||||
}
|
||||
if v := q.FloatAt(1); v != 2 {
|
||||
t.Fatalf("Div value: %v", v)
|
||||
}
|
||||
|
||||
// Weak scalars keep float32 (NEP-50 style).
|
||||
if v := AddI(a, 1).FloatAt(1); v != 3 {
|
||||
t.Fatalf("AddI: %v", v)
|
||||
}
|
||||
if got := AddI(a, 1).Dtype(); got != Float32 {
|
||||
t.Fatalf("AddI dtype: %s", got)
|
||||
}
|
||||
if got := AddF(a, 0.5).Dtype(); got != Float32 {
|
||||
t.Fatalf("AddF keeps float32: %s", got)
|
||||
}
|
||||
if v := AddF(a, 0.25).FloatAt(1); v != 2.25 {
|
||||
t.Fatalf("AddF value: %v", v)
|
||||
}
|
||||
// int arrays still promote to float64 under F scalars.
|
||||
if got := AddF(mustFromInts(t, []int64{1}, 1), 0.5).Dtype(); got != Float {
|
||||
t.Fatalf("int AddF dtype: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFloat32MathFuncs(t *testing.T) {
|
||||
a := mustFromFloat32s(t, []float32{4}, 1)
|
||||
|
||||
sqrt, err := Sqrt(a)
|
||||
if err != nil || sqrt.Dtype() != Float32 {
|
||||
t.Fatalf("Sqrt: %s %v", sqrt, err)
|
||||
}
|
||||
if v := sqrt.FloatAt(0); v != 2 {
|
||||
t.Fatalf("Sqrt value: %v", v)
|
||||
}
|
||||
|
||||
// Rounding keeps the width.
|
||||
r, _ := Floor(mustFromFloat32s(t, []float32{2.7}, 1))
|
||||
if r.Dtype() != Float32 || r.FloatAt(0) != 2 {
|
||||
t.Fatalf("Floor: %s %v", r, r.FloatAt(0))
|
||||
}
|
||||
|
||||
// Abs keeps float32; complex magnitudes still go to float64.
|
||||
ab := Abs(mustFromFloat32s(t, []float32{-2.5}, 1))
|
||||
if ab.Dtype() != Float32 || ab.FloatAt(0) != 2.5 {
|
||||
t.Fatalf("Abs: %s %v", ab, ab.FloatAt(0))
|
||||
}
|
||||
|
||||
tanh, _ := Tanh(mustFromFloat32s(t, []float32{0}, 1))
|
||||
if tanh.Dtype() != Float32 || tanh.FloatAt(0) != 0 {
|
||||
t.Fatalf("Tanh: %s", tanh)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFloat32ReductionsAndMatMul(t *testing.T) {
|
||||
m := mustFromFloat32s(t, []float32{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
|
||||
// Scalar reductions keep the float scalar kind for float32 arrays.
|
||||
sum := Sum(m)
|
||||
if !sum.IsFloat() || sum.Float() != 21 {
|
||||
t.Fatalf("Sum float32: %s", sum)
|
||||
}
|
||||
mn, err := Min(m)
|
||||
if err != nil || !mn.IsFloat() || mn.Float() != 1 {
|
||||
t.Fatalf("Min float32: %s %v", mn, err)
|
||||
}
|
||||
mx, err := Max(m)
|
||||
if err != nil || !mx.IsFloat() || mx.Float() != 6 {
|
||||
t.Fatalf("Max float32: %s %v", mx, err)
|
||||
}
|
||||
mean, err := Mean(m)
|
||||
if err != nil || mean != 3.5 {
|
||||
t.Fatalf("Mean float32: %v %v", mean, err)
|
||||
}
|
||||
f32a := mustFromFloat32s(t, []float32{1, 2}, 2)
|
||||
f32b := mustFromFloat32s(t, []float32{3, 4}, 2)
|
||||
d, err := Dot(f32a, f32b)
|
||||
if err != nil || !d.IsFloat() || d.Float() != 11 {
|
||||
t.Fatalf("Dot float32: %s %v", d, err)
|
||||
}
|
||||
|
||||
sums, err := SumAxis(m, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("SumAxis: %v", err)
|
||||
}
|
||||
if sums.Dtype() != Float32 {
|
||||
t.Fatalf("SumAxis dtype: %s", sums.Dtype())
|
||||
}
|
||||
if v := sums.FloatAt(0); v != 6 {
|
||||
t.Fatalf("SumAxis value: %v", v)
|
||||
}
|
||||
|
||||
maxes, _ := MaxAxis(m, 1)
|
||||
if maxes.Dtype() != Float32 || maxes.FloatAt(1) != 6 {
|
||||
t.Fatalf("MaxAxis: %s", maxes)
|
||||
}
|
||||
|
||||
// Mean is float64 regardless.
|
||||
axisMean, _ := MeanAxis(m, 1)
|
||||
if axisMean.Dtype() != Float {
|
||||
t.Fatalf("MeanAxis dtype: %s", axisMean.Dtype())
|
||||
}
|
||||
|
||||
// MatMul accumulates in float64 and rounds once.
|
||||
w := mustFromFloat32s(t, []float32{1, 2, 3, 4, 5, 6}, 3, 2)
|
||||
p, err := MatMul2D(m, w)
|
||||
if err != nil {
|
||||
t.Fatalf("MatMul: %v", err)
|
||||
}
|
||||
if p.Dtype() != Float32 {
|
||||
t.Fatalf("MatMul dtype: %s", p.Dtype())
|
||||
}
|
||||
// [[1,2,3],[4,5,6]]·[[1,2],[3,4],[5,6]] = [[22,28],[49,64]]
|
||||
if v := p.FloatAt(0); v != 22 {
|
||||
t.Fatalf("MatMul value: %v", v)
|
||||
}
|
||||
if v := p.FloatAt(3); v != 64 {
|
||||
t.Fatalf("MatMul value: %v", v)
|
||||
}
|
||||
|
||||
// Vector shapes keep float32 too.
|
||||
v1 := mustFromFloat32s(t, []float32{1, 2}, 2)
|
||||
mv, err := MatMul2D(m, mustFromFloat32s(t, []float32{1, 1, 1}, 3))
|
||||
if err != nil || mv.Dtype() != Float32 {
|
||||
t.Fatalf("matrix × vector: %s %v", mv, err)
|
||||
}
|
||||
if v := mv.FloatAt(0); v != 6 {
|
||||
t.Fatalf("matrix × vector value: %v", v)
|
||||
}
|
||||
_ = v1
|
||||
}
|
||||
|
||||
func TestFloat32Machinery(t *testing.T) {
|
||||
f := mustFromFloat32s(t, []float32{1, 2, 3, 4}, 2, 2)
|
||||
|
||||
// Identity works for every dtype.
|
||||
id, err := Identity(Float32, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("Identity float32: %v", err)
|
||||
}
|
||||
if id.Dtype() != Float32 || id.FloatAt(0) != 1 || id.FloatAt(4) != 1 {
|
||||
t.Fatalf("Identity float32: %s", id)
|
||||
}
|
||||
// int × int Pow stays int; a negative exponent is an error.
|
||||
pi := mustFromInts(t, []int64{2, 3}, 2)
|
||||
pe := mustFromInts(t, []int64{3, 2}, 2)
|
||||
pow, err := Pow(pi, pe)
|
||||
if err != nil || pow.Dtype() != Int || pow.RawInts()[0] != 8 || pow.RawInts()[1] != 9 {
|
||||
t.Fatalf("Pow int: %s %v", pow, err)
|
||||
}
|
||||
if _, err := Pow(pi, mustFromInts(t, []int64{1, -1}, 2)); err == nil ||
|
||||
!strings.Contains(err.Error(), "negative exponent") {
|
||||
t.Fatalf("Pow negative exponent: %v", err)
|
||||
}
|
||||
// PowI keeps float32.
|
||||
pf, err := PowI(mustFromFloat32s(t, []float32{2, 3}, 2), 2)
|
||||
if err != nil || pf.Dtype() != Float32 || pf.FloatAt(1) != 9 {
|
||||
t.Fatalf("PowI float32: %s %v", pf, err)
|
||||
}
|
||||
|
||||
// Mask selection, Where, comparisons and broadcasting all carry the
|
||||
// dtype.
|
||||
mask, err := GtI(f, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("GtI: %v", err)
|
||||
}
|
||||
sel, _ := Select(f, mask)
|
||||
if sel.Dtype() != Float32 || sel.Len() != 2 {
|
||||
t.Fatalf("Mask: %s", sel)
|
||||
}
|
||||
w, _ := Where(mask, f, mustFromFloat32s(t, []float32{9, 9, 9, 9}, 2, 2))
|
||||
if w.Dtype() != Float32 || w.FloatAt(0) != 9 {
|
||||
t.Fatalf("Where: %s", w)
|
||||
}
|
||||
lt, _ := LtF(f, 3)
|
||||
if got := lt.FloatAt(2); got != 0 {
|
||||
t.Fatalf("LtF on float32: %v", got)
|
||||
}
|
||||
b, _ := Slice(f, 0, 1, 2)
|
||||
if b.Dtype() != Float32 || b.FloatAt(0) != 3 {
|
||||
t.Fatalf("Slice: %s", b)
|
||||
}
|
||||
tt := Transpose(f)
|
||||
if tt.FloatAt(1) != 3 {
|
||||
t.Fatalf("Transpose: %s", tt)
|
||||
}
|
||||
cat, _ := Concat(f, f, 1)
|
||||
if cat.Dtype() != Float32 || cat.Shape()[1] != 4 {
|
||||
t.Fatalf("Concat: %s", cat)
|
||||
}
|
||||
|
||||
// Sort, Clip and elements.
|
||||
s, _ := Sort(mustFromFloat32s(t, []float32{3, 1, 2}, 3))
|
||||
if s.FloatAt(0) != 1 {
|
||||
t.Fatalf("Sort: %s", s)
|
||||
}
|
||||
clip, _ := ClipI(mustFromFloat32s(t, []float32{-5, 3}, 2), 0, 1)
|
||||
if clip.Dtype() != Float32 || clip.FloatAt(0) != 0 || clip.FloatAt(1) != 1 {
|
||||
t.Fatalf("ClipI: %s", clip)
|
||||
}
|
||||
clipF, _ := ClipF(mustFromFloat32s(t, []float32{-5, 3}, 2), 0, 1)
|
||||
if clipF.Dtype() != Float32 {
|
||||
t.Fatalf("ClipF dtype: %s", clipF)
|
||||
}
|
||||
|
||||
// Elements conversions along the ladder.
|
||||
vals, err := f.Elements[float32]()
|
||||
if err != nil || vals[3] != 4 {
|
||||
t.Fatalf("Elements[float32]: %v %v", vals, err)
|
||||
}
|
||||
widened, err := f.Elements[float64]()
|
||||
if err != nil || widened[0] != 1 {
|
||||
t.Fatalf("Elements[float64] from float32: %v %v", widened, err)
|
||||
}
|
||||
if _, err := mustFromFloats(t, []float64{1}, 1).Elements[float32](); err == nil ||
|
||||
!strings.Contains(err.Error(), "cannot narrow float to float32") {
|
||||
t.Fatalf("Elements float to float32: %v", err)
|
||||
}
|
||||
|
||||
// Generator.Float32s stays in range.
|
||||
g := NewGenerator(21)
|
||||
r, err := Float32s(g, 500)
|
||||
if err != nil || r.Dtype() != Float32 {
|
||||
t.Fatalf("Float32s: %s %v", r, err)
|
||||
}
|
||||
for i := range r.Len() {
|
||||
v := r.FloatAt(i)
|
||||
if v < 0 || v >= 1 {
|
||||
t.Fatalf("Float32s out of [0,1): %v", v)
|
||||
}
|
||||
}
|
||||
if _, ok := any(math.NaN()).(float32); ok {
|
||||
_ = ok // keep math imported if unused paths change
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user