diff --git a/examples/coding_agent_rl/generate.py b/examples/coding_agent_rl/generate.py index 25e8b86a70..0f8f3d1cd8 100644 --- a/examples/coding_agent_rl/generate.py +++ b/examples/coding_agent_rl/generate.py @@ -242,7 +242,7 @@ async def generate(args, base_sample: Sample, sampling_params: dict[str, Any]): ) return _abort_result(base_sample, f"exception:{type(e).__name__}", instance_id) finally: - await state.adapter.finish_session(session_id) # idempotent + await state.adapter.drop_session(session_id) # cleanup only, idempotent def _log_timeout_diagnostic(t0: float, instance_id: str) -> None: @@ -282,7 +282,9 @@ def _abort_result(sample: Sample, reason: str, instance_id: str) -> list[Sample] sample.response = "" sample.response_length = 1 sample.loss_mask = [0] + sample.rollout_log_probs = [0.0] sample.reward = 0.0 + sample.remove_sample = True sample.status = Sample.Status.ABORTED sample.metadata = {**(sample.metadata or {}), "abort_reason": reason} logger.warning("[coding_agent_rl] %s aborted: %s", instance_id, reason) diff --git a/examples/coding_agent_rl/run_qwen36_35b_a3b_swe_8nodes.sh b/examples/coding_agent_rl/run_qwen36_35b_a3b_swe_8nodes.sh index 9b701d33b6..26ec97a234 100755 --- a/examples/coding_agent_rl/run_qwen36_35b_a3b_swe_8nodes.sh +++ b/examples/coding_agent_rl/run_qwen36_35b_a3b_swe_8nodes.sh @@ -152,13 +152,13 @@ PERF_ARGS=( ) ALGO_ARGS=( - --advantage-estimator gspo + --advantage-estimator grpo --kl-loss-coef 0.00 --kl-loss-type low_var_kl --kl-coef 0.00 --entropy-coef 0.00 - --eps-clip 1e-4 - --eps-clip-high 2e-4 + --eps-clip 0.2 + --eps-clip-high 0.28 ) OPTIMIZER_ARGS=( diff --git a/slime/agent/adapters/common.py b/slime/agent/adapters/common.py index ad97a57c9a..fa31c2b0a2 100644 --- a/slime/agent/adapters/common.py +++ b/slime/agent/adapters/common.py @@ -237,7 +237,7 @@ async def finish_session( self, sid: str, *, - base_sample=None, + base_sample, reward: float = 0.0, extra_metadata: dict | None = None, wait_timeout: float = 5.0, @@ -264,6 +264,11 @@ async def finish_session( ) return samples + async def drop_session(self, sid: str, *, wait_timeout: float = 5.0) -> None: + await self.shutdown_session(sid, wait_timeout=wait_timeout) + self.store.pop(sid, None) + self.manager.drop_session(sid) + # -- shared request pipeline --------------------------------------------- def _check_turn_cap(self, sid: str) -> web.Response | None: diff --git a/slime/agent/harness/common.py b/slime/agent/harness/common.py index e4c1921512..fe6850d8b5 100644 --- a/slime/agent/harness/common.py +++ b/slime/agent/harness/common.py @@ -151,8 +151,10 @@ async def run_command(sb: Sandbox, *, workdir: str, start_cmd: str, env: dict[st check=False, ) if ec == 0: - exit_code = int((out or "").strip()) - break + exit_code_text = (out or "").strip() + if exit_code_text: + exit_code = int(exit_code_text) + break return exit_code diff --git a/slime/agent/trajectory.py b/slime/agent/trajectory.py index edb130c648..33ec0366d8 100644 --- a/slime/agent/trajectory.py +++ b/slime/agent/trajectory.py @@ -328,6 +328,10 @@ def get_trajectory( self._turn_count.pop(sid, None) return samples + def drop_session(self, sid: str) -> None: + self._trees.pop(sid, None) + self._turn_count.pop(sid, None) + # -------------------- internals ---------------------------------------- def _find_mount_point(self, root: MessageNode, messages: list[dict[str, Any]]) -> tuple[MessageNode, int]: