// Copyright (c) 2026 Petr Balvín (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 }