// Copyright (c) 2026 Petr BalvĂ­n (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) } } } } } } } }