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

使用PyTorch Lightning+HuggingFace时,如何正确保存模型以支持from_pretrained加载?

解决PyTorch Lightning保存模型后用HuggingFace .from_pretrained()加载的问题

问题根源

PyTorch Lightning的ModelCheckpoint保存的.ckpt文件是整个LightningModule的状态字典,其中HuggingFace模型的权重通常嵌套在子键下(比如你定义LightningModule时的self.model)。直接重命名为pytorch_model.bin会让HuggingFace无法识别正确的权重结构,因此触发层重新初始化的警告。

正确保存方案

方案1:从LightningModule中提取HuggingFace模型原生保存

训练完成后,加载.ckpt文件提取内部的HuggingFace模型实例,用HuggingFace自带的save_pretrained()方法保存:

# 假设你的LightningModule定义示例
class MyLightningModule(pl.LightningModule):
    def __init__(self):
        super().__init__()
        self.model = AutoModelForSequenceClassification.from_pretrained("bert-base-uncased")
        # 其他训练组件...

# 加载训练好的ckpt
ckpt_path = "path/to/your/trained_model.ckpt"
lightning_model = MyLightningModule.load_from_checkpoint(ckpt_path)

# 提取并保存符合HuggingFace格式的模型
lightning_model.model.save_pretrained("path/to/hf_compatible_model")

保存后的目录会自动生成pytorch_model.bin和config.json,直接用AutoModel.from_pretrained("path/to/hf_compatible_model")即可正常加载,不会出现初始化警告。

方案2:自定义回调,训练过程中自动保存

如果需要在训练阶段自动保存符合格式的模型,可以自定义Lightning回调:

from pytorch_lightning.callbacks import Callback

class SaveHFModelCallback(Callback):
    def on_validation_epoch_end(self, trainer, pl_module):
        # 验证结束后保存,可根据需求改成on_train_epoch_end等时机
        save_dir = f"hf_model_epoch_{trainer.current_epoch}"
        pl_module.model.save_pretrained(save_dir)

# 训练时添加该回调
trainer = pl.Trainer(callbacks=[ModelCheckpoint(), SaveHFModelCallback()])

方案3:手动处理已有的.ckpt文件

如果已经有现成的.ckpt文件,无需重新训练,可手动提取权重:

import torch

# 加载ckpt的状态字典
ckpt_dict = torch.load("path/to/your/model.ckpt")
# 提取HuggingFace模型的权重(键前缀需匹配你的LightningModule结构,比如"model.")
hf_state_dict = {k.replace("model.", ""): v for k, v in ckpt_dict["state_dict"].items() if k.startswith("model.")}

# 保存为pytorch_model.bin,同时复制原模型的config.json到同一目录
torch.save(hf_state_dict, "path/to/save/pytorch_model.bin")
# 注意:需确保目标目录存在对应的config.json(可从训练时使用的HuggingFace模型目录复制)

验证加载

保存完成后,用以下代码验证是否正常加载:

from transformers import AutoModelForSequenceClassification

model = AutoModelForSequenceClassification.from_pretrained("path/to/hf_compatible_model")
# 可通过对比某层参数值,确认权重是否正确加载

内容的提问来源于stack exchange,提问作者Chan Wing

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 23:20:00