Files

113 lines
4.2 KiB
Go
Raw Permalink 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"
// Spectral solution of the Poisson equation on a periodic domain. The
// Laplacian is diagonal in the Fourier basis: each plane wave
// e^{i(kx·x + ky·y)} is an eigenfunction of −Δ with eigenvalue
// (2πkx/Lx)² + (2πky/Ly)², so solving −Δu = f on the torus is one
// forward FFT, one division per mode and one inverse FFT. The
// convergence is spectral: for a smooth right-hand side the error
// falls faster than any power of the grid spacing, which no finite
// difference stencil matches at comparable cost.
// SolvePoissonPeriodic solves −Δu = f on the periodic square
// [0, Lx] × [0, Ly] sampled on the (rows, cols) grid carried by the
// shape of f, and returns u as a float array of the same shape. Row r
// of f samples y = r·Ly/rows, column c samples x = c·Lx/cols. f must
// be a float64 array (the computation runs in float64 end to end and
// the solution is exact in no narrower dtype), and a float32, int or
// complex f is an error.
//
// A periodic solution exists only when f has zero mean: the zero
// Fourier mode of −Δu vanishes, so a nonzero mean is an error rather
// than a quietly shifted problem (the round-off band around zero is
// accepted and the constant mode of u is set to zero). A rank other
// than 2, a grid dimension below 2 or a non-positive side length is
// an error.
func SolvePoissonPeriodic(f *core.Array, lx, ly float64) (*core.Array, error) {
const name = "SolvePoissonPeriodic"
if f.NDim() != 2 {
return nil, base.Errf("%s: f must be rank 2, got shape %s", name, base.ShapeText(f.Shape()))
}
if f.Dtype() != core.Float {
return nil, base.Errf("%s: f must be a float64 array, got %s", name, f.Dtype())
}
rows, cols := f.Shape()[0], f.Shape()[1]
if rows < 2 || cols < 2 {
return nil, base.Errf("%s: the grid must be at least 2×2, got %d×%d", name, rows, cols)
}
if lx <= 0 || ly <= 0 {
return nil, base.Errf("%s: the side lengths must be positive, got %g and %g", name, lx, ly)
}
spectrum, err := FFT2(f)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
wave := func(n int, i int, l float64) float64 {
k := i
if k > n/2 {
k -= n
}
s := 2 * math.Pi * float64(k) / l
return s * s
}
// A non-finite source would drive the compatibility tolerance NaN
// and slip through its gate, so it is refused up front.
for i := range rows * cols {
if v := f.FloatAt(i); math.IsInf(v, 0) || math.IsNaN(v) {
return nil, base.Errf("%s: f holds the non-finite value %g", name, v)
}
}
largest := 0.0
for _, v := range poissonFloats(f) {
largest = max(largest, math.Abs(v))
}
// The DFT of the constant mode is the plain sum of the samples;
// anything above the accumulated round-off breaks solvability. The
// tolerance is relative to the source: the sum scales with the
// sample magnitude and the element count, so an absolute floor would
// accept a tiny source with a constant mode far above its own size.
meanTolerance := 64 * base.EpsF * float64(rows*cols) * largest
mean := spectrum.ComplexAt(0)
if math.Abs(real(mean)) > meanTolerance || math.Abs(imag(mean)) > meanTolerance {
return nil, base.Errf("%s: f has a nonzero mean %g, the periodic problem has no solution",
name, real(mean)/float64(rows*cols))
}
adjusted := make([]complex128, rows*cols)
for r := range rows {
// Rows sample y = r·ly/rows, columns sample x = c·lx/cols, so
// the row wavenumber carries ly and the column wavenumber lx.
ly2 := wave(rows, r, ly)
for c := range cols {
lambda := ly2 + wave(cols, c, lx)
if lambda == 0 {
continue // the constant mode stays zero
}
adjusted[r*cols+c] = spectrum.ComplexAt(r*cols+c) / complex(lambda, 0)
}
}
spectrumArray, err := core.FromComplexes(adjusted, rows, cols)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
inverted, err := IFFT2(spectrumArray)
if err != nil {
return nil, base.Errf("%s: %w", name, err)
}
out := core.New(f.Dtype(), rows, cols)
vals := out.RawFloats()
for i := range vals {
vals[i] = real(inverted.ComplexAt(i))
}
return out, nil
}