-
Notifications
You must be signed in to change notification settings - Fork 462
fix(textarena): honor reset(seed) so episodes are reproducible #1078
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -8,12 +8,19 @@ | |
|
|
||
| from __future__ import annotations | ||
|
|
||
| import random | ||
| import sys | ||
| import threading | ||
| from typing import Any, Dict, Iterable, List, Optional | ||
| from uuid import uuid4 | ||
|
|
||
| from openenv.core.env_server.interfaces import Environment | ||
|
|
||
| try: | ||
| import numpy as _np | ||
| except Exception: # pragma: no cover - numpy is optional | ||
| _np = None | ||
|
|
||
| try: | ||
| # When running as installed package | ||
| from textarena_env.models import ( | ||
|
|
@@ -38,6 +45,11 @@ | |
| _TEXTARENA_IMPORT_ERROR: Exception | None = None | ||
| _NLTK_DOWNLOADED: bool = False | ||
|
|
||
| # TextArena selects episodes/words via the process-global ``random`` (and | ||
| # sometimes ``numpy``) RNGs. Seeding+selection must be atomic across concurrent | ||
| # sessions, so guard it with a process-wide lock. | ||
| _SEED_LOCK = threading.Lock() | ||
|
|
||
|
|
||
| def _ensure_nltk_data() -> None: | ||
| """Download NLTK data once per process.""" | ||
|
|
@@ -147,7 +159,7 @@ def reset( | |
| if hasattr(env, "full_observations"): | ||
| env.full_observations = {} | ||
|
|
||
| self._ta_env.reset(num_players=self.num_players) | ||
| self._seeded_reset(seed) | ||
|
|
||
| for provider in self._reward_providers: | ||
| provider.reset() | ||
|
|
@@ -203,6 +215,48 @@ def step(self, action: TextArenaAction) -> TextArenaObservation: # type: ignore | |
| def state(self) -> TextArenaState: | ||
| return self._state | ||
|
|
||
| # ------------------------------------------------------------------ | ||
| # Seeding | ||
| # ------------------------------------------------------------------ | ||
| def _seeded_reset(self, seed: Optional[int]) -> None: | ||
| """Reset the underlying TextArena env, honoring ``seed`` for reproducibility. | ||
|
|
||
| TextArena chooses episodes/words with the process-global ``random`` (and | ||
| sometimes ``numpy``) RNGs and does *not* apply ``reset(seed=...)`` to that | ||
| selection. So we seed the global RNGs ourselves immediately before the | ||
| underlying reset. This makes episodes reproducible, which is required, for | ||
| example, for GRPO-style training where every rollout in a group must share | ||
| the same episode (same Wordle word) for the group baseline to be valid. | ||
|
|
||
| A process-wide lock keeps seed+selection atomic across concurrent sessions, | ||
| and we restore the prior RNG state afterwards so that unseeded sessions | ||
| sharing the process are left undisturbed. | ||
| """ | ||
| if seed is None: | ||
| self._ta_reset(seed) | ||
| return | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Unseeded reset skips seed lockMedium Severity
Additional Locations (1)Reviewed by Cursor Bugbot for commit 58723a9. Configure here.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Unseeded reset skips seed lockMedium Severity Unseeded Reviewed by Cursor Bugbot for commit 2131c0f. Configure here. |
||
|
|
||
| with _SEED_LOCK: | ||
| py_state = random.getstate() | ||
| np_state = _np.random.get_state() if _np is not None else None | ||
| try: | ||
| random.seed(seed) | ||
| if _np is not None: | ||
| _np.random.seed(seed) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Large seeds crash with numpyLow Severity When numpy is importable, Reviewed by Cursor Bugbot for commit 2131c0f. Configure here. |
||
| self._ta_reset(seed) | ||
| finally: | ||
| random.setstate(py_state) | ||
| if np_state is not None: | ||
| _np.random.set_state(np_state) | ||
|
|
||
| def _ta_reset(self, seed: Optional[int]) -> None: | ||
| """Call the underlying TextArena reset, forwarding ``seed`` when supported.""" | ||
| try: | ||
| self._ta_env.reset(num_players=self.num_players, seed=seed) | ||
| except TypeError: | ||
| # Older/other TextArena games whose reset does not accept a seed. | ||
| self._ta_env.reset(num_players=self.num_players) | ||
|
|
||
| # ------------------------------------------------------------------ | ||
| # Helpers | ||
| # ------------------------------------------------------------------ | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,79 @@ | ||
| # SPDX-License-Identifier: BSD-3-Clause | ||
|
|
||
| """Reproducibility tests for TextArena environment seed handling. | ||
|
|
||
| TextArena chooses its episode/word with the global ``random`` RNG and does not | ||
| apply ``reset(seed=...)`` to that selection, so the OpenEnv wrapper seeds the | ||
| global RNGs itself. These tests pin that behaviour down: same seed -> same word, | ||
| seed drives selection, and seeding a reset does not disturb the global RNG stream. | ||
| """ | ||
|
|
||
| import random | ||
|
|
||
| import pytest | ||
|
|
||
| # Skip the whole module if the optional textarena dependency is not installed. | ||
| pytest.importorskip("textarena", reason="textarena is not installed") | ||
|
|
||
| from envs.textarena_env.server.environment import TextArenaEnvironment | ||
|
|
||
|
|
||
| def _secret_word(env: TextArenaEnvironment): | ||
| """Best-effort extraction of the hidden Wordle word from the underlying env.""" | ||
| inner = env._ta_env | ||
| while inner is not None: | ||
| game_state = getattr(getattr(inner, "state", None), "game_state", None) | ||
| if isinstance(game_state, dict) and "secret_word" in game_state: | ||
| return game_state["secret_word"] | ||
| inner = getattr(inner, "env", None) | ||
| return None | ||
|
|
||
|
|
||
| @pytest.fixture(scope="module") | ||
| def env(): | ||
| return TextArenaEnvironment(env_id="Wordle-v0", num_players=1) | ||
|
|
||
|
|
||
| def test_same_seed_same_word(env): | ||
| env.reset(seed=123) | ||
| first = _secret_word(env) | ||
| env.reset(seed=123) | ||
| second = _secret_word(env) | ||
| assert first is not None | ||
| assert first == second | ||
|
|
||
|
|
||
| def test_same_seed_across_instances(): | ||
| a = TextArenaEnvironment(env_id="Wordle-v0", num_players=1) | ||
| b = TextArenaEnvironment(env_id="Wordle-v0", num_players=1) | ||
| a.reset(seed=7) | ||
| b.reset(seed=7) | ||
| assert _secret_word(a) is not None | ||
| assert _secret_word(a) == _secret_word(b) | ||
|
|
||
|
|
||
| def test_seed_drives_word_selection(env): | ||
| words = set() | ||
| for seed in range(12): | ||
| env.reset(seed=seed) | ||
| words.add(_secret_word(env)) | ||
| # Combined with the determinism tests above, more than one distinct word | ||
| # confirms the seed actually drives selection rather than being ignored. | ||
| assert len(words) > 1 | ||
|
|
||
|
|
||
| def test_unseeded_reset_still_works(env): | ||
| observation = env.reset() | ||
| assert observation.done is False | ||
| assert _secret_word(env) is not None | ||
|
|
||
|
|
||
| def test_seed_restores_global_rng_state(env): | ||
| random.seed(999) | ||
| baseline = random.random() | ||
|
|
||
| random.seed(999) | ||
| env.reset(seed=1) # a seeded reset must not disturb the global RNG stream | ||
| after = random.random() | ||
|
|
||
| assert baseline == after |


There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
seed=Nonebypasses_SEED_LOCK, so an unseeded reset can run after another thread callsrandom.seed(seed)but before it selects/restores. That consumes the seeded stream, perturbs the unseeded session, and can change the seeded episode. Hold the same lock around every underlying reset; for seeded calls, save/seed/reset/restore while holding it. Also, TextArena 0.7.4 already forwardsseedthroughSinglePlayerState.__init__toState.__init__, which callsrandom.seed(seed), so update the docstring’s contrary claim.