// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package signal import "sourcedock.dev/petrbalvin/tensor/internal/core" import ( "math" "testing" ) // dctNaive evaluates the orthonormal DCT/DST of the given type by the // defining sum, the reference the padded-FFT route must match. func dctNaive(x []float64, kind int, cosine bool) []float64 { n := len(x) y := make([]float64, n) part := math.Cos if !cosine { part = math.Sin } norm := func(k int) float64 { return math.Sqrt(2 / float64(n)) } wIn := func(j int) float64 { return 1 } wOut := func(k int) float64 { return 1 } switch kind { case 1: if cosine { norm = func(int) float64 { return math.Sqrt(2 / float64(n-1)) } wIn = func(j int) float64 { if j == 0 || j == n-1 { return math.Sqrt2 / 2 } return 1 } wOut = wIn } else { norm = func(int) float64 { return math.Sqrt(2 / float64(n+1)) } } case 2: if cosine { norm = func(k int) float64 { if k == 0 { return math.Sqrt(1 / float64(n)) } return math.Sqrt(2 / float64(n)) } } else { norm = func(k int) float64 { if k == n-1 { return math.Sqrt(1 / float64(n)) } return math.Sqrt(2 / float64(n)) } } case 3: wIn = func(j int) float64 { if (cosine && j == 0) || (!cosine && j == n-1) { return math.Sqrt(1 / float64(n)) } return math.Sqrt(2 / float64(n)) } wOut = func(int) float64 { return 1 } norm = func(int) float64 { return 1 } } for k := range n { sum := 0.0 for j := range n { var arg float64 switch kind { case 1: if cosine { arg = math.Pi * float64(j*k) / float64(n-1) } else { arg = math.Pi * float64((j+1)*(k+1)) / float64(n+1) } case 2: if cosine { arg = math.Pi * float64((2*j+1)*k) / float64(2*n) } else { arg = math.Pi * float64((2*j+1)*(k+1)) / float64(2*n) } case 3: if cosine { arg = math.Pi * float64((2*k+1)*j) / float64(2*n) } else { arg = math.Pi * float64((2*k+1)*(j+1)) / float64(2*n) } case 4: arg = math.Pi * float64((2*j+1)*(2*k+1)) / float64(4*n) } sum += wIn(j) * x[j] * part(arg) } y[k] = norm(k) * wOut(k) * sum } return y } // TestDCTDSTAgainstDefinitions checks every type and direction against // the defining sums on a deterministic vector. func TestDCTDSTAgainstDefinitions(t *testing.T) { n := 8 x := make([]float64, n) for i := range n { x[i] = math.Sin(float64(3*i+1)) + 0.25*math.Cos(float64(5*i)) } xa := mustFloats(t, x, n) for kind := 1; kind <= 4; kind++ { if kind == 1 && n < 2 { continue } for _, cosine := range []bool{true, false} { got, err := dctdst(xa, kind, cosine) if err != nil { t.Fatalf("kind %d cosine %v: %v", kind, cosine, err) } want := dctNaive(x, kind, cosine) for k := range n { if math.Abs(got.FloatAt(k)-want[k]) > 1e-12 { t.Fatalf("kind %d cosine %v: y[%d] = %.14g, want %.14g", kind, cosine, k, got.FloatAt(k), want[k]) } } } } } // TestDCTDSTInverses checks the aliasing rule: forward then inverse // returns the input for all eight pairs. func TestDCTDSTInverses(t *testing.T) { n := 11 x := make([]float64, n) for i := range n { x[i] = math.Cos(float64(2*i + 1)) } xa := mustFloats(t, x, n) pairs := []struct { fwd, inv func(*core.Array, int) (*core.Array, error) kind int }{{DCT, IDCT, 1}, {DCT, IDCT, 2}, {DCT, IDCT, 3}, {DCT, IDCT, 4}, {DST, IDST, 1}, {DST, IDST, 2}, {DST, IDST, 3}, {DST, IDST, 4}} for _, p := range pairs { mid, err := p.fwd(xa, p.kind) if err != nil { t.Fatalf("kind %d forward: %v", p.kind, err) } back, err := p.inv(mid, p.kind) if err != nil { t.Fatalf("kind %d inverse: %v", p.kind, err) } for i := range n { if math.Abs(back.FloatAt(i)-x[i]) > 1e-11 { t.Fatalf("kind %d round trip [%d] = %.14g, want %.14g", p.kind, i, back.FloatAt(i), x[i]) } } } } // TestDCTDSTOrthogonality pins the orthonormal claim on the // self-inverse types: applying DCT-I or DCT-IV to the identity's // columns returns a matrix whose Gram is the identity. func TestDCTDSTOrthogonality(t *testing.T) { n := 6 for _, kind := range []int{1, 4} { col := make([]float64, n) gram := make([]float64, n*n) for j := range n { for i := range n { col[i] = 0 if i == j { col[i] = 1 } } y, err := DCT(mustFloats(t, col, n), kind) if err != nil { t.Fatalf("DCT kind %d: %v", kind, err) } for i := range n { gram[i*n+j] = y.FloatAt(i) } } for i := range n { for j := range n { s := 0.0 for l := range n { s += gram[l*n+i] * gram[l*n+j] } want := 0.0 if i == j { want = 1 } if math.Abs(s-want) > 1e-12 { t.Fatalf("DCT-%d Gram[%d][%d] = %.14g, want %.14g", kind, i, j, s, want) } } } } } // TestDCTDSTErrors pins the validation contract. func TestDCTDSTErrors(t *testing.T) { xa := mustFloats(t, []float64{1, 2, 3}, 3) if _, err := DCT(xa, 5); err == nil { t.Fatal("expected an error for kind 5") } if _, err := IDCT(xa, 0); err == nil { t.Fatal("expected an error for kind 0") } if _, err := DST(xa, 7); err == nil { t.Fatal("expected an error for kind 7") } if _, err := IDST(xa, -1); err == nil { t.Fatal("expected an error for kind -1") } if _, err := DCT(mustFloats(t, []float64{1}, 1), 1); err == nil { t.Fatal("expected an error for type I on one point") } rank2, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2) if _, err := DCT(rank2, 2); err == nil { t.Fatal("expected an error for a rank-2 input") } if _, err := DST(mustFloats(t, nil), 2); err == nil { t.Fatal("expected an error for an empty input") } }