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