Skip to content

Commit 108c733

Browse files
authored
[Router Replay]: Improve performance by removing Pydantic validation (#1394)
1 parent 58b119f commit 108c733

9 files changed

Lines changed: 117 additions & 64 deletions

‎docs/reference.md‎

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -181,6 +181,14 @@ class TrajectoryStep(TypedDict):
181181

182182
A single turn in a multi-turn rollout.
183183

184+
### RoutedExpertsPayload
185+
186+
```python
187+
class RoutedExpertsPayload(TypedDict):
188+
data: Any # actually memoryview; kept opaque so Pydantic skips schema validation
189+
shape: list[int]
190+
```
191+
184192
### TrajectoryStepTokens
185193

186194
```python
@@ -192,7 +200,7 @@ class TrajectoryStepTokens(TypedDict):
192200
completion_logprobs: list[float]
193201
overlong_prompt: bool
194202
is_truncated: bool
195-
routed_experts: list[list[list[int]]] | None # [seq_len, layers, topk] to enable router replay
203+
routed_experts: RoutedExpertsPayload | None
196204
multi_modal_data: NotRequired[Any] # renderers.MultiModalData sidecar (pixel_values, placeholder ranges) — set only on multimodal rollouts
197205
```
198206

‎tests/test_client_auth_errors.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -130,6 +130,9 @@ def __init__(self, message: str) -> None:
130130
def __init__(self, message: str) -> None:
131131
self.chat = self._Chat(message)
132132

133+
async def post(self, *args, **kwargs): # noqa: ANN002, ANN003
134+
return await self.chat.completions.create(*args, **kwargs)
135+
133136

134137
@pytest.mark.parametrize(
135138
"error_message",

‎tests/test_openai_chat_completions_token_client.py‎

Lines changed: 21 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
from typing import Any, cast
22

3+
import httpx
34
import pytest
45

56
from verifiers.clients.openai_chat_completions_client import OpenAIChatCompletionsClient
@@ -24,7 +25,25 @@ async def post(
2425
self, path: str, body: dict[str, Any], cast_to: type, **kwargs: Any
2526
) -> Any:
2627
self.calls.append({"path": path, "body": body, "cast_to": cast_to})
27-
return {"ok": True, "path": path, "body": body}
28+
return httpx.Response(
29+
200,
30+
json={
31+
"id": path,
32+
"object": "chat.completion",
33+
"created": 1,
34+
"model": body["model"],
35+
"choices": [
36+
{
37+
"index": 0,
38+
"message": {"role": "assistant", "content": "ok"},
39+
"finish_reason": "stop",
40+
}
41+
],
42+
"ok": True,
43+
"path": path,
44+
"body": body,
45+
},
46+
)
2847

2948

3049
class _PromptIdTestClient(OpenAIChatCompletionsTokenClient):
@@ -270,7 +289,7 @@ async def fake_get_prompt_ids( # noqa: ANN001
270289
state=state,
271290
)
272291

273-
assert response["ok"] is True
292+
assert response.model_extra["ok"] is True
274293
assert len(recording_client.calls) == 1
275294
assert recording_client.calls[0]["path"] == "/chat/completions/tokens"
276295
assert recording_client.calls[0]["body"]["tokens"] == [10, 20]

‎verifiers/clients/openai_chat_completions_client.py‎

Lines changed: 22 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -56,8 +56,10 @@
5656
Usage,
5757
UserMessage,
5858
)
59-
from verifiers.utils.client_utils import setup_openai_client
60-
from verifiers.utils.response_utils import parse_routed_experts
59+
from verifiers.utils.client_utils import (
60+
post_chat_completion_with_routed_experts_sidecar,
61+
setup_openai_client,
62+
)
6163

6264

6365
def handle_openai_overlong_prompt(func):
@@ -300,23 +302,24 @@ def normalize_sampling_args(sampling_args: SamplingArgs):
300302
sampling_args = {**sampling_args, "modalities": ["text"]}
301303

302304
extra_headers = kwargs.pop("extra_headers", None)
303-
305+
request_args = normalize_sampling_args(sampling_args)
306+
extra_body = request_args.pop("extra_body", {})
307+
308+
body: dict[str, Any] = {
309+
"model": model,
310+
"messages": prompt,
311+
**request_args,
312+
**extra_body,
313+
}
304314
if tools:
305-
response = await self.client.chat.completions.create(
306-
model=model,
307-
messages=prompt,
308-
tools=tools,
309-
extra_headers=extra_headers,
310-
**normalize_sampling_args(sampling_args),
311-
)
312-
else:
313-
response = await self.client.chat.completions.create(
314-
model=model,
315-
messages=prompt,
316-
extra_headers=extra_headers,
317-
**normalize_sampling_args(sampling_args),
318-
)
319-
return response
315+
body["tools"] = tools
316+
317+
return await post_chat_completion_with_routed_experts_sidecar(
318+
self.client,
319+
"/chat/completions",
320+
body=body,
321+
extra_headers=extra_headers,
322+
)
320323

321324
async def raise_from_native_response(self, response: OpenAIChatResponse) -> None:
322325
if response is None:
@@ -483,14 +486,13 @@ def parse_tokens(response: OpenAIChatResponse) -> ResponseTokens | None:
483486
completion_logprobs = [token["logprob"] for token in logprobs_content]
484487

485488
choice_extra = choice.model_extra or {}
486-
routed_experts = parse_routed_experts(choice_extra.get("routed_experts"))
487489
return ResponseTokens(
488490
prompt_ids=prompt_ids,
489491
prompt_mask=prompt_mask,
490492
completion_ids=completion_ids,
491493
completion_mask=completion_mask,
492494
completion_logprobs=completion_logprobs,
493-
routed_experts=routed_experts,
495+
routed_experts=choice_extra.get("routed_experts"),
494496
)
495497

496498
def parse_reasoning_content_from_response(

‎verifiers/clients/openai_chat_completions_token_client.py‎

Lines changed: 12 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,6 @@
33

44
from openai import AsyncOpenAI, BaseModel
55
from openai.types.chat import (
6-
ChatCompletion,
76
ChatCompletionAssistantMessageParam,
87
)
98
from openai.types.chat.chat_completion_message_function_tool_call_param import (
@@ -20,6 +19,9 @@
2019
handle_openai_overlong_prompt,
2120
)
2221
from verifiers.types import SamplingArgs, State
22+
from verifiers.utils.client_utils import (
23+
post_chat_completion_with_routed_experts_sidecar,
24+
)
2325

2426

2527
def _has_multimodal_content(messages) -> bool:
@@ -140,20 +142,20 @@ def normalize_sampling_args(sampling_args: SamplingArgs):
140142
)
141143

142144
extra_body = sampling_args.pop("extra_body", {})
143-
body = dict(
144-
model=model,
145-
messages=prompt,
146-
tools=tools,
147-
tokens=prompt_ids,
145+
body = {
146+
"model": model,
147+
"messages": prompt,
148+
"tools": tools,
149+
"tokens": prompt_ids,
148150
**sampling_args,
149151
**extra_body,
150-
)
152+
}
151153

152-
return await self.client.post(
154+
return await post_chat_completion_with_routed_experts_sidecar(
155+
self.client,
153156
"/chat/completions/tokens",
154157
body=body,
155-
cast_to=ChatCompletion,
156-
options={"headers": extra_headers} if extra_headers else {},
158+
extra_headers=extra_headers,
157159
)
158160

159161
async def get_prompt_ids(

‎verifiers/clients/openai_completions_client.py‎

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,6 @@
2525
Usage,
2626
)
2727
from verifiers.utils.client_utils import setup_openai_client
28-
from verifiers.utils.response_utils import parse_routed_experts
2928

3029
OpenAITextMessages = str
3130
OpenAITextResponse = Completion
@@ -170,15 +169,12 @@ def parse_tokens(response: OpenAITextResponse) -> ResponseTokens | None:
170169
)
171170
if completion_logprobs is None:
172171
return None
173-
choice_extra = response.choices[0].model_extra or {}
174-
routed_experts = parse_routed_experts(choice_extra.get("routed_experts"))
175172
return ResponseTokens(
176173
prompt_ids=prompt_ids,
177174
prompt_mask=prompt_mask,
178175
completion_ids=completion_ids,
179176
completion_mask=completion_mask,
180177
completion_logprobs=completion_logprobs,
181-
routed_experts=routed_experts,
182178
)
183179

184180
return Response(

‎verifiers/types.py‎

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -183,13 +183,19 @@ class Usage(CustomBaseModel):
183183
total_tokens: int
184184

185185

186+
class RoutedExpertsPayload(TypedDict):
187+
# Keep the raw response sidecar opaque so Pydantic does not validate memoryview.
188+
data: Any
189+
shape: list[int]
190+
191+
186192
class ResponseTokens(CustomBaseModel):
187193
prompt_ids: list[int]
188194
prompt_mask: list[int]
189195
completion_ids: list[int]
190196
completion_mask: list[int]
191197
completion_logprobs: list[float]
192-
routed_experts: str | None = None # base64 NumPy [seq_len, layers, topk]
198+
routed_experts: RoutedExpertsPayload | None = None
193199
# Renderer-emitted multimodal sidecar (renderers.base.MultiModalData)
194200
# carrying processed pixel_values / placeholder ranges per modality.
195201
# Populated by the renderer client when the rollout went through a
@@ -232,7 +238,7 @@ class TrajectoryStepTokens(TypedDict):
232238
completion_logprobs: list[float]
233239
overlong_prompt: bool
234240
is_truncated: bool
235-
routed_experts: str | None # base64 NumPy [seq_len, layers, topk]
241+
routed_experts: RoutedExpertsPayload | None
236242
# Renderer-emitted multimodal sidecar (renderers.base.MultiModalData)
237243
# carrying processed pixel_values / placeholder ranges per modality.
238244
# ``NotRequired`` because text-only rollouts (and non-renderer client

‎verifiers/utils/client_utils.py‎

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,16 +1,20 @@
11
import json
22
import logging
33
import os
4+
from collections.abc import Mapping
5+
from typing import Any
46
from pathlib import Path
57

68
import httpx
79
from anthropic import AsyncAnthropic
810
from openai import AsyncOpenAI
11+
from openai.types.chat import ChatCompletion
912

1013
from verifiers.types import (
1114
ClientConfig,
1215
EndpointClientConfig,
1316
)
17+
from verifiers.utils.response_utils import strip_routed_experts_data
1418

1519
logger = logging.getLogger(__name__)
1620

@@ -97,6 +101,32 @@ def setup_http_client(config: ClientConfig) -> httpx.AsyncClient:
97101
return _build_http_client(resolved_config, headers)
98102

99103

104+
def parse_chat_completion_with_routed_experts_sidecar(raw: bytes) -> ChatCompletion:
105+
stripped, routed_data = strip_routed_experts_data(raw)
106+
response = ChatCompletion.model_validate_json(stripped)
107+
if routed_data is not None:
108+
choice_extra = response.choices[0].model_extra
109+
assert choice_extra is not None
110+
choice_extra["routed_experts"]["data"] = routed_data
111+
return response
112+
113+
114+
async def post_chat_completion_with_routed_experts_sidecar(
115+
client: AsyncOpenAI,
116+
path: str,
117+
*,
118+
body: dict[str, Any],
119+
extra_headers: Mapping[str, str] | None = None,
120+
) -> ChatCompletion:
121+
raw_response = await client.post(
122+
path,
123+
body=body,
124+
cast_to=httpx.Response,
125+
options={"headers": extra_headers} if extra_headers else {},
126+
)
127+
return parse_chat_completion_with_routed_experts_sidecar(raw_response.content)
128+
129+
100130
def _setup_openai_client_from_resolved(config: ClientConfig) -> AsyncOpenAI:
101131
headers, api_key = _build_headers_and_api_key(config)
102132
return AsyncOpenAI(

‎verifiers/utils/response_utils.py‎

Lines changed: 12 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -1,34 +1,23 @@
1-
import base64
2-
from io import BytesIO
3-
from typing import Any, cast
4-
5-
import numpy as np
6-
71
from verifiers.types import (
82
AssistantMessage,
93
Messages,
104
Response,
115
TrajectoryStepTokens,
126
)
137

8+
ROUTED_EXPERTS_DATA_PREFIX = b'"routed_experts":{"data":"'
149

15-
def parse_routed_experts(raw: Any) -> str | None:
16-
if raw is None:
17-
return None
18-
return cast(str, raw)
19-
20-
21-
def truncate_routed_experts(routed_experts: str | None, seq_len: int) -> str | None:
22-
if routed_experts is None:
23-
return None
2410

25-
array = np.load(BytesIO(base64.b64decode(routed_experts)), allow_pickle=False)
26-
assert array.ndim == 3
27-
assert 0 <= seq_len <= array.shape[0]
11+
def strip_routed_experts_data(raw: bytes) -> tuple[bytes, memoryview | None]:
12+
data_start = raw.find(ROUTED_EXPERTS_DATA_PREFIX)
13+
if data_start < 0:
14+
return raw, None
2815

29-
buffer = BytesIO()
30-
np.save(buffer, np.ascontiguousarray(array[:seq_len]), allow_pickle=False)
31-
return base64.b64encode(buffer.getvalue()).decode("ascii")
16+
data_start += len(ROUTED_EXPERTS_DATA_PREFIX)
17+
data_end = raw.index(b'"', data_start)
18+
routed_data = memoryview(raw)[data_start:data_end]
19+
stripped = raw[:data_start] + raw[data_end:]
20+
return stripped, routed_data
3221

3322

3423
async def parse_response_message(response: Response) -> Messages:
@@ -73,15 +62,11 @@ async def parse_response_tokens(
7362
completion_ids = []
7463
completion_mask = []
7564
completion_logprobs = []
76-
routed_experts = truncate_routed_experts(routed_experts, len(prompt_ids))
7765
elif prompt_len + completion_len > max_seq_len:
7866
is_truncated = True
7967
completion_ids = tokens.completion_ids[: max_seq_len - prompt_len]
8068
completion_mask = tokens.completion_mask[: max_seq_len - prompt_len]
8169
completion_logprobs = tokens.completion_logprobs[: max_seq_len - prompt_len]
82-
routed_experts = truncate_routed_experts(
83-
routed_experts, prompt_len + len(completion_ids)
84-
)
8570
else:
8671
is_truncated = False
8772
else:
@@ -104,4 +89,6 @@ async def parse_response_tokens(
10489
# step. Leaving it on ``response.message.tokens`` too means every
10590
# downstream pass (msgpack, save) has to dedupe the duplicate.
10691
tokens.multi_modal_data = None
92+
if routed_experts is not None:
93+
tokens.routed_experts = None
10794
return out

0 commit comments

Comments
 (0)