跳转到内容

自定义网络 ​

当默认网络不能表达您的结构需求时,可以为图中的某个节点注册自定义网络。其余数据、训练阶段、选模和模型保存流程仍使用统一配置。本页先实现一个温度增量网络,再说明带物理先验和循环状态的用法。

使用场景 ​

先从任务需要表达的信息判断是否需要自定义结构。只调整网络深度或宽度时,可以直接使用内置骨干及其配置;当计算结构发生变化时,再引入自定义网络。

业务需求网络需要表达的内容实际例子
为不同输入设计专用处理方式按变量分组、组合分支或特定激活方式将温度、功率和外扰组织成温度增量预测
保留机理预测,同时学习偏差物理先验与可学习残差无人机利用动力学先验预测运动,再用网络修正
当前观测不足以描述过程跨步循环状态利用连续飞行指令与运动记录学习动态响应

无人机建模使用物理先验与 GRU,预测速度和姿态变化。这样可以保留已有动力学知识,也能学习实际飞行记录中的偏差。是否改善预测,应在相同数据划分和预测窗口下比较。

示例:温度增量网络 ​

设备温度受当前温度、加热功率、风机和外界条件共同影响。下面用一层隐藏层学习温度增量。已知的 heat = 4 × heater 保留在专家函数中,网络负责学习未知的温度响应。

完整实现位于 examples/thermal_workflows/custom_networks/thermal_mlp.py:

python
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 中的片段:

yaml
graph:
  nodes:
    delta_temperature:
      inputs: [temperature, heat, fan, ambient, load]
      network:
        backbone: thermal_tutorial_mlp
        custom_params: {width: 16}
      output_dist: normal

backbone 与注册名一致;自定义参数放在 custom_params 内。output_dist: normal 决定网络输出如何解释成分布,因此网络输出宽度不能简单写死为一个温度值。

将实现放在配置同级的 custom_networks/ 中,revive validate 和 revive train 会自动发现。Python 加载模型前,先运行 revive.discover_project_components(project_dir)。

多步输入与运行状态 ​

普通 MLP 自身没有循环状态;如果输入加入 temperature@[-3:],整张图仍然需要保存历史缓冲区。GRU 等循环结构还需要管理隐藏状态,网络扩展应使用 StatefulNetwork 协议并实现对应的状态接口。

同一轨迹内持续传递状态,新轨迹开始时重新初始化。具体调用见历史窗口和有状态模型部署。

训练与检查结果 ​

在仓库根目录运行。已有示例数据时跳过生成步骤,生成命令会覆盖原示例数据。

bash
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 重建,因此这两个方法必须覆盖构造结构所需的全部参数。

bash
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 是否保存了全部结构参数和映射
连续调用结果异常检查历史缓冲区和隐藏状态是否持续传递、是否在新轨迹重置
训练可用但导出失败检查自定义算子与运行状态的导出支持,参考使用与部署模型