test(distributed): cover the 8-rank read, and pin cross-mode byte parity - #42
Merged
Merged
Conversation
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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
What
Two things this suite could not say before:
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_contiguousissues 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_exactlycovers worlds 2/4/8, but only as planning arithmetic — no read, no collectives.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_8Per strategy, over a pack built to hit all three shard shapes at once:
wide.bf16(8 MiB)world * SHARD_ALIGN_BYTES— every rank owns a real shardragged.fp32world * align— last shard runs long,windowsmust step its shard size downshort.fp32(388 B)world * align— high ranks must skip it identically or the collective sequence desynchronizes and the job hangstest_all_read_modes_agree_byte_for_byteReads 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 accesson 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:meta;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:
No source changes — tests and README only.