feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,79 @@
|
||||
// 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)
|
||||
}
|
||||
Reference in New Issue
Block a user