124 lines
4.1 KiB
Go
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])
|
|
}
|
|
}
|
|
}
|