// Copyright (c) 2026 Petr Balvín (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 } }