test_jirack_generate.py
3.8 KB · 88 lines · python Raw
1 #!/usr/bin/env python
2 # =============================================================================
3 # test_jirack_generate.py
4 # =============================================================================
5 # Функциональная проверка: модель после конвертации жива?
6 #
7 # Загружает сохранённый JiRack-чекпоинт и генерирует текст жадным циклом.
8 # KV-кэша у нашей архитектуры пока нет, поэтому на КАЖДЫЙ новый токен
9 # пересчитывается вся последовательность целиком. Отсюда дефолты:
10 # один промпт и пять токенов. Смотри на печать времени и уже сам решай,
11 # сколько можешь себе позволить.
12 #
13 # При lambda=0 BitLinear -- точный passthrough, так что тут проверяется
14 # ровно одно: корректность архитектуры и загрузки весов.
15 # =============================================================================
16
17 import time
18 import torch
19
20 torch.backends.nnpack.enabled = False # conv1d на этом CPU, чистим лог
21
22 from JiRackDeltaNet_27b import JiRackQwen38ForCausalLM
23
24 CKPT_PATH = "/data/qwen38_27b_checkpoint_migrated.pt"
25 MODEL_ID_FOR_TOKENIZER = "Qwen/Qwen3.8-27B" # тянется только токенизатор
26
27 MAX_NEW_TOKENS = 5
28 PROMPTS = [
29 "The capital of France is",
30 # "2 + 2 =",
31 # "Once upon a time, in a small village,",
32 ]
33
34 # None -- взять lambda из чекпоинта (там 0.0).
35 # Поставь 1.0, чтобы посмотреть на модель в полностью тернарном режиме
36 # БЕЗ дообучения. Ожидаемо будет плохо -- это и есть та деградация,
37 # которую потом лечит warmup.
38 LAMBDA_OVERRIDE = None
39
40
41 @torch.no_grad()
42 def generate_greedy(model, tokenizer, prompt, max_new_tokens):
43 ids = tokenizer(prompt, return_tensors="pt").input_ids
44 print(f" промпт -- {ids.shape[1]} токенов")
45 for step in range(max_new_tokens):
46 t0 = time.time()
47 logits = model(ids)
48 next_id = logits[:, -1, :].argmax(dim=-1, keepdim=True)
49 ids = torch.cat([ids, next_id], dim=1)
50 dt = time.time() - t0
51 piece = tokenizer.decode(next_id[0])
52 print(f" [{step+1}/{max_new_tokens}] {dt:6.1f} c -> {piece!r}")
53 if tokenizer.eos_token_id is not None and next_id.item() == tokenizer.eos_token_id:
54 print(" (EOS)")
55 break
56 return tokenizer.decode(ids[0], skip_special_tokens=True)
57
58
59 def main():
60 from transformers import AutoTokenizer
61
62 print(f"📥 Токенизатор ({MODEL_ID_FOR_TOKENIZER}) ...")
63 tokenizer = AutoTokenizer.from_pretrained(MODEL_ID_FOR_TOKENIZER)
64
65 print(f"📥 Чекпоинт: {CKPT_PATH}")
66 t0 = time.time()
67 model = JiRackQwen38ForCausalLM.load_checkpoint(CKPT_PATH, dtype=torch.bfloat16)
68 model.eval()
69 print(f"✅ загружен за {time.time()-t0:.1f} c")
70
71 if LAMBDA_OVERRIDE is not None:
72 for m in model.modules():
73 if hasattr(m, "set_lambda"):
74 m.set_lambda(LAMBDA_OVERRIDE)
75 print(f"⚙️ lambda принудительно = {LAMBDA_OVERRIDE}")
76
77 for prompt in PROMPTS:
78 print(f"\n── {prompt!r}")
79 text = generate_greedy(model, tokenizer, prompt, MAX_NEW_TOKENS)
80 print(f"\n ИТОГ: {text!r}")
81
82 print("\n👀 Связный текст -- модель жива, конвертация корректна.")
83 print(" Мусор или залипание на одном токене -- разбираемся дальше.")
84
85
86 if __name__ == "__main__":
87 main()
88