Vibrato / run_inference.py
Corolin's picture
feat: initial release of vibrato model artifacts (with LFS for weights and tokenizer)
79c6440
Raw History Blame Contribute Delete
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()