178 lines
6.7 KiB
Go
178 lines
6.7 KiB
Go
// 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.
|
|||
|
|
}
|