PyTorch无需预初始化权重直接从checkpoint加载Transformer模型方法
加载大体积PyTorch Checkpoint内存不足的解决方案
以下两种方法都可以跳过预训练权重/随机权重初始化步骤,全程仅保留一份权重在内存中,大幅降低内存开销:
方法1:基于Meta设备的空模型加载(PyTorch 1.10+支持)
Meta设备上创建的模型仅保留网络结构,不会初始化实际权重,内存占用几乎为0:
import torch from transformers import AutoModel, AutoConfig # 仅加载模型配置,不加载权重 config = AutoConfig.from_pretrained("xlm-roberta-base") # 在meta设备上创建空模型 with torch.device("meta"): model = AutoModel.from_config(config) # 加载checkpoint,可直接map到目标设备避免中间内存占用 checkpoint = torch.load("xlm-roberta-checkpoint.pth", map_location="cuda") # 加载权重到模型,assign=True直接赋值适配meta空模型 model.load_state_dict(checkpoint["model_state_dict"], assign=True) # 开启推理模式 model.eval() model.requires_grad_(False)
注意:assign=True参数需要PyTorch 1.13及以上版本支持,低版本PyTorch建议使用方法2
方法2:Hugging Face内置低内存加载(Transformers 4.20+支持)
Hugging Face官方提供了low_cpu_mem_usage参数,底层封装了Meta设备逻辑,使用更简单:
import torch from transformers import AutoModel, AutoConfig config = AutoConfig.from_pretrained("xlm-roberta-base") # 直接读取checkpoint中的权重 state_dict = torch.load("xlm-roberta-checkpoint.pth")["model_state_dict"] # 低内存模式加载模型,全程仅存一份权重 model = AutoModel.from_config( config, state_dict=state_dict, low_cpu_mem_usage=True ) # 移到推理设备并开启推理模式 model = model.to("cuda").eval().requires_grad_(False)
额外优化建议
- 加载checkpoint时直接通过
map_location参数指定目标设备,避免CPU和GPU之间重复拷贝产生的额外内存占用 - 超大体积checkpoint可以开启
mmap_mode参数:torch.load(PATH, mmap_mode="r"),直接从磁盘映射读取,不需要全量加载到内存 - 如果模型仍超过显存限制,可搭配量化、模型并行等方式进一步压缩体积
内容的提问来源于stack exchange,提问作者AetherPrior
相关产品推荐
相关产品推荐

