Created
May 18, 2026 09:38
-
-
Save bquast/d926dd9c92e124455d5b9f6f9c79c60f to your computer and use it in GitHub Desktop.
MLX script to quantize SmolLM-135M from bf16 to int8
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| import mlx.core as mx | |
| from mlx.utils import tree_flatten | |
| from mlx_lm import load | |
| model, _ = load("HuggingFaceTB/SmolLM2-135M") | |
| def absmax_int8(w): | |
| scale = mx.max(mx.abs(w)) / 127.0 | |
| w_q = mx.round(w / scale).astype(mx.int8) | |
| w_hat = w_q.astype(w.dtype) * scale | |
| return w_q, w_hat, scale | |
| rows, total_fp, total_q = [], 0, 0 | |
| for name, w in tree_flatten(model.parameters()): | |
| if w.ndim < 2 or "weight" not in name: # skip biases, norm scales, embeddings if 1D | |
| continue | |
| _, w_hat, _ = absmax_int8(w) | |
| diff = w - w_hat | |
| rmse = mx.sqrt(mx.mean(diff * diff)).item() | |
| rms = mx.sqrt(mx.mean(w * w)).item() | |
| rel = rmse / (rms + 1e-12) | |
| total_fp += w.nbytes | |
| total_q += w.size # 1 byte per int8 + 2 bytes scale (negligible) | |
| rows.append((name, tuple(w.shape), rmse, rel)) | |
| print(f"FP weights: {total_fp/1e6:7.2f} MB") | |
| print(f"INT8 weights: {total_q /1e6:7.2f} MB") | |
| print(f"Compression: {total_fp/total_q:.2f}x\n") | |
| print(f"{'Layer':<55}{'Shape':<22}{'RMSE':<12}{'rel':<8}") | |
| print("-" * 97) | |
| for name, shape, rmse, rel in rows: | |
| print(f"{name:<55}{str(shape):<22}{rmse:.2e} {rel:.4f}") |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment