Files
tensor/linalg/initialisers_sparse_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

166 lines
3.9 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 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)
}
}
}