Files
tensor/signal/poisson_complex_pins_test.go
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

199 lines
6.0 KiB
Go
Raw Permalink 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 (
"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)
}
}
}