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