Files
tensor/internal/core/mat_kernel_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

193 lines
5.4 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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
}