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]) } } } }