226 lines
5.6 KiB
Go
226 lines
5.6 KiB
Go
// 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")
|
|
}
|
|
}
|