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. **算法优化**:设计更高效的注意力机制

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