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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 23:09:03