"""How well can the experts of layer j be predicted before segment j runs? Take segment j's input hidden state (the residual stream one layer earlier), apply layer j's own router (weights read from the IR) and compare its top-k with the experts the real router picks at the end of segment j, while generating real answers. usage: python predict_experts.py MODEL_DIR""" import sys import numpy as np import openvino as ov from onw.chat import ChatEngine from onw.runtime import SegmentedModel, read_segment def routers(d, meta): """{layer: (W [E,H] f32, gamma [H] f32)} from the 1-token segments' IR.""" core, out = ov.Core(), {} for name, m in sorted(meta["segments"].items()): if m["S"] != 1 or m.get("router_layer") is None: continue model = read_segment(core, d, name, m, False, {}, meta["T"], meta["T"]) tk = next(n for n in model.get_ops() if n.get_type_name() == "TopK") mm = tk.input(0).get_source_output().get_node().input(0).get_source_output().get_node() \ .input(0).get_source_output().get_node() # softmax <- reshape <- matmul assert mm.get_type_name() == "MatMul" w = mm.input(1).get_source_output().get_node() while w.get_type_name() != "Constant": w = w.input(0).get_source_output().get_node() g = mm.input(0).get_source_output().get_node() # xn * gamma gam = next(i.get_source_output().get_node() for i in g.inputs() if i.get_source_output().get_node().get_type_name() == "Constant") out[m["router_layer"]] = (np.asarray(w.get_data(), np.float32), np.asarray(gam.get_data(), np.float32)) return out def main(): d = sys.argv[1] e = ChatEngine(d, "NPU", pld=False) R = routers(d, e.model.meta) print(f"routers read for {len(R)} layers ({sum(w.nbytes / 2 + g.nbytes for w, g in R.values()) / 2**20:.0f} MB as f16)") K = e.model.K stats = {k: [0, 0] for k in (K, K + 4, K + 8, 2 * K)} orig = SegmentedModel._run_segment def spy(self, req, m, names, carry, common, route, S, n, cur, keep=False): L = m.get("router_layer") pred = None if S == 1 and L is not None and L in R and "carry.h" in carry: h = carry["carry.h"].astype(np.float32).reshape(-1) xn = h / np.sqrt((h * h).mean() + 1e-6) * R[L][1] pred = np.argsort(-(R[L][0] @ xn)) out = orig(self, req, m, names, carry, common, route, S, n, cur, keep) if pred is not None: real = set(int(x) for x in out[1][0][0]) for k in stats: stats[k][0] += len(real & set(int(x) for x in pred[:k])) stats[k][1] += len(real) return out SegmentedModel._run_segment = spy for q in ["NPUとGPUの違いを、身近なたとえを使って説明してください。", "次のPython関数に型ヒントを付けてください。\n\ndef add(a, b):\n return a + b", "Explain the difference between a process and a thread."]: e.checkpoint = None list(e.stream_chat([{"role": "user", "content": q}], 96)) for k, (hit, allp) in stats.items(): print(f"predict top-{k:2d}: {hit / allp * 100:5.1f}% of the real top-{K} experts caught") if __name__ == "__main__": main()