Download hako/sources/runtime.py from PowerMachine/HAKO-v1: direct link, hf CLI and curl.
- Browser
- Download file 8.81 kB
-
https://huggingface.co/PowerMachine/HAKO-v1/resolve/main/hako/sources/runtime.py
- Command line
-
hf download hf://PowerMachine/HAKO-v1/hako/sources/runtime.py
-
curl -L -o runtime.py https://huggingface.co/PowerMachine/HAKO-v1/resolve/main/hako/sources/runtime.py
8.81 kB
| """Runtime for frozen ONNX sources: real embedding extraction with fallback. | |
| Theorem T-EXH (reused from LLM01 Phase-1): once features z_i = F(x_i) are | |
| computed, the source graph is detached: dL/dW_src = 0 for all later t. The | |
| runtime enforces this by construction (torch.grad disabled; ORT outputs are | |
| plain numpy). | |
| Primary path: run the real Qwen ONNX generator graph in onnxruntime (int4 | |
| MatMulNBits kernels, CPU) and mean-pool the last hidden state. | |
| Fallback path (exhaustion-grade): frozen embedding-table extractor | |
| z = LayerNorm-free mean-pool of embedding rows (mean over tokens) | |
| which still uses ONLY real frozen source weights (the embedding initializer | |
| of the same model). Both paths are logged to telemetry (runtime_used). | |
| Granite is decomposed but not executed (RAM guard; logged). | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| from pathlib import Path | |
| from typing import List | |
| import numpy as np | |
| log = logging.getLogger("hako.runtime") | |
| class SourceRuntime: | |
| def __init__(self, model_dir: Path, name: str, memmgr=None) -> None: | |
| self.dir = Path(model_dir) | |
| self.name = name | |
| self.memmgr = memmgr | |
| self.sess = None | |
| self.mode = "unloaded" | |
| self.embed_table: np.ndarray | None = None | |
| self._try_session() | |
| # ------------------------------------------------------------- session | |
| def _try_session(self) -> None: | |
| import onnxruntime as ort | |
| model_path = self.dir / "model.onnx" | |
| if not model_path.exists(): | |
| log.warning("[%s] model.onnx missing", self.name) | |
| return | |
| opts = ort.SessionOptions() | |
| opts.intra_op_num_threads = 2 | |
| opts.inter_op_num_threads = 1 | |
| opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL | |
| try: | |
| self.sess = ort.InferenceSession(str(model_path), sess_options=opts, | |
| providers=["CPUExecutionProvider"]) | |
| # PROBE: OGA-style generator graphs are built for the ORT-GenAI | |
| # API and may reject raw session feeds (GQA kernel checks). If a | |
| # minimal 1-token run fails, permanently switch to the frozen | |
| # embedding-table extractor (still real source weights). | |
| if not self._probe_run(): | |
| self.sess = None | |
| self.mode = "embedding" | |
| log.info("[%s] ORT probe failed -> embedding extractor " | |
| "(frozen source weights, memmap)", self.name) | |
| return | |
| self.mode = "ort" | |
| log.info("[%s] ORT session ready (%d inputs)", self.name, | |
| len(self.sess.get_inputs())) | |
| except Exception as exc: | |
| log.warning("[%s] ORT session failed (%s) -> embedding fallback", | |
| self.name, type(exc).__name__) | |
| self.sess = None | |
| self.mode = "embedding" | |
| def _probe_run(self) -> bool: | |
| try: | |
| import numpy as np | |
| feed = {} | |
| for info in self.sess.get_inputs(): | |
| if info.name == "input_ids": | |
| feed[info.name] = np.ones((1, 2), dtype=np.int64) | |
| elif info.name == "attention_mask": | |
| feed[info.name] = np.ones((1, 3), dtype=np.int64) | |
| elif "past" in info.name or "cache" in info.name: | |
| shape = [d if isinstance(d, int) and d > 0 else 1 | |
| for d in info.shape] | |
| shape[2] = 1 | |
| feed[info.name] = np.zeros( | |
| shape, dtype=np.float16 if "16" in str(info.type) | |
| else np.float32) | |
| self.sess.run([o.name for o in self.sess.get_outputs()], feed) | |
| return True | |
| except Exception: | |
| return False | |
| # ------------------------------------------------------- embed fallback | |
| def _load_embed_table(self, tokenizer) -> np.ndarray | None: | |
| if self.embed_table is not None: | |
| return self.embed_table | |
| import onnx | |
| from hako.sources.decompose import _ONNX_DTYPE | |
| model = onnx.load(str(self.dir / "model.onnx"), | |
| load_external_data=False) | |
| for init in model.graph.initializer: | |
| if len(init.dims) == 2 and init.dims[1] in (896, 2048) and \ | |
| init.dims[0] > 30000: | |
| try: | |
| meta = {e.key: e.value for e in init.external_data} | |
| if meta.get("location"): | |
| # memmap the external data (RAM-light, page cache) | |
| dtype = _ONNX_DTYPE.get(init.data_type, np.float32) | |
| path = self.dir / meta["location"].split("/")[-1] | |
| self.embed_table = np.memmap( | |
| path, dtype=dtype, mode="r", | |
| offset=int(meta.get("offset", 0)), | |
| shape=tuple(init.dims)) | |
| else: | |
| from onnx import numpy_helper | |
| self.embed_table = np.asarray( | |
| numpy_helper.to_array(init), dtype=np.float32) | |
| return self.embed_table | |
| except Exception: | |
| continue | |
| return None | |
| def _ids_of(tokenizer, text: str, max_len: int) -> list: | |
| enc = tokenizer.encode(text) | |
| ids = enc.ids if hasattr(enc, "ids") else enc | |
| ids = list(ids)[:max_len] | |
| return ids or [1] | |
| # -------------------------------------------------------------- encode | |
| def encode(self, texts: List[str], tokenizer, max_len: int = 48 | |
| ) -> tuple[np.ndarray, str]: | |
| """Returns (Z (n, d), mode_used).""" | |
| if self.sess is not None: | |
| try: | |
| return self._encode_ort(texts, tokenizer, max_len) | |
| except Exception as exc: | |
| log.warning("[%s] ORT encode failed (%s) -> fallback", | |
| self.name, type(exc).__name__) | |
| self.sess = None | |
| self.mode = "embedding" | |
| return self._encode_embed(texts, tokenizer, max_len) | |
| def _encode_ort(self, texts: List[str], tokenizer, max_len: int) -> tuple: | |
| inputs_info = {i.name: i for i in self.sess.get_inputs()} | |
| ids_all, mask_all = [], [] | |
| for t in texts: | |
| enc = self._ids_of(tokenizer, t, max_len) | |
| ids_all.append(enc) | |
| mask_all.append([1] * len(enc)) | |
| L = max(len(x) for x in ids_all) | |
| ids = np.zeros((len(texts), L), dtype=np.int64) | |
| mask = np.zeros((len(texts), L), dtype=np.int64) | |
| for i, (e, m) in enumerate(zip(ids_all, mask_all)): | |
| ids[i, : len(e)] = e | |
| mask[i, : len(m)] = m | |
| feed = {} | |
| for name, info in inputs_info.items(): | |
| if name in ("input_ids",): | |
| feed[name] = ids | |
| elif name in ("attention_mask",): | |
| feed[name] = mask | |
| elif "past" in name or "cache" in name: | |
| # zeros for KV cache inputs (seq len 1 to satisfy shapes); | |
| # symbolic dims (str/None) coerce to 1 | |
| shape = [d if isinstance(d, int) and d > 0 else 1 | |
| for d in info.shape] | |
| feed[name] = np.zeros(shape, dtype=np.float16 if "16" in | |
| str(info.type) else np.float32) | |
| out_names = [o.name for o in self.sess.get_outputs()] | |
| res = self.sess.run(out_names, feed) | |
| # pick the hidden-state-like output: 3D tensor with last dim 896/2048 | |
| hidden = None | |
| for arr in res: | |
| arr = np.asarray(arr) | |
| if arr.ndim == 3 and arr.shape[-1] in (896, 2048): | |
| hidden = arr.astype(np.float32) | |
| break | |
| if hidden is None: | |
| raise RuntimeError("no hidden-state output found") | |
| pooled = (hidden * mask[:, :, None]).sum(1) / \ | |
| np.maximum(mask.sum(1, keepdims=True), 1) | |
| return pooled, "ort" | |
| def _encode_embed(self, texts: List[str], tokenizer, max_len: int | |
| ) -> tuple: | |
| table = self._load_embed_table(tokenizer) | |
| if table is None: | |
| raise RuntimeError("no embedding table available") | |
| out = np.zeros((len(texts), table.shape[1]), dtype=np.float32) | |
| for i, t in enumerate(texts): | |
| enc = self._ids_of(tokenizer, t, max_len) | |
| enc = [e for e in enc if 0 <= e < table.shape[0]] or [1] | |
| out[i] = np.asarray(table[enc], dtype=np.float32).mean(0) | |
| return out, "embedding" | |
| def close(self) -> None: | |
| self.sess = None | |
| self.embed_table = None | |
| import gc | |
| gc.collect() | |