File size: 10,844 Bytes
de654b2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
import os
import sys
import json
import torch
import numpy as np
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from datasets import load_dataset
from transformers import AutoModelForCausalLM, AutoTokenizer

def main():
    device = "cuda" if torch.cuda.is_available() else "cpu"
    print(f"[Device] Using device: {device}")
    
    # 1. Load model and tokenizer
    print("[Model] Loading model and tokenizer from HF...")
    model_id = "SRJ5035/sw_glu_sw_64_16_8_4_xpert_gpt"
    tokenizer = AutoTokenizer.from_pretrained(model_id, revision="main")
    model = AutoModelForCausalLM.from_pretrained(model_id, revision="main", trust_remote_code=True).eval().to(device)
    
    # 2. Recreate exact validation split
    print("[Data] Loading dataset and building validation split...")
    ds = load_dataset("BabyLM-community/BabyLM-2026-Strict-Small")
    all_text = list(ds['train']['text'])
    
    split = int(len(all_text) * 0.95)
    val_texts = all_text[split:]
    print(f"[Data] Loaded {len(val_texts)} validation sentences.")
    
    # Tokenize
    print("[Data] Tokenizing validation set...")
    all_ids = []
    for t in val_texts:
        if t.strip():
            encoded = tokenizer.encode(t, add_special_tokens=False)
            all_ids.extend(encoded)
            
    total_tokens = len(all_ids)
    print(f"[Data] Total tokens in validation set: {total_tokens:,}")
    
    B, T = 16, 512
    chunk_size = B * T # 8192
    num_batches = total_tokens // chunk_size
    print(f"[Data] Generated {num_batches} batches of size {B}x{T} ({chunk_size} tokens each).")
    
    # 3. Setup hooks to capture routing decisions
    blocks = model.transformer.blocks
    num_layers = len(blocks)
    num_blocks = blocks[0].num_blocks # 4
    print(f"[Model] Found {num_layers} layers, each with {num_blocks} MoE experts.")
    
    # Dict to store statistics per layer
    stats = {}
    for l in range(num_layers):
        stats[l] = {
            "expert_counts": np.zeros(5, dtype=np.int64),
            "cooc_matrix": np.zeros((num_blocks, num_blocks), dtype=np.int64),
            "total_tokens_seen": 0
        }
        
    # Hook implementation
    def get_routing_hook(layer_idx):
        def hook(module, input, output):
            # topk_token_idx is of shape (num_blocks, k_capacity)
            topk_token_idx = module.last_topk_indices
            device = topk_token_idx.device
            k_capacity = topk_token_idx.size(1)
            
            # Since input[0] is x_0 of shape (B, T, D)
            B_batch, T_batch, D_batch = input[0].size()
            n_tokens = B_batch * T_batch
            
            # Count experts per token
            token_expert_counts = torch.zeros(n_tokens, dtype=torch.long, device=device)
            for i in range(num_blocks):
                token_expert_counts.scatter_add_(0, topk_token_idx[i], torch.ones_like(topk_token_idx[i]))
                
            counts = torch.bincount(token_expert_counts, minlength=5).cpu().numpy()
            stats[layer_idx]["expert_counts"] += counts
            
            # Co-occurrence matrix
            M = torch.zeros(n_tokens, num_blocks, dtype=torch.float, device=device)
            for i in range(num_blocks):
                M[topk_token_idx[i], i] = 1.0
                
            cooc = (M.T @ M).cpu().numpy().astype(np.int64)
            stats[layer_idx]["cooc_matrix"] += cooc
            stats[layer_idx]["total_tokens_seen"] += n_tokens
            
        return hook

    # Register hooks
    hooks = []
    for l in range(num_layers):
        h = blocks[l].register_forward_hook(get_routing_hook(l))
        hooks.append(h)
        
    # 4. Run forward passes
    print("[Eval] Running forward passes over the validation set...")
    for idx in range(num_batches):
        batch_ids = all_ids[idx * chunk_size : (idx + 1) * chunk_size]
        input_ids = torch.tensor(batch_ids, dtype=torch.long).view(B, T).to(device)
        with torch.no_grad():
            model(input_ids)
            
        if (idx + 1) % 10 == 0 or (idx + 1) == num_batches:
            print(f"[Eval] Processed batch {idx + 1}/{num_batches}")
            
    # Remove hooks
    for h in hooks:
        h.remove()
        
    print("[Eval] Completed forward passes. Aggregating statistics...")
    
    # 5. Analysis and Calculation
    global_expert_counts = np.zeros(5, dtype=np.int64)
    global_cooc_matrix = np.zeros((num_blocks, num_blocks), dtype=np.int64)
    global_total_tokens = 0
    
    for l in range(num_layers):
        global_expert_counts += stats[l]["expert_counts"]
        global_cooc_matrix += stats[l]["cooc_matrix"]
        global_total_tokens += stats[l]["total_tokens_seen"]
        
    # Convert counts to fractions
    global_expert_frac = global_expert_counts / global_total_tokens
    
    # Calculate mean experts per token
    total_expert_assignments = sum(i * global_expert_counts[i] for i in range(5))
    mean_experts_per_token = total_expert_assignments / global_total_tokens
    
    # Co-occurrence fraction
    global_cooc_frac = global_cooc_matrix / global_total_tokens
    
    print("\n" + "="*50)
    print("ROUTER UTILIZATION SUMMARY")
    print("="*50)
    print(f"Overall fraction of tokens with 0 experts: {global_expert_frac[0]:.4%} ({global_expert_counts[0]:,} tokens)")
    print(f"Overall fraction of tokens with >=1 experts: {sum(global_expert_frac[1:]):.4%}")
    print(f"Overall fraction of tokens with >=3 experts: {sum(global_expert_frac[3:]):.4%}")
    print(f"Mean experts per token (Theoretical expect: ~2.0): {mean_experts_per_token:.4f}")
    print("="*50 + "\n")
    
    # Save stats to json
    os.makedirs("./figures", exist_ok=True)
    
    raw_stats = {
        "global": {
            "expert_counts": global_expert_counts.tolist(),
            "expert_fractions": global_expert_frac.tolist(),
            "mean_experts_per_token": mean_experts_per_token,
            "cooccurrence_matrix": global_cooc_matrix.tolist(),
            "cooccurrence_fractions": global_cooc_frac.tolist()
        },
        "per_layer": {}
    }
    for l in range(num_layers):
        layer_total = stats[l]["total_tokens_seen"]
        layer_frac = stats[l]["expert_counts"] / layer_total
        layer_cooc_frac = stats[l]["cooc_matrix"] / layer_total
        raw_stats["per_layer"][f"layer_{l+1}"] = {
            "expert_counts": stats[l]["expert_counts"].tolist(),
            "expert_fractions": layer_frac.tolist(),
            "cooccurrence_matrix": stats[l]["cooc_matrix"].tolist(),
            "cooccurrence_fractions": layer_cooc_frac.tolist()
        }
        
    with open("router_stats.json", "w") as f:
        json.dump(raw_stats, f, indent=2)
    print("[Stats] Saved raw statistics to 'router_stats.json'.")
    
    # 6. PLOT GENERATION
    # Setup publication-ready matplotlib style
    plt.rcParams.update({
        "font.family": "serif",
        "font.serif": ["Times New Roman", "DejaVu Serif", "Liberation Serif", "serif"],
        "axes.labelsize": 10,
        "font.size": 10,
        "legend.fontsize": 8,
        "xtick.labelsize": 8,
        "ytick.labelsize": 8,
        "figure.titlesize": 11
    })
    
    # 6a. Plot 1: Histogram of expert choices
    fig, ax = plt.subplots(figsize=(4.5, 3.5))
    x_vals = np.arange(5)
    bars = ax.bar(x_vals, global_expert_frac * 100, color='#4f46e5', width=0.6, edgecolor='none')
    ax.set_xlabel("Number of Experts Selected")
    ax.set_ylabel("Fraction of Tokens (%)")
    ax.set_xticks(x_vals)
    ax.set_xticklabels([str(i) for i in x_vals])
    ax.spines['top'].set_visible(False)
    ax.spines['right'].set_visible(False)
    
    # Annotate bars for 0 and >=3 experts
    for i, bar in enumerate(bars):
        val = global_expert_frac[i] * 100
        if i == 0 or i >= 3:
            ax.annotate(f"{val:.2f}%",
                        xy=(bar.get_x() + bar.get_width() / 2, val),
                        xytext=(0, 3),
                        textcoords="offset points",
                        ha='center', va='bottom', fontsize=8, fontweight='bold')
                        
    plt.tight_layout()
    plt.savefig("./figures/router_util_histogram.pdf", dpi=300)
    plt.savefig("./figures/router_util_histogram.png", dpi=300)
    plt.close()
    print("[Plots] Generated Plot 1: router_util_histogram.pdf/.png")
    
    # 6b. Plot 2: Stacked bar chart by layer
    fig, ax = plt.subplots(figsize=(5.5, 3.5))
    layer_indices = np.arange(1, num_layers + 1)
    stacked_data = np.zeros((5, num_layers))
    for l in range(num_layers):
        l_total = stats[l]["total_tokens_seen"]
        stacked_data[:, l] = stats[l]["expert_counts"] / l_total
        
    colors = ['#ef4444', '#f59e0b', '#10b981', '#3b82f6', '#8b5cf6']
    bottoms = np.zeros(num_layers)
    for i in range(5):
        ax.bar(layer_indices, stacked_data[i], bottom=bottoms, label=f"{i} experts", color=colors[i], width=0.5)
        bottoms += stacked_data[i]
        
    ax.set_xlabel("Layer Index")
    ax.set_ylabel("Fraction of Tokens")
    ax.set_xticks(layer_indices)
    ax.set_xticklabels([str(i) for i in layer_indices])
    ax.spines['top'].set_visible(False)
    ax.spines['right'].set_visible(False)
    ax.legend(bbox_to_anchor=(1.02, 1), loc='upper left', frameon=False)
    
    plt.tight_layout()
    plt.savefig("./figures/router_util_by_layer.pdf", dpi=300)
    plt.savefig("./figures/router_util_by_layer.png", dpi=300)
    plt.close()
    print("[Plots] Generated Plot 2: router_util_by_layer.pdf/.png")
    
    # 6c. Plot 3: Heatmap of expert co-occurrence
    fig, ax = plt.subplots(figsize=(4.0, 3.5))
    im = ax.imshow(global_cooc_frac * 100, cmap='Blues', vmin=0, vmax=100)
    
    cbar = fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
    cbar.ax.set_ylabel("Co-occurrence Fraction (%)", rotation=-90, va="bottom", fontsize=8)
    cbar.ax.tick_params(labelsize=8)
    
    ax.set_xticks(np.arange(num_blocks))
    ax.set_yticks(np.arange(num_blocks))
    ax.set_xticklabels([f"Exp {i+1}" for i in range(num_blocks)])
    ax.set_yticklabels([f"Exp {i+1}" for i in range(num_blocks)])
    ax.set_xlabel("Expert Index")
    ax.set_ylabel("Expert Index")
    
    for i in range(num_blocks):
        for j in range(num_blocks):
            val = global_cooc_frac[i, j] * 100
            text_color = "white" if val > 50 else "black"
            ax.text(j, i, f"{val:.2f}%", ha="center", va="center", color=text_color, fontsize=8, fontweight='bold')
            
    plt.tight_layout()
    plt.savefig("./figures/expert_cooccurrence_heatmap.pdf", dpi=300)
    plt.savefig("./figures/expert_cooccurrence_heatmap.png", dpi=300)
    plt.close()
    print("[Plots] Generated Plot 3: expert_cooccurrence_heatmap.pdf/.png")
    
if __name__ == "__main__":
    main()