193 lines
5.4 KiB
Go
193 lines
5.4 KiB
Go
// 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
|
||
}
|