File size: 12,415 Bytes
f664f3f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 | # 面试:推理优化深度
> 本仓库 `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. **算法优化**:设计更高效的注意力机制
> 面试点:投机采样是什么?→ 用小模型快速生成候选,大模型验证并选择,提高生成速度
|