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

268 lines
7.9 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"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// TestHaarReconstruction pins exactness: IDWT(DWT(x)) = x, energy
// conserved, across level counts.
func TestHaarReconstruction(t *testing.T) {
g := core.NewGenerator(3)
vals := make([]float64, 64)
for i := range vals {
vals[i] = g.NormalUnit()
}
x, _ := core.FromFloats(vals, 64)
for _, levels := range []int{1, 3, 6} {
c, err := DWT(x, levels)
if err != nil {
t.Fatalf("DWT(%d): %v", levels, err)
}
eIn, eC := 0.0, 0.0
for i := range 64 {
eIn += x.FloatAt(i) * x.FloatAt(i)
eC += c.FloatAt(i) * c.FloatAt(i)
}
if math.Abs(eIn-eC) > 1e-10*eIn {
t.Fatalf("levels %d: energy %g vs %g", levels, eIn, eC)
}
back, err := IDWT(c, levels)
if err != nil {
t.Fatalf("IDWT(%d): %v", levels, err)
}
for i := range 64 {
if math.Abs(back.FloatAt(i)-x.FloatAt(i)) > 1e-12 {
t.Fatalf("levels %d: reconstruction off at %d", levels, i)
}
}
}
}
// TestHaarStepSignal pins the sparsity promise: a piecewise-constant
// signal has wavelet coefficients concentrated at the jumps.
func TestHaarStepSignal(t *testing.T) {
// The jump sits at index 5: off every dyadic boundary, because a
// step landing exactly on one is invisible to Haar at every level.
vals := make([]float64, 16)
for i := 5; i < 16; i++ {
vals[i] = 1
}
x, _ := core.FromFloats(vals, 16)
c, err := DWT(x, 1)
if err != nil {
t.Fatalf("DWT: %v", err)
}
// Only detail 2 (the pair 4/5) carries the jump.
for i := range 8 {
if i != 2 && math.Abs(c.FloatAt(8+i)) > 1e-12 {
t.Fatalf("detail %d = %g, only the jump pair should fire", i, c.FloatAt(8+i))
}
}
if math.Abs(math.Abs(c.FloatAt(10))-1/math.Sqrt2) > 1e-12 {
t.Fatalf("jump detail = %g, want ±1/√2", c.FloatAt(10))
}
}
// TestCWTMorletRidge pins the time-frequency map: a pure sinusoid's
// Morlet transform peaks at the scale carrying its frequency.
func TestCWTMorletRidge(t *testing.T) {
const (
n = 1024
dt = 0.01
freq = 5.0
omega = 5.0 // Morlet ω₀
)
vals := make([]float64, n)
for i := range n {
vals[i] = math.Sin(2 * math.Pi * freq * dt * float64(i))
}
x, _ := core.FromFloats(vals, n)
scales := make([]float64, 30)
for i := range scales {
scales[i] = 0.01 + 0.01*float64(i)
}
w, err := CWT(x, Morlet, scales, dt)
if err != nil {
t.Fatalf("CWT: %v", err)
}
// Ridge: the scale maximising the mean magnitude. The plain peak
// estimate of the Morlet centre frequency, f ≈ ω₀/(2πa), is within
// the tolerance the test uses.
best := 0.0
bestMag := -1.0
for k := range scales {
m := 0.0
for i := range n {
z := w.ComplexAt(k*n + i)
m += math.Hypot(real(z), imag(z))
}
m /= float64(n)
if m > bestMag {
bestMag, best = m, scales[k]
}
}
want := omega / (2 * math.Pi * freq)
if math.Abs(best-want) > 0.5*want {
t.Fatalf("ridge scale %.4f, expected around %.4f", best, want)
}
}
// TestCWTMexicanHatBump pins the zero-mean bump detector: the CWT of a
// Gaussian bump peaks at the scale matching its width, and a constant
// signal transforms to zero (the wavelet has no DC).
func TestCWTMexicanHatBump(t *testing.T) {
const n = 512
vals := make([]float64, n)
for i := range n {
t0 := (float64(i) - n/2) * 0.05
vals[i] = math.Exp(-t0 * t0 / 2)
}
x, _ := core.FromFloats(vals, n)
scales := []float64{0.5, 1.0, 2.0, 4.0, 8.0}
w, err := CWT(x, MexicanHat, scales, 0.05)
if err != nil {
t.Fatalf("CWT: %v", err)
}
mags := make([]float64, len(scales))
for k := range scales {
m := 0.0
for i := range n {
z := w.ComplexAt(k*n + i)
m = max(m, math.Hypot(real(z), imag(z)))
}
mags[k] = m
}
// The bump has unit width in time units: dt·(width samples); the
// peak response lands at the scale of order one, not the extremes.
if mags[1] < mags[0] || mags[1] < mags[len(scales)-1] {
t.Fatalf("bump response %v peaks at an extreme scale", mags)
}
// DC insensitivity: a constant gives a numerically zero transform
// in the interior, away from both wrap-around edges by more than
// the wavelet's ~4-scale reach (4a/dt = 40 samples here).
consts := make([]float64, 256)
for i := range consts {
consts[i] = 3.5
}
cx, _ := core.FromFloats(consts, 256)
cw, err := CWT(cx, MexicanHat, []float64{1.0}, 0.1)
if err != nil {
t.Fatalf("CWT constant: %v", err)
}
for i := 64; i < 192; i++ {
z := cw.ComplexAt(i)
if math.Hypot(real(z), imag(z)) > 1e-5 {
t.Fatalf("constant leaked into the mexican-hat transform at %d: %g", i, math.Hypot(real(z), imag(z)))
}
}
}
// TestWaveletErrors pins the input gates.
func TestWaveletErrors(t *testing.T) {
x, _ := core.FromFloats([]float64{1, 2, 3, 4}, 4)
if _, err := DWT(x, 0); err == nil {
t.Error("zero levels accepted")
}
if _, err := DWT(x, 3); err == nil {
t.Error("levels beyond log2(n) accepted")
}
odd, _ := core.FromFloats([]float64{1, 2, 3}, 3)
if _, err := DWT(odd, 1); err == nil {
t.Error("odd length accepted")
}
if _, err := CWT(x, "db4", []float64{1}, 0.1); err == nil {
t.Error("unknown wavelet accepted")
}
if _, err := CWT(x, Morlet, []float64{-1}, 0.1); err == nil {
t.Error("negative scale accepted")
}
}
// cwtReference computes one scale of CWT directly in the time domain:
// the output at sample i is the cyclic correlation of the signal with
// the wavelet, sum over j of x[j]·conj(ψ[(j−i) mod 2n]). The sum is
// what the transform's product of spectra evaluates, stated without an
// FFT, and the wavelet is built from the documented rule: its sample j
// sits at t = dt·j and the samples past the midpoint wrap to
// t = dt·(j − 2n), so the wavelet is zero-padded and centred. The
// transform removes the wavelet's DC bin, which the subtraction of the
// mean does here.
func cwtReference(vals []float64, dt, a float64) []complex128 {
n := len(vals)
wlen := 2 * n
const omega0 = 5.0
amp := math.Pow(math.Pi, -0.25) / a
w := make([]complex128, wlen)
mean := complex(0, 0)
for j := range wlen {
tt := dt * float64(j)
if j > wlen/2 {
tt -= dt * float64(wlen) // wrap to the negative half
}
arg := tt / a
w[j] = complex(amp*math.Cos(omega0*arg), amp*math.Sin(omega0*arg)) *
complex(math.Exp(-arg*arg/2), 0)
mean += w[j]
}
mean /= complex(float64(wlen), 0)
row := make([]complex128, n)
for i := range n {
s := complex(0, 0)
for j := range n {
s += complex(vals[j], 0) * conj128(w[((j-i)%wlen+wlen)%wlen]-mean)
}
row[i] = s
}
return row
}
// TestCWTWrapBoundary pins the wavelet's wrap point against the direct
// correlation above. The transform's route differs from the reference
// only in the wavelet sample at j = wlen/2: it is the last sample of
// the non-negative half, at t = +dt·n, and a pass that wraps it too
// puts it at −dt·n instead. That single sample is invisible at small
// scales, where the Gaussian has long since decayed, and decides the
// answer at scales near dt·n, because zeroing the wavelet's DC bin
// leaves every sample's deviation from the mean in every output. The
// scales below span both regimes.
func TestCWTWrapBoundary(t *testing.T) {
const (
n = 64
dt = 1.0
)
vals := make([]float64, n)
for i := range n {
vals[i] = math.Sin(0.3*float64(i)) + 0.25*math.Cos(0.11*float64(i))
}
x := mustFromFloats(t, vals, n)
scales := []float64{float64(n) * dt, float64(n) * dt / 2, 4, 0.5}
out, err := CWT(x, Morlet, scales, dt)
if err != nil {
t.Fatalf("CWT: %v", err)
}
for k, a := range scales {
want := cwtReference(vals, dt, a)
worst, mag := 0.0, 0.0
for i := range n {
got := out.ComplexAt(k*n + i)
d := got - want[i]
if e := math.Hypot(real(d), imag(d)); e > worst {
worst = e
}
if m := math.Hypot(real(want[i]), imag(want[i])); m > mag {
mag = m
}
}
if worst > 1e-11*mag {
t.Fatalf("scale %g: worst deviation from the direct correlation %.6g (scale %.6g), want the wrapped wavelet sample to sit at +dt·n",
a, worst, mag)
}
}
}