NEXORA / nexora /inference.py
devildasdf's picture
Release validated NEXORA research prototype, tiny weights and evidence
12496fc verified
Raw History Blame Contribute Delete
7.31 kB
"""Interchangeable local HF and explicit HTTP chat backends."""
from pathlib import Path
import json
import os
import time
import urllib.request
import urllib.parse
import threading
class HTTPBackend:
def __init__(self, url, model, *, allow_network=False, timeout=60, api_key_env="NEXORA_API_KEY"):
parsed = urllib.parse.urlparse(url)
if parsed.scheme not in {"http", "https"} or parsed.username or parsed.password:
raise ValueError("Expected HTTP(S) URL without embedded credentials")
if parsed.hostname not in {"localhost", "127.0.0.1", "::1"} and not allow_network:
raise PermissionError("Remote inference requires explicit allow_network")
self.url, self.model, self.timeout, self.api_key_env = url.rstrip("/"), model, timeout, api_key_env
def complete(self, messages, schema=None):
body = {"model": self.model, "messages": messages, "temperature": 0, "max_tokens": 512}
if schema:
body["response_format"] = {"type": "json_schema", "json_schema": {"name": "nexora_action", "schema": schema}}
headers = {"Content-Type": "application/json"}
key = os.environ.get(self.api_key_env)
if key:
headers["Authorization"] = "Bearer " + key
request = urllib.request.Request(self.url + "/chat/completions", data=json.dumps(body).encode(), headers=headers)
# Redirects are disabled so an approved local URL cannot redirect prompts elsewhere.
class NoRedirect(urllib.request.HTTPRedirectHandler):
def redirect_request(self, *args, **kwargs):
return None
opener = urllib.request.build_opener(urllib.request.ProxyHandler({}), NoRedirect())
with opener.open(request, timeout=self.timeout) as response:
raw = response.read(2_000_001)
if len(raw) > 2_000_000:
raise ValueError("Inference response too large")
return json.loads(raw)["choices"][0]["message"]["content"]
class HFBackend:
def __init__(self, model_path, *, device="cpu", threads=4, max_new_tokens=256, context_limit=4096):
import torch
from transformers import AutoConfig, AutoTokenizer, AutoModelForCausalLM, AutoModelForImageTextToText
torch.set_num_threads(threads)
self.torch = torch
self.tokenizer = AutoTokenizer.from_pretrained(model_path, local_files_only=True, trust_remote_code=False)
dtype = torch.float32 if device == "cpu" else torch.float16
config = AutoConfig.from_pretrained(model_path, local_files_only=True, trust_remote_code=False)
factory = AutoModelForImageTextToText if config.model_type == "qwen3_5" else AutoModelForCausalLM
self.model = factory.from_pretrained(model_path, local_files_only=True, trust_remote_code=False, dtype=dtype).to(device).eval()
self.max_new_tokens, self.context_limit = max_new_tokens, context_limit
self.last_metrics = {}
def _inputs(self, messages, schema=None):
if schema:
messages = [*messages, {"role": "user", "content": "Return JSON only matching: " + json.dumps(schema)}]
text = self.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, enable_thinking=False)
inputs = self.tokenizer(text, return_tensors="pt").to(self.model.device)
count = inputs["input_ids"].shape[1]
if count + self.max_new_tokens > self.context_limit:
raise ValueError("Request exceeds configured context budget; retrieve less context")
return inputs, count
def complete(self, messages, schema=None):
inputs, count = self._inputs(messages, schema)
start = time.perf_counter()
with self.torch.inference_mode():
result = self.model.generate(**inputs, max_new_tokens=self.max_new_tokens, do_sample=False,
pad_token_id=self.tokenizer.eos_token_id)
if self.model.device.type == "cuda":
self.torch.cuda.synchronize()
elapsed = time.perf_counter() - start
generated = result[0, count:]
self.last_metrics = {"input_tokens": count, "output_tokens": len(generated), "seconds": elapsed,
"tokens_per_second": len(generated) / max(elapsed, 1e-9),
"ttft": None, "note": "Non-streaming adapter; TTFT not measured"}
return self.tokenizer.decode(generated, skip_special_tokens=True)
def stream(self, messages):
"""One request at a time; caller must serialize access to this backend."""
from transformers import TextIteratorStreamer, StoppingCriteria, StoppingCriteriaList
inputs, count = self._inputs(messages)
cancelled = threading.Event()
errors = []
class Stop(StoppingCriteria):
def __call__(self, input_ids, scores, **kwargs):
return cancelled.is_set()
streamer = TextIteratorStreamer(self.tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=120)
def generate():
try:
with self.torch.inference_mode():
self.model.generate(**inputs, max_new_tokens=self.max_new_tokens, do_sample=False,
pad_token_id=self.tokenizer.eos_token_id, streamer=streamer,
stopping_criteria=StoppingCriteriaList([Stop()]))
except Exception as exc:
errors.append(exc)
streamer.on_finalized_text("", stream_end=True)
start, first = time.perf_counter(), None
worker = threading.Thread(target=generate, daemon=True)
worker.start()
chunks = []
try:
for chunk in streamer:
if chunk:
first = time.perf_counter() if first is None else first
chunks.append(chunk)
yield chunk
if errors:
raise errors[0]
finally:
cancelled.set()
worker.join(timeout=120)
if worker.is_alive():
raise RuntimeError("Generation did not stop; restart inference worker")
elapsed = time.perf_counter()-start
output_tokens = len(self.tokenizer.encode("".join(chunks), add_special_tokens=False))
self.last_metrics = {"input_tokens": count, "output_tokens_retokenized": output_tokens,
"seconds": elapsed, "time_to_first_text_seconds": None if first is None else first-start,
"note": "Text chunks can buffer multiple tokens; first text is not exact first-token latency"}
def tiny_generate(checkpoint, prompt, max_new_tokens=64):
import torch
from safetensors.torch import load_file
from .model import NexoraLM, ModelConfig
from .tokenizer import ByteTokenizer
torch.set_num_threads(4)
path = Path(checkpoint)
model = NexoraLM(ModelConfig(**json.loads((path / "config.json").read_text())))
model.load_state_dict(load_file(str(path / "model.safetensors")))
tok = ByteTokenizer()
ids = torch.tensor([[tok.bos_id, *tok.encode(prompt)]])
out = model.generate(ids, max_new_tokens=max_new_tokens)
return tok.decode(out[0].tolist())