140 lines
3.8 KiB
Go
140 lines
3.8 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
// Command wavelets demonstrates the discrete wavelet transform on a
|
|
// denoising task and the continuous transform on a time-frequency
|
|
// task: a clean signal is buried in noise, the detail coefficients are
|
|
// soft-thresholded and the signal rebuilt, then a two-tone signal with
|
|
// an abrupt frequency change is mapped by the CWT so the change is
|
|
// visible in time, not just in frequency.
|
|
//
|
|
// Usage: go run ./examples/wavelets
|
|
package main
|
|
|
|
import (
|
|
"fmt"
|
|
"log"
|
|
"math"
|
|
|
|
"sourcedock.dev/petrbalvin/tensor"
|
|
"sourcedock.dev/petrbalvin/tensor/signal"
|
|
)
|
|
|
|
func main() {
|
|
const n = 1024
|
|
|
|
// A clean decaying sinusoid, buried in noise drawn from the
|
|
// reproducible generator so the run is exactly repeatable.
|
|
g := tensor.NewGenerator(2026)
|
|
noise, err := tensor.Normal(g, n, 0, 0.25)
|
|
if err != nil {
|
|
log.Fatal(err)
|
|
}
|
|
clean := make([]float64, n)
|
|
dirty := make([]float64, n)
|
|
for i := range n {
|
|
x := float64(i) / n
|
|
clean[i] = math.Sin(2*math.Pi*3*x) * math.Exp(-3*x)
|
|
nv, _ := tensor.FloatAt(noise, i)
|
|
dirty[i] = clean[i] + nv
|
|
}
|
|
dirtyArr, err := tensor.FromFloats(dirty, n)
|
|
if err != nil {
|
|
log.Fatal(err)
|
|
}
|
|
|
|
// Decompose, soft-threshold the detail coefficients, rebuild. The
|
|
// threshold sits at twice the noise standard deviation, the level
|
|
// where a noise-only coefficient almost never survives.
|
|
const levels = 5
|
|
coef, err := signal.DWT(dirtyArr, levels)
|
|
if err != nil {
|
|
log.Fatal(err)
|
|
}
|
|
approx := n >> levels
|
|
const threshold = 2 * 0.25
|
|
raw := coef.RawFloats()
|
|
for i := approx; i < len(raw); i++ {
|
|
v := raw[i]
|
|
switch {
|
|
case v > threshold:
|
|
raw[i] = v - threshold
|
|
case v < -threshold:
|
|
raw[i] = v + threshold
|
|
default:
|
|
raw[i] = 0
|
|
}
|
|
}
|
|
denoised, err := signal.IDWT(coef, levels)
|
|
if err != nil {
|
|
log.Fatal(err)
|
|
}
|
|
|
|
mse := func(a []float64) float64 {
|
|
s := 0.0
|
|
for i := range n {
|
|
d := a[i] - clean[i]
|
|
s += d * d
|
|
}
|
|
return s / float64(n)
|
|
}
|
|
fmt.Println("mean squared error against the clean signal:")
|
|
fmt.Printf(" noisy %.6f\n", mse(dirty))
|
|
fmt.Printf(" denoised %.6f\n", mse(denoised.RawFloats()[:n]))
|
|
fmt.Println()
|
|
|
|
// The continuous transform: 512 samples of a signal whose tone
|
|
// jumps from 8 to 32 cycles over the whole run, halfway through.
|
|
// A Morlet scale a responds at omega0/(2*pi*a) cycles per sample,
|
|
// which is omega0*N/(2*pi*a) cycles per record of N = 512 samples,
|
|
// so with omega0 = 5 the two tones live near a = 51 and a = 13;
|
|
// the scalogram ridge must jump between them.
|
|
const m = 512
|
|
chirp := make([]float64, m)
|
|
for i := range m {
|
|
freq := 8.0
|
|
if i >= m/2 {
|
|
freq = 32.0
|
|
}
|
|
chirp[i] = math.Sin(2 * math.Pi * freq * float64(i) / m)
|
|
}
|
|
chirpArr, err := tensor.FromFloats(chirp, m)
|
|
if err != nil {
|
|
log.Fatal(err)
|
|
}
|
|
scales := []float64{4, 8, 13, 16, 26, 32, 51, 64}
|
|
scalogram, err := signal.CWT(chirpArr, signal.Morlet, scales, 1)
|
|
if err != nil {
|
|
log.Fatal(err)
|
|
}
|
|
fmt.Println("CWT ridge: the scale carrying the peak energy in each half")
|
|
// The wavelet of scale 64 spans about 256 samples, so the outer
|
|
// quarters of the run are edge territory; the ridge is read from
|
|
// the interior of each half only.
|
|
const margin = 128
|
|
for _, seg := range []struct {
|
|
label string
|
|
start, stop int
|
|
}{
|
|
{"first half ", margin, m/2 - margin/2},
|
|
{"second half", m/2 + margin/2, m - margin},
|
|
} {
|
|
best := 0
|
|
bestMag := -1.0
|
|
for si := range scales {
|
|
for i := seg.start; i < seg.stop; i++ {
|
|
// The scalogram is (len(scales), m), one complex row
|
|
// per scale; the ridge is the peak magnitude.
|
|
cv, err := tensor.ComplexAt(scalogram, si, i)
|
|
if err != nil {
|
|
log.Fatal(err)
|
|
}
|
|
if a := math.Hypot(real(cv), imag(cv)); a > bestMag {
|
|
best, bestMag = si, a
|
|
}
|
|
}
|
|
}
|
|
fmt.Printf(" %s: scale %.0f\n", seg.label, scales[best])
|
|
}
|
|
}
|