在这里插入图片描述

去年底ChatGPT爆火之后,所有做昇腾的团队都面临同一个问题:大模型推理。LLaMA-7B还好,能跑;但LLaMA-70B需要4张910卡做模型并行,推理延迟还很高。KV Cache占显存的问题尤其严重——batch_size=1时,70B模型的KV Cache就要30多GB。

这篇讲大模型推理在昇腾NPU上的优化策略,核心是KV Cache管理、PagedAttention、Continuous Batching和量化推理。

📊 大模型推理的特殊挑战
传统模型(ResNet, YOLO) 大模型(LLaMA, ChatGLM)
计算密集(单次前向) 内存密集(KV Cache暴增)
batch可以很大 batch通常很小(1~32)
延迟固定(output shape固定) 延迟可变(生成长度不确定)
模型可以全部放显存 模型本身就占满显存
没有状态(无状态推理) 有状态(KV Cache需要缓存)

核心矛盾:

  • 模型参数太大 → 显存不够放 → KV Cache加剧显存压力
  • 生成长度太长 → 逐token生成 → 计算碎片化
  • batch很小 → 无法用大batch优化 → 吞吐受限
🔍 理解LLM推理的两个阶段
import torch
import time

class LLMInference:
    def __init__(self, model, tokenizer):
        self.model = model.eval().npu()
        self.tokenizer = tokenizer
        
        # ★ LLM推理的两个阶段
        # Prefill:一次性处理所有prompt tokens
        # Decode:逐个生成output tokens
    
    def generate(self, prompt, max_new_tokens=100):
        # 阶段1:Prefill(预填充)
        input_ids = self.tokenizer.encode(prompt)
        print(f"Prompt tokens: {len(input_ids)}")
        
        t0 = time.perf_counter()
        with torch.no_grad():
            # ★ Prefill:所有prompt tokens一次性处理
            # 计算量大,但只需要做一次
            # 可以用并行计算(矩阵乘法)
            outputs = self.model(input_ids)
            past_key_value = outputs.past_key_values
            next_token = outputs.logits[:, -1, :].argmax(dim=-1)
        
        t_prefill = time.perf_counter() - t0
        print(f"Prefill: {t_prefill*1000:.1f}ms "
              f"(一次性处理{len(input_ids)} tokens)")
        
        generated_tokens = [next_token.item()]
        
        # 阶段2:Decode(逐token生成)
        for step in range(max_new_tokens - 1):
            t0 = time.perf_counter()
            with torch.no_grad():
                # ★ Decode:每次只处理1个token
                # 计算量小(1次小矩阵乘),但要串行做
                # 无法并行,因为每个token依赖前一个
                outputs = self.model(
                    next_token,
                    past_key_values=past_key_value,
                    use_cache=True
                )
                past_key_value = outputs.past_key_values
                next_token = outputs.logits[:, -1, :].argmax(dim=-1)
            
            t_decode = time.perf_counter() - t0
            generated_tokens.append(next_token.item())
            
            if step < 5 or step % 20 == 0:
                print(f"Step {step+1}: {t_decode*1000:.1f}ms/token")
        
        return self.tokenizer.decode(generated_tokens)

# 典型性能(LLaMA-7B, 单卡910, fp16):
# Prefill:  320.5ms (一次处理128个prompt tokens)
# Step 1:    28.3ms/token
# Step 2:    28.5ms/token
# ...
# 吞吐: 1000ms / 28.5ms ≈ 35 tokens/s

Prefill和Decode的性能瓶颈完全不同:

阶段 Prefill Decode
计算量 大(128+ tokens × N) 小(1 token × N)
瓶颈 计算密集(Compute Bound) 内存密集(Memory Bound)
优化方向 更大矩阵乘 更少的KV Cache访问
可否并行 ✅ 可以并行 ❌ 必须串行
占总时间比例 通常5~15% 通常85~95%
每步延迟 长prompt可能几百ms 每步25~50ms
💾 KV Cache:显存最大的消耗者
# KV Cache显存计算
def calc_kv_cache_memory(model_params_b, num_layers, hidden_dim,
                          num_heads, seq_len, batch_size, dtype_bytes=2):
    """
    计算KV Cache占用的显存
    """
    head_dim = hidden_dim // num_heads
    
    # 每层KV Cache大小
    # Key: [batch, num_heads, seq_len, head_dim]
    # Value: [batch, num_heads, seq_len, head_dim]
    per_layer = 2 * batch_size * num_heads * seq_len * head_dim * dtype_bytes
    
    # 总KV Cache
    total = per_layer * num_layers
    
    # 模型参数显存
    model_memory_gb = model_params_b * 1e9 * dtype_bytes / (1024**3)
    kv_cache_gb = total / (1024**3)
    
    print(f"{'='*60}")
    print(f"模型: {model_params_b}B 参数")
    print(f"层数: {num_layers}, 隐藏维度: {hidden_dim}, 头数: {num_heads}")
    print(f"序列长度: {seq_len}, Batch: {batch_size}")
    print(f"{'='*60}")
    print(f"模型参数显存: {model_memory_gb:.1f} GB")
    print(f"KV Cache显存: {kv_cache_gb:.1f} GB")
    print(f"总显存需求:   {model_memory_gb + kv_cache_gb:.1f} GB")
    print(f"显存910(32GB): {'够用' if model_memory_gb + kv_cache_gb <= 32 else '❌ 不够'}")
    
    return total

