Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

gradmesh

Distributed data-parallel training with the collectives built from scratch: ring all-reduce over raw TCP sockets, no torch.distributed, no MPI, no NCCL. Real processes, real packets, exact sums, and the two questions this repo answers with measurements instead of folklore: is distributed training THE SAME training (yes, provably), and when does it actually pay (three regimes, measured).

When data parallelism pays

The collective, from first principles

The ring all-reduce (gradmesh/ring.py) is the algorithm every GPU cluster runs, in about 120 readable lines: reduce-scatter (N-1 hops, each worker sends a chunk clockwise and adds the chunk arriving), then all-gather (N-1 more hops circulating the finished sums). Per-worker traffic is 2D(N-1)/N bytes, the bandwidth-optimal bound, and the test suite checks my implementation hits it to the byte: 6.0MB for a 4MB gradient at world=4, against 12.0MB for the naive circulate-everything baseline.

Ring all-reduce animated

(Animation: assets/manim_ring.py, video in assets/ring.mp4.)

Correctness is tested the only way that counts for distributed code: spawning 2 through 5 actual OS processes that talk TCP through the loopback, with gradient sizes chosen to not divide evenly, asserting exact sums on every rank (tests/test_ring.py). The whole suite runs in seven seconds.

The three bugs that taught the course

The bug ledger

Building this took one afternoon of writing and one day of debugging, which is the correct ratio for distributed systems and the honest heart of this README:

  1. The mutual sendall. Every worker sent its chunk, then received. Fine until chunks outgrew the kernel's socket buffers, at which point every sendall blocked with no one draining: the whole ring froze at exactly 100,003 elements and never at 1,000. Fix: send on a thread while the main loop receives, so the ring always drains.
  2. The socket that called itself. The nastiest one. My test happened to pick ports above 49152, macOS's ephemeral range. connect() to a neighbor that had not bound yet let the kernel choose OUR OWN port as the source, and TCP's simultaneous-connect rules connected the socket to itself. The worker sat "connected" to its own reflection while its real neighbor waited forever in accept(). I found it by process of elimination: stage logs showed both workers dying inside the constructor, a threads-only control experiment worked, and the only difference left was the port number. The library now rejects any connection where getsockname() == getpeername() and retries.
  3. The sum that would not repeat. At world=3 and 100K elements, 67 values disagreed with the reference by at most 4.8e-7. Not a bug: the ring adds in a different order, float32 addition is not associative, and the test now says so in a comment instead of demanding the impossible. PyTorch's DDP documentation makes the same disclaimer for the same reason.

Distributed training is the same training, provably

On top of the ring: a data-parallel trainer (gradmesh/trainer.py) for a hand-differentiated MLP, so the entire stack from loss gradient to socket write is inspectable. Every worker computes gradients on its shard of the global batch, the ring averages them, everyone steps identically.

The parity receipt

After 60 optimizer steps: all four workers hold bitwise-identical parameters (that is the all-reduce doing its job: identical inputs in, identical outputs everywhere), and the distributed model matches the single-process model to a max difference of 2.98e-08 across all 203,530 parameters, with 138,412 of them exactly equal. The residue is float32 reduction order and nothing else.

When it pays

The hero chart at the top is the finding. Same ring, same trainer, three regimes on four physical cores:

  • Shared BLAS (red): let NumPy's threaded BLAS keep all four cores in the single-process baseline, and data parallelism is strictly a loss at every world size, because extra workers bring no new silicon, only ring tolls and contention. First measurement: world=4 ran at 0.55x.
  • One core per worker, light compute (copper): pin BLAS to one thread each (the cluster-like setting) and a batch-512 step still barely gains: the gradient (3.2MB) is too large relative to the 15ms of compute that produced it, and on a laptop the "network" is the same CPU doing loopback memcpy.
  • One core per worker, heavy compute (green): quadruple the compute per step and the ring finally pays: 1.56x at three workers, declining after, because the sender threads themselves need somewhere to run.

That is the whole economics of data-parallel training in one picture: speedup is bought with the compute-to-communication ratio, comm costs real silicon, and the ideal line is a ceiling you approach only when each worker owns its hardware. It is also why the comm/compute split below is the chart infra teams actually watch:

The toll booth

Running it

pip install numpy                 # the only dependency
python tests/test_ring.py         # 2-5 real processes, exact sums, ~7s
python tests/test_parity.py       # distributed == single-process, ~2s
python benchmarks/scaling.py      # the pinned regime (GRADMESH_PIN=0 for shared,
                                  #  GRADMESH_BATCH=2048 for heavy)
python assets/make_visuals.py     # every chart from the logged results

Where this sits

Fifth in the from-scratch ML systems series, after ember, forge-lm, tinyserve-engine, and quantlab. The pattern here was the strongest yet: the algorithm took an afternoon, the distributed-systems reality took the rest, and every hour of it is in the bug ledger.

About

Distributed data-parallel training from raw TCP sockets: bandwidth-optimal ring all-reduce, a provably-identical-to-single-process trainer, the measured economics of when it pays, and the three bugs that taught the course

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages