Skip to content
Closed
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
28 changes: 26 additions & 2 deletions verifiers/v1/rollout.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,10 @@
import asyncio
import logging
import time
from collections.abc import AsyncIterator
from collections.abc import AsyncIterator, Awaitable
from contextlib import asynccontextmanager
from enum import StrEnum
from typing import Any

from verifiers import __version__
from verifiers.v1.harness import Harness
Expand Down Expand Up @@ -32,11 +33,34 @@
from verifiers.v1.state import state_cls
from verifiers.v1.task import Task
from verifiers.v1.trace import AgentInfo, Trace, TraceTask, VersionInfo
from verifiers.v1.utils.aio import run_shielded
from verifiers.v1.utils.version import verifiers_commit

logger = logging.getLogger(__name__)


async def gather_scoring(*awaitables: Awaitable[Any]) -> list[Any]:
"""Run scoring handlers concurrently and drain every sibling after failure."""
if len(awaitables) == 1 and isinstance(awaitables[0], (list, tuple)):
awaitables = tuple(awaitables[0])
tasks = [asyncio.ensure_future(awaitable) for awaitable in awaitables]
if not tasks:
return []
try:
return await asyncio.gather(*(asyncio.shield(task) for task in tasks))
except BaseException as error:
for task in tasks:
if not task.done():
task.cancel()
try:
await run_shielded(asyncio.gather(*tasks, return_exceptions=True))
except asyncio.CancelledError:
raise
except BaseException:
pass
raise error


class Phase(StrEnum):
PENDING = "pending"
BOOT = "boot"
Expand Down Expand Up @@ -226,7 +250,7 @@ async def run(self) -> Trace:
async with boundary(TaskError, "scoring"):
# Group rewards run later, after the runtime is gone.
await asyncio.wait_for(
asyncio.gather(
gather_scoring(
self.task.score(trace, runtime),
self.harness.score(trace, runtime),
),
Expand Down
Loading