Files
tensor/signal/conv_transpose_test.go
T

121 lines
3.7 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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])
}
}
})
}