| """Hugging Face Inference Endpoint handler for byte-level chat generation.""" |
|
|
| from __future__ import annotations |
|
|
| import json |
| from pathlib import Path |
| from typing import Any |
|
|
| import torch |
| from modeling_harmonic_byte_transformer import ModernByteTransformer |
| from safetensors.torch import load_file |
|
|
| END_SEQUENCE = b"\n<|end|>\n" |
| ROOM_HEADER = ( |
| b"<|room|> scope=direct members=p0,p1\n" |
| b"<|participants|>\n" |
| b"p0=Agent.Ajax\n" |
| b"p1=Agent.Steve\n" |
| ) |
| USER_TURN = b"<|turn|> p0>p1 audience=p0,p1\nAgent.Ajax:\n" |
| ASSISTANT_TURN = b"<|turn|> p1>p0 audience=p0,p1\nAgent.Steve:\n" |
|
|
|
|
| def _serialize_messages(messages: list[dict[str, Any]]) -> bytes: |
| context = bytearray(ROOM_HEADER) |
| for message in messages: |
| role = str(message.get("role", "")).lower() |
| content = str(message.get("content", "")).strip().encode("utf-8", errors="replace") |
| if not content: |
| continue |
| if role == "assistant": |
| context.extend(ASSISTANT_TURN + content + END_SEQUENCE) |
| elif role == "system": |
| context.extend(USER_TURN + b"Context: " + content + END_SEQUENCE) |
| elif role == "user": |
| context.extend(USER_TURN + content + END_SEQUENCE) |
| context.extend(ASSISTANT_TURN) |
| return bytes(context) |
|
|
|
|
| class EndpointHandler: |
| """Load once per replica and serve byte-level prompt or chat requests.""" |
|
|
| def __init__(self, path: str = "") -> None: |
| model_path = Path(path) |
| config = json.loads((model_path / "config.json").read_text(encoding="utf-8")) |
| architecture = config["architecture"] |
| self.model = ModernByteTransformer(**architecture) |
| state = load_file(str(model_path / "model.safetensors"), device="cpu") |
| missing, unexpected = self.model.load_state_dict(state, strict=False) |
| if set(missing) != {"lm_head.weight"} or unexpected: |
| raise RuntimeError( |
| f"weight mismatch: missing={list(missing)}, unexpected={list(unexpected)}" |
| ) |
| self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| self.dtype = torch.bfloat16 if self.device.type == "cuda" else torch.float32 |
| self.model.to(device=self.device, dtype=self.dtype).eval() |
| self.max_seq_len = int(architecture["max_seq_len"]) |
|
|
| @staticmethod |
| def _request_context(data: dict[str, Any]) -> bytes: |
| inputs = data.get("inputs", data.get("messages")) |
| if isinstance(inputs, str): |
| return _serialize_messages([{"role": "user", "content": inputs}]) |
| if isinstance(inputs, list): |
| return _serialize_messages(inputs) |
| raise ValueError("inputs must be a prompt string or a list of role/content messages") |
|
|
| @torch.inference_mode() |
| def __call__(self, data: dict[str, Any]) -> dict[str, Any]: |
| parameters = data.get("parameters") or {} |
| max_new_bytes = min( |
| 1024, |
| max( |
| 1, |
| int(parameters.get("max_new_bytes", parameters.get("max_new_tokens", 256))), |
| ), |
| ) |
| temperature = float(parameters.get("temperature", 0.7)) |
| top_k = max(0, int(parameters.get("top_k", 40))) |
| seed = int(parameters.get("seed", 42)) |
| context = self._request_context(data) |
| generated = bytearray() |
| generator = torch.Generator(device=self.device).manual_seed(seed) |
|
|
| for _ in range(max_new_bytes): |
| window = context[-self.max_seq_len :] |
| x = torch.tensor(list(window), dtype=torch.long, device=self.device).unsqueeze(0) |
| with torch.autocast("cuda", dtype=torch.bfloat16, enabled=self.device.type == "cuda"): |
| logits = self.model(x)[0, -1].float() |
| if temperature <= 0: |
| token = int(torch.argmax(logits).item()) |
| elif top_k: |
| values, indices = torch.topk(logits / temperature, min(top_k, 256)) |
| probabilities = torch.softmax(values, dim=-1) |
| choice = torch.multinomial(probabilities, 1, generator=generator) |
| token = int(indices[choice].item()) |
| else: |
| probabilities = torch.softmax(logits / temperature, dim=-1) |
| token = int(torch.multinomial(probabilities, 1, generator=generator).item()) |
| generated.append(token) |
| context += bytes([token]) |
| if generated.endswith(END_SEQUENCE): |
| del generated[-len(END_SEQUENCE) :] |
| return { |
| "generated_text": generated.decode("utf-8", errors="replace").strip(), |
| "ended": True, |
| "generated_bytes": len(generated), |
| } |
|
|
| return { |
| "generated_text": generated.decode("utf-8", errors="replace").strip(), |
| "ended": False, |
| "generated_bytes": len(generated), |
| } |
|
|