Skip to content
Merged
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
20 changes: 14 additions & 6 deletions agentlightning/server/routes/events.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,17 +170,25 @@ def _to_triplet_format(event: Event) -> Event:


def _dedupe_model_requests_by_prompt_token_ids(events: list[Event]) -> list[Event]:
"""Keep only the last model_request event for each prompt_token_ids key."""
last_index_by_prompt: dict[tuple[Any, ...], int] = {}
"""Keep the last request for each valid, non-empty prompt-token key."""
last_index_by_prompt: dict[tuple[int, ...], int] = {}
kept_indexes: set[int] = set()
for index, event in enumerate(events):
if event.event_type != "model_request":
continue
prompt_token_ids = event.data.get("prompt_token_ids", [])
prompt_key = tuple(prompt_token_ids) if isinstance(prompt_token_ids, list) else ()
last_index_by_prompt[prompt_key] = index
if (
not isinstance(prompt_token_ids, list)
or not prompt_token_ids
or any(type(token_id) is not int for token_id in prompt_token_ids)
):
# Skip deduplication, not the request, when no valid key exists.
kept_indexes.add(index)
continue
last_index_by_prompt[tuple(prompt_token_ids)] = index

last_indexes = set(last_index_by_prompt.values())
return [event for index, event in enumerate(events) if event.event_type != "model_request" or index in last_indexes]
kept_indexes.update(last_index_by_prompt.values())
return [event for index, event in enumerate(events) if event.event_type != "model_request" or index in kept_indexes]


