From f96b559f6795fb04b4aa6617071495dc63fd5e24 Mon Sep 17 00:00:00 2001 From: fzyzcjy Date: Fri, 24 Oct 2025 20:43:56 +0800 Subject: [PATCH] cp --- miles/backends/fsdp_utils/actor.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/miles/backends/fsdp_utils/actor.py b/miles/backends/fsdp_utils/actor.py index e7746410aeb..c700f42b14b 100644 --- a/miles/backends/fsdp_utils/actor.py +++ b/miles/backends/fsdp_utils/actor.py @@ -430,6 +430,9 @@ def train(self, rollout_id: int, rollout_data_ref: Box) -> None: pg_loss, pg_clipfrac = compute_policy_loss(ppo_kl, advantages, self.args.eps_clip, self.args.eps_clip_high) + rollout_log_probs = torch.cat([batch["rollout_log_probs"] for batch in unpacked_batches], dim=0) + rollout_log_probs = rollout_log_probs.to(device=log_probs.device) + # Apply TIS before sample mean calculation if self.args.use_tis: # Initialize TIS variables @@ -444,9 +447,6 @@ def train(self, rollout_id: int, rollout_data_ref: Box) -> None: for batch in unpacked_batches ), "rollout_log_probs must be provided as non-empty torch.Tensor for TIS" - rollout_log_probs = torch.cat([batch["rollout_log_probs"] for batch in unpacked_batches], dim=0) - rollout_log_probs = rollout_log_probs.to(device=log_probs.device) - tis = torch.exp(old_log_probs - rollout_log_probs) ois = (-ppo_kl).exp() tis_clip = torch.clamp(