File size: 3,905 Bytes
e65937c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
"""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"]