使用QLoRA微调LLaMA 2-7B后重载失败:缺失4bit元数据
解决bitsandbytes 4bit量化基座模型保存与加载报错问题
核心问题根源
直接保存应用LoRA或修改词嵌入后的完整模型,会破坏bitsandbytes的4bit量化结构:Linear4bit层的.weight会从bnb.nn.Params4bit被转换为普通torch.nn.Parameter,同时丢失compress_statistics等量化元数据,导致后续加载报错。
正确保存与加载流程
1. 训练前先保存原始4bit量化基座
在添加LoRA或修改词嵌入之前,单独保存未被修改的4bit量化基座,确保量化结构完整:
from transformers import AutoModelForCausalLM, BitsAndBytesConfig import torch # 加载4bit量化基座 bnb_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_use_double_quant=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.bfloat16 ) base_model = AutoModelForCausalLM.from_pretrained( "meta-llama/Llama-2-7b-hf", load_in_4bit=True, device_map="auto", quantization_config=bnb_config ) # 保存原始4bit基座(此时所有Linear4bit层权重均为bnb.nn.Params4bit类型) base_model.save_pretrained("./llama2-7b-4bit-base") custom_tokenizer.save_pretrained("./llama2-7b-4bit-base")
2. 训练后仅保存LoRA适配器
不要保存完整模型,用PEFT库的save_pretrained仅保存LoRA权重:
from peft import LoraConfig, get_peft_model # 配置并包装LoRA模型 lora_config = LoraConfig( r=8, lora_alpha=32, target_modules=["q_proj", "v_proj"], # 根据LLaMA 2结构调整 lora_dropout=0.05, bias="none", task_type="CAUSAL_LM" ) peft_model = get_peft_model(base_model, lora_config) # 执行自适应预训练... # 仅保存LoRA适配器文件 peft_model.save_pretrained("./llama2-7b-arabic-lora")
3. 处理自定义分词器的词嵌入扩展
如果训练时因自定义分词器(63k词汇量)扩展了词嵌入,需单独保存扩展后的embedding层,避免破坏基座的4bit结构:
# 训练时扩展词嵌入后,单独保存embedding权重 torch.save(peft_model.get_input_embeddings().weight, "./extended_embeddings.pt")
4. 加载时的正确步骤
from transformers import AutoModelForCausalLM, BitsAndBytesConfig from peft import PeftModel import torch # 1. 重新加载原始4bit基座 bnb_config = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_use_double_quant=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.bfloat16 ) base_model = AutoModelForCausalLM.from_pretrained( "./llama2-7b-4bit-base", load_in_4bit=True, device_map="auto", quantization_config=bnb_config ) # 2. 替换为扩展后的词嵌入 custom_tokenizer = AutoTokenizer.from_pretrained("./llama2-7b-4bit-base") base_model.resize_token_embeddings(len(custom_tokenizer)) extended_emb = torch.load("./extended_embeddings.pt") base_model.get_input_embeddings().weight.data = extended_emb # 3. 加载LoRA适配器 peft_model = PeftModel.from_pretrained(base_model, "./llama2-7b-arabic-lora") # 验证:检查Linear4bit层权重类型 for name, module in peft_model.named_modules(): if isinstance(module, torch.nn.Linear) and hasattr(module, "weight"): if "4bit" in str(type(module)): print(f"{name} weight type: {type(module.weight)}") # 预期输出:<class 'bitsandbytes.nn.Params4bit'>
内容的提问来源于stack exchange,提问作者orchid Ali
相关产品推荐
相关产品推荐

