159 lines
5.4 KiB
Go
159 lines
5.4 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||
// SPDX-License-Identifier: MIT
|
||
|
||
package signal
|
||
|
||
import (
|
||
"math"
|
||
|
||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
"sourcedock.dev/petrbalvin/tensor/internal/engine"
|
||
)
|
||
|
||
// The short-time Fourier transform: the signal is cut into
|
||
// overlapping segments, each segment is windowed and transformed, and
|
||
// the frames stack into a time-by-frequency picture. The spectrogram
|
||
// is the magnitude squared of it, scaled per frame exactly as Welch's
|
||
// estimate scales its periodograms, so averaging a spectrogram over
|
||
// its frames reproduces WelchPSD on the same segmentation.
|
||
|
||
// STFTOptions tunes the transform. Segment is the length of one
|
||
// frame in samples, Overlap the samples two neighbouring frames share
|
||
// (segment/2 is the usual choice) and Window names the taper: "hann",
|
||
// "hamming" or "box", with "hann" implied by the zero value.
|
||
type STFTOptions struct {
|
||
Segment int
|
||
Overlap int
|
||
Window string
|
||
}
|
||
|
||
// framesOf computes the frame count of a signal of n samples under
|
||
// the segment and overlap geometry, the count Welch's estimate uses.
|
||
func framesOf(n, segment, overlap int) int {
|
||
return (n - overlap) / (segment - overlap)
|
||
}
|
||
|
||
// STFT transforms the vector x frame by frame and returns the
|
||
// complex frames as a (frames × segment) array in row order: row t
|
||
// holds the spectrum of the t-th frame, bin k of frame t is the
|
||
// transform of x[t·hop : t·hop+segment] under the window, with hop =
|
||
// segment − overlap. The Fourier definition treats each windowed
|
||
// segment as one period, which is the standard convention here.
|
||
func STFT(x *core.Array, opts STFTOptions) (*core.Array, error) {
|
||
const name = "STFT"
|
||
if x.NDim() != 1 {
|
||
return 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, base.Errf("%s: complex arrays are not supported", name)
|
||
}
|
||
n := x.Len()
|
||
if opts.Segment < 2 || opts.Segment > n {
|
||
return nil, base.Errf("%s: segment must be between 2 and %d, got %d", name, n, opts.Segment)
|
||
}
|
||
if opts.Overlap < 0 || opts.Overlap >= opts.Segment {
|
||
return nil, base.Errf("%s: overlap must be in [0, %d), got %d", name, opts.Segment, opts.Overlap)
|
||
}
|
||
window := opts.Window
|
||
if window == "" {
|
||
window = "hann"
|
||
}
|
||
taper, err := windowTaper(window, opts.Segment)
|
||
if err != nil {
|
||
return nil, base.Errf("%s: %w", name, err)
|
||
}
|
||
hop := opts.Segment - opts.Overlap
|
||
frames := framesOf(n, opts.Segment, opts.Overlap)
|
||
if frames < 1 {
|
||
return nil, base.Errf("%s: %d samples fill only one %d-sample segment with overlap %d",
|
||
name, n, opts.Segment, opts.Overlap)
|
||
}
|
||
src := widenFloats(x)
|
||
segment := opts.Segment
|
||
out, oerr := core.Zeros(core.Complex, frames, segment)
|
||
if oerr != nil {
|
||
return nil, oerr
|
||
}
|
||
// A frame's windowed samples go straight into its output row and are
|
||
// transformed there: the row is a private, fully overwritten buffer,
|
||
// so the copy a separate scratch segment needed is one pass fewer and
|
||
// the bits land exactly where the copy would have put them.
|
||
rows := out.RawComplexes()[:out.Len()]
|
||
engine.ParallelMin(frames, welchMinSegmentsPerWorker, func(fs, fe int) {
|
||
for f := fs; f < fe; f++ {
|
||
start := f * hop
|
||
z := rows[f*segment : (f+1)*segment]
|
||
for i := range segment {
|
||
z[i] = complex(src[start+i]*taper[i], 0)
|
||
}
|
||
transform(z, -1)
|
||
}
|
||
})
|
||
return out, nil
|
||
}
|
||
|
||
// Spectrogram returns the one-sided power spectrogram of the vector x
|
||
// sampled at fs hertz as a (frames × segment/2+1) array: every frame
|
||
// carries the periodogram of its windowed segment with the same
|
||
// one-sided doubling and window-power normalisation Welch's estimate
|
||
// applies, in units of x²/Hz. Averaging the frames reproduces
|
||
// WelchPSD on the same geometry.
|
||
func Spectrogram(x *core.Array, fs float64, opts STFTOptions) (*core.Array, error) {
|
||
const name = "Spectrogram"
|
||
if !(fs > 0) || math.IsInf(fs, 0) {
|
||
return nil, base.Errf("%s: fs must be positive and finite, got %g", name, fs)
|
||
}
|
||
// The segment geometry is checked before the taper is built: the
|
||
// taper allocates Segment samples, so a hostile segment must be
|
||
// refused here rather than after the allocation. STFT repeats the
|
||
// same checks on its own terms.
|
||
n := x.Len()
|
||
if opts.Segment < 2 || opts.Segment > n {
|
||
return nil, base.Errf("%s: segment must be between 2 and %d, got %d", name, n, opts.Segment)
|
||
}
|
||
if opts.Overlap < 0 || opts.Overlap >= opts.Segment {
|
||
return nil, base.Errf("%s: overlap must be in [0, %d), got %d", name, opts.Segment, opts.Overlap)
|
||
}
|
||
window := opts.Window
|
||
if window == "" {
|
||
window = "hann"
|
||
}
|
||
taper, err := windowTaper(window, opts.Segment)
|
||
if err != nil {
|
||
return nil, base.Errf("%s: %w", name, err)
|
||
}
|
||
wPower := 0.0
|
||
for i := range taper {
|
||
wPower += taper[i] * taper[i]
|
||
}
|
||
if wPower == 0 {
|
||
return nil, base.Errf("%s: the window has zero power", name)
|
||
}
|
||
framesSpec, err := STFT(x, opts)
|
||
if err != nil {
|
||
return nil, base.Errf("%s: %w", name, err)
|
||
}
|
||
frames := framesSpec.Shape()[0]
|
||
segment := opts.Segment
|
||
bins := segment/2 + 1
|
||
scale := 1 / (fs * wPower)
|
||
out, oerr := core.Zeros(core.Float, frames, bins)
|
||
if oerr != nil {
|
||
return nil, oerr
|
||
}
|
||
spec := framesSpec.RawComplexes()[:framesSpec.Len()]
|
||
vals := out.RawFloats()
|
||
for f := range frames {
|
||
for k := range bins {
|
||
mag := base.AbsComplex(spec[f*segment+k])
|
||
p := mag * mag * scale
|
||
if k > 0 && k < segment-k {
|
||
p *= 2
|
||
}
|
||
vals[f*bins+k] = p
|
||
}
|
||
}
|
||
return out, nil
|
||
}
|