Files

170 lines
6.2 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// 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
}