test_jirack_generate.py
| 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 | |