Skip to content

Support sync_ref_model with PEFT by syncing the reference adapter - #6975

Open
behroozazarkhalili wants to merge 7 commits into
mainfrom
fix/3108-sync-ref-model-peft
Open

Support sync_ref_model with PEFT by syncing the reference adapter#6975
behroozazarkhalili wants to merge 7 commits into
mainfrom
fix/3108-sync-ref-model-peft

Conversation

@behroozazarkhalili

@behroozazarkhalili behroozazarkhalili commented Aug 29, 2026

Copy link
Copy Markdown
Collaborator

What

sync_ref_model=True now works with PEFT in GRPOTrainer, DPOTrainer, KTOTrainer and RLOOTrainer. It previously raised NotImplementedError.

Why the guard was too broad

Every one of the four trainers already creates the EMA target a sync needs: a frozen "ref" adapter copied from "default", at grpo_trainer.py:471, dpo_trainer.py:676, kto_trainer.py:724 and rloo_trainer.py:384. The reference forward already selects it via use_adapter(model, adapter_name="ref" if "ref" in model.peft_config else None). The guard rejected that case anyway.

Measured on main before this change, with a LoRA config and beta != 0:

path A (peft_config passed to trainer)  adapters=['default']         sync=True -> NotImplementedError
path B (already-PEFT model passed in)   adapters=['default', 'ref']  sync=True -> NotImplementedError

Path B has the target and was still refused.

The three changes

SyncRefModelCallback. _sync_target_model does zip(target_model.parameters(), model.parameters(), strict=True), which assumes two modules. With PEFT the reference is a second adapter inside the policy model, so there is nothing to zip. The new branch pairs .default. parameters with their .ref. counterparts by name, which is the same mapping the trainers use to initialize the adapter. sync_ref_adapter mirrors the existing sync_target_model ZeRO-3 GatheredParameters wrapper.

The four trainers. Both PEFT entry paths now reach one copy of the create-and-copy block through a needs_ref_adapter flag. On the peft_config path the adapter is created only when a synced reference is requested, because a freshly wrapped adapter is zero-initialized and the base model already serves as a fixed reference; that equivalence breaks as soon as the reference has to move. Each trainer uses its own predicate for "a PEFT reference is in use here", taken from its own elif: args.beta != 0.0 in GRPO, ref_model is None in DPO and KTO, unconditional in RLOO.

