编写任务配置
任务配置把业务中的观测、动作、状态变化和训练目标写成一份 YAML。以倒立摆为例:角度与角速度是观测,电机力矩是动作,世界模型学习施加力矩后的状态变化,策略学习如何让摆杆保持直立。
本页先给出完整配置,再解释每一部分如何对应业务。已经有可运行配置时,可以直接进入设置训练数据、选择建模与控制算法、设置训练流程、评估与选择模型或改善预测与控制效果。
完整配置示例
下面的配置使用 BC 训练世界模型,再使用 PPO 训练控制策略。将它保存为倒立摆任务目录下的 config.yaml,并保留该任务的 reward.py。数据准备见倒立摆。数据文件在执行命令时指定,文件内其他相对路径以配置文件所在目录为基准。
name: pendulum_bc_ppo
version: "2.0"
description: 倒立摆增量式世界模型与 PPO 策略
graph:
transitions: auto
columns:
states:
- {name: cos_theta, min: -1.0, max: 1.0}
- {name: sin_theta, min: -1.0, max: 1.0}
- {name: theta_dot, min: -8.0, max: 8.0}
actions:
- {name: torque, min: -2.0, max: 2.0}
nodes:
actions:
inputs: [states]
network:
backbone: mlp
hidden_dims: [256, 256]
activation: leakyrelu
output_dist: TanhNormal
delta_states:
inputs: [states, actions]
network:
backbone: mlp
hidden_dims: [256, 256]
output_dist: normal
next_states:
inputs: [states, delta_states]
function: builtin.delta_add
data:
train_ratio: 0.8
split_mode: outside_traj
split_seed: 42
batch_size: 256
num_workers: 0
normalization: min_max_scaler
training:
device: auto
seed: 42
log_dir: logs
stages:
- name: env
algorithm: venv.bc
hyperparameters:
epochs: 200
optimizer:
lr: 0.0003
venv_target_nodes: [delta_states]
validation:
selection:
metric: val/rollout/mae
mode: min
one_step:
enabled: true
interval: 1
rollout:
enabled: true
horizon: 50
interval: 10
num_trajectories: 16
force_final: true
- name: policy
algorithm: policy.ppo
inherit_from: env
hyperparameters:
epochs: 300
policy_nodes: [actions]
rollout_horizon: 50
num_rollout_trajs: 256
bpc_type: bc
bpc_weight: 0.01
optimizer:
lr: 0.00004
reward:
path: reward.py
function: get_reward
output:
checkpoint_policy: best_only
tensorboard: true
plotting:
enabled: true
trend: true
interval: 10
onnx:
enabled: true
required: true
opset: 17
io_space: raw
validate: parity
export_dtype: keep
atol: 1.0e-5
rtol: 1.0e-4| 配置块 | 在任务中的作用 |
|---|---|
name、version、description | 任务名称、配置格式版本和说明;version 不等于 SDK 包版本 |
graph | 定义观测、动作、状态变化及各列的业务含义 |
data | 指定数据划分、归一化与读取方式 |
training.stages | 依次学习系统响应和控制动作 |
output | 保存选中的模型、日志和 ONNX 文件 |
本例学习率、轮数和推演长度是倒立摆的一组试验设置。换成其他设备时,应根据采样周期、响应时间和验证误差调整,算法默认值可在算法参数说明中查询。
使用默认参数开始训练
只写任务定义,也可以完成同样的两阶段流程。下面与源码中的 examples/pendulum/config.min.yaml 对应:
graph:
nodes:
actions: {inputs: [states]}
delta_states: {inputs: [states, actions]}
next_states: {inputs: [states, delta_states], function: builtin.delta_add}
columns:
states: [cos_theta, sin_theta, theta_dot]
actions: [{name: torque, min: -2.0, max: 2.0}]
training:
stages:
- {name: venv, algorithm: venv.bc}
- name: policy
algorithm: policy.ppo
inherit_from: venv
hyperparameters:
reward: {path: reward.py, function: get_reward}图:用变量关系、数据列和训练阶段描述一个任务。
网络和优化器等未填写的参数使用默认值。首次接入可以从这份配置开始,再逐项增加需要调整的参数;真实业务配置见任务列表。
定义观测与动作
将业务变量对应到节点
graph.nodes 描述一个时间步内的变量依赖,节点的 name 也是其输出名;映射写法直接以键名作为节点名。inputs 指定该节点可以读取哪些变量。例如,actions 读取当前 states,delta_states 同时读取状态和动作。
在温控业务中,可以对应为当前温度、加热命令和下一周期温度变化。室外温度、生产负荷等无法由控制器决定的量可以作为外部输入;不能把决策时刻尚不可获得的真实未来值放入策略输入。
图中定义了动作网络时,世界模型可以学习历史动作,策略阶段再优化该网络。若动作仅出现在其他节点的 inputs 中,它就是外部条件,推演时需要逐步提供。每个可训练节点还需具备相应标签或可计算的监督目标。
普通网络按 inputs 的声明顺序组织输入;改变顺序会改变特征布局。有状态网络使用框架规定的规范顺序,具体见历史窗口。复用模型时应同时保持列定义与输入布局一致。
指定哪些节点负责控制
training.stages 中每个阶段用 name 标识,用 algorithm 选择算法。inherit_from 引用前置世界模型阶段的名称;策略和控制器在该模型上评估动作。
policy_nodes 指定负责输出动作的节点。倒立摆图中可以唯一推导为 actions,因此最小配置省略了它。多个执行器或复杂中间节点的任务应显式声明,参见多控制节点。
PPO、SAC 和控制器需要 reward 表达业务目标。奖励配置中的 path 是 Python 文件路径,function 是函数名;世界模型和 policy.bc 不要求奖励函数。比如温控奖励可综合温度偏差与能耗,先用业务目标与奖励确定量纲和方向,再组织训练流程。
描述状态变化
学习增量或下一状态
倒立摆使用默认的 raw_delta 约定:网络预测 delta_states,builtin.delta_add 将增量加到当前状态,得到 next_states。graph.transitions: auto 根据 next_ 前缀识别状态转移;名称不遵循该规则时,可以显式填写转移映射。
也可以设置 graph.transition_contract: direct_next,让网络直接预测下一状态。两种写法需要与图结构一致;ADM2 要求显式增量结构。选择依据包括变量尺度、状态变化规律和所选算法,示例见准备数据。
把已知机理写入图中
# graph.nodes 中的一个节点
next_states:
inputs: [states, delta_states]
function: builtin.delta_add声明 function 的节点使用指定函数计算,其余网络节点从数据学习。自定义函数可写为 ./functions.py:compute_heat,并根据计算量的单位设置输入输出空间。解析后的 ref 保存函数引用;用户 YAML 直接把引用字符串写在 function 上。
differentiable 表示函数是否可微。需要沿函数向上游网络传递训练梯度时,应提供可微实现并声明 true;不能用该声明把不可微运算变成可微运算。机理节点和网络节点的组合方式见专家函数,特殊网络结构见自定义网络。
设置类型与边界
列顺序与变量类型
graph.columns 描述每个变量的维度。列的 name 表示业务含义,node 表示所属节点;映射写法可省略 node。列顺序必须与数据数组中的维度顺序一致,例如 states 三列依次为角度余弦、角度正弦和角速度。
type | 使用场景 | 配置重点 |
|---|---|---|
continuous | 温度、速度、力矩等连续量;默认类型 | 根据业务核对单位与范围 |
category | 开关模式、无序动作档位 | 用 values 显式给出类别及顺序 |
discrete | 有序且等间隔的离散数值 | 用边界与 num 定义取值数量 |
例如着陆器的四种发动机指令应作为类别动作:
# graph.columns 片段
columns:
action:
- {name: action, type: category, values: [0, 1, 2, 3]}values 的顺序决定类别编码,调整后需要重新训练并检查部署输入输出。若把档位误写为连续量,模型可能输出业务中不存在的中间值。
按设备能力填写动作范围
min、max 描述物理范围。倒立摆力矩允许范围为 −2 到 2,不能因为一批数据只出现过 −1.5 到 1.5,就把执行器范围改成后者。直接策略的 ONNX 导出要求每个动作维度具备明确的原始及归一化空间边界,预检会检查是否齐全。
可以用数据辅助检查范围:
revive suggest-bounds --config config.yaml --data train.npz --ratio 1.5候选范围按观测中心与半幅生成,已有 min/max 不会被覆盖,类别列不参与;缺失、常量或维度不匹配只做报告。默认命令输出建议,加入 --out config.with-suggestions.yaml 可另存配置。将建议用于动作列前,仍需按设备能力和工艺限制核对。
查看生效参数
revive validate --config config.yaml --train-data train.npz --show-defaults修改配置后,用这个命令检查实际值和来源。REVIVE 合并三类信息,优先级为运行时输入 > 用户 YAML > 算法默认值。
| 来源 | 适合放什么 | 如何核对 |
|---|---|---|
| 用户配置(Source YAML) | 图结构、数据规则、训练流程以及主动调整的参数 | 运行目录中的 config.source.yaml |
| 算法默认值(Defaults) | 未自行配置的网络、优化器和算法选项 | 算法参数说明及默认资源身份 |
| 运行时输入(RuntimeInputs) | 数据文件、日志目录、run ID、随机种子和续训入口 | 本次命令和 config.resolved.yaml |
hyperparameters 只需写需要调整的项,未覆盖项保留默认值;未知字段、层级或类型错误会被拒绝。config.resolved.yaml 记录完整生效配置,比较实验和恢复训练时应一并保存。
CLI 的 --train-data、--val-data、--log-dir、--seed 分别覆盖数据路径、验证路径、日志根目录和训练种子。设备、训练轮数、批量与学习率写在 YAML 中。参数位置不确定时,使用 revive info --explain <字段名> 或查看配置项说明。
检查配置
从任务目录执行:
revive validate --config config.yaml --train-data data/train.npz
revive train --config config.yaml --train-data data/train.npz --run-id pendulum-bc-ppo先核对图的输入依赖、数据列顺序、动作边界、奖励方向与阶段依赖,再运行预检。通过后,用新的 run ID 启动训练,避免覆盖用于比较的结果。
字段分为 basic、advanced、expert 三个阅读层级,可用 revive info --tier basic 查看任务定义相关项。实际必填条件由任务决定:类别列需要 values,函数节点需要函数引用,策略优化需要奖励和环境依赖。数据可推导的维度与算法默认参数通常可以省略。