diff --git a/examples/learn_to_ask/data_prepare/3_rollout_then_evaluate.py b/examples/learn_to_ask/data_prepare/3_rollout_then_evaluate.py index cfbb6503798..7df4056b1cf 100644 --- a/examples/learn_to_ask/data_prepare/3_rollout_then_evaluate.py +++ b/examples/learn_to_ask/data_prepare/3_rollout_then_evaluate.py @@ -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) diff --git a/tests/explorer/rollout_seed_test.py b/tests/explorer/rollout_seed_test.py new file mode 100644 index 00000000000..799707f5447 --- /dev/null +++ b/tests/explorer/rollout_seed_test.py @@ -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="look")] + 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() diff --git a/trinity/common/models/tinker_model.py b/trinity/common/models/tinker_model.py index 5d05db240a2..68d8ae2ad2b 100644 --- a/trinity/common/models/tinker_model.py +++ b/trinity/common/models/tinker_model.py @@ -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), diff --git a/trinity/common/workflows/envs/alfworld/alfworld_workflow.py b/trinity/common/workflows/envs/alfworld/alfworld_workflow.py index 7fb29604a22..e555277002e 100644 --- a/trinity/common/workflows/envs/alfworld/alfworld_workflow.py +++ b/trinity/common/workflows/envs/alfworld/alfworld_workflow.py @@ -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) diff --git a/trinity/common/workflows/workflow.py b/trinity/common/workflows/workflow.py index 25853322fd9..b3648838af2 100644 --- a/trinity/common/workflows/workflow.py +++ b/trinity/common/workflows/workflow.py @@ -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. @@ -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 diff --git a/trinity/explorer/workflow_runner.py b/trinity/explorer/workflow_runner.py index dcdcd2fb5ea..3beb6faeafd 100644 --- a/trinity/explorer/workflow_runner.py +++ b/trinity/explorer/workflow_runner.py @@ -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