File size: 3,025 Bytes
231a673
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
"""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)