Files

286 lines
8.0 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// 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
}