使用模型
训练完成后,您可以从 logs/<run-id>/models/ 加载模型进行推理。两阶段训练生成 env.pt 和 policy.pt;仅训练世界模型时生成 env.pt。以下命令均在 examples/pendulum/ 目录下执行。
单阶段训练使用 world-model-demo,两阶段训练使用 pendulum-quickstart,分别对应上一页的两条训练路径。 如果只运行了 pendulum-smoke,请将相应路径中的运行名称替换为它;smoke 模型仅用于接口检查。
使用世界模型
世界模型可独立使用,无需先加载策略。下面加载上一页单阶段训练的模型,预测施加零力矩后的下一状态。
倒立摆配套图中包含 actions 网络节点。只把 actions 放入 infer_one_step 的输入字典, 不会覆盖这个节点的计算结果。评估指定操作时,应使用 graph.step_raw 的 override_values, 显式固定该步动作;输入动作和返回结果均使用原始物理单位。
import torch
from revive.export import load_env
env = load_env("logs/world-model-demo/models/env.pt")
state = {"states": torch.tensor([1.0, 0.0, 0.0], device=env.device)}
action = torch.tensor([0.0], device=env.device) # 力矩,单位 N·m
with torch.inference_mode():
prediction, _, _ = env.graph.step_raw(
state, override_values={"actions": action}, mode="mode"
)
print(prediction["next_states"])step_raw 返回三个对象:本步节点输出、本步全部变量值、下一步输入状态。 上例只使用第一个返回值中的 next_states。若图中动作本身就是外部输入、没有对应的动作网络节点, 则可以用 env.infer_one_step({"states": ..., "actions": ...}) 直接提供动作。
使用策略并由世界模型评估
以下加载两阶段训练得到的模型:策略先生成力矩,世界模型再按这个力矩预测下一状态。
import torch
from revive.export import load_env, load_policy
policy = load_policy("logs/pendulum-quickstart/models/policy.pt")
env = load_env("logs/pendulum-quickstart/models/env.pt")
# [cos θ, sin θ, θ̇]:摆杆静止于正上方
state = {"states": torch.tensor([1.0, 0.0, 0.0], dtype=torch.float32)}
with torch.inference_mode():
torque = policy.infer(state)["actions"]
outputs, _, _ = env.graph.step_raw(
{"states": state["states"].to(env.device)},
override_values={"actions": torque.to(env.device)},
mode="mode",
)
print(torque) # 策略生成的力矩,具体数值取决于训练结果
print(outputs["next_states"]) # 执行该力矩后的预测状态policy.infer 返回 {节点名: 张量} 字典。键名与决策流图一致,张量位于模型所在的设备上; 不同模型位于不同设备时,应将传给世界模型的状态和动作移动到 env.device。 输入输出使用原始物理单位,无需手工归一化,换算关系随模型文件保存。
批量调用时,将状态和动作分别组织为 [批大小, 状态维度] 与 [批大小, 动作维度]。
连续推演多步
给定一段动作序列
下面沿用上一节已加载的 env,连续 20 步施加零力矩。每步状态由上一步预测推进, 动作由 override_values 中长度为 20 的序列逐步提供:
horizon = 20
action_sequence = [
torch.tensor([[0.0]], dtype=torch.float32, device=env.device)
for _ in range(horizon)
]
with torch.inference_mode():
trajectory = env.graph.rollout_raw(
init_state={
"states": torch.tensor([[1.0, 0.0, 0.0]], dtype=torch.float32, device=env.device)
},
horizon=horizon,
override_values={"actions": action_sequence},
mode="mode",
)
theta_dot_curve = [step["next_states"][:, 2] for step in trajectory]返回列表长度等于 horizon,每个元素包含该步变量的取值。多步示例显式保留批次维度, 动作序列中的每一项形状为 [1, 1]。替换这些力矩值即可比较不同操作方案。 仅在初始状态中填入一个动作,不表示后续各步都会使用该动作。
策略逐步生成动作
如果动作由策略根据当前状态决定,需要每一步重新调用策略,再将动作交给世界模型。 下面继续使用已加载的 policy 与 env,模拟 20 步闭环:
sim_state = {"states": torch.tensor([[1.0, 0.0, 0.0]], device=env.device)}
closed_loop = []
with torch.inference_mode():
for _ in range(20):
closed_action = policy.infer(sim_state)["actions"].to(env.device)
step_outputs, _, next_state = env.graph.step_raw(
sim_state, override_values={"actions": closed_action}, mode="mode"
)
closed_loop.append((closed_action, step_outputs))
sim_state = {"states": next_state["states"]}这段闭环运行在学到的世界模型中,不能替代独立仿真或真实系统中的策略效果评估。 mode="mode" 选择分布的确定性输出;mode="sample" 从预测分布中采样。 预测分布的范围是否可靠,还需单独校准和验证。
推演过程中,未来真实状态只能用于事后计算误差,不能作为下一步的已知状态输入。
使用 ONNX 部署文件
ONNX 用于跨平台部署。标准训练流程要求自动导出并验证 ONNX 模型,成功完成后即可使用 生成的部署文件包。report.md 中可以查看数值一致性验证(parity)的结果:
| 组件 | 验证模式 | 容差 (rtol/atol) | onnxruntime | provider |
| `env` | parity | 0.0001 / 1e-05 | 1.23.2 | CPUExecutionProvider |
| `policy` | parity | 0.0001 / 1e-05 | 1.23.2 | CPUExecutionProvider |parity 检查相同输入下 PyTorch 与 ONNX Runtime 的输出是否在约定容差内一致。 checker 只检查模型结构,因此部署前还需要完成数值一致性验证。
models/policy.onnx.json 是模型说明文件,记录输入名称、维度、动作范围、依赖和验证标识。 它引用的 ONNX 文件保存在 .onnx_generations/<组件>/<身份哈希>/ 中。 部署时应完整复制说明文件及其引用的模型和依赖文件。
Python 的 override_values 是调用时的覆盖选项,不会改写已经导出的 ONNX 输入。 倒立摆完整图中的 actions 由图内网络计算;如果部署要求由外部指定动作,应使用动作作为外部输入的世界模型图, 并按该图重新训练、导出和验证。调用 ONNX 时以 .onnx.json 声明的输入为准。
需要使用其他验证数据重新导出时,可以执行:
revive export \
--artifact logs/pendulum-quickstart/models/policy.pt \
--io-space raw \
--parity-data data/pendulum.npzVerified ONNX bundle 已原子提交: logs/pendulum-quickstart/models
checker=PASS
runtime_parity=PASS (comparisons=2)--parity-data 会启用 --validate parity,不能与 --out 同时使用。 验证后的部署文件包保存在模型所属目录中。指定 --out 并使用 --validate checker 可将诊断文件写到其他位置,但结构检查结果不能作为上线依据。完整流程见 使用与部署模型。
含历史窗口的模型
当实际调用的节点依赖历史窗口或 GRU 隐状态时,需要在连续调用之间保存并回传运行时状态。 策略可通过 policy.is_stateful 和 policy.runtime_manifest 查看其要求;图中无关节点的历史状态不一定属于该策略接口:
初始化状态 → (可选)使用一段历史数据进行预热 → 每步推理并接收新状态 → 保存新状态不同运行序列应分别维护状态。本教程的倒立摆配置无需历史状态,其他模型的使用方法见 有状态推理。
接入自己的业务
将模型接入业务系统时,请确认以下信息:
- 键名换成自己图中的节点名,维度顺序与
columns声明一致; - 上线前确认 ONNX 一致性校验通过(
--validate parity),并把.onnx.json一并交给 部署方; - 策略输出仍应经过一层业务侧的限幅与联锁——模型保证输出落在声明的物理边界内, 但边界之内不等于任何时刻都可执行。
相关文档
后续可参考: