// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package signal import ( "math" "strings" "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // stftTone builds n samples of cos(2π·cycles·i/n). func stftTone(t *testing.T, cycles, n int) *core.Array { t.Helper() vals := make([]float64, n) for i := range n { vals[i] = math.Cos(2 * math.Pi * float64(cycles) * float64(i) / float64(n)) } a, err := core.FromFloats(vals, n) if err != nil { t.Fatalf("FromFloats: %v", err) } return a } // TestSTFTFrameSpectra checks the frame geometry on a tone with an // integer number of cycles per segment under the box window: every // frame is the same spectrum with its peak in the tone's bin. func TestSTFTFrameSpectra(t *testing.T) { const n, segment, cycles = 256, 64, 16 x, err := STFT(stftTone(t, cycles, n), STFTOptions{Segment: segment, Overlap: 0, Window: "box"}) if err != nil { t.Fatalf("STFT: %v", err) } if x.Shape()[0] != n/segment || x.Shape()[1] != segment { t.Fatalf("shape %v, want (%d, %d)", x.Shape(), n/segment, segment) } rows := x.RawComplexes()[:x.Len()] // 16 cycles over 256 samples is 4 cycles per 64-sample segment. perSegment := cycles * segment / n for f := range x.Shape()[0] { peak, peakMag := 0, 0.0 for k := range segment { if mag := math.Hypot(real(rows[f*segment+k]), imag(rows[f*segment+k])); mag > peakMag { peak, peakMag = k, mag } } if peak != perSegment && peak != segment-perSegment { t.Fatalf("frame %d peaks at bin %d, want %d or %d", f, peak, perSegment, segment-perSegment) } } } // TestSpectrogramMatchesWelch pins the scaling contract: averaging // the spectrogram over its frames must reproduce WelchPSD on the same // segmentation, because the spectrogram is that estimate unaveraged. func TestSpectrogramMatchesWelch(t *testing.T) { const n = 512 vals := make([]float64, n) for i := range n { vals[i] = math.Cos(2*math.Pi*13*float64(i)/float64(n)) + 0.25*math.Sin(2*math.Pi*40*float64(i)/float64(n)) } x, err := core.FromFloats(vals, n) if err != nil { t.Fatalf("FromFloats: %v", err) } const fs, segment, overlap, window = 256.0, 128, 64, "hann" opts := STFTOptions{Segment: segment, Overlap: overlap, Window: window} spec, err := Spectrogram(x, fs, opts) if err != nil { t.Fatalf("Spectrogram: %v", err) } freqs, psd, err := WelchPSD(x, fs, segment, overlap, window) if err != nil { t.Fatalf("WelchPSD: %v", err) } if freqs.Len() != spec.Shape()[1] { t.Fatalf("bin counts disagree: spectrogram %d, Welch %d", spec.Shape()[1], freqs.Len()) } frames := float64(spec.Shape()[0]) vals = spec.RawFloats()[:spec.Len()] bins := spec.Shape()[1] for k := range bins { mean := 0.0 for f := range spec.Shape()[0] { mean += vals[f*bins+k] } mean /= frames if math.Abs(mean-psd.RawFloats()[k]) > 1e-12+1e-9*math.Abs(psd.RawFloats()[k]) { t.Fatalf("bin %d: spectrogram mean %.12g vs Welch %.12g", k, mean, psd.RawFloats()[k]) } } } // TestSTFTRefusals checks the geometry and window guards. func TestSTFTRefusals(t *testing.T) { bad := core.New(core.Float, 2, 2) if _, err := STFT(bad, STFTOptions{Segment: 2}); err == nil { t.Fatal("matrix accepted") } one := core.New(core.Float, 16) if _, err := STFT(one, STFTOptions{Segment: 32}); err == nil { t.Fatal("segment past the signal accepted") } if _, err := STFT(one, STFTOptions{Segment: 8, Overlap: 8}); err == nil { t.Fatal("overlap at the segment length accepted") } if _, err := STFT(one, STFTOptions{Segment: 8, Window: "gauss"}); err == nil { t.Fatal("unknown window accepted") } if _, err := Spectrogram(one, -1, STFTOptions{Segment: 8}); err == nil { t.Fatal("negative fs accepted") } } // TestSpectrogramRefusesHostileSegment pins the validation order: the // taper is allocated from Segment, so a hostile segment must be // refused before any allocation happens, not after it. func TestSpectrogramRefusesHostileSegment(t *testing.T) { x := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8}, 8) for _, segment := range []int{-1, 0, 1, 9} { if _, err := Spectrogram(x, 8, STFTOptions{Segment: segment}); err == nil || !strings.Contains(err.Error(), "segment") { t.Errorf("Spectrogram with segment %d: %v", segment, err) } } }