import os import time import numpy as np import onnxruntime as ort from transformers import AutoTokenizer #ONNX_MODEL_PATH = "/root/JiRackTernary1/new/model/jirack_1b.onnx" # В файле chat_jirack_1b.onnx.py ONNX_MODEL_PATH = "jirack_1b_int8.onnx" TOKENIZER_NAME = "meta-llama/Llama-3.2-1B-Instruct" def softmax(x): e_x = np.exp(x - np.max(x)) return e_x / e_x.sum(axis=-1, keepdims=True) def sample_top_p(probs, p=0.9): sorted_indices = np.argsort(probs, axis=-1)[0, ::-1] sorted_probs = probs[0, sorted_indices] cumulative_probs = np.cumsum(sorted_probs) indices_to_remove = cumulative_probs > p indices_to_remove[1:] = indices_to_remove[:-1].copy() indices_to_remove[0] = False sorted_probs[indices_to_remove] = 0 if sorted_probs.sum() > 0: sorted_probs /= sorted_probs.sum() else: return sorted_indices[0] return np.random.choice(sorted_indices, p=sorted_probs) def main(): # Форсируем использование CPU из-за несовместимости версий ROCm providers = ['CPUExecutionProvider'] print(f"🔄 Инициализация ONNX Runtime (CPU Mode)...") try: sess_options = ort.SessionOptions() # Оптимизация для многопоточности на CPU sess_options.intra_op_num_threads = os.cpu_count() sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL session = ort.InferenceSession(ONNX_MODEL_PATH, sess_options=sess_options, providers=providers) print(f"✅ Провайдер: {session.get_providers()[0]}") except Exception as e: print(f"❌ Ошибка инициализации: {e}") return tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_NAME) print("--- ☕ JiRack 1B ONNX Ecommerce Ready (CPU) ---") while True: u = input("\nUser: ") if u.lower() in ["exit", "quit"]: break if not u.strip(): continue # Формат Llama-3.2 prompt = f"<|begin_of_text|><|start_header_id|>user<|end_header_id|>\n\n{u}<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n" inputs = tokenizer(prompt, return_tensors="np") input_ids = inputs["input_ids"].astype(np.int64) print("JiRack (ONNX-CPU): ", end="", flush=True) start_time = time.time() generated_tokens = 0 max_tokens = 256 for _ in range(max_tokens): # Инференс на CPU ort_inputs = {session.get_inputs()[0].name: input_ids[:, -1024:]} logits = session.run(None, ort_inputs)[0] # Сэмплирование (Temperature 0.7 + Top-P 0.9) next_token_logits = logits[:, -1, :] / 0.7 probs = softmax(next_token_logits) next_token_id = sample_top_p(probs, p=0.9) input_ids = np.concatenate([input_ids, [[next_token_id]]], axis=-1) generated_tokens += 1 token_str = tokenizer.decode([next_token_id], skip_special_tokens=True) print(token_str, end="", flush=True) if next_token_id in [tokenizer.eos_token_id, 128001, 128009]: break dt = time.time() - start_time if dt > 0: print(f"\n\n⏱ {generated_tokens} tokens in {dt:.2f}s ({generated_tokens/dt:.1f} tps)") if __name__ == "__main__": main()