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

227 lines
6.0 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 (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// sparseCSRFromEntries builds a SparseCSR straight from (row, col,
// value) triples, going through the public COO path.
func sparseCSRFromEntries(t *testing.T, rows, cols int, entries [][3]float64) *SparseCSR {
t.Helper()
idx := make([]int64, 0, len(entries)*2)
vals := make([]float64, 0, len(entries))
for _, e := range entries {
idx = append(idx, int64(e[0]), int64(e[1]))
vals = append(vals, e[2])
}
indices, ierr := core.FromInts(idx, len(entries), 2)
if ierr != nil {
t.Fatalf("FromInts: %v", ierr)
}
coo, err := core.NewSparseCOO(indices, floatsToArray(vals, []int{len(entries)}), []int{rows, cols})
if err != nil {
t.Fatalf("NewSparseCOO: %v", err)
}
csr, err := CSRFromCOO(coo)
if err != nil {
t.Fatalf("CSRFromCOO: %v", err)
}
return csr
}
// TestSparseMatMulDense checks the sparse-dense product against an
// explicit dense computation on a ragged 4×5 matrix.
func TestSparseMatMulDense(t *testing.T) {
csr := sparseCSRFromEntries(t, 4, 5, [][3]float64{
{0, 0, 2}, {0, 3, -1},
{1, 1, 3}, {1, 4, 0.5},
{2, 0, -1}, {2, 2, 4}, {2, 4, 2},
{3, 3, 1},
})
x := mustFloats(t, []float64{
1, -1, 2, 0, 3,
0.5, 1, 1, 1, -2,
}, 5, 2)
got, err := csr.MatMulDense(x)
if err != nil {
t.Fatalf("MatMulDense: %v", err)
}
if got.Shape()[0] != 4 || got.Shape()[1] != 2 {
t.Fatalf("result shape %v, want (4, 2)", got.Shape())
}
for i := range 4 {
for m := range 2 {
s := 0.0
for j := range 5 {
s += csrValuesAt(csr, i, j) * x.FloatAt(j*2+m)
}
if math.Abs(got.FloatAt(i*2+m)-s) > 1e-14 {
t.Fatalf("out(%d, %d) = %g, want %g", i, m, got.FloatAt(i*2+m), s)
}
}
}
if _, err := csr.MatMulDense(mustFloats(t, []float64{1, 2, 3, 4, 5})); err == nil {
t.Fatal("vector operand: want an error")
}
}
// csrValuesAt reads a CSR entry that may be stored or implicit zero.
func csrValuesAt(c *SparseCSR, i, j int) float64 {
for p := c.RowStart[i]; p < c.RowStart[i+1]; p++ {
if c.ColIdx[p] == j {
return c.Values[p]
}
}
return 0
}
// TestSparseMatMulSparse checks the sparse product against its dense
// equivalent, including cancellation producing an implicit zero and
// the transpose route for the asymmetric case.
func TestSparseMatMulSparse(t *testing.T) {
a := sparseCSRFromEntries(t, 3, 4, [][3]float64{
{0, 0, 1}, {0, 2, 2},
{1, 1, -1}, {1, 3, 4},
{2, 0, 0.5}, {2, 2, 1},
})
b := sparseCSRFromEntries(t, 4, 3, [][3]float64{
{0, 1, 3},
{1, 0, 2}, {1, 2, -1},
{2, 1, 1.5}, {2, 2, 2},
{3, 0, -0.5},
})
ab, err := a.MatMulSparse(b)
if err != nil {
t.Fatalf("MatMulSparse: %v", err)
}
if ab.Rows != 3 || ab.Cols != 3 {
t.Fatalf("product shape %d×%d, want 3×3", ab.Rows, ab.Cols)
}
for i := range 3 {
for j := range 3 {
s := 0.0
for k := range 4 {
s += csrValuesAt(a, i, k) * csrValuesAt(b, k, j)
}
if math.Abs(csrValuesAt(ab, i, j)-s) > 1e-14 {
t.Fatalf("(%d, %d) = %g, want %g", i, j, csrValuesAt(ab, i, j), s)
}
}
}
if _, err := a.MatMulSparse(sparseCSRFromEntries(t, 3, 2, [][3]float64{{0, 0, 1}})); err == nil {
t.Fatal("inner dimension mismatch: want an error")
}
}
// TestSparseTranspose checks the transpose against dense arithmetic
// and against double transposition restoring the original.
func TestSparseTranspose(t *testing.T) {
csr := sparseCSRFromEntries(t, 3, 5, [][3]float64{
{0, 1, 2}, {0, 4, -3},
{1, 0, 1.5},
{2, 1, 4}, {2, 2, -1}, {2, 4, 0.25},
})
at := csr.Transpose()
if at.Rows != 5 || at.Cols != 3 {
t.Fatalf("transpose shape %d×%d, want 5×3", at.Rows, at.Cols)
}
for i := range 5 {
for j := range 3 {
if csrValuesAt(at, i, j) != csrValuesAt(csr, j, i) {
t.Fatalf("transpose(%d, %d) = %g, want %g", i, j,
csrValuesAt(at, i, j), csrValuesAt(csr, j, i))
}
}
}
back := at.Transpose()
for i := range 3 {
for j := range 5 {
if csrValuesAt(back, i, j) != csrValuesAt(csr, i, j) {
t.Fatalf("double transpose(%d, %d) = %g, want %g", i, j,
csrValuesAt(back, i, j), csrValuesAt(csr, i, j))
}
}
}
// The transpose product Aᵀ·A via MatMulSparse must agree with the
// dense Gram matrix.
g, err := at.MatMulSparse(csr)
if err != nil {
t.Fatalf("Aᵀ·A: %v", err)
}
for i := range 5 {
for j := range 5 {
s := 0.0
for k := range 3 {
s += csrValuesAt(csr, k, i) * csrValuesAt(csr, k, j)
}
if math.Abs(csrValuesAt(g, i, j)-s) > 1e-14 {
t.Fatalf("Gram(%d, %d) = %g, want %g", i, j, csrValuesAt(g, i, j), s)
}
}
}
}
// TestSparseOpsLargeCheck runs the three operations against dense
// arithmetic on a banded 40×40 structure, the shape PDE stencils
// produce.
func TestSparseOpsLargeCheck(t *testing.T) {
const n = 40
vals := make([][3]float64, 0, 3*n)
for i := range n {
vals = append(vals, [3]float64{float64(i), float64(i), 4})
if i+1 < n {
vals = append(vals, [3]float64{float64(i), float64(i + 1), -1})
vals = append(vals, [3]float64{float64(i + 1), float64(i), -1})
}
}
a := sparseCSRFromEntries(t, n, n, vals)
x := make([]float64, n*n)
for i := range n {
for j := range n {
x[i*n+j] = math.Sin(float64(i + j + 1))
}
}
xd := floatsToArray(x, []int{n, n})
got, err := a.MatMulDense(xd)
if err != nil {
t.Fatalf("MatMulDense: %v", err)
}
for i := range n {
for j := range n {
s := 0.0
for k := range n {
if v := csrValuesAt(a, i, k); v != 0 {
s += v * x[k*n+j]
}
}
if math.Abs(got.FloatAt(i*n+j)-s) > 1e-12 {
t.Fatalf("(%d, %d): %g vs %g", i, j, got.FloatAt(i*n+j), s)
}
}
}
sq, err := a.MatMulSparse(a)
if err != nil {
t.Fatalf("A·A: %v", err)
}
for i := range n {
for j := range n {
s := 0.0
for k := range n {
if v := csrValuesAt(a, i, k); v != 0 {
if w := csrValuesAt(a, k, j); w != 0 {
s += v * w
}
}
}
if math.Abs(csrValuesAt(sq, i, j)-s) > 1e-12 {
t.Fatalf("A²(%d, %d): %g vs %g", i, j, csrValuesAt(sq, i, j), s)
}
}
}
}