你应该懂的AI大模型(十三) 之 推理框架
在前面的系列文章中,我们已经讨论了模型架构、训练技巧、微调方法等。但真正让大模型从“实验室玩具”变成“生产工具”的关键一环,是推理框架。想象一下,你训练了一个千亿参数的模型,但用户点击“发送”后,却要等30秒才能看到第一个字——这体验是灾难性的。推理框架就是专门解决这类问题的“性能引擎”。### 什么是推理框架?简单来说,推理框架是部署大模型并执行生成请求的软件系统。它负责加载模型权重、管理GPU显存、优化计算图、调度请求,最终以最低延迟和最高吞吐量输出结果。它不负责训练,只负责“跑起来”。为什么不能直接用PyTorch呢?因为PyTorch的推理模式存在几个痛点:-显存浪费:每个请求都要复制一份模型权重。-计算低效:没有对KV Cache、Attention进行专门优化。-无批处理:不同请求无法共享计算。推理框架的核心目标就是解决这三个问题。我们按从易到难的顺序,逐个击破。—## 基础:从零实现一个简单的批处理推理首先,我们理解一下最朴素的方式——单请求推理。此时模型一次只处理一个用户输入。代码如下:pythonimport torchfrom transformers import AutoModelForCausalLM, AutoTokenizer# 加载一个小模型(如GPT-2)model_name = "gpt2"tokenizer = AutoTokenizer.from_pretrained(model_name)model = AutoModelForCausalLM.from_pretrained(model_name).to("cuda")# 单请求生成def generate_single(prompt: str, max_new_tokens: int = 50): inputs = tokenizer(prompt, return_tensors="pt").to("cuda") outputs = model.generate(**inputs, max_new_tokens=max_new_tokens) return tokenizer.decode(outputs[0], skip_special_tokens=True)print(generate_single("The capital of France is"))问题:如果同时有10个用户请求,这个代码只能串行处理,GPU利用率极低。我们改进一下,引入动态批处理(Dynamic Batching)——把多个请求拼成一个批次,共享一次前向计算。pythondef generate_batch(prompts: list[str], max_new_tokens: int = 50): # 将多个prompt编码成同一batch(注意padding) inputs = tokenizer(prompts, return_tensors="pt", padding=True, truncation=True).to("cuda") outputs = model.generate(**inputs, max_new_tokens=max_new_tokens) return [tokenizer.decode(out, skip_special_tokens=True) for out in outputs]# 模拟4个并发请求prompts = [ "What is the weather in Beijing?", "Explain quantum computing briefly.", "Write a haiku about autumn.", "List three benefits of exercise."]results = generate_batch(prompts)for r in results: print(r[:50] + "...")这一步已经比单请求快了不少,因为GPU并行处理多个序列。但这只是“表面功夫”——真正的推理框架会做更多细致优化。—## 进阶:KV Cache 与显存优化大模型生成时,每生成一个新token,都需要重新计算前面所有token的Key和Value向量(用于Attention)。如果不缓存,复杂度是O(n²)。KV Cache就是将这些中间结果缓存下来,把复杂度降为O(n)。下面我们手写一个简化版的KV Cache实现:pythondef inference_with_kv_cache(model, tokenizer, prompt, max_new_tokens=30): inputs = tokenizer(prompt, return_tensors="pt").to("cuda") input_ids = inputs["input_ids"] past_key_values = None # 初始无缓存 generated = input_ids.tolist()[0] for _ in range(max_new_tokens): # 只输入最后一个token,配合past_key_values out = model(input_ids=input_ids, past_key_values=past_key_values, use_cache=True) logits = out.logits[:, -1, :] # 取最后一个位置的logits next_token = torch.argmax(logits, dim=-1).unsqueeze(0) generated.append(next_token.item()) # 更新past_key_values为模型返回的缓存 past_key_values = out.past_key_values input_ids = next_token # 下一轮只输入新token return tokenizer.decode(generated, skip_special_tokens=True)显存优化方面,推理框架常用的技术包括:-模型量化:将FP16降到INT8/INT4,把显存占用缩至1/4甚至1/8。-连续批处理:不再等待整个batch完成,而是动态插入新请求,移除已完成的请求。-PagedAttention(如vLLM):像操作系统分页一样管理KV Cache,减少碎片浪费。—## 高级:主流推理框架实战对比现在,我们介绍两个最主流的开源推理框架:vLLM和TensorRT-LLM。### vLLM:易用性之王vLLM基于PagedAttention,在吞吐量上比HuggingFace Transformers快数倍,且接口极简。python# 安装:pip install vllmfrom vllm import LLM, SamplingParams# 加载模型llm = LLM(model="meta-llama/Llama-2-7b-chat-hf", dtype="float16")# 设置采样参数sampling_params = SamplingParams(temperature=0.8, max_tokens=100)# 批量推理prompts = [ "What is the meaning of life?", "Give me a recipe for chocolate cake.",]outputs = llm.generate(prompts, sampling_params)for output in outputs: print(output.outputs[0].text)vLLM的LLM类自动处理了批处理、缓存、量化等细节。你只需关注业务逻辑。它还支持OpenAI兼容的API服务,可以直接替代/v1/completions接口。### TensorRT-LLM:极致性能NVIDIA的TensorRT-LLM则更底层,它把模型编译成TensorRT引擎,在GPU上运行速度极快,但配置复杂。它常用于需要超低延迟的生产环境(如金融交易、实时对话)。python# 安装:pip install tensorrt_llm# 通常需要先导出ONNX模型,再用trtllm-build命令编译# 伪代码示例:import tensorrt_llm as tllmfrom tensorrt_llm.runtime import ModelRunner# 假设已经编译好了engine文件runner = ModelRunner.from_dir( engine_dir="/path/to/trt_engine", lora_dir=None, is_enc_dec=False, tensor_parallel_size=1, use_gpt_attention_plugin=True,)# 推理output_ids = runner.generate(["Tell me a joke."], max_new_tokens=50)print(output_ids)对比总结:| 框架 | 易用性 | 性能 | 适用场景 ||------|--------|------|----------|| vLLM | ⭐⭐⭐⭐⭐ | ⭐⭐⭐⭐ | 快速部署、高吞吐API || TensorRT-LLM | ⭐⭐ | ⭐⭐⭐⭐⭐ | 极致延迟、定制生产 |—## 实战:部署一个高并发聊天服务结合vLLM,我们可以快速搭建一个能承受高并发的推理服务。下面是一个微型示例:python# 使用FastAPI + vLLMfrom fastapi import FastAPI, Requestfrom vllm import LLM, SamplingParamsimport uvicornapp = FastAPI()llm = LLM(model="mistralai/Mistral-7B-Instruct-v0.2", dtype="float16")@app.post("/chat")async def chat(request: Request): data = await request.json() prompt = data["prompt"] params = SamplingParams(max_tokens=200, temperature=0.7) result = llm.generate([prompt], params) return {"reply": result[0].outputs[0].text}if __name__ == "__main__": uvicorn.run(app, host="0.0.0.0", port=8000)运行这个服务后,你可以用curl测试:bashcurl -X POST http://localhost:8000/chat -H "Content-Type: application/json" -d '{"prompt":"Explain AI in one sentence."}'—## 总结推理框架是大模型落地的“最后一公里”。我们从最基础的批处理讲起,理解了KV Cache和显存优化的必要性,再对比了vLLM和TensorRT-LLM两大主流方案,最后实现了一个可用的推理服务。核心要点:1.动态批处理是提升吞吐量的第一级台阶。2.KV Cache和量化是降低延迟和显存的关键。3. 选择框架时,vLLM适合快速开发,TensorRT-LLM适合极致性能。4. 生产环境还需考虑并发调度、容错、模型热加载等。推理框架的技术仍在快速演进,例如投机采样、并行解码等新方法不断涌现。掌握这些原理,你就能根据业务需求灵活调优,让大模型真正跑得又快又稳。希望这篇文章能帮你建立对推理框架的系统认知。下一期,我们将深入探讨模型量化技术,敬请期待!