170 lines
6.2 KiB
Go
170 lines
6.2 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"
|
||
)
|
||
|
||
// Welch's power spectral density estimate: the signal is cut into
|
||
// overlapping segments, each segment is windowed, transformed and
|
||
// reduced to a periodogram, and the periodograms are averaged. The
|
||
// overlap and the window taper trade variance against spectral
|
||
// leakage; the averaging is what separates Welch from a single
|
||
// periodogram on noisy data.
|
||
|
||
// welchMinSegmentsPerWorker is the smallest per-worker chunk of
|
||
// segments the Welch estimate splits for: below it a chunk's windowing
|
||
// and transforms no longer pay the worker spawn cost.
|
||
const welchMinSegmentsPerWorker = 4
|
||
|
||
// welchScratchMax bounds the per-worker transform scratch the pool
|
||
// keeps, in complex128 entries (1 MiB of memory per buffer at the
|
||
// cap). A bigger segment's scratch is dropped on return and rebuilt by
|
||
// the next borrower, so one huge segment size cannot pin its scratch
|
||
// on every processor; sync.Pool additionally forgets everything at
|
||
// each garbage collection. The estimate allocates at most one buffer
|
||
// per worker per call either way.
|
||
const welchScratchMax = 1 << 16
|
||
|
||
// welchScratch recycles the per-worker transform buffer through the
|
||
// pool: the pointer form keeps the Put from boxing a slice header on
|
||
// every return.
|
||
type welchScratch struct{ z []complex128 }
|
||
|
||
var welchScratchPool = sync.Pool{
|
||
New: func() any { return new(welchScratch) },
|
||
}
|
||
|
||
// WelchPSD estimates the one-sided power spectral density of the
|
||
// vector x sampled at fs hertz. segment is the length of each segment
|
||
// in samples, overlap the number of samples two neighbouring segments
|
||
// share (0 means none; segment/2 is the usual choice), and window
|
||
// names the taper: "hann", "hamming" or "box". The estimate averages
|
||
// (segment/2 + 1) bins from the (n−overlap)/(segment−overlap)
|
||
// segments the signal fills (a signal exactly one segment long is a
|
||
// single-segment periodogram) with the one-sided doubling applied
|
||
// away from DC and the Nyquist bin and the window's power normalising
|
||
// the scale, so a white noise sequence of variance σ² estimates σ²
|
||
// across the band. A short signal, a segment that exceeds it, a
|
||
// negative overlap not below the segment, or an unknown window name
|
||
// is an error.
|
||
func WelchPSD(x *core.Array, fs float64, segment, overlap int, window string) (freqs, psd *core.Array, err error) {
|
||
const name = "WelchPSD"
|
||
if x.NDim() != 1 {
|
||
return nil, nil, base.Errf("%s: the signal must be a vector, got shape %s", name, base.ShapeText(x.Shape()))
|
||
}
|
||
if x.Dtype() == core.Complex {
|
||
return nil, nil, base.Errf("%s: complex arrays are not supported", name)
|
||
}
|
||
n := x.Len()
|
||
if segment < 2 || segment > n {
|
||
return nil, nil, base.Errf("%s: segment must be between 2 and %d, got %d", name, n, segment)
|
||
}
|
||
if overlap < 0 || overlap >= segment {
|
||
return nil, nil, base.Errf("%s: overlap must be in [0, %d), got %d", name, segment, overlap)
|
||
}
|
||
if !(fs > 0) || math.IsInf(fs, 0) {
|
||
return nil, nil, base.Errf("%s: fs must be positive and finite, got %g", name, fs)
|
||
}
|
||
taper, err := windowTaper(window, segment)
|
||
if err != nil {
|
||
return nil, nil, base.Errf("%s: %w", name, err)
|
||
}
|
||
wPower := 0.0
|
||
for i := range segment {
|
||
wPower += taper[i] * taper[i]
|
||
}
|
||
if wPower == 0 {
|
||
return nil, nil, base.Errf("%s: the window has zero power", name)
|
||
}
|
||
// The segments start every segment−overlap samples, so the count is
|
||
// the quotient, not one below it: a signal of exactly one segment
|
||
// fills one segment, not zero.
|
||
segments := (n - overlap) / (segment - overlap)
|
||
if segments < 1 {
|
||
return nil, nil, base.Errf("%s: %d samples fill only one %d-sample segment with overlap %d",
|
||
name, n, segment, overlap)
|
||
}
|
||
bins := segment/2 + 1
|
||
// The magnitude rows come from the pooled scratch: every element
|
||
// is written by the segment pass below before the averaging pass
|
||
// reads it, so a recycled buffer behaves exactly like a fresh one.
|
||
mags := engine.GetFloat64Buf(segments * bins)
|
||
defer engine.PutFloat64Buf(mags)
|
||
src := widenFloats(x)
|
||
// The segments split across workers: every segment writes its own
|
||
// row of per-bin magnitudes and touches no shared state, and the
|
||
// averaging pass below then adds the rows in segment order, so the
|
||
// accumulation sequence over the segments is the serial one
|
||
// unchanged and no addend moves.
|
||
engine.ParallelMin(segments, welchMinSegmentsPerWorker, func(ss, se int) {
|
||
// One scratch segment per worker, borrowed from the pool and
|
||
// fully rewritten by the windowing below for every one of the
|
||
// segments it serves; wrapping it in a core.Array for FFT would
|
||
// copy these same values into a fresh slice only to run the
|
||
// same transform.
|
||
sc := welchScratchPool.Get().(*welchScratch)
|
||
if cap(sc.z) < segment {
|
||
sc.z = make([]complex128, segment)
|
||
} else {
|
||
sc.z = sc.z[:segment]
|
||
}
|
||
z := sc.z
|
||
for s := ss; s < se; s++ {
|
||
start := s * (segment - overlap)
|
||
for i := range segment {
|
||
z[i] = complex(src[start+i]*taper[i], 0)
|
||
}
|
||
transform(z, -1)
|
||
row := mags[s*bins : (s+1)*bins]
|
||
for k := range bins {
|
||
mag := base.AbsComplex(z[k])
|
||
row[k] = mag * mag
|
||
}
|
||
}
|
||
// Retention cap: scratch above welchScratchMax is dropped
|
||
// rather than pooled, so a huge segment size cannot pin its
|
||
// buffer on every processor.
|
||
if cap(sc.z) > welchScratchMax {
|
||
sc.z = nil
|
||
}
|
||
welchScratchPool.Put(sc)
|
||
})
|
||
power := make([]float64, bins)
|
||
for s := range segments {
|
||
row := mags[s*bins : (s+1)*bins]
|
||
for k := range bins {
|
||
power[k] += row[k]
|
||
}
|
||
}
|
||
scale := 1 / (float64(segments) * fs * wPower)
|
||
out, oerr := core.Zeros(core.Float, []int{bins}...)
|
||
if oerr != nil {
|
||
return nil, nil, oerr
|
||
}
|
||
psdVals := out.RawFloats()
|
||
for k := range bins {
|
||
p := power[k] * scale
|
||
if k > 0 && k < segment-k {
|
||
p *= 2 // one-sided doubling between DC and Nyquist
|
||
}
|
||
psdVals[k] = p
|
||
}
|
||
freqsArr, oerr := core.Zeros(core.Float, []int{bins}...)
|
||
if oerr != nil {
|
||
return nil, nil, oerr
|
||
}
|
||
freqVals := freqsArr.RawFloats()
|
||
for k := range bins {
|
||
freqVals[k] = float64(k) * fs / float64(segment)
|
||
}
|
||
return freqsArr, out, nil
|
||
}
|