PyTorch中是否有对应TensorFlow Keras的get_weights与set_weights方法
PyTorch中等效Keras
get_weights()/set_weights()的实现方案 PyTorch没有和TF/Keras的get_weights()、set_weights()完全同名的内置方法,但可以通过原生的state_dict机制轻松实现等效功能,灵活度还更高。
1. 获取模型权重(等效get_weights())
所有继承nn.Module的模型/层都自带state_dict()方法,会返回有序字典,包含所有参数(含可训练权重、BN层的running_mean/running_var等状态)的名称和对应张量:
import torch import torch.nn as nn # 示例模型 class DemoModel(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(10, 5) self.fc2 = nn.Linear(5, 2) model = DemoModel() # 获取全量权重状态 full_weights = model.state_dict() # 仅获取可训练权重(对应Keras get_weights()的默认返回逻辑) trainable_weights = [p.detach() for p in model.parameters()]
2. 设置模型权重(等效set_weights())
调用load_state_dict()方法即可批量加载匹配的权重字典:
# new_weight_dict为和模型结构完全匹配的权重字典 model.load_state_dict(new_weight_dict) # 微调等存在部分key不匹配的场景,可加strict=False忽略不匹配的键 model.load_state_dict(new_weight_dict, strict=False)
3. 自定义封装对齐Keras调用习惯
如果要完全复刻Keras的调用方式,可以给nn.Module扩展两个方法:
def get_weights(self): # 返回numpy格式的权重列表,和Keras输出格式完全一致 return [param.detach().cpu().numpy() for param in self.parameters()] def set_weights(self, weights): # 接收numpy格式的权重列表批量赋值 with torch.no_grad(): for param, w in zip(self.parameters(), weights): param.copy_(torch.from_numpy(w).to(param.device)) # 全局绑定到所有nn.Module子类 nn.Module.get_weights = get_weights nn.Module.set_weights = set_weights # 后续直接和Keras一样调用即可 model.get_weights() model.set_weights(your_weight_list)
注意事项
- 权重赋值时用
torch.no_grad()包裹是为了避免赋值操作被计入计算图,不会影响正常的反向传播逻辑 - 跨CPU/GPU加载权重时,上述
set_weights的封装已经自动适配目标设备,无需额外手动搬移张量
内容的提问来源于stack exchange,提问作者GKozinski
相关产品推荐
相关产品推荐

