Download scripts/inference_example.py from LIJINGHAI111/DomSense: direct link, hf CLI and curl.
- Browser
- Download file 2.55 kB
-
https://huggingface.co/LIJINGHAI111/DomSense/resolve/main/scripts/inference_example.py
- Command line
-
hf download hf://LIJINGHAI111/DomSense/scripts/inference_example.py
-
curl -L -o inference_example.py https://huggingface.co/LIJINGHAI111/DomSense/resolve/main/scripts/inference_example.py
2.55 kB
| """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) | |