feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,165 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user