| """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 |
|
|
| |
| 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) |
|
|
| |
| 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 |
| |
| |
| 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) |
|
|
| |
| 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) |
|
|