166 lines
3.9 KiB
Go
166 lines
3.9 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package linalg
|
||
|
||
import (
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
"testing"
|
||
)
|
||
|
||
func TestTruncatedNormal(t *testing.T) {
|
||
g := core.NewGenerator(42)
|
||
w := core.TruncatedNormal(g, []int{1000}, 0, 1)
|
||
if w.Len() != 1000 {
|
||
t.Errorf("TruncatedNormal len: %d", w.Len())
|
||
}
|
||
// All values must be in [-2, 2].
|
||
for i := range w.RawFloat32s() {
|
||
v := float64(w.RawFloat32s()[i])
|
||
if v < -2 || v > 2 {
|
||
t.Errorf("TruncatedNormal [%d]: %v out of [-2, 2]", i, v)
|
||
}
|
||
}
|
||
}
|
||
|
||
func TestZerosLikeOnesLike(t *testing.T) {
|
||
a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2)
|
||
z := core.ZerosLike(a)
|
||
if z.Shape()[0] != 2 || z.Shape()[1] != 2 {
|
||
t.Errorf("ZerosLike shape: %v", z.Shape())
|
||
}
|
||
for i := range z.RawFloats() {
|
||
if z.RawFloats()[i] != 0 {
|
||
t.Errorf("ZerosLike [%d]: %v, want 0", i, z.RawFloats()[i])
|
||
}
|
||
}
|
||
o := core.OnesLike(a)
|
||
for i := range o.RawFloats() {
|
||
if o.RawFloats()[i] != 1 {
|
||
t.Errorf("OnesLike [%d]: %v, want 1", i, o.RawFloats()[i])
|
||
}
|
||
}
|
||
// Mutating the copy must not affect a.
|
||
z.RawFloats()[0] = 99
|
||
if a.RawFloats()[0] == 99 {
|
||
t.Error("ZerosLike: mutation leaked into source array")
|
||
}
|
||
}
|
||
|
||
func TestSparseRoundTrip(t *testing.T) {
|
||
dense := mustFromFloats(t, []float64{
|
||
0, 1, 0,
|
||
2, 0, 3,
|
||
0, 0, 4,
|
||
}, 3, 3)
|
||
s, err := core.SparseFrom(dense)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if s.NNZ() != 4 {
|
||
t.Errorf("NNZ: %d, want 4", s.NNZ())
|
||
}
|
||
back, err := s.Dense()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
for i := range 9 {
|
||
v, _ := core.FloatAt(back, i/3, i%3)
|
||
orig, _ := core.FloatAt(dense, i/3, i%3)
|
||
if v != orig {
|
||
t.Errorf("sparse round-trip [%d]: got %v, want %v", i, v, orig)
|
||
}
|
||
}
|
||
}
|
||
|
||
func TestSparseFromKeepsDtype(t *testing.T) {
|
||
// SparseFrom on an int array must keep the int dtype through the
|
||
// round trip (the values array is part of the identity).
|
||
di := mustFromInts(t, []int64{0, 5, 0, 7}, 2, 2)
|
||
si, err := core.SparseFrom(di)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if si.Values.Dtype() != core.Int {
|
||
t.Fatalf("SparseFrom int: values dtype %s, want int", si.Values.Dtype())
|
||
}
|
||
back, err := si.Dense()
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !core.Equal(back, di) {
|
||
t.Errorf("int sparse round-trip: got %s, want %s", back, di)
|
||
}
|
||
|
||
df := mustFromFloat32s(t, []float32{0, 1.5, 0, 2.5}, 2, 2)
|
||
sf, err := core.SparseFrom(df)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if sf.Values.Dtype() != core.Float32 {
|
||
t.Fatalf("SparseFrom float32: values dtype %s, want float32", sf.Values.Dtype())
|
||
}
|
||
}
|
||
|
||
func TestSparseMul(t *testing.T) {
|
||
dense := mustFromFloats(t, []float64{
|
||
0, 1, 0,
|
||
2, 0, 3,
|
||
0, 0, 4,
|
||
}, 3, 3)
|
||
s, _ := core.SparseFrom(dense)
|
||
mul, err := core.FromFloats([]float64{10, 20, 30, 40, 50, 60, 70, 80, 90}, 3, 3)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
out, err := core.SpMul(s, mul)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// Only non-zero positions are filled.
|
||
// dense[0,1]=1 * mul[0,1]=20 = 20
|
||
// dense[1,0]=2 * mul[1,0]=40 = 80
|
||
// dense[1,2]=3 * mul[1,2]=60 = 180
|
||
// dense[2,2]=4 * mul[2,2]=90 = 360
|
||
expect := []float64{0, 20, 0, 80, 0, 180, 0, 0, 360}
|
||
for i, w := range expect {
|
||
v, _ := core.FloatAt(out, i/3, i%3)
|
||
if v != w {
|
||
t.Errorf("SpMul [%d]: got %v, want %v", i, v, w)
|
||
}
|
||
}
|
||
}
|
||
|
||
func TestSparseMatMul(t *testing.T) {
|
||
// Sparse 2×3 times dense 3×2.
|
||
indices, err := core.FromInts([]int64{
|
||
0, 0,
|
||
1, 2,
|
||
}, 2, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
values, err := core.FromFloats([]float64{1, 2}, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
s := &core.SparseCOO{Indices: indices, Values: values, Shape: []int{2, 3}}
|
||
dense, _ := core.FromFloats([]float64{
|
||
1, 2,
|
||
3, 4,
|
||
5, 6,
|
||
}, 3, 2)
|
||
out, err := core.SpMatMul(s, dense)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// Row 0: [1, 0, 0] · [[1,2],[3,4],[5,6]] = [1, 2]
|
||
// Row 1: [0, 0, 2] · [[1,2],[3,4],[5,6]] = [10, 12]
|
||
for i, w := range []float64{1, 2, 10, 12} {
|
||
v, _ := core.FloatAt(out, i/2, i%2)
|
||
if v != w {
|
||
t.Errorf("SpMatMul [%d]: got %v, want %v", i, v, w)
|
||
}
|
||
}
|
||
}
|