feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+302
View File
@@ -0,0 +1,302 @@
// 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)
}
}
}
}
}
}
}
}