feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+165
View File
@@ -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)
}
}
}