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
8 changes: 5 additions & 3 deletions nemo_rl/models/generation/vllm/vllm_worker_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -440,10 +440,12 @@ def _clamp_max_tokens(
"""Clamp the request's max output tokens so that input + output <= max_model_len."""
remaining = self.model_config.max_model_len - len(prompt_token_ids)
if remaining <= 0:
raise ValueError(
# preserve the literal "context length" in this message to match Gym's overflow handling
raise VLLMValidationError(
f"Prompt length ({len(prompt_token_ids)}) fills or exceeds "
f"max_model_len ({self.model_config.max_model_len}). "
f"No room for output tokens."
f"this model's maximum context length ({self.model_config.max_model_len}). "
f"No room for output tokens.",
parameter="max_model_len",
)
max_tokens = min(request_max_tokens, remaining)
self._set_max_tokens(request, max_tokens)
Expand Down
70 changes: 69 additions & 1 deletion tests/unit/models/generation/test_vllm_generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -289,7 +289,9 @@ def __init__(self, **kwargs):
self.instances.append(self)

class VLLMValidationError(Exception):
pass
def __init__(self, message="", parameter=None):
super().__init__(message)
self.parameter = parameter

class ToolParserManager:
import_tool_parser = MagicMock()
Expand Down Expand Up @@ -401,6 +403,72 @@ def test_vllm_async_http_server_loads_reasoning_parser_plugin(monkeypatch):
assert "reasoning_parser_plugin" not in openai_serving_chat.instances[0].kwargs


def _nemo_gym_recognizes_context_overflow(
*, status: int, response_content: str
) -> bool:
"""Mirror responses_api_models/vllm_model/app.py overflow detection."""
return status == 400 and (
"context length" in response_content or "max_tokens" in response_content
)


@pytest.mark.asyncio
async def test_vllm_http_server_context_overflow_matches_nemo_gym(monkeypatch):
"""Overflow from _clamp_max_tokens must be HTTP 400 that NeMo Gym recognizes."""
_, _, openai_serving_chat = _install_fake_vllm_openai_modules(monkeypatch)

worker = VllmAsyncGenerationWorkerImpl.__new__(VllmAsyncGenerationWorkerImpl)
worker.cfg = {
"temperature": 1.0,
"top_p": 1.0,
"val_temperature": 0.0,
"val_top_p": 1.0,
"vllm_cfg": {},
}
worker.llm = MagicMock(model_config="model-config", renderer="renderer")
worker.llm_async_engine_args = MagicMock()
worker.llm_async_engine_args.create_model_config.return_value = MagicMock(
served_model_name="served-model", model="model-path"
)

app = _FakeFastAPIApp()
worker._setup_vllm_openai_api_server(app)

max_model_len = 128
renderer = openai_serving_chat.instances[0].kwargs["online_renderer"]
renderer.model_config = types.SimpleNamespace(max_model_len=max_model_len)
serving_chat = openai_serving_chat.instances[0]
chat_handler = next(
handler for path, handler in app.routes if path == "/v1/chat/completions"
)
overflow_prompt = [0] * max_model_len

async def create_chat_completion(request, _raw_request):
renderer._clamp_max_tokens(request, request.max_tokens, overflow_prompt)

serving_chat.create_chat_completion = create_chat_completion
response = await chat_handler(
types.SimpleNamespace(
top_k=-1,
top_p=1.0,
temperature=1.0,
max_tokens=1,
max_completion_tokens=None,
),
MagicMock(),
)

response_content = response.body.decode()
assert response.status_code == 400
assert _nemo_gym_recognizes_context_overflow(
status=response.status_code,
response_content=response_content,
)
error = json.loads(response_content)["error"]
assert error["type"] == "invalid_request_error"
assert error["code"] == 400


def test_nano_v3_reasoning_parser_swaps_reasoning_when_thinking_disabled(
monkeypatch,
):
Expand Down
Loading