121 lines
3.7 KiB
Go
121 lines
3.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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])
|
|
}
|
|
}
|
|
})
|
|
}
|