Download source/src/speculators/data_generation/render_client.py from khazic/spec-b300: direct link, hf CLI and curl.
- Browser
- Download file 3.91 kB
-
https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/data_generation/render_client.py
- Command line
-
hf download hf://khazic/spec-b300/source/src/speculators/data_generation/render_client.py
-
curl -L -o render_client.py https://huggingface.co/khazic/spec-b300/resolve/main/source/src/speculators/data_generation/render_client.py
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.""" | |
| 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"] | |