File size: 7,309 Bytes
12496fc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
"""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())