Skip to content

Instantly share code, notes, and snippets.

@bquast
Created May 18, 2026 09:38
Show Gist options
  • Select an option

  • Save bquast/d926dd9c92e124455d5b9f6f9c79c60f to your computer and use it in GitHub Desktop.

Select an option

Save bquast/d926dd9c92e124455d5b9f6f9c79c60f to your computer and use it in GitHub Desktop.
MLX script to quantize SmolLM-135M from bf16 to int8
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