# ============================================================================== # JiRack 32B Chat (DeepSeek-R1-Distill-Qwen-32B edition, extended tokenizer) # COPYRIGHT (c) 2026 Konstantin Vladimirovich Grabko. # # Mirrors chat_jirack_7b.py, adjusted for the 32B checkpoint: # * VOCAB_SIZE=152064, hidden=5120 (per your verified checkpoint shapes) # * Remember the MKL SIMD dispatch fix if you ever run this on CPU: # export MKL_ENABLE_INSTRUCTIONS=AVX # export MKL_DEBUG_CPU_TYPE=5 # (this CPU only exposes AVX, no AVX2/AVX512 -- MKL crashes with SIGILL # otherwise). On GPU this is not needed. # * 32B is heavy: make sure you actually have the VRAM (bf16 -> ~65GB just # for weights) before loading on CUDA, or run on CPU with the env vars # above (slow, and considerably slower per token than the 1.5B model). # ============================================================================== import os import sys import torch from transformers import AutoTokenizer sys.path.append(os.getcwd()) from JiRackTernaryUltra_1b import JiRackTransformer, JiRackConfig # ========================= EDIT THESE ========================= MODEL_PATH = "/mnt/nfs_share/DeepSeek_1b/ds1p5b_checkpoint_migrated.pt" TOKENIZER_DIR = "/mnt/nfs_share/DeepSeek_1b" NO_THINK = True # True = skip reasoning # your extended tokenizer folder # ================================================================ def load_model(model_path: str): device = "cuda" if torch.cuda.is_available() else "cpu" print(f"🚀 Загрузка модели на устройство: {device.upper()}") config = JiRackConfig() model = JiRackTransformer(config, use_checkpoint=False) print(f"📥 Загрузка весов из {model_path}...") try: ckpt = torch.load(model_path, map_location="cpu", weights_only=False) state_dict = ckpt["model"] if isinstance(ckpt, dict) and "model" in ckpt else ckpt missing, unexpected = model.load_state_dict(state_dict, strict=False) real_missing = [k for k in missing if not k.endswith("lambda_")] if real_missing: print(f"⚠️ Пропущено ключей: {len(real_missing)} -> {real_missing[:10]}") if unexpected: print(f"⚠️ Лишние ключи: {len(unexpected)} -> {unexpected[:10]}") except Exception as e: print(f"❌ Критическая ошибка при загрузке весов: {e}") sys.exit(1) model = model.to(dtype=torch.bfloat16, device=device).eval() model.set_lambda(0.0) # full-precision fast path, no fake-quant at inference if device == "cuda": vram = torch.cuda.memory_allocated(0) / 1024**3 print(f"✅ VRAM занято: {vram:.1f} GB") else: print("⚠️ ВНИМАНИЕ: Запуск 32B на CPU будет ОЧЕНЬ медленным.") print(" Проверь, что выставлены MKL_ENABLE_INSTRUCTIONS=AVX и") print(" MKL_DEBUG_CPU_TYPE=5 перед запуском (см. комментарий в шапке файла).") print("✅ Модель успешно загружена.") return model, device @torch.no_grad() def generate_text(model, tokenizer, input_ids, stop_tokens, max_new_tokens=512, device="cuda"): curr_ids = input_ids.to(device) prompt_len = curr_ids.shape[1] printed = "" temperature = 0.6 top_p = 0.95 repetition_penalty = 1.15 print("JiRack: ", end="", flush=True) for _ in range(max_new_tokens): with torch.autocast(device_type=("cuda" if device == "cuda" else "cpu"), dtype=torch.bfloat16): logits = model(curr_ids) next_token_logits = logits[:, -1, :].float() / temperature for token_id in set(curr_ids[0].tolist()): if next_token_logits[0, token_id] < 0: next_token_logits[0, token_id] *= repetition_penalty else: next_token_logits[0, token_id] /= repetition_penalty sorted_logits, sorted_indices = torch.sort(next_token_logits, descending=True) cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1) sorted_indices_to_remove = cumulative_probs > top_p sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone() sorted_indices_to_remove[..., 0] = 0 next_token_logits[0, sorted_indices[sorted_indices_to_remove]] = -float('Inf') probs = torch.softmax(next_token_logits, dim=-1) next_token = torch.multinomial(probs, num_samples=1) curr_ids = torch.cat([curr_ids, next_token], dim=1) # decode the whole generated tail each step and print only the new part; # this keeps multi-token UTF-8 chars (emoji etc.) intact instead of \ufffd decoded = tokenizer.decode(curr_ids[0, prompt_len:], skip_special_tokens=True) if not decoded.endswith("\ufffd"): print(decoded[len(printed):], end="", flush=True) printed = decoded if next_token.item() in stop_tokens: break print("\n") return curr_ids def main(): if not os.path.exists(MODEL_PATH): print(f"❌ Файл {MODEL_PATH} не найден!") return try: tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_DIR) except Exception as e: print(f"❌ Ошибка токенайзера: {e}") return model, device = load_model(MODEL_PATH) stop_tokens = set() if tokenizer.eos_token_id is not None: stop_tokens.add(tokenizer.eos_token_id) for name in ("<|end_of_sentence|>", "<|endoftext|>", "<|im_end|>"): tid = tokenizer.convert_tokens_to_ids(name) if tid is not None and tid != tokenizer.unk_token_id: stop_tokens.add(tid) print("\n" + "=" * 80) print("✅ JiRack 32B (DeepSeek-R1-Distill-Qwen, extended tokenizer) Ready") print("=" * 80 + "\n") history = [] while True: try: user_input = input("User: ") if user_input.lower() in ["exit", "quit", "q"]: break if not user_input.strip(): continue history.append({"role": "user", "content": user_input}) input_ids = tokenizer.apply_chat_template( history, add_generation_prompt=True, return_tensors="pt", return_dict=False, ) # some transformers versions return a BatchEncoding here regardless; # unwrap it defensively so we always end up with a plain tensor if not torch.is_tensor(input_ids): input_ids = input_ids["input_ids"] if NO_THINK: close_ids = tokenizer.encode("\n\n", add_special_tokens=False, return_tensors="pt") input_ids = torch.cat([input_ids, close_ids], dim=1) curr_ids = generate_text(model, tokenizer, input_ids, stop_tokens, device=device) new_tokens = curr_ids[0, input_ids.shape[1]:] reply = tokenizer.decode(new_tokens, skip_special_tokens=True) history.append({"role": "assistant", "content": reply}) except KeyboardInterrupt: print("\nStopped.") break except Exception as e: print(f"\n❌ Ошибка: {e}") if __name__ == "__main__": main()