Files
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

80 lines
3.1 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package spmd
import (
"sourcedock.dev/petrbalvin/tensor/internal/base"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Span names one rank's contiguous piece of a global axis of data: the
// global length, the piece's first element and one past its last.
type Span struct {
// Global is the axis length every rank's piece is a cut of.
Global int
// Lo is the piece's first global element index; Hi is one past its
// last. The piece is Lo..Hi, possibly empty.
Lo, Hi int
}
// Len returns the number of elements the piece carries along the axis.
func (s Span) Len() int { return s.Hi - s.Lo }
// Partition cuts a global axis of globalN elements into the world's
// contiguous pieces: rank rank's piece is [Lo, Hi). Every piece starts
// and ends on a boundary of the canonical fold partition, which is what
// lets the shards' reductions compose into the single-array fold's
// exact bits, so a program that cuts its data any other way gives up
// that contract. With more ranks than blocks, the pieces beyond the blocks are empty,
// wherever the boundary falls, so rank 0 may hold nothing at all.
//
// A negative globalN, a size below 1 or a rank outside [0, size) is
// refused with an error naming the input. Valid arguments never fail.
func Partition(globalN, size, rank int) (Span, error) {
if globalN < 0 {
return Span{}, base.Errf("spmd: Partition of a negative global length %d", globalN)
}
if size < 1 || rank < 0 || rank >= size {
return Span{}, base.Errf("spmd: Partition of rank %d in a world of %d", rank, size)
}
parts := core.FoldParts(globalN)
return Span{
Global: globalN,
Lo: core.FoldBoundary(globalN, rank*parts/size),
Hi: core.FoldBoundary(globalN, (rank+1)*parts/size),
}, nil
}
// checkSpan verifies that a caller's span is the canonical partition's
// own cut for this rank and that the local slab leads with exactly the
// span's run of the axis. Any other cut is refused by name: the
// bit-identity contract lives on the canonical boundaries.
func (w *World) checkSpan(span Span, local *core.Array) error {
want, err := Partition(span.Global, w.size, w.rank)
if err != nil {
return err
}
if span != want {
return base.Errf("spmd: rank %d holds [%d, %d) of %d, but the canonical partition puts this rank on [%d, %d); cut the data with Partition",
w.rank, span.Lo, span.Hi, span.Global, want.Lo, want.Hi)
}
if local.NDim() == 0 || local.Shape()[0] != span.Len() {
lead := 0
if local.NDim() > 0 {
lead = local.Shape()[0]
}
return base.Errf("spmd: rank %d's slab leads with %d elements for a span of %d", w.rank, lead, span.Len())
}
return nil
}
// checkOneDimSpan is checkSpan for the vector reductions: the shards
// carry one-dimensional arrays, whose length is the span's run.
func (w *World) checkOneDimSpan(local *core.Array, span Span) error {
if local.NDim() != 1 {
return base.Errf("spmd: the sharded product, norm and dot carry 1-D arrays; got %d dimensions", local.NDim())
}
return w.checkSpan(span, local)
}