rushd-agent / code /nci_mlx_patch.py
BinSaqban's picture
Upload code/nci_mlx_patch.py with huggingface_hub
231a673 verified
Raw
History Blame Contribute Delete
3.03 kB
"""NCI Experiment on Rushd-Geo via MLX — monkey-patch approach"""
import os, json, sys, time
os.environ["WANDB_DISABLED"] = "true"
import mlx.core as mx
import numpy as np
from mlx_lm import load
from mlx_lm.models.qwen3_5 import DecoderLayer
MODEL = "/Users/ai/rushd-geo-mlx-4bit"
N_LAYERS = 64
# Monkey-patch: wrap DecoderLayer.__call__ to capture hidden states
original_call = DecoderLayer.__call__
hidden_states_collector = []
def patched_call(self, x, mask=None, cache=None):
global hidden_states_collector
result = original_call(self, x, mask=mask, cache=cache)
hidden_states_collector.append(result)
return result
DecoderLayer.__call__ = patched_call
print(f"Loading model from {MODEL}...", flush=True)
t0 = time.time()
model, tokenizer = load(MODEL)
print(f"Loaded in {time.time()-t0:.1f}s", flush=True)
print(f"Layers: {len(model.layers)}", flush=True)
# Test prompts
prompts = [
("meaningful", "تحليل الوضع الجيوسياسي في الشرق الأوسط بعد اتفاقيات التطبيع."),
("nonsense", "jdska flpz xqwy bnmzx cvbnm lkjh"),
]
def compute_nci(hs_list, layer_a, layer_b):
n_a = hs_list[layer_a].mean(axis=1).squeeze(0)
n_b = hs_list[layer_b].mean(axis=1).squeeze(0)
cos = mx.sum(n_a * n_b) / (mx.linalg.norm(n_a) * mx.linalg.norm(n_b))
return float(cos)
def compute_norm(hs_list, layer):
n = hs_list[layer].mean(axis=1).squeeze(0)
return float(mx.linalg.norm(n))
print(f"\n{'Prompt':<12} {'Norm L6':<10} {'Norm L32':<10} {'NCI(6,32)':<12} {'NCI(6,54)':<12}", flush=True)
print("-" * 60, flush=True)
for label, text in prompts:
hidden_states_collector.clear()
tokens = list(tokenizer.encode(text))
input_ids = mx.array([tokens])
t0 = time.time()
logits = model(input_ids)
elapsed = time.time() - t0
hs = hidden_states_collector # 64 layers
# hs[0] = layer 1 output, hs[63] = layer 64 output
nci_6_32 = compute_nci(hs, 6, 32)
nci_6_54 = compute_nci(hs, 6, 54)
norm_6 = compute_norm(hs, 6)
norm_32 = compute_norm(hs, 32)
print(f"{label:<12} {norm_6:<10.1f} {norm_32:<10.1f} {nci_6_32:<12.4f} {nci_6_54:<12.4f}", flush=True)
print(f"{'':12} tokens={len(tokens)} time={elapsed:.2f}s", flush=True)
# Full NCI profile
print(f"\n=== Full NCI Profile (meaningful) ===", flush=True)
text = "تحليل الوضع الجيوسياسي في الشرق الأوسط."
hidden_states_collector = []
tokens = list(tokenizer.encode(text))
input_ids = mx.array([tokens])
logits = model(input_ids)
hs = hidden_states_collector
print(f"{'Layer':<6} {'NCI(L6,Ln)':<14} {'Norm':<12}", flush=True)
print("-" * 35, flush=True)
for i in range(N_LAYERS):
nci = compute_nci(hs, 6, i) if i < len(hs) else 0
norm_i = compute_norm(hs, i) if i < len(hs) else 0
mark = ""
if i in [0, 3, 6, 12, 18, 24, 30, 32, 40, 48, 54, 60, 63]:
mark = " ★"
print(f"L{i:<4} {nci:<14.4f} {norm_i:<12.1f}{mark}", flush=True)
print("\nDone!", flush=True)