Views
No views yet
1import mlx.core as mx
2import mlx.nn as nn
3import numpy as np
4from mlx_lm import load as mlx_load
5
6# Load model
7model, tokenizer = mlx_load("Qwen/Qwen3-8B-Base")
8
9# Load adapter
10class SwiGLUAdapter(nn.Module):
11 def __init__(self, d_model, d_inner):
12 super().__init__()
13 self.gate_proj = nn.Linear(d_model, d_inner, bias=False)
14 self.up_proj = nn.Linear(d_model, d_inner, bias=False)
15 self.down_proj = nn.Linear(d_inner, d_model, bias=False)
16 def __call__(self, h):
17 return self.down_proj(nn.sigmoid(self.gate_proj(h)) * self.up_proj(h))
18
19adapter = SwiGLUAdapter(4096, 64) # d_model=4096 for 8B
20weights = dict(np.load("ideology_8b_swiglu.npz"))
21adapter.load_weights([(k, mx.array(v)) for k, v in weights.items()])
22
23# Apply: hidden state -> adapter correction -> logits
24h = model.model(tokens)
25h = model.model.norm(h)
26h = h + adapter(mx.stop_gradient(h))
27logits = h @ model.model.embed_tokens.weight.Tideology_8b_swiglu.npz - SwiGLU adapter for Qwen3-8B-Base (786K params)ideology_4b_swiglu.npz - SwiGLU adapter for Qwen3-4B-Base (491K params)ideology_14b_swiglu.npz - SwiGLU adapter for Qwen3-14B-Base (983K params)