JackKozmo29 commited on
Commit
7d03f2d
·
verified ·
1 Parent(s): af0e951

Upload serve_openai_shim.py

Browse files
Files changed (1) hide show
  1. serve_openai_shim.py +75 -0
serve_openai_shim.py ADDED
@@ -0,0 +1,75 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ OpenAI-compatible shim for mailguard-jev-style-1.5b.
3
+ Exposes /v1/chat/completions so any OpenAI client can use the model locally.
4
+
5
+ Usage:
6
+ pip install fastapi uvicorn transformers torch
7
+ python serve_openai_shim.py
8
+ """
9
+
10
+ import json
11
+ import re
12
+ import time
13
+ import uuid
14
+
15
+ import torch
16
+ import uvicorn
17
+ from fastapi import FastAPI
18
+ from fastapi.responses import JSONResponse
19
+ from pydantic import BaseModel
20
+ from transformers import AutoModelForCausalLM, AutoTokenizer
21
+
22
+ MODEL_ID = "JackKozmo29/mailguard-jev-style-1.5b"
23
+
24
+ print(f"Loading {MODEL_ID} ...")
25
+ tok = AutoTokenizer.from_pretrained(MODEL_ID)
26
+ model = AutoModelForCausalLM.from_pretrained(MODEL_ID, dtype=torch.float32).eval()
27
+ print("Model ready.")
28
+
29
+ app = FastAPI()
30
+
31
+
32
+ class Message(BaseModel):
33
+ role: str
34
+ content: str
35
+
36
+
37
+ class ChatRequest(BaseModel):
38
+ model: str = MODEL_ID
39
+ messages: list[Message]
40
+ max_tokens: int = 120
41
+ temperature: float = 0.0
42
+
43
+
44
+ @app.post("/v1/chat/completions")
45
+ def chat(req: ChatRequest):
46
+ messages = [m.model_dump() for m in req.messages]
47
+ prompt = tok.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
48
+ ids = tok(prompt, return_tensors="pt")
49
+ with torch.no_grad():
50
+ out = model.generate(
51
+ **ids,
52
+ max_new_tokens=req.max_tokens,
53
+ do_sample=False,
54
+ pad_token_id=tok.eos_token_id,
55
+ )
56
+ text = tok.decode(out[0][ids["input_ids"].shape[1]:], skip_special_tokens=True)
57
+ m = re.search(r"\{.*\}", text, re.DOTALL)
58
+ content = m.group(0) if m else text.strip()
59
+
60
+ return JSONResponse({
61
+ "id": f"chatcmpl-{uuid.uuid4().hex[:8]}",
62
+ "object": "chat.completion",
63
+ "created": int(time.time()),
64
+ "model": req.model,
65
+ "choices": [{
66
+ "index": 0,
67
+ "message": {"role": "assistant", "content": content},
68
+ "finish_reason": "stop",
69
+ }],
70
+ "usage": {"prompt_tokens": ids["input_ids"].shape[1], "completion_tokens": len(out[0]) - ids["input_ids"].shape[1], "total_tokens": len(out[0])},
71
+ })
72
+
73
+
74
+ if __name__ == "__main__":
75
+ uvicorn.run(app, host="0.0.0.0", port=8000)