113 lines
2.9 KiB
Go
113 lines
2.9 KiB
Go
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])
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|