Add rollout hooks and tests (#25)
This commit is contained in:
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user