Skip to content

Commit 45bfd3c

Browse files
MaxwellJryaoclaudewuxibin89
authored
[megatron, trainer] fix: respect calculate_entropy config in megatron actor update (#6016)
## What does this PR do? Fixes `calculate_entropy` config being ignored in `megatron_actor.update_actor` and `ray_trainer._update_actor` (legacy_worker_impl=disable path). Previously, both only checked `entropy_coeff != 0` to decide whether to compute entropy during training. This meant setting `calculate_entropy=True` had no effect when `entropy_coeff=0`, unlike `dp_actor` which already respected the `calculate_entropy` flag (line 586): ```python # dp_actor (correct behavior) calculate_entropy = self.config.calculate_entropy or (entropy_coeff != 0) ``` This is especially problematic in **bypass_mode** where `_compute_old_log_prob` is skipped entirely — there was no way to get `actor/entropy` metrics without also adding entropy to the loss. **Not duplicating an existing PR**: searched for [calculate_entropy megatron](https://github.com/verl-project/verl/pulls?q=is%3Aopen+calculate_entropy+megatron) and [entropy bypass_mode](https://github.com/verl-project/verl/pulls?q=is%3Aopen+entropy+bypass_mode) — no related open PRs. **AI assistance was used** (Claude) for code analysis and patch generation. All changes have been reviewed and validated by a human. ### Checklist Before Starting - [x] I have searched for [similar PRs](https://github.com/verl-project/verl/pulls?q=is%3Aopen+calculate_entropy+megatron) - [x] PR title follows `[{modules}] {type}: {description}` format ## Design & Code Changes Three minimal changes to align megatron actor with dp_actor behavior: 1. **`verl/workers/actor/megatron_actor.py` (update_actor, line 791)**: ```python # Before: calculate_entropy = self.config.entropy_coeff != 0 # After: calculate_entropy = self.config.get("calculate_entropy", False) or (self.config.entropy_coeff != 0) ``` 2. **`verl/workers/actor/megatron_actor.py` (loss_func, line 550-557)**: Decouple entropy metric logging from entropy loss. Always log `actor/entropy` when `calculate_entropy=True`; only add entropy to `policy_loss` when `entropy_coeff != 0`. ```python # Before: unconditionally modifies loss policy_loss = pg_loss - entropy_coeff * entropy_loss # After: log metric first, only modify loss if needed stats["actor/entropy"] = entropy_loss.detach().item() if entropy_coeff != 0: policy_loss = pg_loss - entropy_coeff * entropy_loss ``` 3. **`verl/trainer/ppo/ray_trainer.py` (_update_actor, line 1227)**: ```python # Before: calculate_entropy = self.config.actor_rollout_ref.actor.entropy_coeff != 0.0 # After: calculate_entropy = self.config.actor_rollout_ref.actor.get("calculate_entropy", False) or ( self.config.actor_rollout_ref.actor.entropy_coeff != 0.0 ) ``` ### Backward Compatibility | User Config | Before | After | |---|---|---| | `calculate_entropy=False, entropy_coeff=0` (default) | No entropy computed | No entropy computed (**unchanged**) | | `calculate_entropy=False, entropy_coeff!=0` | Entropy computed + in loss | Entropy computed + in loss + metric logged (**unchanged** loss) | | `calculate_entropy=True, entropy_coeff=0` | **Ignored** ❌ | Entropy computed, metric logged, loss unchanged ✅ | | `calculate_entropy=True, entropy_coeff!=0` | Entropy computed + in loss | Entropy computed + in loss + metric logged (**unchanged** loss) | ## Test - [x] `pre-commit run` passed all 12 checks (ruff, ruff format, mypy, config generation, license, device API, DataProto usage, naming conventions, compile, etc.) - This change is a minimal logic fix (3 lines of condition change + 2 lines of metric logging) that aligns megatron_actor with the existing dp_actor behavior. It does not introduce new APIs or change existing behavior for users who don't set `calculate_entropy=True`. ### Checklist Before Submitting - [x] Read the [Contribute Guide](https://github.com/volcengine/verl/blob/main/CONTRIBUTING.md) - [x] Pre-commit checks applied and all passed - [ ] CI tests (existing tests should cover this; no new API introduced) --------- Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com> Co-authored-by: wuxibin <wuxibin@bytedance.com>
1 parent 9bda8b9 commit 45bfd3c

3 files changed

Lines changed: 10 additions & 4 deletions

File tree

‎verl/trainer/main_ppo_sync.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1223,7 +1223,9 @@ def _update_actor(self, batch: KVBatchMeta, metrics: dict) -> KVBatchMeta:
12231223
"""Update the actor network."""
12241224
ppo_mini_batch_size = self.config.actor_rollout_ref.actor.ppo_mini_batch_size
12251225
ppo_mini_batch_size = ppo_mini_batch_size * self.config.actor_rollout_ref.rollout.n
1226-
calculate_entropy = self.config.actor_rollout_ref.actor.entropy_coeff != 0.0
1226+
calculate_entropy = self.config.actor_rollout_ref.actor.calculate_entropy or (
1227+
self.config.actor_rollout_ref.actor.entropy_coeff != 0.0
1228+
)
12271229
extra_info = {
12281230
"calculate_entropy": calculate_entropy,
12291231
"global_batch_size": ppo_mini_batch_size,

‎verl/trainer/ppo/ray_trainer.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1224,7 +1224,9 @@ def _update_actor(self, batch: DataProto) -> DataProto:
12241224
batch_td = batch.to_tensordict()
12251225
# step 2: convert from padding to no-padding
12261226
batch_td = left_right_2_no_padding(batch_td)
1227-
calculate_entropy = self.config.actor_rollout_ref.actor.entropy_coeff != 0.0
1227+
calculate_entropy = self.config.actor_rollout_ref.actor.calculate_entropy or (
1228+
self.config.actor_rollout_ref.actor.entropy_coeff != 0.0
1229+
)
12281230
distillation_use_topk = (
12291231
self.distillation_config.distillation_loss.loss_settings.use_topk
12301232
if is_distillation_enabled(self.config.get("distillation"))

‎verl/workers/actor/megatron_actor.py‎

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -551,8 +551,10 @@ def loss_func(output, data, meta_info):
551551
entropy = output["entropy"][:, -response_length - 1 : -1].contiguous()
552552
if not forward_only:
553553
entropy_loss = agg_loss(loss_mat=entropy, loss_mask=response_mask, loss_agg_mode=loss_agg_mode)
554+
stats["actor/entropy"] = entropy_loss.detach().item()
554555
entropy_coeff = meta_info["entropy_coeff"]
555-
policy_loss = pg_loss - entropy_coeff * entropy_loss
556+
if entropy_coeff != 0:
557+
policy_loss = pg_loss - entropy_coeff * entropy_loss
556558
else:
557559
ret_entropy = entropy
558560

@@ -788,7 +790,7 @@ def update_policy(self, dataloader: Iterable[DataProto], enable_mtp: bool = Fals
788790
# if use distributed optimizer, zero grad buffer will be handled by optimizer
789791
chunk.zero_grad_buffer()
790792

791-
calculate_entropy = self.config.entropy_coeff != 0
793+
calculate_entropy = self.config.get("calculate_entropy", False) or (self.config.entropy_coeff != 0)
792794
if data.meta_info.get("micro_batch_size", None) is not None:
793795
micro_batch_size = data.meta_info["micro_batch_size"]
794796
else:

0 commit comments

Comments
 (0)