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