@router.post("/rollouts/{rollout_id}/attempt/{attempt_id}/events", response_model=Event)
Expand Down
11 changes: 8 additions & 3 deletions agentlightning/verl/agl_rollout_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,11 +211,16 @@ def _aligned_image_urls(raw_events: list[Event], n_triplets: int) -> list[list[s
if not any(image_urls for _, _, _, image_urls in requests):
return None

# Server-side _dedupe_model_requests_by_prompt_token_ids: keep last per prompt key.
last_index_by_prompt: dict[tuple[Any, ...], int] = {}
# Mirror the server-side dedupe: keep the last request for each valid prompt key.
# Missing or malformed ids cannot establish that two requests are duplicates.
last_index_by_prompt: dict[tuple[int, ...], int] = {}
kept_indexes: set[int] = set()
for index, (_, prompt_token_ids, _, _) in enumerate(requests):
if not prompt_token_ids or any(type(token_id) is not int for token_id in prompt_token_ids):
kept_indexes.add(index)
continue
last_index_by_prompt[tuple(prompt_token_ids)] = index
kept_indexes = set(last_index_by_prompt.values())
kept_indexes.update(last_index_by_prompt.values())

aligned: list[list[str] | None] = []
for index, (data, _, response_token_ids, image_urls) in enumerate(requests):
Expand Down
43 changes: 43 additions & 0 deletions tests/server/test_endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from __future__ import annotations

import httpx
import pytest
from fastapi.testclient import TestClient

from tests.server.conftest import MODEL_NAME
Expand Down Expand Up @@ -243,6 +244,48 @@ def post_model_request(prompt_token_ids: list[int], response_token_ids: list[int
assert [event["data"]["response_token_ids"] for event in triplet_events[1:]] == [[20], [30]]


@pytest.mark.parametrize("prompt_token_ids", [None, [], [[1]], ["1"], [True], [1.0], "bad-ids", {"token": 1}])
def test_triplet_events_do_not_dedupe_without_valid_prompt_tokens(
client: TestClient,
auth_headers: dict[str, str],
prompt_token_ids: object,
):
rollout = _rollout(client, auth_headers)
rollout_id = rollout["rollout_id"]

def post_model_request(prompt_token_ids: object, response_token_ids: list[int]) -> None:
response = client.post(
f"/api/rollouts/{rollout_id}/attempt/0/events",
json={
"event_type": "model_request",
"data": {
"response": {
"prompt_token_ids": prompt_token_ids,
"choices": [{"token_ids": response_token_ids}],
},
"server": {"model": MODEL_NAME, "version": 3},
},
},
headers=auth_headers,
)
assert response.status_code == 200

post_model_request(prompt_token_ids, [10])
post_model_request(prompt_token_ids, [20])
post_model_request([1, 2], [30])
post_model_request([1, 2], [40])
post_model_request([1, 2], [50])

response = client.get(
f"/api/rollouts/{rollout_id}/events",
params={"format": "triplet"},
headers=auth_headers,
)
assert response.status_code == 200
model_requests = [event for event in response.json() if event["event_type"] == "model_request"]
assert [event["data"]["response_token_ids"] for event in model_requests] == [[10], [20], [50]]


def test_model_endpoints(client: TestClient, auth_headers: dict[str, str]):
created = client.post(
"/api/models",
Expand Down
37 changes: 33 additions & 4 deletions tests/verl/test_agl_rollout_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,7 +115,7 @@ def test_build_completed_rollout_skips_error_and_empty_model_requests() -> None:

def _raw_model_request_event(
urls: list[str],
prompt_token_ids: list[int],
prompt_token_ids: object,
response_token_ids: list[int],
status: str = "success",
http_status: int = 200,
Expand Down Expand Up @@ -145,7 +145,7 @@ def _raw_model_request_event(
)


def _triplet_model_request_event(prompt_token_ids: list[int], response_token_ids: list[int]) -> Event:
def _triplet_model_request_event(prompt_token_ids: object, response_token_ids: list[int]) -> Event:
"""Trimmed (triplet-view) model_request event as stored by the server."""
return _event(
"model_request",
Expand Down Expand Up @@ -189,6 +189,7 @@ def test_build_completed_rollout_aligns_image_urls_with_triplets() -> None:
# Superseded retry (same prompt_token_ids, keep last) with an empty response.
_raw_model_request_event([_IMG], [1, 2], []),
_raw_model_request_event([_IMG], [1, 2], [3, 4]),
_raw_model_request_event([_IMG2], [1, 2], [30, 40]),
# Filtered out: error status and http >= 400 are skipped like the triplet loop.
_raw_model_request_event([_IMG], [5], [6], status="error"),
_raw_model_request_event([_IMG], [50], [60], http_status=500),
Expand All @@ -197,7 +198,7 @@ def test_build_completed_rollout_aligns_image_urls_with_triplets() -> None:
_event("reward", {"value": 1.0}),
]
triplet_events = [
_triplet_model_request_event([1, 2], [3, 4]),
_triplet_model_request_event([1, 2], [30, 40]),
_triplet_model_request_event([7, 8], [9]),
_triplet_model_request_event([10], [11]),
_event("reward", {"value": 1.0}),
Expand All @@ -217,7 +218,35 @@ def test_build_completed_rollout_aligns_image_urls_with_triplets() -> None:

assert completed.triplets is not None
assert len(completed.triplets) == 3
assert [triplet.image_urls for triplet in completed.triplets] == [[_IMG], [_IMG, _IMG2], None]
assert [triplet.image_urls for triplet in completed.triplets] == [[_IMG2], [_IMG, _IMG2], None]


@pytest.mark.parametrize("prompt_token_ids", [None, [], [[1]], ["1"], [True], [1.0], "bad-ids", {"token": 1}])
def test_build_completed_rollout_aligns_images_without_valid_prompt_token_ids(prompt_token_ids: object) -> None:
trimmed_prompt_token_ids = [] if prompt_token_ids is None else prompt_token_ids
raw_events = [
_raw_model_request_event([_IMG], prompt_token_ids, [1]),
_raw_model_request_event([_IMG2], prompt_token_ids, [2]),
]
triplet_events = [
_triplet_model_request_event(trimmed_prompt_token_ids, [1]),
_triplet_model_request_event(trimmed_prompt_token_ids, [2]),
]
manager = _ManagerWithViews(raw_events, triplet_events)

completed = manager._build_completed_rollout(
EnqueuedRollout(
data_id="data-1",
rollout_id="rollout-1",
step=0,
sample_idx_in_step=0,
enqueue_time=0.0,
),
_rollout(),
)

assert completed.triplets is not None
assert [triplet.image_urls for triplet in completed.triplets] == [[_IMG], [_IMG2]]


def test_aligned_image_urls_count_mismatch_returns_none(capsys: pytest.CaptureFixture[str]) -> None:
Expand Down