Files
tensor/internal/core/reshape.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

376 lines
12 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"
// Reshape and padding utilities for the tensor numeric surface. Flatten,
// Squeeze, Unsqueeze and TransposeAxes change the shape without changing
// the element order; Copy returns a fresh copy with the same layout
// (every array is row-major, so Copy is always a real allocation);
// Pad extends arrays along one or more dimensions with constant,
// reflect, replicate or circular modes.
// Flatten returns a copy with the dimensions in [startDim, endDim]
// collapsed into one. Negative indices count from the end (endDim = -1
// means the last dimension). startDim > endDim is an error.
func Flatten(a *Array, startDim, endDim int) (*Array, error) {
ndim := a.NDim()
if startDim < 0 {
startDim += ndim
}
if endDim < 0 {
endDim += ndim
}
if startDim < 0 || startDim >= ndim || endDim < 0 || endDim >= ndim || startDim > endDim {
return nil, errf("Flatten: range [%d, %d] out of bounds for shape %s", startDim, endDim, shapeText(a.shape))
}
flat := 1
for d := startDim; d <= endDim; d++ {
flat *= a.shape[d]
}
newShape := make([]int, 0, ndim-(endDim-startDim))
newShape = append(newShape, a.shape[:startDim]...)
newShape = append(newShape, flat)
newShape = append(newShape, a.shape[endDim+1:]...)
return Reshape(a, newShape...)
}
// Squeeze returns a copy with size-1 dimensions removed. When dim is -1
// every size-1 dimension is dropped; otherwise only that one is (and it
// must have size 1).
func Squeeze(a *Array, dim int) (*Array, error) {
if dim == -1 {
newShape := make([]int, 0, a.NDim())
for _, d := range a.shape {
if d != 1 {
newShape = append(newShape, d)
}
}
if len(newShape) == 0 {
newShape = []int{1}
}
return Reshape(a, newShape...)
}
if dim < 0 || dim >= a.NDim() {
return nil, errf("Squeeze: dimension %d is out of range for shape %s", dim, shapeText(a.shape))
}
if a.shape[dim] != 1 {
return nil, errf("Squeeze: dimension %d has size %d, must be 1", dim, a.shape[dim])
}
newShape := make([]int, 0, a.NDim()-1)
newShape = append(newShape, a.shape[:dim]...)
newShape = append(newShape, a.shape[dim+1:]...)
if len(newShape) == 0 {
newShape = []int{1}
}
return Reshape(a, newShape...)
}
// Unsqueeze returns a copy with a new size-1 dimension inserted at dim.
// dim is in [0, NDim]; negative values count from the end of the result
// rank, so dim = -1 appends the new axis at the end.
func Unsqueeze(a *Array, dim int) (*Array, error) {
ndim := a.NDim() + 1
if dim < 0 {
dim += a.NDim() + 1
}
if dim < 0 || dim >= ndim {
return nil, errf("Unsqueeze: dimension %d is out of range for inserting into shape %s", dim, shapeText(a.shape))
}
newShape := make([]int, 0, ndim)
newShape = append(newShape, a.shape[:dim]...)
newShape = append(newShape, 1)
newShape = append(newShape, a.shape[dim:]...)
return Reshape(a, newShape...)
}
// TransposeAxes returns a copy with the dimensions reordered according to
// dims. dims must be a permutation of [0, NDim); an error names the
// bad permutation. Renamed from Permute to avoid the clash with the
// random-number Fisher-Yates permutation in the generator.
func TransposeAxes(a *Array, dims ...int) (*Array, error) {
if len(dims) != a.NDim() {
return nil, errf("TransposeAxes: needs %d dimensions, got %d", a.NDim(), len(dims))
}
seen := make([]bool, a.NDim())
for _, d := range dims {
if d < 0 || d >= a.NDim() || seen[d] {
return nil, errf("TransposeAxes: invalid permutation %v for shape %s", dims, shapeText(a.shape))
}
seen[d] = true
}
newShape := make([]int, a.NDim())
for i, d := range dims {
newShape[i] = a.shape[d]
}
total := a.Len()
out := &Array{shape: newShape, dt: a.dt}
out.alloc(total)
srcCoord := make([]int, a.NDim())
// The dtype dispatch keeps setFrom's per-element switch out of the
// walk; the destination index still needs the coordinate fold. Bool
// and the narrow integer widths take the setFrom walk directly: the
// same values in the same order, with no per-dtype loop of their own.
switch a.dt {
case Int:
for i := range total {
dst := 0
for k := range a.NDim() {
dst = dst*newShape[k] + srcCoord[dims[k]]
}
out.ints[dst] = a.ints[i]
advanceOdometer(srcCoord, a.shape)
}
case Float16:
for i := range total {
dst := 0
for k := range a.NDim() {
dst = dst*newShape[k] + srcCoord[dims[k]]
}
out.halves[dst] = a.halves[i]
advanceOdometer(srcCoord, a.shape)
}
case Float32:
for i := range total {
dst := 0
for k := range a.NDim() {
dst = dst*newShape[k] + srcCoord[dims[k]]
}
out.floats32[dst] = a.floats32[i]
advanceOdometer(srcCoord, a.shape)
}
case Float:
for i := range total {
dst := 0
for k := range a.NDim() {
dst = dst*newShape[k] + srcCoord[dims[k]]
}
out.floats[dst] = a.floats[i]
advanceOdometer(srcCoord, a.shape)
}
case Complex:
for i := range total {
dst := 0
for k := range a.NDim() {
dst = dst*newShape[k] + srcCoord[dims[k]]
}
out.complexes[dst] = a.complexes[i]
advanceOdometer(srcCoord, a.shape)
}
default:
for i := range total {
dst := 0
for k := range a.NDim() {
dst = dst*newShape[k] + srcCoord[dims[k]]
}
out.setFrom(dst, a, i)
advanceOdometer(srcCoord, a.shape)
}
}
return out, nil
}
// Copy returns a fresh array with the same data, shape and dtype as a.
// Every tensor is already row-major in storage, so Copy is
// always a real allocation: there are no strides to collapse, no
// views to flatten. Useful when a caller wants a guaranteed-new
// buffer they can hand off without worrying about aliasing.
func Copy(a *Array) *Array {
// cloneArray carries every payload the dtype owns; the five-slice
// cloneData leaves the narrow element types empty.
return a.cloneArray()
}
// Pad extends an array along its trailing dimensions. pad is a flat
// sequence of pairs in reverse spatial order:
// - 2-D: pad = (left, right, top, bottom)
// - 3-D: pad = (left, right, top, bottom, front, back)
//
// mode is one of "constant" (fill with value), "reflect" (mirror without
// repeating the edge), "replicate" (repeat the edge), "circular" (wrap
// around). Pad returns an error for an unsupported mode, a malformed
// pad argument, or a complex input under non-constant modes.
func Pad(a *Array, pad []int, mode string, value float64) (*Array, error) {
if len(pad)%2 != 0 {
return nil, errf("Pad: pad must hold a pre/post pair per dimension, got the odd count %d values", len(pad))
}
switch mode {
case "constant", "reflect", "replicate", "circular":
default:
return nil, errf("Pad: unknown mode %q", mode)
}
if mode != "constant" {
// The folding modes index back into the source: an empty
// dimension has no element to mirror, repeat or wrap, and the
// fold loops below would never terminate (or would panic) on
// one.
for d := range a.shape {
if a.shape[d] == 0 {
return nil, errf("Pad: mode %q is not defined for an empty dimension (%d has size 0)", mode, d)
}
}
}
if a.dt == Complex && mode != "constant" {
return nil, errf("Pad: mode %q is not supported for complex arrays", mode)
}
// Pair count must equal rank (each dim gets a left/right).
if len(pad)/2 != a.NDim() {
return nil, errf("Pad: shape %s needs %d pad values, got %d", shapeText(a.shape), a.NDim()*2, len(pad))
}
// Reverse pad so it pairs with shape left-to-right: the pad list is
// laid out from the last dimension to the first (2-D: left, right,
// top, bottom = dim 1 pre/post, dim 0 pre/post). rpad[i*2],
// rpad[i*2+1] is the pre/post pair for dimension i.
rpad := make([]int, a.NDim()*2)
for d := range a.shape {
srcIdx := (a.NDim() - 1 - d) * 2
rpad[2*d] = pad[srcIdx]
rpad[2*d+1] = pad[srcIdx+1]
}
newShape := make([]int, a.NDim())
for d := range a.shape {
// A negative pad has no meaning for the fold below and would
// drive the new shape negative, which alloc would then reject
// with a makeslice panic instead of an error.
if rpad[2*d] < 0 || rpad[2*d+1] < 0 {
return nil, errf("Pad: pad values must be non-negative, got %d and %d for dimension %d",
rpad[2*d], rpad[2*d+1], d)
}
// Bound each extent before adding: hostile pad values must be
// an error, not a wrapped extent that alloc turns into a
// makeslice panic.
extent := a.shape[d] + rpad[2*d]
if extent < rpad[2*d] {
return nil, errf("Pad: pad values overflow dimension %d", d)
}
extent += rpad[2*d+1]
if extent < rpad[2*d+1] {
return nil, errf("Pad: pad values overflow dimension %d", d)
}
newShape[d] = extent
}
if mode == "reflect" {
// A reflection can only fold once: a pad of n or more on an
// axis of length n has no source left to mirror, and folding
// again would hand setFrom a negative offset.
for d := range a.shape {
if a.shape[d] > 1 && (rpad[2*d] >= a.shape[d] || rpad[2*d+1] >= a.shape[d]) {
return nil, errf("Pad: reflect pad %d exceeds dimension %d of length %d",
max(rpad[2*d], rpad[2*d+1]), d, a.shape[d])
}
}
}
// The padded extents are checked like every constructor's shape: a
// product that wraps must be an error, not a tiny allocation paired
// with a huge shape.
total, _, terr := checkedDims(newShape)
if terr != nil {
return nil, base.WrapErr("Pad", terr)
}
out := &Array{shape: newShape, dt: a.dt}
out.alloc(total)
// Fill by source coordinates.
dstCoord := make([]int, a.NDim())
srcCoord := make([]int, a.NDim())
for i := range total {
for d := range a.NDim() {
s := dstCoord[d] - rpad[2*d]
switch mode {
case "constant":
if s < 0 || s >= a.shape[d] {
s = -1 // marker: fill with value
}
case "reflect":
if a.shape[d] == 1 {
s = 0
break
}
for s < 0 {
s = -s
}
for s >= a.shape[d] {
s = 2*a.shape[d] - 2 - s
}
case "replicate":
if s < 0 {
s = 0
}
if s >= a.shape[d] {
s = a.shape[d] - 1
}
case "circular":
// One modulo pair folds any offset into range in
// constant time: Go's % keeps the sign of s, so the
// addition of n normalises the negative side and the
// second % lands in [0, n). The fold loops this
// replaces walked one step at a time, which turned a
// large pad on a short axis into a quadratic fold.
s = ((s % a.shape[d]) + a.shape[d]) % a.shape[d]
default:
return nil, errf("Pad: unknown mode %q", mode)
}
srcCoord[d] = s
}
if mode == "constant" && containsNeg(srcCoord) {
out.setFromValue(i, value)
} else {
off := 0
for d := range srcCoord {
off = off*a.shape[d] + srcCoord[d]
}
out.setFrom(i, a, off)
}
advanceOdometer(dstCoord, newShape)
}
return out, nil
}
// setFromValue sets element i of the result to the given float, widening
// to the array's dtype: a half payload narrows under the
// HalfFromFloat64 contract, and the narrow integer widths and bool take
// the same implicit-store cast filled() carries (Go's conversion through
// int64, v != 0 for bool). Used by Pad's constant mode.
func (a *Array) setFromValue(i int, v float64) {
switch a.dt {
case Int:
a.ints[i] = int64(v)
case Bool:
a.bools[i] = v != 0
case Int8:
a.i8s[i] = int8(int64(v))
case Uint8:
a.u8s[i] = uint8(int64(v))
case Int16:
a.i16s[i] = int16(int64(v))
case Uint16:
a.u16s[i] = uint16(int64(v))
case Int32:
a.i32s[i] = int32(int64(v))
case Uint32:
a.u32s[i] = uint32(int64(v))
case Float16:
a.halves[i] = HalfFromFloat64(v)
case Float32:
a.floats32[i] = float32(v)
case Float:
a.floats[i] = v
case Complex:
a.complexes[i] = complex(v, 0)
default:
// The alloc convention: an unlisted ordinal carries the uint32
// payload, and no validated constructor can produce one.
a.u32s[i] = uint32(int64(v))
}
}
func containsNeg(c []int) bool {
for _, v := range c {
if v < 0 {
return true
}
}
return false
}