feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+128
@@ -0,0 +1,128 @@
|
||||
// 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
|
||||
}
|
||||
Reference in New Issue
Block a user