Download run_inference.py from Corolin/Vibrato: direct link, hf CLI and curl.
- Browser
- Download file 5.91 kB
-
https://huggingface.co/Corolin/Vibrato/resolve/main/run_inference.py
- Command line
-
hf download hf://Corolin/Vibrato/run_inference.py
-
curl -L -o run_inference.py https://huggingface.co/Corolin/Vibrato/resolve/main/run_inference.py
5.91 kB
| # -*- coding: utf-8 -*- | |
| """run_inference.py — Vibrato 开箱即用推理脚本 | |
| 支持两种运行模式: | |
| 1. e5s 满血模式(默认,端到端 ~15ms):结合 e5-small int8 语义底座 + Vibrato e5s 判头 | |
| 2. fast 纯文字极速模式(~3ms,零模型底座):仅加载 Vibrato v6 int8 判头 | |
| 依赖: | |
| pip install onnxruntime tokenizers numpy | |
| 用法: | |
| python run_inference.py # 运行内置测试用例 | |
| python run_inference.py --mode fast # 纯文字极速模式 | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| import sys | |
| import numpy as np | |
| import onnxruntime as ort | |
| from tokenizers import Tokenizer | |
| HERE = os.path.dirname(os.path.abspath(__file__)) | |
| sys.path.insert(0, HERE) | |
| import vib_messages | |
| from battery import FAMILY_IDS, NOUL_IDS | |
| DEMOS = [ | |
| [{"role": "user", "message": "今天累死啦,快抱抱"}], | |
| [{"role": "user", "message": "终于等到你啦~ 今天累死本宝宝了,快抱抱!"}], | |
| [ | |
| {"role": "user", "message": "你刚才为什么要那样敷衍我?给我说清楚!"}, | |
| {"role": "assistant", "message": "抱歉,我没有想敷衍你的意思。"}, | |
| {"role": "user", "message": "算了,说了也没用。"} | |
| ], | |
| [ | |
| {"role": "assistant", "message": "今天过得怎么样呀?", "pad": [0.4, -0.2, 0.1]}, | |
| {"role": "user", "message": "还行吧,就那样,一堆破事"} | |
| ] | |
| ] | |
| def load_e5s(): | |
| vocab = json.load(open(os.path.join(HERE, "vocab.json"), encoding="utf-8")) | |
| tok = Tokenizer.from_file(os.path.join(HERE, "tokenizer", "tokenizer.json")) | |
| so = ort.SessionOptions() | |
| so.intra_op_num_threads = 1 | |
| head = ort.InferenceSession(os.path.join(HERE, "models", "vibrato_e5s_int8.onnx"), so, providers=["CPUExecutionProvider"]) | |
| backbone = ort.InferenceSession(os.path.join(HERE, "models", "e5s_int8.onnx"), so, providers=["CPUExecutionProvider"]) | |
| return vocab, tok, head, backbone | |
| def load_fast(): | |
| vocab = json.load(open(os.path.join(HERE, "vocab.json"), encoding="utf-8")) | |
| so = ort.SessionOptions() | |
| so.intra_op_num_threads = 1 | |
| head = ort.InferenceSession(os.path.join(HERE, "models", "vibrato_v6_int8.onnx"), so, providers=["CPUExecutionProvider"]) | |
| return vocab, head | |
| def get_e5s_feats(backbone, tok, text, max_chars): | |
| n = min(len(text), max_chars) | |
| out = np.zeros((max_chars, 384), dtype=np.float32) | |
| enc = tok.encode(text, add_special_tokens=False) | |
| ids = np.array([enc.ids], dtype=np.int64) | |
| mask = np.ones_like(ids) | |
| hidden = backbone.run(None, {"input_ids": ids, "attention_mask": mask})[0][0] | |
| for (s, e), vec in zip(enc.offsets, hidden): | |
| lo, hi = s, min(e, n) | |
| if hi > lo: | |
| out[lo:hi] = vec | |
| return out | |
| def score_messages(messages, vocab, head, backbone=None, tok=None, max_len=512): | |
| vib_messages.validate_messages(messages) | |
| echo = next((m["pad"] for m in reversed(messages[:-1]) if m["role"] == "user" and "pad" in m), [0.0] * 3) | |
| units = vib_messages.sequence_units(messages, max_window=max_len) | |
| ids, spos, svals, smask = [], [], [], [] | |
| unk = len(vocab) | |
| for u in units: | |
| if u["kind"] == "text": | |
| for ch in u["text"]: | |
| ids.append(vocab.get(ch, unk)) | |
| spos.append(False) | |
| svals.extend([None] * len(u["text"])) | |
| smask.extend([None] * len(u["text"])) | |
| else: | |
| ids.append(0) | |
| spos.append(True) | |
| svals.append(u["vals"]) | |
| smask.append(u["mask"]) | |
| L = len(ids) | |
| z = [0.0] * vib_messages.STATE_DIM | |
| if backbone is not None and tok is not None: | |
| feats = np.zeros((L, 384), dtype=np.float32) | |
| pos = 0 | |
| for u in units: | |
| if u["kind"] == "text": | |
| if u["text"].strip(): | |
| f = get_e5s_feats(backbone, tok, u["text"], L) | |
| n = min(len(u["text"]), L - pos) | |
| feats[pos:pos + n] = f[:n] | |
| pos += len(u["text"]) | |
| else: | |
| pos += 1 | |
| else: | |
| feats = np.zeros((L, 1024), dtype=np.float32) | |
| feed = { | |
| "ids": np.array([ids], dtype=np.int64), | |
| "state_pos": np.array([spos], dtype=np.float32), | |
| "state_vals": np.array([[v if v is not None else z for v in svals]], dtype=np.float32), | |
| "state_mask": np.array([[v if v is not None else z for v in smask]], dtype=np.float32), | |
| "echo_prev": np.array([echo], dtype=np.float32), | |
| "feats": feats[None] | |
| } | |
| o = head.run(None, feed) | |
| pad = o[2][0].tolist() | |
| fam = FAMILY_IDS[int(o[3][0].argmax())] | |
| noul = {q: bool(o[4][0, j, 1] > 0.5) for j, q in enumerate(NOUL_IDS)} | |
| conf = float(o[5][0]) | |
| return pad, fam, noul, conf | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--mode", choices=["e5s", "fast"], default="e5s", help="e5s: 离线端到端满血 (384-dim) | fast: 纯文字极速 (零底座)") | |
| args = parser.parse_args() | |
| print(f"=== Vibrato 推理引擎启动 (模式: {args.mode}) ===") | |
| if args.mode == "e5s": | |
| vocab, tok, head, backbone = load_e5s() | |
| else: | |
| vocab, head = load_fast() | |
| backbone, tok = None, None | |
| for i, msgs in enumerate(DEMOS): | |
| target = msgs[-1]["message"] | |
| pad, fam, noul, conf = score_messages(msgs, vocab, head, backbone=backbone, tok=tok) | |
| print(f"\n[测试例 {i+1}] 输入: 「{target}」") | |
| print(f" -> PAD 情绪分布: Pleasure={pad[0]:+.2f}, Arousal={pad[1]:+.2f}, Dominance={pad[2]:+.2f}") | |
| print(f" -> 情绪族 (Family): {fam}") | |
| print(f" -> 置信度 (Conf): {conf:.2f}") | |
| active_nouls = [k for k, v in noul.items() if v] | |
| print(f" -> 语用标记 (Nouls): {active_nouls if active_nouls else '无特殊标记'}") | |
| if __name__ == "__main__": | |
| main() | |