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

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