247 lines
7.3 KiB
Go
247 lines
7.3 KiB
Go
// 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
|
|
}
|