311 lines
10 KiB
Go
311 lines
10 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package signal
|
||
|
||
import (
|
||
"math"
|
||
"sync"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/engine"
|
||
)
|
||
|
||
// Wavelets. The discrete side carries Haar, the one wavelet
|
||
// whose filter coefficients are exactly ½ and whose reconstruction is
|
||
// exact by construction, at any level the length permits, in the
|
||
// packed [A_L, D_L, …, D_1] layout every wavelet text uses. The
|
||
// continuous side carries the analytic family (Morlet and
|
||
// the Mexican hat) evaluated through the FFT: no filter tables to
|
||
// trust, only formulas, which is the library's way.
|
||
|
||
// DWT returns the discrete Haar wavelet transform of a rank-1 real
|
||
// signal over levels scales (periodic boundary): the coefficients pack
|
||
// as [A_levels, D_levels, D_{levels−1}, …, D_1], the approximation
|
||
// first and each detail band after it. levels must be at least 1 and
|
||
// at most log2(n); energy is conserved (Parseval holds for the
|
||
// orthonormal Haar basis).
|
||
func DWT(x *core.Array, levels int) (*core.Array, error) {
|
||
const name = "DWT"
|
||
n, err := waveletValidate(name, x, levels)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
coef := make([]float64, n)
|
||
for i := range n {
|
||
coef[i] = x.FloatAt(i)
|
||
}
|
||
// In each pass the live block at [0, m) halves: averages land in
|
||
// [0, m/2), details (scaled by 1/sqrt2 for orthonormality) in
|
||
// [m/2, m) of the LIVE block, then shuffle the details to their
|
||
// final resting slot at the end.
|
||
m := n
|
||
for range levels {
|
||
half := m / 2
|
||
avg := make([]float64, half)
|
||
det := make([]float64, half)
|
||
s := math.Sqrt2
|
||
for i := range half {
|
||
avg[i] = (coef[2*i] + coef[2*i+1]) / s
|
||
det[i] = (coef[2*i] - coef[2*i+1]) / s
|
||
}
|
||
copy(coef[:half], avg)
|
||
// Details of this level go to the tail slot reserved for D_{l+1}.
|
||
copy(coef[n-half-(n-m):n-(n-m)], det)
|
||
m = half
|
||
}
|
||
return core.FromFloats(coef, n)
|
||
}
|
||
|
||
// IDWT inverts DWT over the same level count and layout. The
|
||
// coefficients are widened through FloatAt, the accessor DWT reads its
|
||
// input with, so every real dtype DWT accepts inverts here as well.
|
||
func IDWT(coef *core.Array, levels int) (*core.Array, error) {
|
||
const name = "IDWT"
|
||
n, err := waveletValidate(name, coef, levels)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
out := make([]float64, n)
|
||
for i := range n {
|
||
out[i] = coef.FloatAt(i)
|
||
}
|
||
// Undo level by level, innermost (shortest) first.
|
||
m := n >> levels
|
||
for l := levels; l >= 1; l-- {
|
||
half := m
|
||
m = 2 * m
|
||
avg := append([]float64(nil), out[:half]...)
|
||
det := append([]float64(nil), out[n-half-(n-m):n-(n-m)]...)
|
||
s := math.Sqrt2
|
||
for i := range half {
|
||
out[2*i] = (avg[i] + det[i]) / s
|
||
out[2*i+1] = (avg[i] - det[i]) / s
|
||
}
|
||
}
|
||
return core.FromFloats(out, n)
|
||
}
|
||
|
||
// waveletValidate checks the shared contract and returns n.
|
||
func waveletValidate(name string, x *core.Array, levels int) (int, error) {
|
||
if x.NDim() != 1 {
|
||
return 0, base.Errf("%s: needs a rank-1 signal, got shape %s", name, base.ShapeText(x.Shape()))
|
||
}
|
||
if x.Dtype() == core.Complex {
|
||
return 0, base.Errf("%s: complex signals are not supported", name)
|
||
}
|
||
n := x.Len()
|
||
if n == 0 {
|
||
return 0, base.Errf("%s: an empty signal has no transform", name)
|
||
}
|
||
if levels < 1 || (n>>levels)<<levels != n {
|
||
return 0, base.Errf("%s: levels must divide the length exactly, got %d levels for %d samples",
|
||
name, levels, n)
|
||
}
|
||
return n, nil
|
||
}
|
||
|
||
// CWTWavelet names the analysing wavelet of CWT.
|
||
type CWTWavelet string
|
||
|
||
// cwtParallelMinN is the signal length above which one scale's wavelet
|
||
// build plus FFT pair over 2n samples is worth a worker's spawn cost.
|
||
// Below it the CWT scale walk stays on the calling goroutine.
|
||
const cwtParallelMinN = 1 << 10
|
||
|
||
const (
|
||
// Morlet is the complex Morlet with ω₀ = 5, the time-frequency
|
||
// standard.
|
||
Morlet CWTWavelet = "morlet"
|
||
// MexicanHat is the real Ricker wavelet, the zero-mean second
|
||
// derivative of a Gaussian.
|
||
MexicanHat CWTWavelet = "mexicanhat"
|
||
)
|
||
|
||
// CWT returns the continuous wavelet transform of a rank-1 real signal
|
||
// sampled at spacing dt, over the given scales: a (len(scales), n)
|
||
// complex array whose row k is the transform at scales[k]. Morlet
|
||
// yields the full complex transform; MexicanHat is real, so its rows
|
||
// carry the real result with zero imaginary part. The convolution
|
||
// runs through the FFT (zero-padded to 2n), the wavelet is
|
||
// L1-normalised per scale so amplitudes stay comparable across them,
|
||
// and scales must be positive.
|
||
func CWT(x *core.Array, wavelet CWTWavelet, scales []float64, dt float64) (*core.Array, error) {
|
||
const name = "CWT"
|
||
if x.NDim() != 1 {
|
||
return nil, base.Errf("%s: needs a rank-1 signal, got shape %s", name, base.ShapeText(x.Shape()))
|
||
}
|
||
if x.Dtype() == core.Complex {
|
||
return nil, base.Errf("%s: complex signals are not supported", name)
|
||
}
|
||
// +Inf passes a bare > 0 and turns the wavelet's time axis into
|
||
// NaN with no error, so finiteness is part of the gate.
|
||
if !(dt > 0) || math.IsInf(dt, 0) {
|
||
return nil, base.Errf("%s: the sample spacing must be positive and finite, got %g", name, dt)
|
||
}
|
||
n := x.Len()
|
||
if n == 0 {
|
||
return nil, base.Errf("%s: an empty signal has no transform", name)
|
||
}
|
||
if len(scales) == 0 {
|
||
return nil, base.Errf("%s: at least one scale is required", name)
|
||
}
|
||
for _, s := range scales {
|
||
if !(s > 0) || math.IsInf(s, 0) {
|
||
return nil, base.Errf("%s: every scale must be positive and finite, got %g", name, s)
|
||
}
|
||
}
|
||
switch wavelet {
|
||
case Morlet, MexicanHat:
|
||
default:
|
||
return nil, base.Errf("%s: unknown wavelet %q (want morlet or mexicanhat)", name, wavelet)
|
||
}
|
||
|
||
// The signal padded to 2n so the linear convolution has room. The
|
||
// transform runs in place: FFT and IFFT are pure functions of the
|
||
// payload slice, so driving transform directly on scratch the
|
||
// function owns produces exactly the bits the wrappers produced,
|
||
// without their per-call copies.
|
||
sig := make([]complex128, 2*n)
|
||
for i := range n {
|
||
sig[i] = complex(x.FloatAt(i), 0)
|
||
}
|
||
transform(sig, -1)
|
||
|
||
out := make([]complex128, len(scales)*n)
|
||
// transformScale runs the whole per-scale pipeline: the wavelet's
|
||
// cached spectrum, the conjugated product with the signal's
|
||
// spectrum and the return to time. The wavelet of a scale is a pure
|
||
// function of the transform geometry, so its spectrum is served
|
||
// from the same kind of cache the transform's twiddle tables keep;
|
||
// every other scratch is the caller's w, fully rewritten per scale.
|
||
// The per-row arithmetic sequence is the serial one unchanged, so
|
||
// the split cannot move a bit.
|
||
transformScale := func(k int, a float64, w []complex128) {
|
||
spec := cwtWaveletSpectrum(wavelet, n, dt, a)
|
||
for i := range w {
|
||
w[i] = sig[i] * conj128(spec[i])
|
||
}
|
||
transform(w, +1)
|
||
// The inverse scale, applied entry by entry exactly as IFFT
|
||
// applies it, lands straight in this scale's output row.
|
||
row := out[k*n : (k+1)*n]
|
||
scale := complex(float64(2*n), 0)
|
||
for i := range row {
|
||
row[i] = w[i] / scale
|
||
}
|
||
}
|
||
// The scales split across workers once a scale's wavelet build and
|
||
// FFT pair outgrow the worker spawn cost; below that floor the
|
||
// calling goroutine walks every scale itself. Each worker owns one
|
||
// scratch line and rewrites it whole for every scale it serves.
|
||
wlen := 2 * n
|
||
if n >= cwtParallelMinN {
|
||
engine.Parallel(len(scales), func(ks, ke int) {
|
||
w := make([]complex128, wlen)
|
||
for k := ks; k < ke; k++ {
|
||
transformScale(k, scales[k], w)
|
||
}
|
||
})
|
||
} else {
|
||
w := make([]complex128, wlen)
|
||
for k, a := range scales {
|
||
transformScale(k, a, w)
|
||
}
|
||
}
|
||
arr, err := core.ComplexFromArray(out, len(scales), n)
|
||
if err != nil {
|
||
return nil, base.Errf("%s: %w", name, err)
|
||
}
|
||
return arr, nil
|
||
}
|
||
|
||
// cwtSpectrumKey names one cached wavelet spectrum: the analysing
|
||
// wavelet and the exact transform geometry it was built for.
|
||
type cwtSpectrumKey struct {
|
||
wavelet CWTWavelet
|
||
n int
|
||
dt float64
|
||
scale float64
|
||
}
|
||
|
||
// cwtSpectra caches per-scale wavelet spectra after their forward
|
||
// transform, keyed by the geometry that determines them. The spectra
|
||
// are read-only once published. Two caps bound the retention: the
|
||
// twiddleCacheMax policy applies per entry (a transform whose padded
|
||
// length exceeds it still builds its spectrum, it just does not keep
|
||
// it), and the map holds at most cwtSpectrumCacheMax entries, so a
|
||
// caller sweeping unboundedly many distinct scales cannot pin
|
||
// unbounded memory.
|
||
var (
|
||
cwtSpectrumMu sync.RWMutex
|
||
cwtSpectrumKeys = map[cwtSpectrumKey][]complex128{}
|
||
)
|
||
|
||
// cwtSpectrumCacheMax is the entry cap of the wavelet-spectrum cache.
|
||
const cwtSpectrumCacheMax = 128
|
||
|
||
// cwtWaveletSpectrum returns the forward transform of the zero-padded,
|
||
// L1-normalised wavelet of the given scale, its DC bin zeroed for
|
||
// admissibility. Building it repeats the evaluation the per-scale path
|
||
// always ran, sample for sample.
|
||
func cwtWaveletSpectrum(wavelet CWTWavelet, n int, dt, a float64) []complex128 {
|
||
key := cwtSpectrumKey{wavelet: wavelet, n: n, dt: dt, scale: a}
|
||
cwtSpectrumMu.RLock()
|
||
spec, ok := cwtSpectrumKeys[key]
|
||
cwtSpectrumMu.RUnlock()
|
||
if ok {
|
||
return spec
|
||
}
|
||
wlen := 2 * n
|
||
wrap := wlen/2 + 1
|
||
w := make([]complex128, wlen)
|
||
switch wavelet {
|
||
case Morlet:
|
||
// π^{-1/4}·e^{iω₀t/a}·e^{−t²/(2a²)}, L1-normalised by 1/a.
|
||
// Sincos answers the same bits the separate Sin and Cos
|
||
// produce (verified bit-for-bit), so the wavelet is
|
||
// unchanged while the trig work halves.
|
||
const omega0 = 5.0
|
||
amp := math.Pow(math.Pi, -0.25) / a
|
||
for j := range wrap {
|
||
arg := dt * float64(j) / a
|
||
s, c := math.Sincos(omega0 * arg)
|
||
w[j] = complex(amp*c, amp*s) *
|
||
complex(math.Exp(-arg*arg/2), 0)
|
||
}
|
||
period := dt * float64(wlen)
|
||
for j := wrap; j < wlen; j++ {
|
||
arg := (dt*float64(j) - period) / a // wrap to the negative half
|
||
s, c := math.Sincos(omega0 * arg)
|
||
w[j] = complex(amp*c, amp*s) *
|
||
complex(math.Exp(-arg*arg/2), 0)
|
||
}
|
||
case MexicanHat:
|
||
// (2/√3)π^{-1/4}(1−t²/a²)e^{−t²/2a²}, scaled by 1/a.
|
||
c := 2 / math.Sqrt(3) * math.Pow(math.Pi, -0.25)
|
||
for j := range wrap {
|
||
arg := dt * float64(j) / a
|
||
w[j] = complex(c*(1-arg*arg)*math.Exp(-arg*arg/2)/a, 0)
|
||
}
|
||
period := dt * float64(wlen)
|
||
for j := wrap; j < wlen; j++ {
|
||
arg := (dt*float64(j) - period) / a
|
||
w[j] = complex(c*(1-arg*arg)*math.Exp(-arg*arg/2)/a, 0)
|
||
}
|
||
}
|
||
transform(w, -1)
|
||
// Zero the wavelet's DC bin: the truncated tails leave a
|
||
// rounding-level mean, and admissibility demands exactly zero.
|
||
w[0] = 0
|
||
if wlen <= twiddleCacheMax {
|
||
cwtSpectrumMu.Lock()
|
||
if len(cwtSpectrumKeys) < cwtSpectrumCacheMax {
|
||
cwtSpectrumKeys[key] = w
|
||
}
|
||
cwtSpectrumMu.Unlock()
|
||
}
|
||
return w
|
||
}
|