"""Benchmark quantization and dtype options for the AI text detector. Compares fp32, bf16/fp16, and bitsandbytes NF4 4-bit on CPU and Apple Silicon MPS: memory footprint, throughput, and agreement with fp32 reference probabilities. """ import time import torch from transformers import AutoModelForSequenceClassification, AutoTokenizer MODEL_ID = "ShantanuT01/gradient-ai-text-detector" SAMPLES = [ "In today's rapidly evolving digital landscape, organizations must leverage " "synergistic strategies to unlock unprecedented value.", "I went to the store yesterday and bought three apples. They were a bit " "bruised, but the price was right and I figured they would be fine in a pie.", "The committee reviewed the findings over several weeks and concluded that " "no further action was warranted at this time.", "Furthermore, it is important to note that the implementation of such " "measures requires careful consideration of multiple stakeholders.", ] * 2 tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) def footprint_mb(model): return model.get_memory_footprint() / 1024**2 def param_device(model): for param in model.parameters(): return param.device return torch.device("cpu") @torch.no_grad() def probs(model, texts): encoded = tokenizer( texts, return_tensors="pt", truncation=True, max_length=512, padding=True ) device = param_device(model) encoded = {k: v.to(device) for k, v in encoded.items()} logits = model(**encoded).logits.squeeze(-1) return torch.sigmoid(logits).float().cpu().tolist() def bench(name, model, device=None, reference=None): """Run a warm-up batch, then time a full batch of SAMPLES.""" model.eval() moved = "" if device is not None: try: model.to(device) except Exception as exc: moved = f" (stayed {param_device(model).type}: {type(exc).__name__})" try: probs(model, SAMPLES[:2]) except Exception as exc: print(f"{name:20s} FAILED: {type(exc).__name__}: {str(exc)[:120]}") return None start = time.perf_counter() values = probs(model, SAMPLES) elapsed = time.perf_counter() - start delta = max(abs(a - b) for a, b in zip(values, reference)) if reference else None print( f"{name:20s} dev {str(param_device(model)):4s} " f"mem {footprint_mb(model):7.1f} MB " f"{len(SAMPLES)} seqs {elapsed:6.3f}s max|dp| " f"{'n/a' if delta is None else f'{delta:.4f}'}{moved}" ) return values def load(**kwargs): kwargs.setdefault("dtype", torch.float32) return AutoModelForSequenceClassification.from_pretrained(MODEL_ID, **kwargs) def main(): reference = bench("fp32 cpu", load()) bench("fp32 mps", load(), torch.device("mps"), reference) for dtype, label in [(torch.bfloat16, "bf16"), (torch.float16, "fp16")]: bench(f"{label} mps", load(dtype=dtype), torch.device("mps"), reference) from transformers import BitsAndBytesConfig four_bit = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.float32, ) bench("nf4 4-bit cpu", load(quantization_config=four_bit), None, reference) bench( "nf4 4-bit mps", load(quantization_config=four_bit), torch.device("mps"), reference, ) try: int8 = torch.quantization.quantize_dynamic( load(), {torch.nn.Linear}, dtype=torch.qint8 ) bench("dynamic int8 cpu", int8, None, reference) except Exception as exc: print(f"dynamic int8 cpu unsupported: {str(exc)[:100]}") if __name__ == "__main__": main()