Files
tensor/internal/core/float32_test.go
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

335 lines
9.0 KiB
Go
Raw Permalink 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"
"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
}
}