Skip to content

[checkpoint] Refuse a map whose pieces write the same shard address - #8624

Open
0z5a wants to merge 4 commits into
deepspeedai:masterfrom
0z5a:uc/v02-c1-map-validation
Open

0z5a wants to merge 4 commits into
deepspeedai:masterfrom
0z5a:uc/v02-c1-map-validation

Conversation

@0z5a

@0z5a 0z5a commented Sep 22, 2026 •

Copy link
Copy Markdown
Contributor

ParamAffineMap.validate() previously checked element counts without checking where dense pieces write. Two two-element pieces both writing dest[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 raise ValueError consistently 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.
  • Six committed regressions cover singleton-stride overlap and out-of-bounds writes, valid contiguous and strided layouts, element-count errors, and metadata-only validation of a 10^9 x 8 map.
  • All scoped pre-commit checks passed for deepspeed/checkpoint/affine.py and tests/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.

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>
@0z5a
0z5a force-pushed the uc/v02-c1-map-validation branch from 5624904 to 8fb5e28 Compare September 24, 2026 06:03
@Achyuthan-S

Copy link
Copy Markdown
Contributor

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
matches here, so we are looking at the same file.

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.

@0z5a
0z5a marked this pull request as ready for review September 29, 2026 08:54
@0z5a
0z5a requested a review from tjruwase as a code owner September 29, 2026 08:54
@delock
delock self-requested a review September 29, 2026 09:50
@hwchen2017
hwchen2017 requested a balanced review from Copilot September 29, 2026 17:13

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 High severity · 2 Low severity

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.

Comment thread deepspeed/checkpoint/affine.py Outdated
Comment thread deepspeed/checkpoint/affine.py Outdated
Comment thread deepspeed/checkpoint/affine.py
@Achyuthan-S

Copy link
Copy Markdown
Contributor

Confirmed, and the reachability is worth stating precisely.

A (1, 2) piece with dest strides (4, 1) occupies two adjacent addresses, but _dest_interval
rejects it because (4, 1) is not row-major for that shape, so it is skipped and never checked.
Two of them at offset 0: validate() and validate_coverage() both accept, extract returns
[12, 13, 0, 0], rebuild returns [12, 13, 12, 13] from [10, 11, 12, 13].

It is not reachable from the constructors. I scanned every piece replicated_map,
contiguous_split_map, sub_param_map, segmented_map and block_gather_map produce across
single-row, single-column and both partition dims at TP 2 and 4 — none has non-row-major dest
strides, so none is skipped. What makes it worth fixing anyway is from_dict: it takes strides
verbatim out of the file with no validation, and validate() is the guard for exactly that path.
A map this repository never writes is still a map it can be asked to read.

Your suggested fix holds. Dropping size-1 axes before comparing catches the case, still reports
not-dense for a genuinely strided piece like Yuan's o_proj (32, 1) so the check stays quiet
there, and agrees with the current rule on every constructor piece I could generate — so it only
starts checking pieces that were being skipped, and changes nothing that was already checked.

Signed-off-by: 0z5a <dezhen.lu@student.uni-tuebingen.de>
Comment thread deepspeed/checkpoint/affine.py Outdated
@delock

delock commented Oct 8, 2026

Copy link
Copy Markdown
Collaborator

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>

This branch has not been deployed

No deployments
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.

4 participants