Files
tensor/grad/complex_test.go
T

540 lines
14 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"sourcedock.dev/petrbalvin/tensor/internal/core"
"testing"
)
// buildComplex wraps a complex array as a leaf tensor; the losses
// built on it reduce to a real scalar, so Backward has a real seed.
func buildComplex(t *testing.T, vals []complex128, shape ...int) *Tensor {
t.Helper()
a, err := core.FromComplexes(vals, shape...)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
return FromArray(a, true)
}
// numericComplexGrad estimates dL/dRe(z) and dL/dIm(z) by central
// differences; the Wirtinger gradient the graph reports must satisfy
// g = (dL/dRe + i·dL/dIm)/2 element-wise.
func numericComplexGrad(f func(*core.Array) float64, a *core.Array) []complex128 {
n := a.Len()
out := make([]complex128, n)
const h = 1e-6
up := make([]complex128, n)
down := make([]complex128, n)
for i := range n {
base := make([]complex128, n)
for j := range n {
base[j] = a.ComplexAt(j)
}
copy(up, base)
copy(down, base)
up[i] += complex(h, 0)
down[i] -= complex(h, 0)
au, _ := core.FromComplexes(up, a.Shape()...)
ad, _ := core.FromComplexes(down, a.Shape()...)
dRe := (f(au) - f(ad)) / (2 * h)
copy(up, base)
copy(down, base)
up[i] += complex(0, h)
down[i] -= complex(0, h)
au, _ = core.FromComplexes(up, a.Shape()...)
ad, _ = core.FromComplexes(down, a.Shape()...)
dIm := (f(au) - f(ad)) / (2 * h)
out[i] = complex(dRe/2, dIm/2)
}
return out
}
// checkAgainstNumeric compares a Wirtinger gradient with the central-
// difference reference.
func checkAgainstNumeric(t *testing.T, got *core.Array, want []complex128, tol float64) {
t.Helper()
for i, w := range want {
g := got.ComplexAt(i)
if math.Abs(real(g)-real(w)) > tol || math.Abs(imag(g)-imag(w)) > tol {
t.Fatalf("grad[%d] = %v, want %v", i, g, w)
}
}
}
// TestComplexMulGrad pins the Wirtinger adjoint of the element-wise
// product: dz = g·w̄.
func TestComplexMulGrad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i, -0.5 + 0.25i}, 3)
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i, 1 - 3i}, 3)
prod, err := z.Mul(w)
if err != nil {
t.Fatalf("Mul: %v", err)
}
re, err := prod.Real()
if err != nil {
t.Fatalf("Real: %v", err)
}
loss, err := re.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
s += real(a.ComplexAt(i) * w.Data().ComplexAt(i))
}
return s
}
want := numericComplexGrad(lossOf, z.Data())
checkAgainstNumeric(t, z.Grad(), want, 1e-8)
}
// TestComplexDivGrad pins the division adjoint da = g/b̄,
// db = −g·ā/b̄².
func TestComplexDivGrad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i}, 2)
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i}, 2)
q, err := z.Div(w)
if err != nil {
t.Fatalf("Div: %v", err)
}
im, err := q.Imag()
if err != nil {
t.Fatalf("Imag: %v", err)
}
loss, err := im.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
s += imag(a.ComplexAt(i) / w.Data().ComplexAt(i))
}
return s
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(lossOf, z.Data()), 1e-8)
checkAgainstNumeric(t, w.Grad(), numericComplexGrad(func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
s += imag(z.Data().ComplexAt(i) / a.ComplexAt(i))
}
return s
}, w.Data()), 1e-8)
}
// TestComplexQuantumExpectation pins the physics workhorse: the loss
// L = Re(ψ̄·(H·ψ)) with a Hermitian H, whose Wirtinger gradient is
// ψ̄-independent and equals H·ψ... evaluated against central
// differences rather than trust the algebra.
func TestComplexQuantumExpectation(t *testing.T) {
psi := buildComplex(t, []complex128{1 + 0.5i, -0.3 + 0.8i, 0.2 - 1.1i, 0.9 + 0.4i}, 4)
hDense := []complex128{
2, 0.5i, 0, -1,
-0.5i, 3, 1i, 0,
0, -1i, 1.5, 0.5,
-1, 0, 0.5, 2.5,
}
hArr, err := core.FromComplexes(hDense, 4, 4)
if err != nil {
t.Fatalf("FromComplexes: %v", err)
}
h := FromArray(hArr, false)
conj, err := psi.Conj()
if err != nil {
t.Fatalf("Conj: %v", err)
}
// Row vector (1,4) times H·ψ (4,) keeps every MatMul shape legal.
bra, err := conj.Reshape(1, 4)
if err != nil {
t.Fatalf("Reshape: %v", err)
}
hpsi, err := h.MatMul(psi)
if err != nil {
t.Fatalf("MatMul: %v", err)
}
prod, err := bra.MatMul(hpsi)
if err != nil {
t.Fatalf("MatMul: %v", err)
}
// prod is (1,1); Real then Sum flattens to the scalar loss.
r, err := prod.Real()
if err != nil {
t.Fatalf("Real: %v", err)
}
loss, err := r.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
expectation := func(a *core.Array) float64 {
s := 0.0
for i := range 4 {
var acc complex128
for j := range 4 {
acc += hDense[i*4+j] * a.ComplexAt(j)
}
s += real(complex(real(a.ComplexAt(i)), -imag(a.ComplexAt(i))) * acc)
}
return s
}
checkAgainstNumeric(t, psi.Grad(), numericComplexGrad(expectation, psi.Data()), 1e-7)
}
// TestComplexAbs2Grad pins |z|², dz = g·z.
func TestComplexAbs2Grad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, -3 + 0.5i}, 2)
sq, err := z.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// d|z|²/dz̄ = z exactly.
for i := range 2 {
if z.Grad().ComplexAt(i) != z.Data().ComplexAt(i) {
t.Fatalf("grad[%d] = %v, want %v", i, z.Grad().ComplexAt(i), z.Data().ComplexAt(i))
}
}
}
// TestComplexAbsGrad pins the magnitude gradient dz = g·z/(2|z|).
func TestComplexAbsGrad(t *testing.T) {
z := buildComplex(t, []complex128{3 + 4i, -1 + 1i}, 2)
m, err := z.Abs()
if err != nil {
t.Fatalf("Abs: %v", err)
}
loss, err := m.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
expectation := func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
s += cmplxAbs(a.ComplexAt(i))
}
return s
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(expectation, z.Data()), 1e-8)
}
func cmplxAbs(z complex128) float64 { return math.Hypot(real(z), imag(z)) }
// TestComplexBackwardRejectsComplexLoss pins the real-seed contract.
func TestComplexBackwardRejectsComplexLoss(t *testing.T) {
z := buildComplex(t, []complex128{1 + 1i}, 1)
if err := z.Backward(); err == nil {
t.Fatal("Backward accepted a complex output")
}
}
// TestComplexMixedRealLeaf pins the 2·Re narrowing: a real tensor
// multiplied into a complex chain gets the true real gradient.
func TestComplexMixedRealLeaf(t *testing.T) {
xArr, err := core.FromFloats([]float64{1.5, -0.5}, 2)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
x := FromArray(xArr, true)
w := buildComplex(t, []complex128{0.5 - 1i, 2 + 2i}, 2)
prod, err := x.Mul(w)
if err != nil {
t.Fatalf("Mul: %v", err)
}
sq, err := prod.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// L = Σ x²|w|², dL/dx = 2x|w|².
want := []float64{2 * 1.5 * (0.25 + 1), 2 * -0.5 * (4 + 4)}
for i, wv := range want {
if math.Abs(x.Grad().FloatAt(i)-wv) > 1e-12 {
t.Fatalf("x.grad[%d] = %g, want %g", i, x.Grad().FloatAt(i), wv)
}
}
}
// TestComplexSumAxisGrad pins the axis reduction on a complex leaf:
// the backward broadcasts the seed back over the dropped axis, and the
// Wirtinger gradient of a weighted real loss matches central
// differences.
func TestComplexSumAxisGrad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, -0.5 + 0.25i, 0.75 - 1.5i, 2 + 0.5i}, 2, 2)
out, err := z.SumAxis(1)
if err != nil {
t.Fatalf("SumAxis: %v", err)
}
if out.Data().NDim() != 1 || out.Data().Len() != 2 {
t.Fatalf("SumAxis shape = %s, want (2)", prettyShape(out.Data().Shape()))
}
w := buildComplex(t, []complex128{0.3 + 0.4i, -0.6 - 0.1i}, 2)
prod, err := out.Mul(w)
if err != nil {
t.Fatalf("Mul: %v", err)
}
re, err := prod.Real()
if err != nil {
t.Fatalf("Real: %v", err)
}
loss, err := re.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(a *core.Array) float64 {
s := 0.0
for j := range 2 {
var acc complex128
for k := range 2 {
acc += a.ComplexAt(j*2 + k)
}
s += real(acc * w.Data().ComplexAt(j))
}
return s
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(lossOf, z.Data()), 1e-8)
}
// TestComplexMotionOps pins gradient flow through Slice, Concat,
// Reshape and BroadcastTo on complex tensors.
func TestComplexMotionOps(t *testing.T) {
z := buildComplex(t, []complex128{1 + 1i, 2 - 1i, 3 + 2i, 4 - 3i}, 4)
sl, err := z.Slice(0, 1, 3)
if err != nil {
t.Fatalf("Slice: %v", err)
}
sq, err := sl.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// Only the sliced elements receive gradient, and for L = Σ|z|² it
// is exactly z_i.
for i := range 4 {
want := complex(0, 0)
if i == 1 || i == 2 {
want = z.Data().ComplexAt(i)
}
if z.Grad().ComplexAt(i) != want {
t.Fatalf("grad[%d] = %v, want %v", i, z.Grad().ComplexAt(i), want)
}
}
z2 := buildComplex(t, []complex128{0.5 + 0.5i}, 1)
cat, err := z.Concat(z2, 0)
if err != nil {
t.Fatalf("Concat: %v", err)
}
sq2, err := cat.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss2, err := sq2.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss2.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
if z2.Grad().ComplexAt(0) != 0.5+0.5i {
t.Fatalf("concat gradient = %v, want 0.5+0.5i", z2.Grad().ComplexAt(0))
}
}
// TestComplexSumMeanPow pins the complex reducers and integer powers.
func TestComplexSumMeanPow(t *testing.T) {
z := buildComplex(t, []complex128{1 + 2i, 3 - 1i}, 2)
m, err := z.Mean()
if err != nil {
t.Fatalf("Mean: %v", err)
}
if m.Data().ComplexAt(0) != 2+0.5i {
t.Fatalf("mean = %v, want 2+0.5i", m.Data().ComplexAt(0))
}
p, err := z.Pow(3)
if err != nil {
t.Fatalf("Pow: %v", err)
}
// (1+2i)³ = (1+2i)(1+2i)(1+2i) = -11-2i.
if p.Data().ComplexAt(0) != -11-2i {
t.Fatalf("pow = %v, want -11-2i", p.Data().ComplexAt(0))
}
sq, err := p.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
pv := complex(1, 0)
for range 3 {
pv *= a.ComplexAt(i)
}
s += real(pv)*real(pv) + imag(pv)*imag(pv)
}
return s
}, z.Data()), 1e-6)
}
// TestComplexMatMulGrad pins the 2-D complex matmul adjoint against
// central differences.
func TestComplexMatMulGrad(t *testing.T) {
aVals := []complex128{1 + 1i, 2 - 1i, 0.5 + 0i, -1 + 2i}
bVals := []complex128{0.5 - 0.5i, 1 + 1i, -0.5 + 2i, 0.25 - 0.75i}
a := buildComplex(t, aVals, 2, 2)
b := buildComplex(t, bVals, 2, 2)
y, err := a.MatMul(b)
if err != nil {
t.Fatalf("MatMul: %v", err)
}
sq, err := y.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
lossOf := func(av, bv []complex128) float64 {
s := 0.0
for i := range 2 {
for j := range 2 {
var acc complex128
for k := range 2 {
acc += av[i*2+k] * bv[k*2+j]
}
s += real(acc)*real(acc) + imag(acc)*imag(acc)
}
}
return s
}
checkAgainstNumeric(t, a.Grad(), numericComplexGrad(func(arr *core.Array) float64 {
return lossOf(flatComplex(arr), bVals)
}, a.Data()), 1e-7)
checkAgainstNumeric(t, b.Grad(), numericComplexGrad(func(arr *core.Array) float64 {
return lossOf(aVals, flatComplex(arr))
}, b.Data()), 1e-7)
}
func flatComplex(a *core.Array) []complex128 {
out := make([]complex128, a.Len())
for i := range out {
out[i] = a.ComplexAt(i)
}
return out
}
// TestComplexScaleGrad pins Scale on complex tensors.
func TestComplexScaleGrad(t *testing.T) {
z := buildComplex(t, []complex128{1 + 1i, 2 - 1i}, 2)
s, err := z.Scale(2.5)
if err != nil {
t.Fatalf("Scale: %v", err)
}
sq, err := s.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
// d(2.5²|z|²)/dz̄ = 2·2.5²·Re-parts... exact: 6.25·z.
for i := range 2 {
want := 6.25 * z.Data().ComplexAt(i)
got := z.Grad().ComplexAt(i)
if math.Abs(real(got)-real(want)) > 1e-10 || math.Abs(imag(got)-imag(want)) > 1e-10 {
t.Fatalf("grad[%d] = %v, want %v", i, got, want)
}
}
}
// TestComplexExpGrad pins the complex exponential adjoint dz = g·conj(e^z)
// against central differences, and e^z against the polar identity
// e^{x+iy} = e^x(cos y + i sin y).
func TestComplexExpGrad(t *testing.T) {
vals := []complex128{0.3 - 0.2i, -1.1 + 0.7i, 0.05 + 0i}
z := buildComplex(t, vals, 3)
e, err := z.Exp()
if err != nil {
t.Fatalf("Exp: %v", err)
}
for i, v := range vals {
want := complex(math.Exp(real(v)), 0) * complex(math.Cos(imag(v)), math.Sin(imag(v)))
got := e.Data().ComplexAt(i)
if cmplxAbs(got-want) > 1e-14 {
t.Fatalf("exp[%d] = %v, want %v", i, got, want)
}
}
sq, err := e.Abs2()
if err != nil {
t.Fatalf("Abs2: %v", err)
}
loss, err := sq.Sum()
if err != nil {
t.Fatalf("Sum: %v", err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("Backward: %v", err)
}
checkAgainstNumeric(t, z.Grad(), numericComplexGrad(func(a *core.Array) float64 {
s := 0.0
for i := range a.Len() {
vv := a.ComplexAt(i)
ev := complex(math.Exp(real(vv)), 0) * complex(math.Cos(imag(vv)), math.Sin(imag(vv)))
s += real(ev)*real(ev) + imag(ev)*imag(ev)
}
return s
}, z.Data()), 1e-7)
}