Files

193 lines
5.4 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
}