Skip to content
Merged
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
2 changes: 1 addition & 1 deletion docs/source/openreward.md
Original file line number Diff line number Diff line change
Expand Up @@ -167,7 +167,7 @@ spec = OpenRewardSpec("Eigent/SETA", indices=list(range(50, 100))) # range

## How tool binding works

At construction the spec calls the env's `/tools` endpoint to fetch a list of tool specs (each with a name, description, and JSON Schema for arguments). For each tool it generates a Python method on the per-rollout adapter with a typed signature and a docstring derived from the schema. So `transformers.utils.get_json_schema` and TRL's `inspect.getmembers(env, ismethod)` both produce the right tool schema for the model with no per-env wrapper code.
At construction the spec calls the env's `/tools` endpoint to fetch a list of tool specs (each with a name, description, and JSON Schema for arguments). For each tool it generates a Python method on the per-rollout adapter with a typed signature and a docstring derived from the schema. The methods are set on the adapter's class, where TRL's tool collector finds them, so `transformers.utils.get_json_schema` and TRL's `inspect.getmembers(type(env), isfunction)` both produce the right tool schema for the model with no per-env wrapper code.

If a tool description contains characters that aren't safe to splice into Python source, the binder falls back to a sanitized form so binding never fails on real envs.

Expand Down
29 changes: 29 additions & 0 deletions tests/experimental/test_async_grpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -767,6 +767,35 @@ def shout(self, text: str) -> str:
finally:
loop._loop.close()

def test_tool_discovery_does_not_evaluate_properties(self):
# The init-time probe lists an environment's tool methods. `inspect.getmembers` calls `getattr` on every name
# before applying the predicate, so listing the instance would evaluate this property; listing the class
# leaves it inert.
class PropertyEnvironment:
def reset(self, **kwargs): ...

@property
def reward(self) -> float:
raise RuntimeError("`reward` must not be evaluated while discovering tools")

def echo(self, text: str) -> str:
"""Echo the text back.

Args:
text: Text to echo.

Returns:
The text, unchanged.
"""
return text

loop = self._make_loop(PropertyEnvironment)
try:
assert [tool.__name__ for tool in loop._env_tools[None]] == ["echo"]
assert [tool.__name__ for tool in loop.tools] == ["echo"]
finally:
loop._loop.close()

def test_unknown_environment_raises(self):
# An example whose `environment` field doesn't match any configured environment should fail with a clear error
# rather than a bare KeyError mid-rollout. The check fires before any generation, so no vLLM is needed here.
Expand Down
9 changes: 5 additions & 4 deletions tests/experimental/test_harbor.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,12 +127,13 @@ def test_outcome_reward_uses_environment_reward_when_passed(self):
assert _outcome_reward_func(environment_reward=[0.25, 0.75]) == [0.25, 0.75]

def test_fresh_env_reward_is_zero_without_backend(self):
# The trainer discovers tool methods via `inspect.getmembers`, which evaluates properties. A fresh
# env (never `reset`) must expose its tools and return 0.0 from `reward` WITHOUT starting the
# Harbor backend or importing `harbor` (not installed in the trainer env).
# The trainer discovers tool methods on the environment's class (`inspect.getmembers` on the
# instance would evaluate `reward`). A fresh env (never `reset`) must expose its tools and return
# 0.0 from `reward` WITHOUT starting the Harbor backend or importing `harbor` (not installed in
# the trainer env).
import inspect

env = HarborBashEnv()
names = {n for n, _ in inspect.getmembers(env, predicate=inspect.ismethod)}
names = {n for n, _ in inspect.getmembers(type(env), predicate=inspect.isfunction)}
assert {"bash", "reset"} <= names
assert env.reward == 0.0
71 changes: 71 additions & 0 deletions tests/test_grpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -3457,6 +3457,77 @@ def fake_generate(input_ids, **kwargs):
new_param = trainer.model.get_parameter(n)
assert not torch.equal(param, new_param), f"Parameter {n} has not changed."

@pytest.mark.xfail(
condition=Version(transformers.__version__) < Version("5.2.0"),
reason="Environment factory support is not available in transformers versions below 5.2.0",
strict=True,
)
@require_response_parsing
@patch.dict(os.environ, {"TRL_EXPERIMENTAL_SILENCE": "1"})
def test_tool_discovery_does_not_evaluate_environment_properties(self):
# Tool discovery must list the environment's methods on its class. `inspect.getmembers` calls `getattr` on
# every name before applying the predicate, so listing the instance evaluates its properties, and a property
# may be expensive or have side effects (a Harbor env scores the rollout from one). Here `reward` is inert on
# a fresh instance and raises once the environment has been reset: the first batch resets the pooled
# instances, so a second batch that listed the instance would raise before generating anything.
class PropertyEnvironment:
def __init__(self):
self.is_reset = False

def reset(self, **kwargs):
self.is_reset = True

@property
def reward(self) -> float:
if self.is_reset:
raise RuntimeError("`reward` must not be evaluated while discovering tools")
return 0.0

def echo(self, text: str) -> str:
"""
Echo the text back.

Args:
text: Text to echo.

Returns:
The text, unchanged.
"""
return text

