File size: 4,990 Bytes
cd29a1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Таблица префилла на релизной сборке. Замер, а не наследство.

Зачем. В карточке таблица 131К 2496 т/с, 262К 1799, 524К 1140, 1M 648 подана как
замеренная, но квитанции у неё в репозитории нет — сплошная сверка чисел этого не
нашла. Карточка при этом обещает: «каждое число ниже замерено на нашем железе и
имеет квитанцию». Значит либо квитанция, либо число уходит.

Кроме того числа расходятся с сегодняшними прогонами процентов на семнадцать —
конфигурация с тех пор менялась (арена, план позиций, распределитель), так что
старые значения могли просто устареть.

Как мерится. `max_tokens=1`, чтобы декод не размазывал итог: одна ступень декода
стоит около 150 мс, и на коротких длинах это заметная доля. Отдельно пишется
время до первого токена, которое движок сообщает сам, — из него и считается
скорость префилла. Промпт подаётся идентификаторами: построение текстом
повторами даёт не ту длину, на которой думаешь мерить.
"""
from __future__ import annotations

import argparse, json, os, time


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", required=True)
    ap.add_argument("--out", required=True)
    ap.add_argument("--lengths", default="131072,262144,524288,1010000")
    args = ap.parse_args()
    lengths = [int(x) for x in args.lengths.split(",")]

    from vllm import LLM, SamplingParams

    t0 = time.time()
    llm = LLM(model=args.model)                 # без флагов: профиль решает всё
    load_s = round(time.time() - t0, 1)
    cfg = llm.llm_engine.vllm_config
    tok = llm.get_tokenizer()

    unit = ("Полярная станция ведёт наблюдения за дрейфом льда, и записи "
            "хранятся в общем журнале смены. ")
    per = len(tok(unit).input_ids)

    rows = []
    for n in lengths:
        text = unit * (n // max(per - 1, 1) + 64)
        ids = tok(text).input_ids[:n]
        if len(ids) < n:
            rows.append({"target": n, "error": "корпус короче цели"})
            continue
        t1 = time.time()
        out = llm.generate([{"prompt_token_ids": ids}],
                           SamplingParams(temperature=0.0, max_tokens=1))
        wall = time.time() - t1
        m = out[0].metrics
        ttft = None
        if m is not None and getattr(m, "first_token_time", None) and \
           getattr(m, "first_scheduled_time", None):
            ttft = m.first_token_time - m.first_scheduled_time
        rows.append({
            "tokens": len(ids),
            "wall_seconds": round(wall, 2),
            "wall_tokens_per_second": round(len(ids) / wall, 1),
            "time_to_first_token_seconds": round(ttft, 2) if ttft else None,
            "prefill_tokens_per_second": round(len(ids) / ttft, 1) if ttft else None,
        })
        print(json.dumps(rows[-1], ensure_ascii=False), flush=True)

    receipt = {
        "schema": "lomonosov_zenit_prefill_table_v2",
        "what_this_shows": "скорость префилла на релизной сборке, max_tokens=1",
        "supersedes": "таблица карточки 131К 2496 / 262К 1799 / 524К 1140 / 1M 648 — "
                      "без квитанции в репозитории и снятая до правок арены, плана "
                      "позиций и распределителя",
        "load_seconds": load_s,
        "max_model_len": int(cfg.model_config.max_model_len),
        "kv_cache_dtype": str(cfg.cache_config.cache_dtype),
        "arena_bytes": getattr(cfg.cache_config, "kv_cache_memory_bytes", None),
        "chunk": int(cfg.scheduler_config.max_num_batched_tokens),
        "enforce_eager": bool(cfg.model_config.enforce_eager),
        "pytorch_cuda_alloc_conf": os.environ.get("PYTORCH_CUDA_ALLOC_CONF"),
        "position_stages": os.environ.get("ZENIT_POSITION_STAGES"),
        "clocks": "заводские, не фиксировались",
        "rows": rows,
    }
    with open(args.out, "w", encoding="utf-8") as fh:
        json.dump(receipt, fh, ensure_ascii=False, indent=1)
    print(json.dumps(receipt, ensure_ascii=False, indent=1))
    return 0


if __name__ == "__main__":
    raise SystemExit(main())