Add rollout hooks and tests (#25)

This commit is contained in:
Yuge Zhang
2025-09-01 16:33:09 +08:00
committed by GitHub
parent a54c6f39e8
commit b9595b8b0d
3 changed files with 143 additions and 0 deletions
+26
View File
@@ -94,6 +94,32 @@ class LitAgent:
raise ValueError("Runner reference is no longer valid (object has been garbage collected).")
return runner
def on_rollout_start(self, task: Task, runner: AgentRunner, tracer: BaseTracer) -> None:
"""Hook called immediately before a rollout begins.
Args:
task: The :class:`Task` object that will be processed.
runner: The :class:`AgentRunner` managing the rollout.
tracer: The tracer instance associated with the runner.
Subclasses can override this method to implement custom logic such as
logging, metric collection, or resource setup. By default, this is a
no-op.
"""
def on_rollout_end(self, task: Task, rollout: Rollout, runner: AgentRunner, tracer: BaseTracer) -> None:
"""Hook called after a rollout completes.
Args:
task: The :class:`Task` object that was processed.
rollout: The resulting :class:`Rollout` object.
runner: The :class:`AgentRunner` managing the rollout.
tracer: The tracer instance associated with the runner.
Subclasses can override this method for cleanup or additional
logging. By default, this is a no-op.
"""
def training_rollout(self, task: TaskInput, rollout_id: str, resources: NamedResources) -> RolloutRawResult:
"""Defines the agent's behavior for a single training task.
+18
View File
@@ -158,6 +158,11 @@ class AgentRunner(ParallelWorkerBase):
rollout_obj = Rollout(rollout_id=task.rollout_id) # Default empty rollout
try:
try:
self.agent.on_rollout_start(task, self, self.tracer)
except Exception:
logger.exception(f"{self._log_prefix(rollout_id)} Exception during on_rollout_start hook.")
with self.tracer.trace_context(name=f"rollout_{rollout_id}"):
start_time = time.time()
rollout_method = self.agent.training_rollout if task.mode == "train" else self.agent.validation_rollout
@@ -175,6 +180,10 @@ class AgentRunner(ParallelWorkerBase):
except Exception:
logger.exception(f"{self._log_prefix(rollout_id)} Exception during rollout.")
finally:
try:
self.agent.on_rollout_end(task, rollout_obj, self, self.tracer)
except Exception:
logger.exception(f"{self._log_prefix(rollout_id)} Exception during on_rollout_end hook.")
self.client.post_rollout(rollout_obj)
return True
@@ -218,6 +227,11 @@ class AgentRunner(ParallelWorkerBase):
rollout_obj = Rollout(rollout_id=task.rollout_id) # Default empty rollout
try:
try:
self.agent.on_rollout_start(task, self, self.tracer)
except Exception:
logger.exception(f"{self._log_prefix(rollout_id)} Exception during on_rollout_start hook.")
with self.tracer.trace_context(name=f"rollout_{rollout_id}"):
start_time = time.time()
rollout_method = (
@@ -234,6 +248,10 @@ class AgentRunner(ParallelWorkerBase):
except Exception:
logger.exception(f"{self._log_prefix(rollout_id)} Exception during rollout.")
finally:
try:
self.agent.on_rollout_end(task, rollout_obj, self, self.tracer)
except Exception:
logger.exception(f"{self._log_prefix(rollout_id)} Exception during on_rollout_end hook.")
await self.client.post_rollout_async(rollout_obj)
return True
+99
View File
@@ -0,0 +1,99 @@
import pytest
from contextlib import contextmanager
from agentlightning import LitAgent, Task, ResourcesUpdate
from agentlightning.runner import AgentRunner
from agentlightning.tracer import BaseTracer, TripletExporter
class DummyTracer(BaseTracer):
@contextmanager
def trace_context(self, name=None):
yield
def get_last_trace(self):
return []
class DummyClient:
def __init__(self):
self.posted = None
self.polled = False
def poll_next_task(self):
if self.polled:
return None
self.polled = True
return Task(rollout_id="1", input={}, mode="train", resources_id=None)
def get_latest_resources(self):
return ResourcesUpdate(resources_id="r", resources={})
def post_rollout(self, rollout):
self.posted = rollout
class DummyAsyncClient:
def __init__(self):
self.posted = None
self.polled = False
async def poll_next_task_async(self):
if self.polled:
return None
self.polled = True
return Task(rollout_id="1", input={}, mode="train", resources_id=None)
async def get_latest_resources_async(self):
return ResourcesUpdate(resources_id="r", resources={})
async def post_rollout_async(self, rollout):
self.posted = rollout
class HookAgent(LitAgent):
def __init__(self):
super().__init__()
self.start_called = False
self.end_called = False
self.end_rollout = None
def training_rollout(self, task, rollout_id, resources):
return 0.5
async def training_rollout_async(self, task, rollout_id, resources):
return 0.5
def on_rollout_start(self, task, runner, tracer):
self.start_called = True
self.start_task = task
def on_rollout_end(self, task, rollout, runner, tracer):
self.end_called = True
self.end_rollout = rollout
def test_runner_calls_hooks():
agent = HookAgent()
client = DummyClient()
tracer = DummyTracer()
runner = AgentRunner(agent, client, tracer, TripletExporter())
assert runner.run() is True
assert agent.start_called
assert agent.end_called
assert agent.end_rollout.final_reward == 0.5
@pytest.mark.asyncio
async def test_runner_calls_hooks_async():
agent = HookAgent()
client = DummyAsyncClient()
tracer = DummyTracer()
runner = AgentRunner(agent, client, tracer, TripletExporter())
assert await runner.run_async() is True
assert agent.start_called
assert agent.end_called
assert agent.end_rollout.final_reward == 0.5