Files
tensor/signal/transform_cache_pins_test.go
T

113 lines
2.9 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
package signal
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// The fit cache is keyed by the window and the order together; a key
// that lost the order would answer one order's weights for another.
// Each reference is computed with an empty cache, so a collision
// between the two keys cannot hide behind a shared fit row.
func TestSavitzkyGolayCacheKeyCarriesOrder(t *testing.T) {
n := 64
vals := make([]float64, n)
for i := range vals {
x := float64(i) / float64(n-1)
vals[i] = x*x*x - 2*x
}
src, err := core.FromFloats(vals, n)
if err != nil {
t.Fatal(err)
}
run := func() (map[int][]float64, error) {
sgFitKeys = map[sgFitKey]*sgFits{}
got := map[int][]float64{}
for _, order := range []int{1, 2, 3} {
out, err := SavitzkyGolay(src, 7, order)
if err != nil {
return nil, err
}
raw := out.RawFloats()
got[order] = append([]float64(nil), raw...)
}
return got, nil
}
want, err := run()
if err != nil {
t.Fatal(err)
}
got, err := run()
if err != nil {
t.Fatal(err)
}
// The orders must differ from one another on a curved signal,
// otherwise the comparison below pins nothing.
for a := 1; a <= 2; a++ {
for i := range want[a] {
if want[a][i] != want[a+1][i] {
break
}
if i == len(want[a])-1 {
t.Fatalf("orders %d and %d agree everywhere; the fixture does not separate them", a, a+1)
}
}
}
for order := 1; order <= 3; order++ {
for i := range want[order] {
if want[order][i] != got[order][i] {
t.Fatalf("order %d, sample %d: cached run answered %v, fresh-cache run answered %v", order, i, got[order][i], want[order][i])
}
}
}
}
// Sweeping more scales than the cache's entry cap must leave every
// answer unchanged and the cache bounded.
func TestCWTSpectrumCacheCap(t *testing.T) {
n := 32
vals := make([]float64, n)
for i := range vals {
vals[i] = math.Sin(2*math.Pi*float64(i)/8) + 0.25*math.Sin(2*math.Pi*float64(i)/3)
}
src, err := core.FromFloats(vals, n)
if err != nil {
t.Fatal(err)
}
const scales = cwtSpectrumCacheMax + 8
scaleList := make([]float64, scales)
for k := range scaleList {
scaleList[k] = 2 + float64(k)*0.5
}
build := func() [][]complex128 {
cwtSpectrumKeys = map[cwtSpectrumKey][]complex128{}
out, err := CWT(src, Morlet, scaleList, 1)
if err != nil {
t.Fatalf("CWT: %v", err)
}
raw, err := out.ComplexValues("CWT")
if err != nil {
t.Fatalf("ComplexValues: %v", err)
}
rows := make([][]complex128, len(scaleList))
for k := range rows {
rows[k] = append([]complex128(nil), raw[k*n:(k+1)*n]...)
}
return rows
}
want := build()
got := build()
if len(cwtSpectrumKeys) > cwtSpectrumCacheMax {
t.Fatalf("the spectrum cache holds %d entries over the cap %d", len(cwtSpectrumKeys), cwtSpectrumCacheMax)
}
for k := range want {
for i := range want[k] {
if want[k][i] != got[k][i] {
t.Fatalf("scale %v, sample %d: %v vs %v", scaleList[k], i, got[k][i], want[k][i])
}
}
}
}