Files
tensor/signal/wave_bench_test.go
T

447 lines
14 KiB
Go
Raw 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 (
"math"
"math/big"
"sync"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Accuracy cross-check for the chirp-z line transforms the Poisson
// solves run, against a reference that is not the implementation under
// test: the defining trigonometric sums evaluated in 256-bit floating
// point, with π from Machin's formula at the same precision. The
// padded full-length route the solves previously ran survives here as
// the public DST/DCT pair, untouched, so both routes are measured
// against the same exact reference.
const refPrec = 256
// refPiFloor is the magnitude below which a series term is past the
// reference's own precision, so every series here stops on a proven
// bound rather than on an exact zero the wide exponent range may never
// produce.
var refPiFloor = new(big.Float).SetPrec(uint(refPrec)).SetMantExp(big.NewFloat(1).SetPrec(uint(refPrec)), -refPrec-8)
// refPi returns π at refPrec bits, computed once per test binary.
var refPi = sync.OnceValue(func() *big.Float {
prec := uint(refPrec)
atanInv := func(x int64) *big.Float {
xb := new(big.Float).SetPrec(prec).SetInt64(x)
term := new(big.Float).SetPrec(prec).Quo(big.NewFloat(1).SetPrec(prec), xb)
x2 := new(big.Float).SetPrec(prec).Mul(xb, xb)
sum := new(big.Float).SetPrec(prec).Set(term)
for m := 1; m < 4*refPrec; m++ {
term.Quo(term, x2)
d := big.NewFloat(float64(2*m + 1)).SetPrec(prec)
t := new(big.Float).SetPrec(prec).Quo(term, d)
if t.Cmp(refPiFloor) < 0 {
break
}
if m%2 == 0 {
sum.Add(sum, t)
} else {
sum.Sub(sum, t)
}
}
return sum
}
pi := new(big.Float).SetPrec(prec)
pi.Mul(atanInv(5), big.NewFloat(16).SetPrec(prec))
four := new(big.Float).SetPrec(prec).Mul(atanInv(239), big.NewFloat(4).SetPrec(prec))
return pi.Sub(pi, four)
})
// refSin returns sin(z) at refPrec bits by its Taylor series; callers
// keep |z| ≤ 2π, where the terms z^{2m+1}/(2m+1)! pass the precision
// floor in a few dozen steps. The series term carries the factorial
// through its own recurrence, t_m = t_{m−1}·z²/((2m)(2m+1)), so it
// decays factorially rather than as a bare power.
func refSin(z *big.Float) *big.Float {
prec := uint(refPrec)
term := new(big.Float).SetPrec(prec).Set(z)
sum := new(big.Float).SetPrec(prec).Set(z)
z2 := new(big.Float).SetPrec(prec).Mul(z, z)
floor := new(big.Float).SetPrec(prec).SetMantExp(big.NewFloat(1).SetPrec(prec), -int(refPrec)-8)
for m := 1; m < 4*refPrec; m++ {
term.Mul(term, z2)
term.Quo(term, big.NewFloat(float64(2*m*(2*m+1))).SetPrec(prec))
// term and the denominator are non-negative, so the term is;
// the floor test is a plain comparison.
if term.Sign() == 0 || term.Cmp(floor) < 0 {
break
}
if m%2 == 0 {
sum.Add(sum, term)
} else {
sum.Sub(sum, term)
}
}
return sum
}
// refCos returns cos(z) at refPrec bits by its Taylor series, the
// terms carrying the factorial through t_m = t_{m−1}·z²/((2m−1)(2m)).
func refCos(z *big.Float) *big.Float {
prec := uint(refPrec)
term := big.NewFloat(1).SetPrec(prec)
sum := big.NewFloat(1).SetPrec(prec)
z2 := new(big.Float).SetPrec(prec).Mul(z, z)
floor := new(big.Float).SetPrec(prec).SetMantExp(big.NewFloat(1).SetPrec(prec), -int(refPrec)-8)
for m := 1; m < 4*refPrec; m++ {
term.Mul(term, z2)
term.Quo(term, big.NewFloat(float64((2*m-1)*(2*m))).SetPrec(prec))
if term.Sign() == 0 || term.Cmp(floor) < 0 {
break
}
if m%2 == 0 {
sum.Add(sum, term)
} else {
sum.Sub(sum, term)
}
}
return sum
}
// refLineTransform evaluates the defining sum of the orthonormal
// DST-I (sine) or DCT-I (cosine) of x at refPrec bits. The trig
// argument ((j+1)(k+1)) mod 2(n+1) is reduced in exact integers
// first, so no series argument exceeds 2π.
func refLineTransform(x []float64, sine bool) []*big.Float {
prec := uint(refPrec)
n := len(x)
pi := refPi()
out := make([]*big.Float, n)
xb := make([]*big.Float, n)
for j := range n {
xb[j] = big.NewFloat(x[j]).SetPrec(prec)
}
if sine {
den := float64(n + 1)
sqrt2 := new(big.Float).SetPrec(prec).Sqrt(big.NewFloat(2).SetPrec(prec))
norm := new(big.Float).SetPrec(prec).Quo(sqrt2,
new(big.Float).SetPrec(prec).Sqrt(big.NewFloat(den).SetPrec(prec)))
period := 2 * (n + 1)
for k := range n {
sum := new(big.Float).SetPrec(prec)
for j := range n {
r := ((j + 1) * (k + 1)) % period
arg := new(big.Float).SetPrec(prec).Quo(
new(big.Float).SetPrec(prec).Mul(pi, big.NewFloat(float64(r)).SetPrec(prec)),
big.NewFloat(den).SetPrec(prec))
sum.Add(sum, new(big.Float).SetPrec(prec).Mul(xb[j], refSin(arg)))
}
out[k] = new(big.Float).SetPrec(prec).Mul(norm, sum)
}
return out
}
den := float64(n - 1)
sqrt2 := new(big.Float).SetPrec(prec).Sqrt(big.NewFloat(2).SetPrec(prec))
half := new(big.Float).SetPrec(prec).Quo(sqrt2, big.NewFloat(2).SetPrec(prec))
norm := new(big.Float).SetPrec(prec).Quo(sqrt2,
new(big.Float).SetPrec(prec).Sqrt(big.NewFloat(den).SetPrec(prec)))
period := 2 * (n - 1)
wj := make([]*big.Float, n)
for j := range n {
if j == 0 || j == n-1 {
wj[j] = half
} else {
wj[j] = big.NewFloat(1).SetPrec(prec)
}
}
for k := range n {
sum := new(big.Float).SetPrec(prec)
for j := range n {
r := (j * k) % period
arg := new(big.Float).SetPrec(prec).Quo(
new(big.Float).SetPrec(prec).Mul(pi, big.NewFloat(float64(r)).SetPrec(prec)),
big.NewFloat(den).SetPrec(prec))
term := new(big.Float).SetPrec(prec).Mul(xb[j], wj[j])
sum.Add(sum, term.Mul(term, refCos(arg)))
}
wk := big.NewFloat(1).SetPrec(prec)
if k == 0 || k == n-1 {
wk = half
}
out[k] = new(big.Float).SetPrec(prec).Mul(norm,
new(big.Float).SetPrec(prec).Mul(wk, sum))
}
return out
}
// lineFixture builds a deterministic line of length n with mixed
// magnitudes and signs.
func lineFixture(n int) []float64 {
x := make([]float64, n)
state := uint64(0x9e3779b97f4a7c15) ^ uint64(n)
for i := range n {
state ^= state << 13
state ^= state >> 7
state ^= state << 17
v := float64(int64(state)%2000)/1000 - 1
x[i] = v * math.Pow(2, float64((i%7)-3))
}
return x
}
// legacyLine runs the pre-chirp route for one line: the public DST or
// DCT at kind 1, which carries the padded full-length arithmetic the
// solves used to run per line.
func legacyLine(t *testing.T, x []float64, sine bool) []float64 {
arr, err := core.FromFloats(x, len(x))
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
var out *core.Array
if sine {
out, err = DST(arr, 1)
} else {
out, err = DCT(arr, 1)
}
if err != nil {
t.Fatalf("public transform: %v", err)
}
vals := make([]float64, out.Len())
for i := range vals {
vals[i] = out.FloatAt(i)
}
return vals
}
// TestPoissonLineTransformAccuracy holds both line-transform routes
// against the exact trigonometric sums at 256 bits: the chirp-z route
// the solves run now, and the padded full-length route the public
// DST/DCT still runs. The verdict is printed for every length and
// family; the assertion keeps the chirp route at the rounding floor
// and never far behind the padded route it replaced.
func TestPoissonLineTransformAccuracy(t *testing.T) {
for _, tc := range []struct {
n int
sine bool
label string
}{
{8, true, "dst1-8"},
{31, true, "dst1-31"},
{64, true, "dst1-64"},
{8, false, "dct1-8"},
{31, false, "dct1-31"},
{64, false, "dct1-64"},
} {
x := lineFixture(tc.n)
want := refLineTransform(x, tc.sine)
wantMax := 0.0
for _, w := range want {
f, _ := w.Float64()
wantMax = max(wantMax, math.Abs(f))
}
if wantMax == 0 {
t.Fatalf("%s: reference is identically zero", tc.label)
}
scale := 0.0
for _, v := range x {
scale = max(scale, math.Abs(v))
}
errOf := func(got []float64) float64 {
worst := 0.0
for k := range tc.n {
g, _ := want[k].Float64()
worst = max(worst, math.Abs(got[k]-g))
}
return worst / (wantMax)
}
legacy := legacyLine(t, x, tc.sine)
legacyErr := errOf(legacy)
chirp := make([]float64, tc.n)
newLineTransformPlan(tc.n, tc.sine).apply(chirp, x)
chirpErr := errOf(chirp)
t.Logf("%s (max|x| = %.3g): padded route %.3g, chirp route %.3g, relative to max|transform|",
tc.label, scale, legacyErr, chirpErr)
if chirpErr > 1e-13 {
t.Fatalf("%s: chirp route relative error %.3g above the 1e-13 floor", tc.label, chirpErr)
}
if chirpErr > max(4*legacyErr, 4e-15) {
t.Fatalf("%s: chirp route error %.3g against padded route %.3g, beyond the 4x margin",
tc.label, chirpErr, legacyErr)
}
}
}
// TestPoissonLineTransformAgreementAtSolveSizes measures how far the
// chirp route sits from the padded route at the grid lengths the
// benchmarks and the solvers actually use. Two correct routes to the
// same sums disagree at rounding level; the bound pins that, and the
// log records the movement for the report.
func TestPoissonLineTransformAgreementAtSolveSizes(t *testing.T) {
for _, tc := range []struct {
n int
sine bool
label string
}{
{255, true, "dst1-255"},
{256, true, "dst1-256"},
{510, true, "dst1-510"},
{511, true, "dst1-511"},
{256, false, "dct1-256"},
{510, false, "dct1-510"},
} {
x := lineFixture(tc.n)
legacy := legacyLine(t, x, tc.sine)
chirp := make([]float64, tc.n)
newLineTransformPlan(tc.n, tc.sine).apply(chirp, x)
legacyMax := 0.0
for _, v := range legacy {
legacyMax = max(legacyMax, math.Abs(v))
}
delta := 0.0
for k := range tc.n {
delta = max(delta, math.Abs(chirp[k]-legacy[k]))
}
rel := delta / max(legacyMax, 1e-300)
t.Logf("%s: max|chirp − padded| = %.3g, relative %.3g", tc.label, delta, rel)
if rel > 1e-11 {
t.Fatalf("%s: routes disagree at relative %.3g, above the 1e-11 rounding bound", tc.label, rel)
}
}
}
// poissonLegacyTransform runs the pre-chirp block transform through
// the public DST/DCT pair, exactly as the solves' old pipeline did.
func poissonLegacyTransform(t *testing.T, f *core.Array, row0, col0, rows, cols int, sine bool) []float64 {
srcVals := poissonFloats(f)
srcCols := f.Shape()[1]
work := make([]float64, rows*cols)
apply := func(in []float64) []float64 { return legacyLine(t, in, sine) }
row := make([]float64, cols)
for r := range rows {
for c := range cols {
row[c] = srcVals[(r+row0)*srcCols+(c+col0)]
}
if sine && cols == 1 {
work[r*cols] = row[0]
continue
}
copy(work[r*cols:(r+1)*cols], apply(row))
}
col := make([]float64, rows)
for c := range cols {
for r := range rows {
col[r] = work[r*cols+c]
}
if sine && rows == 1 {
continue
}
out := apply(col)
for r := range rows {
work[r*cols+c] = out[r]
}
}
return work
}
// TestPoissonSolveRoutesAgree compares the shipped solves against the
// same solves rebuilt on the public DST/DCT pipeline, on deterministic
// grids of two sizes. The routes answer the same mathematics through
// different rounding; the drift must stay at rounding level relative to
// the solution's own scale, and the measured value is logged.
func TestPoissonSolveRoutesAgree(t *testing.T) {
for _, n := range []int{32, 64} {
fVals := make([]float64, n*n)
state := uint64(0x123456789abcdef) ^ uint64(n)
for i := range n * n {
state ^= state << 13
state ^= state >> 7
state ^= state << 17
fVals[i] = float64(int64(state)%2000)/1000 - 1
}
f, err := core.FromFloats(fVals, n, n)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
got, err := SolvePoissonDirichlet(f, 1, 1)
if err != nil {
t.Fatalf("SolvePoissonDirichlet: %v", err)
}
// The legacy pipeline on the same input: interior DST-I pair
// with the same eigenvalue division the solve runs at lx = ly
// = 1, where hx = hy = 1/(n−1).
spectrum := poissonLegacyTransform(t, f, 1, 1, n-2, n-2, true)
interiorC, interiorR := n-2, n-2
h := 1 / float64(n-1)
eigX := make([]float64, interiorC+1)
for kx := 1; kx <= interiorC; kx++ {
eigX[kx] = 2 / (h * h) * (1 - math.Cos(math.Pi*float64(kx)/float64(interiorC+1)))
}
eigY := make([]float64, interiorR+1)
for ky := 1; ky <= interiorR; ky++ {
eigY[ky] = 2 / (h * h) * (1 - math.Cos(math.Pi*float64(ky)/float64(interiorR+1)))
}
for ky := 1; ky <= interiorR; ky++ {
row := (ky - 1) * interiorC
for kx := 1; kx <= interiorC; kx++ {
spectrum[row+kx-1] /= eigX[kx] + eigY[ky]
}
}
back := poissonLegacyTransform(t, mustFloats(t, spectrum, interiorR, interiorC), 0, 0, interiorR, interiorC, true)
gotVals := poissonFloats(got)
scale := 0.0
for _, v := range back {
scale = max(scale, math.Abs(v))
}
delta := 0.0
for r := range interiorR {
for c := range interiorC {
delta = max(delta, math.Abs(gotVals[(r+1)*n+(c+1)]-back[r*interiorC+c]))
}
}
rel := delta / max(scale, 1e-300)
t.Logf("Dirichlet %dx%d: max|chirp solve − padded solve| = %.3g, relative %.3g", n, n, delta, rel)
if rel > 1e-11 {
t.Fatalf("Dirichlet %dx%d: solve routes disagree at relative %.3g", n, n, rel)
}
// The solve is deterministic: the same input answers the same
// bits on a second call.
again, err := SolvePoissonDirichlet(f, 1, 1)
if err != nil {
t.Fatalf("second SolvePoissonDirichlet: %v", err)
}
av := poissonFloats(again)
for i := range gotVals {
if av[i] != gotVals[i] {
t.Fatalf("Dirichlet %dx%d: repeated solve moved a bit at %d", n, n, i)
}
}
}
}
// BenchmarkPoissonLineTransform measures one line transform of the
// lengths the 512×512 Dirichlet and 256×256 Neumann solves carry,
// plan construction included, which is what every line of those grids
// pays.
func BenchmarkPoissonLineTransform(b *testing.B) {
for _, tc := range []struct {
n int
sine bool
name string
}{
{510, true, "dst1-510"},
{256, false, "dct1-256"},
} {
x := make([]float64, tc.n)
for i := range x {
x[i] = math.Sin(float64(i)) + 0.25*math.Cos(3*float64(i))
}
dst := make([]float64, tc.n)
b.Run(tc.name, func(b *testing.B) {
b.ReportAllocs()
for b.Loop() {
newLineTransformPlan(tc.n, tc.sine).apply(dst, x)
}
})
}
}