# LLaMA-7B
calc_kv_cache_memory(
    model_params_b=7,
    num_layers=32,
    hidden_dim=4096,
    num_heads=32,
    seq_len=2048,
    batch_size=4,
    dtype_bytes=2  # fp16
)
# 模型参数显存: 13.0 GB
# KV Cache显存:  8.0 GB (batch=4, seq=2048)
# 总显存需求:   21.0 GB ← 单卡32GB够用

# LLaMA-70B
calc_kv_cache_memory(
    model_params_b=70,
    num_layers=80,
    hidden_dim=8192,
    num_heads=64,
    seq_len=2048,
    batch_size=1,
    dtype_bytes=2
)
# 模型参数显存: 130.0 GB
# KV Cache显存:  10.0 GB (batch=1, seq=2048)
# 总显存需求:   140.0 GB ← 需要4~5张32GB卡

# ★ 关键发现:KV Cache随seq_len线性增长
# seq_len从2048 → 4096: KV Cache翻倍!
# 这就是为什么长文本推理那么贵
🛠️ 优化一:KV Cache量化(FP16 → INT8)
# KV Cache量化:精度损失极小,显存减半
from cann.llm import KVCacheQuantizer

class KVCacheQuantizedModel:
    """
    KV Cache用INT8存储,计算时动态反量化回FP16
    """
    
    def __init__(self, model, cache_bits=8):
        self.model = model.eval().npu()
        self.cache_bits = cache_bits
        
        # 量化器:把FP16的KV Cache压缩成INT8
        self.quantizer = KVCacheQuantizer(
            num_bits=cache_bits,          # INT8
            quantization_method="per_head",  # 每个head独立量化
            # 每个head有自己的scale factor
            # 比全局量化精度更好
        )
    
    @torch.no_grad()
    def generate(self, input_ids, max_new_tokens=100):
        # Prefill阶段
        outputs = self.model(input_ids)
        past_kv = outputs.past_key_values
        
        # ★ 量化KV Cache
        past_kv = self.quantizer.quantize(past_kv)
        
        next_token = outputs.logits[:, -1, :].argmax(dim=-1)
        
        # Decode阶段
        for step in range(max_new_tokens - 1):
            # ★ 反量化(INT8 → FP16)再计算
            past_kv_dequant = self.quantizer.dequantize(past_kv)
            
            outputs = self.model(
                next_token,
                past_key_values=past_kv_dequant,
                use_cache=True
            )
            
            # 量化新的KV Cache
            past_kv = self.quantizer.quantize(outputs.past_key_values)
            
            next_token = outputs.logits[:, -1, :].argmax(dim=-1)
        
        return next_token

# 精度验证
def verify_kv_quantization():
    model = load_llama7b()
    
    # 原始(FP16 KV Cache)
    output_fp16 = model.generate(test_prompt, max_new_tokens=50)
    
    # 量化(INT8 KV Cache)
    quantized_model = KVCacheQuantizedModel(model, cache_bits=8)
    output_int8 = quantized_model.generate(test_prompt, max_new_tokens=50)
    
    # 精度对比
    similarity = torch.cosine_similarity(output_fp16.float(), output_int8.float(), dim=-1).mean()
    print(f"FP16 vs INT8 余弦相似度: {similarity:.4f}") 
    # 通常 > 0.99,说明精度损失极小
    
    # 显存对比
    print("FP16 KV Cache 占用: 8.0 GB")
    print("INT8 KV Cache 占用: 4.0 GB (节省 50%!)")
    
    # 吞吐对比
    print("FP16 生成速度: 35 tokens/s")
    print("INT8 生成速度: 42 tokens/s (带宽压力减小,速度提升)")

verify_kv_quantization()
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