316 lines
9.9 KiB
Go
316 lines
9.9 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package spmd
|
|
|
|
import (
|
|
"encoding/binary"
|
|
"math"
|
|
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
)
|
|
|
|
// The wire form of an array is a self-describing frame payload: one
|
|
// dtype byte, one dimension-count byte, the extents as little-endian
|
|
// int64, then the elements as little-endian fixed-width raw bits in
|
|
// row-major order. Every multi-byte field is written little-endian by
|
|
// explicit conversion, never by copying the payload's memory, so the
|
|
// form is the same on every architecture the library builds for.
|
|
// Floating-point payloads keep their exact bit patterns, NaN payloads
|
|
// and signed zeros included.
|
|
|
|
// dtypeWidth is the wire width of one element per dtype.
|
|
func dtypeWidth(dt core.Dtype) int {
|
|
switch dt {
|
|
case core.Int8, core.Uint8, core.Bool:
|
|
return 1
|
|
case core.Int16, core.Uint16, core.Float16:
|
|
return 2
|
|
case core.Int32, core.Uint32, core.Float32:
|
|
return 4
|
|
case core.Int, core.Float:
|
|
return 8
|
|
case core.Complex:
|
|
return 16
|
|
}
|
|
return 0
|
|
}
|
|
|
|
// encodeArray appends the wire form of a to dst and returns the grown
|
|
// slice.
|
|
func encodeArray(dst []byte, a *core.Array) ([]byte, error) {
|
|
return encodePart(dst, a, a.Shape(), 0, a.Len())
|
|
}
|
|
|
|
// encodeHead appends just the head of a wire form: dtype, dimension
|
|
// count and extents. A rank that holds the head can name its piece of
|
|
// the data before any payload flows.
|
|
func encodeHead(dst []byte, dt core.Dtype, shape []int) ([]byte, error) {
|
|
if dtypeWidth(dt) == 0 {
|
|
return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
|
|
}
|
|
if len(shape) > 255 {
|
|
return nil, base.Errf("spmd: an array of %d dimensions exceeds the wire form's 255", len(shape))
|
|
}
|
|
dst = append(dst, byte(dt), byte(len(shape)))
|
|
for _, d := range shape {
|
|
if d < 0 {
|
|
return nil, base.Errf("spmd: a wire form cannot name a negative extent %d", d)
|
|
}
|
|
dst = binary.LittleEndian.AppendUint64(dst, uint64(d))
|
|
}
|
|
return dst, nil
|
|
}
|
|
|
|
// decodeHead reads a wire form's head: the dtype and the shape. It
|
|
// validates only the head's own length; the payload against the shape
|
|
// is decodeArray's check.
|
|
func decodeHead(wire []byte) (core.Dtype, []int, error) {
|
|
if len(wire) < 2 {
|
|
return 0, nil, base.Errf("spmd: a wire form needs at least two bytes, got %d", len(wire))
|
|
}
|
|
dt := core.Dtype(wire[0])
|
|
if dtypeWidth(dt) == 0 {
|
|
return 0, nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
|
|
}
|
|
ndim := int(wire[1])
|
|
if len(wire) < 2+8*ndim {
|
|
return 0, nil, base.Errf("spmd: a wire form for %d dimensions is short by %d bytes", ndim, 2+8*ndim-len(wire))
|
|
}
|
|
shape := make([]int, ndim)
|
|
for i := range shape {
|
|
v := binary.LittleEndian.Uint64(wire[2+8*i:])
|
|
if v > math.MaxInt {
|
|
return 0, nil, base.Errf("spmd: an extent of %d does not fit this machine's int", v)
|
|
}
|
|
shape[i] = int(v)
|
|
}
|
|
return dt, shape, nil
|
|
}
|
|
|
|
// decodeWire builds an array from its complete wire form, head and
|
|
// payload together. Every extent is bounded before anything is
|
|
// allocated: an overflowing shape or a payload that does not carry
|
|
// exactly the named elements is an error, never a partial answer.
|
|
func decodeWire(wire []byte) (*core.Array, error) {
|
|
dt, shape, err := decodeHead(wire)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return decodeArray(dt, shape, wire[2+8*len(shape):])
|
|
}
|
|
|
|
// partHeadLen is the byte length of a one-dimensional part's wire
|
|
// head, the shape encodePart writes for `[]int{count}`: one dtype
|
|
// byte, one dimension-count byte, and the single uint64 extent.
|
|
const partHeadLen = 2 + 8
|
|
|
|
// encodePart appends the wire form of a contiguous element range of a,
|
|
// presented under the given shape: the head names the shape, the
|
|
// payload carries elements [first, first+count) of a. Only the array's
|
|
// own elements take part: a rebased view's payload may run past its
|
|
// element count, so every payload walk stops at Len.
|
|
func encodePart(dst []byte, a *core.Array, shape []int, first, count int) ([]byte, error) {
|
|
dt := a.Dtype()
|
|
width := dtypeWidth(dt)
|
|
if width == 0 {
|
|
return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
|
|
}
|
|
if len(shape) > 255 {
|
|
return nil, base.Errf("spmd: an array of %d dimensions exceeds the wire form's 255", len(shape))
|
|
}
|
|
if first < 0 || count < 0 || first+count > a.Len() {
|
|
return nil, base.Errf("spmd: element range [%d, %d) is outside the array's %d elements", first, first+count, a.Len())
|
|
}
|
|
for _, d := range shape {
|
|
if d < 0 {
|
|
return nil, base.Errf("spmd: a wire form cannot name a negative extent %d", d)
|
|
}
|
|
}
|
|
dst = append(dst, byte(dt), byte(len(shape)))
|
|
for _, d := range shape {
|
|
dst = binary.LittleEndian.AppendUint64(dst, uint64(d))
|
|
}
|
|
start := len(dst)
|
|
switch dt {
|
|
case core.Int:
|
|
for _, v := range a.RawInts()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint64(dst, uint64(v))
|
|
}
|
|
case core.Float:
|
|
for _, v := range a.RawFloats()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(v))
|
|
}
|
|
case core.Float32:
|
|
for _, v := range a.RawFloat32s()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint32(dst, math.Float32bits(v))
|
|
}
|
|
case core.Float16:
|
|
for _, v := range a.RawHalves()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint16(dst, v)
|
|
}
|
|
case core.Complex:
|
|
for _, v := range a.RawComplexes()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(real(v)))
|
|
dst = binary.LittleEndian.AppendUint64(dst, math.Float64bits(imag(v)))
|
|
}
|
|
case core.Bool:
|
|
for _, v := range a.RawBools()[first : first+count] {
|
|
b := byte(0)
|
|
if v {
|
|
b = 1
|
|
}
|
|
dst = append(dst, b)
|
|
}
|
|
case core.Int8:
|
|
for _, v := range a.RawInt8s()[first : first+count] {
|
|
dst = append(dst, byte(v))
|
|
}
|
|
case core.Uint8:
|
|
for _, v := range a.RawUint8s()[first : first+count] {
|
|
dst = append(dst, v)
|
|
}
|
|
case core.Int16:
|
|
for _, v := range a.RawInt16s()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint16(dst, uint16(v))
|
|
}
|
|
case core.Uint16:
|
|
for _, v := range a.RawUint16s()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint16(dst, v)
|
|
}
|
|
case core.Int32:
|
|
for _, v := range a.RawInt32s()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint32(dst, uint32(v))
|
|
}
|
|
case core.Uint32:
|
|
for _, v := range a.RawUint32s()[first : first+count] {
|
|
dst = binary.LittleEndian.AppendUint32(dst, v)
|
|
}
|
|
}
|
|
if got := len(dst) - start; got != count*width {
|
|
return nil, base.Errf("spmd: array of dtype %s encoded %d bytes for %d elements", dt, got, count)
|
|
}
|
|
return dst, nil
|
|
}
|
|
|
|
// decodeArray builds an array from the wire payload of a dtype and
|
|
// shape the caller has read. The payload must carry exactly the
|
|
// elements the shape names at the dtype's width; the caller has
|
|
// already bounded payload by the world's message ceiling.
|
|
func decodeArray(dt core.Dtype, shape []int, payload []byte) (*core.Array, error) {
|
|
width := dtypeWidth(dt)
|
|
if width == 0 {
|
|
return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
|
|
}
|
|
n, ok := elementCount(shape)
|
|
if !ok {
|
|
return nil, base.Errf("spmd: shape %v overflows the element count", shape)
|
|
}
|
|
if int64(len(payload)) != int64(n)*int64(width) {
|
|
return nil, base.Errf("spmd: wire payload of %d bytes does not carry %d elements of %s", len(payload), n, dt)
|
|
}
|
|
switch dt {
|
|
case core.Int:
|
|
vals := make([]int64, n)
|
|
for i := range vals {
|
|
vals[i] = int64(binary.LittleEndian.Uint64(payload[i*8:]))
|
|
}
|
|
return core.FromInts(vals, shape...)
|
|
case core.Float:
|
|
vals := make([]float64, n)
|
|
for i := range vals {
|
|
vals[i] = math.Float64frombits(binary.LittleEndian.Uint64(payload[i*8:]))
|
|
}
|
|
return core.FromFloats(vals, shape...)
|
|
case core.Float32:
|
|
vals := make([]float32, n)
|
|
for i := range vals {
|
|
vals[i] = math.Float32frombits(binary.LittleEndian.Uint32(payload[i*4:]))
|
|
}
|
|
return core.FromFloat32s(vals, shape...)
|
|
case core.Float16:
|
|
vals := make([]uint16, n)
|
|
for i := range vals {
|
|
vals[i] = binary.LittleEndian.Uint16(payload[i*2:])
|
|
}
|
|
return core.HalvesFromArray(vals, shape...)
|
|
case core.Complex:
|
|
vals := make([]complex128, n)
|
|
for i := range vals {
|
|
re := math.Float64frombits(binary.LittleEndian.Uint64(payload[i*16:]))
|
|
im := math.Float64frombits(binary.LittleEndian.Uint64(payload[i*16+8:]))
|
|
vals[i] = complex(re, im)
|
|
}
|
|
return core.FromComplexes(vals, shape...)
|
|
case core.Bool:
|
|
vals := make([]bool, n)
|
|
for i := range vals {
|
|
switch payload[i] {
|
|
case 0:
|
|
case 1:
|
|
vals[i] = true
|
|
default:
|
|
return nil, base.Errf("spmd: bool wire byte %d at index %d is neither 0 nor 1", payload[i], i)
|
|
}
|
|
}
|
|
return core.FromBools(vals, shape...)
|
|
case core.Int8:
|
|
vals := make([]int8, n)
|
|
for i := range vals {
|
|
vals[i] = int8(payload[i])
|
|
}
|
|
return core.FromInt8s(vals, shape...)
|
|
case core.Uint8:
|
|
vals := make([]uint8, n)
|
|
for i := range vals {
|
|
vals[i] = payload[i]
|
|
}
|
|
return core.FromUint8s(vals, shape...)
|
|
case core.Int16:
|
|
vals := make([]int16, n)
|
|
for i := range vals {
|
|
vals[i] = int16(binary.LittleEndian.Uint16(payload[i*2:]))
|
|
}
|
|
return core.FromInt16s(vals, shape...)
|
|
case core.Uint16:
|
|
vals := make([]uint16, n)
|
|
for i := range vals {
|
|
vals[i] = binary.LittleEndian.Uint16(payload[i*2:])
|
|
}
|
|
return core.FromUint16s(vals, shape...)
|
|
case core.Int32:
|
|
vals := make([]int32, n)
|
|
for i := range vals {
|
|
vals[i] = int32(binary.LittleEndian.Uint32(payload[i*4:]))
|
|
}
|
|
return core.FromInt32s(vals, shape...)
|
|
case core.Uint32:
|
|
vals := make([]uint32, n)
|
|
for i := range vals {
|
|
vals[i] = binary.LittleEndian.Uint32(payload[i*4:])
|
|
}
|
|
return core.FromUint32s(vals, shape...)
|
|
}
|
|
return nil, base.Errf("spmd: dtype %s cannot travel the wire", dt)
|
|
}
|
|
|
|
// elementCount is the product of the shape, reported as false when any
|
|
// extent is negative or the product overflows an int.
|
|
func elementCount(shape []int) (int, bool) {
|
|
n := 1
|
|
for _, d := range shape {
|
|
if d < 0 {
|
|
return 0, false
|
|
}
|
|
if d == 0 {
|
|
return 0, true
|
|
}
|
|
if n > math.MaxInt/d {
|
|
return 0, false
|
|
}
|
|
n *= d
|
|
}
|
|
return n, true
|
|
}
|