如何加载AutoModelForCausalLM时跳过权重初始化以节省资源?
无需初始化权重加载预训练Transformer模型的优雅方案
针对你提出的需求,以下是几种标准、优雅的方法,可跳过AutoModelForCausalLM等Transformers类的权重初始化,直接加载预训练权重以节省时间和内存:
方法1:使用Transformers官方参数low_cpu_mem_usage
这是Transformers库原生支持的优化参数,内部会自动跳过模型权重的初始化,直接从预训练文件加载权重,是最简便的方案。
from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained( args.model_path, low_cpu_mem_usage=True )
该参数不仅能跳过初始化,还会优化内存占用,适合大模型加载场景,完全替代你之前的hack方法。
方法2:结合skip_init与模型类实例化
虽然skip_init不能直接作用于AutoModelForCausalLM,但可以通过获取具体模型类的方式间接使用:
from transformers import AutoConfig, AutoModelForCausalLM from torch.nn.utils import skip_init import torch # 获取模型配置 config = AutoConfig.from_pretrained(args.model_path) # 获取对应的具体模型类(比如LlamaForCausalLM) model_class = AutoModelForCausalLM.get(config.model_type) # 用skip_init跳过初始化创建模型实例 model = skip_init(model_class, config=config) # 加载预训练权重 model.load_state_dict(torch.load(f"{args.model_path}/pytorch_model.bin"))
此方法通过skip_init直接跳过模型层的初始化逻辑,再手动加载预训练权重,适合需要自定义模型实例化流程的场景。
方法3:重写模型的_init_weights方法
如果需要更灵活的控制,可以临时重写模型类的权重初始化方法,避免手动替换全局torch初始化函数:
from transformers import AutoModelForCausalLM # 临时保存原初始化方法 original_init_weights = AutoModelForCausalLM._init_weights # 替换为无操作的初始化方法 AutoModelForCausalLM._init_weights = lambda self, module: None # 加载模型,此时不会执行权重初始化 model = AutoModelForCausalLM.from_pretrained(args.model_path) # 恢复原初始化方法 AutoModelForCausalLM._init_weights = original_init_weights
这种方式比你之前替换全局torch初始化函数的hack更精准,只会影响目标模型类的初始化逻辑。
内容的提问来源于stack exchange,提问作者Poe Dator
相关产品推荐
相关产品推荐

