// Copyright (c) 2026 Petr BalvĂ­n (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]) } } }