"""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 @staticmethod 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()