Files
tensor/signal/nufft_test.go
T

120 lines
3.4 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 (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
import (
"math"
"math/cmplx"
"testing"
)
// nufftCase builds a deterministic nonuniform problem: coordinates
// spread over the band, complex values with structure at several
// scales.
func nufftCase(count int) ([]float64, []complex128) {
x := make([]float64, count)
c := make([]complex128, count)
for j := range count {
x[j] = -0.5 + float64(j)*0.97/float64(count)
c[j] = cmplx.Exp(complex(0.3*float64(j%7), 0.9*float64(j%5)))
}
return x, c
}
// directDFT evaluates the type-1 sum the slow honest way.
func directDFT(x []float64, c []complex128, n int) []complex128 {
out := make([]complex128, n)
for k := range n {
s := complex(0, 0)
for j := range x {
s += c[j] * cmplx.Exp(complex(0, 2*math.Pi*float64(k)*x[j]))
}
out[k] = s
}
return out
}
// TestNUFFTType1MatchesDirectDFT checks the gridded transform against
// the direct sum over the whole output band: the Gaussian kernel's
// guarantee is a handful of relative digits at every frequency, not
// just the low ones.
func TestNUFFTType1MatchesDirectDFT(t *testing.T) {
const count, n = 64, 64
x, c := nufftCase(count)
xa, cerr := core.FromFloats(x, count)
if cerr != nil {
t.Fatal(cerr)
}
ca, cerr := core.FromComplexes(c, count)
if cerr != nil {
t.Fatal(cerr)
}
got, err := NUFFTType1(xa, ca, n)
if err != nil {
t.Fatalf("NUFFTType1: %v", err)
}
want := directDFT(x, c, n)
worst, scale := 0.0, 0.0
for k := range n {
d := base.AbsComplex(got.ComplexAt(k) - want[k])
worst = math.Max(worst, d)
scale = math.Max(scale, base.AbsComplex(want[k]))
}
if worst/scale > 1e-4 {
t.Fatalf("worst relative error %.3g, want under 1e-4", worst/scale)
}
}
// TestNUFFTType1LargeGrid covers an output grid wider than the sample
// count, the regime interferometric imaging actually runs.
func TestNUFFTType1LargeGrid(t *testing.T) {
const count, n = 24, 96
x := make([]float64, count)
c := make([]complex128, count)
for j := range count {
x[j] = -0.5 + (float64(j)*0.77+0.1)/float64(count)
c[j] = complex(math.Cos(float64(j)), math.Sin(float64(2*j)))
}
xa, _ := core.FromFloats(x, count)
ca, _ := core.FromComplexes(c, count)
got, err := NUFFTType1(xa, ca, n)
if err != nil {
t.Fatalf("NUFFTType1: %v", err)
}
want := directDFT(x, c, n)
worst := 0.0
for k := range n {
worst = math.Max(worst, base.AbsComplex(got.ComplexAt(k)-want[k]))
}
if worst > 1e-4 {
t.Fatalf("worst absolute error %.3g, want under 1e-4", worst)
}
}
// TestNUFFTType1Errors pins the validation contract.
func TestNUFFTType1Errors(t *testing.T) {
xa := mustFloats(t, []float64{-0.4, 0.1, 0.3}, 3)
ca, _ := core.FromComplexes([]complex128{1, 1i, 2}, 3)
if _, err := NUFFTType1(xa, ca, 0); err == nil {
t.Fatal("expected an error for a zero output size")
}
ca2, _ := core.FromComplexes([]complex128{1, 1i}, 2)
if _, err := NUFFTType1(xa, ca2, 8); err == nil {
t.Fatal("expected an error for mismatched lengths")
}
outside := mustFloats(t, []float64{-0.6, 0.1, 0.3}, 3)
if _, err := NUFFTType1(outside, ca, 8); err == nil {
t.Fatal("expected an error for a coordinate below −1/2")
}
atEdge := mustFloats(t, []float64{-0.4, 0.1, 0.5}, 3)
if _, err := NUFFTType1(atEdge, ca, 8); err == nil {
t.Fatal("expected an error for a coordinate at +1/2")
}
}