ldov
/

File size: 4,135 Bytes
e46c127
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Accuracy + throughput of an openjev checkpoint: transformers (eval.py's NLIScorer) vs the SGLang server.

    ./serve_sglang.sh ckpt/qwen3.5-0.8b-nli-v2s-long &
    python bench_sglang.py --ckpt ckpt/qwen3.5-0.8b-nli-v2s-long --out results/sglang/bench.json

Both backends go through the same eval.py task code; only `predict` is timed (no dataset loading)."""
import argparse
import json
import os
import time
from types import SimpleNamespace

import numpy as np
import torch

import eval as E
from sglang_client import OpenJevSGLang


class Timed:
    def __init__(self, predict, tok, max_len):
        self._predict, self.tok, self.max_len = predict, tok, max_len
        self.reset()

    def reset(self):
        self.sec, self.pairs, self.tokens = 0.0, 0, 0

    def predict(self, pairs):
        t = time.perf_counter()
        out = self._predict(pairs)
        if torch.cuda.is_available():
            torch.cuda.synchronize()
        self.sec += time.perf_counter() - t
        self.pairs += len(pairs)
        texts = [E_TEMPLATE.format(premise=p.strip(), hypothesis=h.strip()) for p, h in pairs]
        self.tokens += sum(min(len(x), self.max_len) for x in self.tok(texts)["input_ids"])
        return out


def run(scorer, tasks, mc_items, mnli_n):
    res = {}
    for t in tasks:
        scorer.reset()
        if t == "mnli":
            r = E.eval_mnli(scorer, mnli_n)
            acc = float(np.mean([v["acc"] for v in r.values()]))
        elif t in E.NLI_SETS:
            acc = E.eval_nli_set(scorer, t)["acc"]
        else:
            acc = E.eval_mc(scorer, mc_items[t])["rerank_acc"]
        res[t] = {"acc": acc, "pairs": scorer.pairs, "tokens": scorer.tokens, "sec": round(scorer.sec, 2),
                  "pairs_per_s": round(scorer.pairs / scorer.sec, 1), "tok_per_s": round(scorer.tokens / scorer.sec)}
        print(t, res[t], flush=True)
    return res


if __name__ == "__main__":
    ap = argparse.ArgumentParser()
    ap.add_argument("--ckpt", required=True)
    ap.add_argument("--url", default="http://127.0.0.1:30000")
    ap.add_argument("--out", required=True)
    ap.add_argument("--tasks", nargs="+", default=["mnli", "anli_r1", "anli_r2", "anli_r3", "wanli", "scitail", "control",
                                                   "arc_challenge", "hellaswag"])
    ap.add_argument("--backends", nargs="+", default=["sglang", "hf"])
    ap.add_argument("--mnli-n", type=int, default=None)
    ap.add_argument("--mc-n", type=int, default=2000, help="hellaswag / mmlu subsample")
    ap.add_argument("--bs", type=int, default=32)
    ap.add_argument("--max-len", type=int, default=4096)
    ap.add_argument("--workers", type=int, default=16)
    a = ap.parse_args()
    from transformers import AutoConfig, AutoTokenizer

    tok = AutoTokenizer.from_pretrained(a.ckpt)
    E_TEMPLATE = AutoConfig.from_pretrained(a.ckpt).nli_template
    margs = SimpleNamespace(mc_n=a.mc_n, chess_n=500, fewshot=5)
    mc_items = {t: E.MC_TASKS[t](margs) for t in a.tasks if t in E.MC_TASKS}
    results = {}
    for b in a.backends:
        print(f"\n===== {b}")
        if b == "hf":
            model = E.NLIScorer(a.ckpt, bs=a.bs, max_len=a.max_len)
            predict = model.predict
        else:
            predict = OpenJevSGLang(a.url, template=E_TEMPLATE, bs=a.bs, workers=a.workers).predict
        predict([("warm", "up")] * a.bs)
        results[b] = run(Timed(predict, tok, a.max_len), a.tasks, mc_items, a.mnli_n)
    os.makedirs(os.path.dirname(a.out) or ".", exist_ok=True)
    json.dump(results, open(a.out, "w"), indent=2)
    print("\n| task | pairs | avg tok | " + " | ".join(f"{b} acc | {b} pairs/s" for b in a.backends) + " | speedup |")
    print("|---|---|---|" + "---|---|" * len(a.backends) + "---|")
    for t in a.tasks:
        r0 = results[a.backends[0]][t]
        row = " | ".join(f"{results[b][t]['acc']:.4f} | {results[b][t]['pairs_per_s']}" for b in a.backends)
        sp = r0["pairs_per_s"] / results[a.backends[-1]][t]["pairs_per_s"] if len(a.backends) > 1 else 1.0
        print(f"| {t} | {r0['pairs']} | {r0['tokens'] // r0['pairs']} | {row} | {sp:.1f}x |")