Skip to content

FIX set_adapter leaves the new adapter frozen after a merge under FSDP2 - #3895

Open
rupeshpoojary9 wants to merge 2 commits into
huggingface:mainfrom
rupeshpoojary9:fix/fsdp2-reshard-merged-submodule
Open

rupeshpoojary9 wants to merge 2 commits into
huggingface:mainfrom
rupeshpoojary9:fix/fsdp2-reshard-merged-submodule

Conversation

@rupeshpoojary9

Copy link
Copy Markdown
Contributor

Fixes #3872.

_reshard_fsdp_modules() (#3839) skips resharding an FSDP unit if any tuner layer inside it is merged, since resharding would drop the merge. set_adapter() called that reshard before base_model.set_adapter() got a chance to unmerge the layer, so on a model with a merged adapter, the reshard was skipped (correctly, nothing could be done yet), but then the new adapter's requires_grad=True landed on a transient unsharded copy instead of the FSDP2 shard. The next reshard (naturally triggered by the following forward/backward) drops that copy, leaving the new adapter frozen again.

@hivaze's repro shows it precisely: after merge_adapter() then set_adapter("other"), the first training step trains other through the stale unsharded copy, the second step gathers from the shard, where other is still frozen, and crashes with element 0 of tensors does not require grad and does not have a grad_fn.

Fix: unmerge first when anything is merged, then reshard, so the requires_grad change that follows always lands on the shard. base_model.set_adapter() already unmerges a merged layer on its own, this just makes sure that happens before the reshard attempt instead of after.

Verified against @hivaze's reproducer at both world_size=1 and world_size=2 (CPU/gloo): crashes on main, passes with this fix. Added a regression test mirroring the same scenario to TestFsdp2UnshardedParams (confirmed it fails without the fix, passes with it). All 13 existing tests in that class still pass. Ran the broader test_custom_models.py suite too; the only failures I saw (AdaMSS/FRoD-related, pre-existing) are identical with and without this change, confirmed by running the same tests on unpatched main.

Scoped this to set_adapter() only, matching what's actually been demonstrated broken. set_requires_grad() doesn't unmerge at all today (by design, it's a lower-level call), so it's unaffected by this change and would need separate analysis if that's also wanted.

_reshard_fsdp_modules() skips resharding an FSDP unit that still holds a
merge, since resharding would drop it. set_adapter() called that before
unmerging, so by the time base_model.set_adapter() unmerges the layer and
sets requires_grad on the new adapter, the reshard already ran (or was
skipped) and the change lands on a transient unsharded copy instead of
the FSDP2 shard. The next reshard drops it, leaving the new adapter
frozen again.

Unmerge first when anything is merged, then reshard, so the later
requires_grad change always lands on the shard.

Fixes huggingface#3872

@BenjaminBossan BenjaminBossan left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Thanks for the PR @rupeshpoojary9, I can confirm that the test fails without the fix and passes with it. I have a comment, please check.

@hivaze If you could also double check, that would be great.

Comment thread src/peft/peft_model.py
# change that follows would then land on a transient unsharded copy that the next reshard silently
# drops, leaving the newly active adapter frozen (#3872). base_model.set_adapter() below already
# unmerges a merged layer before setting it active, but only after the reshard attempt already ran.
self.base_model.unmerge_adapter()

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

So previously, set_adapter would call BaseTuner.set_adapter, which detects if any adapter is merged, warns, then unmerges said adapter. Now, we unmerge before going through BaseTuner.set_adapter, i.e. there is nothing to unmerge anymore. This means where previously, users would get a warning, there is no warning anymore. IMO we should still warn about this, the warning needs to be moved here. Please also add a test for the general case of set_adapter + merged modules which checks for the warning, there is no test for this so far.

Moving the unmerge ahead of the reshard meant base_model.set_adapter()
never saw a merged layer anymore, so the "Adapter cannot be set when the
model is merged" warning silently stopped firing. Emit it from the new
call site instead, and add a test for the general (non-FSDP2) case,
there wasn't one before.
@rupeshpoojary9

Copy link
Copy Markdown
Contributor Author

Good catch, pushed eb0007b: moved the warning to the new call site, and added a test for the general (non-FSDP2) case since there wasn't one. Confirmed the test fails without the warning (DID NOT WARN) and passes with it.

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

_reshard_fsdp_modules skips the whole FSDP unit if any adapter inside it is merged, not just the merged one

2 participants