// Copyright (c) 2026 Petr BalvĂ­n (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 }