from fastapi import FastAPI, HTTPException from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel import onnxruntime as ort import numpy as np import time from transformers import AutoTokenizer import os app = FastAPI(title="JiRack 1B ONNX Server") app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) # ========================= CONFIG ========================= ONNX_MODEL_PATH = "jirack_orca_1b_int8.onnx" TOKENIZER_NAME = "meta-llama/Llama-3.2-1B-Instruct" MAX_TOKENS = 512 TEMPERATURE = 0.7 TOP_P = 0.9 # Load model print("🔄 Loading ONNX model...") sess_options = ort.SessionOptions() 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=['CPUExecutionProvider']) tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_NAME) print(f"✅ Model loaded | CPU Threads: {os.cpu_count()}") class ChatRequest(BaseModel): messages: list model: str = "jirack" temperature: float = TEMPERATURE max_tokens: int = MAX_TOKENS top_p: float = TOP_P def softmax(x): x = x - np.max(x) return np.exp(x) / np.sum(np.exp(x)) def sample_top_p(probs, p=0.9): sorted_indices = np.argsort(probs)[::-1] sorted_probs = probs[sorted_indices] cumulative_probs = np.cumsum(sorted_probs) sorted_probs[cumulative_probs > p] = 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) @app.post("/v1/chat/completions") async def chat_completions(request: ChatRequest): try: # Get last user message user_content = request.messages[-1]["content"] prompt = f"<|begin_of_text|><|start_header_id|>user<|end_header_id|>\n\n{user_content}<|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) generated_tokens = [] start_time = time.time() for _ in range(request.max_tokens): # Take last 1024 tokens ort_inputs = {session.get_inputs()[0].name: input_ids[:, -1024:]} logits = session.run(None, ort_inputs)[0] # shape: (1, seq_len, vocab) next_token_logits = logits[0, -1, :] / request.temperature probs = softmax(next_token_logits) next_token_id = sample_top_p(probs, request.top_p) input_ids = np.concatenate([input_ids, [[next_token_id]]], axis=-1) token_str = tokenizer.decode([next_token_id], skip_special_tokens=True) generated_tokens.append(token_str) if next_token_id in [tokenizer.eos_token_id, 128001, 128009]: break response_text = "".join(generated_tokens).strip() return { "id": f"chatcmpl-{int(time.time())}", "object": "chat.completion", "created": int(time.time()), "model": request.model, "choices": [{ "index": 0, "message": {"role": "assistant", "content": response_text}, "finish_reason": "stop" }], "usage": { "prompt_tokens": len(input_ids[0]), "completion_tokens": len(generated_tokens), "total_tokens": len(input_ids[0]) + len(generated_tokens) } } except Exception as e: print(f"ERROR: {e}") raise HTTPException(status_code=500, detail=str(e)) @app.get("/health") async def health(): return {"status": "ok", "model": "JiRack 1B ONNX"} if __name__ == "__main__": import uvicorn uvicorn.run( app, host="0.0.0.0", port=7869, timeout_keep_alive=600, workers=1, log_level="info" )