feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+576
@@ -0,0 +1,576 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package spmd
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
||||
)
|
||||
|
||||
// Default bounds of a networked world: one collective's wait, one
|
||||
// frame's payload, and the number of frames a link may queue. All
|
||||
// three exist so that a stuck or hostile peer is an error, never a
|
||||
// hang and never an allocation.
|
||||
const (
|
||||
defaultTimeout = 10 * time.Minute
|
||||
defaultMaxMessage = int64(16) << 30
|
||||
linkQueue = 64
|
||||
// pendingCap bounds how many unmatched frames one recvFrom may
|
||||
// hold: a peer flooding frames nobody asked for is an error,
|
||||
// never an unbounded allocator.
|
||||
pendingCap = 4096
|
||||
)
|
||||
|
||||
// Options bounds a networked world. An in-process world from Launch
|
||||
// takes none of them: its links are channels, and progress is the
|
||||
// program's own business, as it is in MPI.
|
||||
type Options struct {
|
||||
// Timeout bounds one collective's wait on the network: dialing,
|
||||
// the handshake and every send and receive carry it as a
|
||||
// deadline, refreshed each time a frame moves. Zero means the
|
||||
// default of ten minutes; a negative value means no deadline at
|
||||
// all.
|
||||
Timeout time.Duration
|
||||
// MaxMessage is the largest frame payload the world accepts, in
|
||||
// bytes. A peer announcing more is refused before any allocation.
|
||||
// Zero means the default of 16 GiB.
|
||||
MaxMessage int64
|
||||
}
|
||||
|
||||
func (o Options) timeout() time.Duration {
|
||||
switch {
|
||||
case o.Timeout > 0:
|
||||
return o.Timeout
|
||||
case o.Timeout < 0:
|
||||
return 0 // no deadline
|
||||
default:
|
||||
return defaultTimeout
|
||||
}
|
||||
}
|
||||
|
||||
func (o Options) maxMessage() int64 {
|
||||
if o.MaxMessage > 0 {
|
||||
return o.MaxMessage
|
||||
}
|
||||
return defaultMaxMessage
|
||||
}
|
||||
|
||||
// World is one rank's end of an SPMD world: its place in it, the links
|
||||
// to the other ranks and the collectives. A World is driven by one
|
||||
// goroutine: like an MPI rank, it never runs two collectives at once.
|
||||
type World struct {
|
||||
rank int
|
||||
size int
|
||||
timeout time.Duration
|
||||
maxMessage int64
|
||||
networked bool
|
||||
|
||||
peers []*peer // peers[r] is the link to rank r; nil for this rank
|
||||
closers []io.Closer
|
||||
|
||||
// The hub's readers and writers. The writers are waited on before
|
||||
// an orderly close, so the last collective's frames are on the
|
||||
// wire before the connections go away; the readers are waited on
|
||||
// after, once those connections have broken their blocking reads.
|
||||
drainWg sync.WaitGroup
|
||||
pumpWg sync.WaitGroup
|
||||
|
||||
done chan struct{} // closed when the world failed or left
|
||||
closeDone sync.Once
|
||||
failErr error
|
||||
}
|
||||
|
||||
// Rank returns this rank's index, from 0 to Size-1.
|
||||
func (w *World) Rank() int { return w.rank }
|
||||
|
||||
// Size returns the number of ranks in the world.
|
||||
func (w *World) Size() int { return w.size }
|
||||
|
||||
// Close ends a networked world: whatever the collectives queued is
|
||||
// written to the wire, then the connections close and the other ranks
|
||||
// see the departure as their next receive failing. On an in-process
|
||||
// world it is a no-op, because Launch tears the world down.
|
||||
func (w *World) Close() error {
|
||||
w.leave()
|
||||
w.drainWg.Wait()
|
||||
var first error
|
||||
for _, c := range w.closers {
|
||||
if err := c.Close(); err != nil && first == nil {
|
||||
first = err
|
||||
}
|
||||
}
|
||||
w.pumpWg.Wait()
|
||||
return first
|
||||
}
|
||||
|
||||
// leave closes done without recording a failure: the ordinary exit of
|
||||
// the rank's program.
|
||||
func (w *World) leave() {
|
||||
w.closeDone.Do(func() { close(w.done) })
|
||||
}
|
||||
|
||||
// fail records the world's first failure, closes done so that every
|
||||
// wait on this world wakes, and closes the networked links. It returns
|
||||
// the failure an outside caller should see.
|
||||
func (w *World) fail(err error) error {
|
||||
w.closeDone.Do(func() {
|
||||
w.failErr = err
|
||||
close(w.done)
|
||||
for _, c := range w.closers {
|
||||
c.Close()
|
||||
}
|
||||
})
|
||||
return w.status()
|
||||
}
|
||||
|
||||
// status is the entry check every public operation makes: a world that
|
||||
// failed, or a rank whose program already returned, answers with an
|
||||
// error and never with data.
|
||||
func (w *World) status() error {
|
||||
select {
|
||||
case <-w.done:
|
||||
if w.failErr != nil {
|
||||
return base.Errf("spmd: rank %d world is in a failed state: %v", w.rank, w.failErr)
|
||||
}
|
||||
return base.Errf("spmd: rank %d world has left", w.rank)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// deadline is the absolute time one collective may run to on a
|
||||
// networked world, refreshed each time a frame moves; the zero time
|
||||
// means no deadline, which is what an in-process world always gets.
|
||||
func (w *World) deadline() time.Time {
|
||||
if !w.networked || w.timeout <= 0 {
|
||||
return time.Time{}
|
||||
}
|
||||
return time.Now().Add(w.timeout)
|
||||
}
|
||||
|
||||
// depart closes the peer's gone channel: the peer closed its side of
|
||||
// the link, so no further frame will ever come from it.
|
||||
func (p *peer) depart() {
|
||||
select {
|
||||
case <-p.gone:
|
||||
default:
|
||||
close(p.gone)
|
||||
}
|
||||
}
|
||||
|
||||
// hasLeft reports whether the peer closed its side of the link.
|
||||
func (p *peer) hasLeft() bool {
|
||||
if p.gone == nil {
|
||||
return false
|
||||
}
|
||||
select {
|
||||
case <-p.gone:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// sendTo carries one frame to rank r. On an in-process world it walks
|
||||
// the direct channel; on a networked world it walks this rank's own
|
||||
// link to the hub, because rank 0 routes every frame by its
|
||||
// destination.
|
||||
func (w *World) sendTo(r int, tag uint8, payload []byte) error {
|
||||
if err := w.status(); err != nil {
|
||||
return err
|
||||
}
|
||||
m := message{tag: tag, from: w.rank, dest: r, data: payload}
|
||||
p := w.peers[r]
|
||||
if p != nil && p.outbox != nil {
|
||||
select {
|
||||
case p.outbox <- m:
|
||||
// On a buffered networked link a queued frame is not yet
|
||||
// a delivered frame, so a peer that left between the two
|
||||
// dropped it. An in-process handoff is the delivery
|
||||
// itself: the receiver taking the frame and then leaving
|
||||
// is its own healthy business.
|
||||
if w.networked && p.hasLeft() {
|
||||
return w.fail(base.Errf("spmd: rank %d sending to rank %d: rank %d left the world", w.rank, r, r))
|
||||
}
|
||||
return nil
|
||||
case <-p.gone:
|
||||
return w.fail(base.Errf("spmd: rank %d sending to rank %d: rank %d left the world", w.rank, r, r))
|
||||
case <-w.done:
|
||||
return w.status()
|
||||
}
|
||||
}
|
||||
// A networked rank that is not the hub owns one connection, and
|
||||
// every word it sends rides it; the hub reads the destination.
|
||||
if err := w.peers[0].conn.writeFrame(m, w.deadline()); err != nil {
|
||||
return w.fail(base.Errf("spmd: rank %d sending to rank %d: %v", w.rank, r, err))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// recvFrom returns the payload of the next frame rank r sent under the
|
||||
// wanted tag, holding frames that arrived earlier for other ranks or
|
||||
// tags until the collective asks for them: no answer ever depends on
|
||||
// the arrival order. Any failure fails the world.
|
||||
func (w *World) recvFrom(r int, tag uint8) ([]byte, error) {
|
||||
if err := w.status(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
p := w.peers[r]
|
||||
if w.networked && w.rank != 0 {
|
||||
// One stream carries every rank's words to a rank that is not
|
||||
// the hub, so one pending stash serves them all.
|
||||
p = w.peers[0]
|
||||
}
|
||||
for i, m := range p.pending {
|
||||
if m.from == r && m.tag == tag {
|
||||
p.pending = append(p.pending[:i], p.pending[i+1:]...)
|
||||
return m.data, nil
|
||||
}
|
||||
}
|
||||
for {
|
||||
var m message
|
||||
var err error
|
||||
switch {
|
||||
case p.inbox != nil:
|
||||
select {
|
||||
case m = <-p.inbox:
|
||||
case <-p.gone:
|
||||
// The peer left: whatever it queued before leaving is
|
||||
// still in the buffer and still counts; only an empty
|
||||
// buffer means the frames will never come.
|
||||
for {
|
||||
select {
|
||||
case m = <-p.inbox:
|
||||
if m.from == r && m.tag == tag {
|
||||
return m.data, nil
|
||||
}
|
||||
p.pending = append(p.pending, m)
|
||||
continue
|
||||
default:
|
||||
}
|
||||
return nil, w.fail(base.Errf("spmd: rank %d receiving from rank %d: rank %d left the world", w.rank, r, r))
|
||||
}
|
||||
case <-w.done:
|
||||
return nil, w.status()
|
||||
}
|
||||
default:
|
||||
// This rank drives its single connection; the hub has
|
||||
// already routed whatever was not for it.
|
||||
m, err = p.conn.readFrame(w.maxMessage, w.deadline())
|
||||
if err != nil {
|
||||
return nil, w.fail(base.Errf("spmd: rank %d receiving from rank %d: %v", w.rank, r, w.readErr(err)))
|
||||
}
|
||||
if m.dest != w.rank {
|
||||
return nil, w.fail(base.Errf("spmd: rank %d got a frame addressed to rank %d", w.rank, m.dest))
|
||||
}
|
||||
}
|
||||
if m.from == r && m.tag == tag {
|
||||
return m.data, nil
|
||||
}
|
||||
p.pending = append(p.pending, m)
|
||||
if len(p.pending) >= pendingCap {
|
||||
return nil, w.fail(base.Errf("spmd: rank %d holds %d unmatched frames against rank %d", w.rank, len(p.pending), r))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// readErr names a networked read failure for what it is: a peer that
|
||||
// closed or dropped its connection.
|
||||
func (w *World) readErr(err error) error {
|
||||
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
return errors.New("the peer closed its connection")
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// pump reads one networked hub link for the world's lifetime: frames
|
||||
// for rank 0 join the link's inbox, frames for anybody else join that
|
||||
// rank's outbox unchanged. The hub is a post office, never an
|
||||
// interpreter: what a frame carries is the collectives' business. A
|
||||
// frame bound for another rank reads into a buffer from the routed
|
||||
// frame pool, whose ownership rides the message through the outbox
|
||||
// channel to the link's one drain. A peer that closes its connection
|
||||
// has left the world, which is its own business too; only a broken or
|
||||
// unreadable link fails this world.
|
||||
func (w *World) pump(p *peer) {
|
||||
for {
|
||||
m, err := p.conn.readFrameInto(w.maxMessage, w.deadline(), takeFrameBuffer)
|
||||
if err != nil {
|
||||
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) || errors.Is(err, net.ErrClosed) {
|
||||
p.depart()
|
||||
return
|
||||
}
|
||||
w.fail(base.Errf("spmd: rank 0 reading from rank %d: %v", p.rank, w.readErr(err)))
|
||||
return
|
||||
}
|
||||
if m.from != p.rank || m.dest < 0 || m.dest >= w.size {
|
||||
w.fail(base.Errf("spmd: rank 0 got a mislabelled frame from rank %d", p.rank))
|
||||
return
|
||||
}
|
||||
var out chan message
|
||||
if m.dest == 0 {
|
||||
out = p.inbox
|
||||
} else {
|
||||
out = w.peers[m.dest].outbox
|
||||
}
|
||||
select {
|
||||
case out <- m:
|
||||
case <-w.done:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// drain writes one networked hub link for the world's lifetime: it is
|
||||
// the only goroutine that ever writes the connection, so the frames
|
||||
// the collective logic and the routed traffic send share one ordered
|
||||
// stream without a lock. It is also the single consumer of the
|
||||
// routed frames the pump queues, which makes it the one place a
|
||||
// pooled payload is returned: writeFrame is the payload's last
|
||||
// reader, so the buffer goes back the moment the write returns,
|
||||
// whatever the answer was. When the world ends it writes out whatever
|
||||
// the last collectives queued before it leaves, so an orderly close
|
||||
// never drops a frame that was sent. A write against a link whose peer
|
||||
// already left is not failed, because the departure was the peer's own
|
||||
// clean act; any other write error is, named at once rather than left
|
||||
// to surface later as a hang.
|
||||
func (w *World) drain(p *peer) {
|
||||
write := func(m message) bool {
|
||||
err := p.conn.writeFrame(m, w.deadline())
|
||||
if m.pooled {
|
||||
routedFrames.retire(m.data)
|
||||
}
|
||||
if err == nil {
|
||||
return true
|
||||
}
|
||||
if p.hasLeft() || w.ended() {
|
||||
return false
|
||||
}
|
||||
w.fail(base.Errf("spmd: rank 0 writing to rank %d: %v", p.rank, err))
|
||||
return false
|
||||
}
|
||||
for {
|
||||
select {
|
||||
case m := <-p.outbox:
|
||||
if !write(m) {
|
||||
return
|
||||
}
|
||||
case <-w.done:
|
||||
for {
|
||||
select {
|
||||
case m := <-p.outbox:
|
||||
if !write(m) {
|
||||
return
|
||||
}
|
||||
default:
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ended reports whether the world's done channel has closed.
|
||||
func (w *World) ended() bool {
|
||||
select {
|
||||
case <-w.done:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// startPumps launches the hub's readers and writers, one of each per
|
||||
// link. They live as long as the world does.
|
||||
func (w *World) startPumps() {
|
||||
for r := 1; r < w.size; r++ {
|
||||
p := w.peers[r]
|
||||
w.pumpWg.Go(func() { w.pump(p) })
|
||||
w.drainWg.Go(func() { w.drain(p) })
|
||||
}
|
||||
}
|
||||
|
||||
// Launch runs the same function on size ranks of one process, one
|
||||
// goroutine per rank, over in-process links: the same collectives, the
|
||||
// same answers and the same rules as a networked world, which makes it
|
||||
// the development and test surface of the package. The errors of the
|
||||
// ranks that failed come back joined in rank order, so the report is
|
||||
// deterministic and the rank that caused the trouble is in it; a rank
|
||||
// whose function panics fails its world, and the panic comes back as
|
||||
// that rank's error.
|
||||
func Launch(size int, fn func(w *World) error) error {
|
||||
if size < 1 {
|
||||
return base.Errf("spmd: a world needs at least one rank, got %d", size)
|
||||
}
|
||||
worlds := make([]*World, size)
|
||||
for r := range worlds {
|
||||
worlds[r] = &World{
|
||||
rank: r,
|
||||
size: size,
|
||||
done: make(chan struct{}),
|
||||
peers: make([]*peer, size),
|
||||
}
|
||||
}
|
||||
// One channel per direction of every pair: the sender's outbox is
|
||||
// the receiver's inbox, so a frame crosses without a middleman and
|
||||
// a receive wakes the moment its rank's world leaves.
|
||||
for r := range worlds {
|
||||
for q := r + 1; q < size; q++ {
|
||||
rToQ := make(chan message, linkQueue)
|
||||
qToR := make(chan message, linkQueue)
|
||||
worlds[r].peers[q] = &peer{rank: q, inbox: qToR, outbox: rToQ, gone: worlds[q].done}
|
||||
worlds[q].peers[r] = &peer{rank: r, inbox: rToQ, outbox: qToR, gone: worlds[r].done}
|
||||
}
|
||||
}
|
||||
errs := make([]error, size)
|
||||
var wg sync.WaitGroup
|
||||
for r := range worlds {
|
||||
wg.Go(func() {
|
||||
defer func() {
|
||||
if p := recover(); p != nil {
|
||||
errs[r] = base.Errf("spmd: rank %d panicked: %v", r, p)
|
||||
worlds[r].fail(errs[r])
|
||||
}
|
||||
worlds[r].leave()
|
||||
}()
|
||||
errs[r] = fn(worlds[r])
|
||||
})
|
||||
}
|
||||
wg.Wait()
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// Listen builds the rank 0 end of a networked world: it listens on the
|
||||
// address until every other rank has joined, assigning ranks in dial
|
||||
// order. Connections that do not say the spmd handshake are closed
|
||||
// and skipped, so they cannot take a rank's place.
|
||||
func Listen(addr string, size int, opts Options) (*World, error) {
|
||||
ln, err := net.Listen("tcp", addr)
|
||||
if err != nil {
|
||||
return nil, base.Errf("spmd: listening on %s: %v", addr, err)
|
||||
}
|
||||
return listen(ln, size, opts)
|
||||
}
|
||||
|
||||
// listen assembles the rank 0 world on a ready listener; the tests use
|
||||
// it to hand over a listener whose address they already know.
|
||||
func listen(ln net.Listener, size int, opts Options) (*World, error) {
|
||||
if size < 1 {
|
||||
ln.Close()
|
||||
return nil, base.Errf("spmd: a world needs at least one rank, got %d", size)
|
||||
}
|
||||
timeout := opts.timeout()
|
||||
w := &World{
|
||||
rank: 0,
|
||||
size: size,
|
||||
timeout: timeout,
|
||||
maxMessage: opts.maxMessage(),
|
||||
networked: true,
|
||||
done: make(chan struct{}),
|
||||
peers: make([]*peer, size),
|
||||
}
|
||||
w.closers = append(w.closers, ln)
|
||||
if tcp, ok := ln.(*net.TCPListener); ok && timeout > 0 {
|
||||
tcp.SetDeadline(time.Now().Add(timeout))
|
||||
}
|
||||
deadline := w.deadline()
|
||||
for joined := 1; joined < size; joined++ {
|
||||
conn, err := ln.Accept()
|
||||
if err != nil {
|
||||
return nil, w.fail(base.Errf("spmd: rank 0 accepting rank %d on %s: %v", joined, ln.Addr(), err))
|
||||
}
|
||||
if err := readHello(conn, deadline); err != nil {
|
||||
// A connection that does not say the handshake is not one
|
||||
// of ours; it takes no rank's place.
|
||||
conn.Close()
|
||||
joined--
|
||||
continue
|
||||
}
|
||||
if err := sendWelcome(conn, size, joined); err != nil {
|
||||
conn.Close()
|
||||
return nil, w.fail(base.Errf("spmd: rank 0 welcoming rank %d: %v", joined, err))
|
||||
}
|
||||
w.peers[joined] = &peer{
|
||||
rank: joined,
|
||||
conn: newTCPLink(conn),
|
||||
inbox: make(chan message, linkQueue),
|
||||
outbox: make(chan message, linkQueue),
|
||||
gone: make(chan struct{}),
|
||||
}
|
||||
w.closers = append(w.closers, conn)
|
||||
}
|
||||
// Every rank has its link; nobody else joins this world. The
|
||||
// listener was closers[0], and its job is done.
|
||||
ln.Close()
|
||||
w.closers = w.closers[1:]
|
||||
w.startPumps()
|
||||
return w, nil
|
||||
}
|
||||
|
||||
// Join builds the other ranks' end of a networked world: it dials the
|
||||
// listening rank 0, which answers with the world's size and this
|
||||
// connection's rank. The world is a star over rank 0's listener, so
|
||||
// one address is the whole world's knowledge; rank 0 routes every
|
||||
// frame to its destination.
|
||||
func Join(addr string, opts Options) (*World, error) {
|
||||
timeout := opts.timeout()
|
||||
d := net.Dialer{Timeout: timeout}
|
||||
conn, err := d.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
return nil, base.Errf("spmd: rank dialling %s: %v", addr, err)
|
||||
}
|
||||
w := &World{
|
||||
timeout: timeout,
|
||||
maxMessage: opts.maxMessage(),
|
||||
networked: true,
|
||||
done: make(chan struct{}),
|
||||
closers: []io.Closer{conn},
|
||||
}
|
||||
if err := sendHello(conn); err != nil {
|
||||
return nil, w.fail(base.Errf("spmd: rank saying hello to %s: %v", addr, err))
|
||||
}
|
||||
size, rank, err := readWelcome(conn, w.deadline())
|
||||
if err != nil {
|
||||
return nil, w.fail(base.Errf("spmd: rank reading %s's welcome: %v", addr, err))
|
||||
}
|
||||
if rank == 0 {
|
||||
return nil, w.fail(base.Errf("spmd: the listener at %s answered a joining connection with rank 0", addr))
|
||||
}
|
||||
w.rank, w.size = rank, size
|
||||
w.peers = make([]*peer, size)
|
||||
w.peers[0] = &peer{rank: 0, conn: newTCPLink(conn)}
|
||||
return w, nil
|
||||
}
|
||||
|
||||
// Barrier blocks until every rank of the world has reached it. It
|
||||
// carries no data and no arithmetic, so there is nothing in it to be
|
||||
// anything but deterministic.
|
||||
func (w *World) Barrier() error {
|
||||
if err := w.status(); err != nil {
|
||||
return err
|
||||
}
|
||||
if w.rank == 0 {
|
||||
for r := 1; r < w.size; r++ {
|
||||
if _, err := w.recvFrom(r, tagBarrier); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for r := 1; r < w.size; r++ {
|
||||
if err := w.sendTo(r, tagBarrierAck, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err := w.sendTo(0, tagBarrier, nil); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := w.recvFrom(0, tagBarrierAck)
|
||||
return err
|
||||
}
|
||||
Reference in New Issue
Block a user