LangChain加载LLaMA类模型显存占用过高,如何优化?
问题描述
尝试从本地加载基于LLaMA的量化模型(以wizardLM-7B-GPTQ-4bit-128g为例,其他同类模型也存在相同问题),使用LangChain结合HuggingFace Pipeline时,8GB显存被完全占满。该模型可在KoboldAI、text-generation-webui等框架正常加载,说明存在可行方案。需实现从HuggingFace下载的wizardLM-7B-GPTQ-4bit-128g模型通过LangChain在Python中运行,GPU为AMD® Radeon RX 6600(8GB显存,ROCM 5.4.2 + PyTorch)。
原代码如下:
from langchain.llms import HuggingFacePipeline from langchain import PromptTemplate, LLMChain import torch from transformers import LlamaTokenizer, LlamaForCausalLM, LlamaConfig, pipeline torch.cuda.set_device(torch.device("cuda:0")) PATH = './models/wizardLM-7B-GPTQ-4bit-128g' config = LlamaConfig.from_json_file(f'{PATH}/config.json') base_model = LlamaForCausalLM(config=config).half() torch.cuda.empty_cache() tokenizer = LlamaTokenizer.from_pretrained( pretrained_model_name_or_path=PATH, low_cpu_mem_usage=True, local_files_only=True ) torch.cuda.empty_cache() pipe = pipeline( "text-generation", model=base_model, tokenizer=tokenizer, batch_size=1, device=0, max_length=100, temperature=0.6, top_p=0.95, repetition_penalty=1.2 )
解决方案
核心问题是原代码未正确加载GPTQ量化模型,而是初始化了一个空的半精度模型,导致显存占用异常。需使用专门的GPTQ加载工具,结合显存优化选项实现模型加载。
关键优化点
- 安装适配ROCM环境的
auto-gptq库 - 使用
AutoGPTQForCausalLM加载GPTQ量化模型,启用device_map='auto'自动分配显存 - 开启
low_cpu_mem_usage=True减少CPU内存占用,避免不必要的显存预分配 - 加载模型时自动读取内置量化配置,确保正确识别4bit权重
修改后的代码
from langchain.llms import HuggingFacePipeline from langchain import PromptTemplate, LLMChain import torch from transformers import AutoTokenizer, pipeline from auto_gptq import AutoGPTQForCausalLM # 设定设备 device = "cuda:0" if torch.cuda.is_available() else "cpu" PATH = './models/wizardLM-7B-GPTQ-4bit-128g' # 加载GPTQ量化模型 model = AutoGPTQForCausalLM.from_quantized( PATH, device=device, use_triton=False, # ROCm环境下禁用triton use_safetensors=True, trust_remote_code=True, low_cpu_mem_usage=True, quantize_config=None # 自动读取模型内的量化配置 ) # 加载tokenizer tokenizer = AutoTokenizer.from_pretrained( PATH, low_cpu_mem_usage=True, local_files_only=True, trust_remote_code=True ) # 创建text-generation pipeline,启用显存优化 pipe = pipeline( "text-generation", model=model, tokenizer=tokenizer, batch_size=1, device=device, max_length=100, temperature=0.6, top_p=0.95, repetition_penalty=1.2, torch_dtype=torch.float16, low_cpu_mem_usage=True ) # 测试LangChain集成 prompt = PromptTemplate( input_variables=["question"], template="Q: {question}\nA:" ) llm_chain = LLMChain(prompt=prompt, llm=HuggingFacePipeline(pipeline=pipe)) response = llm_chain.run("什么是大语言模型?") print(response)
额外优化建议
- 若GPTQ加载出现问题,可改用
bitsandbytes的8bit量化:加载模型时添加load_in_8bit=True参数,但GPTQ 4bit显存占用更低 - 加载模型前执行
torch.cuda.empty_cache()清理无用显存 - 确保
transformers版本在4.30.0以上,auto-gptq版本适配当前PyTorch和ROCM环境
内容的提问来源于stack exchange,提问作者someone
相关产品推荐
相关产品推荐

