从PyTorch的.pt模型导出state_dict checkpoint的方法问询
从RadTTS预训练.pt文件提取state_dict的方法
1. 加载.pt文件
用torch.load()加载预训练模型文件,指定map_location="cpu"可避免GPU环境依赖问题:
import torch # 替换为你的.pt文件实际路径 loaded_data = torch.load("radtts_pretrained.pt", map_location="cpu")
2. 确认文件结构
.pt文件内容通常有两种形式:要么是直接的模型state_dict,要么是包含state_dict的复合字典(比如训练时保存的checkpoint会附带优化器参数等)。先打印结构确认:
if isinstance(loaded_data, dict): print(loaded_data.keys())
3. 提取目标state_dict
根据打印结果选择对应操作:
- 如果
loaded_data本身就是state_dict(所有值均为torch.Tensor类型),直接使用即可 - 如果字典里存在
state_dict或model这类键,提取对应值:
# 示例:假设键名为"state_dict",根据实际打印结果调整键名 desired_state_dict = loaded_data.get("state_dict", loaded_data)
4. 处理前缀问题(可选)
分布式训练保存的模型,state_dict的键可能带有module.前缀,若目标库要求无此前缀,可批量移除:
desired_state_dict = {k.replace("module.", ""): v for k, v in desired_state_dict.items()}
5. 保存提取后的state_dict(可选)
如果需要单独保存提取好的state_dict:
torch.save(desired_state_dict, "extracted_radtts_state_dict.pt")
内容的提问来源于stack exchange,提问作者NicoCaldo
相关产品推荐
相关产品推荐

