面试:推理优化深度
本仓库
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)
为什么需要流式生成?
- 用户体验:实时看到生成结果
- 早停:用户可以在生成完成前停止
- 调试:实时观察生成过程
面试点:流式生成如何实现?→ 通过回调函数,在每步生成后返回当前 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
为什么有条件限制?
seq_len > 1:单 token 无需注意力not self.causal or past_key_value is None:Flash Attention 对 causal mask 支持有限attention_mask is None:Flash Attention 不支持自定义 mask
面试点:Flash-Attention 对推理有什么好处?→ 减少 HBM 读写,降低延迟
Q9. 批量推理优化
问题
逐条推理效率低,需要批量处理。
解决方案
- Padding:将不同长度的序列 padding 到统一长度
- Dynamic Batching:根据序列长度动态调整 batch size
- 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 精度显存占用高,推理速度慢。
解决方案
- INT8 量化:将权重从 FP16 量化到 INT8
- INT4 量化:将权重从 FP16 量化到 INT4
- GPTQ:基于二阶信息的量化方法
- 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. 推理服务设计
问题
如何设计高并发的推理服务?
解决方案
- 动态批处理:根据请求到达时间动态组 batch
- 请求排队:使用消息队列管理请求
- 负载均衡:将请求分发到多个 GPU
- 模型并行:将模型拆分到多个 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. 推理优化的未来方向
当前瓶颈
- 内存墙:显存带宽限制推理速度
- 计算墙:GPU 计算能力限制吞吐量
- 延迟墙:逐 token 生成限制响应速度
未来方向
- 投机采样:用小模型预测,大模型验证
- 模型并行:将模型拆分到多个 GPU
- 硬件优化:使用专用推理芯片
- 算法优化:设计更高效的注意力机制
面试点:投机采样是什么?→ 用小模型快速生成候选,大模型验证并选择,提高生成速度