129 lines
4.0 KiB
Go
129 lines
4.0 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"
|
|
)
|
|
|
|
// Concat joins u after t along an existing axis, the differentiable
|
|
// inverse of Slice, and the building block that lets recurrent layers
|
|
// assemble per-step outputs into one sequence core.
|
|
|
|
// Concat returns the tensors joined along the given existing dimension.
|
|
// Every other dimension must agree. The backward routes each side its
|
|
// own span of the incoming gradient along that dimension, narrowed to
|
|
// the side's dtype first: a real side of a real/complex join receives
|
|
// 2·Re(g), the rule every other mixed-dtype op applies.
|
|
func (t *Tensor) Concat(u *Tensor, dim int) (*Tensor, error) {
|
|
if err := t.checkDiff("Concat"); err != nil {
|
|
return nil, err
|
|
}
|
|
if err := u.checkDiff("Concat"); err != nil {
|
|
return nil, err
|
|
}
|
|
out, err := core.Concat(t.data, u.data, dim)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
at, au := t.data, u.data
|
|
return binaryResult("Concat", t, u, out, func(g gradSlot, dst *[2]gradSlot, ar *gradArena) error {
|
|
dt := at.Dtype()
|
|
du := au.Dtype()
|
|
// One narrowing per side before either span is copied: the
|
|
// concat's own dtype promotes along the ladder, so a complex
|
|
// gradient reaches a real operand in a mixed join and must
|
|
// narrow by 2·Re exactly as the leaf commit would. The narrowed
|
|
// arrays also put both copies on the dtype-matched raw path.
|
|
gA, err := narrowGradient(g, dt)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
gB, err := narrowGradient(g, du)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
shA, shB := at.Shape(), au.Shape()
|
|
outer, inner, err := outerInner(shA, dim)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
spanA := shA[dim]
|
|
total := spanA + shB[dim]
|
|
da := gradSlot{arr: ar.borrowGrad(dt, shA), sh: shA}
|
|
// Every row contributes one contiguous inner run, so matching
|
|
// dtypes collapse the triple loop to a raw slice move per row.
|
|
fastA := !gA.arr.Strided() && gA.arr.Dtype() == dt && dt != core.Int
|
|
for o := range outer {
|
|
for i := range spanA {
|
|
d := o*spanA*inner + i*inner
|
|
s := o*total*inner + i*inner
|
|
if fastA {
|
|
copySegRaw(da.arr, gA.arr, d, s, inner)
|
|
continue
|
|
}
|
|
for j := range inner {
|
|
copyElem(da.arr, d+j, gA.arr, s+j)
|
|
}
|
|
}
|
|
}
|
|
db := gradSlot{arr: ar.borrowGrad(du, shB), sh: shB}
|
|
spanB := shB[dim]
|
|
fastB := !gB.arr.Strided() && gB.arr.Dtype() == du && du != core.Int
|
|
for o := range outer {
|
|
for i := range spanB {
|
|
d := o*spanB*inner + i*inner
|
|
s := o*total*inner + (spanA+i)*inner
|
|
if fastB {
|
|
copySegRaw(db.arr, gB.arr, d, s, inner)
|
|
continue
|
|
}
|
|
for j := range inner {
|
|
copyElem(db.arr, d+j, gB.arr, s+j)
|
|
}
|
|
}
|
|
}
|
|
dst[0], dst[1] = da, db
|
|
return nil
|
|
}), nil
|
|
}
|
|
|
|
// copySegRaw moves n elements from src at sOff to dst at dOff through
|
|
// the raw payloads. The caller checks dtype equality and contiguity;
|
|
// the per-element values, and so the bits, are the ones copyElem
|
|
// writes one accessor call at a time.
|
|
func copySegRaw(dst, src *core.Array, dOff, sOff, n int) {
|
|
switch src.Dtype() {
|
|
case core.Float32:
|
|
copy(dst.RawFloat32s()[dOff:dOff+n], src.RawFloat32s()[sOff:sOff+n])
|
|
case core.Float:
|
|
copy(dst.RawFloats()[dOff:dOff+n], src.RawFloats()[sOff:sOff+n])
|
|
case core.Complex:
|
|
copy(dst.RawComplexes()[dOff:dOff+n], src.RawComplexes()[sOff:sOff+n])
|
|
default:
|
|
copy(dst.RawInts()[dOff:dOff+n], src.RawInts()[sOff:sOff+n])
|
|
}
|
|
}
|
|
|
|
// outerInner splits a shape into the products of the dimensions before
|
|
// and after dim, the strides a flat row-major walk needs when only one
|
|
// axis is being split or joined.
|
|
func outerInner(shape []int, dim int) (int, int, error) {
|
|
if len(shape) == 0 {
|
|
return 0, 0, errf("Concat: cannot concatenate a scalar")
|
|
}
|
|
if dim < 0 || dim >= len(shape) {
|
|
return 0, 0, errf("Concat: dimension %d is out of range for shape %v", dim, shape)
|
|
}
|
|
outer := 1
|
|
for d := range dim {
|
|
outer *= shape[d]
|
|
}
|
|
inner := 1
|
|
for d := dim + 1; d < len(shape); d++ {
|
|
inner *= shape[d]
|
|
}
|
|
return outer, inner, nil
|
|
}
|