diff --git a/tests/experimental/examples/deepswe_dist/deepswe_test.py b/tests/experimental/examples/deepswe_dist/deepswe_test.py new file mode 100644 index 000000000..aba9986da --- /dev/null +++ b/tests/experimental/examples/deepswe_dist/deepswe_test.py @@ -0,0 +1,218 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from pathlib import Path +from unittest import mock + +from absl.testing import absltest +from examples.deepswe import deepswe_data +from examples.deepswe import swe_env +from tunix.experimental.examples.deepswe_dist import deepswe +from tunix.experimental.rl.agentic import registry + + +class DeepSWEDistTest(absltest.TestCase): + + def test_dataset_loader_reuses_recipe_implementation(self): + dataset = object() + with mock.patch.object( + deepswe_data, "create_dataset", return_value=dataset + ) as mock_create_dataset: + result = deepswe.load_deepswe_dataset( + dataset_name="custom/deepswe", + dataset_split="validation", + cache_dir="/tmp/deepswe-cache", + shuffle=True, + seed=123, + ) + + self.assertIs(result, dataset) + mock_create_dataset.assert_called_once_with( + dataset_name="custom/deepswe", + dataset_split="validation", + cache_dir="/tmp/deepswe-cache", + shuffle=True, + seed=123, + ) + + def test_distributed_deepswe_reuses_recipe_data_agent_and_env(self): + package_dir = Path(deepswe.__file__).parent + self.assertFalse((package_dir / "swe_agent.py").exists()) + self.assertFalse((package_dir / "swe_env.py").exists()) + self.assertFalse((package_dir / "deepswe_data.py").exists()) + self.assertIs(deepswe.deepswe_data, deepswe_data) + self.assertIs(deepswe.swe_env, swe_env) + self.assertTrue(issubclass(deepswe.DeepSWEEnv, swe_env.SWEEnv)) + source = Path(deepswe.__file__).read_text(encoding="utf-8") + self.assertIn("from examples.deepswe import swe_agent", source) + + def test_registered_components(self): + self.assertIs( + registry.ENV_REGISTRY.get(deepswe.DEEPSWE_ENV_NAME), deepswe.DeepSWEEnv + ) + self.assertIs( + registry.AGENT_REGISTRY.get(deepswe.DEEPSWE_AGENT_NAME), + deepswe.DeepSWEAgent, + ) + + def test_build_prompt_item_carries_env_and_agent_config(self): + item = deepswe.build_prompt_item( + entry={ + "instance_id": "repo__issue-1", + "problem_statement": "Fix this bug.", + }, + prompt_idx=0, + max_turns=3, + max_response_length=128, + episode_timeout_secs=300, + temperature=0.7, + top_p=0.9, + top_k=20, + step_timeout_secs=30, + reward_timeout_secs=40, + overlong_filter=True, + env_backend="kubernetes", + use_agent_sandbox=True, + batch_size=8, + num_generations=8, + max_warmpool_replicas=6, + scaffold="r2egym", + env_verbose=True, + ) + + self.assertEqual(item["prompt_id"], "repo__issue-1__0") + self.assertEqual(item["prompt"], "Fix this bug.") + self.assertEqual(item["generation_kwargs"]["max_generation_steps"], 128) + self.assertEqual(item["generation_kwargs"]["max_response_length"], 128) + self.assertEqual(item["metadata"]["instance_id"], "repo__issue-1") + self.assertEqual(item["metadata"]["prefix_hash"], "repo__issue-1") + self.assertEqual(item["metadata"]["episode_timeout"], 300) + self.assertTrue(item["metadata"]["overlong_filter"]) + env_config = item["metadata"]["env_config"] + self.assertEqual(env_config["entry"]["instance_id"], "repo__issue-1") + self.assertEqual(env_config["backend"], "kubernetes") + self.assertTrue(env_config["use_agent_sandbox"]) + self.assertEqual(env_config["batch_size"], 8) + self.assertEqual(env_config["group_size"], 8) + self.assertEqual(env_config["max_warmpool_replicas"], 6) + self.assertEqual(item["metadata"]["agent_config"], {"scaffold": "r2egym"}) + + def test_iter_prompt_items_recycles_dataset(self): + dataset = [ + {"instance_id": "task-1", "problem_statement": "first"}, + {"instance_id": "task-2", "problem_statement": "second"}, + ] + + items = list( + deepswe.iter_prompt_items( + dataset=dataset, + max_steps=2, + batch_size=2, + max_turns=3, + max_response_length=128, + episode_timeout_secs=300, + temperature=1.0, + top_p=1.0, + top_k=None, + step_timeout_secs=30, + reward_timeout_secs=40, + overlong_filter=True, + env_backend="kubernetes", + use_agent_sandbox=True, + num_generations=2, + max_warmpool_replicas=None, + scaffold="r2egym", + env_verbose=False, + ) + ) + + self.assertEqual( + [item["prompt_id"] for item in items], + ["task-1__0", "task-2__1", "task-1__2", "task-2__3"], + ) + self.assertLen({item["prompt_id"] for item in items}, 4) + + def test_sandbox_fleet_loads_full_dataset_from_env(self): + dataset = [ + {"instance_id": "task-1", "problem_statement": "first"}, + {"instance_id": "task-2", "problem_statement": "second"}, + ] + with mock.patch.dict( + "os.environ", + { + "DATASET_NAME": "custom/deepswe", + "DATASET_SPLIT": "validation", + "DATASET_CACHE_DIR": "/tmp/deepswe-cache", + "SHUFFLE": "false", + "SEED": "123", + "SANDBOX_MAX_CONCURRENCY": "7", + }, + ): + with mock.patch.object( + deepswe, "load_deepswe_dataset", return_value=dataset + ) as mock_load: + with mock.patch.object( + deepswe.swe_env, + "_get_global_fleet", + side_effect=RuntimeError("not initialized"), + ): + with mock.patch.object( + deepswe.swe_env, "_init_global_fleet", return_value="fleet" + ) as mock_init: + # pylint: disable=protected-access + fleet = deepswe._init_sandbox_fleet_from_env( + {"instance_id": "fallback"}, + group_size=2, + batch_size=3, + max_warmpool_replicas=5, + ) + # pylint: enable=protected-access + + self.assertEqual(fleet, "fleet") + mock_load.assert_called_once_with( + dataset_name="custom/deepswe", + dataset_split="validation", + dataset_path="", + cache_dir="/tmp/deepswe-cache", + shuffle=False, + seed=123, + ) + mock_init.assert_called_once_with( + tasks=dataset, + max_concurrency=7, + num_generations=2, + batch_size=3, + max_warmpool_replicas=5, + ) + + def test_sandbox_fleet_reuses_process_global_before_loading_dataset(self): + with mock.patch.object( + deepswe.swe_env, "_get_global_fleet", return_value="existing" + ): + with mock.patch.object(deepswe, "load_deepswe_dataset") as mock_load: + # pylint: disable=protected-access + fleet = deepswe._init_sandbox_fleet_from_env( + {"instance_id": "unused"}, + group_size=8, + batch_size=8, + max_warmpool_replicas=None, + ) + # pylint: enable=protected-access + + self.assertEqual(fleet, "existing") + mock_load.assert_not_called() + + +if __name__ == "__main__": + absltest.main() diff --git a/tunix/experimental/examples/deepswe_dist/README.md b/tunix/experimental/examples/deepswe_dist/README.md index 27ada32c5..774da1bbc 100644 --- a/tunix/experimental/examples/deepswe_dist/README.md +++ b/tunix/experimental/examples/deepswe_dist/README.md @@ -1,12 +1,14 @@ # Distributed DeepSWE GRPO Pipeline -This example is the first DeepSWE-specific version of the experimental -distributed RL pipeline. It follows the same control-plane shape as the -distributed GSM8K example: +This example ports the non-experimental `examples/deepswe` recipe to the +experimental distributed RL control plane. It reuses the recipe's +`examples/deepswe/deepswe_data.py`, `examples/deepswe/swe_agent.py`, and +`examples/deepswe/swe_env.py` directly; only the distributed registry and +request wiring, orchestration, and launchers live in this directory. 1. `run_deepswe_dist.py` runs the CPU orchestrator. -2. `../common/run_rollout_node.py` runs a rollout worker configured with - DeepSWE's `SWEEnv` and `SWEAgent`. +2. `../common/run_rollout_node.py` runs a rollout worker configured with the + original recipe's `swe_env.SWEEnv` and `swe_agent.SWEAgent` implementations. 3. The trainer worker is reused from `../common/run_trainer_node.py` because it is already a generic PeftTrainer V2 worker. diff --git a/tunix/experimental/examples/deepswe_dist/deepswe.py b/tunix/experimental/examples/deepswe_dist/deepswe.py index 882d8ca76..74f05d6ec 100644 --- a/tunix/experimental/examples/deepswe_dist/deepswe.py +++ b/tunix/experimental/examples/deepswe_dist/deepswe.py @@ -19,19 +19,19 @@ from collections.abc import Iterator import json import logging +import os +import threading from typing import Any -import numpy as np -from tunix.experimental.rl.agentic import registry - from examples.deepswe import deepswe_data -from examples.deepswe import swe_agent from examples.deepswe import swe_env - +import numpy as np +from tunix.experimental.rl.agentic import registry DEEPSWE_ENV_NAME = "deepswe_env" DEEPSWE_AGENT_NAME = "deepswe_agent" -DEFAULT_DATASET_NAME = "R2E-Gym/R2E-Gym-Subset" +DEFAULT_DATASET_NAME = "R2E-Gym/R2E-Gym-V1" +_SANDBOX_INIT_LOCK = threading.Lock() def normalize_example_value(value: Any) -> Any: @@ -56,7 +56,7 @@ def as_text(value: Any) -> str: def _jsonify_lists(entry: dict[str, Any]) -> dict[str, Any]: - """Matches the legacy DeepSWE recipe's heterogeneous dataset normalization.""" + """Normalizes heterogeneous DeepSWE dataset values for distributed rollout.""" normalized = {} for key, value in entry.items(): value = normalize_example_value(value) @@ -82,7 +82,6 @@ def load_deepswe_dataset( dataset = load_from_disk(dataset_path) if isinstance(dataset, DatasetDict): dataset = dataset[dataset_split] - dataset = dataset.map(_jsonify_lists, keep_in_memory=True) if shuffle: dataset = dataset.shuffle(seed=seed) return dataset @@ -124,19 +123,25 @@ def build_prompt_item( prompt_idx: int, max_turns: int, max_response_length: int, + episode_timeout_secs: int, temperature: float, top_p: float | None, top_k: int | None, step_timeout_secs: int, reward_timeout_secs: int, + overlong_filter: bool, env_backend: str, use_agent_sandbox: bool, + batch_size: int, + num_generations: int, + max_warmpool_replicas: int | None, scaffold: str, env_verbose: bool, ) -> dict[str, Any]: """Builds one StandardRLProgram prompt item for a DeepSWE task.""" problem = _problem_statement(entry) - prompt_id = as_text(entry.get("instance_id") or f"deepswe_{prompt_idx}") + instance_id = as_text(entry.get("instance_id") or "deepswe") + prompt_id = f"{instance_id}__{prompt_idx}" env_config = { "entry": entry, "prompt_id": prompt_id, @@ -145,6 +150,9 @@ def build_prompt_item( "reward_timeout": reward_timeout_secs, "backend": env_backend, "use_agent_sandbox": use_agent_sandbox, + "batch_size": batch_size, + "group_size": num_generations, + "max_warmpool_replicas": max_warmpool_replicas, "scaffold": scaffold, "verbose": env_verbose, } @@ -155,15 +163,20 @@ def build_prompt_item( "max_turns": max_turns, "generation_kwargs": { "max_generation_steps": max_response_length, + # The collector uses this as the episode-wide token budget, while + # max_generation_steps is the per-call sampler limit. + "max_response_length": max_response_length, "temperature": temperature, "top_p": top_p, "top_k": top_k, "return_logprobs": True, }, "metadata": { - "instance_id": prompt_id, + "instance_id": instance_id, "problem_statement": problem, - "prefix_hash": prompt_id, + "prefix_hash": instance_id, + "episode_timeout": episode_timeout_secs, + "overlong_filter": overlong_filter, "env_config": env_config, "agent_config": agent_config, }, @@ -177,13 +190,17 @@ def iter_prompt_items( batch_size: int, max_turns: int, max_response_length: int, + episode_timeout_secs: int, temperature: float, top_p: float | None, top_k: int | None, step_timeout_secs: int, reward_timeout_secs: int, + overlong_filter: bool, env_backend: str, use_agent_sandbox: bool, + num_generations: int, + max_warmpool_replicas: int | None, scaffold: str, env_verbose: bool, ) -> Iterator[dict[str, Any]]: @@ -198,18 +215,93 @@ def iter_prompt_items( prompt_idx=prompt_idx, max_turns=max_turns, max_response_length=max_response_length, + episode_timeout_secs=episode_timeout_secs, temperature=temperature, top_p=top_p, top_k=top_k, step_timeout_secs=step_timeout_secs, reward_timeout_secs=reward_timeout_secs, + overlong_filter=overlong_filter, env_backend=env_backend, use_agent_sandbox=use_agent_sandbox, + batch_size=batch_size, + num_generations=num_generations, + max_warmpool_replicas=max_warmpool_replicas, scaffold=scaffold, env_verbose=env_verbose, ) +def _env_bool(name: str, default: bool) -> bool: + value = os.getenv(name) + if value is None: + return default + return value.lower() not in ("0", "false", "no", "off") + + +def _env_int(name: str, default: int) -> int: + value = os.getenv(name) + if value is None or value == "": + return default + return int(value) + + +def _sandbox_tasks_from_env() -> list[dict[str, Any]]: + dataset = load_deepswe_dataset( + dataset_name=os.getenv("DATASET_NAME", DEFAULT_DATASET_NAME), + dataset_split=os.getenv("DATASET_SPLIT", "train"), + dataset_path=os.getenv("DATASET_PATH", ""), + cache_dir=os.getenv("DATASET_CACHE_DIR") or None, + shuffle=_env_bool("SHUFFLE", True), + seed=_env_int("SEED", 42), + ) + return [_entry_at(dataset, i) for i in range(len(dataset))] + + +def _init_sandbox_fleet_from_env( + entry: dict[str, Any], + group_size: int, + batch_size: int, + max_warmpool_replicas: int | None, +) -> Any: + """Initializes DeepSWE's process-wide SandboxFleet from rollout metadata.""" + with _SANDBOX_INIT_LOCK: + try: + return swe_env._get_global_fleet() # pylint: disable=protected-access + except RuntimeError: + pass + + max_concurrency = _env_int( + "SANDBOX_MAX_CONCURRENCY", + _env_int("ROLLOUT_MAX_CONCURRENCY", group_size), + ) + try: + tasks = _sandbox_tasks_from_env() + except Exception: # pylint: disable=broad-exception-caught + logging.exception( + "Failed to load the full DeepSWE dataset for SandboxFleet. Falling " + "back to the current task only." + ) + tasks = [entry] + logging.info( + "Initializing DeepSWE SandboxFleet in rollout worker with %d task(s) " + "(max_concurrency=%d, batch_size=%d, num_generations=%d, " + "max_warmpool_replicas=%s).", + len(tasks), + max_concurrency, + batch_size, + group_size, + max_warmpool_replicas, + ) + return swe_env._init_global_fleet( # pylint: disable=protected-access + tasks=tasks, + max_concurrency=max_concurrency, + num_generations=group_size, + batch_size=batch_size, + max_warmpool_replicas=max_warmpool_replicas, + ) + + @registry.register_env(DEEPSWE_ENV_NAME) class DeepSWEEnv(swe_env.SWEEnv): """Registry adapter that lets RolloutWorker construct SWEEnv per request.""" @@ -220,33 +312,26 @@ def __init__( prompt_id: str = "", group_index: int = 0, group_size: int = 1, + batch_size: int = 1, + max_warmpool_replicas: int | None = None, policy_version: int = 0, - group_id: Any = None, - pair_index: int | None = None, **kwargs: Any, ): entry = dict(entry or kwargs.pop("task", {}) or {}) if prompt_id and "instance_id" not in entry: entry["instance_id"] = prompt_id - if group_id is None: - group_id = prompt_id or None - if pair_index is None: - pair_index = group_index if kwargs.get("use_agent_sandbox") and kwargs.get("fleet") is None: - logging.info( - "Initializing DeepSWE SandboxFleet in rollout worker " - "(max_concurrency=%s).", - group_size, - ) - kwargs["fleet"] = swe_env._init_global_fleet( # pylint: disable=protected-access - tasks=[entry], - max_concurrency=group_size, + kwargs["fleet"] = _init_sandbox_fleet_from_env( + entry, + group_size=group_size, + batch_size=batch_size, + max_warmpool_replicas=max_warmpool_replicas, ) super().__init__( entry=entry, - group_id=group_id, - pair_index=pair_index, + group_id=prompt_id or None, + pair_index=group_index, **kwargs, ) self.task = { @@ -259,7 +344,17 @@ def __init__( @registry.register_agent(DEEPSWE_AGENT_NAME) -class DeepSWEAgent(swe_agent.SWEAgent): - """Registry adapter for the legacy DeepSWE XML-tool agent.""" +class DeepSWEAgent: + """Loads the recipe Agent only inside the rollout worker process.""" + + name = DEEPSWE_AGENT_NAME + + def __init__(self, **kwargs: Any): + # r2egym is a rollout-only dependency, so keep it out of the CPU + # orchestrator process that also imports this registry module for data. + from examples.deepswe import swe_agent # pylint: disable=g-import-not-at-top + + self._agent = swe_agent.SWEAgent(**kwargs) - name = DEEPSWE_AGENT_NAME \ No newline at end of file + def __getattr__(self, name: str) -> Any: + return getattr(self._agent, name)