Files
tensor/grad/spectral.go
T

247 lines
7.3 KiB
Go
Raw 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 grad
import (
"sourcedock.dev/petrbalvin/tensor/internal/core"
"sourcedock.dev/petrbalvin/tensor/signal"
)
// Spectral autograd: the Fourier transforms as graph nodes,
// so deconvolution, spectral de-noising and frequency-domain fitting
// differentiate end to end. The adjoint of the unnormalised forward
// DFT y = F·z is dz = Fᴴ·g = n·IFFT(g) in the Wirtinger convention
// (the conjugate transpose falls out of dz = 2Re[ḡᵀ·dy] exactly the
// way the MatMul adjoint does); the inverse transform is its mirror.
// Real inputs flow through unchanged: signal.FFT widens them to
// complex, and the engine's complex-to-real narrowing (2·Re) is
// precisely the adjoint of that widening.
// FFT is the forward discrete Fourier transform of a rank-1 tensor;
// the backward multiplies the incoming gradient by Fᴴ, which is the
// inverse transform scaled by n.
func (t *Tensor) FFT() (*Tensor, error) {
if err := t.checkDiff("FFT"); err != nil {
return nil, err
}
if t.data.NDim() != 1 {
return nil, errf("autograd FFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.FFT(t.data)
if err != nil {
return nil, err
}
scale := complex(float64(t.data.Len()), 0)
sh := t.data.Shape()
return t.unaryResult("FFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
inv, err := signal.IFFT(g.arr)
if err != nil {
return err
}
cs := inv.RawComplexes()
for i := range cs {
cs[i] *= scale
}
dst[0] = gradSlot{arr: inv, sh: sh}
return nil
}), nil
}
// IFFT is the inverse transform of a rank-1 tensor; the backward runs
// the forward transform scaled by 1/n.
func (t *Tensor) IFFT() (*Tensor, error) {
if err := t.checkDiff("IFFT"); err != nil {
return nil, err
}
if t.data.NDim() != 1 {
return nil, errf("autograd IFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.IFFT(t.data)
if err != nil {
return nil, err
}
scale := complex(1/float64(t.data.Len()), 0)
sh := t.data.Shape()
return t.unaryResult("IFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
fwd, err := signal.FFT(g.arr)
if err != nil {
return err
}
cs := fwd.RawComplexes()
for i := range cs {
cs[i] *= scale
}
dst[0] = gradSlot{arr: fwd, sh: sh}
return nil
}), nil
}
// FFT2 is the 2-D forward transform; the backward is the 2-D inverse
// scaled by H·W, the total element count.
func (t *Tensor) FFT2() (*Tensor, error) {
if err := t.checkDiff("FFT2"); err != nil {
return nil, err
}
if t.data.NDim() != 2 {
return nil, errf("autograd FFT2: needs a rank-2 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.FFT2(t.data)
if err != nil {
return nil, err
}
scale := complex(float64(t.data.Len()), 0)
sh := t.data.Shape()
return t.unaryResult("FFT2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
inv, err := signal.IFFT2(g.arr)
if err != nil {
return err
}
cs := inv.RawComplexes()
for i := range cs {
cs[i] *= scale
}
dst[0] = gradSlot{arr: inv, sh: sh}
return nil
}), nil
}
// IFFT2 is the 2-D inverse transform; the backward is the 2-D forward
// scaled by 1/(H·W).
func (t *Tensor) IFFT2() (*Tensor, error) {
if err := t.checkDiff("IFFT2"); err != nil {
return nil, err
}
if t.data.NDim() != 2 {
return nil, errf("autograd IFFT2: needs a rank-2 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.IFFT2(t.data)
if err != nil {
return nil, err
}
scale := complex(1/float64(t.data.Len()), 0)
sh := t.data.Shape()
return t.unaryResult("IFFT2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
fwd, err := signal.FFT2(g.arr)
if err != nil {
return err
}
cs := fwd.RawComplexes()
for i := range cs {
cs[i] *= scale
}
dst[0] = gradSlot{arr: fwd, sh: sh}
return nil
}), nil
}
// RFFT is the real-input half-spectrum transform. The input must be a
// rank-1 real tensor; the backward folds the incoming half-spectrum
// gradient into dx = 2·Re(F_halfᴴ·g), evaluated by one padded forward
// FFT so the cost matches the forward transform. The factor 2 lands
// only on the mirrored bins through the zero padding, exactly the
// combinatorics the derivation gives.
func (t *Tensor) RFFT() (*Tensor, error) {
if err := t.checkDiff("RFFT"); err != nil {
return nil, err
}
if isComplexArr(t.data) {
return nil, errf("autograd RFFT: needs a real tensor, got %s", t.data.Dtype())
}
if t.data.NDim() != 1 {
return nil, errf("autograd RFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.RFFT(t.data)
if err != nil {
return nil, err
}
in := t.data
return t.unaryResult("RFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
n := in.Len()
half := n/2 + 1
// The adjoint needs Σ_{k<half} g_k·e^{+2πijk/n}, a +sign DFT
// of the zero-padded gradient: conj(FFT(conj(·))).
pad := make([]complex128, n)
for k := range half {
pad[k] = conj(g.arr.ComplexAt(k))
}
padArr, err := core.ComplexFromArray(pad, n)
if err != nil {
return err
}
spec, err := signal.FFT(padArr)
if err != nil {
return err
}
sh := in.Shape()
dx := gradSlot{arr: ar.borrowGrad(in.Dtype(), sh), sh: sh}
// The transform's output and dx are both freshly allocated and
// dense, so the doubled real part is taken from the payload.
ss := spec.RawComplexes()
if in.Dtype() == core.Float32 {
ds := dx.arr.RawFloat32s()[:dx.arr.Len()]
for j := range n {
ds[j] = float32(2 * real(ss[j]))
}
} else {
ds := dx.arr.RawFloats()[:dx.arr.Len()]
for j := range n {
ds[j] = 2 * real(ss[j])
}
}
dst[0] = dx
return nil
}), nil
}
// IRFFT is the inverse half-spectrum transform: a rank-1 complex
// tensor of n/2+1 bins into a real signal of length n. The backward
// widens the real gradient to a full forward FFT and halves the
// self-mirrored bins (DC, and Nyquist when n is even): dIn_k =
// FFT(g)_k/n on the ordinary bins and half of that on the mirrored
// bins, the transpose of the Hermitian extension the forward performs.
func (t *Tensor) IRFFT(n int) (*Tensor, error) {
if !isComplexArr(t.data) {
return nil, errf("autograd IRFFT: needs a complex tensor, got %s", t.data.Dtype())
}
if t.data.NDim() != 1 {
return nil, errf("autograd IRFFT: needs a rank-1 tensor, got shape %s", prettyShape(t.data.Shape()))
}
out, err := signal.IRFFT(t.data, n)
if err != nil {
return nil, err
}
in := t.data
return t.unaryResult("IRFFT", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
gs := make([]complex128, n)
for j := range n {
gs[j] = complex(g.arr.FloatAt(j), 0)
}
gArr, err := core.ComplexFromArray(gs, n)
if err != nil {
return err
}
spec, err := signal.FFT(gArr)
if err != nil {
return err
}
half := in.Len()
full := complex(float64(n), 0)
sh := []int{half}
dx := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh}
ds := dx.arr.RawComplexes()[:half]
ss := spec.RawComplexes()
ds[0] = ss[0] / (2 * full)
for k := 1; k < half; k++ {
if n%2 == 0 && k == half-1 {
// Nyquist mirrors itself.
ds[k] = ss[k] / (2 * full)
continue
}
ds[k] = ss[k] / full
}
dst[0] = dx
return nil
}), nil
}