Files

129 lines
4.2 KiB
Go
Raw Permalink 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"
"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)
}
}
}