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

170 lines
6.2 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"
)
// 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
}