Files

226 lines
5.6 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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")
}
}