feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+128
View File
@@ -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
}