SIGN IN SIGN UP

[FlexCheckpoint] Fix nested state_dict load overwritten by loop variable (#79707)

* [FlexCheckpoint] Fix nested state_dict load overwritten by loop variable

`load_state_dict_impl` flattens the caller's `state_dict`, loads into the
flat view, then walks `mapping` to write the tensors back into the caller's
nested layout:

    tmp = state_dict
    for key in keys[:-1]:
        tmp = tmp[key]
    tmp[keys[-1]] = flat_state_dict[flat_key]

Two loops in between reuse the name `state_dict` for a per-file tensor dict.
Python leaks loop variables into the enclosing scope, so by the time the walk
runs, `state_dict` refers to the last checkpoint file's contents instead of
the caller's dict. Its keys are flat physical tensor names, so descending
into a nested target raises `KeyError`:

    dist.load_state_dict(
        {"layer": {"weight": w}}, path,
        aoa_config={"aoa_statements": ["layer.weight -> layer.weight"]},
    )
    # KeyError: 'layer'

A flat target does not descend and so does not raise, but the write lands in
the file dict that is about to be freed, leaving the restore a no-op.

The path is reached whenever the local-resume fast path is skipped -- that is
with `safetensors=True`, an `aoa_config`, or a resharding load -- and only on
ranks that actually read a file, since an empty `source_state_dict` leaves
the loop body unexecuted. That makes the failure rank dependent and hard to
diagnose.

Rename both loop variables to `file_tensors` and add regression tests for
nested, deeply nested, and flat targets.

* [FlexCheckpoint] Cover replica eviction in nested state_dict load test

Add test_loads_checkpoint_with_replica_entries, which rewrites the saved
storage metadata so one shard is recorded twice (replica_id 0 and 1) while a
second shard is homed only in the duplicated file. The rank then has to read
both files, so the load actually walks the replica eviction branch that the
previous tests never reached.
L
Liumengyuan committed
d5070ae1caecd3e2db65799023b664fb0468cae0
Parent: bfa2247
Committed by GitHub <noreply@github.com> on 8/31/2026, 2:20:10 AM