176 lines
5.7 KiB
Go
176 lines
5.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package spmd
|
||
|
|
|
||
|
|
import (
|
||
|
|
"bytes"
|
||
|
|
"math"
|
||
|
|
"testing"
|
||
|
|
|
||
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||
|
|
)
|
||
|
|
|
||
|
|
// The wire form is judged by one rule: the array that comes out carries
|
||
|
|
// the exact bits of the array that went in, shape and dtype included.
|
||
|
|
func TestWireRoundTrip(t *testing.T) {
|
||
|
|
nanLow := math.Float64frombits(0x7ff8000000000001) // NaN, low payload bit set
|
||
|
|
nanHigh := math.Float64frombits(0xfff8deadbeef0000)
|
||
|
|
negZero := math.Copysign(0, -1)
|
||
|
|
cases := []*core.Array{
|
||
|
|
mk(core.FromFloats([]float64{1, -2.5, 0, negZero, math.Inf(1), math.Inf(-1), nanLow, nanHigh, math.MaxFloat64, math.SmallestNonzeroFloat64}, 10)),
|
||
|
|
mk(core.FromFloat32s([]float32{1.5, -0.25, float32(negZero), float32(math.Inf(1)), float32(nanLow)}, 5)),
|
||
|
|
mk(core.FromComplexes([]complex128{1 + 2i, complex(negZero, nanHigh), complex(0, negZero)}, 3)),
|
||
|
|
mk(core.FromInts([]int64{math.MaxInt64, math.MinInt64, -1, 0, 42}, 5)),
|
||
|
|
mk(core.FromBools([]bool{true, false, true, true}, 4)),
|
||
|
|
mk(core.FromInt8s([]int8{math.MinInt8, math.MaxInt8, -1, 0}, 4)),
|
||
|
|
mk(core.FromUint8s([]uint8{0, 255, 128}, 3)),
|
||
|
|
mk(core.FromInt16s([]int16{math.MinInt16, math.MaxInt16, -1}, 3)),
|
||
|
|
mk(core.FromUint16s([]uint16{0, 65535, 32768}, 3)),
|
||
|
|
mk(core.FromInt32s([]int32{math.MinInt32, math.MaxInt32, -1}, 3)),
|
||
|
|
mk(core.FromUint32s([]uint32{0, 4294967295, 2147483648}, 3)),
|
||
|
|
// Float16 raw halves: denormals, infinities, NaN payloads, the
|
||
|
|
// whole span the payload can hold.
|
||
|
|
mk(core.HalvesFromArray([]uint16{0x0001, 0x03ff, 0x7bff, 0x7c00, 0xfc00, 0x7e00, 0x7eaa}, 7)),
|
||
|
|
// Empty and multi-dimensional shapes.
|
||
|
|
mk(core.FromFloats(nil, 0)),
|
||
|
|
mk(core.FromInts([]int64{1, 2, 3, 4, 5, 6}, 2, 3)),
|
||
|
|
mk(core.FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12}, 2, 1, 2, 3)),
|
||
|
|
}
|
||
|
|
for i, a := range cases {
|
||
|
|
wire, err := encodeArray(nil, a)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("case %d (%s): %v", i, a.Dtype(), err)
|
||
|
|
}
|
||
|
|
b, err := decodeWire(wire)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("case %d (%s): %v", i, a.Dtype(), err)
|
||
|
|
}
|
||
|
|
if b.Dtype() != a.Dtype() {
|
||
|
|
t.Fatalf("case %d: dtype %s came back as %s", i, a.Dtype(), b.Dtype())
|
||
|
|
}
|
||
|
|
if !sameShape(a.Shape(), b.Shape()) {
|
||
|
|
t.Fatalf("case %d: shape %v came back as %v", i, a.Shape(), b.Shape())
|
||
|
|
}
|
||
|
|
if !sameBits(a, b) {
|
||
|
|
t.Fatalf("case %d (%s): the round trip moved bits", i, a.Dtype())
|
||
|
|
}
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestWireRefusesTheHostile is the hostile-reader rule: a form that
|
||
|
|
// does not carry exactly what its head names is an error, never a short
|
||
|
|
// read and never a partial answer.
|
||
|
|
func TestWireRefusesTheHostile(t *testing.T) {
|
||
|
|
a := mk(core.FromFloats([]float64{1, 2, 3, 4}, 2, 2))
|
||
|
|
wire, err := encodeArray(nil, a)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
payload := wire[2+8*2:]
|
||
|
|
if _, err := decodeWire(wire[:len(wire)-1]); err == nil {
|
||
|
|
t.Fatal("a short payload decoded")
|
||
|
|
}
|
||
|
|
if _, err := decodeWire(append(bytes.Clone(wire), 0)); err == nil {
|
||
|
|
t.Fatal("a long payload decoded")
|
||
|
|
}
|
||
|
|
if _, err := decodeWire(wire[:1]); err == nil {
|
||
|
|
t.Fatal("a headless form decoded")
|
||
|
|
}
|
||
|
|
if _, err := decodeWire(wire[:6]); err == nil {
|
||
|
|
t.Fatal("a truncated head decoded")
|
||
|
|
}
|
||
|
|
big := []int{math.MaxInt32, math.MaxInt32}
|
||
|
|
if _, err := decodeArray(core.Float, big, nil); err == nil {
|
||
|
|
t.Fatal("an overflowing shape decoded")
|
||
|
|
}
|
||
|
|
if _, err := decodeArray(core.Float, []int{2, -2}, payload); err == nil {
|
||
|
|
t.Fatal("a negative extent decoded")
|
||
|
|
}
|
||
|
|
badBool := mk(core.FromBools([]bool{true, false}, 2))
|
||
|
|
bw, err := encodeArray(nil, badBool)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
bp := bytes.Clone(bw)
|
||
|
|
bp[2+8] = 2
|
||
|
|
if _, err := decodeWire(bp); err == nil {
|
||
|
|
t.Fatal("a bool byte of 2 decoded")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// mk builds a fixture or panics; the fixtures are package-level
|
||
|
|
// literals, so a panic lands in the test that declared them.
|
||
|
|
func mk(a *core.Array, err error) *core.Array {
|
||
|
|
if err != nil {
|
||
|
|
panic(err)
|
||
|
|
}
|
||
|
|
return a
|
||
|
|
}
|
||
|
|
|
||
|
|
// sameBits compares two arrays' raw payloads bit for bit, dtype by
|
||
|
|
// dtype. NaN payloads and signed zeros are part of the contract.
|
||
|
|
func sameBits(a, b *core.Array) bool {
|
||
|
|
if a.Len() != b.Len() || a.Dtype() != b.Dtype() {
|
||
|
|
return false
|
||
|
|
}
|
||
|
|
switch a.Dtype() {
|
||
|
|
case core.Int:
|
||
|
|
return slicesEqual(a.RawInts()[:a.Len()], b.RawInts()[:b.Len()])
|
||
|
|
case core.Float:
|
||
|
|
x, y := a.RawFloats()[:a.Len()], b.RawFloats()[:b.Len()]
|
||
|
|
for i := range x {
|
||
|
|
if math.Float64bits(x[i]) != math.Float64bits(y[i]) {
|
||
|
|
return false
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return true
|
||
|
|
case core.Float32:
|
||
|
|
x, y := a.RawFloat32s()[:a.Len()], b.RawFloat32s()[:b.Len()]
|
||
|
|
for i := range x {
|
||
|
|
if math.Float32bits(x[i]) != math.Float32bits(y[i]) {
|
||
|
|
return false
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return true
|
||
|
|
case core.Float16:
|
||
|
|
return slicesEqual(a.RawHalves()[:a.Len()], b.RawHalves()[:b.Len()])
|
||
|
|
case core.Complex:
|
||
|
|
x, y := a.RawComplexes()[:a.Len()], b.RawComplexes()[:b.Len()]
|
||
|
|
for i := range x {
|
||
|
|
if math.Float64bits(real(x[i])) != math.Float64bits(real(y[i])) ||
|
||
|
|
math.Float64bits(imag(x[i])) != math.Float64bits(imag(y[i])) {
|
||
|
|
return false
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return true
|
||
|
|
case core.Bool:
|
||
|
|
return slicesEqual(a.RawBools()[:a.Len()], b.RawBools()[:b.Len()])
|
||
|
|
case core.Int8:
|
||
|
|
return slicesEqual(a.RawInt8s()[:a.Len()], b.RawInt8s()[:b.Len()])
|
||
|
|
case core.Uint8:
|
||
|
|
return slicesEqual(a.RawUint8s()[:a.Len()], b.RawUint8s()[:b.Len()])
|
||
|
|
case core.Int16:
|
||
|
|
return slicesEqual(a.RawInt16s()[:a.Len()], b.RawInt16s()[:b.Len()])
|
||
|
|
case core.Uint16:
|
||
|
|
return slicesEqual(a.RawUint16s()[:a.Len()], b.RawUint16s()[:b.Len()])
|
||
|
|
case core.Int32:
|
||
|
|
return slicesEqual(a.RawInt32s()[:a.Len()], b.RawInt32s()[:b.Len()])
|
||
|
|
case core.Uint32:
|
||
|
|
return slicesEqual(a.RawUint32s()[:a.Len()], b.RawUint32s()[:b.Len()])
|
||
|
|
}
|
||
|
|
return false
|
||
|
|
}
|
||
|
|
|
||
|
|
func slicesEqual[T comparable](a, b []T) bool {
|
||
|
|
if len(a) != len(b) {
|
||
|
|
return false
|
||
|
|
}
|
||
|
|
for i := range a {
|
||
|
|
if a[i] != b[i] {
|
||
|
|
return false
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return true
|
||
|
|
}
|