def reward_func(completions, **kwargs):
# Never reads `environments`, so nothing but tool discovery can touch the `reward` property.
return [1.0] * len(completions)

training_args = GRPOConfig(
output_dir=self.tmp_dir,
per_device_train_batch_size=3, # reduce the batch size to reduce memory usage
num_generations=3, # reduce the number of generations to reduce memory usage
max_steps=2, # the first batch resets the pooled instances, the second re-lists their tools
report_to="none",
)

trainer = GRPOTrainer(
model="trl-internal-testing/tiny-Qwen3MoeForCausalLM",
reward_funcs=reward_func,
args=training_args,
environment_factory=PropertyEnvironment,
)

def fake_generate(input_ids, **kwargs):
# "I won't increment<|im_end|>" — no tool call, so one generation round per step.
completion_ids = torch.tensor(
[[40, 2765, 944, 16252, 151645]] * input_ids.shape[0], device=input_ids.device
)
return torch.cat([input_ids, completion_ids], dim=-1)

with patch.object(trainer.model, "generate", side_effect=fake_generate):
trainer.train()

assert trainer.state.log_history[-1]["train_loss"] is not None
# The regression is only exercised if the instances the last batch listed had really been reset before it.
assert all(environment.is_reset for environment in trainer.environments)

@pytest.mark.xfail(
condition=Version(transformers.__version__) < Version("5.2.0"),
reason="Environment factory support is not available in transformers versions below 5.2.0",
Expand Down
9 changes: 5 additions & 4 deletions trl/experimental/async_grpo/async_rollout_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -414,14 +414,15 @@ def __init__(
instance = factory()
has_reset = False
methods = []
for member_name, member in inspect.getmembers(instance, predicate=inspect.ismethod):
# List on the class: getmembers on the instance evaluates properties
for member_name, _ in inspect.getmembers(type(instance), predicate=inspect.isfunction):
if member_name == "reset":
has_reset = True
elif member_name == "get_reward":
if type(instance) not in self._env_reward_types:
self._env_reward_types.append(type(instance))
elif not member_name.startswith("_"):
methods.append(member)
methods.append(getattr(instance, member_name))
if not has_reset:
raise ValueError(
"Each environment instance returned by `environment_factory` must define a callable `reset`."
Expand Down Expand Up @@ -574,8 +575,8 @@ async def _generate_loop(self, stop_event: asyncio.Event) -> None:
methods = []
if environment is not None:
methods = [
member
for member_name, member in inspect.getmembers(environment, predicate=inspect.ismethod)
getattr(environment, member_name)
for member_name, _ in inspect.getmembers(type(environment), predicate=inspect.isfunction)
if member_name not in ("reset", "get_reward") and not member_name.startswith("_")
]
tool_dict = {tool.__name__: tool for tool in self._standalone_tools + methods}
Expand Down
7 changes: 4 additions & 3 deletions trl/experimental/harbor/_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,9 +87,10 @@ def _exec(self, command: str, timeout: int = 180) -> str:
def reward(self) -> float:
# Submission = the agent wrote /workdir/answer.txt during the rollout; the verifier reads it.
# Computed once, lazily, on first read (TRL reads this after the rollout via reward_funcs).
# A fresh env that was never `reset` (e.g. the trainer probing tool methods via
# `inspect.getmembers`, which evaluates properties) has no sandbox/task to verify — return 0.0
# without invoking the verifier, which would start the Harbor backend and import `harbor`.
# A fresh env that was never `reset` has no sandbox/task to verify — return 0.0 without invoking
# the verifier, which would start the Harbor backend and import `harbor`. The trainer discovers
# tool methods on the class, so it never reads this property; the guard is for user code that
# touches a fresh env directly.
if self._env is None:
return 0.0
if self._reward is _NO_REWARD:
Expand Down
10 changes: 6 additions & 4 deletions trl/trainer/grpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -649,13 +649,14 @@ def __init__(
has_reset = False
has_reward = False
methods = []
for member_name, member in inspect.getmembers(instance, predicate=inspect.ismethod):
# List on the class: getmembers on the instance evaluates properties
for member_name, _ in inspect.getmembers(type(instance), predicate=inspect.isfunction):
if member_name == "reset":
has_reset = True
elif member_name == "get_reward":
has_reward = True
elif not member_name.startswith("_"):
methods.append(member)
methods.append(getattr(instance, member_name))
if not has_reset:
raise ValueError(
"Each environment instance returned by `environment_factory` must define a callable `reset`."
Expand Down Expand Up @@ -2275,9 +2276,10 @@ def _generate_and_score_completions(
for i in range(len(inputs)):
methods = []
if self.environments:
environment = self.environments[i]
methods = [
member
for member_name, member in inspect.getmembers(self.environments[i], predicate=inspect.ismethod)
getattr(environment, member_name)
for member_name, _ in inspect.getmembers(type(environment), predicate=inspect.isfunction)
if member_name not in ("reset", "get_reward") and not member_name.startswith("_")
]
sync_tool_dict, async_tool_dict = {}, {}
Expand Down
Loading