File size: 2,552 Bytes
f73f9b3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 | """DomSense / SystemOne 决策模型的推理示例。
这个脚本演示 Hugging Face 标准加载方式:
1. from_pretrained 加载 config + safetensors 权重
2. 用模型对「场景 + 若干选项」做多步前瞻决策
3. 展示选择概率、转移概率、期望回报与置信度
用法:
python scripts/inference_example.py [--model saved_model_dir]
无需联网;模型使用内置 lightweight 字符编码器,权重自包含。
"""
import argparse
import sys
from pathlib import Path
import torch
REPO_ROOT = Path(__file__).resolve().parent.parent.parent
sys.path.insert(0, str(REPO_ROOT))
from huggingface_repo import SystemOneModelForDecision
SCENARIOS = [
("医疗场景:患者持续低烧三天伴随咳嗽,选择下一步处置", ["建议自行服药观察", "立即前往发热门诊", "多喝水并休息"]),
("金融场景:你有一笔闲置资金希望一年内保值增值,风险承受力中等", ["存入活期存款", "购买稳健型理财", "全部投入高波动股票"]),
]
def main(model_dir: str):
print(f"加载模型: {model_dir}")
model = SystemOneModelForDecision.from_pretrained(model_dir)
model.eval()
total, trainable = model.count_parameters()
print(f"总参数量: {total:,} | 可训练: {trainable:,}")
for text, actions in SCENARIOS:
print("\n" + "=" * 60)
print(f"场景: {text}")
print(f"选项: {actions}")
with torch.no_grad():
out = model.forward([text], action_texts=[actions])
probs = out.choice_probs[0].tolist()
choices = out.choices[0]
conf = out.confidences[0].item()
exp_ret = out.expected_returns[0]
print(f" 选择: 选项{choices.item() + 1} ({actions[choices.item()]}) 置信度={conf:.4f}")
print(f" 各选项概率: {[round(p, 4) for p in probs]}")
if out.transition_probs is not None:
tp = out.transition_probs[0].tolist()
print(f" 转移概率(逐动作×结局): {[[round(x, 4) for x in row] for row in tp]}")
print(f" Q值: {[round(q, 4) for q in out.q_values[0].tolist()]}")
print(f" 期望回报: {round(exp_ret[choices.item()].item(), 4) if exp_ret.ndim else exp_ret}")
print("\n成功:标准 from_pretrained 加载 + 决策推理完成 ✅")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model", default=str(REPO_ROOT / "huggingface_repo" / "saved_model"))
args = parser.parse_args()
main(args.model)
|