447 lines
14 KiB
Go
447 lines
14 KiB
Go
// 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)
|
||
}
|
||
})
|
||
}
|
||
}
|