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`)
```python
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`)
```python
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。
### 本仓库优化
```python
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 采样
```python
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)
```python
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 采样
```python
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`)
```python
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`)
```python
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`)
```python
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`)
```python
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`)
```python
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**:激活感知的量化方法
### 本仓库的量化支持
```python
# 通过 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/`)
```python
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. 推理与训练的区别
### 训练
- 需要反向传播
- 需要梯度存储
- 需要优化器状态
- 显存占用高
### 推理
- 只需要前向传播
- 不需要梯度存储
- 不需要优化器状态
- 显存占用低
### 本仓库的切换
```python
# 训练时
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. **算法优化**:设计更高效的注意力机制
> 面试点:投机采样是什么?→ 用小模型快速生成候选,大模型验证并选择,提高生成速度