nemo_rl.algorithms.ppo#
Module Contents#
Classes#
Configuration for asynchronous PPO training. |
|
Functions#
Apply optional mismatch masking and return the advantage mask and metrics. |
|
Main entry point for running PPO algorithm. |
|
Select complete prompt groups with non-trivial reward distributions. |
|
Create and return an advantage estimator based on configuration. |
|
Aggregate value-model metrics under the |
|
Run PPO training algorithm. |
|
Return the collector lead without crossing the safe warmup frontier. |
|
Keep frozen-policy rollouts valid through their safe training frontier. |
|
Run PPO while a background collector fills a replay buffer. |
|
Run validation on the validation dataset. |
Data#
API#
- nemo_rl.algorithms.ppo.TokenizerType#
‘TypeVar(…)’
- class nemo_rl.algorithms.ppo.AsyncPPOConfig#
Bases:
pydantic.BaseModelConfiguration for asynchronous PPO training.
- enabled: bool#
False
- max_trajectory_age_steps: int#
‘Field(…)’
- warmup_generation_lead_steps: int | None#
‘Field(…)’
- in_flight_weight_updates: bool#
False
- recompute_kv_cache_after_weight_updates: bool#
False
- drop_incomplete_targets_on_restore: bool#
False
- validate_settings() nemo_rl.algorithms.ppo.AsyncPPOConfig#
- property resolved_warmup_generation_lead_steps: int#
Resolve the optional warmup generation lead.
- class nemo_rl.algorithms.ppo.PPOConfig#
Bases:
pydantic.BaseModel- num_prompts_per_step: int#
32
- num_generations_per_prompt: int#
16
- max_num_epochs: int#
100000
- max_num_steps: int#
100000
- max_rollout_turns: int#
1
- val_period: int#
20
- val_batch_size: int#
256
- val_at_start: bool#
True
- val_at_end: bool#
False
- max_val_samples: int#
256
- skip_reference_policy_logprobs_calculation: bool#
True
- seed: int#
42
- overlong_filtering: bool#
False
- use_dynamic_sampling: bool#
False
- dynamic_sampling_max_gen_batches: int#
10
- batch_multiplier: float#
1.0
- ppo_epochs: int#
4
- critic_ppo_epochs: int#
4
- reward_shaping: nemo_rl.algorithms.reward_functions.RewardShapingConfig#
‘Field(…)’
- reward_scaling: nemo_rl.algorithms.grpo.RewardScalingConfig#
‘Field(…)’
- adv_estimator: nemo_rl.algorithms.advantage_estimator.GAEConfig#
‘Field(…)’
- policy_training_start_step: int#
0
- warm_start_value_checkpoint: str | None#
None
- seq_logprob_error_threshold: float | None#
None
- invalid_tool_call_advantage: float | None#
None
- malformed_thinking_advantage: float | None#
None
- async_ppo: nemo_rl.algorithms.ppo.AsyncPPOConfig | None#
‘Field(…)’
- validate_epoch() nemo_rl.algorithms.ppo.PPOConfig#
- validate_async_warmup() nemo_rl.algorithms.ppo.PPOConfig#
- class nemo_rl.algorithms.ppo.PPOSaveState#
Bases:
typing.TypedDict- consumed_samples: int#
None
- current_step: int#
None
- current_epoch: int#
None
- total_steps: int#
None
- total_valid_tokens: int#
None
- val_reward: NotRequired[float]#
None
- nemo_rl.algorithms.ppo._default_ppo_save_state() nemo_rl.algorithms.ppo.PPOSaveState#
- nemo_rl.algorithms.ppo._apply_ppo_seq_logprob_error_masking(
- train_data: nemo_rl.distributed.batched_data_dict.BatchedDataDict,
- rewards: torch.Tensor,
- seq_logprob_error_threshold: float | None,
Apply optional mismatch masking and return the advantage mask and metrics.
- class nemo_rl.algorithms.ppo.MasterConfig#
Bases:
pydantic.BaseModel- policy: nemo_rl.models.policy.PolicyConfig#
None
- value: nemo_rl.models.value.ValueConfig#
None
- loss_fn: nemo_rl.algorithms.loss.ClippedPGLossConfig#
None
- value_loss_fn: nemo_rl.algorithms.loss.loss_functions.MseValueLossConfig#
None
- env: dict[str, Any]#
None
- data: nemo_rl.data.DataConfig#
None
- ppo: nemo_rl.algorithms.ppo.PPOConfig#
None
- logger: nemo_rl.utils.logger.LoggerConfig#
None
- cluster: nemo_rl.distributed.virtual_cluster.ClusterConfig#
None
- checkpointing: nemo_rl.utils.checkpoint.CheckpointingConfig#
None
- telemetry: Optional[nemo_rl.telemetry.config.TelemetryConfig]#
None
- nemo_rl.algorithms.ppo.setup(
- master_config: nemo_rl.algorithms.ppo.MasterConfig,
- tokenizer: nemo_rl.algorithms.ppo.TokenizerType,
- dataset: nemo_rl.data.datasets.AllTaskProcessedDataset,
- val_dataset: Optional[nemo_rl.data.datasets.AllTaskProcessedDataset],
- processor: Optional[transformers.AutoProcessor] = None,
Main entry point for running PPO algorithm.
- Returns:
tuple of (policy, policy_generation, value_model, clusters, dataloader, val_dataloader, loss_fn, value_loss_fn, logger, checkpointer, ppo_save_state, master_config).
- nemo_rl.algorithms.ppo.dynamic_sampling(
- repeated_batch: nemo_rl.distributed.batched_data_dict.BatchedDataDict[nemo_rl.data.interfaces.DatumSpec],
- std: torch.Tensor,
- baseline: torch.Tensor,
- dynamic_sampling_num_gen_batches: int,
- master_config: nemo_rl.algorithms.ppo.MasterConfig,
- timer: nemo_rl.utils.timer.Timer,
- batch_cache: nemo_rl.distributed.batched_data_dict.BatchedDataDict[nemo_rl.data.interfaces.DatumSpec] = None,
- is_trivial_prompt_distribution: torch.Tensor | None = None,
Select complete prompt groups with non-trivial reward distributions.
Exact reward equality determines triviality, independently of floating-point standard-deviation noise. Every rollout for a prompt is kept or discarded together. If the current batch has fewer non-trivial prompt groups than the required batch size, defined as num_prompts_per_step * num_generations_per_prompt, we store it in the batch_cache to be used in later iterations. If the current batch has more non-trivial prompt groups than the required batch size, the batch is sliced to ensure batch size is num_prompts_per_step * num_generations_per_prompt. is_batch_complete is set to False to indicate that the current batch is not enough to meet the required batch size. This is used as a signal in the training loop to continue sampling or proceed to training. This approach is based on the dynamic sampling algorithm from the DAPO paper: https://arxiv.org/pdf/2503.14476.
- Parameters:
repeated_batch (BatchedDataDict[DatumSpec]) – The current batch of data containing prompts, responses, rewards, baselines, and std.
std (torch.Tensor) – Tensor representing the standard deviation for each prompt group.
baseline (torch.Tensor) – Baseline values for each prompt group.
dynamic_sampling_num_gen_batches (int) – Number of generation batches processed at the current step.
master_config (MasterConfig) – Configuration containing PPO and policy settings.
batch_cache (BatchedDataDict[DatumSpec], optional) – Cache storing previously selected non-trivial prompt groups.
is_trivial_prompt_distribution (torch.Tensor, optional) – Exact-equality mask for each sample’s full prompt reward group. Trivial groups are filtered all-or-nothing.
- Returns:
A tuple containing: - repeated_batch (BatchedDataDict[DatumSpec]): Updated batch with selected prompts. - is_batch_complete (bool): Indicates if the batch has enough non-trivial samples for training. - batch_cache (BatchedDataDict[DatumSpec]): Updated cache for future iterations.
- Return type:
tuple
- nemo_rl.algorithms.ppo._create_advantage_estimator(
- master_config: nemo_rl.algorithms.ppo.MasterConfig,
Create and return an advantage estimator based on configuration.
PPO’s training loop consumes a
(advantages, returns)pair from a value-model-based estimator, so onlygaeandraw_rewardare supported here. Group-relative estimators like GRPO / Reinforce++ are not compatible with PPO’s loop and live ingrpo.py.- Parameters:
master_config – The master configuration dictionary.
- Returns:
A
GeneralizedAdvantageEstimatororRawRewardAdvantageEstimatorinstance.- Raises:
ValueError – If the advantage estimator name is not recognized.
- nemo_rl.algorithms.ppo.CRITIC_LOSS_KEY#
‘critic/loss’
- nemo_rl.algorithms.ppo.CRITIC_TEED_METRICS#
()
- nemo_rl.algorithms.ppo._compute_critic_metrics(
- value_results: dict[str, Any],
Aggregate value-model metrics under the
critic/namespace.
- nemo_rl.algorithms.ppo.ppo_train(
- policy: nemo_rl.models.policy.interfaces.ColocatablePolicyInterface,
- policy_generation: Optional[nemo_rl.models.generation.interfaces.GenerationInterface],
- value_model: nemo_rl.models.value.interfaces.ValueInterface,
- dataloader: torchdata.stateful_dataloader.StatefulDataLoader,
- val_dataloader: Optional[torchdata.stateful_dataloader.StatefulDataLoader],
- tokenizer: nemo_rl.algorithms.ppo.TokenizerType,
- loss_fn: nemo_rl.algorithms.loss.interfaces.LossFunction,
- value_loss_fn: nemo_rl.algorithms.loss.interfaces.LossFunction,
- task_to_env: dict[str, nemo_rl.environments.interfaces.EnvironmentInterface],
- val_task_to_env: Optional[dict[str, nemo_rl.environments.interfaces.EnvironmentInterface]],
- logger: nemo_rl.utils.logger.Logger,
- checkpointer: nemo_rl.utils.checkpoint.CheckpointManager,
- ppo_save_state: nemo_rl.algorithms.ppo.PPOSaveState,
- master_config: nemo_rl.algorithms.ppo.MasterConfig,
Run PPO training algorithm.
Based on the grpo_train loop with PPO-specific modifications:
Value model inference and training (actor-critic)
GAE advantage estimation with value bootstrap
Multiple actor and critic training steps per rollout
Configurable policy training start epoch
- nemo_rl.algorithms.ppo._async_ppo_generation_lead_steps(
- *,
- step: int,
- policy_training_start_step: int,
- max_trajectory_age_steps: int,
- warmup_generation_lead_steps: int,
Return the collector lead without crossing the safe warmup frontier.
- nemo_rl.algorithms.ppo._async_ppo_buffer_max_age(
- *,
- step: int,
- policy_training_start_step: int,
- max_trajectory_age_steps: int,
- warmup_generation_lead_steps: int,
Keep frozen-policy rollouts valid through their safe training frontier.
- nemo_rl.algorithms.ppo.async_ppo_train(
- policy: nemo_rl.models.policy.interfaces.ColocatablePolicyInterface,
- policy_generation: Optional[nemo_rl.models.generation.interfaces.GenerationInterface],
- value_model: nemo_rl.models.value.interfaces.ValueInterface,
- dataloader: torchdata.stateful_dataloader.StatefulDataLoader,
- val_dataloader: Optional[torchdata.stateful_dataloader.StatefulDataLoader],
- tokenizer: nemo_rl.algorithms.ppo.TokenizerType,
- loss_fn: nemo_rl.algorithms.loss.interfaces.LossFunction,
- value_loss_fn: nemo_rl.algorithms.loss.interfaces.LossFunction,
- task_to_env: dict[str, nemo_rl.environments.interfaces.EnvironmentInterface],
- val_task_to_env: Optional[dict[str, nemo_rl.environments.interfaces.EnvironmentInterface]],
- logger: nemo_rl.utils.logger.Logger,
- checkpointer: nemo_rl.utils.checkpoint.CheckpointManager,
- ppo_save_state: nemo_rl.algorithms.ppo.PPOSaveState,
- master_config: nemo_rl.algorithms.ppo.MasterConfig,
Run PPO while a background collector fills a replay buffer.
- nemo_rl.algorithms.ppo.validate(
- policy_generation: nemo_rl.models.generation.interfaces.GenerationInterface,
- val_dataloader: Optional[torchdata.stateful_dataloader.StatefulDataLoader],
- tokenizer,
- val_task_to_env: Optional[dict[str, nemo_rl.environments.interfaces.EnvironmentInterface]],
- step: int,
- master_config: nemo_rl.algorithms.ppo.MasterConfig,
- logger: Optional[nemo_rl.utils.logger.Logger] = None,
Run validation on the validation dataset.