// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import "testing" // The products below are pinned against the scalar sum computed in this // file, not against another kernel: every output element receives its // addends over p ascending whatever panel, unroll or column band carries // it, so each walk must reproduce that order exactly. The samples are // integer-valued, so the comparison is exact. // TestComplexMatMulPanelWalks pins the complex product's four-row panel // and its single-row tail: five rows make the panel run once and the // tail once. func TestComplexMatMulPanelWalks(t *testing.T) { a := mustFromComplexes(t, []complex128{ 1, -2, 3, 4, -5, 6, 7, -8, 9, 10, }, 5, 2) b := mustFromComplexes(t, []complex128{ 1, 0, 2, 0, 1, -3, }, 2, 3) got, err := MatMul2D(a, b) if err != nil { t.Fatalf("MatMul complex 5x2 by 2x3: %v", err) } if sh := got.Shape(); len(sh) != 2 || sh[0] != 5 || sh[1] != 3 { t.Fatalf("MatMul complex shape: %v, want [5 3]", sh) } want := naiveMatMulCells(a.RawComplexes(), b.RawComplexes(), 5, 2, 3) for i := range want { if got.RawComplexes()[i] != want[i] { t.Fatalf("MatMul complex cell %d = %v, want %v", i, got.RawComplexes()[i], want[i]) } } // Four rows exactly: the panel with no remainder row. c := mustFromComplexes(t, []complex128{ 2, 1, 0, -3, 4, 5, 6, -7, }, 4, 2) d := mustFromComplexes(t, []complex128{ 1, 2, 3, -1, }, 2, 2) got4, err := MatMul2D(c, d) if err != nil { t.Fatalf("MatMul complex 4x2 by 2x2: %v", err) } want4 := naiveMatMulCells(c.RawComplexes(), d.RawComplexes(), 4, 2, 2) for i := range want4 { if got4.RawComplexes()[i] != want4[i] { t.Fatalf("MatMul complex panel cell %d = %v, want %v", i, got4.RawComplexes()[i], want4[i]) } } } // TestComplexMatVecRows pins the two-row complex unroll of the // matrix-vector product: five rows make the unroll run twice and the // single-row tail once. func TestComplexMatVecRows(t *testing.T) { m := mustFromComplexes(t, []complex128{ 1, 2, -3, 4, -5, 6, 7, 8, 9, -1, 2, 3, 4, 5, -6, }, 5, 3) v := mustFromComplexes(t, []complex128{2, -1, 3}, 3) got, err := MatMul2D(m, v) if err != nil { t.Fatalf("MatMul complex 5x3 by 3: %v", err) } if sh := got.Shape(); len(sh) != 1 || sh[0] != 5 { t.Fatalf("MatMul complex matrix-vector shape: %v, want [5]", sh) } want := naiveMatVec(m.RawComplexes(), v.RawComplexes(), 5, 3) for i := range want { if got.RawComplexes()[i] != want[i] { t.Fatalf("MatMul complex matrix-vector row %d = %v, want %v", i, got.RawComplexes()[i], want[i]) } } } // TestVecMatColumnWalks pins the vector-matrix product's column walks. // The float64 payload accumulates four columns at a time in registers // while the band is no wider than vecMatColBandMax and falls back to the // accumulator slots above it; the complex payload always walks the // slots. Every column sums its addends over p ascending. func TestVecMatColumnWalks(t *testing.T) { const k = 3 v := mustFromFloats(t, []float64{2, -3, 5}, k) // Four and six columns take the register walk (six with a // remainder), nine columns the slot walk. for _, cols := range []int{4, 6, 9} { vals := make([]float64, k*cols) for i := range vals { vals[i] = float64(i%7) - 3 + 0.5 } m := mustFromFloats(t, vals, k, cols) got, err := MatMul2D(v, m) if err != nil { t.Fatalf("MatMul vector by %dx%d: %v", k, cols, err) } if sh := got.Shape(); len(sh) != 1 || sh[0] != cols { t.Fatalf("MatMul vector by %dx%d shape: %v, want [%d]", k, cols, sh, cols) } want := naiveVecMat(v.RawFloats(), m.RawFloats(), k, cols) for j := range want { if got.RawFloats()[j] != want[j] { t.Fatalf("MatMul vector by %dx%d column %d = %v, want %v", k, cols, j, got.RawFloats()[j], want[j]) } } } // The complex payload takes the slot walk for every band. cv := mustFromComplexes(t, []complex128{2, -1, 3}, k) cvals := make([]complex128, k*4) for i := range cvals { cvals[i] = complex(float64(i%5)-2, float64(i%3)-1) } cm := mustFromComplexes(t, cvals, k, 4) cgot, err := MatMul2D(cv, cm) if err != nil { t.Fatalf("MatMul complex vector by %dx4: %v", k, err) } cwant := naiveVecMat(cv.RawComplexes(), cm.RawComplexes(), k, 4) for j := range cwant { if cgot.RawComplexes()[j] != cwant[j] { t.Fatalf("MatMul complex vector-matrix column %d = %v, want %v", j, cgot.RawComplexes()[j], cwant[j]) } } } // naiveMatMulCells returns the n×m product of an n×k matrix and a k×m // matrix, each cell summed over p ascending. func naiveMatMulCells[T complex128 | float64](a, b []T, n, k, m int) []T { out := make([]T, n*m) for i := range n { for j := range m { var s T for p := range k { s += a[i*k+p] * b[p*m+j] } out[i*m+j] = s } } return out } // naiveMatVec returns the n-vector an n×k matrix multiplies a k-vector // into, each row summed over p ascending. func naiveMatVec[T complex128 | float64](a, v []T, n, k int) []T { out := make([]T, n) for i := range n { var s T for p := range k { s += a[i*k+p] * v[p] } out[i] = s } return out } // naiveVecMat returns the m-vector a k-vector multiplies a k×m matrix // into, each column summed over p ascending. func naiveVecMat[T complex128 | float64](v, m []T, k, cols int) []T { out := make([]T, cols) for j := range cols { var s T for p := range k { s += v[p] * m[p*cols+j] } out[j] = s } return out }