Skip to content

Commit 9cea5a6

Browse files
committed
Simplify NeMo Gym V1 adapter
1 parent 1a9085a commit 9cea5a6

2 files changed

Lines changed: 22 additions & 99 deletions

File tree

verifiers/v1/tasksets/nemo_gym/taskset.py

Lines changed: 20 additions & 91 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,6 @@
77

88
from __future__ import annotations
99

10-
import copy
1110
import json
1211
from collections.abc import AsyncIterator, Iterator
1312
from contextlib import asynccontextmanager
@@ -19,7 +18,7 @@
1918
from pydantic import Field
2019

2120
from verifiers.v1.decorators import reward
22-
from verifiers.v1.dialects.responses import ResponsesDialect
21+
from verifiers.v1.dialects.responses import ResponsesDialect, messages_to_wire
2322
from verifiers.v1.mcp import SharedToolsetConfig, Toolset
2423
from verifiers.v1.state import State
2524
from verifiers.v1.task import Task, TaskConfig, TaskData
@@ -44,7 +43,7 @@ class NeMoGymTaskConfig(TaskConfig):
4443

4544

4645
class NeMoGymConfig(TasksetConfig):
47-
dataset_path: Path | None = None
46+
dataset_path: Path
4847
"""JSONL rows containing ``responses_create_params`` and verifier metadata."""
4948

