You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch Lightning中如何从Checkpoint加载嵌套模型?

解决方案

核心问题是加载主模型checkpoint时,PrimaryModel的__init__仍会强制初始化NestedModel并要求传入checkpoint路径,但主模型的state_dict已经包含了mlp1、mlp2的参数,无需再从NestedModel的checkpoint加载。以下是两种简洁的修改方式:

方案一:提前定义mlp结构,按需加载NestedModel参数

直接在PrimaryModel的__init__中定义与NestedModel完全一致的mlp1、mlp2结构,仅当提供NestedModel checkpoint路径时,才将其参数加载到mlp中。加载主模型时无需传入该路径,主模型的state_dict会自动覆盖mlp参数。

修改后的PrimaryModel代码:

import torch.nn as nn
from pytorch_lightning import LightningModule

class PrimaryModel(LightningModule):
    def __init__(self, nested_model_ckpt_path=None, **kwargs):
        super().__init__()
        # 定义与NestedModel中完全一致的mlp1、mlp2结构
        # 示例结构,需替换为你实际的NestedModel.mlp1/mlp2结构
        self.mlp1 = nn.Sequential(
            nn.Linear(in_features=256, out_features=128),
            nn.ReLU(),
            nn.Linear(128, 64)
        )
        self.mlp2 = nn.Sequential(
            nn.Linear(in_features=64, out_features=32),
            nn.ReLU(),
            nn.Linear(32, 16)
        )

        # 仅当提供路径时,加载NestedModel的参数到mlp中
        if nested_model_ckpt_path is not None:
            nested_model = NestedModel(nested_model_ckpt_path, **kwargs)
            self.mlp1.load_state_dict(nested_model.mlp1.state_dict())
            self.mlp2.load_state_dict(nested_model.mlp2.state_dict())
        
        # 其他主模型初始化逻辑...

加载主模型时直接调用:

model = PrimaryModel.load_from_checkpoint(primary_model_ckpt_path)

方案二:添加参数控制NestedModel初始化

给PrimaryModel的__init__添加开关参数,加载主模型时关闭NestedModel的初始化逻辑,直接创建空的mlp结构,由主模型checkpoint填充参数。

修改后的PrimaryModel代码:

class PrimaryModel(LightningModule):
    def __init__(self, nested_model_ckpt_path=None, load_nested_model=True, **kwargs):
        super().__init__()
        if load_nested_model:
            # 训练/从头初始化时,加载NestedModel
            assert nested_model_ckpt_path is not None, "必须传入nested_model_ckpt_path以加载NestedModel"
            nested_model = NestedModel(nested_model_ckpt_path, **kwargs)
            self.mlp1 = nested_model.mlp1
            self.mlp2 = nested_model.mlp2
        else:
            # 加载主模型checkpoint时,直接创建与NestedModel一致的mlp结构
            self.mlp1 = nn.Sequential(...)  # 匹配NestedModel.mlp1结构
            self.mlp2 = nn.Sequential(...)  # 匹配NestedModel.mlp2结构
        
        # 其他主模型初始化逻辑...

加载主模型时传入开关参数:

model = PrimaryModel.load_from_checkpoint(
    primary_model_ckpt_path,
    load_nested_model=False
)

关键注意事项

  • 确保PrimaryModel中定义的mlp1、mlp2结构与NestedModel中的完全一致,否则state_dict参数无法正确匹配加载。
  • 若NestedModel的mlp结构复杂,可将其结构封装为单独的函数,在PrimaryModel和NestedModel中复用,避免重复代码。

内容的提问来源于stack exchange,提问作者malfonsoarquimea

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.20 19:05:42