File size: 7,239 Bytes
c9a012a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
"""
Проверка: попали ли новые кастомные токены (включая 4 tool-calling токена
<|tool_call_start|>/<|tool_call_end|>/<|tool_result_start|>/<|tool_result_end|>)
в зарезервированные слоты эмбеддинг-матрицы чекпоинта JiRack, или токенизатор
физически расширился поверх неё.

Ничего передавать не нужно -- просто:

    python3 check_tokenizer_reserve.py

Скрипт сам обходит известные места на сервере и находит:
  - папку токенизатора (пробует несколько имён и папок)
  - .pt чекпоинт модели 8B Precision (первый подходящий найденный файл)

Если хочешь всё же указать пути вручную -- можно как раньше:
    python3 check_tokenizer_reserve.py <путь_к_токенизатору> <путь_к_checkpoint.pt>
"""
import sys
import os
import glob

EMB_KEY = "token_emb.weight"
HEAD_KEY = "lm_head.weight"

TOOL_TOKENS = ["<|tool_call_start|>", "<|tool_call_end|>", "<|tool_result_start|>", "<|tool_result_end|>"]

TOKENIZER_SEARCH_DIRS = [
    ".",
    "./qwen_ji_router_tokenizer",
    "./ji_precision_tokenizer",
    "/mnt/nfs_share/Qweb2_5_Tokenizer/qwen_ji_router_tokenizer",
    "/mnt/nfs_share/Qweb2_5_Tokenizer/ji_precision_tokenizer",
    "/mnt/nfs_share/Qweb2_5_Tokenizer",
]

CHECKPOINT_SEARCH_GLOBS = [
    "/mnt/nfs_share/JiRackPrecision_8b/*.pt",
    "/mnt/nfs_share/JiRackPrecision_8b/**/*.pt",
    "./*.pt",
]


def looks_like_tokenizer_dir(d):
    if not os.path.isdir(d):
        return False
    names = os.listdir(d)
    return any(n in names for n in ("tokenizer_config.json", "tokenizer.json", "vocab.json"))


def find_tokenizer_dir():
    for d in TOKENIZER_SEARCH_DIRS:
        if looks_like_tokenizer_dir(d):
            return d
    return None


def find_checkpoint():
    for pattern in CHECKPOINT_SEARCH_GLOBS:
        hits = sorted(glob.glob(pattern, recursive=True))
        if hits:
            return hits[0]
    return None


def resolve_paths(argv):
    tok_path = argv[1] if len(argv) > 1 else None
    ckpt_path = argv[2] if len(argv) > 2 else None

    if not tok_path:
        tok_path = find_tokenizer_dir()
        if tok_path:
            print(f"(токенизатор не указан -- нашёл автоматически: {tok_path})")
        else:
            print("!!! Не смог автоматически найти папку токенизатора нигде из известных мест.")
            print("    Передай путь вручную первым аргументом.")
            sys.exit(1)

    if not ckpt_path:
        ckpt_path = find_checkpoint()
        if ckpt_path:
            print(f"(чекпоинт не указан -- нашёл автоматически: {ckpt_path})")
        else:
            print("(чекпоинт не найден автоматически -- сверка с реальной моделью будет пропущена)")

    return tok_path, ckpt_path


def main():
    from transformers import AutoTokenizer

    tok_path, ckpt_path = resolve_paths(sys.argv)

    print(f"\nЗагружаю токенизатор: {tok_path}")
    tok = AutoTokenizer.from_pretrained(tok_path)

    vocab = tok.get_vocab()
    all_ids = list(vocab.values())
    max_id = max(all_ids)
    required_rows = max_id + 1

    print(f"  len(tokenizer)      = {len(tok)}")
    print(f"  max token id        = {max_id}")
    print(f"  required emb rows   = {required_rows}  (max_id + 1)")

    if len(all_ids) != len(set(all_ids)):
        print("  !!! ОШИБКА: два разных токена делят один и тот же id. "
              "Токенизатор сломан, надо пересобрать до дальнейших шагов.")
        sys.exit(1)
    print("  все id уникальны -- ок")

    added = tok.get_added_vocab()
    if added:
        a_ids = sorted(added.values())
        print(f"  добавленных спецтокенов: {len(added)}, id диапазон [{a_ids[0]} .. {a_ids[-1]}]")
    else:
        print("  ВНИМАНИЕ: get_added_vocab() пуст -- проверь, тот ли это токенизатор.")

    print("\n--- Проверка tool-calling токенов ---")
    all_tool_tokens_present = True
    for t in TOOL_TOKENS:
        if t in vocab:
            print(f"    {t:25s} id={vocab[t]:6d}  [найден]")
        else:
            all_tool_tokens_present = False
            print(f"    {t:25s} НЕ НАЙДЕН в словаре! -> это не свежий токенизатор, "
                  f"перегенерируй через make_tokenizator.py")

    if not ckpt_path:
        print("\n(Чекпоинт не найден -- сверки с реальной моделью не будет.)")
        return

    import torch
    print(f"\nЗагружаю чекпоинт: {ckpt_path}")
    ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
    sd = ckpt["model"] if "model" in ckpt else ckpt

    if EMB_KEY not in sd or HEAD_KEY not in sd:
        print(f"  !!! Не нашёл ключи {EMB_KEY!r}/{HEAD_KEY!r} в state_dict. "
              f"Ключи в наличии (первые 10): {list(sd.keys())[:10]}")
        sys.exit(1)

    have_rows, hidden = sd[EMB_KEY].shape
    head_rows, head_hidden = sd[HEAD_KEY].shape
    print(f"  {EMB_KEY}:  [{have_rows}, {hidden}]")
    print(f"  {HEAD_KEY}: [{head_rows}, {head_hidden}]")

    if have_rows != head_rows:
        print(f"  !!! Несогласованность: token_emb {have_rows} строк, lm_head {head_rows} строк. "
              f"Это чинить раньше, чем что-либо ещё.")
        sys.exit(1)

    print("\n================ ВЕРДИКТ ================")
    if required_rows <= have_rows:
        spare = have_rows - required_rows
        print(f"RESIZE НЕ НУЖЕН.")
        print(f"   Токенизатору нужно {required_rows} строк, в чекпоинте есть {have_rows}.")
        print(f"   Все добавленные токены умещаются в резерв ({spare} строк ещё свободно).")
        print(f"   VOCAB_SIZE в JiRackTernaryUltra_*.py / JiRackPrecision_*.py оставить как есть: {have_rows}")
    else:
        extra = required_rows - have_rows
        print(f"RESIZE ОБЯЗАТЕЛЕН.")
        print(f"   Токенизатору нужно {required_rows} строк, в чекпоинте только {have_rows}.")
        print(f"   Не хватает {extra} строк -> нужно расширять token_emb.weight и lm_head.weight.")

    if not all_tool_tokens_present:
        print("\nПРИМЕЧАНИЕ: не все tool-calling токены найдены -- перегенерируй токенизатор прежде чем доверять вердикту выше.")


if __name__ == "__main__":
    main()