# -*- 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()