feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,246 @@
|
||||
// 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
|
||||
}
|
||||
Reference in New Issue
Block a user