Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

166 lines
4.1 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Command deconv recovers a sharp image from a blurred, noisy
// observation by gradient descent through the Fourier transform: the
// convolution runs as a spectral product, the loss differentiates
// through FFT2 and IFFT2, and Tikhonov regularisation keeps the noise
// from winning. Astronomical PSF deconvolution in a page of gradient,
// the workload the spectral autograd exists for.
//
// Usage: go run ./examples/deconv
package main
import (
"fmt"
"log"
"math"
"sourcedock.dev/petrbalvin/tensor"
)
func main() {
const (
n = 32
sigma = 6.0
)
// The truth: one off-centre Gaussian source.
truth := make([]float64, n*n)
for r := range n {
for c := range n {
dx := float64(c) - 20.0
dy := float64(r) - 12.0
truth[r*n+c] = math.Exp(-(dx*dx + dy*dy) / (2 * 2.5 * 2.5))
}
}
// The PSF: a wider Gaussian, the blur to undo.
psf := make([]float64, n*n)
for r := range n {
for c := range n {
dx := float64(c) - n/2
dy := float64(r) - n/2
psf[r*n+c] = math.Exp(-(dx*dx + dy*dy) / (2 * sigma * sigma))
}
}
obsArr, err := tensor.FromFloats(truth, n, n)
if err != nil {
log.Fatal(err)
}
psfArr, err := tensor.FromFloats(psf, n, n)
if err != nil {
log.Fatal(err)
}
// The blur runs as a spectral product; the constants ride the same
// graph nodes without requiring grad.
kT := tensor.FromArray(psfArr, false)
kF, err := kT.FFT2()
if err != nil {
log.Fatal(err)
}
obsT := tensor.FromArray(obsArr, false)
spec, err := obsT.FFT2()
if err != nil {
log.Fatal(err)
}
prod, err := spec.Mul(kF)
if err != nil {
log.Fatal(err)
}
blurred, err := prod.IFFT2()
if err != nil {
log.Fatal(err)
}
noise := tensor.NewGenerator(42)
noisy := make([]float64, n*n)
for i := range noisy {
noisy[i] = real(blurred.Data().ComplexAt(i)) + 0.01*noise.NormalUnit()
}
obsC, err := tensor.FromFloats(noisy, n, n)
if err != nil {
log.Fatal(err)
}
obsComplex, err := tensor.Astype(obsC, tensor.Complex)
if err != nil {
log.Fatal(err)
}
obsTt := tensor.FromArray(obsComplex, false)
// The spectral loss is a stiff quadratic (curvature ~ |K|² per
// frequency), exactly the landscape plain gradient descent crawls
// on and Newton-CG eats: the CG solve rides the Hessian-vector
// product through the same FFT chain, two backward passes per
// iteration, no dense Hessian ever formed.
xArr, err := tensor.FromFloats(make([]float64, n*n), n, n)
if err != nil {
log.Fatal(err)
}
objective := func(x *tensor.Tensor) (*tensor.Tensor, error) {
spec, err := x.FFT2()
if err != nil {
return nil, err
}
prod, err := spec.Mul(kF)
if err != nil {
return nil, err
}
model, err := prod.IFFT2()
if err != nil {
return nil, err
}
resid, err := model.Sub(obsTt)
if err != nil {
return nil, err
}
dataTerm, err := resid.Abs2()
if err != nil {
return nil, err
}
regTerm, err := x.Abs2()
if err != nil {
return nil, err
}
reg, err := regTerm.Scale(1e-3)
if err != nil {
return nil, err
}
both, err := dataTerm.Add(reg)
if err != nil {
return nil, err
}
return both.Sum()
}
solution, loss, err := tensor.MinimiseNewtonCG(objective, xArr,
tensor.NewtonCGOptions{Tolerance: 1e-8, MaxIterations: 80})
if err != nil {
log.Fatal(err)
}
xArr = solution
fmt.Printf("newton-cg converged, loss = %.6e\n", loss)
peakOf := func(a *tensor.Array) (int, int, float64) {
best := math.Inf(-1)
br, bc := 0, 0
for r := range n {
for c := range n {
if v := a.FloatAt(r*n + c); v > best {
best, br, bc = v, r, c
}
}
}
return br, bc, best
}
tr, tc, tv := peakOf(obsArr)
br, bc, bv := peakOf(obsC)
rr, rc, rv := peakOf(xArr)
fmt.Printf("truth peak at (%d, %d), height %.3f\n", tr, tc, tv)
fmt.Printf("blurred observation peak at (%d, %d), height %.3f\n", br, bc, bv)
fmt.Printf("recovered peak at (%d, %d), height %.3f (final loss %.3e)\n", rr, rc, rv, loss)
if rr != tr || rc != tc {
log.Fatal("the recovery did not localise the source")
}
if math.Abs(rv-tv) > 0.25*tv {
log.Fatal("the recovery did not restore the source height")
}
}