onw / predict_experts.py
ryugyosoft's picture
v0.8.0: faster SSD streaming without more memory - prompt prefetch into cache slots, predicted next-layer expert prefetch while decoding (helper thread, separate read queue), single-copy reads
ee52468 verified
Raw History Blame Contribute Delete
3.35 kB
"""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()