TensorFlow2.0扩充词表后如何正确加载嵌入层权重
问题场景
- 权重保存调用接口:
model.save_weights('weights_file') - 权重加载调用接口:
model.load_weights('weights_file') - 触发报错场景:接入新数据开展训练时,嵌入层对应词表大小提升,加载权重后调用
model.fit()抛出维度不兼容错误,核心报错信息:
ValueError: Shapes (31, 5) and (15, 5) are incompatible
根因确认:新模型嵌入层初始化形状为
(31, 5),加载的历史权重对应嵌入层形状为(15, 5),二者形状不匹配,同时优化器存储的历史动量等槽位变量也和新层维度不兼容,最终触发报错。
可行解决方案
- 跳过不匹配权重加载
调用权重加载接口时传入by_name=True和skip_mismatch=True参数,自动跳过形状不匹配的层权重及对应优化器槽位变量,嵌入层新增维度的参数会使用默认初始化值填充,调用方式如下:model.load_weights('weights_file', by_name=True, skip_mismatch=True)
该方式会清空优化器的历史训练状态,加载完成后建议先用较小学习率训练2-3个epoch做warm up,再恢复原学习率,避免训练震荡。 - 手动迁移嵌入层旧权重
该方式可以完整保留旧词表对应的已训练词向量,不会丢失之前的训练效果,适合词表增量更新场景,操作逻辑如下:- 读取旧权重文件中存储的嵌入层权重矩阵
- 初始化新模型的嵌入层权重矩阵,将旧权重复制到新矩阵中对应旧词表的位置,新增词表对应的行用随机初始化或预训练词向量填充
- 将拼接完成的权重矩阵赋值给新模型的嵌入层,其余形状匹配的层正常加载权重即可
参考代码:
# 读取旧嵌入层权重,替换成实际的嵌入层名称 old_emb_weight = old_model.get_layer("embedding").get_weights()[0] # 获取新模型嵌入层的初始化权重 new_emb_layer = new_model.get_layer("embedding") new_emb_weight = new_emb_layer.get_weights()[0] # 迁移旧权重到新矩阵 new_emb_weight[:old_emb_weight.shape[0], :] = old_emb_weight new_emb_layer.set_weights([new_emb_weight]) # 加载其余匹配层的权重 new_model.load_weights("weights_file", by_name=True, skip_mismatch=True) - 重新初始化优化器
如果不需要保留优化器的历史训练状态,加载完可匹配的层权重后,重新实例化优化器,再次调用model.compile()配置损失函数、评估指标即可,编译后优化器会根据新的模型参数形状重新创建槽位变量,不会再触发形状不匹配报错。
内容的提问来源于stack exchange,提问作者GBPU
相关产品推荐
相关产品推荐

