184 lines
7.3 KiB
Python
184 lines
7.3 KiB
Python
# ==============================================================================
|
|
# 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 <think> 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("</think>\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()
|