Skip to content

test(distributed): cover the 8-rank read, and pin cross-mode byte parity - #42

Merged
jfischoff merged 1 commit into
mainfrom
test/distributed-read-world-8
Aug 17, 2026
Merged

jfischoff merged 1 commit into
mainfrom
test/distributed-read-world-8

Conversation

@jfischoff

Copy link
Copy Markdown
Contributor

What

Two things this suite could not say before:

  1. The 8-rank read works. Every distributed test here runs at _WORLD = 2, but the shape this path ships in is an 8-GPU node — and world size is not a free parameter of the sharded reader. It sets the shard plan (_shard_range / _plan_windows), how many ranks fall off the end of a short block with an empty shard, and how many collectives _replicate_contiguous issues per block. At world 2 a block only ever splits into "rank 0 takes it all" or "one shard each", so the interior ranks that exist only at larger worlds, and the geometric shard step-down in _plan_windows, were never exercised end to end. test_plan_windows_tiles_blocks_exactly covers worlds 2/4/8, but only as planning arithmetic — no read, no collectives.

  2. The read modes agree with each other. The existing sharded tests compare each mode against the source state. Nothing compared the modes against each other, which is the claim a caller actually leans on when it flips distributed_sharded.

Both new tests run on gloo/CPU like the rest of the module, so the extra ranks cost processes, not GPUs.

test_sharded_read_is_byte_exact_at_world_8

Per strategy, over a pack built to hit all three shard shapes at once:

block why
wide.bf16 (8 MiB) ≥ world * SHARD_ALIGN_BYTES — every rank owns a real shard
ragged.fp32 not a multiple of world * align — last shard runs long, windows must step its shard size down
short.fp32 (388 B) smaller than world * align — high ranks must skip it identically or the collective sequence desynchronizes and the job hangs

test_all_read_modes_agree_byte_for_byte

Reads one pack four ways in a single process — whole-pack, broadcast, sharded/contiguous, sharded/windows — and requires identical bytes on every rank.

Why now

A downstream app hit a CUDA illegal memory access on its first forward after switching to a FlashPack load. The sharded read was named as the cause and the caller reverted off it. It was not the read. On 8×H200 with a 28.58 GB production bf16 pack:

  • all four modes returned byte-identical payloads on all 8 ranks;
  • the model built from them matched the loader it replaced on all 1095 parameters, with zero tensors left on meta;
  • the fault still reproduced with distributed_sharded=False, with FlashPack disabled entirely, and finally on an unmodified checkout of the app's main branch — i.e. it was pre-existing and unrelated.

Nothing in this suite could have said that quickly, so clearing flashpack cost a multi-hour bisect on real hardware. These tests make it a two-second answer, and the README now carries the triage order (compare modes → diff the built model → only then look at the caller).

Test plan

Run on an 8×H200 node, torch 2.11, gloo/CPU ranks:

$ pytest tests/test_distributed_load.py -k 'world_8 or modes_agree'
3 passed

$ pytest tests/test_distributed_load.py
22 passed in 148.20s

No source changes — tests and README only.

Every distributed test in this module runs at `_WORLD = 2`, but the shape
this path actually ships in is an 8-GPU node -- and world size is not a free
parameter of the sharded reader. It sets the shard plan (`_shard_range` /
`_plan_windows`), how many ranks fall off the end of a short block with an
empty shard, and how many collectives `_replicate_contiguous` issues per
block. At world 2 a block only ever splits into "rank 0 takes it all" or
"one shard each", so the interior ranks that exist only at larger worlds,
and the geometric shard step-down in `_plan_windows`, were never exercised
end to end. `test_plan_windows_tiles_blocks_exactly` checks worlds 2/4/8,
but only as planning arithmetic -- no read, no collectives.

Two additions, both on gloo/CPU like the rest of the module (extra ranks
cost processes, not GPUs):

- `test_sharded_read_is_byte_exact_at_world_8`, per strategy, over a pack
  built to hit all three shard shapes at once: a block every rank shards, a
  ragged block whose remainder forces the windows step-down, and a block
  smaller than `world * align` that the high ranks must skip identically or
  the collective sequence desynchronizes.

- `test_all_read_modes_agree_byte_for_byte`, which reads one pack four ways
  in a single process -- whole-pack, broadcast, sharded/contiguous,
  sharded/windows -- and requires identical bytes. The existing sharded
  tests compare against the source state; this compares the modes against
  each other, which is the claim callers actually rely on when they flip
  `distributed_sharded`.

Motivated by a downstream investigation where a CUDA illegal-memory-access
in a model was attributed to the sharded read and the caller reverted off
it. It was not the read: on 8xH200 with a 28.58 GB production bf16 pack, all
four modes returned byte-identical payloads on all 8 ranks, the model built
from them matched the loader it replaced on all 1095 parameters with zero
meta tensors, and the fault still reproduced with `distributed_sharded=False`
and with flashpack removed entirely. Nothing in the suite could have said
that quickly, so the accusation cost a multi-hour hardware bisect. These
tests make it a two-second answer, and the README now carries the triage
order.

Validated on an 8xH200 node (gloo/CPU ranks, torch 2.11):
`pytest tests/test_distributed_load.py -k 'world_8 or modes_agree'` -> 3 passed.
@jfischoff
jfischoff merged commit 0db65e2 into main Aug 17, 2026
6 checks passed
@jfischoff
jfischoff deleted the test/distributed-read-world-8 branch August 17, 2026 17:31
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant