Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
343 changes: 24 additions & 319 deletions examples/deepswe/swe_agent.py

Large diffs are not rendered by default.

216 changes: 166 additions & 50 deletions examples/deepswe/swe_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,10 +2,17 @@
import json
import logging
import os
import re
import threading
import time
from typing import Any, Optional, cast
import numpy as np

try:
from examples.deepswe import template as template_mod
except ImportError:
import template as template_mod # pytype: disable=import-error

_GLOBAL_FLEET = None
_FLEET_LOCK = threading.Lock()
_PATCH_LOCK = threading.Lock()
Expand Down Expand Up @@ -79,26 +86,54 @@ def _patched_start_container(
logging.debug("[SandboxFleet] r2egym in-memory patch note: %s", e)


def _normalize_tasks_for_fleet(tasks: Any) -> list[Any]:
def _get_image_rewrite_fn(image_rewrite: Any | None = None) -> Any | None:
"""Retrieve or construct the image rewrite function from prefix if configured."""
if image_rewrite is not None:
return image_rewrite
if os.getenv("IMAGE_REWRITE_PREFIX"):
prefix = os.environ["IMAGE_REWRITE_PREFIX"].rstrip("/")
return lambda img: f"{prefix}/{img.split('/')[-1]}"
return None


def _normalize_tasks_for_fleet(
tasks: Any, scaffold: str = "r2egym"
) -> list[Any]:
"""Normalize heterogeneous dataset entries into Task objects for SandboxFleet."""
TaskCls = None
try:
from agent_sandbox_rl import Task # pytype: disable=import-error
if isinstance(Task, type):
TaskCls = Task
except ImportError:
return list(tasks)
pass

if TaskCls is None:
from dataclasses import dataclass

@dataclass
class TaskCls: # pytype: disable=reimported
id: str
image: str
metadata: dict[str, Any]

del scaffold
normalized = []
for item in tasks:
if isinstance(item, Task):
if hasattr(item, "image") and hasattr(item, "id"):
normalized.append(item)
elif isinstance(item, dict):
img = item.get("docker_image") or item.get("image", "default")
img = (
item.get("docker_image")
or item.get("image", "default")
)
if isinstance(img, (list, np.ndarray)):
img = img[0] if len(img) > 0 else "default"
t_id = item.get("instance_id") or item.get("id") or img
if isinstance(t_id, (list, np.ndarray)):
t_id = t_id[0] if len(t_id) > 0 else "default"
normalized.append(
Task(id=str(t_id), image=str(img), metadata={"ds": item})
TaskCls(id=str(t_id), image=str(img), metadata={"ds": item})
)
else:
normalized.append(item)
Expand All @@ -111,6 +146,8 @@ def _init_global_fleet(
num_generations: int = 8,
batch_size: int = 8,
max_warmpool_replicas: int | None = None,
scaffold: str = "r2egym",
image_rewrite: Any | None = None,
) -> Any:
"""Initialize the process-wide SandboxFleet instance once upfront."""
global _GLOBAL_FLEET
Expand All @@ -126,7 +163,6 @@ def _init_global_fleet(
FleetConfig,
SandboxFleet,
Task,
TemplateSpec,
)
except ImportError as e:
raise ImportError(
Expand All @@ -144,30 +180,44 @@ def _init_global_fleet(
effective_max_concurrent = max(
max_concurrency, batch_size * num_generations * 2
)
fleet_cfg = FleetConfig(
clusters=[

template = template_mod.get_template(scaffold, node_sel)

fleet_kwargs = {
"clusters": [
ClusterConfig(
name="default",
namespace=fleet_ns,
node_selector=node_sel,
in_cluster=in_cluster,
)
],
max_concurrent=effective_max_concurrent,
window_size=batch_size,
max_warmpool_size=max_warmpool_replicas
if max_warmpool_replicas is not None
else num_generations,
warm_per_task=True,
)
"max_concurrent": effective_max_concurrent,
"window_size": batch_size,
"max_warmpool_size": (
max_warmpool_replicas
if max_warmpool_replicas is not None
else num_generations
),
"warm_per_task": True,
}
if template is not None:
fleet_kwargs["template"] = template
fleet_cfg = FleetConfig(**fleet_kwargs)
fleet_inst = SandboxFleet(fleet_cfg)
image_rewrite_fn = _get_image_rewrite_fn(image_rewrite)
fleet_inst._image_rewrite_fn = image_rewrite_fn
if tasks is not None:
fleet_inst.load_tasks(_normalize_tasks_for_fleet(tasks))
normalized_tasks = _normalize_tasks_for_fleet(tasks, scaffold=scaffold)
if image_rewrite_fn is not None:
fleet_inst.load_tasks(normalized_tasks, image_rewrite=image_rewrite_fn)
else:
fleet_inst.load_tasks(normalized_tasks)
msg = (
f"[SandboxFleet] Initializing pipelined fleet in namespace={fleet_ns}"
f" (max_concurrent={effective_max_concurrent},"
f" window_size={batch_size},"
f" max_warmpool_replicas={max_warmpool_replicas if max_warmpool_replicas is not None else num_generations},"
f" max_warmpool_replicas={fleet_kwargs['max_warmpool_size']},"
" warm_per_task=True)..."
)
logging.info(msg)
Expand Down Expand Up @@ -201,12 +251,18 @@ def __init__(
num_generations: int = 8,
batch_size: int = 8,
max_warmpool_replicas: int | None = None,
scaffold: str = "r2egym",
image_rewrite: Any | None = None,
):
self.dataset_iter = iter(dataset)
self.num_generations = num_generations
self.batch_size = batch_size
self.max_warmpool_replicas = max_warmpool_replicas
self.scaffold = scaffold
self.fleet = fleet or _get_global_fleet()
self.image_rewrite = _get_image_rewrite_fn(
image_rewrite or getattr(self.fleet, "_image_rewrite_fn", None)
)
self.current_batch = None
self.next_batch = None
self.prev_batch_images: list[str] = []
Expand Down Expand Up @@ -246,12 +302,14 @@ def _extract_images(self, batch: Any) -> list[str]:
if isinstance(item, dict) and item.get("docker_image")
]

# Safely decode/stringify all elements
# Safely decode/stringify all elements and apply rewrite if configured
rewrite_fn = getattr(self, "image_rewrite", None) or _get_image_rewrite_fn()
str_images = []
for img in raw_images:
str_images.append(
img.decode("utf-8") if hasattr(img, "decode") else str(img)
)
s = img.decode("utf-8") if hasattr(img, "decode") else str(img)
if rewrite_fn is not None:
s = rewrite_fn(s)
str_images.append(s)

return list(dict.fromkeys(str_images))

Expand Down Expand Up @@ -341,7 +399,25 @@ def _teardown_global_fleet() -> None:
r2egym = cast(Any, None)
EnvArgs = cast(Any, None)
RepoEnv = cast(Any, None)
Action = cast(Any, None)
Action = None


class _ActionFallback:
"""Minimal Action parser fallback when r2egym is not installed."""

def __init__(self, function_name: str, parameters: dict[str, str]):
self.function_name = function_name
self.parameters = parameters

@classmethod
def from_string(cls, action_str: str) -> "_ActionFallback":
fn_match = re.search(r"<function\s*=\s*([^>]+)>", action_str)
function_name = fn_match.group(1).strip() if fn_match else ""
pattern = r"<parameter\s*=\s*([^>]+)>(.*?)</parameter>"
param_matches = re.findall(pattern, action_str, flags=re.DOTALL)
params = {k.strip(): v.strip() for k, v in param_matches}
return cls(function_name, params)


from tunix.rl.agentic.environments.base_environment import BaseTaskEnv, EnvStepResult

Expand Down Expand Up @@ -410,7 +486,7 @@ def __init__(
backend: Backend to use for the environment.
delete_image: Whether to delete the Docker image after closing.
verbose: Verbose output toggle.
scaffold: Scaffold tool set ('r2egym' or 'sweagent').
scaffold: Scaffold tool set ('r2egym', 'sweagent', or 'openhands').
max_steps: Maximum interaction steps.
use_agent_sandbox: If True, strictly forces SandboxFleet and
AgentSandboxRuntime.
Expand All @@ -432,12 +508,9 @@ def __init__(
assert scaffold in [
"r2egym",
"sweagent",
], f"Invalid scaffold: {scaffold}, must be one of ['r2egym', 'sweagent']"
"openhands",
], f"Invalid scaffold: {scaffold}, must be one of ['r2egym', 'sweagent', 'openhands']"
super().__init__(max_steps=max_steps)

if not hasattr(self, "extra_kwargs"):
self.extra_kwargs = {}

self.extra_kwargs["group_id"] = group_id
self.extra_kwargs["pair_index"] = pair_index

Expand All @@ -446,29 +519,63 @@ def _initial_observation(self) -> Any:
if self.use_agent_sandbox:
_patch_r2egym_for_agent_sandbox()
from agent_sandbox_rl import Task # pytype: disable=import-error
from agent_sandbox_rl.adapters.r2egym import ( # pytype: disable=import-error
make_fleet_repo_env,
r2egym_command_files,
)

fleet = self.fleet or _get_global_fleet()
msg = (
"[SWEEnv] Acquiring SandboxHandle from SandboxFleet and"
" constructing FleetRepoEnv!"
"[SWEEnv] Acquiring SandboxHandle from SandboxFleet!"
)
logging.info(msg)
task = Task(
id=str(
self.entry.get(
"instance_id", self.entry.get("docker_image", "default")
)
),
image=self.entry.get("docker_image", "default"),
metadata={"ds": self.entry},
task_id = str(
self.entry.get(
"instance_id", self.entry.get("docker_image", "default")
)
)
self.handle = fleet.acquire(task)
# TODO(wuhao): Revisit command_files once other harnesses (such as OpenHands) are supported.
cmd_files = r2egym_command_files()
task = None
if hasattr(fleet, "tasks") and fleet.tasks:
for t in fleet.tasks:
if t.id == task_id:
task = t
break
if task is None:
task_img = self.entry.get("docker_image", "default")
if isinstance(task_img, (list, np.ndarray)):
task_img = task_img[0] if len(task_img) > 0 else "default"
task_img_str = str(task_img)
rewrite_fn = _get_image_rewrite_fn(
getattr(fleet, "_image_rewrite_fn", None)
)
if rewrite_fn:
task_img_str = rewrite_fn(task_img_str)
task = Task(
id=task_id,
image=task_img_str,
metadata={"ds": self.entry},
)
max_acquire_retries = 5
for attempt in range(max_acquire_retries):
try:
self.handle = fleet.acquire(task)
break
except Exception as e:
if attempt < max_acquire_retries - 1:
logging.warning(
"[SWEEnv] fleet.acquire failed (attempt %d/%d): %s; retrying in %ds...",
attempt + 1,
max_acquire_retries,
e,
5 * (attempt + 1),
)
time.sleep(5 * (attempt + 1))
else:
raise
from agent_sandbox_rl.adapters.r2egym import ( # pytype: disable=import-error
make_fleet_repo_env,
r2egym_command_files,
)
if self.scaffold in ("sweagent", "openhands"):
cmd_files = SWEAGENT_COMMAND_FILES
else:
cmd_files = r2egym_command_files()
self.env = make_fleet_repo_env(self.handle, command_files=cmd_files)
else:
# Initialize standard local Docker RepoEnv
Expand All @@ -486,12 +593,13 @@ def _initial_observation(self) -> Any:
)
if self.scaffold == "r2egym":
self.env.add_commands(R2EGYM_COMMAND_FILES)
else:
elif self.scaffold in ("sweagent", "openhands"):
self.env.add_commands(SWEAGENT_COMMAND_FILES)
else:
self.env.reset()
if self.env is not None:
self.env.reset()

self.final_reward_fn = self.env.compute_reward # pytype: disable=attribute-error
self.final_reward_fn = self.env.compute_reward
self.total_steps = 0

# Polls docker runtime to get task instruction.
Expand All @@ -500,14 +608,22 @@ def _initial_observation(self) -> Any:
def _step_impl(self, action: Any) -> EnvStepResult:
global Action
if Action is None:
from r2egym.agenthub.action import Action # pytype: disable=import-error
try:
from r2egym.agenthub.action import Action # pytype: disable=import-error
except ImportError:
Action = _ActionFallback
if isinstance(action, str):
action_obj = Action.from_string(action)
else:
action_obj = action

if not action_obj.function_name:
return EnvStepResult(observation="", reward=0, done=False, info={})
return EnvStepResult(
observation="",
reward=0,
done=False,
info={"max_steps": self.max_steps},
)

# RepoEnv always returns 0 reward, must be evaluated by DockerRuntime.
if not self.env:
Expand Down
Loading
Loading