Download openai_api.py from tchbcb/samai-4b: direct link, hf CLI and curl.
- Browser
- Download file 10.6 kB
-
https://huggingface.co/tchbcb/samai-4b/resolve/main/openai_api.py
- Command line
-
hf download hf://tchbcb/samai-4b/openai_api.py
-
curl -L -o openai_api.py https://huggingface.co/tchbcb/samai-4b/resolve/main/openai_api.py
10.6 kB
| # -*- coding: utf-8 -*- | |
| """s4_openai_api.py — samai-4b OpenAI 兼容 API (Flask, port 7862, key=1234). | |
| Endpoints: | |
| GET /v1/models | |
| POST /v1/chat/completions {model, messages, temperature, top_p, max_tokens, | |
| stream, enable_thinking} (think -> reasoning_content) | |
| GET /health | |
| Auth: Authorization: Bearer 1234 (或 x-api-key: 1234) | |
| 示例: | |
| curl http://127.0.0.1:7862/v1/chat/completions \\ | |
| -H "Authorization: Bearer 1234" -H "Content-Type: application/json" \\ | |
| -d '{"model":"samai-4b","messages":[{"role":"user","content":"你好"}]}' | |
| 协议: 默认强制思考 (enable_thinking=true, 与 SFT 训练格式一致); | |
| 思考文本放 message.reasoning_content (DeepSeek-R1 风格), 正文放 message.content。 | |
| 上下文: config.max_position_embeddings (200k) 为硬上限, 超限返回 400。 | |
| 注意 T4 16GB 显存下 KV 实际可服务约 4~5 万 token 长输入。 | |
| """ | |
| import json | |
| import os | |
| import time | |
| import threading | |
| import uuid | |
| import torch | |
| from flask import Flask, request, jsonify, Response | |
| from transformers import (AutoModelForCausalLM, AutoTokenizer, | |
| StoppingCriteria, StoppingCriteriaList) | |
| MODEL_DIR = os.environ.get("S4_MODEL_DIR", "/content/samai-4b-sft") | |
| PORT = int(os.environ.get("S4_API_PORT", "7862")) | |
| API_KEY = os.environ.get("S4_API_KEY", "1234") | |
| MODEL_NAME = "samai-4b" | |
| DEFAULT_MAX_TOKENS = 1024 | |
| HARD_MAX_TOKENS = 8192 | |
| EOS_IDS = [1] | |
| STATE = {"loaded": False, "error": None, "t0": time.time()} | |
| LOCK = threading.Lock() | |
| MODEL = {"tok": None, "m": None, "ctx": 200000} | |
| app = Flask(__name__) | |
| class AntiLoop(StoppingCriteria): | |
| """末尾片段(L=3..16 token)连续重复 >=3 次 -> 复读退化, 提前截断.""" | |
| def __call__(self, input_ids, scores, **kwargs): | |
| ids = input_ids[0].tolist() | |
| tail = ids[-64:] | |
| if len(tail) < 9: | |
| return False | |
| for L in range(3, 17): | |
| if len(tail) < 3 * L: | |
| break | |
| seg = tail[-L:] | |
| if seg == tail[-2 * L:-L] == tail[-3 * L:-2 * L]: | |
| return True | |
| return False | |
| def build(): | |
| global MODEL, STATE | |
| print("[build] loading tokenizer ...", flush=True) | |
| tok = AutoTokenizer.from_pretrained(MODEL_DIR) | |
| print("[build] loading model fp16 -> cuda ...", flush=True) | |
| m = AutoModelForCausalLM.from_pretrained( | |
| MODEL_DIR, trust_remote_code=True, dtype=torch.float16).cuda().eval() | |
| ctx = int(getattr(m.config, "max_position_embeddings", 200000)) | |
| MODEL.update(tok=tok, m=m, ctx=ctx) | |
| STATE["loaded"] = True | |
| print(f"[build] DONE ctx={ctx}", flush=True) | |
| MARKS = ["<think>", "</think>", "<|User|>", "<|Bot|>", "<|System|>", "<|Tool|>", | |
| "<|start▁of▁sentence|>", "<|end▁of▁sentence|>", "<|▁pad▁|>", | |
| "<|start▁of▁text|>", "<|end▁of▁text|>", "<unk>"] | |
| def clean(s): | |
| for mk in MARKS: | |
| s = s.replace(mk, "") | |
| return s.strip() | |
| def split_think(text): | |
| text = text.split("<|end▁of▁sentence|>")[0] | |
| if "<think>" in text and "</think>" in text: | |
| _, b = text.split("<think>", 1) | |
| th, rest = b.split("</think>", 1) | |
| return th.strip(), clean(rest) | |
| if "<think>" in text: | |
| return text.split("<think>", 1)[1].strip(), "" | |
| if "</think>" in text: | |
| th, rest = text.split("</think>", 1) | |
| return th.strip(), clean(rest) | |
| return "", clean(text) | |
| def build_inputs(msgs, enable_think): | |
| tok = MODEL["tok"] | |
| text = tok.apply_chat_template(msgs, add_generation_prompt=True, | |
| tokenize=False, enable_thinking=enable_think) | |
| text += "\n" | |
| return tok(text, return_tensors="pt", return_dict=True, | |
| add_special_tokens=False) | |
| def err(message, etype="invalid_request_error", code=400): | |
| return jsonify({"error": {"message": message, "type": etype, | |
| "param": None, "code": code}}), code | |
| def auth_ok(): | |
| h = request.headers.get("Authorization", "") | |
| if h == f"Bearer {API_KEY}": | |
| return True | |
| if request.headers.get("x-api-key") == API_KEY: | |
| return True | |
| return False | |
| def gen_once(msgs, temperature, top_p, max_new, enable_think, stream_hints=None): | |
| """单次生成, 返回 (text, n_in, n_out, stopped, loop, steps).""" | |
| tok, model = MODEL["tok"], MODEL["m"] | |
| enc = build_inputs(msgs, enable_think) | |
| n_in = enc["input_ids"].shape[1] | |
| if n_in + max_new > MODEL["ctx"]: | |
| raise ValueError( | |
| f"prompt {n_in} + max_tokens {max_new} > context {MODEL['ctx']}") | |
| enc = {k: v.to(model.device) for k, v in enc.items()} | |
| do_sample = temperature and temperature > 0 | |
| gen_kw = dict(max_new_tokens=max_new, do_sample=do_sample, | |
| repetition_penalty=1.05, pad_token_id=2, eos_token_id=EOS_IDS, | |
| stopping_criteria=StoppingCriteriaList([AntiLoop()])) | |
| if do_sample: | |
| gen_kw["temperature"] = float(temperature) | |
| gen_kw["top_p"] = float(top_p) | |
| n_log = len(model._ponder_log) | |
| t0 = time.time() | |
| with torch.no_grad(): | |
| out = model.generate(**enc, **gen_kw) | |
| el = round(time.time() - t0, 1) | |
| text = tok.decode(out[0][n_in:], skip_special_tokens=False) | |
| entries = model._ponder_log[n_log:] | |
| steps = None | |
| if entries: | |
| steps = round(sum(e.get("steps_mean", e.get("executed", 0)) or 0 | |
| for e in entries) / len(entries), 2) | |
| n_out = int(out.shape[1] - n_in) | |
| stopped = int(out[0][-1]) in EOS_IDS | |
| loop = (not stopped) and (n_out < max_new) | |
| if stream_hints is not None: | |
| stream_hints.update(elapsed=el, steps=steps) | |
| return text, n_in, n_out, stopped, loop, steps | |
| def to_openai(cid, model_name, reply, think, finish, n_in, n_out, steps=None, | |
| elapsed=None, stream=False): | |
| msg = {"role": "assistant", "content": reply} | |
| if think: | |
| msg["reasoning_content"] = think | |
| body = { | |
| "id": f"chatcmpl-{cid}", "object": "chat.completion", | |
| "created": int(time.time()), "model": model_name, | |
| "choices": [{"index": 0, | |
| "message": msg, | |
| "finish_reason": finish}], | |
| "usage": {"prompt_tokens": n_in, "completion_tokens": n_out, | |
| "total_tokens": n_in + n_out}, | |
| } | |
| if steps is not None: | |
| body["samai_ponder_steps"] = steps | |
| if elapsed is not None: | |
| body["samai_elapsed_s"] = elapsed | |
| return body | |
| def list_models(): | |
| if not auth_ok(): | |
| return err("Invalid API key", code=401) | |
| return jsonify({"object": "list", "data": [ | |
| {"id": MODEL_NAME, "object": "model", "created": int(STATE["t0"]), | |
| "owned_by": "samai"}]}) | |
| def chat_completions(): | |
| if not auth_ok(): | |
| return err("Invalid API key", code=401) | |
| if not STATE["loaded"]: | |
| return err("model loading", "server_error", 503) | |
| d = request.get_json(force=True) | |
| msgs_in = d.get("messages") or [] | |
| msgs = [{"role": m.get("role", "user"), "content": str(m.get("content", ""))} | |
| for m in msgs_in | |
| if m.get("role") in ("system", "user", "assistant") and | |
| m.get("content") is not None] | |
| if not msgs: | |
| return err("messages is empty") | |
| enable_think = bool(d.get("enable_thinking", True)) | |
| temperature = d.get("temperature", 0.6 if enable_think else 1.0) | |
| temperature = 0.0 if temperature is None else float(temperature) | |
| top_p = float(d.get("top_p", 0.95)) | |
| max_tokens = int(d.get("max_tokens") or DEFAULT_MAX_TOKENS) | |
| max_tokens = max(1, min(HARD_MAX_TOKENS, max_tokens)) | |
| stream = bool(d.get("stream", False)) | |
| model_name = d.get("model") or MODEL_NAME | |
| with LOCK: | |
| try: | |
| hints = {} | |
| text, n_in, n_out, stopped, loop, steps = gen_once( | |
| msgs, temperature, top_p, max_tokens, enable_think, hints) | |
| except ValueError as e: | |
| return err(str(e), code=400) | |
| except Exception as e: # noqa: BLE001 | |
| import traceback | |
| traceback.print_exc() | |
| return err(repr(e)[:300], "internal_error", 500) | |
| think, reply = split_think(text) | |
| finish = "stop" if stopped else ("length" if not loop else "stop") | |
| cid = uuid.uuid4().hex[:24] | |
| if not stream: | |
| body = to_openai(cid, model_name, reply, think, finish, | |
| n_in, n_out, steps, hints.get("elapsed")) | |
| return jsonify(body) | |
| # ---- SSE 流式 (生成完整成文后切片推送; reasoning_content 先行) ---- | |
| def sse(): | |
| def chunk(delta, fr=None): | |
| c = {"id": f"chatcmpl-{cid}", "object": "chat.completion.chunk", | |
| "created": int(time.time()), "model": model_name, | |
| "choices": [{"index": 0, "delta": delta, "finish_reason": fr}]} | |
| return f"data: {json.dumps(c, ensure_ascii=False)}\n\n" | |
| yield chunk({"role": "assistant"}) | |
| if think: | |
| yield chunk({"reasoning_content": think}) | |
| step = 8 | |
| for i in range(0, len(reply), step): | |
| yield chunk({"content": reply[i:i + step]}) | |
| last = {"content": ""} | |
| if steps is not None: | |
| last["ponder_steps"] = steps | |
| yield chunk(last, fr=finish) | |
| u = {"prompt_tokens": n_in, "completion_tokens": n_out, | |
| "total_tokens": n_in + n_out} | |
| yield ("data: " + json.dumps( | |
| {"id": f"chatcmpl-{cid}", "object": "chat.completion.chunk", | |
| "created": int(time.time()), "model": model_name, | |
| "choices": [], "usage": u}, ensure_ascii=False) + "\n\n") | |
| yield "data: [DONE]\n\n" | |
| return Response(sse(), mimetype="text/event-stream", | |
| headers={"Cache-Control": "no-cache", | |
| "X-Accel-Buffering": "no"}) | |
| def health(): | |
| return jsonify({"loaded": STATE["loaded"], "error": STATE["error"], | |
| "service": "samai-4b-openai-api", "port": PORT, | |
| "ctx": MODEL["ctx"], | |
| "uptime_s": round(time.time() - STATE["t0"])}) | |
| def index(): | |
| return jsonify({"service": "samai-4b OpenAI-compatible API", | |
| "endpoints": ["/v1/models", "/v1/chat/completions", | |
| "/health"], | |
| "auth": "Authorization: Bearer <key>"}) | |
| if __name__ == "__main__": | |
| threading.Thread(target=build, daemon=True).start() | |
| app.run(host="0.0.0.0", port=PORT, threaded=True) | |