Skip to content

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

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

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

Conversation

@0z5a

@0z5a 0z5a commented Sep 22, 2026 •

Copy link
Copy Markdown
Contributor

Incremental review: Files changed against master.

Summary

  • ParamAffineMap.validate() compared element counts, which cannot see where those elements land. A map can pass it while two of one rank's pieces write the same shard address and a third range is never written at all.
  • Walk the destination addresses of the pieces that pack densely instead. Bounded by the piece count, not the shard size, and silent about ranks whose pieces stride -- an interval cannot describe those.
  • The guard is in affine.py, with committed regressions in test_affine_shard_map.py.

The counterexample

Four elements [10, 11, 12, 13] into a (4,) shard, with two pieces both writing the front:

piece A: source[0:2] -> dest[0:2]
piece B: source[2:4] -> dest[0:2]

source coverage      [0, 1, 2, 3]      complete
piece numels         2 + 2 = 4         equals the shard
destination writes   [2, 2, 0, 0]      two addresses twice, two never

Run against the real module at 1eb56d29 (deepspeed/checkpoint/affine.py, sha256 21701805c23abc303028e95049f1a423fad4d50fe27ac3f43f9356229d218892), both validate() and validate_coverage() accept this map. validate_coverage is the thorough one, and it is not enough either: it checks that every element of the parameter is covered by someone and that each piece's holders are honest. Neither question is about the shard's own address space, which is what a reader or a writer actually indexes.

Downstream that means rebuild returns a tensor whose tail came from nowhere, and restore or a direct transfer reads a shard as though it were complete.

What the guard does

For each rank, every piece whose destination strides are the row-major strides of its own shape occupies one contiguous run of the shard. Those runs must lie inside the shard and must not overlap.

It stops there deliberately. Where some piece of a rank strides through the shard -- a column split, Yuan's o_proj -- a gap between the runs is not yet a gap in the shard, so the check says nothing rather than guessing. And where every piece is dense, no separate hole check is needed: runs that are in bounds, non-overlapping and sum to the shard size cannot leave a gap. That is why there is no hole assertion here.

Validation

test_affine_shard_map.py at this PR's head with the guard
pre-existing tests 60 passed 62 passed
test_double_written_shard_is_refused fails (map accepted) passes
test_destination_outside_the_shard_is_refused fails (map accepted) passes
6 builders asserted still valid (replicated, uneven row, column, sub-params, GPTBigCode segments, Yuan gather) n/a pass
test_validate_never_walks_elements (per-element helpers bound to raisers, then a 10^9 x 8 map validated) n/a passes

The two refusals failing at base is the point: the same commit that adds them fails without the change, so the guard is what makes them pass. Run on CPU (Python 3.12.13, torch 2.13.0+cu130). The wider tests/unit/checkpoint directory was run on a pristine 1eb56d29 worktree and on this branch; failing and erroring test names are identical, 241 either way, none introduced. Two environmental causes on both trees: multi-rank cases cannot run under CUDA_VISIBLE_DEVICES="", and this environment ships pytest 9.1.0 while dev requirements pin <8.4.0, so the repository's distributed harness never registers its fixtures.

One behavior change to be aware of

validate() raised AssertionError on a numel mismatch. It now raises ValueError, because a map that overwrites part of its own shard is not a condition worth tying to an interpreter flag that python -O clears. Every call site is inside the checkpoint path, where a silent pass is the failure mode being fixed. No test in the repository relied on the assertion type.

Scope

  • No general overlap decision for arbitrary strided views. That needs interval algebra over each rank's address set, and the layouts in tree do not need it yet.
  • Not a validation framework, and not a dependency for anything else: affine_transfer.py in [UC] Plan a scale-1 affine shard-to-shard transfer without rebuilding the parameter #8623 proves coverage over boxes on its own.
  • No change to the restore path's fallback policy -- what a loader should do when a map is present but invalid is a separate question with its own call sites, and mixing it in here would hide which of the two a failure came from.
  • No GPU run: the guard is pure Python over integers.

Refs #8230, #8252

Follow-up at 1a7ddbe: six additional committed regression cases pass on the existing 0z5a Python environment. The full file reports 51 passed and 8 failures in Yuan/BigCode integration cases; the same eight cases fail at the prior PR head (45 passed, six new tests deselected), due to the installed DeepSpeed module-injection API differing from this PR source. No environment packages were changed.

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 on lines +153 to +155
if piece.numel == 0 or piece.dest_strides != _row_major_strides(piece.shape):
return None
return piece.dest_offset, piece.dest_offset + piece.numel
Comment thread deepspeed/checkpoint/affine.py Outdated
Comment on lines +198 to +202
Counting elements is not the same as placing them. Two pieces can both write the head of
a shard and leave its tail unwritten while summing to the shard size exactly, so the
totals agree and nothing here reports a problem. This also walks the destination addresses
of every piece that packs densely, which bounds the check by the piece count rather than
the shard size.
Comment on lines +207 to +210
if held != expected:
raise ValueError(f'Rank {rank} holds a shard of {expected} elements but its pieces '
f'account for {held}.')
self._validate_destinations(rank, pieces, expected)
@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>

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.

3 participants