Repository navigation
Conversation
validate() compared per-rank element counts, so a map can pass it while two of a rank's pieces write the same destination and another range is never written at all. Rebuild, restore and any direct transfer all read such a shard as if it were complete. Walk the destination addresses of the pieces that pack densely instead, which bounds the check by the piece count rather than the shard size, and leave a rank holding strided pieces alone, since intervals cannot describe where they land. Signed-off-by: 0z5a <dezhen.lu@student.uni-tuebingen.de>
Signed-off-by: 0z5a <dezhen.lu@student.uni-tuebingen.de>
5624904 to
8fb5e28
Compare
|
Confirmed against master, and the downstream consequence is sharper than "a tail that came from nowhere" — it is wrong values, not absent ones. I built your counterexample: validate() and validate_coverage() both accept it, extract returns [12, 13, 0, 0] (piece B overwrites A; the tail is never written), and rebuild returns [12, 13, 12, 13] from an original of [10, 11, 12, 13]. So a map like this round-trips a checkpoint to different values with every check green. Your sha256 for affine.py at 1eb56d2 Your reading of why validate_coverage does not catch it is right: source coverage is complete, and the collision is in the shard's address space, which it never looks at. The pigeonhole argument in _validate_destinations holds — intervals in bounds, pairwise disjoint, lengths summing to the shard size must tile it, so a dense rank needs no separate hole check. Stopping at strided pieces rather than guessing is the right call. On AssertionError -> ValueError: agreed, and it makes the module consistent rather than just safer. Every other guard in affine.py already raises ValueError; the assert was the outlier. One gap, and it is distinct from the fallback-policy question you scoped out. extract() never calls validate() at all — only rebuild() does. So on the restore path the guard does not fire, invalid or not: on your patched tree, extract still returns [12, 13, 0, 0] without complaint. That is not "what should a loader do when a map is invalid" but "the check never runs there." validate() is O(pieces) rather than O(elements), so extract could afford it. Happy to add that side in #8622 if you would rather keep this PR to validate(). Last thing, on reproducibility rather than correctness: with the tests taken back out in the second commit, the validation table cannot be checked from the diff. I get 53 pre-existing tests in test_affine_shard_map.py at base where you report 60 — the file is byte-identical to 1eb56d2 here and nothing is skipped, so it is probably pytest 9.1.0 versus 7.4.3, but it is worth reconciling since the 62 is the number that shows the guard is load-bearing. Committing them the way you did on #8623 would settle it. |
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
Singleton-axis layouts can bypass overlap detection, and the described regression tests are not committed.
Review effort: Balanced
Findings: 1
Open (3)
What changed in this PR
Adds destination-overlap validation for affine checkpoint shard maps.
Changes:
- Detects overlapping or out-of-bounds dense destination intervals.
- Replaces assertion-based size validation with
ValueError.
| File | Description |
|---|---|
deepspeed/checkpoint/affine.py |
Adds dense destination interval validation. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
|
Confirmed, and the reachability is worth stating precisely. A (1, 2) piece with dest strides (4, 1) occupies two adjacent addresses, but _dest_interval It is not reachable from the constructors. I scanned every piece replicated_map, Your suggested fix holds. Dropping size-1 axes before comparing catches the case, still reports |
Signed-off-by: 0z5a <dezhen.lu@student.uni-tuebingen.de>
|
Hi @0z5a, thank you for your PR, I have left my comments, can you take a look? Thanks! |
Signed-off-by: 0z5a <dezhen.lu@student.uni-tuebingen.de>


ParamAffineMap.validate()previously checked element counts without checking where dense pieces write. Two two-element pieces both writingdest[0:2]satisfy the count for a four-element shard while overwriting the head and leaving the tail unwritten.Check dense destination intervals for bounds and overlap using piece metadata. Ignore strides on singleton axes when identifying contiguous footprints, so
(1, 2)pieces with destination strides(4, 1)receive the same validation. Genuinely strided layouts retain their existing checks. Invalid shard element counts raiseValueErrorconsistently with other structural errors.The code docstrings describe the validation contract; the motivation and counterexample are kept here, as requested in review.
Validation:
tests/unit/checkpoint/test_affine_shard_map.py: 59 passed, using the real repository modules on macOS CPU, Python 3.12.13, PyTorch 2.14.1.10^9 x 8map.deepspeed/checkpoint/affine.pyandtests/unit/checkpoint/test_affine_shard_map.py; repeated for the revised docstrings.Full upstream CPU and formatting workflows require maintainer approval (
action_required, no jobs executed). The passing local checks do not establish a full CI pass.Refs #8230, #8252.