自定义网络
当默认网络不能表达您的结构需求时,可以为图中的某个节点注册自定义网络。其余数据、训练阶段、选模和模型保存流程仍使用统一配置。本页先实现一个温度增量网络,再说明带物理先验和循环状态的用法。
使用场景
先从任务需要表达的信息判断是否需要自定义结构。只调整网络深度或宽度时,可以直接使用内置骨干及其配置;当计算结构发生变化时,再引入自定义网络。
| 业务需求 | 网络需要表达的内容 | 实际例子 |
|---|---|---|
| 为不同输入设计专用处理方式 | 按变量分组、组合分支或特定激活方式 | 将温度、功率和外扰组织成温度增量预测 |
| 保留机理预测,同时学习偏差 | 物理先验与可学习残差 | 无人机利用动力学先验预测运动,再用网络修正 |
| 当前观测不足以描述过程 | 跨步循环状态 | 利用连续飞行指令与运动记录学习动态响应 |
无人机建模使用物理先验与 GRU,预测速度和姿态变化。这样可以保留已有动力学知识,也能学习实际飞行记录中的偏差。是否改善预测,应在相同数据划分和预测窗口下比较。
示例:温度增量网络
设备温度受当前温度、加热功率、风机和外界条件共同影响。下面用一层隐藏层学习温度增量。已知的 heat = 4 × heater 保留在专家函数中,网络负责学习未知的温度响应。
完整实现位于 examples/thermal_workflows/custom_networks/thermal_mlp.py:
import torch.nn as nn
from revive.core.networks.base import ConfigurableNetwork
from revive.core.networks.mapping import FeatureMapping
from revive.core.networks.registry import register_network
@register_network("thermal_tutorial_mlp")
class ThermalMLP(ConfigurableNetwork):
def __init__(self, in_features, out_features, width=16, feature_mapping=None):
super().__init__(feature_mapping)
if width < 1:
raise ValueError("width must be positive")
self.in_features, self.out_features, self.width = in_features, out_features, width
self.net = nn.Sequential(
nn.Linear(in_features, width), nn.Tanh(), nn.Linear(width, out_features)
)
def forward(self, x):
return self.net(x)
@classmethod
def build_kwargs(cls, params, ctx):
unknown = set(params.custom_params) - {"width"}
if unknown:
raise ValueError(f"unknown ThermalMLP parameters: {sorted(unknown)}")
return {
**super().build_kwargs(params, ctx),
"width": int(params.custom_params.get("width", 16)),
}
def get_config(self):
return {
"in_features": self.in_features,
"out_features": self.out_features,
"width": self.width,
"feature_mapping": self.mapping.to_dict() if self.mapping else None,
}
@classmethod
def from_config(cls, config):
config = dict(config)
mapping = config.pop("feature_mapping", None)
return cls(**config, feature_mapping=FeatureMapping.from_dict(mapping) if mapping else None)输入、输出与构建参数
forward 接收拼接后的 processed 特征,即框架处理后的网络输入。in_features 由实际输入布局确定;out_features 由输出分布需要的参数宽度确定,应直接使用框架传入的值。
width 决定隐藏层宽度。build_kwargs 从 custom_params 读取它,并拒绝未知参数;get_config 与 from_config 保存和重建网络结构,权重由模型加载流程恢复。需要按业务变量读取特征位置时,使用 FeatureMapping。
配置网络节点
以下是完整示例 examples/thermal_workflows/config.yaml 中的片段:
graph:
nodes:
delta_temperature:
inputs: [temperature, heat, fan, ambient, load]
network:
backbone: thermal_tutorial_mlp
custom_params: {width: 16}
output_dist: normalbackbone 与注册名一致;自定义参数放在 custom_params 内。output_dist: normal 决定网络输出如何解释成分布,因此网络输出宽度不能简单写死为一个温度值。
将实现放在配置同级的 custom_networks/ 中,revive validate 和 revive train 会自动发现。Python 加载模型前,先运行 revive.discover_project_components(project_dir)。
多步输入与运行状态
普通 MLP 自身没有循环状态;如果输入加入 temperature@[-3:],整张图仍然需要保存历史缓冲区。GRU 等循环结构还需要管理隐藏状态,网络扩展应使用 StatefulNetwork 协议并实现对应的状态接口。
同一轨迹内持续传递状态,新轨迹开始时重新初始化。具体调用见历史窗口和有状态模型部署。
训练与检查结果
在仓库根目录运行。已有示例数据时跳过生成步骤,生成命令会覆盖原示例数据。
python examples/thermal_workflows/prepare_data.py
revive validate --config examples/thermal_workflows/config.yaml
revive train --config examples/thermal_workflows/config.yaml --run-id thermal_network查看 示例目录下的 logs/thermal_network/report.md 中的多步预测误差。models/env.pt 是选中的世界模型。修改 width 后用新的 run ID 训练,在相同数据划分和窗口下比较精度与耗时。
自定义网络至少应检查三个环节:前向输出形状正确;训练时参数得到更新;保存并重新加载后,相同输入给出一致输出。业务效果则通过未参与训练的轨迹评估,关注温度变化趋势、外扰响应和长时误差。
模型加载与导出
部署环境需要保留自定义网络代码及其依赖,并在加载前完成注册。网络参数通过 get_config / from_config 重建,因此这两个方法必须覆盖构造结构所需的全部参数。
revive export --artifact examples/thermal_workflows/logs/thermal_network/models/env.pt \
--project examples/thermal_workflows导出后核对 PyTorch 与 ONNX 的输出。包含循环状态时,按模型 manifest 的初始化和逐步调用接口传递状态。修改输入变量、列宽或历史长度会改变模型结构,应重新训练并检查导出结果。
使用混合网络时的正则项
内置 ParallelHybridNet 可以组合不同计算分支。其负门控熵正则鼓励混合权重保留多条分支,默认系数为 0.1;需要关闭时显式设为 0。训练日志分别记录 train/supervised_loss 和 train/auxiliary_loss。如果业务要求固定优先使用某条机理关系,应在图结构中直接表达该关系。
常见问题
| 现象 | 处理方法 |
|---|---|
| 找不到注册名 | 检查目录位置、装饰器名称和组件发现步骤 |
| 自定义参数被拒绝 | 将参数放到 custom_params,并在 build_kwargs 中显式读取 |
| 输出维度不匹配 | 使用框架传入的 out_features,核对输出分布 |
| 加载时无法构建网络 | 检查 get_config / from_config 是否保存了全部结构参数和映射 |
| 连续调用结果异常 | 检查历史缓冲区和隐藏状态是否持续传递、是否在新轨迹重置 |
| 训练可用但导出失败 | 检查自定义算子与运行状态的导出支持,参考使用与部署模型 |