120 lines
3.4 KiB
Go
120 lines
3.4 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/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")
|
||
}
|
||
}
|