File size: 3,728 Bytes
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
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
"""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()