khazic's picture
Archive three-epoch run: logs and provenance part 3
e65937c verified
Raw History Blame Contribute Delete
3.91 kB
"""Client for vLLM's ``/v1/chat/completions/render`` endpoint.
Only ``token_ids`` are requested: the loss mask is derived from the boundary
between two renders, not the server's ``assistant_tokens_mask`` (which needs
``{% generation %}`` tags). Renders come from the vLLM instance the pipeline
already runs, so one tokenizer feeds the mask, hidden states, and serving.
"""
import logging
import os
from http import HTTPStatus
from typing import Literal
import httpx
from speculators.data_generation.vllm_client import InvalidResponseError, with_retries
logging.getLogger("httpx").setLevel(logging.WARNING)
DEFAULT_RENDER_TIMEOUT = 30
_client: httpx.Client | None = None
_client_pid: int | None = None
def _post(url: str, *, json: dict, timeout: float) -> httpx.Response:
"""POST through a per-process pooled client.
Boundary derivation issues a couple of renders per assistant turn, tens of
thousands per dataset. Module-level ``httpx.post`` builds a new client
(SSL context, CA bundle) and a new TCP connection for every call, ~2-3 ms
of overhead each plus one throwaway socket. A pooled client removes both.
The client is keyed by PID because ``datasets.map`` workers are forked
processes: pooled sockets inherited across a fork would be shared with the
parent and corrupt each other's responses.
"""
global _client, _client_pid # noqa: PLW0603
pid = os.getpid()
if _client is None or _client_pid != pid:
_client = httpx.Client()
_client_pid = pid
return _client.post(url, json=json, timeout=timeout)
# 4xx that mean "retry", not "your request is wrong".
TRANSIENT_STATUSES = frozenset(
{HTTPStatus.REQUEST_TIMEOUT, HTTPStatus.TOO_MANY_REQUESTS}
)
class RenderError(Exception):
"""Non-200, retry-eligible response from the render endpoint."""
@with_retries
def render_conversation(
endpoint: str,
messages: list[dict],
*,
add_generation_prompt: bool,
tools: list[dict] | None = None,
chat_template_kwargs: dict | None = None,
truncate_prompt_tokens: int | None = None,
truncation_side: Literal["left", "right"] | None = None,
timeout: float = DEFAULT_RENDER_TIMEOUT,
) -> list[int]:
"""POST to ``/v1/chat/completions/render`` and return the token ids.
``truncate_prompt_tokens`` and ``truncation_side`` are optional so callers
that need the complete render can omit them. Preprocessing uses right-side
truncation to keep the rendered prefix and assistant boundary in-window.
"""
url = f"{endpoint.rstrip('/')}/v1/chat/completions/render"
body = {
"messages": messages,
"add_generation_prompt": add_generation_prompt,
}
if tools is not None:
body["tools"] = tools
if chat_template_kwargs is not None:
body["chat_template_kwargs"] = chat_template_kwargs
if truncate_prompt_tokens is not None:
body["truncate_prompt_tokens"] = truncate_prompt_tokens
if truncation_side is not None:
body["truncation_side"] = truncation_side
resp = _post(url, json=body, timeout=timeout)
if (
HTTPStatus.BAD_REQUEST <= resp.status_code < HTTPStatus.INTERNAL_SERVER_ERROR
and resp.status_code not in TRANSIENT_STATUSES
):
# Deterministic client error (bad request, wrong URL) -- retrying wastes
# requests without changing the outcome. InvalidResponseError short-
# circuits @with_retries (see vllm_client._handle_retry_error).
raise InvalidResponseError(
f"Render endpoint returned {resp.status_code}: {resp.text}"
)
if resp.status_code != HTTPStatus.OK:
raise RenderError(f"Render endpoint returned {resp.status_code}: {resp.text}")
data = resp.json()
if "token_ids" not in data:
raise RenderError(f"Render endpoint response missing 'token_ids': {data}")
return data["token_ids"]