// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import "sourcedock.dev/petrbalvin/tensor/internal/base" import "math" // Concatenation and stacking. Concat joins along an existing // dimension, Stack along a new one; both copy, the result dtype promotes // per the ladder, and every shape disagreement names both shapes. // Concat joins b after a along the given existing dimension. Every other // dimension must agree. func Concat(a, b *Array, dim int) (*Array, error) { if a.NDim() != b.NDim() { return nil, errf("Concat: ranks differ, %s and %s", shapeText(a.shape), shapeText(b.shape)) } if dim < 0 || dim >= a.NDim() { return nil, errf("Concat: dimension %d is out of range for shape %s", dim, shapeText(a.shape)) } for d := range a.shape { if d != dim && a.shape[d] != b.shape[d] { return nil, errf("Concat: shapes %s and %s disagree outside dimension %d", shapeText(a.shape), shapeText(b.shape), dim) } } newShape := a.Shape() // Bound the joined length before adding it: a wrapped sum would pair // a negative dimension with an unvalidated allocation. if b.shape[dim] > math.MaxInt-a.shape[dim] { return nil, errf("Concat: joining along dimension %d overflows: %d + %d", dim, a.shape[dim], b.shape[dim]) } newShape[dim] = a.shape[dim] + b.shape[dim] total, _, terr := checkedDims(newShape) if terr != nil { return nil, base.WrapErr("Concat", terr) } dt := promote(a.dt, b.dt) out := &Array{shape: newShape, dt: dt} out.alloc(total) if dt == a.dt && dt == b.dt && a.isContiguous() && b.isContiguous() { // No promotion: every outer position contributes one run of each // operand, so the join is two block copies per position rather // than a converted element per slot. copyRun dispatches every // dtype the package stores, the narrow payloads included. tail := 1 for d := dim + 1; d < len(newShape); d++ { tail *= newShape[d] } head := 1 for d := range dim { head *= newShape[d] } na, nb := a.shape[dim], b.shape[dim] for h := range head { base := h * newShape[dim] * tail copyRun(out, a, base, h*na*tail, na*tail) copyRun(out, b, base+na*tail, h*nb*tail, nb*tail) } return out, nil } if a.isContiguous() && b.isContiguous() { // Promotion: every outer position still contributes one run of // each operand, exactly as the no-promotion path above lays the // join out, so the join is two bulk conversions per position // rather than a setConverted dispatch per element. The // conversions are the ones setConverted performs, value for // value, and the run layout is the same walk's. tail := 1 for d := dim + 1; d < len(newShape); d++ { tail *= newShape[d] } head := 1 for d := range dim { head *= newShape[d] } na, nb := a.shape[dim], b.shape[dim] for h := range head { base := h * newShape[dim] * tail concatConvert(out, a, base, h*na*tail, na*tail) concatConvert(out, b, base+na*tail, h*nb*tail, nb*tail) } return out, nil } coord := make([]int, len(newShape)) for i := range total { src := a c := coord[dim] if c >= a.shape[dim] { src = b c -= a.shape[dim] } // Flat offset into the source: same walk with the local dim // coordinate and the source's own dimension sizes. off := 0 for d := range coord { cd := coord[d] if d == dim { cd = c } off = off*src.shape[d] + cd } out.setConverted(i, src, off) advanceOdometer(coord, newShape) } return out, nil } // concatConvert converts n elements of a contiguous src, read from // srcOff, into dst's dtype at dstOff, the exact conversions setConverted // performs for the same pair: the source's payload is read where it // lies and the promoted store is a plain loop, so no element pays a // second dispatch or a stride resolution. The conversion arms mirror // setConverted arm for arm. func concatConvert(dst, src *Array, dstOff, srcOff, n int) { switch dst.dt { case Float: d := dst.floats[dstOff : dstOff+n] switch src.dt { case Float: copy(d, src.floats[srcOff:srcOff+n]) case Float32: p := src.floats32[srcOff : srcOff+n] for i := range d { d[i] = float64(p[i]) } case Float16: p := src.halves[srcOff : srcOff+n] for i := range d { d[i] = HalfToFloat64(p[i]) } case Int: p := src.ints[srcOff : srcOff+n] for i := range d { d[i] = float64(p[i]) } default: // Bool and the narrow integers: floatAt widens each exactly. for i := range d { d[i] = src.floatAt(srcOff + i) } } case Complex: d := dst.complexes[dstOff : dstOff+n] if src.dt == Complex { copy(d, src.complexes[srcOff:srcOff+n]) return } for i := range d { d[i] = src.complexAt(srcOff + i) } case Int: d := dst.ints[dstOff : dstOff+n] switch src.dt { case Int: // Exact: the float64 detour rounds above 2^53. copy(d, src.ints[srcOff:srcOff+n]) case Float: p := src.floats[srcOff : srcOff+n] for i := range d { d[i] = int64(p[i]) } case Float16: p := src.halves[srcOff : srcOff+n] for i := range d { d[i] = int64(HalfToFloat64(p[i])) } case Float32: p := src.floats32[srcOff : srcOff+n] for i := range d { d[i] = int64(p[i]) } default: for i := range d { d[i] = src.intAt(srcOff + i) } } case Float32: d := dst.floats32[dstOff : dstOff+n] if src.dt == Float32 { copy(d, src.floats32[srcOff:srcOff+n]) return } for i := range d { d[i] = src.float32At(srcOff + i) } case Float16: d := dst.halves[dstOff : dstOff+n] for i := range d { d[i] = HalfFromFloat64(src.floatAt(srcOff + i)) } default: // The narrow integer destinations: setConverted casts the // widened intAt read down, the implicit-store semantics every // promoted join target has always carried. for i := range n { dst.setConverted(dstOff+i, src, srcOff+i) } } } // Stack joins b after a along a new dimension inserted at dim; the two // arrays must have identical shapes. func Stack(a, b *Array, dim int) (*Array, error) { if !sameShape(a.shape, b.shape) { return nil, errf("Stack: shapes %s and %s must be identical", shapeText(a.shape), shapeText(b.shape)) } if dim < 0 || dim > a.NDim() { return nil, errf("Stack: dimension %d is out of range for inserting into shape %s", dim, shapeText(a.shape)) } newShape := make([]int, 0, a.NDim()+1) newShape = append(newShape, a.shape[:dim]...) newShape = append(newShape, 2) newShape = append(newShape, a.shape[dim:]...) total := 1 for _, d := range newShape { total *= d } dt := promote(a.dt, b.dt) out := &Array{shape: newShape, dt: dt} out.alloc(total) tail := 1 for d := dim; d < a.NDim(); d++ { tail *= a.shape[d] } head := 1 for d := range dim { head *= a.shape[d] } if dt == a.dt && dt == b.dt && a.isContiguous() && b.isContiguous() { // The stacked axis is a factor of two: each outer position holds // one run of each operand, so the join is two block copies. // copyRun dispatches every dtype the package stores, the narrow // payloads included. for h := range head { copyRun(out, a, h*2*tail, h*tail, tail) copyRun(out, b, h*2*tail+tail, h*tail, tail) } return out, nil } if a.isContiguous() && b.isContiguous() { // Promotion: each outer position still holds one run of each // operand, so the join is two bulk conversions per position // rather than a setConverted dispatch per element, the // conversions concatConvert keeps. for h := range head { concatConvert(out, a, h*2*tail, h*tail, tail) concatConvert(out, b, h*2*tail+tail, h*tail, tail) } return out, nil } coord := make([]int, len(newShape)) for i := range total { src := a if coord[dim] == 1 { src = b } // Flat offset into the source: drop the stacked coordinate. off := 0 for d := range coord { if d == dim { continue } off = off*src.shape[stackDim(d, dim)] + coord[d] } out.setConverted(i, src, off) advanceOdometer(coord, newShape) } return out, nil } // stackDim maps a new-shape dimension to the source-shape dimension, // skipping the inserted axis. func stackDim(d, dim int) int { if d > dim { return d - 1 } return d }