Skip to content
Open
Show file tree
Hide file tree
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
1 change: 1 addition & 0 deletions agentlightning/config/server.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ key: ""
default_proxy:
model_name: "Qwen/Qwen2.5-7B-Instruct"
include_log_probs: True
include_routed_experts: False
train:
temperature: 1
val:
Expand Down
17 changes: 17 additions & 0 deletions agentlightning/server/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@ def __init__(self, default_proxy: Mapping[str, Any]) -> None:
self._train_temperature = float(default_proxy["train"]["temperature"])
self._val_temperature = float(default_proxy["val"]["temperature"])
self._include_log_probs = bool(default_proxy.get("include_log_probs", True))
self._include_routed_experts = bool(default_proxy.get("include_routed_experts", False))

@property
def model_name(self) -> str:
Expand All @@ -69,6 +70,8 @@ def prepare_body(self, body: dict[str, Any], mode: str) -> dict[str, Any]:
}
if self._include_log_probs:
prepared["logprobs"] = True
if self._include_routed_experts:
prepared["return_routed_experts"] = True
return prepared
if mode == "val":
prepared = {
Expand Down Expand Up @@ -126,6 +129,7 @@ async def forward_request(
response_body = (
response.json() if response.headers.get("content-type", "").startswith("application/json") else {}
)
routed_experts = _pop_routed_experts(response_body)

_capture_event(
rollout_id=rollout_id,
Expand All @@ -137,13 +141,24 @@ async def forward_request(
http_status=response.status_code,
status=_status_from_http_status(response.status_code),
retry_count=int(response.extensions.get("agl_retry_count", 0)),
routed_experts=routed_experts,
)
return JSONResponse(content=response_body, status_code=response.status_code)
finally:
if pause_state is not None:
await _dec_inflight(pause_state)


def _pop_routed_experts(response_body: dict[str, Any]) -> Any:
choices = response_body.get("choices")
if not isinstance(choices, list):
return None
for choice in choices:
if isinstance(choice, dict) and (routed_experts := choice.pop("routed_experts", None)) is not None:
return routed_experts
return None


async def _send_upstream_with_retries(
*,
client: httpx.AsyncClient,
Expand Down Expand Up @@ -216,6 +231,7 @@ def _capture_event(
http_status: int,
status: str,
retry_count: int,
routed_experts: Any = None,
) -> None:
record_event(
rollout_id,
Expand All @@ -231,6 +247,7 @@ def _capture_event(
"http_status": http_status,
"status": status,
"retry_count": retry_count,
"routed_experts": routed_experts,
"usage": _extract_usage(response_body),
"finish_reason": _extract_finish_reason(response_body),
},
Expand Down
47 changes: 43 additions & 4 deletions agentlightning/server/routes/events.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,11 +21,39 @@ def _not_found(rollout_id: str) -> HTTPException:
return HTTPException(status_code=404, detail=f"Rollout not found: {rollout_id}")


def _drop_superseded_routed_experts(events: list[Event], event_type: str, data: dict[str, Any]) -> None:
"""Keep only the last route tensor in each mergeable trajectory."""
if event_type != "model_request" or data.get("routed_experts") is None:
Comment on lines +24 to +26
return

current = _trim_model_request(data)
prompt_ids = current["prompt_token_ids"]
for index in range(len(events) - 1, -1, -1):
previous = events[index]
if previous.event_type != "model_request":
continue
previous_data = _trim_model_request(previous.data)
previous_context = previous_data["prompt_token_ids"] + previous_data["response_token_ids"]
same_prompt = prompt_ids == previous_data["prompt_token_ids"]
extends_context = prompt_ids[: len(previous_context)] == previous_context
if (same_prompt or extends_context) and "routed_experts" in previous.data:
data_without_routes = dict(previous.data)
data_without_routes.pop("routed_experts")
events[index] = previous.model_copy(update={"data": data_without_routes})
return


def record_event(rollout_id: str, attempt_id: str, event_type: str, data: dict[str, Any]) -> Event:
"""Append a single event for an existing rollout."""
if rollout_id not in _rollouts:
raise _not_found(rollout_id)

rid_events = _events[rollout_id]
if attempt_id not in rid_events:
rid_events[attempt_id] = []
attempt_events = rid_events[attempt_id]
_drop_superseded_routed_experts(attempt_events, event_type, data)

event = Event(
event_type=event_type,
rollout_id=rollout_id,
Expand All @@ -34,10 +62,7 @@ def record_event(rollout_id: str, attempt_id: str, event_type: str, data: dict[s
data=data,
)

rid_events = _events[rollout_id]
if attempt_id not in rid_events:
rid_events[attempt_id] = []
rid_events[attempt_id].append(event)
attempt_events.append(event)
return event


Expand Down Expand Up @@ -109,6 +134,7 @@ def _trim_model_request(data: dict[str, Any]) -> dict[str, Any]:
prompt_token_ids: list[int] = []
response_token_ids: list[int] = []
response_log_probs: list[float] | None = None
routed_experts = data.get("routed_experts")

if isinstance(resp, dict):
prompt_token_ids = resp.get("prompt_token_ids", [])
Expand Down Expand Up @@ -136,6 +162,8 @@ def _trim_model_request(data: dict[str, Any]) -> dict[str, Any]:
"response_log_probs": response_log_probs,
"server": {"model": srv.get("model"), "version": srv.get("version")},
}
if routed_experts is not None:
trimmed["routed_experts"] = routed_experts
for key in ("http_status", "status"):
if key in data:
trimmed[key] = data[key]
Expand Down Expand Up @@ -169,6 +197,14 @@ def _to_triplet_format(event: Event) -> Event:
return event


def _without_routed_experts(event: Event) -> Event:
if event.event_type != "model_request" or "routed_experts" not in event.data:
return event
data = dict(event.data)
data.pop("routed_experts")
return event.model_copy(update={"data": data})


def _dedupe_model_requests_by_prompt_token_ids(events: list[Event]) -> list[Event]:
"""Keep the last request for each valid, non-empty prompt-token key."""
last_index_by_prompt: dict[tuple[int, ...], int] = {}
Expand Down Expand Up @@ -202,6 +238,7 @@ async def query_events(
rollout_id: str,
event_type: str | None = None,
format: str | None = Query(None, description="Set to 'triplet' to trim events for RL training"),
include_routed_experts: bool = True,
) -> list[Event]:
"""Query events for the default rollout attempt."""
events = _query_events(
Expand All @@ -211,4 +248,6 @@ async def query_events(
if format == "triplet":
events = [_to_triplet_format(e) for e in events]
events = _dedupe_model_requests_by_prompt_token_ids(events)
if not include_routed_experts:
events = [_without_routed_experts(event) for event in events]
return events
33 changes: 28 additions & 5 deletions agentlightning/verl/agl_rollout_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,14 @@ def _to_native(obj: Any) -> Any:
return obj


def _without_routed_experts(event: Event) -> Event:
if event.event_type != "model_request" or "routed_experts" not in event.data:
return event
data = dict(event.data)
data.pop("routed_experts")
return event.model_copy(update={"data": data})


# [multimodal-patch] Extract image URLs from OpenAI-style chat messages, in order of
# appearance (ported from agent-lightning v0.3.0 TripletAdapter.extract_prompt_image_urls).
def _extract_image_urls_from_messages(messages: Any) -> list[str]:
Expand Down Expand Up @@ -339,9 +347,22 @@ def _record_lifecycle_timestamps(enqueued_rollout: EnqueuedRollout, rollout: Rol
if state in TERMINAL_STATES:
enqueued_rollout.finished_at = updated_at

def _get_events(self, rollout_id: str, *, event_type: str | None = None, format: str | None = None) -> list[Event]:
def _get_events(
self,
rollout_id: str,
*,
event_type: str | None = None,
format: str | None = None,
include_routed_experts: bool = True,
) -> list[Event]:
params = {
key: value for key, value in {"event_type": event_type, "format": format}.items() if value is not None
key: value
for key, value in {
"event_type": event_type,
"format": format,
"include_routed_experts": include_routed_experts,
}.items()
if value is not None
}
response = self.client.get(f"/api/rollouts/{rollout_id}/events", params=params)
response.raise_for_status()
Expand Down Expand Up @@ -407,7 +428,7 @@ def _create_rollouts(self, data: dict[str, Any], *, is_train: bool) -> list[Enqu
]

def _fetch_rollout_events(self, rollout_id: str) -> tuple[list[Event], list[Event]]:
raw_events = self._get_events(rollout_id)
raw_events = self._get_events(rollout_id, include_routed_experts=False)
triplet_events = self._get_events(rollout_id, format="triplet")
return raw_events, triplet_events

Expand All @@ -425,7 +446,7 @@ def _run_succeeded_hook(self, rollout: Rollout) -> None:
return
attempt_id = rollout.status.last_attempt_id or "unknown"
trace_event_helper = _TraceEventHelper()
raw_events = self._get_events(rollout.rollout_id)
raw_events = self._get_events(rollout.rollout_id, include_routed_experts=False)
events_by_attempt = self._events_by_attempt(raw_events, attempt_id)
try:
self._hooks.on_succeeded(rollout, events_by_attempt, trace_event_helper)
Expand Down Expand Up @@ -466,6 +487,7 @@ def _build_completed_rollout(self, enqueued_rollout: EnqueuedRollout, rollout: R
response={
"token_ids": response_token_ids,
"log_probs": data.get("response_log_probs"),
"routed_experts": data.get("routed_experts"),
},
reward=None,
metadata={"server": data.get("server", {})},
Expand Down Expand Up @@ -494,6 +516,7 @@ def _build_completed_rollout(self, enqueued_rollout: EnqueuedRollout, rollout: R
finished_at = enqueued_rollout.finished_at
if finished_at is None:
finished_at = rollout.status.updated_at
diagnostic_triplet_events = [_without_routed_experts(event).model_dump() for event in triplet_events]
return CompletedRollout(
rollout_id=enqueued_rollout.rollout_id,
data_id=enqueued_rollout.data_id,
Expand All @@ -507,7 +530,7 @@ def _build_completed_rollout(self, enqueued_rollout: EnqueuedRollout, rollout: R
triplets=triplets,
metadata=metadata,
events=[event.model_dump() for event in raw_events],
triplet_events=[event.model_dump() for event in triplet_events],
triplet_events=diagnostic_triplet_events,
rollout_state=rollout.status.state,
error_message=rollout.status.error_message,
)
Expand Down
Loading