371 lines
9.8 KiB
Go
371 lines
9.8 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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)
|
|
}
|
|
}
|
|
}
|