使用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
相关产品推荐
相关产品推荐

