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
|
|||
|
|
}
|