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

124 lines
4.1 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package linalg
import (
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/engine"
)
// sparseCSCMatVecReference accumulates the product from the stored
// entries directly. The test matrices carry small integer values, so
// every sum is exact whatever order the terms arrive in and the
// reference's answer pins MatVec bit for bit.
func sparseCSCMatVecReference(c *SparseCSC, x []float64) []float64 {
y := make([]float64, c.Rows)
for j, v := range x {
for p := c.ColStart[j]; p < c.ColStart[j+1]; p++ {
y[c.RowIdx[p]] += c.Values[p] * v
}
}
return y
}
// TestSparseCSCMatVecSplitMatchesReference drives the split path of
// the CSC product: a banded 131072-row matrix clears the split's work
// and width budgets, and its answer must match the exact reference.
// A band is the shape the split exists for, three entries per column.
func TestSparseCSCMatVecSplitMatchesReference(t *testing.T) {
// The split is a decision of the worker policy, whose default
// follows the machine's CPU count, so the policy is pinned here:
// the guard below states the matrix's shape, and the split is
// driven on a single-core machine exactly as on a many-core one.
prev := engine.SetNumWorkers(4)
defer engine.SetNumWorkers(prev)
const n = 131072
c := &SparseCSC{Rows: n, Cols: n, ColStart: make([]int, n+1)}
for j := range n {
c.ColStart[j] = len(c.RowIdx)
for _, r := range []int{j - 1, j, j + 1} {
if r < 0 || r >= n {
continue
}
c.RowIdx = append(c.RowIdx, r)
c.Values = append(c.Values, float64(r%7+1))
}
}
c.ColStart[n] = len(c.RowIdx)
if cscMatVecBlocks(c.Rows, c.Cols, len(c.Values)) <= 1 {
t.Fatal("the banded matrix does not reach the split, the test would not drive it")
}
xv := make([]float64, n)
for i := range xv {
xv[i] = float64(i%11 + 1)
}
out, err := c.MatVec(mustFloats(t, xv, n))
if err != nil {
t.Fatalf("MatVec: %v", err)
}
want := sparseCSCMatVecReference(c, xv)
got := out.RawFloats()
for i := range want {
if got[i] != want[i] {
t.Fatalf("element %d = %v, want %v", i, got[i], want[i])
}
}
}
// TestSparseCSCMatVecNonCanonicalFallsBack pins the canonicality
// check: a matrix above the split's budgets whose rows descend inside
// a column must answer the exact product anyway, through the serial
// column walk the split falls back to. The exported fields make such
// a matrix constructible by any caller, and the serial walk answered
// it before the split existed.
func TestSparseCSCMatVecNonCanonicalFallsBack(t *testing.T) {
// The same pinned policy as above: the guard needs the split
// reachable, and only then does the canonicality check hand the
// matrix to the fallback.
prev := engine.SetNumWorkers(4)
defer engine.SetNumWorkers(prev)
const n = 8192
c := &SparseCSC{Rows: n, Cols: n, ColStart: make([]int, n+1)}
for j := range n {
c.ColStart[j] = len(c.RowIdx)
if j == 100 {
// Column 100 stores its two entries in descending row
// order, and the two rows land in different blocks: the
// block walk's column skip reads the first entry's row as
// the column's lower bound, drops the pair from the lower
// block, and the early exit drops it from the upper one.
c.RowIdx = append(c.RowIdx, 5000, 100)
c.Values = append(c.Values, 1, 3)
continue
}
c.RowIdx = append(c.RowIdx, j)
c.Values = append(c.Values, float64(j%5+1))
}
c.ColStart[n] = len(c.RowIdx)
if cscMatVecBlocks(c.Rows, c.Cols, len(c.Values)) <= 1 {
t.Fatal("the matrix does not reach the split, the fallback would not be exercised")
}
if cscRowsAscending(c.ColStart, c.RowIdx) {
t.Fatal("the matrix reads as canonical, the fallback would not be exercised")
}
xv := make([]float64, n)
for i := range xv {
xv[i] = float64(i%9 + 1)
}
out, err := c.MatVec(mustFloats(t, xv, n))
if err != nil {
t.Fatalf("MatVec: %v", err)
}
want := sparseCSCMatVecReference(c, xv)
got := out.RawFloats()
for i := range want {
if got[i] != want[i] {
t.Fatalf("element %d = %v, want %v", i, got[i], want[i])
}
}
}