Files
tensor/spmd/wire.go
T
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

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
}