199 lines
6.0 KiB
Go
199 lines
6.0 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package signal
|
||
|
||
import (
|
||
"math"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
"testing"
|
||
)
|
||
|
||
// anisoGrid samples f(x, y) = sin(2x)·sin(y) on a rows×cols grid over
|
||
// [0, Lx]×[0, Ly], x along columns, y along rows.
|
||
func anisoGrid(t *testing.T, rows, cols int, lx, ly float64) *core.Array {
|
||
t.Helper()
|
||
vals := make([]float64, rows*cols)
|
||
for r := range rows {
|
||
y := float64(r) * ly / float64(rows)
|
||
for c := range cols {
|
||
x := float64(c) * lx / float64(cols)
|
||
vals[r*cols+c] = math.Sin(2*x) * math.Sin(y)
|
||
}
|
||
}
|
||
a, err := core.FromFloats(vals, rows, cols)
|
||
if err != nil {
|
||
t.Fatalf("FromFloats: %v", err)
|
||
}
|
||
return a
|
||
}
|
||
|
||
// TestPoissonAnisotropicDomain pins the axis-length pairing: on a
|
||
// domain with Lx ≠ Ly the eigenvalues must pair (2πk/Lx)² for columns
|
||
// with (2πm/Ly)² for rows. The swapped pairing the code used to have
|
||
// inverts the aspect ratio.
|
||
func TestPoissonAnisotropicDomain(t *testing.T) {
|
||
const rows, cols = 24, 32
|
||
const lx, ly = 4 * math.Pi, 2 * math.Pi
|
||
f := anisoGrid(t, rows, cols, lx, ly)
|
||
u, err := SolvePoissonPeriodic(f, lx, ly)
|
||
if err != nil {
|
||
t.Fatalf("SolvePoissonPeriodic: %v", err)
|
||
}
|
||
// The analytic solution of −Δu = f on the periodic domain is
|
||
// u = f/λ: f = sin(2x)·sin(y) carries the modes 2πk/Lx = 2 (k = 4)
|
||
// and 2πm/Ly = 1 (m = 1), so λ = 2² + 1² = 5.
|
||
const lam = 5
|
||
maxErr := 0.0
|
||
for r := range rows {
|
||
for c := range cols {
|
||
want := f.FloatAt(r*cols+c) / lam
|
||
maxErr = max(maxErr, math.Abs(u.FloatAt(r*cols+c)-want))
|
||
}
|
||
}
|
||
if maxErr > 1e-10 {
|
||
t.Fatalf("max error %g, want below 1e-10 (aspect ratio mixed up)", maxErr)
|
||
}
|
||
}
|
||
|
||
// TestPoissonComplexRejects pins the dtype contract.
|
||
func TestPoissonComplexRejects(t *testing.T) {
|
||
a, err := core.FromComplexes([]complex128{1i, 2, 3, 4}, 2, 2)
|
||
if err != nil {
|
||
t.Fatalf("FromComplexes: %v", err)
|
||
}
|
||
if _, err := SolvePoissonPeriodic(a, 1, 1); err == nil {
|
||
t.Fatal("SolvePoissonPeriodic accepted complex input")
|
||
}
|
||
}
|
||
|
||
// TestPoolingComplexRejects pins the dtype guard on the pooling
|
||
// surface: complex input must be an error, not a FloatAt panic.
|
||
func TestPoolingComplexRejects(t *testing.T) {
|
||
mk := func(shape ...int) *core.Array {
|
||
n := 1
|
||
for _, d := range shape {
|
||
n *= d
|
||
}
|
||
vals := make([]complex128, n)
|
||
for i := range vals {
|
||
vals[i] = complex(float64(i), 1)
|
||
}
|
||
a, _ := core.FromComplexes(vals, shape...)
|
||
return a
|
||
}
|
||
in2 := mk(1, 1, 4, 4)
|
||
if _, err := MaxPool2D(in2, 2, 2, 0); err == nil {
|
||
t.Error("MaxPool2D accepted complex input")
|
||
}
|
||
if _, err := AvgPool2D(in2, 2, 2, 0, true); err == nil {
|
||
t.Error("AvgPool2D accepted complex input")
|
||
}
|
||
if _, err := AdaptiveMaxPool2D(in2, 2, 2); err == nil {
|
||
t.Error("AdaptiveMaxPool2D accepted complex input")
|
||
}
|
||
if _, err := AdaptiveAvgPool2D(in2, 2, 2); err == nil {
|
||
t.Error("AdaptiveAvgPool2D accepted complex input")
|
||
}
|
||
if _, err := GlobalAvgPool2D(in2); err == nil {
|
||
t.Error("GlobalAvgPool2D accepted complex input")
|
||
}
|
||
in1 := mk(1, 1, 8)
|
||
if _, err := MaxPool1D(in1, 2, 2, 0); err == nil {
|
||
t.Error("MaxPool1D accepted complex input")
|
||
}
|
||
if _, err := AdaptiveAvgPool1D(in1, 2); err == nil {
|
||
t.Error("AdaptiveAvgPool1D accepted complex input")
|
||
}
|
||
if _, err := GlobalMaxPool1D(in1); err == nil {
|
||
t.Error("GlobalMaxPool1D accepted complex input")
|
||
}
|
||
in3 := mk(1, 1, 4, 4, 4)
|
||
if _, err := MaxPool3D(in3, [3]int{2, 2, 2}, [3]int{2, 2, 2}, [3]int{0, 0, 0}); err == nil {
|
||
t.Error("MaxPool3D accepted complex input")
|
||
}
|
||
if _, err := AdaptiveAvgPool3D(in3, 2, 2, 2); err == nil {
|
||
t.Error("AdaptiveAvgPool3D accepted complex input")
|
||
}
|
||
if _, err := GlobalMaxPool3D(in3); err == nil {
|
||
t.Error("GlobalMaxPool3D accepted complex input")
|
||
}
|
||
}
|
||
|
||
// TestTransformsComplexRejects pins the dtype guard on Welch, Lomb-
|
||
// Scargle, DCT/DST and the stencils.
|
||
func TestTransformsComplexRejects(t *testing.T) {
|
||
v, _ := core.FromComplexes([]complex128{1 + 1i, 2, 3 + 2i, 4, 5, 6, 7, 8}, 8)
|
||
if _, _, err := WelchPSD(v, 1.0, 4, 1, "hann"); err == nil {
|
||
t.Error("WelchPSD accepted complex input")
|
||
}
|
||
ts, _ := core.FromFloats([]float64{0, 1, 2, 3}, 4)
|
||
if _, _, err := LombScargle(ts, v, 0.1, 1, 8); err == nil {
|
||
t.Error("LombScargle accepted complex values")
|
||
}
|
||
if _, err := DCT(v, 2); err == nil {
|
||
t.Error("DCT accepted complex input")
|
||
}
|
||
if _, err := DST(v, 1); err == nil {
|
||
t.Error("DST accepted complex input")
|
||
}
|
||
if _, err := Gradient1D(v, 0.5); err == nil {
|
||
t.Error("Gradient1D accepted complex input")
|
||
}
|
||
v2, _ := core.FromComplexes([]complex128{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3)
|
||
if _, err := Laplacian(v2, 1, 1); err == nil {
|
||
t.Error("Laplacian accepted complex input")
|
||
}
|
||
}
|
||
|
||
// TestFFTDoesNotMutateComplexInput pins the immutability contract: the
|
||
// 1-D transforms must not overwrite the caller's complex payload, which
|
||
// ComplexValues aliasing used to let happen.
|
||
func TestFFTDoesNotMutateComplexInput(t *testing.T) {
|
||
vals := []complex128{1 + 2i, 3 - 1i, -0.5 + 0.5i, 2}
|
||
a, err := core.FromComplexes(vals, 4)
|
||
if err != nil {
|
||
t.Fatalf("FromComplexes: %v", err)
|
||
}
|
||
if _, err := FFT(a); err != nil {
|
||
t.Fatalf("FFT: %v", err)
|
||
}
|
||
for i, want := range vals {
|
||
if a.ComplexAt(i) != want {
|
||
t.Fatalf("FFT mutated input[%d] = %v, want %v", i, a.ComplexAt(i), want)
|
||
}
|
||
}
|
||
if _, err := IFFT(a); err != nil {
|
||
t.Fatalf("IFFT: %v", err)
|
||
}
|
||
for i, want := range vals {
|
||
if a.ComplexAt(i) != want {
|
||
t.Fatalf("IFFT mutated input[%d] = %v, want %v", i, a.ComplexAt(i), want)
|
||
}
|
||
}
|
||
// Double transform stays correct: FFT twice then IFFT twice is a
|
||
// round trip, only meaningful when neither call clobbers the input.
|
||
f1, err := FFT(a)
|
||
if err != nil {
|
||
t.Fatalf("FFT: %v", err)
|
||
}
|
||
f2, err := FFT(f1)
|
||
if err != nil {
|
||
t.Fatalf("FFT: %v", err)
|
||
}
|
||
i1, err := IFFT(f2)
|
||
if err != nil {
|
||
t.Fatalf("IFFT: %v", err)
|
||
}
|
||
i2, err := IFFT(i1)
|
||
if err != nil {
|
||
t.Fatalf("IFFT: %v", err)
|
||
}
|
||
for i, want := range vals {
|
||
got := i2.ComplexAt(i)
|
||
if math.Abs(real(got)-real(want)) > 1e-12 || math.Abs(imag(got)-imag(want)) > 1e-12 {
|
||
t.Fatalf("round trip[%d] = %v, want %v", i, got, want)
|
||
}
|
||
}
|
||
}
|