Download chat_server.py from tchbcb/samai-4b: direct link, hf CLI and curl.
- Browser
- Download file 14.4 kB
-
https://huggingface.co/tchbcb/samai-4b/resolve/main/chat_server.py
- Command line
-
hf download hf://tchbcb/samai-4b/chat_server.py
-
curl -L -o chat_server.py https://huggingface.co/tchbcb/samai-4b/resolve/main/chat_server.py
14.4 kB
| # -*- coding: utf-8 -*- | |
| """s4_server.py — samai-4b 网页聊天demo (Colab T4, port 7861) [v4, 移植自 r18 v3.2] | |
| 协议适配 Spark 模板: | |
| - force_think=True(默认): 生成提示 = ...<|Bot|><think> + "\\n" (与 SFT 训练格式一致) | |
| - force_think=False: 生成提示 = ...<|Bot|></think> + "\\n" (跳过思考) | |
| - eos=[1] (<|end▁of▁sentence|>); decode 不剥特殊token, 手工清理模板标记 | |
| 继承: AntiLoop 复读截断 / best-of-n 投票 / 强制思考开关 / ponder 步数展示 | |
| 端点: GET / | GET /health | POST /chat {message, history, force_think, n_votes} | |
| """ | |
| import json, os, re, threading, time, uuid | |
| from collections import Counter | |
| import torch | |
| from flask import Flask, request, jsonify, Response | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, StoppingCriteria, StoppingCriteriaList | |
| MODEL_DIR = "/content/samai-4b-sft" | |
| PORT = 7861 | |
| MAX_NEW_THINK = 320 | |
| MAX_NEW_AUTO = 192 | |
| MAX_PROMPT_TOKENS = 900 | |
| EOS_IDS = [1] | |
| STATE = {"loaded": False, "error": None, "ckpt": os.path.basename(MODEL_DIR), | |
| "t0": time.time()} | |
| LOCK = threading.Lock() | |
| JOBS = {} | |
| JOBS_MU = threading.Lock() | |
| MODEL = {"tok": None, "m": None} | |
| 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) | |
| model = AutoModelForCausalLM.from_pretrained( | |
| MODEL_DIR, trust_remote_code=True, dtype=torch.float16).cuda().eval() | |
| MODEL["tok"], MODEL["m"] = tok, model | |
| STATE["loaded"] = True | |
| print("[build] DONE", type(model).__name__, flush=True) | |
| MARKS = ["<think>", "</think>", "<|User|>", "<|Bot|>", "<|System|>", "<|Tool|>", | |
| "<|start▁of▁sentence|>", "<|end▁of▁sentence|>", "<|▁pad▁|>", | |
| "<|start▁of▁text|>", "<|end▁of▁text|>", "<unk>"] | |
| def build_inputs(msgs, force_think): | |
| tok = MODEL["tok"] | |
| text = tok.apply_chat_template(msgs, add_generation_prompt=True, | |
| tokenize=False, enable_thinking=force_think) | |
| # force=True: ...<|Bot|><think> + "\n"; force=False: ...<|Bot|></think> + "\n" | |
| text += "\n" | |
| return tok(text, return_tensors="pt", return_dict=True, add_special_tokens=False) | |
| def split_think(text): | |
| text = text.split("<|end▁of▁sentence|>")[0] | |
| def clean(s): | |
| for mk in MARKS: | |
| s = s.replace(mk, "") | |
| return s.strip() | |
| if "<think>" in text and "</think>" in text: | |
| a, b = text.split("<think>", 1) | |
| th, rest = b.split("</think>", 1) | |
| return th.strip(), clean(a + 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 health(): | |
| gpu = "" | |
| try: | |
| import subprocess | |
| r = subprocess.run(["nvidia-smi", "--query-gpu=memory.used", | |
| "--format=csv,noheader"], capture_output=True, text=True) | |
| gpu = r.stdout.strip().splitlines()[0] if r.stdout.strip() else "" | |
| except Exception: | |
| pass | |
| return jsonify({"loaded": STATE["loaded"], "error": STATE["error"], | |
| "ckpt": STATE["ckpt"], | |
| "uptime_s": round(time.time() - STATE["t0"]), | |
| "gpu_mem_used": gpu}) | |
| def chat(): | |
| if not STATE["loaded"]: | |
| return jsonify({"error": "model loading"}), 503 | |
| d = request.get_json(force=True) | |
| msg = (d.get("message") or "").strip() | |
| history = d.get("history") or [] | |
| force = bool(d.get("force_think", True)) | |
| try: | |
| n_votes = max(1, min(5, int(d.get("n_votes") or 1))) | |
| except Exception: | |
| n_votes = 1 | |
| if not msg: | |
| return jsonify({"error": "empty"}), 400 | |
| msgs = [m for m in history if m.get("role") in ("user", "assistant") and m.get("content")] | |
| msgs = msgs[-8:] + [{"role": "user", "content": msg}] | |
| jid = uuid.uuid4().hex[:12] | |
| with JOBS_MU: | |
| JOBS[jid] = {"status": "running", "reply": "", "think": "", "steps": None, | |
| "elapsed_s": 0, "error": None} | |
| if len(JOBS) > 32: | |
| for k in [k for k, v in JOBS.items() if v["status"] != "running"][:-16]: | |
| JOBS.pop(k, None) | |
| threading.Thread(target=run_job, args=(jid, msgs, force, n_votes), | |
| daemon=True).start() | |
| return jsonify({"job_id": jid}) | |
| def run_job(jid, msgs, force_think, n_votes=1): | |
| job = JOBS[jid] | |
| tok, model = MODEL["tok"], MODEL["m"] | |
| with LOCK: | |
| try: | |
| enc = build_inputs(msgs, force_think) | |
| n_in = enc["input_ids"].shape[1] | |
| if n_in > MAX_PROMPT_TOKENS: | |
| msgs = msgs[-4:] | |
| enc = build_inputs(msgs, force_think) | |
| n_in = enc["input_ids"].shape[1] | |
| enc = {k: v.to(model.device) for k, v in enc.items()} | |
| max_new = MAX_NEW_THINK if force_think else MAX_NEW_AUTO | |
| temp = 0.6 if force_think else 1.0 | |
| def gen_once(): | |
| n_log = len(model._ponder_log) | |
| t0 = time.time() | |
| with torch.no_grad(): | |
| out = model.generate(**enc, max_new_tokens=max_new, | |
| do_sample=True, temperature=temp, top_p=0.95, | |
| repetition_penalty=1.05, | |
| pad_token_id=2, | |
| eos_token_id=EOS_IDS, | |
| stopping_criteria=StoppingCriteriaList([AntiLoop()])) | |
| 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) | |
| new_tokens = int(out.shape[1] - n_in) | |
| stopped = int(out[0][-1]) in EOS_IDS | |
| loop_hit = (not stopped) and (new_tokens < max_new) | |
| return text, steps, new_tokens, stopped, loop_hit | |
| def vote_key(reply): | |
| nums = re.findall(r"\d[\d,]*(?:\.\d+)?", (reply or "").replace(",", "")) | |
| if nums: | |
| return "n:" + nums[-1] | |
| return "t:" + (reply or "").strip()[:40] | |
| t0 = time.time() | |
| votes = None | |
| if n_votes <= 1: | |
| text, steps, new_tokens, stopped, loop = gen_once() | |
| think, reply = split_think(text) | |
| else: | |
| rs = [] | |
| for _ in range(n_votes): | |
| text_i, st, nt, sp, lp = gen_once() | |
| th, rp = split_think(text_i) | |
| rs.append({"think": th, "reply": rp, "steps": st, | |
| "new_tokens": nt, "stopped": sp, "loop": lp}) | |
| cnt = Counter(vote_key(r["reply"]) for r in rs) | |
| bk = cnt.most_common(1)[0][0] | |
| best = next(r for r in rs if vote_key(r["reply"]) == bk) | |
| think, reply = best["think"], best["reply"] | |
| steps, new_tokens = best["steps"], best["new_tokens"] | |
| stopped, loop = best["stopped"], best["loop"] | |
| votes = {(k[2:] if k[:2] in ("n:", "t:") else k): c | |
| for k, c in cnt.most_common()} | |
| el = round(time.time() - t0, 1) | |
| job.update({"status": "done", "reply": reply or text[:400], | |
| "think": think, "steps": steps, "elapsed_s": el, | |
| "new_tokens": new_tokens, "stopped": stopped, "loop": loop, | |
| "votes": votes, | |
| "mode": ("think" if force_think else "auto") | |
| + ("x%d" % n_votes if n_votes > 1 else ""), | |
| "temp": temp}) | |
| except Exception as e: | |
| import traceback | |
| traceback.print_exc() | |
| job.update({"status": "error", "error": repr(e)[:300]}) | |
| def result(): | |
| jid = request.args.get("id", "") | |
| with JOBS_MU: | |
| job = JOBS.get(jid) | |
| if job is None: | |
| return jsonify({"error": "unknown job"}), 404 | |
| return jsonify(dict(job)) | |
| PAGE = """<!doctype html><html lang="zh"><head><meta charset="utf-8"> | |
| <meta name="viewport" content="width=device-width,initial-scale=1"> | |
| <title>samai-4b · chat</title><style> | |
| :root{--bg:#0f1115;--card:#171a21;--line:#262b36;--fg:#e6e9ef;--dim:#8b93a5;--acc:#7c5bff;--ok:#3fbf7f} | |
| *{box-sizing:border-box}body{margin:0;background:var(--bg);color:var(--fg); | |
| font-family:-apple-system,'PingFang SC','Microsoft YaHei',sans-serif;display:flex;flex-direction:column;height:100vh} | |
| header{padding:12px 18px;border-bottom:1px solid var(--line);display:flex;align-items:center;gap:10px;flex-wrap:wrap} | |
| header b{font-size:15px}.pill{font-size:12px;padding:2px 10px;border-radius:999px;border:1px solid var(--line);color:var(--dim)} | |
| .pill.ok{color:var(--ok);border-color:var(--ok)} | |
| label.pill{cursor:pointer;user-select:none;display:flex;align-items:center;gap:5px} | |
| label.pill input{accent-color:#7c5bff} | |
| #log{flex:1;overflow-y:auto;padding:18px;display:flex;flex-direction:column;gap:14px} | |
| .msg{max-width:82%;padding:10px 14px;border-radius:14px;line-height:1.6;white-space:pre-wrap;word-break:break-word;font-size:14.5px} | |
| .u{align-self:flex-end;background:var(--acc);color:#fff;border-bottom-right-radius:4px} | |
| .a{align-self:flex-start;background:var(--card);border:1px solid var(--line);border-bottom-left-radius:4px} | |
| .meta{align-self:flex-start;font-size:11.5px;color:var(--dim)} | |
| details{margin-top:6px}summary{cursor:pointer;color:var(--dim);font-size:12px} | |
| details pre{white-space:pre-wrap;font-size:12.5px;color:var(--dim);margin:6px 0 0} | |
| footer{border-top:1px solid var(--line);padding:12px;display:flex;gap:10px} | |
| textarea{flex:1;background:var(--card);border:1px solid var(--line);color:var(--fg);border-radius:10px; | |
| padding:10px 12px;font-size:14.5px;resize:none;height:52px;font-family:inherit} | |
| button{background:var(--acc);border:0;color:#fff;border-radius:10px;padding:0 22px;font-size:14.5px;cursor:pointer} | |
| button:disabled{opacity:.5}</style></head><body> | |
| <header><b>samai-4b · pnet-dMoE (Spark-X2.5 骨干)</b><span class="pill" id="st">loading…</span><span class="pill" id="ck"></span> | |
| <label class="pill" title="开: 强制思考 T=0.6 (数学/推理稳) · 关: 跳过思考直接答 T=1.0"> | |
| <input type="checkbox" id="ft" checked>💭 强制思考</label> | |
| <label class="pill" title="best-of-n: 采样5次提取答案取众数 (数学/计算题稳, 耗时≈×5)"> | |
| <input type="checkbox" id="vb">🎯 投票×5</label></header> | |
| <div id="log"><div class="meta">samai-4b v4 · Spark-X2.5-4B + Pondernet(8专家/后8层) SFT · 强制思考 T=0.6 max320 · 反复读截断 · 🎯投票×5 · eos=[1]</div></div> | |
| <footer><textarea id="in" placeholder="说点什么… (Enter 发送)"></textarea><button id="go">发送</button></footer> | |
| <script> | |
| const log=document.getElementById('log'),inp=document.getElementById('in'),go=document.getElementById('go'), | |
| st=document.getElementById('st'),ck=document.getElementById('ck'),ft=document.getElementById('ft'), | |
| vb=document.getElementById('vb');let hist=[],busy=false; | |
| function esc(s){return s.replace(/[&<>]/g,c=>({'&':'&','<':'<','>':'>'}[c]))} | |
| function add(cls,html){const d=document.createElement('div');d.className=cls;d.innerHTML=html;log.appendChild(d);log.scrollTop=1e9;return d} | |
| async function poll(){try{const r=await fetch('/health');const j=await r.json(); | |
| if(j.loaded){st.textContent='ready';st.className='pill ok';ck.textContent=j.ckpt;} | |
| else{st.textContent=j.error?('error: '+j.error):'loading model…';}}catch(e){st.textContent='offline'}setTimeout(poll,3000)}poll(); | |
| async function send(){if(busy)return;const m=inp.value.trim();if(!m)return;busy=true;go.disabled=true;inp.value=''; | |
| add('u',esc(m));const w=add('a meta','思考中…'); | |
| try{const r=await fetch('/chat',{method:'POST',headers:{'Content-Type':'application/json'}, | |
| body:JSON.stringify({message:m,history:hist,force_think:ft.checked,n_votes:(vb.checked?5:1)})});const j=await r.json(); | |
| if(j.error||!j.job_id){w.remove();add('a meta','⚠ '+esc(j.error||'no job'));busy=false;go.disabled=false;return} | |
| let res=null;for(let i=0;i<400;i++){await new Promise(s=>setTimeout(s,1500)); | |
| const rr=await fetch('/result?id='+j.job_id);res=await rr.json(); | |
| if(res.status!=='running')break;} | |
| w.remove(); | |
| if(res.error){add('a meta','⚠ '+esc(res.error))}else{ | |
| let h=esc(res.reply||''); | |
| if(res.think)h='<details><summary>💭 思考过程</summary><pre>'+esc(res.think)+'</pre></details>'+h; | |
| add('a',h); | |
| add('meta','🧠 '+(res.mode&&res.mode.indexOf('think')===0?'强制思考 T=0.6':'自动 T=1.0')+' · '+res.steps+' 步 · '+res.elapsed_s+'s · '+(res.new_tokens||'')+' tokens' | |
| +(res.loop?' · ⚠反循环截断':(res.stopped?'':' · ⚠长度截断')) | |
| +(res.votes?' · 🎯 '+Object.entries(res.votes).map(p=>p[0]+'×'+p[1]).join(' / '):'')); | |
| hist.push({role:'user',content:m},{role:'assistant',content:res.reply||''});hist=hist.slice(-8);} | |
| }catch(e){w.remove();add('a meta','⚠ '+e)}busy=false;go.disabled=false;inp.focus()} | |
| go.onclick=send;inp.addEventListener('keydown',e=>{if(e.key==='Enter'&&!e.shiftKey){e.preventDefault();send()}}); | |
| </script></body></html>""" | |
| def index(): | |
| return Response(PAGE, mimetype="text/html") | |
| if __name__ == "__main__": | |
| threading.Thread(target=build, daemon=True).start() | |
| app.run(host="0.0.0.0", port=PORT, threaded=True) | |