Files
tensor/signal/pooling2d_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

195 lines
5.2 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package signal
import "sourcedock.dev/petrbalvin/tensor/internal/core"
import (
"math"
"testing"
)
func TestConv2D(t *testing.T) {
// 1×1×3×3 input, 1×1×2×2 kernel (basic case).
in, _ := core.FromFloats([]float64{
1, 2, 3,
4, 5, 6,
7, 8, 9,
}, 1, 1, 3, 3)
ker, _ := core.FromFloats([]float64{
1, 0,
0, 1,
}, 1, 1, 2, 2)
got, err := Conv2D(in, ker, nil, 1, 0)
if err != nil {
t.Fatal(err)
}
// 2×2 output: conv at (oh, ow) = sum(in[oh..oh+2, ow..ow+2] * ker).
// (0,0) = 1*1 + 2*0 + 4*0 + 5*1 = 6
// (0,1) = 2*1 + 3*0 + 5*0 + 6*1 = 8
// (1,0) = 4*1 + 5*0 + 7*0 + 8*1 = 12
// (1,1) = 5*1 + 6*0 + 8*0 + 9*1 = 14
for i, w := range []float64{6, 8, 12, 14} {
v, _ := core.FloatAt(got, 0, 0, i/2, i%2)
if v != w {
t.Errorf("conv2d [%d]: got %v, want %v", i, v, w)
}
}
// With stride=2, output is 1×1.
got, err = Conv2D(in, ker, nil, 2, 0)
if err != nil {
t.Fatal(err)
}
if got.Shape()[2] != 1 || got.Shape()[3] != 1 {
t.Errorf("conv2d stride 2: shape = %v", got.Shape())
}
// With bias.
bias, _ := core.FromFloats([]float64{1}, 1)
got, err = Conv2D(in, ker, bias, 1, 0)
if err != nil {
t.Fatal(err)
}
v, _ := core.FloatAt(got, 0, 0, 0, 0)
if v != 7 {
t.Errorf("conv2d+bias [0,0,0,0]: got %v, want 7", v)
}
// Rank error.
badIn, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
if _, err := Conv2D(badIn, ker, nil, 1, 0); err == nil {
t.Error("conv2d: expected error for wrong input rank")
}
}
func TestMaxPool2D(t *testing.T) {
in, _ := core.FromFloats([]float64{
1, 3, 2, 4,
5, 6, 7, 8,
9, 2, 3, 1,
4, 0, 6, 2,
}, 1, 1, 4, 4)
got, err := MaxPool2D(in, 2, 2, 0)
if err != nil {
t.Fatal(err)
}
// 2×2 output. Each 2×2 window picks the max.
// TL: max(1,3,5,6) = 6; TR: max(2,4,7,8) = 8
// BL: max(9,2,4,0) = 9; BR: max(3,1,6,2) = 6
for i, w := range []float64{6, 8, 9, 6} {
v, _ := core.FloatAt(got, 0, 0, i/2, i%2)
if v != w {
t.Errorf("maxpool2d [%d]: got %v, want %v", i, v, w)
}
}
// With stride=1 and kernel=3: 2×2 output, overlapping.
got, err = MaxPool2D(in, 3, 1, 0)
if err != nil {
t.Fatal(err)
}
if got.Shape()[2] != 2 || got.Shape()[3] != 2 {
t.Errorf("maxpool2d 3x1: shape = %v", got.Shape())
}
// With padding=1, kernel=3, stride=1: output is (H+2*1-3)/1+1 = H.
got, err = MaxPool2D(in, 3, 1, 1)
if err != nil {
t.Fatal(err)
}
if got.Shape()[2] != 4 || got.Shape()[3] != 4 {
t.Errorf("maxpool2d 3x1+pad: shape = %v", got.Shape())
}
// Rank error.
bad, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
if _, err := MaxPool2D(bad, 2, 2, 0); err == nil {
t.Error("maxpool2d: expected error for wrong rank")
}
}
func TestAvgPool2D(t *testing.T) {
in, _ := core.FromFloats([]float64{
1, 2,
3, 4,
}, 1, 1, 2, 2)
got, err := AvgPool2D(in, 2, 2, 0, true)
if err != nil {
t.Fatal(err)
}
v, _ := core.FloatAt(got, 0, 0, 0, 0)
if v != 2.5 {
t.Errorf("avgpool2d: got %v, want 2.5", v)
}
// Without counting padding: same result when no padding.
got, err = AvgPool2D(in, 2, 2, 0, false)
if err != nil {
t.Fatal(err)
}
v, _ = core.FloatAt(got, 0, 0, 0, 0)
if v != 2.5 {
t.Errorf("avgpool2d (no pad): got %v, want 2.5", v)
}
}
func TestAvgPool2DCountIncludePad(t *testing.T) {
// 1×1×2×2 with kernel 3, stride 1, padding 1. The output is 2×2.
// Window (0,0) covers input rows/cols -1..1, so rows/cols 0..1 are
// real, the rest padded: 4 real cells of sum 10. includePad divides
// by 9 (10/9 ≈ 1.111), excludePad by 4 (10/4 = 2.5).
in, _ := core.FromFloats([]float64{
1, 2,
3, 4,
}, 1, 1, 2, 2)
inc, err := AvgPool2D(in, 3, 1, 1, true)
if err != nil {
t.Fatal(err)
}
exc, err := AvgPool2D(in, 3, 1, 1, false)
if err != nil {
t.Fatal(err)
}
for i := range 4 {
vi, _ := core.FloatAt(inc, 0, 0, i/2, i%2)
ve, _ := core.FloatAt(exc, 0, 0, i/2, i%2)
// The includePad divisor is the full kernel (9); the excludePad
// divisor is the number of real cells, which is 4 for all four
// windows of a 2×2 input with kernel 3 and padding 1.
if math.Abs(vi-10.0/9) > 1e-9 {
t.Errorf("includePad [%d]: got %v, want %v", i, vi, 10.0/9)
}
if math.Abs(ve-2.5) > 1e-9 {
t.Errorf("excludePad [%d]: got %v, want 2.5", i, ve)
}
}
}
func TestAdaptiveAvgPool2D(t *testing.T) {
// 1×1×4×4 gives 1×1×2×2.
in, _ := core.FromFloats([]float64{
1, 2, 3, 4,
5, 6, 7, 8,
9, 10, 11, 12,
13, 14, 15, 16,
}, 1, 1, 4, 4)
got, err := AdaptiveAvgPool2D(in, 2, 2)
if err != nil {
t.Fatal(err)
}
// Top-left 2x2: avg(1,2,5,6) = 3.5
v, _ := core.FloatAt(got, 0, 0, 0, 0)
if math.Abs(v-3.5) > 1e-9 {
t.Errorf("adaptive avgpool [0,0,0,0]: got %v, want 3.5", v)
}
// Top-right 2x2: avg(3,4,7,8) = 5.5
v, _ = core.FloatAt(got, 0, 0, 0, 1)
if math.Abs(v-5.5) > 1e-9 {
t.Errorf("adaptive avgpool [0,0,0,1]: got %v, want 5.5", v)
}
// Invalid output size.
if _, err := AdaptiveAvgPool2D(in, 0, 1); err == nil {
t.Error("adaptive avgpool: expected error for outputH=0")
}
// Rank error.
bad, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
if _, err := AdaptiveAvgPool2D(bad, 2, 2); err == nil {
t.Error("adaptive avgpool: expected error for wrong rank")
}
}