feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,120 @@
|
||||
// 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])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user