// Copyright (c) 2026 Petr Balvín (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) }