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

从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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 00:43:16