Files
tensor/grad/complex.go
T

504 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 (
"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
}