omni / docs /interview /inference.md
chenbhao's picture
refactor(docs): restructure docs with modular organization
f664f3f
|
Raw
History Blame Contribute Delete
12.4 kB

面试:推理优化深度

本仓库 src/models/lm/model.py 的推理实现,覆盖 KV Cache、采样策略、流式生成

0. 推理流程概览

输入 prompt
    │
    ▼
Token Embedding
    │
    ▼
┌──────────────────────────────┐
│  Prefill 阶段                │  处理整个 prompt
│  (一次性计算所有 token)       │
└──────────────┬───────────────┘
       │
       ▼
┌──────────────────────────────┐
│  Decode 阶段                 │  逐 token 生成
│  (KV Cache 加速)             │
└──────────────┬───────────────┘
       │
       ▼
    输出序列

Q1. KV Cache 是什么?为什么快?

核心思想

缓存已算的 K/V,每步只算新 token 的 Q 与已有 K/V 做注意力,避免重算。

本仓库实现(src/core/attention.py:39-42

def forward(self, x, start_pos, freqs_cos, freqs_sin, mask=None):
    # 计算 Q/K/V
    xq = self.q_norm(self.q_proj(x))
    xk = self.k_norm(self.k_proj(x))
    xv = self.v_proj(x)
    
    # 应用 RoPE
    xq, xk = apply_rotary_pos_emb(xq, xk, freqs_cos, freqs_sin)
    
    # KV Cache 简单拼接实现
    if past_key_value is not None:
        xk = torch.cat([past_key_value[0], xk], dim=2)
        xv = torch.cat([past_key_value[1], xv], dim=2)
    past_key_value = (xk, xv)
    
    # 计算注意力
    attn_output = F.scaled_dot_product_attention(xq, xk, xv, attn_mask=mask)
    return attn_output, past_key_value

显存节省

假设:

  • batch_size=B, seq_len=S, num_layers=L
  • num_kv_heads=H, head_dim=D
  • 精度=fp16(2 bytes)

无 KV Cache

  • 每步计算量 = B × S² × L × H × D
  • 显存 = B × S × L × H × D × 2 bytes

有 KV Cache

  • 每步计算量 = B × S × L × H × D
  • 显存 = B × S × L × H × D × 2 bytes(但只需算一次)

面试点:KV Cache 为什么能加速?→ 避免重复计算 K/V,每步只需算新 token 的 Q


Q2. Prefill vs Decode 阶段

Prefill 阶段

  • 处理整个 prompt
  • 一次性计算所有 token 的 K/V
  • 计算量大,但只需做一次

Decode 阶段

  • 逐 token 生成
  • 每步只算新 token 的 Q
  • 计算量小,但需要很多步

本仓库实现(src/models/lm/model.py:73-113

def generate(self, input_ids, max_new_tokens=200, temperature=0.6, top_k=5, top_p=0.8):
    for _ in range(max_new_tokens):
        # Prefill 阶段:处理整个 prompt
        if idx_cond.shape[1] > 1:
            logits, _ = self(idx_cond)
        # Decode 阶段:只处理最后一个 token
        else:
            logits, _ = self(idx_cond, start_pos=start_pos)
        
        # 采样
        logits = logits[:, -1, :] / temperature
        idx_next = self.sample(logits, top_k=top_k, top_p=top_p)
        
        # 更新序列
        idx_cond = torch.cat([idx_cond, idx_next], dim=1)
        start_pos += 1

面试点:为什么 Prefill 和 Decode 要分开处理?→ Prefill 可以并行处理所有 token,Decode 只能逐 token 处理


Q3. Logits to Keep 优化(src/models/lm/model.py:65-66

问题

在生成时,我们只需要最后一个 token 的 logits,但标准实现会计算所有 token 的 logits。

本仓库优化

def forward(self, input_ids, ..., logits_to_keep=1):
    # 仅计算最后 N 个 token 的 logits
    hidden_states = hidden_states[:, -logits_to_keep:]
    logits = self.lm_head(hidden_states)

显存节省

假设 vocab_size=64000,seq_len=32768:

  • 不优化:32768 × 64000 × 2 bytes ≈ 4GB
  • 优化后:1 × 64000 × 2 bytes ≈ 128KB

面试点:为什么可以只算最后 N 个?→ 生成时只需要最后一个 token 的 logits 来采样下一个 token


Q4. 采样策略

Top-k 采样

def top_k_logits(logits, k):
    # 只保留概率最高的 k 个 token
    values, indices = torch.topk(logits, k)
    # 其他 token 设为 -inf
    logits[logits < values[:, -1:]] = float('-inf')
    return logits

Top-p 采样(Nucleus Sampling)

def top_p_logits(logits, p):
    # 按概率排序
    sorted_logits, sorted_indices = torch.sort(logits, descending=True)
    cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
    
    # 移除累积概率超过 p 的 token
    sorted_indices_to_remove = cumulative_probs > p
    sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
    sorted_indices_to_remove[..., 0] = 0
    
    # 恢复原始顺序
    indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove)
    logits[indices_to_remove] = float('-inf')
    return logits

Temperature 采样

def temperature_scale(logits, temperature):
    # temperature < 1: 分布更尖锐,更确定
    # temperature > 1: 分布更平滑,更随机
    return logits / temperature

面试点:Top-k 和 Top-p 的区别?→ Top-k 保留固定数量的 token,Top-p 保留累积概率达到 p 的 token


Q5. Repetition Penalty

问题

生成时可能重复之前的 token,导致输出质量下降。

本仓库实现(src/models/lm/model.py:73-113

def generate(self, input_ids, ..., repetition_penalty=1.2):
    # 对已生成的 token 施加惩罚
    if input_ids.shape[1] > 1:
        # 计算每个 token 的出现次数
        for i in range(input_ids.shape[1]):
            token_id = input_ids[0, i].item()
            logits[0, token_id] /= repetition_penalty

为什么用除法?

  • 除以惩罚系数,降低已生成 token 的概率
  • 保持概率分布的相对顺序

面试点:repetition_penalty=1.0 表示什么?→ 无惩罚,等于没有使用 repetition penalty


Q6. 流式生成(Streaming)

本仓库实现(src/models/lm/model.py:73-113

def generate(self, input_ids, ..., stream_callback=None):
    for _ in range(max_new_tokens):
        # ... 前向传播 ...
        
        # 流式回调
        if stream_callback:
            stream_callback(idx_next)
        
        # 更新序列
        idx_cond = torch.cat([idx_cond, idx_next], dim=1)

为什么需要流式生成?

  1. 用户体验:实时看到生成结果
  2. 早停:用户可以在生成完成前停止
  3. 调试:实时观察生成过程

面试点:流式生成如何实现?→ 通过回调函数,在每步生成后返回当前 token


Q7. GQA 对推理的影响

KV Cache 节省

假设:

  • num_attention_heads = 8
  • num_key_value_heads = 4
  • head_dim = 64

MHA(Multi-Head Attention)

  • KV Cache = 2 × B × S × L × 8 × 64 × 2 bytes

GQA(Grouped-Query Attention)

  • KV Cache = 2 × B × S × L × 4 × 64 × 2 bytes

节省:50%

本仓库实现(src/core/attention.py:13-16

class Attention(nn.Module):
    def __init__(self, config):
        self.n_local_heads = config.num_attention_heads      # 8
        self.n_local_kv_heads = config.num_key_value_heads  # 4
        self.n_rep = self.n_local_heads // self.n_local_kv_heads  # 2 倍复制

面试点:GQA 如何减少 KV Cache?→ KV 头数从 n_heads 减到 n_kv_heads,KV Cache 减少 n_heads/n_kv_heads 倍


Q8. Flash-Attention 对推理的影响

核心思想

IO 感知的注意力计算,避免物化完整的 N×N 注意力矩阵。

本仓库的 Flash-Attention 条件(src/core/attention.py:28, 44

if (seq_len > 1 and 
    (not self.causal or past_key_value is None) and 
    attention_mask is None):
    # 使用 Flash Attention

为什么有条件限制?

  1. seq_len > 1:单 token 无需注意力
  2. not self.causal or past_key_value is None:Flash Attention 对 causal mask 支持有限
  3. attention_mask is None:Flash Attention 不支持自定义 mask

面试点:Flash-Attention 对推理有什么好处?→ 减少 HBM 读写,降低延迟


Q9. 批量推理优化

问题

逐条推理效率低,需要批量处理。

解决方案

  1. Padding:将不同长度的序列 padding 到统一长度
  2. Dynamic Batching:根据序列长度动态调整 batch size
  3. Continuous Batching:不等待整个 batch 完成,动态添加新请求

本仓库的批量推理(src/models/lm/model.py:73-113

def generate(self, input_ids, ...):
    # 支持 batch_size > 1
    for _ in range(max_new_tokens):
        logits, _ = self(idx_cond)
        # ... 采样 ...

面试点:Padding 的缺点是什么?→ 浪费计算资源,短序列需要 padding 到长序列长度


Q10. 量化推理

问题

FP16 精度显存占用高,推理速度慢。

解决方案

  1. INT8 量化:将权重从 FP16 量化到 INT8
  2. INT4 量化:将权重从 FP16 量化到 INT4
  3. GPTQ:基于二阶信息的量化方法
  4. AWQ:激活感知的量化方法

本仓库的量化支持

# 通过 config.dtype 控制精度
config.dtype = 'float16'  # FP16
config.dtype = 'bfloat16'  # BF16

面试点:量化会损失多少精度?→ 取决于量化方法和位数,INT8 通常损失很小,INT4 可能有明显损失


Q11. 推理显存估算

KV Cache 显存

假设:

  • batch_size=B, seq_len=S, num_layers=L
  • num_kv_heads=H, head_dim=D
  • 精度=fp16(2 bytes)

KV Cache 显存 = 2 × B × S × L × H × D × 2 bytes

模型参数显存

假设:

  • hidden_size=d, num_layers=L
  • 精度=fp16(2 bytes)

模型参数显存 = 12 × L × d × d × 2 bytes(Q/K/V/O 四个投影矩阵)

总显存

总显存 ≈ 模型参数显存 + KV Cache 显存

面试点:如何估算推理显存?→ 模型参数显存 + KV Cache 显存


Q12. 推理延迟估算

Prefill 延迟

假设:

  • batch_size=B, seq_len=S, hidden_size=d
  • FLOPS = 2 × B × S² × d

Decode 延迟

假设:

  • batch_size=B, hidden_size=d
  • FLOPS = 2 × B × S × d

总延迟

总延迟 ≈ Prefill 延迟 + Decode 延迟 × 生成长度

面试点:如何优化推理延迟?→ 使用 Flash-Attention、KV Cache、量化等方法


Q13. 推理服务设计

问题

如何设计高并发的推理服务?

解决方案

  1. 动态批处理:根据请求到达时间动态组 batch
  2. 请求排队:使用消息队列管理请求
  3. 负载均衡:将请求分发到多个 GPU
  4. 模型并行:将模型拆分到多个 GPU

本仓库的推理服务(src/serve/

class RealtimeSession:
    def __init__(self, model):
        self.model = model
        self.vad = SileroVAD()  # 语音活动检测
    
    def process(self, audio):
        # 1. VAD 检测
        if not self.vad.detect(audio):
            return None
        
        # 2. 推理
        output = self.model.generate(audio)
        
        return output

面试点:如何提高推理吞吐量?→ 使用动态批处理、模型并行、量化等方法


Q14. 推理与训练的区别

训练

  • 需要反向传播
  • 需要梯度存储
  • 需要优化器状态
  • 显存占用高

推理

  • 只需要前向传播
  • 不需要梯度存储
  • 不需要优化器状态
  • 显存占用低

本仓库的切换

# 训练时
model.train()
loss = model(input_ids, labels=labels)

# 推理时
model.eval()
with torch.no_grad():
    logits = model(input_ids)

面试点:为什么推理时要 torch.no_grad()?→ 节省显存,避免存储梯度


Q15. 推理优化的未来方向

当前瓶颈

  1. 内存墙:显存带宽限制推理速度
  2. 计算墙:GPU 计算能力限制吞吐量
  3. 延迟墙:逐 token 生成限制响应速度

未来方向

  1. 投机采样:用小模型预测,大模型验证
  2. 模型并行:将模型拆分到多个 GPU
  3. 硬件优化:使用专用推理芯片
  4. 算法优化:设计更高效的注意力机制

面试点:投机采样是什么?→ 用小模型快速生成候选,大模型验证并选择,提高生成速度