perf(genrm): reduce response copying and group cleanup overhead - #3625
yuhezhang-ai wants to merge 9 commits into
Conversation
Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
3bafd27 to
05baf2c
Compare
|
/ok to test 05baf2c |
| cohort = watermark.active_cohort if watermark is not None else None | ||
| if cohort is None or cohort.group_attempt >= new_attempt: | ||
| return | ||
| async with cohort.lock: |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:552-568
None of the code that runs while holding cohort.lock or _cohort_registry_lock awaits anything, except for acquiring cohort.lock itself. That is true on main as well as in this PR.
asyncio.Lock.acquire() returns without suspending when the lock is free (CPython 3.13 locks.py). So in practice this lock is never contended, and a cancellation cannot land between retiring the old attempt and committing the new watermark.
Two review experiments found no contention on either main or this branch:
- a stress run with 200 groups, random judge delays, and client cancellations
- 300 randomized interleavings of the final member arriving alongside a replacement request
The two replacement-race tests in test_cohort_transition_races.py reproduce the bug on main only because they substitute an ObservedLock and hold it from the test.
The change is still reasonable as hardening. Could the PR description present it that way, rather than as a fix for a race that happens today?
The whole state machine now depends on one rule: no cohort.lock or _cohort_registry_lock section may await anything other than acquiring cohort.lock.
Nothing enforces that rule today.
If someone later adds an await inside _publish_verify_cohort(), every group's arrival would wait on the registry lock, and the interleavings these tests simulate would become real.
Could you record the rule in a comment and add a test that fails when it is broken?
| async with cohort.lock: | |
| async with cohort.lock: | |
| # Safe to wait here while holding the registry lock only because no cohort.lock section awaits. | |
| # Adding an await under cohort.lock would make every group's arrival wait on this one group. |
The test below replaces both locks with a lock that records whether its holder let the event loop run while holding it. It then drives a randomized workload with duplicates, cancellations, judge failures, pruning, and a shutdown while groups are in flight.
It passes on this branch.
Adding await asyncio.sleep(0) inside _publish_verify_cohort()'s lock makes it fail.
Proposed tests/test_lock_invariant.py
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""Cohort and registry lock bodies must never suspend.
The state machine relies on this: a lock that is never held across an await is never contended,
so waiters, phase changes, and registry updates are atomic with respect to the event loop.
"""
import asyncio
import random
from fastapi import HTTPException
import resources_servers.genrm_compare.app as genrm
from nemo_gym.judge import JudgeError
from resources_servers.genrm_compare.tests.test_cohort_lifecycle import member
class NoYieldLock(asyncio.Lock):
"""Records an acquisition that is still held when the loop next runs a ready callback."""
def __init__(self, violations: list[str]):
super().__init__()
self._generation = 0
self._violations = violations
async def acquire(self) -> bool:
await super().acquire()
self._generation += 1
generation = self._generation
holder = asyncio.current_task()
asyncio.get_running_loop().call_soon(self._check, generation, holder)
return True
def _check(self, generation: int, holder: asyncio.Task | None) -> None:
if self.locked() and self._generation == generation:
self._violations.append(repr(holder.get_stack(limit=3) if holder else None))
async def test_detector_flags_a_suspending_body():
violations = []
lock = NoYieldLock(violations)
async with lock:
pass
await asyncio.sleep(0)
assert not violations
async with lock:
await asyncio.sleep(0)
assert len(violations) == 1
async def test_lock_bodies_never_suspend_under_randomized_workload(server, monkeypatch):
violations: list[str] = []
init = genrm._CohortState.__init__
def instrumented_init(self, *args, **kwargs):
init(self, *args, **kwargs)
self.lock = NoYieldLock(violations)
monkeypatch.setattr(genrm._CohortState, "__init__", instrumented_init)
server._cohort_registry_lock = NoYieldLock(violations)
server.config.num_rollouts_per_prompt = 4
server.config.max_terminal_cohorts = 8
rng = random.Random(0)
async def compare(*args, response_objs, **kwargs):
await asyncio.sleep(rng.random() * 0.005)
if rng.random() < 0.1:
raise JudgeError("judge down")
return [1.0] * len(response_objs), {}, [], []
server._run_compare = compare
outcomes: dict[str, int] = {}
async def one(group: int, attempt: int, index: int) -> None:
await asyncio.sleep(rng.random() * 0.02)
task = asyncio.create_task(server.verify(member(index, group=f"g{group}", attempt=attempt)))
if rng.random() < 0.2:
await asyncio.sleep(rng.random() * 0.005)
task.cancel()
try:
await task
outcome = "ok"
except (asyncio.CancelledError, HTTPException, JudgeError) as error:
outcome = type(error).__name__
outcomes[outcome] = outcomes.get(outcome, 0) + 1
requests = [one(g, a, i) for g in range(100) for a in range(3) for i in range(4) for _ in range(rng.randint(1, 2))]
await asyncio.gather(*requests)
# Shut down while collecting and evaluating cohorts still own waiters and tasks.
blocked = asyncio.Event()
async def blocked_compare(*args, response_objs, **kwargs):
await blocked.wait()
return [1.0] * len(response_objs), {}, [], []
server._run_compare = blocked_compare
in_flight = [asyncio.create_task(one(g, 0, i)) for g in range(100, 140) for i in range(rng.randint(1, 4))]
await asyncio.sleep(0.05)
await asyncio.gather(server.aclose(), *in_flight)
await asyncio.sleep(0) # let the last sentinels run
assert not violations, violations[:3]
assert outcomes.get("ok", 0) > 100 and outcomes.get("CancelledError", 0) > 50, outcomes| results = await asyncio.wait_for(asyncio.gather(first, last, return_exceptions=True), 0.5) | ||
| assert isinstance(results[0], HTTPException) and results[0].status_code == 503 | ||
| assert "evaluation task factory failed" in results[0].detail | ||
| assert isinstance(results[1], (HTTPException, asyncio.CancelledError)) |
There was a problem hiding this comment.
resources_servers/genrm_compare/tests/test_cohort_transition_races.py:82-127
This test would go away if the task-startup handler is removed, as suggested in the comment on app.py.
If the handler stays, note that the disconnect=True case never delivers a cancellation to the request.
Task startup fails without any await, so _fail_verify_cohort_locked() completes the request's future before verify() reaches await asyncio.shield(future).
That await then returns immediately, and the request finishes with a 503 in the same event-loop step.
The cancel() scheduled with call_soon runs afterwards and returns False.
Running an instrumented copy of this test showed results[1] is an HTTPException in both cases.
As a result, the two parametrizations test the same path.
The assertion that accepts asyncio.CancelledError hides this.
In that case, please remove the disconnect parameter or assert the outcome that actually happens:
| assert isinstance(results[1], (HTTPException, asyncio.CancelledError)) | |
| assert isinstance(results[1], HTTPException) and results[1].status_code == 503 |
| waiters=[future], | ||
| ) | ||
| if cohort.collection_timeout_task is None: | ||
| del response_obj # Only the cohort owns the compact scoring payload while this request waits. |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:391-408
The waiting verify() call still holds the full request body until its reward returns, because _verify_response() echoes body.response.
On main, member.body pointed at that same object rather than a copy.
So storing the compact dict does not reduce memory while a request is waiting.
The real saving is in two other places.
First, evaluation no longer calls model_dump() on every member and holds those copies for the whole judging call.
Second, a member whose client disconnected no longer keeps its full body alive.
A review measurement of an answer with 65,536 generated token IDs found that model_dump() retains about 1.1 MiB per answer.
That matches the 8,192-answer benchmark, but only because every group in that benchmark is judged at the same time.
In production, the saving scales with the number of answers being judged at once, not the number held in memory.
Could the PR description say this more precisely?
The del response_obj also frees only a small text dict, so it and its comment can be removed.
| except Exception as error: | ||
| logger.exception("GenRM cohort task startup failed for %s", prompt_key) | ||
| # Complete the failure before releasing the lock. A disconnect | ||
| # must not leave a group with neither a timer nor a judging task. |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:410-433
This try/except runs only if asyncio.create_task() raises.
On a running event loop with no custom task factory, it does not raise.
In CPython 3.13, BaseEventLoop.create_task() fails only when the loop is closed, and verify() is itself running on that loop.
Shutdown does not create that situation either.
aclose() fails the open groups before it cancels their tasks, so later arrivals see a failed group and never start a task.
The coro.close() guard added to _own_task() covers the same case.
All six test cases that reach these paths do so by monkeypatching create_task to fail:
test_failed_evaluation_start_releases_peers_even_with_disconnect(both parameters)test_failed_collection_start_fails_registered_membertest_task_start_failure_releases_registered_waiters_and_closes_coroutine(both parameters)test_judge_task_start_failure_releases_all_http_waiters
Could you remove both handlers and those tests?
They add branches that production never reaches, and the repository guidelines ask to avoid code for hypothetical failures.
On a copy of this branch, removing them deletes 13 lines from app.py and about 150 test lines, and the rest of the suite passes.
The description's claims about task-startup failures would then need to be dropped too.
If you prefer to keep the handler, please fix its comment.
The try block contains no await, so a disconnect cannot interrupt it.
The real reason to fail the group here is that failing it while the lock is still held leaves no point where a cancellation could interrupt the cleanup.
| return | ||
| excess = len(watermarks) - limit | ||
| evicted = [] | ||
| for group_id, watermark in watermarks.items(): |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:769-820
Each attempt record (watermark) is timestamped at its group's last arrival, but the group's result is timestamped when judging finishes.
With a slow judge, the watermark therefore expires before the result it protects.
In that window, a request for an older attempt of the same group is no longer recognized as stale.
It is accepted as new work, starts a fresh group, and waits out the collection timeout.
The window is as long as the judging time, so up to cohort_evaluation_timeout_s.
The test below reproduces this on the current head: the stale attempt-0 request is accepted instead of receiving a 409.
The same ordering also affects pruning.
A group that is still collecting or judging cannot be evicted, but its watermark stays at the front of the eviction order.
If such groups are past the retention time and any finished group exists behind them, every arrival walks all of them.
A counting script visited 1,001 watermarks per prune with 1,000 such groups.
Both problems go away with a simpler structure, and it also removes active_cohort and _active_group_count:
- Keep a separate
OrderedDictof idle groups, meaning groups whose latest attempt has finished. Stamp each entry when that attempt finishes or is replayed. - Remove a group from that order when it starts a new attempt, and prune only from the front of it.
- To retire an older attempt, look up its cohort by its
group_id::…::group_attempt::Nkey instead of following a pointer.
A watermark then lives at least as long as the newest result it protects.
Pruning visits only the entries it evicts.
The manually maintained counter and its bookkeeping are no longer needed.
This was checked on a copy of this branch:
- The genrm_compare suite passes, and the lock-invariant test from the other comment finds no violations.
- A randomized state checker with 200 seeds found no inconsistency between active groups, idle groups, and watermarks. Mutated variants of the change fail it.
- 17 existing tests failed, all because they read
active_cohortor_active_group_countdirectly. Each has a behavior-level rewrite that passes.
Tested diff for app.py (also includes the late-duplicate fix from the other comment)
diff --git a/resources_servers/genrm_compare/app.py b/resources_servers/genrm_compare/app.py
index b41d217..7c0e9ec 100644
--- a/resources_servers/genrm_compare/app.py
+++ b/resources_servers/genrm_compare/app.py
@@ -124,8 +124,6 @@ class _GroupAttemptWatermark:
latest_attempt: int
prompt_digest: str
- updated_at: float
- active_cohort: Optional[_CohortState] = None
class GenRMCompareConfig(BaseResourcesServerConfig):
@@ -306,11 +304,11 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
config: GenRMCompareConfig
_verify_cohorts: Dict[str, _CohortState] = PrivateAttr(default_factory=dict)
- # Terminal records are ordered by completion; watermarks by last accepted request.
+ # Terminal records are ordered by completion.
_terminal_cohorts: OrderedDict[str, _CohortState] = PrivateAttr(default_factory=OrderedDict)
- _latest_group_attempts: OrderedDict[str, _GroupAttemptWatermark] = PrivateAttr(default_factory=OrderedDict)
- # Cache the count so pruning an all-active registry requires no scan.
- _active_group_count: int = PrivateAttr(default=0)
+ _latest_group_attempts: Dict[str, _GroupAttemptWatermark] = PrivateAttr(default_factory=dict)
+ # Only groups whose latest attempt is terminal can be evicted, ordered by completion or last replay.
+ _idle_groups: OrderedDict[str, float] = PrivateAttr(default_factory=OrderedDict)
_cohort_registry_lock: asyncio.Lock = PrivateAttr(default_factory=asyncio.Lock)
_cohort_tasks: set[asyncio.Task] = PrivateAttr(default_factory=set)
@@ -503,7 +501,6 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
if self._closed:
raise HTTPException(status_code=503, detail="GenRM server is shutting down")
self._prune_terminal_cohorts()
- now = time.monotonic()
watermark = self._latest_group_attempts.get(body.group_id)
if watermark is not None and watermark.prompt_digest != prompt_digest:
raise HTTPException(
@@ -526,15 +523,20 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
)
# Commit the new identity only after older work is retired.
# Cancellation while acquiring its lock leaves the old watermark intact.
- watermark = _GroupAttemptWatermark(
+ self._latest_group_attempts[body.group_id] = _GroupAttemptWatermark(
latest_attempt=body.group_attempt,
prompt_digest=prompt_digest,
- updated_at=time.monotonic(),
)
- self._latest_group_attempts[body.group_id] = watermark
- else:
- watermark.updated_at = now
- self._latest_group_attempts.move_to_end(body.group_id)
+ elif prompt_key not in self._verify_cohorts:
+ # This attempt already finished and its tombstone was evicted.
+ # Recreating it would wait out the collection deadline or judge a duplicate group.
+ raise HTTPException(
+ status_code=409,
+ detail=(
+ f"GenRM group {body.group_id!r} attempt {body.group_attempt} already finished and its "
+ f"result is no longer retained; retry with a higher {GROUP_ATTEMPT_KEY_NAME}"
+ ),
+ )
cohort = self._verify_cohorts.get(prompt_key)
if cohort is None:
@@ -545,8 +547,11 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
group_attempt=body.group_attempt,
)
self._verify_cohorts[prompt_key] = cohort
- watermark.active_cohort = cohort
- self._active_group_count += 1
+ self._idle_groups.pop(body.group_id, None)
+ elif cohort.terminal_at is not None:
+ # A replay keeps the fence for this finished attempt alive.
+ self._idle_groups.pop(body.group_id, None)
+ self._idle_groups[body.group_id] = time.monotonic()
return cohort
async def _supersede_older_group_attempts(
@@ -557,8 +562,10 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
) -> None:
"""Release waiters and payloads owned by older active attempts."""
watermark = self._latest_group_attempts.get(group_id)
- cohort = watermark.active_cohort if watermark is not None else None
- if cohort is None or cohort.group_attempt >= new_attempt:
+ if watermark is None:
+ return
+ cohort = self._verify_cohorts.get(self._group_cohort_key(group_id, watermark.latest_attempt))
+ if cohort is None:
return
async with cohort.lock:
# The final member may have started judging while this lock was contended.
@@ -775,16 +782,13 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
if cohort.group_id is not None or cohort.phase == "failed":
self._terminal_cohorts[cohort.key] = cohort
watermark = self._latest_group_attempts.get(cohort.group_id)
- if watermark is not None and watermark.active_cohort is cohort:
- watermark.active_cohort = None
- self._active_group_count -= 1
+ if watermark is not None and watermark.latest_attempt == cohort.group_attempt:
+ self._idle_groups[cohort.group_id] = cohort.terminal_at
def _prune_terminal_cohorts(self) -> None:
"""Evict the oldest eligible records, stopping at the first retained deadline.
- Active attempts keep their watermarks even past the retention deadline.
- They may be skipped at the head, but terminal records need no full scan
- or sort on every arrival.
+ Active groups are absent from the idle order, so they keep their watermarks without being scanned.
"""
now = time.monotonic()
ttl = self.config.cohort_result_ttl_s
@@ -799,25 +803,13 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
if self._verify_cohorts.get(key) is cohort:
self._verify_cohorts.pop(key)
- watermarks = self._latest_group_attempts
- # If every watermark owns active work, nothing can be evicted.
- prunable = len(watermarks) - self._active_group_count
- if prunable <= 0:
- return
- excess = len(watermarks) - limit
- evicted = []
- for group_id, watermark in watermarks.items():
- expired = ttl is not None and now - watermark.updated_at >= ttl
- if not expired and excess <= 0:
+ while self._idle_groups:
+ group_id, idle_since = next(iter(self._idle_groups.items()))
+ expired = ttl is not None and now - idle_since >= ttl
+ if not expired and len(self._latest_group_attempts) <= limit:
break
- if watermark.active_cohort is None:
- evicted.append(group_id)
- excess -= 1
- prunable -= 1
- if prunable == 0:
- break
- for group_id in evicted:
- watermarks.pop(group_id)
+ self._idle_groups.popitem(last=False)
+ del self._latest_group_attempts[group_id]
def setup_webserver(self) -> FastAPI:
app = super().setup_webserver()
@@ -851,7 +843,7 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
self._verify_cohorts.clear()
self._terminal_cohorts.clear()
self._latest_group_attempts.clear()
- self._active_group_count = 0
+ self._idle_groups.clear()
def _get_verify_cohort_key(
self,
@@ -861,7 +853,7 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
) -> str:
"""Return an attempt-scoped key so replacement cohorts cannot mix with old responses."""
if body.group_id is not None:
- return f"group_id::{body.group_id}::group_attempt::{body.group_attempt}"
+ return self._group_cohort_key(body.group_id, body.group_attempt)
prompt_key = get_prompt_key_from_input(input_messages, principle)
if body.task_index is not None:
@@ -872,6 +864,10 @@ class GenRMCompareResourcesServer(SimpleResourcesServer):
logical_key = prompt_key
return f"{logical_key}::group_attempt::{body.group_attempt}"
+ @staticmethod
+ def _group_cohort_key(group_id: str, group_attempt: int) -> str:
+ return f"group_id::{group_id}::group_attempt::{group_attempt}"
+
async def _run_compare(
self,
conversation_history: List[Dict[str, str]],
Regression test for the stale-attempt window
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import asyncio
from types import SimpleNamespace
import pytest
from fastapi import HTTPException
import resources_servers.genrm_compare.app as genrm
from resources_servers.genrm_compare.tests.test_cohort_lifecycle import member
async def test_stale_attempt_stays_fenced_while_newer_result_is_retained(server, monkeypatch):
clock = [100.0]
monkeypatch.setattr(genrm, "time", SimpleNamespace(monotonic=lambda: clock[0]))
server.config.cohort_result_ttl_s = 10
release = asyncio.Event()
async def compare(*, response_objs, **kwargs):
await release.wait()
return [3.0, 3.0], {}, [], []
server._run_compare = compare
latest = [asyncio.create_task(server.verify(member(i, attempt=1))) for i in range(2)]
await asyncio.sleep(0.01)
clock[0] = 108 # Judging finishes eight seconds after the last arrival.
release.set()
assert [r.reward for r in await asyncio.gather(*latest)] == [3.0, 3.0]
clock[0] = 111 # The attempt-1 result is still retained until 118.
with pytest.raises(HTTPException) as error:
await asyncio.wait_for(server.verify(member(0, attempt=0)), 0.1)
assert error.value.status_code == 409 and "superseded" in error.value.detail| if not cohort.members: | ||
| cohort.conversation_history = _input_to_conversation_history(input_messages) | ||
| cohort.principle = principle | ||
| except Exception as error: |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:391-402
If converting one member's input fails, this fails the whole group.
Every peer then receives a 503, and callers usually treat 503 as retryable, even though the failure comes from the input.
For a legacy group without _ng_group_id, the failed key also stays fenced until its time-to-live expires.
main produced the same outcome later, when evaluation failed, and validated HTTP input cannot reach this path.
So this is a small point.
It differs from the description's statement that rejected requests do not fail an existing group, though.
Could the description mention this case?
The alternative is to reject only the offending request with a 422 and leave the group collecting.
This path also logs twice, once from logger.exception() and once from the warning inside _fail_verify_cohort_locked(). One of the two is probably enough.
| """ | ||
| reasoning, answer = extract_from_response_obj(response) | ||
| return { | ||
| "id": response.get("id") if isinstance(response, dict) else response.id, |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:604-617
The docstring says this keeps only scoring text, but the compact dict also keeps id.
Nothing in production reads that id. Only test_app.py and test_response_compaction.py read it.
Could you either note that id is kept for diagnostics, or remove it and have the tests check the reasoning and answer text instead?
| else: | ||
| watermark.updated_at = now | ||
| self._latest_group_attempts.move_to_end(body.group_id) |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:506-549
A late exact duplicate for an attempt that already finished can hang for the full collection timeout.
Here is the sequence:
- Group
a, attempt 0, completes. Its tombstone and its watermark are both retained. - The tombstone is evicted, but the watermark survives. This happens because a replay refreshes the watermark but not the tombstone, or because a retried group uses two tombstones but only one watermark.
- A late duplicate for
a, attempt 0, arrives. The watermark says attempt 0 is current, so theelsebranch refreshes it. - No cohort exists under the key, so a new collecting cohort is created for an attempt that already finished.
The request then waits cohort_collection_timeout_s (1800 s by default) and gets a 503 saying the group did not collect enough rollouts.
If every slot is replayed, the group is judged a second time and new rewards are published for the old attempt.
A script reproduced all of these cases with max_terminal_cohorts=2 and with a short TTL.
main behaves the same way, but this PR rewrites exactly this retention and watermark code, so it seems like the right place to close it.
An active cohort is never removed from _verify_cohorts.
So when the attempt matches and no cohort exists under the key, the attempt must have finished and been evicted.
Rejecting it immediately fixes both outcomes:
| else: | |
| watermark.updated_at = now | |
| self._latest_group_attempts.move_to_end(body.group_id) | |
| elif prompt_key not in self._verify_cohorts: | |
| # This attempt already finished and its tombstone was evicted. | |
| # Recreating it would wait out the collection timeout or judge the group again. | |
| raise HTTPException( | |
| status_code=409, | |
| detail=( | |
| f"GenRM group {body.group_id!r} attempt {body.group_attempt} already finished and its " | |
| f"result is no longer retained; retry with a higher {GROUP_ATTEMPT_KEY_NAME}" | |
| ), | |
| ) | |
| else: | |
| watermark.updated_at = now | |
| self._latest_group_attempts.move_to_end(body.group_id) |
The rejection happens before updated_at and move_to_end(), so rejected duplicates do not keep the watermark alive.
A first arrival for a new attempt still takes the first branch, and it registers its cohort in the same critical section.
With the idle-group restructuring in the pruning comment, the else branch goes away and only the elif remains.
With this change the existing suite passes, and the regression test below fails on the current head with a timeout.
Proposed regression test
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from fastapi import HTTPException
import resources_servers.genrm_compare.app as genrm
from resources_servers.genrm_compare.tests.test_cohort_lifecycle import member
@pytest.fixture
def clock(monkeypatch):
now = [100.0]
monkeypatch.setattr(genrm, "time", SimpleNamespace(monotonic=lambda: now[0]))
return now
async def test_late_duplicate_after_tombstone_expiry_is_rejected_without_new_cohort(server, clock):
server.config.cohort_result_ttl_s = 10
server._run_single_comparison = AsyncMock(return_value=(3.0, 3.0, 3.5))
await asyncio.gather(*(server.verify(member(i, group="a")) for i in range(2)))
clock[0] = 108
assert (await server.verify(member(0, group="a"))).reward == 3 # Refreshes only the watermark.
clock[0] = 111
calls = server._run_single_comparison.await_count
with pytest.raises(HTTPException) as error:
await asyncio.wait_for(server.verify(member(0, group="a")), 0.1)
assert error.value.status_code == 409
assert not server._verify_cohorts
assert server._run_single_comparison.await_count == calls
# A newer attempt for the same group still starts normally.
rewards = await asyncio.gather(*(server.verify(member(i, group="a", attempt=1)) for i in range(2)))
assert [r.reward for r in rewards] == [3.0, 3.0]| def _response_digest(response: Any) -> str: | ||
| """Hash the exact response payload whose tokens will receive the reward.""" |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:620-625
This PR's goal is to cut per-request CPU and memory, and this digest is now the largest remaining per-request cost.
It serializes the full response with json.dumps(), including every token ID and logprob, on the event loop.
For an answer with 65,536 generated token IDs, 16,384 prompt IDs, and 65,536 logprobs, it measured about 23 ms per request.
About 19 ms of that is json.dumps() formatting the logprob floats.
Across 8,192 answers, that is roughly 150 s during which the event loop cannot serve any other request.
orjson is already a core nemo-gym dependency, and orjson.dumps(..., option=orjson.OPT_SORT_KEYS) brings the digest down to about 5 ms.
Sorting keys matters here.
NeMoGymResponse allows extra fields, so two equal responses can list their extras or metadata keys in a different order.
A digest based on model_dump_json() would then differ, and a legitimate retry would get a false 409.
orjson raises on lone surrogates and integers wider than 64 bits, which test_response_digest_accepts_surrogates covers, so it needs a fallback to the current path.
The body of _response_digest() would become:
"""Hash the exact response payload whose tokens will receive the reward."""
payload = response.model_dump(mode="json") if hasattr(response, "model_dump") else response
try:
# orjson formats the token log-prob floats much faster than json.dumps.
canonical = orjson.dumps(payload, option=orjson.OPT_SORT_KEYS)
except TypeError:
# orjson rejects lone surrogates and integers wider than 64 bits.
canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True).encode("utf-8")
return hashlib.sha256(canonical).hexdigest()This also needs import orjson at the top of the file.
The digest is only compared within one process, so changing its byte format is safe.
Two things were checked with this change:
- Reordered extras still produce equal digests, and model and dict inputs still match.
- The existing suite passes.
Moving the digest to a thread would not help much, because JSON serialization holds the GIL.
Keep latest-attempt retention aligned with completion, reject retries whose retained result expired, and evict the newest result with its attempt record. Prune idle groups without scanning active ones. Remove unsupported task-start failure branches and artificial contention tests; enforce non-suspending bookkeeping and accelerate canonical response fingerprints with compatibility fallbacks. Suggested-by: Ananth Subramaniam Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
|
/ok to test 3aaa569 |
|
Thanks Ananth—addressed all nine comments in
For input-conversion failures, I kept the existing whole-group failure behavior and documented the exception you identified, rather than introducing the alternative 422 behavior. Validation passed: 269 GenRM tests, 253 related core tests, real-model GPU checks covering scoring and retries, and CI. The existing Ray shutdown issue remains documented separately. Could you take another look when you have time? |
| if not expired and len(self._latest_group_attempts) <= limit: | ||
| break | ||
| self._idle_groups.popitem(last=False) | ||
| watermark = self._latest_group_attempts.pop(group_id) | ||
| # Replays refresh idle order, but not result expiry. Count eviction | ||
| # can therefore select a group whose newest result is still retained. | ||
| # Remove that result too, rather than expose it without its attempt fence. | ||
| key = self._group_cohort_key(group_id, watermark.latest_attempt) | ||
| self._terminal_cohorts.pop(key, None) | ||
| self._verify_cohorts.pop(key, None) |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:810-822
This loop stops evicting only when len(self._latest_group_attempts) <= limit, and that count includes groups that are still collecting or judging.
Active groups are never evicted, so once they reach max_terminal_cohorts, every group that finishes is evicted at the next arrival.
The new lines below the check also delete that group's saved result.
An exact retry of a group that just finished then finds neither an attempt record nor a result.
A single-member retry starts a new group and waits out cohort_collection_timeout_s.
A whole-group retry is judged a second time.
This is the outcome the late-duplicate fix set out to remove.
Two regression tests with max_terminal_cohorts=2 and three collecting groups reproduce both cases on this head, and both pass on main and on 05baf2c.
main also counted active groups here, but it capped results separately, so exact retries still got their saved score.
The fix has two parts.
First, count only finished groups against the limit.
Active groups are already bounded by in-flight work, and the config docstring describes max_terminal_cohorts as a cap on finished groups.
Second, stop deleting the result when its attempt record is evicted.
With the change suggested on _supersede_older_group_attempts(), the only result a group can still have is its newest one, so keeping it replayable exposes nothing stale.
| if not expired and len(self._latest_group_attempts) <= limit: | |
| break | |
| self._idle_groups.popitem(last=False) | |
| watermark = self._latest_group_attempts.pop(group_id) | |
| # Replays refresh idle order, but not result expiry. Count eviction | |
| # can therefore select a group whose newest result is still retained. | |
| # Remove that result too, rather than expose it without its attempt fence. | |
| key = self._group_cohort_key(group_id, watermark.latest_attempt) | |
| self._terminal_cohorts.pop(key, None) | |
| self._verify_cohorts.pop(key, None) | |
| # Active groups do not count toward the limit. | |
| # They are bounded by in-flight work and must not evict groups that just finished. | |
| if not expired and len(self._idle_groups) <= limit: | |
| break | |
| self._idle_groups.popitem(last=False) | |
| # A retained result stays replayable after its group's attempt record is evicted. | |
| self._latest_group_attempts.pop(group_id) |
The README paragraph about count eviction would change to match: active groups do not count toward max_terminal_cohorts, and an exact replay still returns a retained result after its attempt record is evicted.
The combined diff for all three fixes in this review, with tests, is in the review summary.
| cohort = self._verify_cohorts.get(self._group_cohort_key(group_id, watermark.latest_attempt)) | ||
| if cohort is None: | ||
| return | ||
| async with cohort.lock: | ||
| # Registry operations acquire cohort locks, so no cohort lock section | ||
| # may suspend: one slow group must not block every group's arrivals. | ||
| self._fail_verify_cohort_locked( | ||
| cohort, | ||
| (f"GenRM group {group_id!r} attempt {cohort.group_attempt} was superseded by attempt {new_attempt}"), | ||
| expected_phase=old_phase, | ||
| f"GenRM group {group_id!r} attempt {cohort.group_attempt} was superseded by attempt {new_attempt}", | ||
| ) |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:551-563
When a completed attempt is superseded, _fail_verify_cohort_locked() does nothing because the cohort is already terminal.
Its saved result therefore stays in _terminal_cohorts and _verify_cohorts.
While the group's attempt record exists, a request for the old attempt gets a 409 before any lookup, so this result is unreachable.
It still uses one of the max_terminal_cohorts slots, and a new test shows it pushing out another group's replayable result.
Once the group's attempt record is evicted, a request for the old attempt finds no record.
It creates a new record at the old attempt number, finds the old completed result, and replays that stale reward without calling the judge.
main has the same gap, so this is not a regression, but the fix is two lines.
Dropping the superseded result also means each group keeps at most one result, which makes the cleanup change in the other comment safe.
| cohort = self._verify_cohorts.get(self._group_cohort_key(group_id, watermark.latest_attempt)) | |
| if cohort is None: | |
| return | |
| async with cohort.lock: | |
| # Registry operations acquire cohort locks, so no cohort lock section | |
| # may suspend: one slow group must not block every group's arrivals. | |
| self._fail_verify_cohort_locked( | |
| cohort, | |
| (f"GenRM group {group_id!r} attempt {cohort.group_attempt} was superseded by attempt {new_attempt}"), | |
| expected_phase=old_phase, | |
| f"GenRM group {group_id!r} attempt {cohort.group_attempt} was superseded by attempt {new_attempt}", | |
| ) | |
| key = self._group_cohort_key(group_id, watermark.latest_attempt) | |
| cohort = self._verify_cohorts.get(key) | |
| if cohort is None: | |
| return | |
| async with cohort.lock: | |
| # Registry operations acquire cohort locks, so no cohort lock section | |
| # may suspend: one slow group must not block every group's arrivals. | |
| self._fail_verify_cohort_locked( | |
| cohort, | |
| f"GenRM group {group_id!r} attempt {cohort.group_attempt} was superseded by attempt {new_attempt}", | |
| ) | |
| # The new attempt record rejects this attempt from now on. | |
| # Drop its result too, so evicting that record later cannot expose the result to a stale request. | |
| self._terminal_cohorts.pop(key, None) | |
| self._verify_cohorts.pop(key, None) |
The pops run after the lock is released with no await in between, so they happen together with the new attempt record.
| try: | ||
| if isinstance(response, dict): | ||
| # Direct Python inputs can contain non-finite numbers. Unlike | ||
| # validated JSON-mode models, these have not normalized them to null. | ||
| pending = [payload] | ||
| while pending: | ||
| value = pending.pop() | ||
| if isinstance(value, dict): | ||
| pending.extend(value.values()) | ||
| elif isinstance(value, (list, tuple)): | ||
| pending.extend(value) | ||
| elif isinstance(value, float) and not isfinite(value): | ||
| raise TypeError("preserve non-finite Python input") | ||
| canonical = orjson.dumps(payload, option=orjson.OPT_SORT_KEYS) | ||
| except TypeError: | ||
| # Preserve support for lone surrogates, wide integers and non-finite | ||
| # Python inputs; orjson otherwise rejects or normalizes these values. | ||
| canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True).encode("utf-8") | ||
| return hashlib.sha256(canonical).hexdigest() |
There was a problem hiding this comment.
resources_servers/genrm_compare/app.py:615-636
orjson writes NaN, inf, and -inf as null.
My earlier suggestion missed this, and the scan added here only covers dict inputs.
The comment says validated models have already turned non-finite values into null, but they have not.
model_dump(mode="json") keeps nan and inf as floats in typed fields such as generation_log_probs.
So on the HTTP path, where the response is always a model, NaN, inf, -inf, and None in the same position produce the same digest.
A retry for slot 0 carrying +inf where the original carried -inf was accepted as an exact replay and returned the saved reward, instead of a 409.
FastAPI does accept NaN and Infinity literals in a JSON body.
Gym's own client serializes with orjson and cannot send them, so only other clients can trigger this, and it is unlikely in practice.
It does break the digest's purpose of detecting conflicting retries, and a model and an equal dict now hash differently when they contain NaN.
The dict-only scan has two further problems.
It walks every element in Python, which makes dict inputs slower than before the change.
It never finishes on a dict that contains itself, where the old json.dumps() raised ValueError: Circular reference detected.
A version that checks every payload after orjson succeeds fixes all three problems:
| try: | |
| if isinstance(response, dict): | |
| # Direct Python inputs can contain non-finite numbers. Unlike | |
| # validated JSON-mode models, these have not normalized them to null. | |
| pending = [payload] | |
| while pending: | |
| value = pending.pop() | |
| if isinstance(value, dict): | |
| pending.extend(value.values()) | |
| elif isinstance(value, (list, tuple)): | |
| pending.extend(value) | |
| elif isinstance(value, float) and not isfinite(value): | |
| raise TypeError("preserve non-finite Python input") | |
| canonical = orjson.dumps(payload, option=orjson.OPT_SORT_KEYS) | |
| except TypeError: | |
| # Preserve support for lone surrogates, wide integers and non-finite | |
| # Python inputs; orjson otherwise rejects or normalizes these values. | |
| canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True).encode("utf-8") | |
| return hashlib.sha256(canonical).hexdigest() | |
| try: | |
| canonical = orjson.dumps(payload, option=orjson.OPT_SORT_KEYS) | |
| except TypeError: | |
| # orjson rejects lone surrogates, integers wider than 64 bits, and cycles. | |
| # json.dumps handles the first two and raises a clear error for cycles. | |
| canonical = None | |
| # orjson writes NaN and infinities as null, so a retry differing only in those values would look identical. | |
| # The scan runs only after orjson succeeds, so the payload is acyclic. | |
| if canonical is None or _contains_non_finite(payload): | |
| canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=True).encode("utf-8") | |
| return hashlib.sha256(canonical).hexdigest() |
It needs this module-level helper, which checks numeric arrays such as log-probabilities in C through sum():
def _contains_non_finite(value: Any) -> bool:
"""Return whether an acyclic JSON-like value contains NaN or an infinity."""
pending = [value]
while pending:
value = pending.pop()
if isinstance(value, dict):
pending.extend(value.values())
elif isinstance(value, (list, tuple)):
try:
# Check numeric arrays, such as token log-probabilities, in C.
if not isfinite(sum(value, 0.0)):
return True
except (TypeError, OverflowError):
pending.extend(value)
elif isinstance(value, float) and not isfinite(value):
return True
return FalseThese timings are for a response with 65,536 generated token IDs and log-probabilities and 16,384 prompt IDs:
| Version | Model input | Dict input |
|---|---|---|
main |
22.5 ms | 20.0 ms |
| This head | 4.5 ms | 15.8 ms |
| This suggestion | 5.2 ms | 3.6 ms |
With this change:
- NaN,
inf, and-infstay distinct. - Models and equal dicts produce the same digest.
- A cyclic dict raises the same
ValueErroras before.
test_digest_preserves_nonstandard_python_values only covers dicts. The combined tests add a model case and a cyclic-dict case.
| server._prune_terminal_cohorts() | ||
| assert not server._verify_cohorts | ||
| assert not server._terminal_cohorts | ||
| assert all(c.phase in ("completed", "failed") for c in server._verify_cohorts.values()) |
There was a problem hiding this comment.
resources_servers/genrm_compare/tests/test_cohort_storage.py:332-334
This assertion always passes, because line 332 already asserts that server._verify_cohorts is empty.
It looks like a leftover from replacing the _active_group_count checks.
The same replacement left redundant assertions at line 241, line 396, and test_http_lifecycle.py lines 211-214.
Could you remove them, or replace them with a check that says something new, such as assert set(server._idle_groups) <= set(server._latest_group_attempts)?
| assert await asyncio.wait_for(started.get(), 1) == 0 | ||
| latest = [asyncio.create_task(server.verify(member(i, attempt=1, response_id=f"new-{i}"))) for i in range(2)] | ||
| try: | ||
| assert await asyncio.wait_for(started.get(), 1) == 1 |
There was a problem hiding this comment.
resources_servers/genrm_compare/tests/test_response_compaction.py:83-86
Before the id field was removed, this test checked that the first judge call carried the old attempt's answer and the second carried the new attempt's answer.
It now only counts calls, so it would still pass if the wrong attempt's members reached the judge.
Could it compare the answer text in started_calls[0] and started_calls[1], the way test_app.py now does?
|
Summary for the review on Thanks for the quick and thorough update. Eight of the nine earlier comments are fully resolved:
Accepting the whole-group failure for input-conversion errors makes sense, because validated HTTP requests cannot reach that path and This round found three problems in the new commit. The inline comments explain each one.
The three fixes interact, so they were applied together to a copy of this branch:
Combined test changes (adapted and new tests)--- a/resources_servers/genrm_compare/tests/test_app.py 2026-09-28 15:56:00.800801487 -0700
+++ b/resources_servers/genrm_compare/tests/test_app.py 2026-09-28 16:03:28.393103649 -0700
@@ -678,7 +678,8 @@
assert [result.reward for result in first_attempt] == [1.0, 2.0]
assert [result.reward for result in replacement_attempt] == [3.0, 4.0]
assert run_compare.await_count == 2
- assert len(server._verify_cohorts) == 2
+ # The superseded attempt's tombstone is dropped; the watermark fences it.
+ assert list(server._verify_cohorts) == ["group_id::completed-group::group_attempt::1"]
async def test_partial_old_attempt_does_not_mix_with_completed_replacement(self, config, monkeypatch: MonkeyPatch):
config = config.model_copy(update={"num_rollouts_per_prompt": 2})
@@ -701,6 +702,7 @@
)
)
await asyncio.sleep(0)
+ old_cohort = next(iter(server._verify_cohorts.values()))
replacement = await asyncio.gather(
server.verify(
self._verify_request(
@@ -726,12 +728,11 @@
assert [result.reward for result in replacement] == [3.0, 4.0]
assert [result.group_attempt for result in replacement] == [1, 1]
- assert len(server._verify_cohorts) == 2
+ assert len(server._verify_cohorts) == 1
assert len(old_result) == 1
assert isinstance(old_result[0], HTTPException)
assert old_result[0].status_code == 503
assert "superseded by attempt 1" in str(old_result[0].detail)
- old_cohort = next(cohort for cohort in server._verify_cohorts.values() if cohort.group_attempt == 0)
assert old_cohort.phase == "failed"
assert all(member.response_obj is None and not member.waiters for member in old_cohort.members.values())
--- a/resources_servers/genrm_compare/tests/test_cohort_lifecycle.py 2026-09-28 15:56:00.800869818 -0700
+++ b/resources_servers/genrm_compare/tests/test_cohort_lifecycle.py 2026-09-28 16:03:28.393334233 -0700
@@ -192,11 +192,11 @@
server._run_compare = compare
old = [asyncio.create_task(server.verify(member(i))) for i in range(2)]
await asyncio.wait_for(started.wait(), 1)
+ retired = next(iter(server._verify_cohorts.values()))
new = await asyncio.gather(*(server.verify(member(i, attempt=1)) for i in range(2)))
old_results = await asyncio.gather(*old, return_exceptions=True)
assert all(isinstance(r, HTTPException) and r.status_code == 503 for r in old_results)
assert [r.reward for r in new] == [1.0, 2.0]
- retired = next(c for c in server._verify_cohorts.values() if c.group_attempt == 0)
assert retired.phase == "failed" and not retired.rewards
--- a/resources_servers/genrm_compare/tests/test_cohort_storage.py 2026-09-28 15:56:00.800895366 -0700
+++ b/resources_servers/genrm_compare/tests/test_cohort_storage.py 2026-09-28 16:04:34.885711278 -0700
@@ -222,7 +222,7 @@
assert list(server._latest_group_attempts) == ["a"]
-async def test_count_eviction_uses_completion_order_and_protects_active_attempts(server, clock):
+async def test_count_eviction_uses_completion_order_and_ignores_active_attempts(server, clock):
server.config.cohort_result_ttl_s = None
server.config.max_terminal_cohorts = 1
server._run_single_comparison = AsyncMock(return_value=(3.0, 3.0, 3.5))
@@ -231,9 +231,11 @@
clock[0] = 101
await complete(server, "b")
clock[0] = 102
+ await complete(server, "c")
server._prune_terminal_cohorts()
- assert "a" in server._latest_group_attempts
- assert "b" not in server._latest_group_attempts
+ # The active group neither counts toward the cap nor loses its record.
+ assert set(server._latest_group_attempts) == {"a", "c"}
+ assert (await server.verify(member(0, group="c"))).reward == 3
await server.verify(member(1, group="a"))
await a
server._prune_terminal_cohorts()
@@ -353,13 +355,17 @@
await first
-async def test_new_attempt_preserves_completed_record_and_owns_active_watermark(server):
+async def test_new_attempt_retires_completed_record_and_owns_active_watermark(server):
server._run_single_comparison = AsyncMock(return_value=(3.0, 3.0, 3.5))
await complete(server, "a")
completed = next(iter(server._verify_cohorts.values()))
first = asyncio.create_task(server.verify(member(0, group="a", attempt=1)))
await asyncio.sleep(0)
assert completed.phase == "completed" and completed.rewards == {0: 3.0, 1: 3.0}
+ assert completed not in server._verify_cohorts.values()
+ with pytest.raises(HTTPException) as error:
+ await server.verify(member(0, group="a"))
+ assert error.value.status_code == 409
active = server._verify_cohorts[server._group_cohort_key("a", 1)]
assert active is not completed and active.group_attempt == 1
assert_cohort_indices_match(server)
--- a/resources_servers/genrm_compare/tests/test_response_digest.py 2026-09-28 15:56:00.801036417 -0700
+++ b/resources_servers/genrm_compare/tests/test_response_digest.py 2026-09-28 16:03:40.009669627 -0700
@@ -39,3 +39,22 @@
digest = GenRMCompareResourcesServer._response_digest
assert digest({"value": value}) == digest({"value": value})
assert digest({"value": value}) != digest({"value": None})
+
+
+def test_digest_keeps_nonfinite_model_values_distinct_and_matches_dictionary():
+ # HTTP JSON bodies may carry NaN and Infinity; a retry that changes one must still conflict.
+ digest = GenRMCompareResourcesServer._response_digest
+ digests = set()
+ for value in (float("nan"), float("inf"), float("-inf"), -0.5):
+ response = training_member(0).response.model_copy(deep=True)
+ response.output[1].generation_log_probs = [value, *response.output[1].generation_log_probs[1:]]
+ assert digest(response) == digest(response.model_dump(mode="json"))
+ digests.add(digest(response))
+ assert len(digests) == 4
+
+
+def test_digest_rejects_cyclic_dictionary():
+ payload = {"id": "response"}
+ payload["self"] = payload
+ with pytest.raises(ValueError, match="Circular reference"):
+ GenRMCompareResourcesServer._response_digest(payload)
--- a/resources_servers/genrm_compare/tests/test_retention_boundaries.py 2026-09-28 15:56:00.801069231 -0700
+++ b/resources_servers/genrm_compare/tests/test_retention_boundaries.py 2026-09-28 16:11:27.931621967 -0700
@@ -59,9 +59,13 @@
assert (await server.verify(member(0, group="a"))).reward == 3
clock[0] = 111
else:
+ # Failed legacy groups hold result slots without attempt records, so
+ # the count cap evicts a's result while its attempt record remains.
server.config.max_terminal_cohorts = 2
- await complete(server, "b")
- await complete(server, "b", attempt=1)
+ server.config.cohort_collection_timeout_s = 0.01
+ for legacy_attempt in range(2):
+ with pytest.raises(HTTPException):
+ await server.verify(member(0, group=None, attempt=legacy_attempt))
calls = judge.await_count
for _ in range(2):
outcomes = await asyncio.wait_for(
@@ -73,9 +77,10 @@
assert [r.reward for r in await complete(server, "a", attempt=1)] == [3, 3]
-async def test_count_eviction_removes_newest_result_with_its_attempt_record(server, clock):
+async def test_count_eviction_keeps_retained_result_replayable(server, clock):
server.config.max_terminal_cohorts = 2
- server._run_single_comparison = AsyncMock(return_value=(3.0, 3.0, 3.5))
+ judge = AsyncMock(return_value=(3.0, 3.0, 3.5))
+ server._run_single_comparison = judge
await complete(server, "a")
clock[0] = 101
await complete(server, "b", attempt=1)
@@ -84,12 +89,58 @@
clock[0] = 103
await complete(server, "c")
server._prune_terminal_cohorts()
- # Replaying a refreshed its attempt record only. Evicting b's older record
- # must not leave its newer result exposed without stale-attempt protection.
+ # Replaying a moved it behind b, so the cap evicts b's attempt record while b's result is newer than a's.
assert "b" not in server._latest_group_attempts
- assert not any(c.group_id == "b" for c in server._verify_cohorts.values())
- assert len(server._latest_group_attempts) <= 2
- assert len(server._terminal_cohorts) <= 2
+ calls = judge.await_count
+ replay = await asyncio.wait_for(complete(server, "b", attempt=1), 0.1)
+ assert [r.reward for r in replay] == [3, 3]
+ assert judge.await_count == calls
+
+
+@pytest.mark.parametrize("whole_group", [False, True])
+async def test_active_groups_do_not_evict_a_finished_result(server, whole_group):
+ server.config.max_terminal_cohorts = 2
+ judge = AsyncMock(return_value=(3.0, 3.0, 3.5))
+ server._run_single_comparison = judge
+ active = [asyncio.create_task(server.verify(member(0, group=f"active-{i}"))) for i in range(3)]
+ await asyncio.sleep(0)
+ try:
+ assert [r.reward for r in await complete(server, "done")] == [3, 3]
+ calls = judge.await_count
+ indices = range(2) if whole_group else [0]
+ retry = await asyncio.wait_for(asyncio.gather(*(server.verify(member(i, group="done")) for i in indices)), 0.1)
+ assert [r.reward for r in retry] == [3] * len(indices)
+ assert judge.await_count == calls
+ finally:
+ for task in active:
+ task.cancel()
+ await asyncio.gather(*active, return_exceptions=True)
+
+
+async def test_superseded_result_neither_uses_a_slot_nor_replays_after_eviction(server, clock):
+ server.config.max_terminal_cohorts = 3
+ server.config.cohort_collection_timeout_s = 0.05
+ judge = AsyncMock(return_value=(3.0, 3.0, 3.5))
+ server._run_single_comparison = judge
+ await complete(server, "x1")
+ await complete(server, "x2")
+ clock[0] = 101
+ await complete(server, "a")
+ await complete(server, "a", attempt=1)
+ # a's superseded result must not displace x1's result.
+ assert (await server.verify(member(0, group="x1"))).reward == 3
+ clock[0] = 102
+ await server.verify(member(0, group="x2"))
+ clock[0] = 103
+ await complete(server, "y")
+ server._prune_terminal_cohorts()
+ # Replays moved a to the front of the eviction order, ahead of results completed before a.
+ assert "a" not in server._latest_group_attempts
+ calls = judge.await_count
+ with pytest.raises(HTTPException) as error:
+ await server.verify(member(0, group="a"))
+ assert error.value.status_code == 503
+ assert judge.await_count == calls
async def test_rejected_expired_duplicate_does_not_extend_attempt_retention(server, clock):
Two smaller test comments point out assertions that became vacuous or weaker when Design visualization for this revision: https://terryk.gitlab-master-pages.nvidia.com/nemo-html/ansubramania/gym-pr-reviews/pr-3625/design-3aaa569b1063.html · compare revisions: https://terryk.gitlab-master-pages.nvidia.com/nemo-html/ansubramania/gym-pr-reviews/pr-3625/index.html Generated by Claude Code |
Count only finished groups against the attempt-record cap and keep independently retained results replayable. Drop obsolete results when tracked attempts are superseded. Preserve non-finite values in response fingerprints for model and dictionary inputs, reject cyclic dictionaries, and strengthen retry, HTTP, and scoring-text regression coverage. Suggested-by: Ananth Subramaniam Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
|
/ok to test e5d02fe |
Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
|
/ok to test 416ccd7 |
Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
|
/ok to test 8c5b1f4 |
Signed-off-by: Yuhe Zhang <yuhez@nvidia.com>
|
/ok to test 3194f34 |
|
Thanks for the review, @ananthsub! I’ve addressed your comments and added a small fix so an invalid retry cannot block later valid retries. All 280 GenRM tests and CI checks pass, and a fresh GPU run confirmed that correct retries receive their saved scores without additional judging. The PR description is updated—could you take another look? |
Summary
GenRM judges several answers to the same question as a group. This follow-up to merged #3351 reduces the CPU and memory overhead around that judging:
Addresses #3026's response-copying and cleanup work. It also adds regression coverage for #3027's comparison-completion race; the scheduling fix itself came from merged #2814. Scoring and the policy of waiting for every answer before judging are unchanged.
Retry behavior and compatibility
An exact retry of the current attempt receives its saved score while that result is retained. If Gym still remembers the attempt but has already removed its result, the retry now receives 409 immediately, without starting another group or judge call. Recovery requires the caller to dispatch a complete group with a higher shared
_ng_group_attempt.Time-based retention of the newest-attempt record starts when the group finishes. Active groups keep their attempt records and do not count toward the finished-group limit. Results and finished-group attempt records have separate count limits, so an exact retry can still receive a retained result after its attempt record is evicted. Accepting a newer tracked attempt removes the obsolete result.
Retention remains bounded. After an attempt record is evicted or the server restarts, callers must enforce accepted attempt identity themselves and use fresh group IDs for unrelated work. This is an in-memory cache, not durable recovery.
/compare, including dictionary-shaped direct Python inputs. New members are compacted once; history is converted once per group.Validation
Current head:
3194f3427b7edd71db485712bb621bd49dad2a7d.gym env test +entrypoint=resources_servers/genrm_compare +should_validate_data=truein the prepared Python 3.13.14 environment.pre-commit run --all-files, whitespace checks and DCO sign-off passed. Independent Codex and Claude Opus 5.5 reviews found no blocking regressions in the latest fix.[1.5, 1.0]and full responses, with zero additional judge calls. All four model request/reply pairs were audited, generations ended normally, judge scores parsed strictly, and source hashes matched this commit.Earlier validation at
e5d02fe27covered 253 core serialization/judge tests and a broader real-model run: 16 fresh Qwen3-0.6B answers and 38 comparisons through the TP8 235B GenRM judge. That run checked N=16 scoring/response echo and N=2 replay, active-group pressure, supersession, expiry rejection and recovery with a higher attempt. All 54 raw model request/reply pairs were audited. These are prior-run results, not reruns on the current head.Full RL training, production throughput and replica scaling remain untested.
Performance evidence and limits
On runtime head
e5d02fe27, a synthetic response with 65,536 generated token IDs/log probabilities and 16,384 prompt IDs took 24.8 ms → 6.0 ms to fingerprint as a validated model, and 22.9 ms → 4.1 ms as a dictionary, compared with the previous standard-library serialization path. These are medians of ten alternating measurements after warmup on one CPU node. They measure fingerprinting only, not HTTP handling, model inference or training throughput.An additional synthetic check included nested expert-routing records (32 layers and eight expert IDs per generated token). Model fingerprinting took 79.2 → 58.8 ms at 4,096 generated tokens and 338.7 → 255.6 ms at 16,384 tokens, versus the original standard-library path. The non-finite scan alone cost 23.7 and 90.1 ms, respectively: the improvement remains, but deeply nested metadata still warrants profiling on the real workload before further optimization. These are ten-sample alternating medians after warmup, with 16,384 prompt tokens; no production throughput was measured.
Earlier end-to-end CPU preparation measurements compared #3351 (
6a7f30b70) with the earlier #3625 implementation (3bafd2779):Those are historical medians of two fresh-process runs, not a rerun of the current head. Every group was held at the simulated judge together, and memory included the original requests. The saving mainly removes extra serialized copies during judging; connected requests still retain their original answers while waiting. Disconnected requests can release their full responses while compact scoring inputs remain. The large case is a synthetic stress workload, not a production length distribution.
We have not tested production RL throughput. Early judging remains separate work; the earlier small experiment was inconclusive.
Internal evidence
Latest invalid-retry fix, regression checks, independent review and fresh model validation:
Earlier independent review, wording follow-up and nested-metadata CPU measurements:
Runtime fixes, regression checks, CPU measurement and real-model validation:
Historical copying/memory measurements and model evidence: