80 lines
3.1 KiB
Go
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)
|
|
}
|