// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package signal import ( "math" "testing" ) // convTranspose2DReference computes the transpose convolution straight // from its scatter definition, each output element summing the input // positions that map onto it in the (input channel, kernel row, // kernel column) order, the bias added last. It is the test's own // arithmetic, not a second kernel: the loops mirror the definition // the documentation states. func convTranspose2DReference(in []float64, k []float64, bias []float64, n, cIn, hIn, wIn, cOut, kH, kW, 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) 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 { ih := oh + padding - kh if ih < 0 || ih%stride != 0 { continue } ih /= stride if ih >= hIn { continue } for kw := range kW { iw := ow + padding - kw if iw < 0 || iw%stride != 0 { continue } iw /= stride if iw >= wIn { continue } sum += in[((b*cIn+ic)*hIn+ih)*wIn+iw] * k[((ic*cOut+oc)*kH+kh)*kW+kw] } } } if bias != nil { sum += bias[oc] } out[((b*cOut+oc)*hOut+oh)*wOut+ow] = sum } } } } return out } // TestConvTranspose2DChannelBlocking pins the channel-blocked kernel // against the scatter definition on shapes the single-channel test // cannot reach: several input and output channels (so the kernel's // channel stride matters), a stride above one, a bias, and an output // wider than one channel block, which leaves a partial block at the // end. func TestConvTranspose2DChannelBlocking(t *testing.T) { seed := 1 rand := func() float64 { seed = (7*seed + 3) % 97 return float64(seed)/97 - 0.5 } fill := func(n int) []float64 { v := make([]float64, n) for i := range v { v[i] = rand() } return v } t.Run("several channels, stride two, bias", func(t *testing.T) { const n, cIn, cOut, k = 2, 2, 3, 2 const hIn, wIn, stride, padding = 3, 3, 2, 1 in := fill(n * cIn * hIn * wIn) kk := fill(cIn * cOut * k * k) bias := fill(cOut) got, err := ConvTranspose2D(mustFromFloats(t, in, n, cIn, hIn, wIn), mustFromFloats(t, kk, cIn, cOut, k, k), mustFromFloats(t, bias, cOut), stride, padding) if err != nil { t.Fatalf("ConvTranspose2D: %v", err) } want := convTranspose2DReference(in, kk, bias, n, cIn, hIn, wIn, cOut, k, k, stride, padding) outF := got.RawFloats() for i := range want { if math.Float64bits(outF[i]) != math.Float64bits(want[i]) { t.Fatalf("element %d = %v, want %v", i, outF[i], want[i]) } } }) t.Run("partial channel block", func(t *testing.T) { // An output row of 1025 caps one block at three channels, so // four output channels run as a full block and a partial one. const n, cIn, cOut, k = 1, 1, 4, 3 const hIn, wIn, stride, padding = 1025, 1025, 1, 1 in := fill(n * cIn * hIn * wIn) kk := fill(cIn * cOut * k * k) bias := fill(cOut) got, err := ConvTranspose2D(mustFromFloats(t, in, n, cIn, hIn, wIn), mustFromFloats(t, kk, cIn, cOut, k, k), mustFromFloats(t, bias, cOut), stride, padding) if err != nil { t.Fatalf("ConvTranspose2D: %v", err) } want := convTranspose2DReference(in, kk, bias, n, cIn, hIn, wIn, cOut, k, k, stride, padding) outF := got.RawFloats() for i := range want { if math.Float64bits(outF[i]) != math.Float64bits(want[i]) { t.Fatalf("element %d = %v, want %v", i, outF[i], want[i]) } } }) }