batmac's picture
Keep classifier head in fp32 so the checkpoint loads on Linux CPU
c0f8664 verified
Raw History Blame Contribute Delete
3.22 kB
"""Save a 4-bit NF4 checkpoint of the detector for fast repeated loading.
Writes to models/gradient-ai-text-detector-4bit by default. The quantized
checkpoint is roughly 650 MB instead of 1.7 GB.
Usage:
.venv/bin/python scripts/quantize.py [output_dir]
"""
import sys
from pathlib import Path
import torch
from transformers import (
AutoModelForSequenceClassification,
AutoTokenizer,
BitsAndBytesConfig,
)
MODEL_ID = "ShantanuT01/gradient-ai-text-detector"
DEFAULT_OUT = Path(__file__).parent.parent / "models" / "gradient-ai-text-detector-4bit"
SAMPLE = "In today's rapidly evolving digital landscape, organizations must unlock value."
# bitsandbytes' packed CPU kernel asserts that each quantized weight's output
# dimension divides evenly by its block size. The classifier head is [1, 1024],
# so it must stay in fp32 or the checkpoint cannot run on Linux CPU (Spaces).
SKIP_MODULES = ["classifier"]
BLOCK = 64
def check_cpu_compatible(model):
"""Fail loudly rather than ship a checkpoint only macOS can load."""
quantized = [
module
for module in model.modules()
if type(module).__name__ == "Linear4bit"
]
offenders = sorted(
{module.out_features for module in quantized if module.out_features % BLOCK}
)
if offenders:
raise SystemExit(
f"quantized layers with out_features {offenders} are not divisible by "
f"{BLOCK}; add them to SKIP_MODULES"
)
return len(quantized)
def main(out_dir):
out_dir = Path(out_dir)
out_dir.mkdir(parents=True, exist_ok=True)
model = AutoModelForSequenceClassification.from_pretrained(
MODEL_ID,
quantization_config=BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float32,
llm_int8_skip_modules=SKIP_MODULES,
),
)
model.eval()
model.to(torch.device("cpu"))
print(f"quantized {check_cpu_compatible(model)} linear layers")
print(f"kept in fp32: {SKIP_MODULES} ({model.classifier.weight.dtype})")
model.save_pretrained(out_dir, safe_serialization=True)
AutoTokenizer.from_pretrained(MODEL_ID).save_pretrained(out_dir)
AutoTokenizer.from_pretrained(MODEL_ID, use_fast=False).save_pretrained(out_dir)
(out_dir / ".gitattributes").write_text(
"*.safetensors filter=lfs diff=lfs merge=lfs -text\n"
"*.bin filter=lfs diff=lfs merge=lfs -text\n"
"*.spm filter=lfs diff=lfs merge=lfs -text\n"
"*.model filter=lfs diff=lfs merge=lfs -text\n"
)
size_mb = sum(f.stat().st_size for f in out_dir.rglob("*") if f.is_file()) / 1024**2
print(f"saved quantized checkpoint to {out_dir} ({size_mb:,.0f} MB)")
reloaded = AutoModelForSequenceClassification.from_pretrained(out_dir)
reloaded.eval()
reloaded.to(torch.device("cpu"))
tokenizer = AutoTokenizer.from_pretrained(out_dir)
with torch.no_grad():
logits = reloaded(**tokenizer(SAMPLE, return_tensors="pt")).logits.squeeze(-1)
print(f"reload check P(AI) = {torch.sigmoid(logits).item():.4f}")
if __name__ == "__main__":
main(sys.argv[1] if len(sys.argv) > 1 else DEFAULT_OUT)