504 lines
14 KiB
Go
504 lines
14 KiB
Go
// 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
|
|||
|
|
}
|