samai-4b / openai_api.py
tchbcb's picture
200k ctx (rope_scaling hook) + openai-compatible api (key 1234) + bootstrap v2 (chat 7861 / api 7862) + model card
85708e3 verified
Raw History Blame Contribute Delete
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
@app.route("/v1/models")
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"}]})
@app.route("/v1/chat/completions", methods=["POST"])
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"})
@app.route("/health")
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"])})
@app.route("/")
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)