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: 2 additions & 0 deletions examples/learn_to_ask/data_prepare/3_rollout_then_evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@
"prompt_learn2ask",
os.path.join(os.path.dirname(__file__), "..", "workflow", "prompt_learn2ask.py"),
)
if spec is None or spec.loader is None:
raise ImportError("Cannot load the learn-to-ask prompt module.")
prompt_module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(prompt_module)

Expand Down
96 changes: 96 additions & 0 deletions tests/explorer/rollout_seed_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
# -*- coding: utf-8 -*-
"""Regression tests for sampling seeds in repeated rollouts."""

import importlib.util
import unittest
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch

import torch

from trinity.common.config import Config, InferenceModelConfig
from trinity.common.experience import Experience
from trinity.common.workflows.envs.alfworld.alfworld_workflow import (
StepWiseAlfworldWorkflow,
)
from trinity.common.workflows.workflow import Task, Workflow
from trinity.explorer.workflow_runner import WorkflowRunner


class RepeatIndexWorkflow(Workflow):
"""Expose the run context through an experience without model inference."""

can_reset = True

def reset(self, task):
self.task = task

def run(self):
return [
Experience(
tokens=torch.tensor([1, 2]),
prompt_length=1,
reward=float(self.repeat_index),
)
]


class RolloutSeedTest(unittest.IsolatedAsyncioTestCase):
async def test_repeat_indices_include_shard_offset_in_every_runner_mode(self):
for mode in ("sequential", "asynchronous", "multi-threading"):
with self.subTest(mode=mode):
config = Config()
config.explorer.concurrent_mode = mode
config.explorer.rollout_model.enable_history = False
model = MagicMock()
model.clean_workflow_state = AsyncMock()
with patch(
"trinity.explorer.workflow_runner.Allocator.get_model", return_value=model
):
runner = WorkflowRunner(config, rollout_model_id=0, runner_id=0)
task = Task(workflow=RepeatIndexWorkflow)

# Simulate a group split into two runner assignments.
first = await runner._run_task(task, repeat_times=2, run_id_base=0)
second = await runner._run_task(task, repeat_times=2, run_id_base=2)
self.assertTrue(first.status.ok)
self.assertTrue(second.status.ok)
experiences = first.experiences + second.experiences
self.assertEqual([exp.reward for exp in experiences], [0, 1, 2, 3])
self.assertEqual([exp.eid.run for exp in experiences], [0, 1, 2, 3])

def test_alfworld_repeats_have_distinct_reproducible_seeds(self):
model = MagicMock()
model.chat.return_value = [SimpleNamespace(response_text="<action>look</action>")]
with patch.object(StepWiseAlfworldWorkflow, "_setup_environment"):
workflow = StepWiseAlfworldWorkflow(model=model, task=Task(raw_task={}))
workflow.env = MagicMock()
workflow.env.step.return_value = ("room", 0, False, {})

seeds = []
for repeat_index in (0, 1, 7, 0):
workflow.repeat_index = repeat_index
workflow.observation = "room"
workflow.memory = []
workflow.step(0)
seeds.append(model.chat.call_args.kwargs["seed"])
self.assertEqual(seeds, [1, 2, 8, 1])

@unittest.skipUnless(importlib.util.find_spec("tinker"), "Tinker SDK is not installed")
async def test_tinker_only_uses_explicit_request_seeds(self):
from trinity.common.models.tinker_model import TinkerModel

# Avoid a Ray actor and network service; exercise request construction.
model = TinkerModel.__new__(TinkerModel)
model.config = InferenceModelConfig(seed=42, max_response_tokens=16)
model.model = SimpleNamespace(sample_async=AsyncMock())
for kwargs, expected_seed in (({}, None), ({"seed": 0}, 0), ({"seed": 7}, 7)):
with self.subTest(kwargs=kwargs):
await model._generate_internal({"prompt_token_ids": [1, 2]}, **kwargs)
params = model.model.sample_async.call_args.kwargs["sampling_params"]
self.assertEqual(params["seed"], expected_seed)
self.assertEqual(params["max_tokens"], 16)


if __name__ == "__main__":
unittest.main()
4 changes: 3 additions & 1 deletion trinity/common/models/tinker_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,9 @@ async def _generate_internal(self, prompt: dict, **kwargs) -> types.SampleRespon
assert self.model is not None
sampling_params = {
"max_tokens": kwargs.get("max_tokens", self.config.max_response_tokens),
"seed": kwargs.get("seed", self.config.seed),
# A fixed per-request seed replays identical prompts across GRPO
# repeats. Leave sampling unseeded unless the caller requests a seed.
"seed": kwargs.get("seed"),
"temperature": kwargs.get("temperature", 1.0),
"top_k": kwargs.get("top_k", -1),
"top_p": kwargs.get("top_p", 1),
Expand Down
5 changes: 3 additions & 2 deletions trinity/common/workflows/envs/alfworld/alfworld_workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -279,8 +279,9 @@ def step(self, step_num: int) -> bool:
env_state_hash = hashlib.sha256(format_obs.encode()).hexdigest()
self.memory.append({"role": "user", "content": format_obs})

# Get action from the model
responses = self.model.chat(self.memory)
# Repeated episodes start from the same prompt. Give each repeat a
# distinct, reproducible request seed, including across runner shards.
responses = self.model.chat(self.memory, seed=self.repeat_index + 1)
response_text = responses[0].response_text
self.memory.append({"role": "assistant", "content": response_text})
action = parse_action(response_text)
Expand Down
4 changes: 4 additions & 0 deletions trinity/common/workflows/workflow.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,9 @@ class Workflow:
Attributes:
auxiliary_model_wrappers: List of ModelWrapper instances for auxiliary models.
auxiliary_models: List of OpenAI clients (sync or async based on is_async) for auxiliary models.
repeat_index: Zero-based run index within the task's full repeat group.
The runner sets this before each non-repeatable workflow run, including
the offset when the group is split across runners.
"""

can_reset: bool = False # whether the workflow can be reset with a new task. If true, `reset()` must be implemented.
Expand All @@ -108,6 +111,7 @@ def __init__(
else:
self.auxiliary_models = [m.get_openai_client() for m in auxiliary_models]
self.run_id_base = 0
self.repeat_index = 0
self.logger = get_logger(__name__)

@property
Expand Down
1 change: 1 addition & 0 deletions trinity/explorer/workflow_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -257,6 +257,7 @@ async def _execute_single_run(
run_id_base: int,
) -> Tuple[bool, List[Experience], Optional[Dict[str, float]], Optional[str]]:
st = time.time()
workflow.repeat_index = run_id_base + run_index
await self.model_wrapper.clean_workflow_state()
self.runner_state["workflow_id"] = f"{task.batch_id}/{task.task_id}/{run_index}"
self.runner_state["terminate_time"] = None
Expand Down
Loading