Repository navigation
_reshard_fsdp_modules skips the whole FSDP unit if any adapter inside it is merged, not just the merged one #3872
Description
Activity
Thanks, that's right, and it's on purpose. FSDP2 reshards a whole module at once, so
v_projcan't be resharded without dropping the merge inq_proj, and the later unmerge would then subtract the delta from the base weights. The docs say it ("FSDP2 modules that shard parameters of a merged adapter layer are not resharded, so unmerge before calling these methods"), but they could say more plainly that it covers every adapter layer in that module.The same skip causes a worse case in
set_adapter()on a merged model. PEFT unmerges there for you, but only after the reshard was skipped, so the newrequires_gradlands on the unsharded copies. Aftermerge_adapter()andset_adapter("other"), the shards still havedefaulttrainable andotherfrozen. So the first step trainsother, and with only the adapters trainable the second raiseselement 0 of tensors does not require grad. Callingunmerge_adapter()beforeset_adapter()avoids it, so unmerging before the reshard inPeftModel.set_adapter()should fix it. I checked both on one and two CPU ranks.We don't merge adapters under FSDP2 ourselves, so we won't send a PR for this one.
Appreciate the deeper read, that's a more serious failure mode than the one I filed. Looked at where the unmerge actually happens: it's in the module-level
set_adapter()helper intuners_utils.py(if module.merged: warnings.warn(...); module.unmerge()), called fromLoraModel.set_adapter→PeftModel.set_adapter, which runs afterself._reshard_fsdp_modules()today. So by the time any layer unmerges, the reshard attempt already happened (and was skipped, since the layer was still merged at that point). Moving the unmerge ahead of the reshard call inPeftModel.set_adapter(), and the same indisable_adapter()'s exit path sinceenable_adapter_layers()goes through the same helper, should close it. I'd guard it behind "is anything actually merged" first, sinceunmerge_adapter()warns "Already unmerged, nothing to do" per layer otherwise, that'd be a UX regression on every normal (non-merged) call.I tried reproducing the crash itself (merge_adapter → set_adapter("other") → two training steps with an explicit reshard between them) on CPU/gloo, both 1 and 2 ranks, and couldn't trigger it, the params never diverged from the shard in my setup. Would you be willing to share the exact script/sequence you used? Want to validate the fix against the actual failure before sending a PR rather than guess at the repro.
Sure, script below. It needs two things: a forward before
merge_adapter(), so the merge lands on the unsharded copies, andreshard_after_forward=False. Without that first forward, or withreshard_after_forward=True, all three steps trainother, so your run may have missed one of them.repro.py
"""set_adapter() on a merged model under FSDP2 (reshard_after_forward=False): which adapter gets gradients per step.""" import warnings import torch import torch.distributed as dist from torch import nn from torch.distributed.device_mesh import init_device_mesh from torch.distributed.fsdp import fully_shard from torch.distributed.tensor import DTensor from peft import LoraConfig, get_peft_model warnings.simplefilter("ignore") dist.init_process_group("gloo") rank, world = dist.get_rank(), dist.get_world_size() mesh = init_device_mesh("cpu", (world,)) class MLP(nn.Module): def __init__(self): super().__init__() self.lin0 = nn.Linear(10, 20) self.relu = nn.ReLU() self.lin1 = nn.Linear(20, 2) def forward(self, X): return self.lin1(self.relu(self.lin0(X))) def run(unmerge_first): torch.manual_seed(0) cfg = LoraConfig(target_modules=["lin0", "lin1"], init_lora_weights=False) model = get_peft_model(MLP(), cfg) model.add_adapter("other", cfg) fully_shard(model.base_model.model.lin0, mesh=mesh, reshard_after_forward=False) fully_shard(model, mesh=mesh, reshard_after_forward=False) shards = {n: p for n, p in model.named_parameters() if ".lora_" in n} assert all(isinstance(p, DTensor) for p in shards.values()) X = torch.arange(90).view(9, 10).float() with torch.no_grad(): model(X) model.merge_adapter() if unmerge_first: model.unmerge_adapter() model.set_adapter("other") out = [] for step in (1, 2, 3): for p in shards.values(): p.grad = None try: model(X).sum().backward() except RuntimeError as e: out.append(f"step {step}: RuntimeError: {e}") break with_grad = sorted({n.split(".lora_")[1].split(".")[1] for n, p in shards.items() if p.grad is not None}) out.append(f"step {step}: grads on {with_grad}") if rank == 0: print(f"[world={world}] unmerge before set_adapter={unmerge_first}: " + "; ".join(out), flush=True) run(unmerge_first=False) run(unmerge_first=True) dist.destroy_process_group()
On main (532a05d) with torch 2.11,
torchrun --nproc_per_node=1and then--nproc_per_node=2print:[world=1] unmerge before set_adapter=False: step 1: grads on ['other']; step 2: RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn [world=1] unmerge before set_adapter=True: step 1: grads on ['other']; step 2: grads on ['other']; step 3: grads on ['other'] [world=2] unmerge before set_adapter=False: step 1: grads on ['other']; step 2: RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn [world=2] unmerge before set_adapter=True: step 1: grads on ['other']; step 2: grads on ['other']; step 3: grads on ['other']Step 1 trains
otherthrough the unsharded copies, its backward reshards, and step 2 gathers from shards whereotheris still frozen.If the fix also touches the
disable_adapter()exit, note that #3800 has a proposed change there.Thanks the two of you for bringing attention to this.
I agree that the documentation could put more stress on the fact that any submodule being merged skips the resharding. Note, though, that finding one sub-module merged and another unmerged is not a very common occurrence.
About this fix: Sounds good. Please include tests along the line of the provided reproducer.
/peft-triage approved
PR up: #3895. Verified against your reproducer at world_size=1 and 2, crashes on main, passes with the fix. Added a regression test mirroring it to TestFsdp2UnshardedParams, confirmed it fails without the fix first.
System Info
main (post #3839,
src/peft/peft_model.py::_reshard_fsdp_modules)Who can help?
@hivaze @BenjaminBossan
Reproduction
_reshard_fsdp_modules()(added in #3839) skips calling.reshard()on an FSDP2 module if any tuner layer merged into that module's subtree, to avoid dropping merged weights:The skip check works at FSDP-module granularity. If an FSDP unit wraps more than one target module (the common case: a decoder layer wrapped as one FSDP unit contains
q_proj,v_proj, etc.), merging just one of them marks the entire FSDP unit as "contains a merged layer" and skips resharding it, including for every other, still-unmerged adapter in that same unit. That reintroduces exactly the bug #3839 fixed, just for the sibling adapter instead of the merged one.Proof that the reshard is actually skipped (no multi-GPU needed, this part of the defect doesn't depend on world size, it's a pure graph-traversal issue):
Output:
v_proj's adapter is never merged, but it loses its reshard entirely onceq_proj(a different target module in the same FSDP unit) gets merged, for as long as that merge persists. I don't have multi-GPU hardware to additionally confirm the downstreamrequires_gradstaleness end-to-end atworld_size > 1(that specific parameter-swap behavior needs real sharding to observe), but the skippedreshard()call itself is the root cause #3839 is built around, and it's independent of world size.Expected behavior
The skip should be scoped to the specific tuner layer(s) that are merged, not the whole FSDP module. One option: instead of skipping
fsdp_module.reshard()outright, only exclude it when the merged tuner layer's parameters can't be separated from the rest of the unit's sharding (which may be the common case forfully_shard), in which case the limitation should at least be called out explicitly (today's doc note for #3839 doesn't mention this partial-merge case). A more complete fix would reshard everything except the specific merged submodule's parameters, if FSDP2's API allows resharding at that granularity, or otherwise document that merging any target module inside a shared FSDP unit disables the fix for every other adapter in that unit until unmerged.