Skip to content

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

Description

@rupeshpoojary9

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:

for fsdp_module in fsdp_modules:
    if merged_modules:
        sharded_modules, stack = set(), [fsdp_module]
        while stack:
            module = stack.pop()
            sharded_modules.add(module)
            stack.extend(child for child in module.children() if child not in fsdp_modules)
        if not sharded_modules.isdisjoint(merged_modules):
            continue
    fsdp_module.reshard()

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):

import torch
import torch.distributed as dist
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.fsdp import fully_shard
from transformers import LlamaConfig, LlamaForCausalLM
from peft import LoraConfig, get_peft_model

dist.init_process_group(backend="gloo", init_method="tcp://127.0.0.1:29515", rank=0, world_size=1)
mesh = init_device_mesh("cpu", (1,))
torch.manual_seed(0)
config = LlamaConfig(vocab_size=128, hidden_size=64, intermediate_size=128,
    num_hidden_layers=1, num_attention_heads=4, num_key_value_heads=2, tie_word_embeddings=False)
model = LlamaForCausalLM(config)
model = get_peft_model(model, LoraConfig(r=8, target_modules=["q_proj", "v_proj"], init_lora_weights=False))
layer = model.base_model.model.model.layers[0]
fully_shard(layer, mesh=mesh, reshard_after_forward=False)
fully_shard(model, mesh=mesh, reshard_after_forward=False)

q_proj = layer.self_attn.q_proj
reshard_calls = []
orig_reshard = layer.reshard
def spy_reshard(*a, **kw):
    reshard_calls.append(layer)
    return orig_reshard(*a, **kw)
layer.reshard = spy_reshard

model.set_requires_grad(adapter_names="default", requires_grad=True)
print(f"reshard calls on `layer` before any merge: {len(reshard_calls)}")  # 1

reshard_calls.clear()
q_proj.merge()  # v_proj is untouched, still unmerged
model.set_requires_grad(adapter_names="default", requires_grad=True)
print(f"reshard calls on `layer` after q_proj (only) is merged: {len(reshard_calls)}")  # 0, bug

Output:

reshard calls on `layer` before any merge: 1
reshard calls on `layer` after q_proj (only) is merged: 0

v_proj's adapter is never merged, but it loses its reshard entirely once q_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 downstream requires_grad staleness end-to-end at world_size > 1 (that specific parameter-swap behavior needs real sharding to observe), but the skipped reshard() 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 for fully_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.

Activity

  1. hivaze commented on Oct 3, 2026

    @hivaze
    Contributor

    Thanks, that's right, and it's on purpose. FSDP2 reshards a whole module at once, so v_proj can't be resharded without dropping the merge in q_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 new requires_grad lands on the unsharded copies. After merge_adapter() and set_adapter("other"), the shards still have default trainable and other frozen. So the first step trains other, and with only the adapters trainable the second raises element 0 of tensors does not require grad. Calling unmerge_adapter() before set_adapter() avoids it, so unmerging before the reshard in PeftModel.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.

  2. rupeshpoojary9 commented on Oct 3, 2026

    @rupeshpoojary9
    ContributorAuthor

    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 in tuners_utils.py (if module.merged: warnings.warn(...); module.unmerge()), called from LoraModel.set_adapter → PeftModel.set_adapter, which runs after self._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 in PeftModel.set_adapter(), and the same in disable_adapter()'s exit path since enable_adapter_layers() goes through the same helper, should close it. I'd guard it behind "is anything actually merged" first, since unmerge_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.

  3. hivaze commented on Oct 4, 2026

    @hivaze
    Contributor

    Sure, script below. It needs two things: a forward before merge_adapter(), so the merge lands on the unsharded copies, and reshard_after_forward=False. Without that first forward, or with reshard_after_forward=True, all three steps train other, 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=1 and then --nproc_per_node=2 print:

    [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 other through the unsharded copies, its backward reshards, and step 2 gathers from shards where other is still frozen.

    If the fix also touches the disable_adapter() exit, note that #3800 has a proposed change there.

  4. BenjaminBossan commented on Oct 5, 2026

    @BenjaminBossan
    Member

    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

  5. rupeshpoojary9 commented on Oct 5, 2026

    @rupeshpoojary9
    ContributorAuthor

    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.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions