268 lines
7.9 KiB
Go
268 lines
7.9 KiB
Go
// 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)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|