"""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), }