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
相关产品推荐
相关产品推荐

