129 lines
4.2 KiB
Go
129 lines
4.2 KiB
Go
// 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)
|
|
}
|
|
}
|
|
}
|