// Copyright (c) 2026 Petr Balvín (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