// Copyright (c) 2026 Petr Balvín (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) } }) } }