Text Classification
Transformers
Safetensors
English
deberta-v2
ai-generated-text-detection
4-bit precision
bitsandbytes
nf4
quantization
text-embeddings-inference
Instructions to use batmac/gradient-ai-text-detector-4bit with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use batmac/gradient-ai-text-detector-4bit with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-classification", model="batmac/gradient-ai-text-detector-4bit")# Load model directly from transformers import AutoTokenizer, AutoModelForSequenceClassification tokenizer = AutoTokenizer.from_pretrained("batmac/gradient-ai-text-detector-4bit") model = AutoModelForSequenceClassification.from_pretrained("batmac/gradient-ai-text-detector-4bit", device_map="auto") - Notebooks
- Google Colab
- Kaggle
File size: 3,222 Bytes
8465953 c0f8664 8465953 c0f8664 8465953 c0f8664 8465953 c0f8664 8465953 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 | """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)
|