Files
tensor/grad/spectral_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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)
}
}
}