// Copyright (c) 2026 Petr Balvín (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 }