5049
task: NeMoGymTaskConfig = NeMoGymTaskConfig()
@@ -77,7 +76,7 @@ class _NeMoGymToolset(Toolset[SharedToolsetConfig, NeMoGymState]):
7776
rollout session on every request.
7877
"""
7978

80-
TOOL_PREFIX = "nemo_gym"
79+
TOOL_PREFIX = None
8180

8281
def _register(self, mcp: FastMCP) -> None:
8382
server = mcp._mcp_server
@@ -124,8 +123,6 @@ async def list_tools(self) -> list[MCPTool]:
124123
async def call_tool(self, name: str, arguments: dict[str, Any]) -> CallToolResult:
125124
from mcp.types import CallToolResult, TextContent
126125

127-
if name not in self.state.tool_names:
128-
raise ValueError(f"unknown NeMo Gym tool: {name}")
129126
if self.state.mcp_url is not None:
130127
async with self._upstream() as session:
131128
return await session.call_tool(name, arguments)
@@ -157,54 +154,25 @@ def _trace_to_nemo_response(
157154
f"NeMo Gym scoring requires exactly one trace branch, got {len(branches)}"
158155
)
159156

160-
known_names = set(tool_names)
161-
known_names.update(
157+
known_names = set(tool_names) | {
162158
spec["name"]
163159
for spec in responses_create_params.get("tools") or []
164160
if spec.get("type") == "function" and isinstance(spec.get("name"), str)
165-
)
166-
aliases = {
167-
name: {
168-
name,
169-
f"nemo_gym_{name}",
170-
f"nemo_gym__{name}",
171-
f"mcp__nemo_gym__{name}",
172-
}
173-
for name in known_names
174-
}
175-
response_item_types = {
176-
"reasoning",
177-
"message",
178-
"function_call",
179-
"mcp_call",
180-
"mcp_list_tools",
181-
"mcp_approval_request",
182161
}
162+
response_item_types = {"reasoning", "message", "function_call"}
183163
output: list[dict[str, Any]] = []
184164
started = False
185165

186166
for node in branches[0].nodes:
187167
message = node.message
188168
if isinstance(message, AssistantMessage) and node.sampled:
189169
started = True
190-
provider_items = message.provider_state or []
191-
if provider_items and all(
170+
if message.provider_state and not all(
192171
isinstance(item, dict) and item.get("type") in response_item_types
193-
for item in provider_items
172+
for item in message.provider_state
194173
):
195-
items = copy.deepcopy(provider_items)
196-
for item in items:
197-
if item.get("type") != "function_call":
198-
continue
199-
emitted_name = str(item.get("name"))
200-
for raw_name in sorted(known_names, key=len, reverse=True):
201-
if emitted_name in aliases[raw_name]:
202-
item["name"] = raw_name
203-
break
204-
output.extend(items)
205-
continue
206-
207-
if message.reasoning_content:
174+
message = message.model_copy(update={"provider_state": None})
175+
if message.reasoning_content and not message.provider_state:
208176
output.append(
209177
{
210178
"id": f"rs_{trace.id}_{len(output)}",
@@ -217,49 +185,24 @@ def _trace_to_nemo_response(
217185
],
218186
}
219187
)
220-
for call in message.tool_calls or []:
221-
name = call.name
222-
for raw_name in sorted(known_names, key=len, reverse=True):
223-
if name in aliases[raw_name]:
224-
name = raw_name
225-
break
226-
output.append(
227-
{
228-
"id": call.id,
229-
"type": "function_call",
230-
"call_id": call.id,
231-
"name": name,
232-
"arguments": call.arguments,
233-
"status": "completed",
234-
}
235-
)
236-
if message.content:
237-
output.append(
238-
{
239-
"id": f"msg_{trace.id}_{len(output)}",
240-
"type": "message",
241-
"role": "assistant",
242-
"status": "completed",
243-
"content": [
244-
{
245-
"type": "output_text",
246-
"text": message.content,
247-
"annotations": [],
248-
}
249-
],
250-
}
251-
)
188+
output.extend(dict(item) for item in messages_to_wire([message]))
252189
elif started and isinstance(message, ToolMessage):
253190
output.append(
254191
{
255-
"id": f"fco_{trace.id}_{len(output)}",
256192
"type": "function_call_output",
257193
"call_id": message.tool_call_id,
258194
"output": content_text(message.content),
259-
"status": "completed",
260195
}
261196
)
262197

198+
for item in output:
199+
if item.get("type") != "function_call":
200+
continue
201+
name = str(item.get("name", ""))
202+
bare_name = name.removeprefix("_")
203+
if bare_name in known_names:
204+
item["name"] = bare_name
205+
263206
model = responses_create_params.get("model") or "verifiers"
264207
return {
265208
"id": f"resp_{trace.id}",
@@ -344,29 +287,15 @@ class NeMoGymTaskset(Taskset[NeMoGymTask, NeMoGymConfig]):
344287
tools = (_NeMoGymToolset,)
345288

346289
def load(self) -> Iterator[NeMoGymTask]:
347-
if self.config.dataset_path is None:
348-
raise ValueError("NeMoGymConfig.dataset_path is required")
349290
path = self.config.dataset_path.expanduser().resolve()
350291
dialect = ResponsesDialect()
351292
count = 0
352293

353294
with path.open(encoding="utf-8") as dataset:
354-
for line_number, line in enumerate(dataset, start=1):
295+
for line in dataset:
355296
if not line.strip():
356297
continue
357-
try:
358-
row = json.loads(line)
359-
except json.JSONDecodeError as exc:
360-
raise ValueError(
361-
f"invalid JSON in {path} line {line_number}: {exc.msg}"
362-
) from exc
363-
if not isinstance(row, dict) or not isinstance(
364-
row.get("responses_create_params"), dict
365-
):
366-
raise ValueError(
367-
f"{path} line {line_number} must be an object with "
368-
"responses_create_params"
369-
)
298+
row = json.loads(line)
370299
params = row["responses_create_params"]
371300
prompt, _ = dialect.parse_request(params)
372301
yield NeMoGymTask(

verifiers/v1/tasksets/nemo_gym_weather/taskset.py

Lines changed: 2 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,6 @@
66
``NEMO_GYM_PORT`` is not 8000.
77
"""
88

9-
from collections.abc import Iterator
109
from pathlib import Path
1110

1211
from verifiers.v1.taskset import Taskset
@@ -18,13 +17,8 @@
1817

1918

2019
class NeMoGymWeatherConfig(NeMoGymConfig):
21-
dataset_path: Path | None = None
20+
dataset_path: Path = Path(__file__).with_name("example.jsonl")
2221

2322

2423
class NeMoGymWeatherTaskset(NeMoGymTaskset, Taskset[NeMoGymTask, NeMoGymWeatherConfig]):
25-
def load(self) -> Iterator[NeMoGymTask]:
26-
if self.config.dataset_path is None:
27-
self.config = self.config.model_copy(
28-
update={"dataset_path": Path(__file__).with_name("example.jsonl")}
29-
)
30-
return super().load()
24+
pass

0 commit comments

Comments
 (0)