Files
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

311 lines
10 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"
"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
}