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