113 lines
4.2 KiB
Go
113 lines
4.2 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"
|
||
|
||
// 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
|
||
}
|