nemo_rl.algorithms.ppo#

Module Contents#

Classes#

AsyncPPOConfig

Configuration for asynchronous PPO training.

PPOConfig

PPOSaveState

MasterConfig

Functions#

_default_ppo_save_state

_apply_ppo_seq_logprob_error_masking

Apply optional mismatch masking and return the advantage mask and metrics.

setup

Main entry point for running PPO algorithm.

dynamic_sampling

Select complete prompt groups with non-trivial reward distributions.

_create_advantage_estimator

Create and return an advantage estimator based on configuration.

_compute_critic_metrics

Aggregate value-model metrics under the critic/ namespace.

ppo_train

Run PPO training algorithm.

_async_ppo_generation_lead_steps

Return the collector lead without crossing the safe warmup frontier.

_async_ppo_buffer_max_age

Keep frozen-policy rollouts valid through their safe training frontier.

async_ppo_train

Run PPO while a background collector fills a replay buffer.

validate

Run validation on the validation dataset.

Data#

API#

nemo_rl.algorithms.ppo.TokenizerType#

‘TypeVar(…)’

class nemo_rl.algorithms.ppo.AsyncPPOConfig#

Bases: pydantic.BaseModel

Configuration 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,
) → tuple[torch.Tensor, dict[str, float | int]]#

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,
) → tuple[nemo_rl.models.policy.interfaces.ColocatablePolicyInterface, Optional[nemo_rl.models.generation.interfaces.GenerationInterface], nemo_rl.models.value.interfaces.ValueInterface, tuple[nemo_rl.distributed.virtual_cluster.RayVirtualCluster, nemo_rl.distributed.virtual_cluster.RayVirtualCluster], torchdata.stateful_dataloader.StatefulDataLoader, Optional[torchdata.stateful_dataloader.StatefulDataLoader], nemo_rl.algorithms.loss.ClippedPGLossFn, nemo_rl.algorithms.loss.loss_functions.MseValueLossFn, nemo_rl.utils.logger.Logger, nemo_rl.utils.checkpoint.CheckpointManager, nemo_rl.algorithms.ppo.PPOSaveState, nemo_rl.algorithms.ppo.MasterConfig]#

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,
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[nemo_rl.data.interfaces.DatumSpec]#

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 only gae and raw_reward are supported here. Group-relative estimators like GRPO / Reinforce++ are not compatible with PPO’s loop and live in grpo.py.

Parameters:

master_config – The master configuration dictionary.

Returns:

A GeneralizedAdvantageEstimator or RawRewardAdvantageEstimator instance.

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],
) → 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,
) → None#

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,
) → 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,
) → 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,
) → None#

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,
) → tuple[dict[str, Any], dict[str, Any]]#

Run validation on the validation dataset.