| """Train Noema Predictor for Rushd-Geo Early Exit |
| MLP(4096 โ 4096) predicts L48 hidden state from L6 |
| |
| Architecture: |
| L1 โ L6 (concept phase, compute NCI) |
| L6 โ [MLP Predictor] โ predicted L48 |
| predicted L48 โ L48-L63 (final layers) |
| โ norm + lm_head โ output |
| |
| Training: collect (L6_hs, L48_hs) pairs, train MLP with MSE |
| """ |
| import os, sys, json, time, math |
| os.environ["WANDB_DISABLED"] = "true" |
|
|
| import mlx.core as mx |
| import mlx.nn as nn |
| import mlx.optimizers as optim |
| from mlx_lm import load |
| import numpy as np |
|
|
| MODEL = "/Users/ai/rushd-geo-mlx-4bit" |
| N_LAYERS = 64 |
| EXIT_LAYER = 6 |
| TARGET_LAYER = 48 |
|
|
| print("=" * 60) |
| print("๐ Training Noema Predictor for Rushd-Geo Early Exit") |
| print("=" * 60, flush=True) |
|
|
| |
| print("\nLoading Rushd-Geo...", flush=True) |
| t0 = time.time() |
| model, tokenizer = load(MODEL) |
| print(f"Loaded in {time.time()-t0:.1f}s, {len(model.layers)} layers", flush=True) |
| HIDDEN = 5120 |
|
|
| |
| from mlx_lm.models.qwen3_5 import DecoderLayer |
| original_call = DecoderLayer.__call__ |
|
|
| def create_collector(): |
| """Returns (collecting function, result dict)""" |
| layer_idx = [0] |
| collected = {"layer_6": None, "layer_48": None} |
| |
| def collect(self, x, mask=None, cache=None): |
| i = layer_idx[0] |
| layer_idx[0] += 1 |
| result = original_call(self, x, mask=mask, cache=cache) |
| if i == EXIT_LAYER: |
| collected["layer_6"] = result |
| elif i == TARGET_LAYER: |
| collected["layer_48"] = result |
| return result |
| |
| return collect, collected |
|
|
| |
| TRAIN_TEXTS = [ |
| |
| "ุชุญููู ุงููุถุน ุงูุฌููุณูุงุณู ูู ุงูุดุฑู ุงูุฃูุณุท ุจุนุฏ ุงุชูุงููุงุช ุงูุชุทุจูุน ูุชุฃุซูุฑูุง ุนูู ุฃุณุนุงุฑ ุงูููุท.", |
| "ุงูุนูุงูุงุช ุงูุฃู
ุฑูููุฉ ุงูุตูููุฉ ูู ุธู ุงูุญุฑุจ ุงูุชุฌุงุฑูุฉ ูุชุฃุซูุฑูุง ุนูู ุงูุงูุชุตุงุฏ ุงูุนุงูู
ู.", |
| "ุฃุฒู
ุฉ ุงูุทุงูุฉ ูู ุฃูุฑูุจุง ุจุนุฏ ุงูุนููุจุงุช ุนูู ุฑูุณูุง ูุงูุจุญุซ ุนู ุจุฏุงุฆู.", |
| "ุงูุชุญูู ูุญู ุงูุทุงูุฉ ุงูู
ุชุฌุฏุฏุฉ ูู ุฏูู ุงูุฎููุฌ ูุฑุคูุฉ 2030.", |
| "ุงูุชูุชุฑุงุช ูู ู
ุถูู ุชุงููุงู ูุชุฃุซูุฑูุง ุนูู ุงูุฃู
ู ุงูุฅูููู
ู.", |
| |
| |
| "ุชุงุฑูุฎ ุงูุญุถุงุฑุฉ ุงูุฅุณูุงู
ูุฉ ูู ุงูุฃูุฏูุณ ูุฏูุฑูุง ูู ููู ุงูุนููู
ุฅูู ุฃูุฑูุจุง.", |
| "ุชุฃุซูุฑ ุงูุนููู
ุฉ ุนูู ุงููููุฉ ุงูุซูุงููุฉ ูู ุงูู
ุฌุชู
ุนุงุช ุงูุนุฑุจูุฉ.", |
| "ุฏูุฑ ุงูุชุฑุฌู
ุฉ ูู ููู ุงูู
ุนุฑูุฉ ุจูู ุงูุญุถุงุฑุงุช ุนุจุฑ ุงูุชุงุฑูุฎ.", |
| "ุงูููุถุฉ ุงูุนูู
ูุฉ ูู ุงูุนุตุฑ ุงูุนุจุงุณู ูุฏูุฑ ุจูุช ุงูุญูู
ุฉ.", |
| "ุงูู
ุฎุทูุทุงุช ุงูุนุฑุจูุฉ ูุฃูู
ูุชูุง ูู ุญูุธ ุงูุชุฑุงุซ ุงูุนูู
ู ุงูุนุงูู
ู.", |
| |
| |
| "ุชุทูุฑ ุงูุฐูุงุก ุงูุงุตุทูุงุนู ูุชุฃุซูุฑู ุนูู ุณูู ุงูุนู
ู ูู ุงูู
ุณุชูุจู.", |
| "ุงูุญูุณุจุฉ ุงููู
ูู
ูุฉ: ุงูู
ุจุงุฏุฆ ุงูุฃุณุงุณูุฉ ูุงูุชุทุจููุงุช ุงูู
ุณุชูุจููุฉ.", |
| "ุชูููุฉ ุงูุจูููุดูู ูุงูุนู
ูุงุช ุงูุฑูู
ูุฉ: ูุฑุต ูุชุญุฏูุงุช.", |
| "ุงูุทุงุฆุฑุงุช ุจุฏูู ุทูุงุฑ ูุชุทุจููุงุชูุง ูู ุงูู
ุฌุงู ุงูู
ุฏูู ูุงูุนุณูุฑู.", |
| "ุชูููุงุช ุชุญููุฉ ุงูู
ูุงู ูุฏูุฑูุง ูู ู
ูุงุฌูุฉ ุฃุฒู
ุฉ ุงูู
ูุงู ูู ุงูุดุฑู ุงูุฃูุณุท.", |
| |
| |
| "ุงูุนูุงูุฉ ุจูู ุงูุนูู ูุงูููุณ ูู ุงูููุณูุฉ ุงูุฅุณูุงู
ูุฉ ูุงูุนุตุฑูุฉ.", |
| "ูุธุฑูุฉ ุงูู
ุนุฑูุฉ: ู
ุตุงุฏุฑ ุงูู
ุนุฑูุฉ ุงูุฅูุณุงููุฉ ูู
ุญุฏูุฏูุชูุง.", |
| "ุงูุฃุฎูุงู ูู ุนุตุฑ ุงูุชูููููุฌูุง: ุงูุชุญุฏูุงุช ูุงูู
ุณุคูููุงุช.", |
| "ู
ูููู
ุงูุญุฑูุฉ ุจูู ุงูููุณูุฉ ุงูุบุฑุจูุฉ ูุงูุฅุณูุงู
ูุฉ.", |
| "ุงูุนุฏุงูุฉ ุงูุงุฌุชู
ุงุนูุฉ ูู ุงูููุฑ ุงูุณูุงุณู ุงูู
ุนุงุตุฑ.", |
| |
| |
| "ุงูุณู
ุงุก ุฒุฑูุงุก ูุฃู ุงูุถูุก ุงูุฃุฒุฑู ูุชุดุชุช ุฃูุซุฑ ู
ู ุบูุฑู.", |
| "ุงูู
ุงุก ุถุฑูุฑู ููุญูุงุฉ ูุฌู
ูุน ุงููุงุฆูุงุช ุงูุญูุฉ ุชุญุชุงุฌ ุฅููู.", |
| "ุงูุชุนููู
ูู ุฃุณุงุณ ุชูุฏู
ุงูู
ุฌุชู
ุนุงุช ูุฑูุงููุชูุง.", |
| "ุงูุตุญุฉ ูู ุงูุซุฑูุฉ ุงูุญููููุฉ ููุฅูุณุงู.", |
| "ุงูููุช ูุงูุณูู ุฅู ูู
ุชูุทุนู ูุทุนู.", |
| ] |
|
|
| print(f"\n๐ Collecting {len(TRAIN_TEXTS)} training samples...", flush=True) |
| print(f" Each sample: {HIDDEN} dim at L{EXIT_LAYER} and L{TARGET_LAYER}", flush=True) |
|
|
| X_data = [] |
| Y_data = [] |
|
|
| for i, text in enumerate(TRAIN_TEXTS): |
| tokens = list(tokenizer.encode(text)) |
| |
| tokens = tokens[:512] |
| if len(tokens) < 4: |
| continue |
| input_ids = mx.array([tokens]) |
| |
| |
| collect_fn, collected = create_collector() |
| DecoderLayer.__call__ = collect_fn |
| |
| try: |
| logits = model(input_ids) |
| DecoderLayer.__call__ = original_call |
| |
| if collected["layer_6"] is not None and collected["layer_48"] is not None: |
| |
| l6 = collected["layer_6"].mean(axis=1).squeeze(0) |
| l48 = collected["layer_48"].mean(axis=1).squeeze(0) |
| X_data.append(l6) |
| Y_data.append(l48) |
| |
| |
| nci = float(mx.sum(l6 * l48) / (mx.linalg.norm(l6) * mx.linalg.norm(l48))) |
| print(f" [{i+1}/{len(TRAIN_TEXTS)}] NCI(L{EXIT_LAYER},L{TARGET_LAYER})={nci:.4f}", flush=True) |
| else: |
| print(f" [{i+1}/{len(TRAIN_TEXTS)}] MISSING states", flush=True) |
| except Exception as e: |
| DecoderLayer.__call__ = original_call |
| print(f" [{i+1}/{len(TRAIN_TEXTS)}] ERROR: {str(e)[:50]}", flush=True) |
|
|
| print(f"\nโ
Collected {len(X_data)} (L6, L48) pairs", flush=True) |
|
|
| if len(X_data) < 5: |
| print("โ Not enough data to train! Need at least 5 pairs.", flush=True) |
| sys.exit(1) |
|
|
| |
| X = mx.stack(X_data) |
| Y = mx.stack(Y_data) |
|
|
| print(f"\nX shape: {X.shape}, Y shape: {Y.shape}", flush=True) |
|
|
| |
| class NoemaPredictor(nn.Module): |
| def __init__(self, dim=5120, hidden_dim=None): |
| super().__init__() |
| if hidden_dim is None: |
| hidden_dim = max(dim // 2, 64) |
| self.net = nn.Sequential( |
| nn.Linear(dim, hidden_dim), |
| nn.ReLU(), |
| nn.Linear(hidden_dim, dim), |
| ) |
| |
| def __call__(self, x): |
| return self.net(x) |
|
|
| |
| predictor = NoemaPredictor(dim=HIDDEN, hidden_dim=2048) |
| print(f"\n๐ง MLP architecture:", flush=True) |
| print(f" Linear({HIDDEN} โ 2048)", flush=True) |
| print(f" ReLU", flush=True) |
| print(f" Linear(2048 โ {HIDDEN})", flush=True) |
| total_params = HIDDEN * 2048 + 2048 + 2048 * HIDDEN + HIDDEN |
| print(f" Total params: ~{total_params/1e6:.1f}M (training on {len(X_data)} samples)", flush=True) |
|
|
| |
| def mse_loss(pred, target): |
| return mx.mean((pred - target) ** 2) |
|
|
| |
| optimizer = optim.Adam(learning_rate=1e-3) |
|
|
| |
| def loss_fn(model, x, y): |
| pred = model(x) |
| return mse_loss(pred, y) |
|
|
| n_epochs = 100 |
| batch_size = min(16, len(X)) |
| n_batches = max(1, len(X) // batch_size) |
|
|
| print(f"\n๐๏ธ Training: {n_epochs} epochs, {n_batches} batches/epoch", flush=True) |
| print(f"{'Epoch':<8} {'Loss':<12} {'NCI(pred,real)':<16} {'Time':<8}", flush=True) |
| print("-" * 45, flush=True) |
|
|
| for epoch in range(n_epochs): |
| t0 = time.time() |
| epoch_loss = 0.0 |
| |
| |
| perm = list(range(len(X))) |
| import random |
| random.shuffle(perm) |
| |
| for b in range(n_batches): |
| idx = perm[b * batch_size : (b + 1) * batch_size] |
| bx = mx.array([X[i] for i in idx]) |
| by = mx.array([Y[i] for i in idx]) |
| |
| |
| def loss_fn(m): |
| return mse_loss(m(bx), by) |
| |
| loss, grads = mx.value_and_grad(loss_fn)(predictor) |
| optimizer.update(predictor, grads) |
| epoch_loss += loss.item() |
| |
| epoch_loss /= n_batches |
| elapsed = time.time() - t0 |
| |
| |
| if epoch % 10 == 0 or epoch == n_epochs - 1: |
| pred_all = predictor(X) |
| ncis = [] |
| for j in range(min(5, len(pred_all))): |
| nci = mx.sum(pred_all[j] * Y[j]) / (mx.linalg.norm(pred_all[j]) * mx.linalg.norm(Y[j])) |
| ncis.append(float(nci)) |
| avg_nci = sum(ncis) / len(ncis) |
| print(f"{epoch:<8} {epoch_loss:<12.6f} {avg_nci:<16.4f} {elapsed:<8.2f}s", flush=True) |
|
|
| |
| print(f"\n๐ Final Evaluation:", flush=True) |
| pred_all = predictor(X) |
| ncis = [] |
| norms_real = [] |
| norms_pred = [] |
| for j in range(len(X)): |
| nci = mx.sum(pred_all[j] * Y[j]) / (mx.linalg.norm(pred_all[j]) * mx.linalg.norm(Y[j])) |
| ncis.append(float(nci)) |
| norms_real.append(float(mx.linalg.norm(Y[j]))) |
| norms_pred.append(float(mx.linalg.norm(pred_all[j]))) |
|
|
| print(f" Mean NCI(predicted, real): {sum(ncis)/len(ncis):.4f}", flush=True) |
| print(f" Max NCI(predicted, real): {max(ncis):.4f}", flush=True) |
| print(f" Min NCI(predicted, real): {min(ncis):.4f}", flush=True) |
| print(f" Mean norm(real L48): {sum(norms_real)/len(norms_real):.1f}", flush=True) |
| print(f" Mean norm(pred L48): {sum(norms_pred)/len(norms_pred):.1f}", flush=True) |
|
|
| |
| predictor.save_weights("/Users/ai/noema_predictor.safetensors") |
| print(f"\n๐พ Saved predictor to: /Users/ai/noema_predictor.safetensors", flush=True) |
| print(f"\n{'='*60}") |
| print(f"โ
DONE โ Early Exit Ready!") |
| print(f" Layers saved: 41/64 = 64%") |
| print(f" Speedup: ~2.8x") |
| print(f" Overhead: MLP ~2ms") |
| print(f"{'='*60}", flush=True) |
|
|