Files
tensor/signal/nan_propagation_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

178 lines
6.7 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"
"strings"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// NaN-propagation pins: max pools that silently dropped NaN, empty
// windows behind oversized padding, estimators that published NaN
// without an error, and the odd-length inverse real FFT with no test
// at all.
// TestMaxPoolNaNPropagates: v > best reads false against NaN, so the
// max pools answered the largest finite neighbour and hid the NaN; a
// NaN in a window now wins the comparison and propagates.
func TestMaxPoolNaNPropagates(t *testing.T) {
lane := mustFromFloats(t, []float64{1, 3, math.NaN(), 2}, 1, 1, 4)
got, err := MaxPool1D(lane, 2, 2, 0)
if err != nil {
t.Fatal(err)
}
if v0, _ := core.FloatAt(got, 0, 0, 0); v0 != 3 {
t.Errorf("MaxPool1D finite window: got %v, want 3", v0)
}
if v1, _ := core.FloatAt(got, 0, 0, 1); !math.IsNaN(v1) {
t.Errorf("MaxPool1D window with NaN: got %v, want NaN", v1)
}
sq := mustFromFloats(t, []float64{
math.NaN(), 2, 3,
4, 5, 6,
7, 8, 9,
}, 1, 1, 3, 3)
got2, err := MaxPool2D(sq, 2, 1, 0)
if err != nil {
t.Fatal(err)
}
// The top-left window holds the NaN; the bottom-right one does not.
if v, _ := core.FloatAt(got2, 0, 0, 0, 0); !math.IsNaN(v) {
t.Errorf("MaxPool2D window with NaN: got %v, want NaN", v)
}
if v, _ := core.FloatAt(got2, 0, 0, 1, 1); v != 9 {
t.Errorf("MaxPool2D finite window: got %v, want 9", v)
}
// The adaptive and global variants share the walk.
got3, err := AdaptiveMaxPool2D(sq, 2, 2)
if err != nil {
t.Fatal(err)
}
if v, _ := core.FloatAt(got3, 0, 0, 0, 0); !math.IsNaN(v) {
t.Errorf("AdaptiveMaxPool2D window with NaN: got %v, want NaN", v)
}
got4, err := GlobalMaxPool1D(mustFromFloats(t, []float64{1, 2, 3, math.NaN()}, 1, 1, 4))
if err != nil {
t.Fatal(err)
}
if v, _ := core.FloatAt(got4, 0, 0, 0); !math.IsNaN(v) {
t.Errorf("GlobalMaxPool1D with NaN: got %v, want NaN", v)
}
}
// TestPoolPaddingBelowKernelRefused: padding at or above the kernel
// leaves output windows entirely inside the padding, where a max
// answered −Inf and an average 0; such configurations are refused.
func TestPoolPaddingBelowKernelRefused(t *testing.T) {
lane := mustFromFloats(t, []float64{1, 2, 3, 4}, 1, 1, 4)
if _, err := MaxPool1D(lane, 1, 1, 2); err == nil || !strings.Contains(err.Error(), "below the kernel") {
t.Errorf("MaxPool1D padding 2 kernel 1: err = %v", err)
}
if _, err := AvgPool1D(lane, 1, 1, 2, true); err == nil || !strings.Contains(err.Error(), "below the kernel") {
t.Errorf("AvgPool1D padding 2 kernel 1: err = %v", err)
}
// padding == kernel empties the first window too, so it is
// refused all the same.
sq := mustFromFloats(t, make([]float64, 9), 1, 1, 3, 3)
if _, err := MaxPool2D(sq, 2, 1, 2); err == nil || !strings.Contains(err.Error(), "below the kernel") {
t.Errorf("MaxPool2D padding equal to the kernel: err = %v", err)
}
cube := mustFromFloats(t, make([]float64, 8), 1, 1, 2, 2, 2)
if _, err := MaxPool3D(cube, [3]int{2, 1, 1}, [3]int{1, 1, 1}, [3]int{2, 0, 0}); err == nil || !strings.Contains(err.Error(), "below the kernel") {
t.Errorf("MaxPool3D oversized padding: err = %v", err)
}
// Valid configurations are untouched: padding below the kernel
// still pools, and every window keeps at least one real sample.
if _, err := MaxPool1D(lane, 2, 2, 1); err != nil {
t.Errorf("MaxPool1D padding 1 kernel 2 refused: %v", err)
}
if _, err := MaxPool2D(sq, 3, 1, 1); err != nil {
t.Errorf("MaxPool2D padding 1 kernel 3 refused: %v", err)
}
}
// TestCorrelatePeriodogramNonFiniteRefused: a non-finite sample drove
// the estimators' normalisers NaN and published NaN estimates with no
// error; the inputs are refused up front now.
func TestCorrelatePeriodogramNonFiniteRefused(t *testing.T) {
clean := mustFromFloats(t, []float64{1, 2, 3, 4, 5}, 5)
nan := mustFromFloats(t, []float64{1, 2, math.NaN(), 4, 5}, 5)
inf := mustFromFloats(t, []float64{1, 2, math.Inf(1), 4, 5}, 5)
if _, err := Autocorrelate(nan, 2); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Errorf("Autocorrelate with NaN: err = %v", err)
}
if _, err := PartialAutocorrelate(nan, 2); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Errorf("PartialAutocorrelate with NaN: err = %v", err)
}
if _, err := CrossCorrelate(clean, inf); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Errorf("CrossCorrelate with Inf: err = %v", err)
}
times := mustFromFloats(t, []float64{0, 1, 2, 3, math.NaN()}, 5)
vals := mustFromFloats(t, []float64{1, -1, 1, -1, 1}, 5)
if _, _, err := LombScargle(times, vals, 0.1, 0.4, 8); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Errorf("LombScargle with NaN times: err = %v", err)
}
timesClean := mustFromFloats(t, []float64{0, 1, 2, 3, 4}, 5)
valsBad := mustFromFloats(t, []float64{1, -1, math.NaN(), -1, 1}, 5)
if _, _, err := LombScargle(timesClean, valsBad, 0.1, 0.4, 8); err == nil || !strings.Contains(err.Error(), "non-finite") {
t.Errorf("LombScargle with NaN values: err = %v", err)
}
// Clean inputs still estimate.
if _, err := Autocorrelate(clean, 2); err != nil {
t.Errorf("Autocorrelate clean: %v", err)
}
}
// TestIRFFTOddLengthRoundTrip: the odd-length inverse real FFT had no
// coverage; the spectrum must match the naive DFT and the round trip
// must restore the series.
func TestIRFFTOddLengthRoundTrip(t *testing.T) {
for _, n := range []int{5, 7, 9} {
vals := make([]float64, n)
g := core.NewGenerator(int64(n))
for i := range vals {
f, _ := core.Floats(g, 1)
v, _ := core.FloatAt(f, 0)
vals[i] = v*10 - 5
}
src := mustFromFloats(t, vals, n)
spec, err := RFFT(src)
if err != nil {
t.Fatalf("RFFT(%d): %v", n, err)
}
complexVals := make([]complex128, n)
for i := range vals {
complexVals[i] = complex(vals[i], 0)
}
naive := naiveDFT(complexVals, -1)
half := n/2 + 1
for k := range half {
got, _ := core.ComplexAt(spec, k)
if math.Abs(real(got)-real(naive[k])) > 1e-9 || math.Abs(imag(got)-imag(naive[k])) > 1e-9 {
t.Fatalf("RFFT(%d)[%d]: got %v, want %v", n, k, got, naive[k])
}
}
back, err := IRFFT(spec, n)
if err != nil {
t.Fatalf("IRFFT(%d): %v", n, err)
}
if back.Len() != n {
t.Fatalf("IRFFT(%d) length %d, want %d", n, back.Len(), n)
}
for i := range n {
got, _ := core.FloatAt(back, i)
if math.Abs(got-vals[i]) > 1e-9 {
t.Fatalf("IRFFT(%d)[%d]: got %.14g, want %.14g", n, i, got, vals[i])
}
}
}
// The bin count of an odd spectrum is (n+1)/2, so the documented
// default n = 2·(half−1) cannot recover an odd length: odd
// round trips must pass n explicitly, which is exactly the path
// checked above.
}