// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "math" "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // spectralWeights builds a deterministic complex weight vector used to // fold a spectrum into a real scalar loss. func spectralWeights(n int, seed int) []complex128 { g := core.NewGenerator(int64(seed)) w := make([]complex128, n) for i := range n { w[i] = complex(g.NormalUnit(), g.NormalUnit()) } return w } // complexLossOf runs the op chain and folds the result into the real // scalar Σ Re(w·y) + Σ|y|²/len, the same fold the graph's foldReal // builds from Mul and Real, so oracle and graph define one loss. func complexLossOf(t *testing.T, op func(*Tensor) (*Tensor, error), w []complex128) func(*core.Array) float64 { t.Helper() return func(a *core.Array) float64 { y, err := op(FromArray(a, false)) if err != nil { t.Fatalf("forward: %v", err) } s := 0.0 for i := range y.Data().Len() { z := y.Data().ComplexAt(i) s += real(w[i%len(w)] * z) s += (real(z)*real(z) + imag(z)*imag(z)) / float64(y.Data().Len()) } return s } } // realLossOf is complexLossOf for chains that end in a real tensor. func realLossOf(t *testing.T, op func(*Tensor) (*Tensor, error), w []complex128) func(*core.Array) float64 { t.Helper() return func(a *core.Array) float64 { y, err := op(FromArray(a, false)) if err != nil { t.Fatalf("forward: %v", err) } s := 0.0 for i := range y.Data().Len() { v := y.Data().FloatAt(i) s += real(w[i%len(w)])*v + v*v/float64(y.Data().Len()) } return s } } // backwardComplex runs op on a fresh tensor over vals and returns the // leaf after Backward. func backwardComplex(t *testing.T, op func(*Tensor) (*Tensor, error), vals []complex128, shape ...int) *Tensor { t.Helper() a, err := core.FromComplexes(vals, shape...) if err != nil { t.Fatalf("FromComplexes: %v", err) } xt := FromArray(a, true) y, err := op(xt) if err != nil { t.Fatalf("forward: %v", err) } // Fold to a real scalar so Backward has its seed. w := spectralWeights(y.Data().Len(), 11) var loss *Tensor loss, err = foldReal(t, y, w) if err != nil { t.Fatalf("fold: %v", err) } if err := loss.Backward(); err != nil { t.Fatalf("Backward: %v", err) } return xt } // foldReal reduces a complex tensor to Σ Re(w̄·y) + Σ|y|²/len through // the graph ops, and a real tensor to the analogous real fold. func foldReal(t *testing.T, y *Tensor, w []complex128) (*Tensor, error) { t.Helper() n := y.Data().Len() wa, err := core.FromComplexes(w, y.Data().Shape()...) if err != nil { return nil, err } wt := FromArray(wa, false) if isComplexArr(y.Data()) { prod, err := y.Mul(wt) if err != nil { return nil, err } re, err := prod.Real() if err != nil { return nil, err } s1, err := re.Sum() if err != nil { return nil, err } sq, err := y.Abs2() if err != nil { return nil, err } s2, err := sq.Sum() if err != nil { return nil, err } s2s, err := s2.Scale(1 / float64(n)) if err != nil { return nil, err } return s1.Add(s2s) } prod, err := y.Mul(wt) if err != nil { return nil, err } re, err := prod.Real() if err != nil { return nil, err } s1, err := re.Sum() if err != nil { return nil, err } sq, err := y.Abs2() if err != nil { return nil, err } s2, err := sq.Sum() if err != nil { return nil, err } s2s, err := s2.Scale(1 / float64(n)) if err != nil { return nil, err } return s1.Add(s2s) } // TestGradFFTWirtinger pins the FFT adjoint against central // differences on a power-of-two and a Bluestein length. func TestGradFFTWirtinger(t *testing.T) { for _, n := range []int{8, 12} { g := core.NewGenerator(int64(n)) vals := make([]complex128, n) for i := range n { vals[i] = complex(g.NormalUnit(), g.NormalUnit()) } op := func(x *Tensor) (*Tensor, error) { return x.FFT() } xt := backwardComplex(t, op, vals, n) loss := complexLossOf(t, op, spectralWeights(n, 11)) checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7) } } // TestGradIFFTWirtinger pins the IFFT adjoint. func TestGradIFFTWirtinger(t *testing.T) { n := 10 g := core.NewGenerator(3) vals := make([]complex128, n) for i := range n { vals[i] = complex(g.NormalUnit(), g.NormalUnit()) } op := func(x *Tensor) (*Tensor, error) { return x.IFFT() } xt := backwardComplex(t, op, vals, n) loss := complexLossOf(t, op, spectralWeights(n, 11)) checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7) } // TestGradFFT2Wirtinger pins the 2-D adjoint. func TestGradFFT2Wirtinger(t *testing.T) { rows, cols := 3, 4 g := core.NewGenerator(5) vals := make([]complex128, rows*cols) for i := range vals { vals[i] = complex(g.NormalUnit(), g.NormalUnit()) } op := func(x *Tensor) (*Tensor, error) { return x.FFT2() } xt := backwardComplex(t, op, vals, rows, cols) loss := complexLossOf(t, op, spectralWeights(rows*cols, 11)) checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7) } // TestGradFFTRealInput pins the 2·Re narrowing path: a real leaf under // a complex FFT node. func TestGradFFTRealInput(t *testing.T) { n := 8 g := core.NewGenerator(9) vals := make([]float64, n) for i := range n { vals[i] = g.NormalUnit() } op := func(x *Tensor) (*Tensor, error) { return x.FFT() } a, err := core.FromFloats(vals, n) if err != nil { t.Fatalf("FromFloats: %v", err) } xt := FromArray(a, true) y, err := op(xt) if err != nil { t.Fatalf("forward: %v", err) } loss, err := foldReal(t, y, spectralWeights(n, 11)) if err != nil { t.Fatalf("fold: %v", err) } if err := loss.Backward(); err != nil { t.Fatalf("Backward: %v", err) } lossOf := realToComplexLoss(t, op, n) want := numericGrad(lossOf, a) for i := range n { if math.Abs(xt.Grad().FloatAt(i)-want[i]) > 1e-6 { t.Fatalf("grad[%d] = %g, want %g", i, xt.Grad().FloatAt(i), want[i]) } } } // realToComplexLoss adapts a chain over real input for numericGrad: it // re-runs the forward and folds the (complex) output into a scalar. func realToComplexLoss(t *testing.T, op func(*Tensor) (*Tensor, error), n int) func(*core.Array) float64 { t.Helper() w := spectralWeights(n, 11) return func(a *core.Array) float64 { y, err := op(FromArray(a, false)) if err != nil { t.Fatalf("forward: %v", err) } s := 0.0 for i := range n { z := y.Data().ComplexAt(i) s += real(w[i] * z) s += (real(z)*real(z) + imag(z)*imag(z)) / float64(n) } return s } } // TestGradRFFT pins the half-spectrum adjoint, even and odd lengths. func TestGradRFFT(t *testing.T) { for _, n := range []int{8, 9} { g := core.NewGenerator(int64(n * 2)) vals := make([]float64, n) for i := range n { vals[i] = g.NormalUnit() } a, err := core.FromFloats(vals, n) if err != nil { t.Fatalf("FromFloats: %v", err) } xt := FromArray(a, true) y, err := xt.RFFT() if err != nil { t.Fatalf("RFFT: %v", err) } w := spectralWeights(y.Data().Len(), 11) var loss *Tensor loss, err = foldReal(t, y, w) if err != nil { t.Fatalf("fold: %v", err) } if err := loss.Backward(); err != nil { t.Fatalf("Backward: %v", err) } lossOf := func(a *core.Array) float64 { yp, err := FromArray(a, false).RFFT() if err != nil { t.Fatalf("RFFT: %v", err) } s := 0.0 for i := range yp.Data().Len() { z := yp.Data().ComplexAt(i) s += real(w[i] * z) s += (real(z)*real(z) + imag(z)*imag(z)) / float64(yp.Data().Len()) } return s } want := numericGrad(lossOf, a) for i := range n { if math.Abs(xt.Grad().FloatAt(i)-want[i]) > 1e-6 { t.Fatalf("n=%d grad[%d] = %g, want %g", n, i, xt.Grad().FloatAt(i), want[i]) } } } } // TestGradIRFFT pins the half-spectrum inverse adjoint, on an even // length (Nyquist half-weight bin) and an odd one (the last bin is an // ordinary mirrored bin with the full 1/n weight). func TestGradIRFFT(t *testing.T) { for _, n := range []int{8, 9} { half := n/2 + 1 g := core.NewGenerator(21) vals := make([]complex128, half) for i := range half { vals[i] = complex(g.NormalUnit(), g.NormalUnit()) } op := func(x *Tensor) (*Tensor, error) { return x.IRFFT(n) } xt := backwardComplex(t, op, vals, half) // numericGrad over complex perturbations, folded through the real // output. w := spectralWeights(n, 11) loss := func(a *core.Array) float64 { y, err := FromArray(a, false).IRFFT(n) if err != nil { t.Fatalf("IRFFT: %v", err) } s := 0.0 for i := range n { v := y.Data().FloatAt(i) s += real(w[i])*v + v*v/float64(n) } return s } checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(loss, xt.Data()), 1e-7) } } // TestGradFFTRoundtripIdentity pins the composition: gradient through // IFFT∘FFT must arrive unchanged (Fᴴ·(1/n)F = I). func TestGradFFTRoundtripIdentity(t *testing.T) { n := 8 vals := make([]complex128, n) g := core.NewGenerator(4) for i := range n { vals[i] = complex(g.NormalUnit(), g.NormalUnit()) } a, err := core.FromComplexes(vals, n) if err != nil { t.Fatalf("FromComplexes: %v", err) } xt := FromArray(a, true) f, err := xt.FFT() if err != nil { t.Fatalf("FFT: %v", err) } fi, err := f.IFFT() if err != nil { t.Fatalf("IFFT: %v", err) } w := spectralWeights(n, 7) loss, err := foldReal(t, fi, w) if err != nil { t.Fatalf("fold: %v", err) } if err := loss.Backward(); err != nil { t.Fatalf("Backward: %v", err) } // dL/dy at the roundtrip output, propagated through both adjoints, // must equal dL/dy itself: the fold's Wirtinger gradient is // w̄/2 + y/n (Mul+Real contributes w̄/2, Abs2/n contributes y/n). for i := range n { dy := conj(w[i])/2 + fi.Data().ComplexAt(i)/complex(float64(n), 0) got := xt.Grad().ComplexAt(i) if cmplxAbs(got-dy) > 1e-8 { t.Fatalf("grad[%d] = %v, want %v", i, got, dy) } } }