// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "sourcedock.dev/petrbalvin/tensor/internal/core" ) // Complex differentiation. The graph accepts complex128 // tensors alongside float64/float32, with the Wirtinger convention the // optimiser ecosystem settled on: Backward seeds a REAL scalar loss // (a complex output is rejected with an error telling the caller to // reduce first), and the gradient a complex leaf accumulates is // ∂L/∂z̄, the direction gradient descent steps along. Under that // convention the adjoint of a holomorphic op y = f(z) is // dz += g·conj(f′(z)), so every conjugation below sits exactly where // the calculus puts it. // // A real tensor inside a complex graph narrows the incoming complex // gradient by 2·Re: for a real variable x, dL/dx = 2·Re(∂L/∂x̄), and // the factor also cancels the ½ the Real backward contributes, so // mixed graphs compose exactly. // checkDiff is checkFloat plus complex: the ops that can differentiate // complex inputs validate with it. func (t *Tensor) checkDiff(name string) error { switch t.data.Dtype() { case core.Float, core.Float32, core.Complex: return nil } return errf("autograd: %s needs a float, float32 or complex tensor, got %s", name, t.data.Dtype()) } // isComplexArr reports whether a holds complex128 data. func isComplexArr(a *core.Array) bool { return a.Dtype() == core.Complex } // eitherComplex reports whether either operand is complex. func eitherComplex(a, b *core.Array) bool { return isComplexArr(a) || isComplexArr(b) } // conjArray returns the element-wise conjugate. Real arrays come back // unchanged (their conjugate is themselves), so mixed-dtype adjoints // can call it unconditionally. func conjArray(a *core.Array) *core.Array { if !isComplexArr(a) { return a } out := zeros(core.Complex, a.Shape()) cs := out.RawComplexes() if a.Strided() { for i := range cs { cs[i] = conj(a.ComplexAt(i)) } return out } as := a.RawComplexes() for i := range cs { z := as[i] cs[i] = complex(real(z), -imag(z)) } return out } func conj(z complex128) complex128 { return complex(real(z), -imag(z)) } // copyElem copies one element between gradient arrays of the same // dtype; the callers narrow the incoming gradient to the operand's // dtype with narrowGradient before the copy, so a mixed real/complex // pair never reaches here and a real destination never reads a complex // payload. A complex destination reads through ComplexAt, which serves // a strided source too. func copyElem(dst *core.Array, di int, src *core.Array, si int) { if dst.Dtype() == core.Complex { dst.RawComplexes()[di] = src.ComplexAt(si) return } dst.SetFloatAt(di, src.FloatAt(si)) } // narrowGradient converts a gradient to the dtype of the tensor it // accumulates into. Complex to real takes 2·Re (the real-tensor rule // above); everything else routes through Astype. func narrowGradient(g gradSlot, dt core.Dtype) (gradSlot, error) { if g.arr.Dtype() == dt { return g, nil } sh := g.sh if sh == nil { sh = g.arr.Shape() } if g.arr.Dtype() == core.Complex && dt != core.Complex { out := zeros(dt, sh) gs := g.arr.RawComplexes() if dt == core.Float32 && !g.arr.Strided() { os := out.RawFloat32s() for i := range os { os[i] = float32(2 * real(gs[i])) } return gradSlot{arr: out, sh: sh}, nil } if dt == core.Float && !g.arr.Strided() { os := out.RawFloats() for i := range os { os[i] = 2 * real(gs[i]) } return gradSlot{arr: out, sh: sh}, nil } // out is freshly allocated and dense, so a real destination // takes its payload directly; the complex source keeps the // accessor read that rebases a strided index. switch dt { case core.Float32: os := out.RawFloat32s() for i := range os { os[i] = float32(2 * real(g.arr.ComplexAt(i))) } return gradSlot{arr: out, sh: sh}, nil case core.Float: os := out.RawFloats() for i := range os { os[i] = 2 * real(g.arr.ComplexAt(i)) } return gradSlot{arr: out, sh: sh}, nil } for i := range g.arr.Len() { out.SetFloatAt(i, 2*real(g.arr.ComplexAt(i))) } return gradSlot{arr: out, sh: sh}, nil } c, err := core.Astype(g.arr, dt) if err != nil { return gradSlot{}, err } return gradSlot{arr: c, sh: sh}, nil } // scalarComplex builds a 1-element complex array holding z. func scalarComplex(z complex128) *core.Array { out := zeros(core.Complex, []int{1}) out.RawComplexes()[0] = z return out } // fillComplex returns a complex array shaped like a with every element // set to z. func fillComplex(a *core.Array, z complex128) *core.Array { out := zeros(core.Complex, a.Shape()) cs := out.RawComplexes() for i := range cs { cs[i] = z } return out } // Conj returns the element-wise complex conjugate. The conjugate is // anti-holomorphic: its ∂/∂z̄ adjoint conjugates the incoming // gradient (dz = conj(g)), which is what makes ⟨ψ|H|ψ⟩ come out as // Hψ rather than only its real part. func (t *Tensor) Conj() (*Tensor, error) { if err := t.checkDiff("Conj"); err != nil { return nil, err } a := t.data out := conjArray(t.data) if !isComplexArr(t.data) { // conj of a real tensor is a copy, so the graph needs its own // node data, not the operand alias. out = cloneReal(t.data) } return t.unaryResult("Conj", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) if !isComplexArr(g.arr) { c, err := copyGradSlot(ar, g, sh) if err != nil { return err } dst[0] = c return nil } n := g.arr.Len() da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} cs := da.arr.RawComplexes()[:n] gs := g.arr.RawComplexes()[:n] for i := range cs { z := gs[i] cs[i] = complex(real(z), -imag(z)) } dst[0] = da return nil }), nil } // cloneReal copies a real array (the graph never aliases operands). func cloneReal(a *core.Array) *core.Array { out := zeros(a.Dtype(), a.Shape()) switch { case a.Strided(): for i := range a.Len() { out.SetFloatAt(i, a.FloatAt(i)) } case a.Dtype() == core.Float32: copy(out.RawFloat32s(), a.RawFloat32s()) case a.Dtype() == core.Float: copy(out.RawFloats(), a.RawFloats()) default: copy(out.RawInts(), a.RawInts()) } return out } // Real returns the real part of each element as a float tensor. The // complex backward halves the gradient (∂Re z/∂z̄ = ½), which the // 2·Re narrowing at any real destination cancels exactly. func (t *Tensor) Real() (*Tensor, error) { if err := t.checkDiff("Real"); err != nil { return nil, err } if !isComplexArr(t.data) { // Real of a real tensor is a copy with its own storage. out := cloneReal(t.data) a := t.data return t.unaryResult("Real", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { c, err := copyGradSlot(ar, g, gradShape(g, a)) if err != nil { return err } dst[0] = c return nil }), nil } out := zeros(core.Float, t.data.Shape()) if t.data.Strided() { for i := range t.data.Len() { out.SetFloatAt(i, real(t.data.ComplexAt(i))) } } else { cs := t.data.RawComplexes() os := out.RawFloats() for i := range os { os[i] = real(cs[i]) } } // The shape is captured now: nothing may be read off the input at // backward time, or a ReplaceWith in between would change it. shape := t.data.Shape() return t.unaryResult("Real", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { da := gradSlot{arr: ar.borrowGrad(core.Complex, shape), sh: shape} cs := da.arr.RawComplexes()[:da.arr.Len()] if g.arr.Strided() || g.arr.Dtype() != core.Float { for i := range cs { cs[i] = complex(g.arr.FloatAt(i)/2, 0) } dst[0] = da return nil } gs := g.arr.RawFloats() for i := range cs { cs[i] = complex(gs[i]/2, 0) } dst[0] = da return nil }), nil } // Imag returns the imaginary part of each element as a float tensor; // the complex backward scales by i/2 (∂Im z/∂z̄ = i/2). func (t *Tensor) Imag() (*Tensor, error) { if !isComplexArr(t.data) { return nil, errf("autograd: Imag needs a complex tensor, got %s", t.data.Dtype()) } out := zeros(core.Float, t.data.Shape()) if t.data.Strided() { for i := range t.data.Len() { out.SetFloatAt(i, imag(t.data.ComplexAt(i))) } } else { cs := t.data.RawComplexes() os := out.RawFloats() for i := range os { os[i] = imag(cs[i]) } } shape := t.data.Shape() return t.unaryResult("Imag", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { da := gradSlot{arr: ar.borrowGrad(core.Complex, shape), sh: shape} cs := da.arr.RawComplexes()[:da.arr.Len()] if g.arr.Strided() || g.arr.Dtype() != core.Float { for i := range cs { cs[i] = complex(0, g.arr.FloatAt(i)/2) } dst[0] = da return nil } gs := g.arr.RawFloats() for i := range cs { cs[i] = complex(0, gs[i]/2) } dst[0] = da return nil }), nil } // Abs2 returns |z|² of each element, a real tensor. The complex // backward is dz = g·z (∂|z|²/∂z̄ = z); the real input path is the // square with its 2x backward, keeping the operand's width (a float32 // input squares in float64 and stays float32, exactly as Pow // does). func (t *Tensor) Abs2() (*Tensor, error) { if err := t.checkDiff("Abs2"); err != nil { return nil, err } if !isComplexArr(t.data) { return t.squareGraph() } a := t.data out := zeros(core.Float, a.Shape()) if a.Strided() { for i := range a.Len() { z := a.ComplexAt(i) out.SetFloatAt(i, real(z)*real(z)+imag(z)*imag(z)) } } else { // Bound the walk by the destination's length: a rebased view's // payload may run longer than its element count. as := a.RawComplexes() os := out.RawFloats() for i := range os { z := as[i] os[i] = real(z)*real(z) + imag(z)*imag(z) } } return t.unaryResult("Abs2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} cs := da.arr.RawComplexes()[:da.arr.Len()] if a.Strided() || g.arr.Strided() || g.arr.Dtype() != core.Float { for i := range cs { cs[i] = complex(g.arr.FloatAt(i), 0) * a.ComplexAt(i) } dst[0] = da return nil } as, gs := a.RawComplexes(), g.arr.RawFloats() for i := range cs { cs[i] = complex(gs[i], 0) * as[i] } dst[0] = da return nil }), nil } // squareGraph is the real-input branch of Abs2: y = x², dx = 2x·g.arr. // The output keeps the operand's width, squared in float64 and rounded // once, exactly as Pow does, so Abs2 and Pow(2) agree on dtype // and value for a float32 operand. func (t *Tensor) squareGraph() (*Tensor, error) { a := t.data out := zeros(a.Dtype(), a.Shape()) switch { case a.Dtype() == core.Float32 && !a.Strided(): as, os := a.RawFloat32s(), out.RawFloat32s() for i := range os { v := float64(as[i]) os[i] = float32(v * v) } case a.Strided(): for i := range out.Len() { v := a.FloatAt(i) out.SetFloatAt(i, v*v) } default: as, os := a.RawFloats(), out.RawFloats() for i := range os { v := as[i] os[i] = v * v } } return t.unaryResult("Abs2", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) // dx = 2·x·g with the staged chain's rounding: the product // forms first and the doubling multiplies it, per element. if !a.Strided() && !g.arr.Strided() && a.Dtype() == g.arr.Dtype() && a.Len() == g.arr.Len() { n := a.Len() switch a.Dtype() { case core.Float: da := gradSlot{arr: ar.borrowGrad(core.Float, sh), sh: sh} as, gs, ds := a.RawFloats()[:n], g.arr.RawFloats()[:n], da.arr.RawFloats()[:n] for i := range ds { ds[i] = (as[i] * gs[i]) * 2 } dst[0] = da return nil case core.Float32: da := gradSlot{arr: ar.borrowGrad(core.Float32, sh), sh: sh} as, gs, ds := a.RawFloat32s()[:n], g.arr.RawFloat32s()[:n], da.arr.RawFloat32s()[:n] for i := range ds { p := float32(float64(as[i]) * float64(gs[i])) ds[i] = float32(float64(p) * 2) } dst[0] = da return nil } } da, err := core.Mul(a, g.arr) if err != nil { return err } dst[0] = gradSlot{arr: core.MulI(da, 2), sh: sh} return nil }), nil } // Abs returns the absolute value of each element: complex input yields // float magnitudes with dz = g·z/(2|z|) (zero at the origin, the // subgradient). The real branch lives beside the other real kernels in // tensor.go and dispatches here for complex input. func (t *Tensor) absComplex() (*Tensor, error) { a := t.data out := core.Abs(a) return t.unaryResult("Abs", out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error { sh := gradShape(g, a) da := gradSlot{arr: ar.borrowGrad(core.Complex, sh), sh: sh} cs := da.arr.RawComplexes()[:da.arr.Len()] if a.Strided() || g.arr.Strided() || out.Strided() || g.arr.Dtype() != core.Float || out.Dtype() != core.Float { for i := range cs { z := a.ComplexAt(i) m := out.FloatAt(i) if m == 0 { continue } cs[i] = complex(g.arr.FloatAt(i)/(2*m), 0) * z } dst[0] = da return nil } as, gs, os := a.RawComplexes(), g.arr.RawFloats(), out.RawFloats() for i := range cs { m := os[i] if m == 0 { continue } cs[i] = complex(gs[i]/(2*m), 0) * as[i] } dst[0] = da return nil }), nil } // powComplexGrad builds the Wirtinger backward of y = zⁿ: // dz = g·n·conj(z)ⁿ⁻¹, assembled by repeated conjugate multiplication // (the exponent is a small integer; a loop beats a general power). // sh is the shape the incoming gradient carries, or the operand's own // on the legacy sweep path (gradShape). func powComplexGrad(ar *gradArena, g gradSlot, a *core.Array, n int64, sh []int) *core.Array { da := ar.borrowGrad(core.Complex, sh) cs := da.RawComplexes()[:da.Len()] if a.Strided() || g.arr.Strided() { for i := range cs { term := complex(1, 0) for range n - 1 { term *= conj(a.ComplexAt(i)) } cs[i] = complex(float64(n), 0) * g.arr.ComplexAt(i) * term } return da } as, gs := a.RawComplexes(), g.arr.RawComplexes() for i := range cs { term := complex(1, 0) for range n - 1 { term *= conj(as[i]) } cs[i] = complex(float64(n), 0) * gs[i] * term } return da } // powComplexForward raises each complex element to a non-negative // integer power by repeated multiplication. func powComplexForward(a *core.Array, n int64) *core.Array { out := zeros(core.Complex, a.Shape()) cs := out.RawComplexes() if a.Strided() { for i := range cs { p := complex(1, 0) for range n { p *= a.ComplexAt(i) } cs[i] = p } return out } as := a.RawComplexes() for i := range cs { p := complex(1, 0) for range n { p *= as[i] } cs[i] = p } return out }