初始化项目,由ModelHub XC社区提供模型
Model: CMSManhattan/JiRackUltra_1b Source: Original Platform
This commit is contained in:
183
chat_jirack_1b.py
Normal file
183
chat_jirack_1b.py
Normal file
@@ -0,0 +1,183 @@
|
||||
# ==============================================================================
|
||||
# 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()
|
||||
Reference in New Issue
Block a user