376 lines
12 KiB
Go
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
|
||
|
|
}
|