227 lines
6.0 KiB
Go
227 lines
6.0 KiB
Go
// 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)
|
||
}
|
||
}
|
||
}
|
||
}
|