Repository navigation
FIX set_adapter leaves the new adapter frozen after a merge under FSDP2 - #3895
rupeshpoojary9 wants to merge 2 commits into
Conversation
_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
left a comment
There was a problem hiding this comment.
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.
| # 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() |
There was a problem hiding this comment.
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.
|
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. |
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 beforebase_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'srequires_grad=Truelanded 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()thenset_adapter("other"), the first training step trainsotherthrough the stale unsharded copy, the second step gathers from the shard, whereotheris still frozen, and crashes withelement 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_gradchange 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=1andworld_size=2(CPU/gloo): crashes onmain, passes with this fix. Added a regression test mirroring the same scenario toTestFsdp2UnshardedParams(confirmed it fails without the fix, passes with it). All 13 existing tests in that class still pass. Ran the broadertest_custom_models.pysuite 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 unpatchedmain.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.