Files
tensor/grad/complex_test.go
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

540 lines
14 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
}