博客
关于我
pytorch从预训练权重加载完全相同的层
阅读量:798 次
发布时间:2023-03-04

本文共 749 字,大约阅读时间需要 2 分钟。

加载预训练模型并初始化新模型权重

在PyTorch中加载预训练模型并初始化新模型的过程中,我们可以参考以下步骤:

加载预训练模型

首先,从指定路径加载预训练模型的状态字典:

saved = torch.load(cfg.Transfer, map_location=device)
old_state_dict = saved['state_dict']

初始化新模型

创建一个空的字典用于存储新模型的权重:

new_state_dict = {}

处理模型层

遍历当前模型的状态字典,并将旧模型中与新模型层形状匹配的权重复制到新模型中:

all_layer = len(model.state_dict())
num = 0
for key in model.state_dict():
if key in old_state_dict and old_state_dict[key].shape == model.state_dict()[key].shape:
new_state_dict[key] = old_state_dict[key]
num += 1
else:
new_state_dict[key] = model.state_dict()[key]

加载新模型部分权重

使用load_state_dict方法加载新模型的部分权重:

model.load_state_dict(new_state_dict)

输出结果

打印加载的层数信息:

print("从预训练模型中加载了{}/{}层".format(num, all_layer))

这个方法可以帮助开发者在保持模型灵活性的同时,充分利用预训练模型的优势。

转载地址:http://clxfk.baihongyu.com/

你可能感兴趣的文章
Postgres Docker版本安装mysql_fdw 插件
查看>>
Postgres invalid command \N数据恢复处理
查看>>
Postgres like 模糊查询匹配集合
查看>>
Postgres 自定义函数内实现 in 操作符的递归查询
查看>>
Postgres 返回当前时间前后指定天数的集合
查看>>
postgres--vacuum
查看>>
postgres--wal
查看>>
postgres--流复制
查看>>
postgres10配置huge_pages
查看>>
PostgreSQL 10.0 preview sharding增强 - pushdown 增强
查看>>
PostgreSQL 10.0 preview 变化 - pg_xlog,pg_clog,pg_log目录更名为pg_wal,pg_xact,log
查看>>
PostgreSQL 10.1 手册_部分 II. SQL 语言_第 15章 并行查询_15.2. 何时会用到并行查询?...
查看>>
PostgreSQL 10.1 手册_部分 II. SQL 语言_第 9 章 函数和操作符_9.23. 行和数组比较
查看>>
PostgreSQL 10.1 手册_部分 III. 服务器管理_第 21 章 数据库角色
查看>>
Qt开发——网络编程UDP网络广播软件之服务器端
查看>>
Postgresql 12.9如何配置允许远程连接
查看>>
PostgreSQL 9.6 同步多副本 与 remote_apply事务同步级别 应用场景分析
查看>>
Postgresql CopyManager 流式批量数据入库
查看>>
PostgreSQL cube 插件 - 多维空间对象
查看>>
PostgreSQL Daily Maintenance - cluster table
查看>>