The guard. It now fires only when no reference adapter could be created: peft<0.20.0 with a LoRA config using target_parameters (peft#3340), or a standalone ref_model passed alongside a PEFT policy, whose parameters do not pair with the adapter's. In both cases the reference cannot track the policy, so there is nothing to sync, and the message says which case applies instead of describing PEFT as unsupported.

Bias-bearing LoRA is rejected, not worked around. A bias other than "none" trains bias terms that live in the base model, so disabling the adapter does not recover a fixed reference: measured on the tiny Qwen2 model, bias="all" makes 6 base biases trainable and the adapter-disabled logits move when they move (PEFT warns about the same thing). PEFT also refuses a second bias-bearing adapter, so no frozen copy is possible. All four trainers raise a ValueError at construction, with or without sync_ref_model; an earlier revision of this PR warned and fell back to the base model, which would have trained against a drifting reference.

vLLM weight sync. The sync loops skipped the "ref" adapter's LoRA matrices through the tuner prefix but not its modules_to_save copy: lm_head.modules_to_save.ref.weight was streamed to vLLM under that name. The four loops (server and colocate, in VLLMGeneration and OnlineDPO) now skip every modules_to_save copy except "default". Reachable through GRPO and RLOO with peft_config, sync_ref_model=True, use_vllm=True and modules_to_save; verified against the PEFT state dict of the tiny model rather than a live vLLM server.

Verification

Red then green, per trainer. Reverting only the source and keeping the four new tests makes all four fail with NotImplementedError; with the source they pass.

check result
new test_train_with_sync_ref_model_and_peft x4, with fix 4 passed
same 4 tests, source reverted to main 4 failed, NotImplementedError
existing test_train_with_sync_ref_model x4 4 passed
construction matrix, 4 trainers x 2 paths x sync on/off 16/16
DPO and KTO sync tests: after training, the trainer's reference log probs differ from the base model's with adapters disabled pass; a tree whose reference path ignores the "ref" adapter fails both
bias="all" rejected at construction x4 4 passed

End to end on a real two-step GRPO run with ref_model_mixup_alpha=0.6: with sync_ref_model=True the eight .ref. parameters move by 59.73 in total; with it off there are no .ref. parameters and nothing moves.

The callback change alone is measurable in isolation: at alpha=1.0 the summed absolute difference between the "default" and "ref" LoRA B matrices goes from 6144.0 to 0.0, meaning the reference becomes the policy exactly. Before the change the same call left it at 6144.0.

Memory

sync=False on the peft_config path still yields ['default'] in all four trainers, so a LoRA user who does not ask for a synced reference allocates nothing extra. Asking for one costs an adapter, not a second copy of the model.

Resolves #3108


Note

Medium Risk
Touches reference-model semantics and parameter sync across four RLHF trainers and vLLM export; incorrect adapter pairing could silently break KL/reference terms, though coverage is heavy on edge-case naming and integration tests.

Overview
sync_ref_model=True now works with PEFT in DPO, KTO, GRPO, and RLOO. Instead of always erroring, trainers can add a frozen "ref" adapter when sync is requested and SyncRefModelCallback EMA-updates matched "default" / "ref" parameters (including trainable_token_indices and nested modules_to_save names), with a DeepSpeed ZeRO-3 gather path parallel to the standalone ref_model sync.

Trainer setup changes: the "ref" adapter is created only when sync needs a moving reference (not for every LoRA run); copying into "ref" uses the same name-pairing logic as the callback rather than a naive .default..ref. replace. LoRA bias other than "none" is rejected at init because the reference cannot stay fixed. The old blanket PEFT NotImplementedError is limited to cases where no "ref" adapter exists (e.g. peft<0.20.0 + target_parameters, or PEFT policy plus standalone ref_model).

vLLM weight sync skips non-default modules_to_save adapter copies so the frozen "ref" head is not pushed to vLLM. DPO/KTO docs describe the updated sync_ref_model + PEFT constraints.

Reviewed by Cursor Bugbot for commit 04fdedd. Bugbot is set up for automated code reviews on this repo. Configure here.

`sync_ref_model=True` raised NotImplementedError on every PEFT model in
GRPO, DPO, KTO and RLOO, on the grounds that PEFT keeps no separate
`ref_model` to synchronize. All four trainers already build what a sync
needs: a frozen "ref" adapter copied from "default", created when an
already-PEFT model is passed in. The guard rejected that case too, so it
was broader than the limitation it described.

SyncRefModelCallback now handles it. The existing `_sync_target_model`
zips the parameters of two modules, which cannot express a reference
that lives inside the policy model, so the PEFT branch pairs each
"default" parameter with its "ref" counterpart by name, reusing the
mapping that initializes the adapter. The ZeRO-3 gather wrapper mirrors
the one already beside it.

The trainers also create the "ref" adapter when the caller passes
`peft_config` and asks for a synced reference. Without that, the same
config would work or raise depending only on how the model was
constructed. A caller who leaves `sync_ref_model=False` still gets no
extra adapter, so nobody pays for one they did not ask for.

The guard now fires only when no reference adapter could be created at
all, which is `peft<0.20.0` with a LoRA config using `target_parameters`.

Resolves #3108
@bot-ci-comment

Copy link
Copy Markdown

The docs for this PR live here. All of your documentation changes will be reflected on that endpoint. The docs are available until 30 days after the last update.

@Jokasa7 Jokasa7 left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

I traced the new EMA path through PEFT 0.20.0's actual parameter layouts and all four duplicated trainer branches. Default LoRA matrices follow the intended path, but two supported LoRA configurations remain either silently unsynchronized or fail while constructing the reference adapter; details inline. I used AI assistance for source navigation and verified both findings against the exact PR head and exact upstream methods with executable reproductions.

Comment thread trl/trainer/callbacks.py Outdated
# `_sync_target_model`'s parameter-wise zip of two modules does not apply. Pair each `"default"` parameter with
# its `"ref"` counterpart by name instead; this is the same mapping used to initialize the `"ref"` adapter.
for name, param in model.named_parameters():
if ".default." in name:

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

[P2] Map PEFT state structurally instead of matching .default.

This predicate misses valid parameters from LoraConfig(trainable_token_indices=...). PEFT 0.20.0 stores those deltas in a ParameterDict under names ending in ...trainable_tokens_delta.default, so there is no trailing dot after the adapter name. The identical initialization loops in the four trainers and this callback therefore never copy or EMA-update the corresponding .ref delta. Executing the exact-head method with a default delta of 7 and a ref delta of 0 left ref at 0 even with alpha=1.0, while reference forwards explicitly select "ref". The same hard-coding also copies default rather than model.active_adapter when an already-PEFT model is training a custom adapter. Please map through PEFT adapter state dicts, or otherwise handle terminal adapter keys and derive the source adapter explicitly; a trainable_token_indices regression test would catch the silent stale-reference case.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Confirmed. Fixed in be926d5.

On a tiny Qwen2 with LoraConfig(trainable_token_indices=[0, 1]) the parameter is base_model.model.model.embed_tokens.token_adapter.trainable_tokens_delta.default. It ends in .default with nothing after it, so ".default." in name is false and one of five adapter parameter pairs was skipped. Setting the policy delta to 7.0 and the reference to 0.0, then running the callback with alpha=1.0, left the reference at 0.0.

Getting the replacement right took three attempts, and the two wrong ones are worth recording because they fail in opposite directions:

  1. Matching "default" as a path component fixes the terminal case but matches a base-model parameter of that name. PEFT reserves neither default nor ref, so a model owning base_model.model.default produced base_model.model.ref, which nothing provides, and get_parameter raised during construction.
  2. Taking the last such component fixes that but breaks modules_to_save. PEFT wraps the saved module in a ModuleDict keyed by adapter, so a saved module with its own default child reads ...block.modules_to_save.default.default.weight, where the adapter key is the first component. The last one belongs to the base module, so the parameter was skipped and its reference never moved, reintroducing exactly the defect you reported.

What actually distinguishes an adapter parameter is the container: PEFT keys them by adapter name inside a ModuleDict (LoRA matrices, modules_to_save) or a ParameterDict (the token deltas). The scan now walks the default components from the end and takes the first whose container also holds a ref key, which additionally declines a base ModuleDict that happens to have a default entry and no counterpart.

Five tests in TestSyncRefModelCallbackAdapterPairing cover the terminal token delta, a base module and parameter named default, a modules_to_save child of the same name, a base ModuleDict with no ref key, and the EMA equation itself. Each of the four ways to get the rule wrong fails one of them: reverting to the last-occurrence rule, dropping the ref membership test, dropping the container-type test, and replacing the EMA update with a plain copy.

Comment thread trl/trainer/dpo_trainer.py Outdated
# model is the reference. That equivalence only holds while the reference stays fixed: with
# `sync_ref_model=True` the reference has to track the policy, which requires parameters of its own to
# move. So in that case create the "ref" adapter here as well.
needs_ref_adapter = args.sync_ref_model and ref_model is None

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

[P2] Handle bias-bearing LoRA before creating the ref adapter

This new peft_config path also reaches model.add_adapter("ref", default_config) for valid LoraConfig(bias="all") and bias="lora_only" configurations. PEFT 0.20.0 rejects that second adapter in BaseTuner._check_new_adapter_config with supports only 1 adapter with bias, which I reproduced using the exact upstream method. All four trainers fail here before the new "ref" not in peft_config guard, so the guard's message that only old target_parameters configurations lack a sync target is incomplete. Since these bias parameters are shared rather than adapter-local, please either reject this configuration explicitly before adding "ref" with an accurate explanation, or use a synchronization design with an independent target for it.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Confirmed. Fixed in be926d5.

LoraConfig(bias="all") and bias="lora_only" both raise ValueError: LoraModel supports only 1 adapter with bias from add_adapter("ref", ...), while bias="none" succeeds. The call sat after the guard, so those configurations failed while the trainer was still being constructed.

The guard now skips the reference adapter for them, in the same shape as the existing target_parameters skip. With sync_ref_model=True that lands on the existing check for a missing "ref" adapter, so the configuration is rejected with an explanation rather than a PEFT stack trace. Skipping at this layer rather than raising here is deliberate: needs_ref_adapter also covers a pretrained PEFT model without synchronization, where a fixed base-model reference is supported and raising would break a working setup.

You were right that the guard's message was incomplete, and it was worse than incomplete once bias could reach it. It named target_parameters and peft<0.20.0 as the only cause, which would have been an actively misleading diagnosis for a bias config. It now names both causes and both remedies.

On your point about rejecting versus falling back: the reference log probs do come from the base model, which is fixed and cannot track the policy, so asking for both a bias-bearing adapter and sync_ref_model=True is refused rather than silently trained against a frozen reference. The regression test asserts the raise rather than a successful run.

…ring LoRA

The "default" to "ref" pairing matched the adapter name as the substring
".default.", which requires a path component after it. `trainable_token_indices`
stores its deltas in a `ParameterDict` keyed by adapter, so those parameter names
end in ".default" with nothing following. The predicate skipped them, and the
reference copy of the token deltas stayed frozen while the policy moved. Matching
"default" as a path component covers both placements: on a tiny Qwen2 with
`trainable_token_indices=[0, 1]`, one of five adapter parameter pairs went
unsynced before, none after.

A LoRA config with `bias` other than "none" trains bias terms shared with the
base model, and PEFT permits one such adapter per model: `add_adapter("ref", ...)`
raises `LoraModel supports only 1 adapter with bias` for both "all" and
"lora_only", while "none" succeeds. That call sat after the guard, so those
configurations failed during trainer construction. The guard now skips the
reference adapter for them, which routes into the existing `sync_ref_model`
rejection rather than a PEFT stack trace.

That rejection named `target_parameters` and `peft<0.20.0` as the only cause,
which is wrong once bias reaches it, so the message now names both.

All five call sites move together, as the trainers require.
…ameters

The eight tests added in the previous commit build a `LoraConfig`, which is
imported only when PEFT is installed, and they carried no `@require_peft` while
every neighbouring PEFT test does. All eight failed the no-optional-dependency
CI lane with `NameError: name 'LoraConfig' is not defined`. A local run cannot
show this, because PEFT is installed there.

Pairing "default" with "ref" on the path component alone was also wrong in both
directions. PEFT reserves neither name, so a base model owning a parameter called
`default` produced a reference name nothing provides:
`base_model.model.default -> base_model.model.ref` raised `AttributeError`
during construction. Taking the last such component instead was wrong the other
way: `modules_to_save` wraps the saved module in a `ModuleDict` keyed by adapter,
so a saved module with its own `default` child reads
`...modules_to_save.default.default.weight`, whose adapter key is the first
component. The last one belongs to the base module, so the parameter was skipped
and its reference never moved, which is the defect this branch set out to fix.

Adapter parameters are the ones PEFT keys by adapter name inside a `ModuleDict`
or a `ParameterDict`, so the scan now walks the candidates from the end and takes
the first whose container also holds a "ref" key. That covers the LoRA matrices,
the `trainable_token_indices` deltas and `modules_to_save` alike, and it declines
a base container that happens to hold a `default` key with no counterpart.

Five callback tests cover the cases: a base module and parameter named `default`,
a `modules_to_save` child of the same name, a base `ModuleDict` without a "ref"
key, the terminal token delta, and the EMA equation itself. Each of the four ways
to get the rule wrong fails one of them.
…f the vLLM sync, and correct the docs

A LoRA config with `bias` other than `"none"` trains bias terms that
live in the base model. The previous commit turned PEFT's crash on the
second adapter into a warning saying the reference would come from the
base model with adapters disabled. That reference is not fixed: with
`bias="all"` six base biases train, and the adapter-disabled logits
move when they move (PEFT warns about the same thing). All four
trainers now raise a `ValueError` at construction, with or without
`sync_ref_model`, and the `NotImplementedError` for a missing "ref"
adapter no longer lists bias as a cause. DPO and KTO also name the
case it had left out: a standalone `ref_model` passed alongside a PEFT
policy, whose parameters do not pair with the adapter's.

The vLLM weight sync skipped the "ref" adapter's LoRA matrices through
the tuner prefix but not its `modules_to_save` copy:
`lm_head.modules_to_save.ref.weight` was streamed to vLLM under that
name. The four sync loops (server and colocate, in `VLLMGeneration`
and OnlineDPO) now skip every `modules_to_save` copy except the
"default" one. Reachable since the previous commit through GRPO and
RLOO with `peft_config`, `sync_ref_model=True`, `use_vllm=True` and
`modules_to_save`.

RLOO created the "ref" adapter for `beta=0.0` and then raised that
there is nothing to synchronize; its guard now carries GRPO's
`beta != 0.0` clause.

The DPO and KTO sync tests checked only that the adapter's parameters
moved. They now also compute the reference log probs through the
trainer's own path after training and require them to differ from the
base model's, taken with adapters disabled at the same moment: a fresh
"ref" adapter reproduces the base model exactly, so only a synced one
can differ. A copy of the tree whose reference path ignores the "ref"
adapter fails that assertion in both trainers. (A before/after
comparison was tried first and rejected: enabling gradient
checkpointing during `train()` shifts the reference forward by 1e-3 on
identical weights, which a same-moment comparison does not see.) The callbacks fixture said it owned a
base parameter named "default" but registered "default_bias"; it now
owns one, and the test asserts the sync leaves it alone.

The DPO and KTO docs still said `sync_ref_model=True` is unsupported
with PEFT models; they now describe the "ref" adapter copy and the
three cases that are rejected.

@cursor cursor Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Cursor Bugbot has reviewed your changes using default effort and found 2 potential issues.

Fix All in Cursor

❌ Bugbot Autofix is OFF. To automatically fix reported issues with cloud agents, enable autofix in the Cursor dashboard.

Reviewed by Cursor Bugbot for commit 2d45dc4. Configure here.

if "original_module" in name:
# When module to save, remove its prefix and discard the original module, as well as the copies held by
# other adapters (such as the frozen "ref" one); vLLM sees only the "default" copy
if "original_module" in name or (".modules_to_save." in name and ".modules_to_save.default." not in name):

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

vLLM sync leaks ref token deltas

Medium Severity

The new vLLM skip only drops non-default modules_to_save copies. Adapter-keyed PEFT state that does not use that prefix, such as trainable_tokens_delta.ref, is still streamed under its PEFT name. That path is now reachable with peft_config, sync_ref_model=True, and use_vllm=True, so weight sync can push names vLLM does not own.

Additional Locations (2)
Fix in Cursor Fix in Web

Reviewed by Cursor Bugbot for commit 2d45dc4. Configure here.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

Confirmed, and the leak was wider than the ref copy. On a tiny Qwen2 with LoraConfig(trainable_token_indices=[0, 1]) plus a ref adapter, the filter let three names through: model.embed_tokens.token_adapter.weight, ...trainable_tokens_delta.default and ...trainable_tokens_delta.ref. Only the last one is new in this PR; the other two mean trainable_token_indices never worked with vLLM sync. Fixed in 04fdedd in all four loops: PEFT merges the deltas into the embedding before the sync (merge_adapter covers TrainableTokensLayer), so the deltas are skipped and the token_adapter. prefix is stripped, which puts the merged weight under its base name. The same probe now reports no leaked name for plain LoRA, trainable_token_indices and modules_to_save.

Comment thread trl/trainer/dpo_trainer.py
…keep trainable-token state out of the vLLM sync

The bias check ran only when a "ref" adapter was about to be created, so with sync_ref_model off a LoRA config with bias="all" was still accepted although the docs and the error message said otherwise. The reason the message gives holds on both paths: the trained biases live in the base model, so disabling the adapter does not give a fixed reference whether or not the reference is synced. The check now runs whenever the trainer takes its reference from the PEFT model, on the peft_config path and on the pretrained-adapter path, in all four trainers. The bias tests cover sync on and off.

With trainable_token_indices the vLLM sync pushed three names vLLM does not own: the wrapped embedding weight under its token_adapter prefix and the per-adapter deltas for "default" and "ref". Only the "ref" copy is new in this PR; the other two mean the feature never worked with vLLM. PEFT merges the deltas into the embedding before the sync, so the four loops now skip trainable_tokens_delta and strip the token_adapter prefix, and the merged weight lands under its base name. Measured on a tiny Qwen2 with a fake process group: three leaked names before, none after.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Sync Reference Model is incompatible with LoRA in GRPO Trainer

2 participants