// Copyright (c) 2026 Petr Balvín (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)< 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 }