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