286 lines
8.0 KiB
Go
286 lines
8.0 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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
|
|
}
|