166 lines
4.1 KiB
Go
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")
|
|
}
|
|
}
|