Files
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

1247 lines
36 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
// Package tensor provides immutable, shape-checked numeric arrays for
// Go.
//
// tensor is a general-purpose scientific computing library: it adds no
// dependency at all. The founding rules: values are immutable, shapes never
// broadcast silently, and elements are int64 and float64 first.
package core
import (
"fmt"
"math"
"slices"
"strconv"
"strings"
)
// Dtype is the element type of an Array. Mixing dtypes in one operation
// promotes along the ladder int to float16 to float32 to float64 to
// complex128. The constants' numeric values are
// part of the recorded oracle contract and stay fixed: float16 slots
// into the ladder at its own value, and promote runs through
// dtypeRank rather than through the ordinals.
type Dtype uint8
const (
// Int is the int64 element type.
Int Dtype = iota
// Float32 is the float32 element type: the ML memory and bandwidth
// dtype.
Float32
// Float is the float64 element type, the default real dtype.
Float
// Complex is the complex128 element type.
Complex
// Float16 is the IEEE 754 binary16 element type: the half payload
// holds uint16 bit patterns, widened exactly on every read.
// It sits between Int and Float32 on the promotion ladder (see
// dtypeRank), one step below float32 in both precision and range.
Float16
// Bool is the boolean element type: logical vectors for masking,
// Where and masked reads. It carries no arithmetic of its own; a
// binary arithmetic operation on Bool operands is a loud error.
Bool
// Int8 is the int8 element type: signed bytes.
Int8
// Uint8 is the uint8 element type: byte payloads for images, masks
// and the byte classes of the file formats.
Uint8
// Int16 is the int16 element type.
Int16
// Uint16 is the uint16 element type.
Uint16
// Int32 is the int32 element type.
Int32
// Uint32 is the uint32 element type.
Uint32
)
// String renders the dtype as it appears in diagnostics: "bool",
// "int", "int8", "uint8", "int16", "uint16", "int32", "uint32",
// "float16", "float32", "float", "complex".
func (d Dtype) String() string {
switch d {
case Bool:
return "bool"
case Float16:
return "float16"
case Float32:
return "float32"
case Float:
return "float"
case Complex:
return "complex"
case Int8:
return "int8"
case Uint8:
return "uint8"
case Int16:
return "int16"
case Uint16:
return "uint16"
case Int32:
return "int32"
case Uint32:
return "uint32"
default:
return "int"
}
}
// Array is an immutable, shape-checked numeric array: every
// operation returns a new Array and writes neither its receiver nor its
// arguments. Elements are int64, float16 (uint16 half bit patterns),
// float32, float64 or complex128, any number of
// dimensions is allowed, and storage is row-major.
//
// A result may be a read-only view sharing another array's storage.
// Such a view's payload is rebased to the view's origin, so element 0
// of the view is payload[0] and element i is payload[i]: the package
// never sets strides on an array it produces, which makes physIndex
// the identity for every array from a public constructor. The strides
// field is defensive machinery for strided arrays constructed directly
// in tests; kernels that read payloads raw rely on it staying nil.
type Array struct {
shape []int
dt Dtype
ints []int64
halves []uint16
floats32 []float32
floats []float64
complexes []complex128
bools []bool
i8s []int8
u8s []uint8
i16s []int16
u16s []uint16
i32s []int32
u32s []uint32
// strides is nil for a contiguous array and non-nil only for a
// read-only view; it is never written through.
strides []int
}
// Shape returns a copy of the array's dimensions.
func (a *Array) Shape() []int {
out := make([]int, len(a.shape))
copy(out, a.shape)
return out
}
// Dtype returns the element type.
func (a *Array) Dtype() Dtype { return a.dt }
// Len returns the number of elements: the product of the shape, which
// for a view is shorter than the storage it aliases.
func (a *Array) Len() int {
n := 1
for _, d := range a.shape {
n *= d
}
return n
}
// isContiguous reports whether payload[i] is element i. Only a
// contiguous array may be handed to code that reads payload windows
// directly.
func (a *Array) isContiguous() bool { return a.strides == nil }
// physIndex maps a logical row-major flat index to the payload index.
// A contiguous array maps identity, so payload[i] is element i; a view
// decomposes the index into coordinates and applies its strides.
func (a *Array) physIndex(flat int) int {
if a.strides == nil {
return flat
}
off := 0
for d := a.NDim() - 1; d >= 0; d-- {
off += (flat % a.shape[d]) * a.strides[d]
flat /= a.shape[d]
}
return off
}
// materialise returns an array owning a contiguous payload. A
// contiguous array is returned as is; a strided view is copied into a
// fresh dense array: the boundary a kernel calls when it needs raw
// payload windows rather than accessor reads.
func (a *Array) materialise() *Array {
if a.strides == nil {
return a
}
n := a.Len()
out := &Array{shape: a.Shape(), dt: a.dt}
out.alloc(n)
// The dtype dispatch sits outside the loop: the per-element work is
// the physIndex rebasing, not a switch re-evaluated n times.
switch a.dt {
case Int:
for i := range n {
out.ints[i] = a.ints[a.physIndex(i)]
}
case Float16:
for i := range n {
out.halves[i] = a.halves[a.physIndex(i)]
}
case Float32:
for i := range n {
out.floats32[i] = a.floats32[a.physIndex(i)]
}
case Float:
for i := range n {
out.floats[i] = a.floats[a.physIndex(i)]
}
case Complex:
for i := range n {
out.complexes[i] = a.complexes[a.physIndex(i)]
}
case Bool:
for i := range n {
out.bools[i] = a.bools[a.physIndex(i)]
}
case Int8:
for i := range n {
out.i8s[i] = a.i8s[a.physIndex(i)]
}
case Uint8:
for i := range n {
out.u8s[i] = a.u8s[a.physIndex(i)]
}
case Int16:
for i := range n {
out.i16s[i] = a.i16s[a.physIndex(i)]
}
case Uint16:
for i := range n {
out.u16s[i] = a.u16s[a.physIndex(i)]
}
case Int32:
for i := range n {
out.i32s[i] = a.i32s[a.physIndex(i)]
}
default:
for i := range n {
out.u32s[i] = a.u32s[a.physIndex(i)]
}
}
return out
}
// NDim returns the number of dimensions.
func (a *Array) NDim() int { return len(a.shape) }
// FromInts builds an int array of the given shape from vals, copying them:
// later changes to vals never reach the array.
func FromInts(vals []int64, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
ints := make([]int64, len(vals))
copy(ints, vals)
return &Array{shape: sh, dt: Int, ints: ints}, nil
}
// FromFloat32s builds a float32 array of the given shape from vals,
// copying them.
func FromFloat32s(vals []float32, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
floats32 := make([]float32, len(vals))
copy(floats32, vals)
return &Array{shape: sh, dt: Float32, floats32: floats32}, nil
}
// FromFloat32Slice aliases vals into a float32 array of the given
// shape WITHOUT copying. The returned array is valid only while the
// caller leaves vals untouched: reads reflect later writes, so the
// pattern fits build-then-consume windows (generation caches, staging
// slabs) and nothing else. Alias everything or nothing: no offset,
// no strides.
func FromFloat32Slice(vals []float32, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
return &Array{shape: sh, dt: Float32, floats32: vals}, nil
}
// FromFloatSlice aliases vals into a float array of the given shape
// without copying: the float64 twin of FromFloat32Slice with the same
// build-then-consume contract.
func FromFloatSlice(vals []float64, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
return &Array{shape: sh, dt: Float, floats: vals}, nil
}
// FromFloats builds a float array of the given shape from vals, copying
// them.
func FromFloats(vals []float64, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
floats := make([]float64, len(vals))
copy(floats, vals)
return &Array{shape: sh, dt: Float, floats: floats}, nil
}
// FromComplexes builds a complex array of the given shape from vals,
// copying them.
func FromComplexes(vals []complex128, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
complexes := make([]complex128, len(vals))
copy(complexes, vals)
return &Array{shape: sh, dt: Complex, complexes: complexes}, nil
}
// Zeros builds an array of the given dtype and shape filled with zeros.
// The dtype must be one of the five element types the package stores;
// anything else is an error, not a payload every diagnostic renders
// as int.
func Zeros(dt Dtype, shape ...int) (*Array, error) {
return filled(dt, shape, 0, 0, 0)
}
// Ones builds an array of the given dtype and shape filled with ones.
func Ones(dt Dtype, shape ...int) (*Array, error) {
return filled(dt, shape, 1, 1, 1)
}
// FullI builds an int array of the given shape filled with v.
func FullI(v int64, shape ...int) (*Array, error) {
return filled(Int, shape, v, float64(v), complex(float64(v), 0))
}
// FullF builds a float array of the given shape filled with v.
func FullF(v float64, shape ...int) (*Array, error) {
return filled(Float, shape, int64(v), v, complex(v, 0))
}
// FullF32s builds a float32 array of the given shape filled with v.
func FullF32s(v float32, shape ...int) (*Array, error) {
return filled(Float32, shape, int64(v), float64(v), complex(float64(v), 0))
}
// FullF16 builds a float16 array of the given shape filled with v,
// narrowed to the nearest half under the HalfFromFloat64 contract.
func FullF16(v float64, shape ...int) (*Array, error) {
return filled(Float16, shape, int64(v), v, complex(v, 0))
}
// FullC builds a complex array of the given shape filled with v.
func FullC(v complex128, shape ...int) (*Array, error) {
return filled(Complex, shape, int64(real(v)), real(v), v)
}
func filled(dt Dtype, shape []int, iv int64, fv float64, cv complex128) (*Array, error) {
// filled is the throat every public filling constructor (Zeros,
// Ones, FullI, FullF, FullF16, FullF32s, FullC) and everything built
// on Zeros passes through, so the dtype check lives here once: a
// caller's Dtype(99) is an error, not a payload every diagnostic
// renders as int. The switch is over the element types the package
// stores, and the kernels that build arrays through alloc directly
// never pay it.
switch dt {
case Int, Float16, Float32, Float, Complex, Bool, Int8, Uint8, Int16, Uint16, Int32, Uint32:
default:
return nil, errf("unknown element type %d", uint8(dt))
}
total, sh, err := checkedDims(shape)
if err != nil {
return nil, err
}
out := &Array{shape: sh, dt: dt}
out.alloc(total)
switch dt {
case Int:
for i := range out.ints {
out.ints[i] = iv
}
case Float16:
hv := HalfFromFloat64(fv)
for i := range out.halves {
out.halves[i] = hv
}
case Float32:
f32v := float32(fv)
for i := range out.floats32 {
out.floats32[i] = f32v
}
case Float:
for i := range out.floats {
out.floats[i] = fv
}
case Complex:
for i := range out.complexes {
out.complexes[i] = cv
}
case Bool:
bv := iv != 0
for i := range out.bools {
out.bools[i] = bv
}
case Int8:
v := int8(iv)
for i := range out.i8s {
out.i8s[i] = v
}
case Uint8:
v := uint8(iv)
for i := range out.u8s {
out.u8s[i] = v
}
case Int16:
v := int16(iv)
for i := range out.i16s {
out.i16s[i] = v
}
case Uint16:
v := uint16(iv)
for i := range out.u16s {
out.u16s[i] = v
}
case Int32:
v := int32(iv)
for i := range out.i32s {
out.i32s[i] = v
}
default:
v := uint32(iv)
for i := range out.u32s {
out.u32s[i] = v
}
}
return out, nil
}
// Range builds the int array start, start+1, …, stop-1; start >= stop
// yields an empty array.
func Range(start, stop int64) (*Array, error) {
return RangeBy(start, stop, 1)
}
// RangeBy builds the int array start, start+step, …, staying below stop
// for a positive step and above it for a negative one. A zero step is an
// error, and the walk stops early rather than looping when the next value
// would overflow int64.
func RangeBy(start, stop, step int64) (*Array, error) {
if step == 0 {
return nil, errf("RangeBy: step cannot be zero")
}
var vals []int64
for v := start; (step > 0 && v < stop) || (step < 0 && v > stop); {
vals = append(vals, v)
next := v + step
if (step > 0 && next < v) || (step < 0 && next > v) {
break // the addition wrapped; nothing sane remains
}
v = next
}
return FromInts(vals, len(vals))
}
// Equal reports whether two arrays have the same dtype, shape and values.
// The dtype is part of the identity: an int 1 does not equal a float 1.0.
// Floats compare with ==, so arrays holding NaN are never equal.
func Equal(a, b *Array) bool {
if a == b {
return true
}
if a == nil || b == nil {
return false
}
if a.dt != b.dt || !sameShape(a.shape, b.shape) {
return false
}
// Compare exactly the arrays' own elements: a rebased view's payload
// may run past its element count, and those invisible tail slots
// must not influence equality. slices.Equal compares with ==, so
// NaN never equals NaN, the documented Equal semantics. A strided
// view's payload is not in element order, so both sides are
// materialised first; a contiguous array is returned unchanged.
a = a.materialise()
b = b.materialise()
n := a.Len()
switch a.dt {
case Int:
return slices.Equal(a.ints[:n], b.ints[:n])
case Float16:
// Half payloads compare by value, not by bits: +0.0 and -0.0
// compare equal the way every float dtype's == does, and two
// NaNs never compare equal.
return slices.EqualFunc(a.halves[:n], b.halves[:n], func(x, y uint16) bool {
return HalfToFloat64(x) == HalfToFloat64(y)
})
case Float32:
return slices.Equal(a.floats32[:n], b.floats32[:n])
case Float:
return slices.Equal(a.floats[:n], b.floats[:n])
case Complex:
return slices.Equal(a.complexes[:n], b.complexes[:n])
case Bool:
return slices.Equal(a.bools[:n], b.bools[:n])
case Int8:
return slices.Equal(a.i8s[:n], b.i8s[:n])
case Uint8:
return slices.Equal(a.u8s[:n], b.u8s[:n])
case Int16:
return slices.Equal(a.i16s[:n], b.i16s[:n])
case Uint16:
return slices.Equal(a.u16s[:n], b.u16s[:n])
case Int32:
return slices.Equal(a.i32s[:n], b.i32s[:n])
default:
return slices.Equal(a.u32s[:n], b.u32s[:n])
}
}
// String renders the dtype, the shape and the values, as in
// "int (2, 2) [1, 2, 3, 4]". Arrays of three or more dimensions wrap
// the values per trailing dimension so the structure is readable:
// "float (2, 2, 2) [[1, 2, 3, 4], [5, 6, 7, 8]]". It is a debugging
// aid, not a format.
func (a *Array) String() string {
var sb strings.Builder
fmt.Fprintf(&sb, "%s %s ", a.dt, shapeText(a.shape))
if a.NDim() >= 3 {
// writeSlice emits the full bracket structure, including the
// outermost pair.
a.writeSlice(&sb, 0, 0)
} else {
sb.WriteByte('[')
a.writeValues(&sb)
sb.WriteByte(']')
}
return sb.String()
}
// writeValues renders the flat values with brackets per dimension for
// arrays of rank ≥ 3; ranks 1 and 2 stay flat, matching the compact
// diagnostic format the tests and examples rely on.
func (a *Array) writeValues(sb *strings.Builder) {
nd := a.NDim()
if nd <= 2 {
for i := range a.Len() {
if i > 0 {
sb.WriteString(", ")
}
writeElem(sb, a, i)
}
return
}
// Recursively emit each slice along the leading dimension.
a.writeSlice(sb, 0, 0)
}
// writeSlice emits the elements of a[coord...] with a bracket per
// remaining dimension. Dimensions of size 1 are transparent: a shape
// like (1, 1, 3) renders as [7, 9, 11], not [[[7, 9, 11]]].
func (a *Array) writeSlice(sb *strings.Builder, dim, flat int) {
if a.shape[dim] == 1 && dim < a.NDim()-1 {
a.writeSlice(sb, dim+1, flat)
return
}
sb.WriteByte('[')
block := 1
for d := dim + 1; d < a.NDim(); d++ {
block *= a.shape[d]
}
for i := range a.shape[dim] {
if i > 0 {
sb.WriteString(", ")
}
if dim == a.NDim()-1 {
writeElem(sb, a, flat+i)
} else {
a.writeSlice(sb, dim+1, flat+i*block)
}
}
sb.WriteByte(']')
}
// writeElem renders one element in its dtype's format.
func writeElem(sb *strings.Builder, a *Array, i int) {
if a.strides != nil {
i = a.physIndex(i)
}
switch a.dt {
case Int:
fmt.Fprintf(sb, "%d", a.ints[i])
case Float16:
// Printed through the exact float64 widening, the same 'g'
// shortest-round-trip formatting float32 and float use.
sb.WriteString(strconv.FormatFloat(HalfToFloat64(a.halves[i]), 'g', -1, 64))
case Float32:
sb.WriteString(strconv.FormatFloat(float64(a.floats32[i]), 'g', -1, 32))
case Float:
sb.WriteString(strconv.FormatFloat(a.floats[i], 'g', -1, 64))
case Complex:
fmt.Fprintf(sb, "%v", a.complexes[i])
case Bool:
sb.WriteString(strconv.FormatBool(a.bools[i]))
case Int8:
fmt.Fprintf(sb, "%d", a.i8s[i])
case Uint8:
fmt.Fprintf(sb, "%d", a.u8s[i])
case Int16:
fmt.Fprintf(sb, "%d", a.i16s[i])
case Uint16:
fmt.Fprintf(sb, "%d", a.u16s[i])
case Int32:
fmt.Fprintf(sb, "%d", a.i32s[i])
default:
fmt.Fprintf(sb, "%d", a.u32s[i])
}
}
// checkedDims validates a shape: at least one dimension, none negative,
// and an element count that fits in an int. It returns the element count
// and a private copy of the shape.
func checkedDims(shape []int) (int, []int, error) {
if len(shape) == 0 {
return 0, nil, errf("an array needs at least one dimension")
}
total := 1
for _, d := range shape {
if d < 0 {
return 0, nil, errf("dimensions must be zero or greater, got %d", d)
}
// Bound each factor before multiplying it in: a product that wraps
// would otherwise agree with a caller's own wrapped arithmetic and
// silently pair a huge declared shape with a tiny allocation.
if d != 0 && total > math.MaxInt/d {
return 0, nil, errf("the shape %s holds more elements than fit in an index", shapeText(shape))
}
total *= d
}
sh := make([]int, len(shape))
copy(sh, shape)
return total, sh, nil
}
// shapeFor validates a shape for a payload of exactly n values and
// returns a private copy of it: the shared front half of every From…
// constructor. The shape must be well formed and hold all n values.
func shapeFor(shape []int, n int) ([]int, error) {
total, sh, err := checkedDims(shape)
if err != nil {
return nil, err
}
if total != n {
return nil, errf("%d values do not fill the shape %s", n, shapeText(sh))
}
return sh, nil
}
// sameShape reports whether two shapes are identical.
func sameShape(a, b []int) bool { return slices.Equal(a, b) }
// floatAt returns element i as float64, widening every element type the
// package stores; each widening of an integer, a bool, a half or a
// float32 is exact. It must not be called on complex arrays; callers
// dispatch on the promoted dtype first.
func (a *Array) floatAt(i int) float64 {
if a.strides != nil {
i = a.physIndex(i)
}
switch a.dt {
case Float:
return a.floats[i]
case Float16:
return HalfToFloat64(a.halves[i])
case Float32:
return float64(a.floats32[i])
case Bool:
if a.bools[i] {
return 1
}
return 0
case Int8:
return float64(a.i8s[i])
case Uint8:
return float64(a.u8s[i])
case Int16:
return float64(a.i16s[i])
case Uint16:
return float64(a.u16s[i])
case Int32:
return float64(a.i32s[i])
case Uint32:
return float64(a.u32s[i])
default:
return float64(a.ints[i])
}
}
// float32At returns element i as float32, narrowing float64 elements,
// used where a Float32 result must be written from a wider computation.
func (a *Array) float32At(i int) float32 {
if a.dt == Float32 {
if a.strides != nil {
i = a.physIndex(i)
}
return a.floats32[i]
}
return float32(a.floatAt(i))
}
// complexAt returns element i as complex128, converting every real
// element exactly.
func (a *Array) complexAt(i int) complex128 {
if a.strides != nil {
i = a.physIndex(i)
}
switch a.dt {
case Complex:
return a.complexes[i]
case Float:
return complex(a.floats[i], 0)
case Float16:
return complex(HalfToFloat64(a.halves[i]), 0)
default:
return complex(a.floatAt(i), 0)
}
}
// intAt returns element i as int64, resolving the stride table a
// test-side view carries. The integer-class dtypes widen exactly; bool
// reads 0 or 1; a float element casts the way an implicit store into an
// int destination has always cast. The elementwise fallback reads its
// operands through it, and setConverted reads narrow sources through it.
func (a *Array) intAt(i int) int64 {
if a.strides != nil {
i = a.physIndex(i)
}
switch a.dt {
case Int:
return a.ints[i]
case Bool:
if a.bools[i] {
return 1
}
return 0
case Int8:
return int64(a.i8s[i])
case Uint8:
return int64(a.u8s[i])
case Int16:
return int64(a.i16s[i])
case Uint16:
return int64(a.u16s[i])
case Int32:
return int64(a.i32s[i])
case Uint32:
return int64(a.u32s[i])
case Float16:
return int64(HalfToFloat64(a.halves[i]))
case Float32:
return int64(a.floats32[i])
case Float:
return int64(a.floats[i])
default:
return int64(real(a.complexes[i]))
}
}
// cloneData returns a deep copy of the array's own elements for the
// five legacy payloads (int64, half, float32, float64, complex128):
// for a view that is the view's contents, not the storage it aliases.
// The narrow element types clone through cloneArray, which carries
// their payloads; a narrow array reaching this function has no payload
// field here, so the default arm copies from its nil complex payload
// and fabricates zero complex values, which is why every caller
// dispatches to cloneArray before it reaches this function. A
// contiguous payload copies with a single memcpy per dtype; only a
// strided view pays the per-element physIndex walk. Exactly one of the
// five slices is non-nil, matching the array's legacy dtype.
func (a *Array) cloneData() ([]int64, []uint16, []float32, []float64, []complex128) {
if a.strides == nil {
// A contiguous array's own elements are payload[:Len]; a
// rebased view's payload may run further, so the copy is
// bounded by n rather than trusting the payload length.
n := a.Len()
switch a.dt {
case Int:
ints := make([]int64, n)
copy(ints, a.ints)
return ints, nil, nil, nil, nil
case Float16:
halves := make([]uint16, n)
copy(halves, a.halves)
return nil, halves, nil, nil, nil
case Float32:
floats32 := make([]float32, n)
copy(floats32, a.floats32)
return nil, nil, floats32, nil, nil
case Float:
floats := make([]float64, n)
copy(floats, a.floats)
return nil, nil, nil, floats, nil
default:
complexes := make([]complex128, n)
copy(complexes, a.complexes)
return nil, nil, nil, nil, complexes
}
}
n := a.Len()
switch a.dt {
case Int:
ints := make([]int64, n)
for i := range n {
ints[i] = a.ints[a.physIndex(i)]
}
return ints, nil, nil, nil, nil
case Float16:
halves := make([]uint16, n)
for i := range n {
halves[i] = a.halves[a.physIndex(i)]
}
return nil, halves, nil, nil, nil
case Float32:
floats32 := make([]float32, n)
for i := range n {
floats32[i] = a.floats32[a.physIndex(i)]
}
return nil, nil, floats32, nil, nil
case Float:
floats := make([]float64, n)
for i := range n {
floats[i] = a.floats[a.physIndex(i)]
}
return nil, nil, nil, floats, nil
default:
complexes := make([]complex128, n)
for i := range n {
complexes[i] = a.complexes[a.physIndex(i)]
}
return nil, nil, nil, nil, complexes
}
}
// alloc prepares the payload for total elements of the array's dtype.
func (a *Array) alloc(total int) {
switch a.dt {
case Int:
a.ints = make([]int64, total)
case Float16:
a.halves = make([]uint16, total)
case Float32:
a.floats32 = make([]float32, total)
case Float:
a.floats = make([]float64, total)
case Complex:
a.complexes = make([]complex128, total)
case Bool:
a.bools = make([]bool, total)
case Int8:
a.i8s = make([]int8, total)
case Uint8:
a.u8s = make([]uint8, total)
case Int16:
a.i16s = make([]int16, total)
case Uint16:
a.u16s = make([]uint16, total)
case Int32:
a.i32s = make([]int32, total)
default:
a.u32s = make([]uint32, total)
}
}
// setFrom copies element srcFlat of src into element dstFlat of a. The
// dtypes must match; the destination is always a freshly allocated,
// contiguous array, while the source may be a view.
func (a *Array) setFrom(dstFlat int, src *Array, srcFlat int) {
srcFlat = src.physIndex(srcFlat)
switch a.dt {
case Int:
a.ints[dstFlat] = src.ints[srcFlat]
case Float16:
a.halves[dstFlat] = src.halves[srcFlat]
case Float32:
a.floats32[dstFlat] = src.floats32[srcFlat]
case Float:
a.floats[dstFlat] = src.floats[srcFlat]
case Complex:
a.complexes[dstFlat] = src.complexes[srcFlat]
case Bool:
a.bools[dstFlat] = src.bools[srcFlat]
case Int8:
a.i8s[dstFlat] = src.i8s[srcFlat]
case Uint8:
a.u8s[dstFlat] = src.u8s[srcFlat]
case Int16:
a.i16s[dstFlat] = src.i16s[srcFlat]
case Uint16:
a.u16s[dstFlat] = src.u16s[srcFlat]
case Int32:
a.i32s[dstFlat] = src.i32s[srcFlat]
default:
a.u32s[dstFlat] = src.u32s[srcFlat]
}
}
// canStore reports whether an element of dtype src may be stored in a
// destination of dtype dst. The ladder runs up (int to float16 to
// float32 to float to complex) and the descending directions the
// library performs are float to int, float to float32, float to float16
// and complex to float (the real part), all as Astype documents;
// narrowing complex into int, float16 or float32 is the pair Astype
// rejects, and every store respects that.
func canStore(dst, src Dtype) bool {
if src == Complex {
// Narrowing complex into any integer destination is the pair
// Astype rejects, and every store respects that; complex to
// float keeps the real part, and complex to bool reads against
// zero, the same test Astype's bool target applies.
return dst != Int && dst != Float16 && dst != Float32 &&
dst != Int8 && dst != Uint8 &&
dst != Int16 && dst != Uint16 && dst != Int32 && dst != Uint32
}
return true
}
// setConverted stores element srcFlat of src in a's dtype, converting
// along the ladder, used when Concat and Stack promote their operands
// and when Scatter mixes dtypes. The read goes through the source's own
// accessor, so a source above the destination on the ladder never reads
// a payload the source does not have; the rejected narrowings are
// decided by canStore before the walk.
func (a *Array) setConverted(dstFlat int, src *Array, srcFlat int) {
switch a.dt {
case Int:
if src.dt == Int {
// Exact: the float64 detour rounds above 2^53.
a.ints[dstFlat] = src.ints[src.physIndex(srcFlat)]
return
}
a.ints[dstFlat] = int64(src.floatAt(srcFlat))
case Float16:
a.halves[dstFlat] = HalfFromFloat64(src.floatAt(srcFlat))
case Float32:
a.floats32[dstFlat] = src.float32At(srcFlat)
case Float:
if src.dt == Complex {
// Complex to float keeps the real part, as Astype documents.
a.floats[dstFlat] = real(src.complexAt(srcFlat))
return
}
a.floats[dstFlat] = src.floatAt(srcFlat)
case Complex:
a.complexes[dstFlat] = src.complexAt(srcFlat)
case Bool:
a.bools[dstFlat] = src.boolAt(srcFlat)
case Int8:
a.i8s[dstFlat] = int8(src.intAt(srcFlat))
case Uint8:
a.u8s[dstFlat] = uint8(src.intAt(srcFlat))
case Int16:
a.i16s[dstFlat] = int16(src.intAt(srcFlat))
case Uint16:
a.u16s[dstFlat] = uint16(src.intAt(srcFlat))
case Int32:
a.i32s[dstFlat] = int32(src.intAt(srcFlat))
default:
// Uint32 and any unlisted ordinal: the widened read casts down,
// the implicit-store semantics every promoted Concat and Scatter
// target has always carried.
a.u32s[dstFlat] = uint32(src.intAt(srcFlat))
}
}
// dtypeRank returns a dtype's position on the numeric tower: bool
// below the integer widths below int64 below float16 below float32
// below float64 below complex128. Mixed signedness integer pairs are
// not resolved by rank alone; promote resolves the whole integer class
// through intPromote in narrowdt.go, and only cross-class pairs walk
// this ladder.
func dtypeRank(d Dtype) int {
switch d {
case Bool:
return 0
case Int8, Uint8:
return 1
case Int16, Uint16:
return 2
case Int32, Uint32:
return 3
case Float16:
return 5
case Float32:
return 6
case Float:
return 7
case Complex:
return 8
default:
return 4
}
}
// promote returns the common dtype under the library's numeric tower.
// Across classes the higher ladder position decides, the rule the
// int-to-float-to-complex promotions have always walked; inside the
// integer class the result is the smallest integer dtype whose value
// range contains both operands, so a mixed signedness pair widens
// instead of losing its negative half (int8 with uint8 answers int16).
func promote(a, b Dtype) Dtype {
if a == b {
return a
}
if intClass(a) && intClass(b) {
return intPromote[intClassIndex(a)][intClassIndex(b)]
}
if dtypeRank(a) >= dtypeRank(b) {
return a
}
return b
}
// Strided reports whether the array carries a non-trivial stride
// layout. Dense arrays answer false.
func (a *Array) Strided() bool { return a.strides != nil }
// RawFloats returns the array's float64 payload directly: element i
// of the array sits at payload index i, views included, because the
// package never sets strides. Treat the slice as read-only; only
// freshly allocated arrays an owner writes through it are safe to
// mutate.
func (a *Array) RawFloats() []float64 { return a.floats }
// RawFloat32s returns the float32 payload with the RawFloats
// contract.
func (a *Array) RawFloat32s() []float32 { return a.floats32 }
// RawInts returns the int64 payload with the RawFloats contract.
func (a *Array) RawInts() []int64 { return a.ints }
// RawComplexes returns the complex128 payload with the RawFloats
// contract.
func (a *Array) RawComplexes() []complex128 { return a.complexes }
// FloatsFromArray builds an array that takes ownership of vals: no
// copy is made, so the caller must not touch the slice afterwards.
// The value count must fill the shape exactly.
func FloatsFromArray(vals []float64, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
return &Array{shape: sh, dt: Float, floats: vals}, nil
}
// ComplexFromArray builds an array that takes ownership of vals,
// with the FloatsFromArray contract.
func ComplexFromArray(vals []complex128, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
return &Array{shape: sh, dt: Complex, complexes: vals}, nil
}
// IntsFromArray builds an array that takes ownership of vals, with
// the FloatsFromArray contract.
func IntsFromArray(vals []int64, shape ...int) (*Array, error) {
sh, err := shapeFor(shape, len(vals))
if err != nil {
return nil, err
}
return &Array{shape: sh, dt: Int, ints: vals}, nil
}
// New allocates a zeroed array of the given dtype and shape, the
// no-error constructor for internally derived shapes: the count and
// shape come from existing arrays, so an invalid argument is a bug in
// the caller, answered by a nil array.
func New(dt Dtype, shape ...int) *Array {
a, err := Zeros(dt, shape...)
if err != nil {
return nil
}
return a
}
// ComplexValues returns the array's elements as complex values,
// reading a complex payload directly and converting everything else.
// A contiguous complex array shares its payload, so the slice is a
// read-only alias; a strided view is copied out element by element.
func (a *Array) ComplexValues(name string) ([]complex128, error) {
if a.dt == Complex {
if a.strides == nil {
// A rebased view's payload runs past its element count, and
// only the prefix bounded by Len is the view's own data.
return a.complexes[:a.Len()], nil
}
out := make([]complex128, a.Len())
for i := range out {
out[i] = a.complexes[a.physIndex(i)]
}
return out, nil
}
if a.NDim() != 1 {
return nil, errf("%s: needs a 1-D array, got shape %s", name, shapeText(a.shape))
}
if a.Len() == 0 {
return nil, errf("%s: an empty array has no transform", name)
}
out := make([]complex128, a.Len())
for i := range out {
out[i] = a.complexAt(i)
}
return out, nil
}
// ComplexAt returns element i as complex128, converting real
// elements; every real dtype widens exactly, bool reads 0/1.
func (a *Array) ComplexAt(i int) complex128 {
if a.strides != nil {
i = a.physIndex(i)
}
switch a.dt {
case Complex:
return a.complexes[i]
case Float:
return complex(a.floats[i], 0)
case Float16:
return complex(HalfToFloat64(a.halves[i]), 0)
case Float32:
return complex(float64(a.floats32[i]), 0)
default:
return complex(a.floatAt(i), 0)
}
}
// FloatAt returns element i widened to float64. It is the numeric
// read primitive the derivative packages build on; for typed access
// prefer IntAt/FloatAt-by-name variants and Elements[E].
func (a *Array) FloatAt(i int) float64 { return a.floatAt(i) }
// SetFloatAt sets element i from v, converting to the array's dtype,
// the numeric write primitive complementing FloatAt.
func (a *Array) SetFloatAt(i int, v float64) { a.setFromValue(i, v) }
// CopyRows returns a new 2+-D array holding the rows of a at the given
// leading-dimension indices, preserving every other dimension and the
// element type. An index outside the leading dimension is an error, as
// in every other selecting entry.
func (a *Array) CopyRows(idx []int) (*Array, error) {
for _, r := range idx {
if r < 0 || r >= a.shape[0] {
return nil, errf("CopyRows: row %d out of range for %d rows", r, a.shape[0])
}
}
rowLen := 1
for _, d := range a.shape[1:] {
rowLen *= d
}
out := &Array{shape: append([]int{len(idx)}, a.shape[1:]...), dt: a.dt}
out.alloc(len(idx) * rowLen)
if a.strides == nil {
// Contiguous rows copy whole: one memcpy per selected row, no
// per-element dispatch. Views fall through to setFrom, which
// rebases each flat index through the strides.
switch a.dt {
case Int:
for i, r := range idx {
copy(out.ints[i*rowLen:(i+1)*rowLen], a.ints[r*rowLen:(r+1)*rowLen])
}
case Float16:
for i, r := range idx {
copy(out.halves[i*rowLen:(i+1)*rowLen], a.halves[r*rowLen:(r+1)*rowLen])
}
case Float32:
for i, r := range idx {
copy(out.floats32[i*rowLen:(i+1)*rowLen], a.floats32[r*rowLen:(r+1)*rowLen])
}
case Float:
for i, r := range idx {
copy(out.floats[i*rowLen:(i+1)*rowLen], a.floats[r*rowLen:(r+1)*rowLen])
}
case Complex:
for i, r := range idx {
copy(out.complexes[i*rowLen:(i+1)*rowLen], a.complexes[r*rowLen:(r+1)*rowLen])
}
case Bool:
for i, r := range idx {
copy(out.bools[i*rowLen:(i+1)*rowLen], a.bools[r*rowLen:(r+1)*rowLen])
}
case Int8:
for i, r := range idx {
copy(out.i8s[i*rowLen:(i+1)*rowLen], a.i8s[r*rowLen:(r+1)*rowLen])
}
case Uint8:
for i, r := range idx {
copy(out.u8s[i*rowLen:(i+1)*rowLen], a.u8s[r*rowLen:(r+1)*rowLen])
}
case Int16:
for i, r := range idx {
copy(out.i16s[i*rowLen:(i+1)*rowLen], a.i16s[r*rowLen:(r+1)*rowLen])
}
case Uint16:
for i, r := range idx {
copy(out.u16s[i*rowLen:(i+1)*rowLen], a.u16s[r*rowLen:(r+1)*rowLen])
}
case Int32:
for i, r := range idx {
copy(out.i32s[i*rowLen:(i+1)*rowLen], a.i32s[r*rowLen:(r+1)*rowLen])
}
default:
for i, r := range idx {
copy(out.u32s[i*rowLen:(i+1)*rowLen], a.u32s[r*rowLen:(r+1)*rowLen])
}
}
return out, nil
}
for i, r := range idx {
for j := range rowLen {
out.setFrom(i*rowLen+j, a, r*rowLen+j)
}
}
return out, nil
}
// Bytes returns the payload interpreted as raw bytes, used by the
// archive writers to embed non-numeric payloads. Only int arrays carry
// byte payloads; other element types return nil.
func (a *Array) Bytes() []byte {
if a.dt != Int {
return nil
}
out := make([]byte, a.Len())
if a.strides == nil {
// Contiguous payload: the first n slots are the array's own
// elements (a rebased view's payload may run further, so the
// walk is bounded by n, not the payload length).
for i := range out {
out[i] = byte(a.ints[i])
}
return out
}
for i := range out {
out[i] = byte(a.ints[a.physIndex(i)])
}
return out
}
// FromBytes wraps raw bytes as an int64 array, the inverse of Bytes.
func FromBytes(b []byte) (*Array, error) {
vals := make([]int64, len(b))
for i, c := range b {
vals[i] = int64(c)
}
return FromInts(vals, len(b))
}