Clipped Importance Sampling Policy Optimization (CISPO)¶
CISPO (Clipped Importance Sampling Policy Optimization) is a
GRPO specialization that clips
importance weights directly and uses them to scale a log-prob objective.
CISPO uses the same group-based advantage calculation as GRPO, however, the objective function is closer to that of REINFORCE, multiplying the log-probability term of the function by a scaled importance ratio. A stop gradient is applied to the importance ratio, meaning the ratio is treated as a constant that scales each token’s contribution to the overall policy gradient.
In AgileRL, CISPO can be used for single-turn reasoning tasks or multi-turn agentic finetuning. In the multi-turn case, rollouts are still treated as a bandit problem, with environment generated tokens masked and reward signal calculated from cumulative episode reward.
Example¶
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
from agilerl.algorithms import CISPO
model = AutoModelForCausalLM.from_pretrained(
"Qwen/Qwen2.5-3B",
torch_dtype=torch.bfloat16,
device_map="auto",
)
tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen2.5-3B")
agent = CISPO(
actor_network=model,
pad_token_id=tokenizer.eos_token_id,
pad_token=tokenizer.eos_token,
device="cuda" if torch.cuda.is_available() else "cpu",
batch_size=8,
group_size=8,
)
Training and Usage¶
Use CISPO anywhere you would use GRPO in AgileRL training loops, such
as train_llm_rollout. Single-turn reasoning is the max_turns=1 case
of the same function.
from agilerl.llm_envs import RolloutHarness
from agilerl.training.llm import train_llm_rollout
def reward_fn(completion: str, answer: str, question: str) -> float:
del question
return float(answer.lower() in completion.lower())
# 1) Single-turn / reasoning datasets (a prompt dataset is just an env, max_turns=1)
class QADataset:
"""Single-turn dataset env: serve a question on reset, score it on step."""
def __init__(self, questions, answers, reward_fn, prompt_builder,
test_questions=None, test_answers=None):
self.questions, self.answers = questions, answers
self.test_questions, self.test_answers = test_questions, test_answers
self.reward_fn, self.prompt_builder = reward_fn, prompt_builder
self._cursor, self._split = 0, ""
@property
def dataset_size(self) -> int:
return len(self.questions)
def reset(self, seed=None, *, row_index=None, evaluation=None):
if evaluation and self.test_questions is not None:
qs, ans, split = self.test_questions, self.test_answers, "eval"
else:
qs, ans, split = self.questions, self.answers, "train"
if row_index is None:
if split != self._split:
self._cursor, self._split = 0, split
row_index, self._cursor = self._cursor, self._cursor + 1
self._q, self._a = qs[row_index % len(qs)], ans[row_index % len(ans)]
return self.prompt_builder(self._q), {}
def step(self, action):
return "", float(self.reward_fn(action, self._a, self._q)), True, False, {}
env_factory = lambda: RolloutHarness.local(
QADataset(
questions=["2+2?", "Capital of France?"],
answers=["4", "Paris"],
reward_fn=reward_fn,
prompt_builder=lambda question: f"Q: {question}\nA:",
test_questions=["3+3?"],
test_answers=["6"],
),
tokenizer,
max_turns=1,
pad_id=tokenizer.eos_token_id,
apply_chat_template=True,
max_model_len=1024,
)
trained_pop = train_llm_rollout(
pop=[agent],
max_turns=1,
env_factory=env_factory,
max_steps=2000,
evaluation_interval=50,
)
# 2) Multi-turn text environments (each rollout drives its own in-process env)
class ToyRolloutEnv:
def reset(self, seed=None):
del seed
return "Start: What is 2+2?", {}
def step(self, action: str):
reward = 1.0 if "4" in action else 0.0
return "Done.", reward, True, False, {"correct": bool(reward)}
env_factory = lambda: RolloutHarness.local(
ToyRolloutEnv(),
tokenizer,
max_turns=4,
pad_id=tokenizer.eos_token_id,
max_model_len=1024,
max_output_tokens=128,
)
trained_pop = train_llm_rollout(
pop=[agent],
max_turns=4,
env_factory=env_factory,
max_steps=2000,
evaluation_interval=50,
)
Saving and Loading Agents¶
To save an agent, use the save_llm_checkpoint function:
from agilerl.utils.utils import save_llm_checkpoint
save_llm_checkpoint(agent, "path/to/checkpoint")
Parameters¶
- class agilerl.algorithms.cispo.CISPO(pad_token_id: int, pad_token: str, model_name: str | None = None, actor_network: PreTrainedModel | PeftModel | None = None, model_config: dict[str, Any] | None = None, hp_config: HyperparameterConfig | None = None, index: int = 0, batch_size: int = 16, beta: float = 0.001, lr: float = 5e-07, clip_coef: float | tuple[float, float] = 0.2, max_grad_norm: float = 0.1, update_epochs: int = 1, group_size: int = 8, temperature: float = 0.9, repetition_penalty: float = 1.0, top_p: float = 1.0, top_k: int = 50, min_p: float = 0.0, offload_trainer_during_rollout: bool = True, calc_position_embeddings: bool = True, micro_batch_size_per_gpu: int | None = None, mini_batch_size: int | None = None, max_output_tokens: int | None = None, min_output_tokens: int | None = None, max_model_len: int | None = 1024, max_row_tokens: int | None = None, hf_generate_chunk_size: int | None = None, lora_config: LoraConfig | None = None, cosine_lr_schedule_config: CosineLRScheduleConfig | None = None, fsdp_config: FSDPConfig | None = None, device: str | torch.device | None = None, wrap: bool = True, clone: bool = False, vllm_config: VLLMConfig | None = None, seed: int = 42, gradient_checkpointing: bool = True, torch_compiler: str | None = None, use_liger_loss: bool = True, use_kl_advantage_shaping: bool = False, adv_norm: str = 'mean_only', importance_sampling_level: Literal['token', 'turn', 'trajectory'] | None = None, advantage_granularity: Literal['auto', 'trajectory', 'turn'] = 'auto', use_separate_reference_adapter: bool = True, whiten_advantages: bool = False, adv_clip_range: float | None = None, filter_zero_adv: bool = False, adv_filter_eps: float = 0.0, turn_advantage_trajectory_fallback: bool = True, cast_logprobs_to_fp32: bool = True, chunk_rows: int | None = None, quantization_config: BitsAndBytesConfig | None = None, activation_offload: bool = False, moe_lora_recompute: bool | None = None, lora_target_scope: str | None = None, vllm_importance_sampling_correction: bool = True, vllm_importance_sampling_cap: float = 2.0, vllm_max_logprob_gap: float = 0.1, vllm_max_clip_fraction: float = 0.02, use_sequence_packing: bool = False, loss_norm: Literal['micro_batch', 'accumulation_window', 'episode'] = 'accumulation_window', profiling_config: ProfilingConfig | None = None, old_logprobs_source: Literal['trainer', 'rollout'] = 'trainer', off_policy_token_mask_bounds: tuple[float, float] | None = (0.5, 5.0), off_policy_sequence_mask_threshold: float | None = 0.03, use_bias_correction_kl: bool = True, kl_clamp: float | None = 10.0)¶
CISPO loss variant of
agilerl.algorithms.grpo.GRPOClamps importance weights from above only;
clip_coef_mindoes not apply to this objective.Paper: https://arxiv.org/abs/2506.13585
- clone(index: int | None = None, wrap: bool = True) Self¶
Create a clone of the algorithm.
QLoRA clones rebuild the base via
from_pretrainedand transfer only adapter (+ value head) weights. FSDP2 clones copy a rank-0 CPU full state dict onto a fresh CPU actor, then shard. The dense full model is never placed on GPU.- Parameters:
index (int | None, optional) – The index of the clone, defaults to None
wrap (bool, optional) – Unused. Clones always call
wrap_models(). Kept so tournament / multi-frequency can passwrap=False.
- Returns:
A clone of the algorithm
- Return type:
- configure_batch_size_per_process(batch_size: int, micro_batch_size_per_gpu: int | None, mini_batch_size: int | None, group_size: int = 1) None¶
Derive per-process batch sizes and gradient accumulation steps.
batch_sizeis the global collect size (prompt groups for GRPO-family). Each data-parallel replica holds(batch_size / dp_size) * group_sizesamples, wheredp_sizeis the process-group world folded by tensor-parallel degree (world_size / tp). Unsetmini_batch_sizeusesmicro_batch_size_per_gpuwhen the class default is"micro_batch"(RL rollout algorithms) and that is set, else the per-rank collect. Unsetmicro_batch_size_per_gpuuses the mini-batch.gradient_accumulation_stepsismini_batch_size / micro_batch_size_per_gpu.
- static copy_attributes(agent: IndividualT, clone: IndividualT, exclude: Iterable[str] = ()) IndividualT¶
Copy the non-evolvable attributes of the algorithm to a clone.
- Parameters:
agent (EvolvableAlgorithm) – The algorithm to copy attributes from.
clone (EvolvableAlgorithm) – The clone of the algorithm.
exclude (Iterable[str]) – Attribute names to leave on
clone/agent.
- Returns:
The clone of the algorithm.
- Return type:
- property current_lr_critic: float¶
Critic learning rate the next
learncall trains with;lris the peak withoutlr_critic.
- eval_policy_network_ids() set[int]¶
Return the id of every evaluation network in the agent’s policy group.
- evolvable_attributes(networks_only: bool = False) dict[str, Any]¶
Return the attributes related to the evolvable networks in the algorithm. Includes attributes that are either EvolvableModule or ModuleDict objects, as well as the optimizers associated with the networks.
- finalize_training_step(num_steps: int) None¶
Close the agent’s training block, storing any captured GraMa scores.
- Parameters:
num_steps (int) – Number of steps taken during the training step.
- Returns:
None.
- Return type:
None
- property fitness: list[float | ndarray[tuple[int, ...], dtype[_ScalarType_co]]]¶
Fitness history (scalars, or per-sub-agent rows for multi-agent).
- get_action(obs: list[RolloutPrompt] | RolloutPrompt, training: bool = True, repeat_prompts: bool = True, *args: Any, **kwargs: Any) ActionResult¶
Return generated completions for each prompt (GRPO groups when training).
- Parameters:
obs (LLMObsType) – List of HF-style prompt dicts (this implementation mutates them).
training (bool) – If
True, generate with training sampling settings.repeat_prompts (bool) – If
Trueandtraining=True, duplicate each promptself.group_sizetimes (legacy GRPO grouped mode). IfFalse, treat the batch as already expanded trajectories.
- Returns:
An
ActionResultof completion token IDs, per-sequence action masks, and (when captured) per-completion vLLM sampling logprobs for the mismatch correction.- Return type:
ActionResult
- static get_action_dim(action_space: Space | list[Space] | dict[str, Space]) int | dict[str, int] | tuple[int | dict[str, int], ...]¶
Return the dimension of the action space as it pertains to the underlying networks (i.e. the output size of the networks).
- get_eval_modules(cloning: bool = True) tuple[dict[str, EvolvableModule], dict[str, EvolvableModule]]¶
Get the offsprings of all of the evaluation modules in the individual.
- Parameters:
cloning (bool, optional) – Whether to clone each evaluation module before returning it, defaults to True.
- Returns:
Tuple of offspring policy and the rest of the evaluation modules
- Return type:
tuple[dict[str, EvolvableModule], dict[str, EvolvableModule]]
- get_lr_names() list[str | tuple[str, str]]¶
Return the learning-rate attribute name(s) of each optimizer.
- get_policy() EvolvableModuleProtocol¶
Return the policy network of the algorithm.
- static get_state_dim(observation_space: Space | list[Space] | dict[str, Space]) tuple[int, ...] | dict[str, tuple[int, ...]] | tuple[tuple[int, ...] | dict[str, tuple[int, ...]], ...]¶
Return the dimension of the state space as it pertains to the underlying networks (i.e. the input size of the networks).
- property hp_config: HyperparameterConfig¶
Return the hyperparameter configuration for Evo-HPO mutations.
- init_training_step(capture_grama: bool = False) None¶
Open the agent’s training block: metrics tracking, and GraMa capture.
Hooks are registered afresh each cycle, so they follow the agent through architecture mutations, checkpoint reloads and accelerator re-wrapping. Opening a block implicitly closes one that an earlier call left open.
- Parameters:
capture_grama (bool) – Whether to register GraMa capture hooks for this training step. Defaults to False since the LLM finetuners never run ReGraMa.
- Returns:
None.
- Return type:
None
- static inspect_attributes(agent: EvolvableAlgorithmProtocol | AgentWrapperProtocol[Any], input_args_only: bool = False, exclude: Iterable[str] = ()) dict[str, Any]¶
Inspect and retrieve the attributes of the current object, excluding attributes related to the underlying evolvable networks (i.e. EvolvableModule, torch.optim.Optimizer) and with an option to include only the attributes that are input arguments to the constructor.
- Parameters:
input_args_only (bool) – If True, only include attributes that are input arguments to the constructor. Defaults to False.
exclude (Iterable[str], optional) – Extra attribute names to drop from the result, on top of the standard exclusions below. For a caller-specific reason to leave an attribute out of its own view.
- Returns:
A dictionary of attribute names and their values.
- Return type:
- learn(experiences: tuple[list[Tensor] | Tensor, list[Tensor] | Tensor, Tensor], turn_ids: Tensor | None = None, sampling_logps: list[Tensor | None] | None = None, pixel_values: Tensor | None = None, pixel_image_counts: Sequence[int] | None = None, episode_segments: list[EpisodeSegments | None] | None = None, image_token_id: int | None = None) dict[str, float]¶
Update agent network parameters to learn from experiences.
- Parameters:
experiences (LLMRolloutExperiences) –
(token_ids, action_masks, rewards)stacked batch. Forimportance_sampling_level="turn"with per-turn rewards,rewardsis(batch, max_turns); otherwise it is one scalar per trajectory (per-turn rewards are summed to the episode return).sampling_logps (list[torch.Tensor | None] | None) – Optional per-row flat vLLM sampling logprobs (one 1-D tensor per trajectory, generated tokens only; concatenated across turns for multi-turn) for the sampling-mismatch correction. Parallel to the stacked
token_idsrows.Nonedisables the correction for this update.pixel_values (torch.Tensor | None) – Optional per-sample vision tensors aligned with the stacked
token_idsbatch.turn_ids (torch.Tensor | None) –
(batch, seq_len-1)turn index per action token (-1for non-action tokens), aligned with the action mask. Required when the resolved advantage granularity is"turn"(per-turn group-relative advantages need per-turn rewards). Also consumed by turn-level importance-ratio pooling whenimportance_sampling_level="turn". Ignored when neither applies.episode_segments (list[EpisodeSegments | None] | None) – Optional segment layout per trajectory, parallel to the stacked
token_idsrows (Nonefor an unsegmented trajectory). Each segment trains as its own row with its episode’s advantage. The update keeps the optimizer steps of the unsegmented batch: each step accumulates a window of segment rows, padded so every rank runs the same micro-batches, and every action token of a window weighs the same across ranks unlessloss_norm="episode"weighs every episode the same.image_token_id (int | None) – Token id the VL forward scatters one image feature row into. Required with
pixel_valuesandepisode_segments, to cut the filler rows down to one image.
- Returns:
Dict with averaged
loss,kl(NaN on the fused path atbeta == 0.0),clipfracandcompletion_length(plus per-learn advantage stats, the update-loopentropy/kl_ref/kl_old/is_*diagnostics, averagedgrad_norm_pre/grad_norm_post,learn_phase_<phase>_swall seconds per learn phase, and thevllm_is_*sampling-mismatch metrics when the correction is active).- Return type:
- classmethod load(path: str, device: str | device = 'cpu', accelerator: Accelerator | None = None) Self¶
Load an algorithm from a checkpoint.
- Parameters:
path (string) – Location to load checkpoint from.
device (str, optional) – Device to load the algorithm on, defaults to ‘cpu’
accelerator (Accelerator | None, optional) – Accelerator object for distributed computing, defaults to None
- Returns:
An instance of the algorithm
- Return type:
- load_checkpoint(path: str, load_optimizer: bool = False, overwrite_reference_adapter: bool | None = None, overwrite_critic_adapter: bool = False, restore_config: bool = True, restore_hyperparameters: bool = True) None¶
Load adapter weights and algorithm state from a checkpoint directory.
Adapter roles restored on load:
actor— the trained policy. Always loaded.reference— loaded from the checkpoint’sreference/adapter when it has one; otherwise the checkpoint’sactoris copied ontoreferenceso SFT -> DPO -> GRPO chains work out of the box.critic— loaded from the checkpoint’scritic/adapter when it has one, otherwise left at its fresh LoRA init. Setoverwrite_critic_adapterto seed it from the actor.
The checkpoint’s LoRA config must match the live algorithm’s config; a mismatch raises
ValueError(re-create the agent with the checkpoint’s LoRA config to load it).The same flow applies to plain, DDP and FSDP2 runs:
- lora_only=T -> PEFT adapter dirs are loaded into the live
adapters.
- lora_only=F -> the full actor state_dict is restored from
attributes.pt.
When
load_optimizer=Truethe optimizer state is restored fromattributes.pt; if the checkpoint contains no optimizer state (saved withsave_optimizer=False), aUserWarningis emitted and a freshly-initialised optimizer is used. The LR schedule always resumes at the checkpoint’s learn step, with this instance’slr/lr_critic(after any hyperparameter restore) as its peaks.- Parameters:
path (str) – Directory containing a checkpoint written by
save_checkpoint().load_optimizer (bool) – If
Truealso load the optimizer state so training can resume.overwrite_reference_adapter (bool | None) – Copy the checkpoint’s
actorontoreferenceeven when it has areference/adapter.Nonecopies only when the checkpoint has no reference adapter.overwrite_critic_adapter (bool) – Seed
criticfrom the checkpoint’sactor.restore_config (bool) – If
False, keep this instance’s algorithm settings and restore only training state: step count, scores, fitness,reference_update_tracker,rngand, withrestore_hyperparameters, the registry’s mutable hyperparameters.restore_hyperparameters (bool) – If
False, keep this instance’s values for the registry’s mutable hyperparameters (e.g.lr). Only a run that mutates them needs the checkpoint’s values.
- load_weights(path: str, overwrite_reference_adapter: bool | None = None, overwrite_critic_adapter: bool = False) None¶
Load only the LoRA adapters (and value head) from a checkpoint directory.
- Parameters:
path (str) – Directory containing a checkpoint written by
save_checkpoint().overwrite_reference_adapter (bool | None) – See
load_checkpoint().overwrite_critic_adapter (bool) – See
load_checkpoint().
- classmethod population(size: int, device: str | device = 'cpu', resume_from_checkpoint: str | None = None, **kwargs: Any) list[Self]¶
Create a population of LLM algorithms.
Builds agent 0 fully (loading the model from disk), then clones for agents 1..N. Under FSDP2 / QLoRA this uses adapter-only
clone(); otherwise the actor is copied viaclone_llm().- Parameters:
- Returns:
A list of LLM algorithms.
- Return type:
list[LLMAlgorithm]
- preprocess_observation(observation: Tensor | TensorDict | tuple[Tensor, ...] | dict[str, Tensor]) Tensor | TensorDict | tuple[Tensor, ...] | dict[str, Tensor]¶
Preprocess observations (dummy) for forward pass through neural network.
- recompile() None¶
Recompile evolvable modules with
torch.compile.Iterates over
evolvable_attributesand compiles each one. Skipped for distributed runs, matching_initialize_actors().
- register_mutation_hook(hook: LambdaType | MethodType) None¶
Register a hook to be executed after a mutation is performed on the algorithm.
- Parameters:
hook (MutationHook) – The hook to be executed after mutation.
- register_network_group(group: NetworkGroup) None¶
Set the evaluation network for the algorithm.
- Parameters:
name (str) – The name of the evaluation network.
- reinit_optimizers(optimizer: OptimizerConfig | None = None) None¶
Reinitialize the optimizers of an algorithm. If no optimizer is passed, all optimizers are reinitialized.
- Parameters:
optimizer (OptimizerConfig | None, optional) – The optimizer to reinitialize, defaults to None, in which case all optimizers are reinitialized.
- save_checkpoint(path: str, lora_only: bool = True, save_optimizer: bool = True) None¶
Save adapter weights and algorithm state to a directory.
AgileRL never persists base-model weights when
lora_only=Truefor LLM algorithms: a checkpoint is a directory containing<adapter>/adapter_model.safetensors+adapter_config.json— one subdirectory per adapter inselected_adapters(alwaysactor, plusreference/criticwhen those adapters are configured). Written only whenlora_only=True.attributes.pt— algorithm hyperparameters, plus (optionally) the actor state dict and/or optimizer state dict. Always present.
The same format is written for plain, DDP and FSDP2 runs:
- lora_only=T, save_optimizer=T -> PEFT adapter dirs on disk +
optimizer state in
attributes.pt
lora_only=T, save_optimizer=F -> PEFT adapter dirs only lora_only=F, save_optimizer=T -> full actor state_dict +
optimizer state in
attributes.ptlora_only=F, save_optimizer=F -> full actor state_dict in
attributes.ptFSDP2-sharded parameters and optimizer state are gathered to full tensors before saving, so checkpoints are rank-count independent.
- Parameters:
path (str) – Directory to write the checkpoint into.
lora_only (bool) – If
True(default) only adapter weights are written to disk viasave_pretrained; the base model is shared across checkpoints and not serialised. IfFalse, the full actor state dict is persisted intoattributes.pt.save_optimizer (bool) – If
True(default) also persist the optimizer and LR scheduler state inattributes.ptso training can resume.
- property scores: list[float | list[float]]¶
Per-episode scores (per-group score rows for multi-agent metrics).
- select_adapter(adapter_name: str) Generator[None, None, None]¶
Temporarily switch adapter; restores the actor adapter on exit.
- Parameters:
adapter_name (str) – Name of the adapter to activate (“actor”, “critic”, “reference”).
- set_reference_policy(reference_update_tracker: int) None¶
Update the reference policy when the tracker advances past the stored value.
Base weights are immutable in AgileRL’s LoRA-only training: with
use_separate_reference_adapter=Truethe actor adapter is copied onto thereferenceadapter; without one the implicit reference (the base model with adapters disabled) cannot move, so the update request is acknowledged with a one-time warning and the KL anchor stays the initial policy.- Parameters:
reference_update_tracker (int) – The reference policy update tracker
- set_training_mode(training: bool) None¶
Set the training mode of the algorithm.
- Parameters:
training (bool) – If True, set the algorithm to training mode.
- snapshot_checkpoint(lora_only: bool = True, save_optimizer: bool = True) LLMCheckpointSnapshot¶
Copy the checkpoint of
save_checkpoint()into host memory.All ranks must call this together (FSDP2 gathers are collective). The snapshot holds the weights and optimizer state of this step: later training does not change it, so
LLMCheckpointSnapshot.write()may run on another thread while training continues.- Parameters:
lora_only (bool) – See
save_checkpoint().save_optimizer (bool) – See
save_checkpoint().
- Returns:
Host copy of the checkpoint; empty off the main process.
- Return type:
LLMCheckpointSnapshot
- test(env: RolloutHarness, loop: int = 1, *args: Any, **kwargs: Any) ndarray¶
Return fitness (test) score of the llm on the test sub-set.
- Parameters:
env (RolloutHarness) – Tokenized rollout episode environment (single- or multi-turn).
loop (int) – Number of outer test iterations (episodes).
- Returns:
Zero-dimensional array holding the mean per-step reward, which is also recorded in the agent’s fitness history.
- Return type:
np.ndarray
- to_device(*experiences: Tensor | TensorDict | tuple[Tensor, ...] | dict[str, Tensor]) tuple[Tensor | TensorDict | tuple[Tensor, ...] | dict[str, Tensor], ...]¶
Move experiences to the device.
- unrolled_eval_networks() list[tuple[str | None, Module]]¶
Return the agent’s evaluation networks as (network_id, network) pairs.
- update_existing_adapter(checkpoint_dir: str, adapter_name: str) None¶
Overwrite weights of an existing adapter in-place without creating new parameters.
- Parameters:
checkpoint_dir (str) – Checkpoint directory
adapter_name (str.) – Adapter name
- Returns:
None
- Return type:
None
- static update_lr(optimizer: OptimizerWrapper, lr: float | tuple[float, float | None], scheduler_config: CosineLRScheduleConfig | None = None, schedule_step: int = 0) LambdaLR | None¶
Set the peak learning rate of each param group and rebuild the schedule.
- Parameters:
optimizer (OptimizerWrapper) – LLM optimizer.
lr (float | tuple[float, float | None]) – Learning rate value, or actor/critic pair; a
Nonecritic uses the actor rate.scheduler_config (CosineLRScheduleConfig | None) – Scheduler configuration;
Noneholdslrconstant.schedule_step (int) – Learn steps the schedule has already taken.
- Returns:
Scheduler at
schedule_stepwhenscheduler_configis set.- Return type:
LambdaLR | None
- use_adapter(adapter_name: str) None¶
Switch the active PEFT adapter, handling all side-effects.
For “reference”: switches adapter and freezes reference params (never trained). For all others: switches adapter and restores requires_grad=True on all training adapter LoRA params so distributed gradient hooks keep firing.
- Parameters:
adapter_name (str) – Name of the adapter to activate (“actor”, “critic”, “reference”).
- wrap_models() None¶
Prepare the actor for training.
Places the actor (dense
.to(device)or FSDP2 shard) then builds the optimizer and LR scheduler on those parameters. FSDP2 wraps each transformer block with activation checkpointing beforefully_shardwhengradient_checkpointingis on. Data-parallelprepare_actoruses HuggingFacegradient_checkpointing_enable.