HAKO-v1 / hako /sources /runtime.py
PowerMachine's picture
HAKO upload: hako/sources/runtime.py
137f9e4 verified
Raw History Blame Contribute Delete
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
@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()