303 lines
9.0 KiB
Go
303 lines
9.0 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package signal
|
||
|
|
|
||
|
|
import (
|
||
|
|
"testing"
|
||
|
|
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
|
|
)
|
||
|
|
|
||
|
|
// The fused interiors, the phase-lattice gather and the span tables
|
||
|
|
// are reach-and-order refactors of the per-tap walks: every output
|
||
|
|
// element must keep the exact accumulated bits. The fixtures are
|
||
|
|
// small integers, so every product and sum is exact in float64 and the
|
||
|
|
// comparison against the references below is exact regardless of
|
||
|
|
// summation order; the references walk the documented shapes and the
|
||
|
|
// (cIn, k...) tap order.
|
||
|
|
|
||
|
|
func pinInts(n int, seed int) []float64 {
|
||
|
|
v := make([]float64, n)
|
||
|
|
x := seed
|
||
|
|
for i := range n {
|
||
|
|
x = (x*1103515245 + 12345) & 0x7fffffff
|
||
|
|
v[i] = float64(x%7 - 3)
|
||
|
|
}
|
||
|
|
return v
|
||
|
|
}
|
||
|
|
|
||
|
|
// refConv1D is the documented forward pass: out[n][oc][ol] is the bias
|
||
|
|
// plus the (ic, kL)-ordered tap sum of in[n][ic][ol*stride + kL -
|
||
|
|
// padding] against ker[oc][ic][kL].
|
||
|
|
func refConv1D(in []float64, n, cIn, lIn int, ker []float64, cOut, kL int, bias []float64, stride, padding int) []float64 {
|
||
|
|
lOut := (lIn+2*padding-kL)/stride + 1
|
||
|
|
out := make([]float64, n*cOut*lOut)
|
||
|
|
for b := range n {
|
||
|
|
for oc := range cOut {
|
||
|
|
for ol := range lOut {
|
||
|
|
sum := 0.0
|
||
|
|
for ic := range cIn {
|
||
|
|
for kl := range kL {
|
||
|
|
il := ol*stride + kl - padding
|
||
|
|
if il < 0 || il >= lIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
sum += in[b*cIn*lIn+ic*lIn+il] * ker[oc*cIn*kL+ic*kL+kl]
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if bias != nil {
|
||
|
|
sum += bias[oc]
|
||
|
|
}
|
||
|
|
out[b*cOut*lOut+oc*lOut+ol] = sum
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return out
|
||
|
|
}
|
||
|
|
|
||
|
|
// refConv3D is the documented NCDHW forward pass in (ic, kD, kH, kW)
|
||
|
|
// tap order.
|
||
|
|
func refConv3D(in []float64, n, cIn, dIn, hIn, wIn int, ker []float64, cOut, kD, kH, kW int, bias []float64, stride, padD, padH, padW int) []float64 {
|
||
|
|
dOut := (dIn+2*padD-kD)/stride + 1
|
||
|
|
hOut := (hIn+2*padH-kH)/stride + 1
|
||
|
|
wOut := (wIn+2*padW-kW)/stride + 1
|
||
|
|
out := make([]float64, n*cOut*dOut*hOut*wOut)
|
||
|
|
plane := dIn * hIn * wIn
|
||
|
|
for b := range n {
|
||
|
|
for oc := range cOut {
|
||
|
|
for od := range dOut {
|
||
|
|
for oh := range hOut {
|
||
|
|
for ow := range wOut {
|
||
|
|
sum := 0.0
|
||
|
|
for ic := range cIn {
|
||
|
|
for kd := range kD {
|
||
|
|
id := od*stride + kd - padD
|
||
|
|
if id < 0 || id >= dIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
for kh := range kH {
|
||
|
|
ih := oh*stride + kh - padH
|
||
|
|
if ih < 0 || ih >= hIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
for kw := range kW {
|
||
|
|
iw := ow*stride + kw - padW
|
||
|
|
if iw < 0 || iw >= wIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
sum += in[b*cIn*plane+ic*plane+id*hIn*wIn+ih*wIn+iw] *
|
||
|
|
ker[oc*cIn*kD*kH*kW+ic*kD*kH*kW+kd*kH*kW+kh*kW+kw]
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if bias != nil {
|
||
|
|
sum += bias[oc]
|
||
|
|
}
|
||
|
|
out[b*cOut*dOut*hOut*wOut+oc*dOut*hOut*wOut+od*hOut*wOut+oh*wOut+ow] = sum
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return out
|
||
|
|
}
|
||
|
|
|
||
|
|
// refConvTranspose2D is the documented gather: out[n][oc][oh][ow] is
|
||
|
|
// the bias plus the (ic, kH, kW) tap sum over the pairs whose reversed
|
||
|
|
// column lands in the input, reading in[n][ic][ih][iw] against
|
||
|
|
// ker[ic][oc][kH][kW].
|
||
|
|
func refConvTranspose2D(in []float64, n, cIn, hIn, wIn int, ker []float64, cOut, kH, kW int, bias []float64, stride, padding int) []float64 {
|
||
|
|
hOut := (hIn-1)*stride - 2*padding + kH
|
||
|
|
wOut := (wIn-1)*stride - 2*padding + kW
|
||
|
|
out := make([]float64, n*cOut*hOut*wOut)
|
||
|
|
plane := hIn * wIn
|
||
|
|
for b := range n {
|
||
|
|
for oc := range cOut {
|
||
|
|
for oh := range hOut {
|
||
|
|
for ow := range wOut {
|
||
|
|
sum := 0.0
|
||
|
|
for ic := range cIn {
|
||
|
|
for kh := range kH {
|
||
|
|
hn := oh + padding - kh
|
||
|
|
if hn%stride != 0 {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
ih := hn / stride
|
||
|
|
if ih < 0 || ih >= hIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
for kw := range kW {
|
||
|
|
wn := ow + padding - kw
|
||
|
|
if wn%stride != 0 {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
iw := wn / stride
|
||
|
|
if iw < 0 || iw >= wIn {
|
||
|
|
continue
|
||
|
|
}
|
||
|
|
sum += in[b*cIn*plane+ic*plane+ih*wIn+iw] *
|
||
|
|
ker[ic*cOut*kH*kW+oc*kH*kW+kh*kW+kw]
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
if bias != nil {
|
||
|
|
sum += bias[oc]
|
||
|
|
}
|
||
|
|
out[b*cOut*hOut*wOut+oc*hOut*wOut+oh*wOut+ow] = sum
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return out
|
||
|
|
}
|
||
|
|
|
||
|
|
func pinConvEqual(t *testing.T, name string, got, want []float64) {
|
||
|
|
t.Helper()
|
||
|
|
if len(got) != len(want) {
|
||
|
|
t.Fatalf("%s: length %d, want %d", name, len(got), len(want))
|
||
|
|
}
|
||
|
|
for i := range want {
|
||
|
|
if got[i] != want[i] {
|
||
|
|
t.Fatalf("%s[%d] = %v, want %v", name, i, got[i], want[i])
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestConv1DFusedInteriorMatchesReference(t *testing.T) {
|
||
|
|
cases := []struct {
|
||
|
|
name string
|
||
|
|
n, cIn, cOut, lIn, kL int
|
||
|
|
stride, padding int
|
||
|
|
}{
|
||
|
|
{"fused3-long", 2, 3, 2, 3000, 3, 1, 1},
|
||
|
|
{"fused7", 1, 2, 3, 200, 7, 1, 3},
|
||
|
|
{"fused2-verylong", 1, 1, 1, 5000, 2, 1, 0},
|
||
|
|
{"stride2", 2, 2, 2, 100, 5, 2, 2},
|
||
|
|
{"span-extreme", 1, 1, 1, 1, 8, 1, 5},
|
||
|
|
}
|
||
|
|
for _, tc := range cases {
|
||
|
|
in := pinInts(tc.n*tc.cIn*tc.lIn, 1)
|
||
|
|
ker := pinInts(tc.cOut*tc.cIn*tc.kL, 2)
|
||
|
|
bias := pinInts(tc.cOut, 3)
|
||
|
|
ain, err := core.FromFloats(in, tc.n, tc.cIn, tc.lIn)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("%s input: %v", tc.name, err)
|
||
|
|
}
|
||
|
|
aker, err := core.FromFloats(ker, tc.cOut, tc.cIn, tc.kL)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("%s kernel: %v", tc.name, err)
|
||
|
|
}
|
||
|
|
abias, err := core.FromFloats(bias, tc.cOut)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("%s bias: %v", tc.name, err)
|
||
|
|
}
|
||
|
|
out, oerr := Conv1D(ain, aker, abias, tc.stride, tc.padding, 1)
|
||
|
|
if oerr != nil {
|
||
|
|
t.Fatalf("%s: %v", tc.name, oerr)
|
||
|
|
}
|
||
|
|
want := refConv1D(in, tc.n, tc.cIn, tc.lIn, ker, tc.cOut, tc.kL, bias, tc.stride, tc.padding)
|
||
|
|
pinConvEqual(t, tc.name, out.RawFloats(), want)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestConv3DFusedInteriorMatchesReference(t *testing.T) {
|
||
|
|
cases := []struct {
|
||
|
|
name string
|
||
|
|
n, cIn, cOut int
|
||
|
|
dIn, hIn, wIn int
|
||
|
|
kD, kH, kW int
|
||
|
|
stride int
|
||
|
|
padD, padH, padW int
|
||
|
|
}{
|
||
|
|
{"fused3", 2, 2, 3, 6, 7, 8, 3, 3, 3, 1, 1, 1, 1},
|
||
|
|
{"mixed-k", 1, 2, 2, 4, 5, 6, 2, 3, 4, 1, 0, 1, 1},
|
||
|
|
{"fused7", 1, 1, 1, 8, 8, 8, 7, 7, 7, 1, 3, 3, 3},
|
||
|
|
}
|
||
|
|
for _, tc := range cases {
|
||
|
|
in := pinInts(tc.n*tc.cIn*tc.dIn*tc.hIn*tc.wIn, 11)
|
||
|
|
ker := pinInts(tc.cOut*tc.cIn*tc.kD*tc.kH*tc.kW, 12)
|
||
|
|
ain, err := core.FromFloats(in, tc.n, tc.cIn, tc.dIn, tc.hIn, tc.wIn)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("%s input: %v", tc.name, err)
|
||
|
|
}
|
||
|
|
aker, err := core.FromFloats(ker, tc.cOut, tc.cIn, tc.kD, tc.kH, tc.kW)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("%s kernel: %v", tc.name, err)
|
||
|
|
}
|
||
|
|
out, oerr := Conv3D(ain, aker, nil, tc.stride, [3]int{tc.padD, tc.padH, tc.padW}, [3]int{1, 1, 1})
|
||
|
|
if oerr != nil {
|
||
|
|
t.Fatalf("%s: %v", tc.name, oerr)
|
||
|
|
}
|
||
|
|
want := refConv3D(in, tc.n, tc.cIn, tc.dIn, tc.hIn, tc.wIn, ker, tc.cOut, tc.kD, tc.kH, tc.kW, nil, tc.stride, tc.padD, tc.padH, tc.padW)
|
||
|
|
pinConvEqual(t, tc.name, out.RawFloats(), want)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestConvTranspose2DPhaseLatticeMatchesReference(t *testing.T) {
|
||
|
|
cases := []struct {
|
||
|
|
name string
|
||
|
|
n, cIn, cOut, hIn, wIn int
|
||
|
|
kH, kW int
|
||
|
|
stride, padding int
|
||
|
|
}{
|
||
|
|
{"lattice2", 2, 2, 3, 5, 6, 3, 3, 2, 1},
|
||
|
|
{"lattice3", 1, 2, 2, 4, 5, 2, 2, 3, 2},
|
||
|
|
{"unit-stride", 2, 3, 2, 6, 6, 3, 3, 1, 1},
|
||
|
|
{"tall-k", 1, 1, 2, 3, 3, 5, 4, 2, 0},
|
||
|
|
}
|
||
|
|
for _, tc := range cases {
|
||
|
|
in := pinInts(tc.n*tc.cIn*tc.hIn*tc.wIn, 21)
|
||
|
|
ker := pinInts(tc.cIn*tc.cOut*tc.kH*tc.kW, 22)
|
||
|
|
bias := pinInts(tc.cOut, 23)
|
||
|
|
ain, err := core.FromFloats(in, tc.n, tc.cIn, tc.hIn, tc.wIn)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("%s input: %v", tc.name, err)
|
||
|
|
}
|
||
|
|
aker, err := core.FromFloats(ker, tc.cIn, tc.cOut, tc.kH, tc.kW)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("%s kernel: %v", tc.name, err)
|
||
|
|
}
|
||
|
|
abias, err := core.FromFloats(bias, tc.cOut)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("%s bias: %v", tc.name, err)
|
||
|
|
}
|
||
|
|
out, oerr := ConvTranspose2D(ain, aker, abias, tc.stride, tc.padding)
|
||
|
|
if oerr != nil {
|
||
|
|
t.Fatalf("%s: %v", tc.name, oerr)
|
||
|
|
}
|
||
|
|
want := refConvTranspose2D(in, tc.n, tc.cIn, tc.hIn, tc.wIn, ker, tc.cOut, tc.kH, tc.kW, bias, tc.stride, tc.padding)
|
||
|
|
pinConvEqual(t, tc.name, out.RawFloats(), want)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestConvRowSpansFusedWindowAvoidsInvalidTaps(t *testing.T) {
|
||
|
|
// The fused interior takes its window bounds from spans[0].lo and
|
||
|
|
// spans[kW-1].hi, so an invalid tap must leave an empty window and
|
||
|
|
// the window must never cover a tap that reaches nothing: those
|
||
|
|
// are the two states the per-tap walk skips by construction.
|
||
|
|
for kW := 2; kW <= 8; kW++ {
|
||
|
|
for padW := 0; padW <= kW+1; padW++ {
|
||
|
|
for _, stride := range []int{1, 2, 3} {
|
||
|
|
for wIn := 1; wIn <= 9; wIn++ {
|
||
|
|
for wOut := 1; wOut <= 12; wOut++ {
|
||
|
|
spans := convRowSpans(kW, padW, stride, wIn, wOut)
|
||
|
|
fa, fb := spans[0].lo, spans[kW-1].hi
|
||
|
|
for kw := range kW {
|
||
|
|
if !spans[kw].valid && spans[kw].hi > spans[kw].lo {
|
||
|
|
t.Fatalf("kW=%d padW=%d stride=%d wIn=%d wOut=%d: tap %d invalid with window [%d, %d)",
|
||
|
|
kW, padW, stride, wIn, wOut, kw, spans[kw].lo, spans[kw].hi)
|
||
|
|
}
|
||
|
|
if fb > fa && !spans[kw].valid {
|
||
|
|
t.Fatalf("kW=%d padW=%d stride=%d wIn=%d wOut=%d: fused window [%d, %d) covers invalid tap %d",
|
||
|
|
kW, padW, stride, wIn, wOut, fa, fb, kw)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|