Files
tensor/signal/stft.go
T

159 lines
5.4 KiB
Go
Raw 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"
"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
}