File size: 15,961 Bytes
c1a46f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# EasyTranslate 任务分工文档

## 项目概述

**项目名称**: EasyTranslate — 基于 Transformer 的英中翻译系统  
**课程**: 研究生 NLP 期末作业  
**团队规模**: 5 人  
**技术栈**: PyTorch + HuggingFace Transformers + Flash Attention 2 + RoPE + LoRA  

---

## 📂 项目结构

```
UCAS-EasyTranslate/
├── configs/
│   ├── default_config.yaml          # 主配置文件
│   └── deepspeed_config.json        # DeepSpeed 分布式训练配置
├── src/easytranslate/
│   ├── data/                        # [Person A] 数据模块
│   │   ├── dataset.py               # 数据集加载
│   │   ├── tokenizer.py             # 分词器
│   │   ├── preprocessing.py         # 数据预处理
│   │   └── collator.py              # 数据整理 & 动态批处理
│   ├── model/                       # [Person B] 模型模块
│   │   ├── transformer.py           # Transformer 主模型
│   │   ├── encoder.py               # 编码器
│   │   ├── decoder.py               # 解码器
│   │   ├── attention.py             # 注意力机制 (标准 + Flash Attention)
│   │   ├── positional.py            # 位置编码 (Sinusoidal + RoPE)
│   │   └── finetune.py              # 预训练模型微调 (NLLB + LoRA)
│   ├── training/                    # [Person C] 训练模块
│   │   ├── trainer.py               # 训练器 (完整训练循环)
│   │   ├── optimizer.py             # 优化器 & 学习率调度
│   │   └── loss.py                  # 损失函数 (标签平滑)
│   ├── evaluation/                  # [Person D] 评估模块
│   │   ├── metrics.py               # 评估指标 (BLEU, COMET, chrF)
│   │   ├── decoding.py              # 解码策略 (贪心, Beam Search, 采样)
│   │   └── evaluator.py             # 评估器
│   └── utils/                       # [Person E] 工具模块
│       ├── config.py                # 配置管理
│       ├── seed.py                  # 随机种子
│       └── logging.py               # 日志管理
├── scripts/
│   ├── train.py                     # 训练入口
│   ├── evaluate.py                  # 评估入口
│   ├── translate.py                 # 翻译推理 (CLI + Web UI)
│   ├── run_experiments.py           # [Person E] 实验运行器
│   └── visualize.py                 # [Person E] 可视化分析
├── tests/
│   ├── test_data.py                 # 数据模块测试
│   ├── test_model.py                # 模型模块测试
│   ├── test_training.py             # 训练模块测试
│   └── test_evaluation.py           # 评估模块测试
├── requirements.txt
├── setup.py
└── README.md
```

---

## 👥 人员分工

### Person A — 数据工程师

**负责文件**: `src/easytranslate/data/` 目录下所有文件  
**预计工作量**: ~800 行代码  
**截止日期建议**: 第 1-2 周

#### 核心任务

| 优先级 | 文件 | 任务 | 要点 |
|--------|------|------|------|
| P0 | `dataset.py` | `TranslationDataset.__getitem__` | 实现 tokenize + padding + teacher forcing 输入构造 |
| P0 | `dataset.py` | `load_wmt_dataset` | 使用 HuggingFace datasets 加载 WMT19 zh-en |
| P0 | `tokenizer.py` | `TokenizerWrapper` | 统一分词器接口 (encode/decode/vocab_size/special_tokens) |
| P0 | `tokenizer.py` | `train_bpe_tokenizer` | 使用 HuggingFace tokenizers 库训练 BPE |
| P0 | `tokenizer.py` | `build_tokenizer` | 根据配置构建分词器 |
| P1 | `preprocessing.py` | `clean_text` | Unicode 标准化 + 控制字符去除 |
| P1 | `preprocessing.py` | `filter_by_length` | 按长度和长度比过滤 |
| P1 | `preprocessing.py` | `preprocess_pipeline` | 完整预处理流水线 |
| P1 | `collator.py` | `TranslationCollator` | batch padding + mask 生成 |
| P2 | `collator.py` | `DynamicBatchSampler` | 按 token 数动态构建 batch |
| P2 | `dataset.py` | `load_opus_dataset` | OPUS-100 数据集加载 |
| P2 | `dataset.py` | `load_custom_dataset` | 自定义语料加载 |

#### 技术参考
- HuggingFace datasets: https://huggingface.co/docs/datasets
- HuggingFace tokenizers: https://huggingface.co/docs/tokenizers
- WMT19 zh-en: `datasets.load_dataset("wmt19", "zh-en")`

#### 验收标准
- [ ] `pytest tests/test_data.py` 全部通过
- [ ] 能够成功加载 WMT19 数据集并完成预处理
- [ ] BPE 分词器训练成功,encode/decode 往返一致
- [ ] DataLoader 能正常迭代,batch 维度正确

---

### Person B — 模型架构师

**负责文件**: `src/easytranslate/model/` 目录下所有文件  
**预计工作量**: ~1000 行代码  
**截止日期建议**: 第 1-2 周

#### 核心任务

| 优先级 | 文件 | 任务 | 要点 |
|--------|------|------|------|
| P0 | `attention.py` | `MultiHeadAttention` | 标准多头注意力: QKV 投影 + scaled dot-product + mask |
| P0 | `attention.py` | `FlashMultiHeadAttention` | 使用 F.scaled_dot_product_attention (PyTorch 2.0+) |
| P0 | `positional.py` | `SinusoidalPositionalEncoding` | 经典正弦余弦位置编码 |
| P0 | `positional.py` | `RotaryPositionalEmbedding` | RoPE 旋转位置编码 (前沿技术!) |
| P0 | `encoder.py` | `TransformerEncoderLayer/Encoder` | Pre-LayerNorm Encoder |
| P0 | `decoder.py` | `TransformerDecoderLayer/Decoder` | Pre-LayerNorm Decoder (self-attn + cross-attn + FFN) |
| P0 | `transformer.py` | `TransformerTranslationModel` | 完整 Enc-Dec 模型: embedding + encoder + decoder + projection |
| P1 | `transformer.py` | `encode / decode_step` | 推理用的编码和单步解码 |
| P1 | `finetune.py` | `load_pretrained_model` | 加载 NLLB/mBART 预训练模型 |
| P1 | `finetune.py` | `setup_lora` | 使用 PEFT 库配置 LoRA 微调 |

#### 关键技术点

1. **Pre-LayerNorm** (比 Post-LN 训练更稳定):
   ```
   x → LayerNorm → Attention → Add(x) → LayerNorm → FFN → Add
   ```

2. **RoPE 旋转位置编码** (LLaMA/GPT-NeoX 使用):
   ```python
   q' = q * cos(θ) + rotate_half(q) * sin(θ)
   k' = k * cos(θ) + rotate_half(k) * sin(θ)
   # attention(q', k') 自动编码相对位置信息
   ```

3. **Flash Attention 2** (PyTorch 2.0+ 原生支持):
   ```python
   F.scaled_dot_product_attention(Q, K, V, attn_mask, dropout_p, is_causal)
   ```

4. **LoRA 微调** (参数高效):
   ```python
   from peft import LoraConfig, get_peft_model, TaskType
   config = LoraConfig(r=16, lora_alpha=32, target_modules=["q_proj", "v_proj"], task_type=TaskType.SEQ_2_SEQ_LM)
   model = get_peft_model(model, config)
   ```

#### 验收标准
- [ ] `pytest tests/test_model.py` 全部通过
- [ ] 模型前向传播输出维度正确: `[B, T, vocab_size]`
- [ ] Flash Attention 和标准 Attention 输出一致
- [ ] RoPE 编码能正确应用到 Q, K
- [ ] LoRA 模型可训练参数量远小于全量参数

---

### Person C — 训练工程师

**负责文件**: `src/easytranslate/training/` 目录下所有文件  
**预计工作量**: ~700 行代码  
**截止日期建议**: 第 2-3 周 (依赖 Person A, B)

#### 核心任务

| 优先级 | 文件 | 任务 | 要点 |
|--------|------|------|------|
| P0 | `loss.py` | `LabelSmoothedCrossEntropyLoss` | 带标签平滑的交叉熵损失 + 忽略 padding |
| P0 | `optimizer.py` | `build_optimizer` | 构建 AdamW,支持参数分组 |
| P0 | `optimizer.py` | `build_scheduler` | Cosine with warmup / Inverse sqrt 调度器 |
| P0 | `trainer.py` | `Trainer.__init__` | 初始化训练环境 (设备, 精度, 分布式) |
| P0 | `trainer.py` | `_train_one_epoch` | 单 epoch 训练循环 (混合精度 + 梯度累积) |
| P0 | `trainer.py` | `_validate` | 验证循环 |
| P1 | `trainer.py` | `train` | 主训练循环 (多 epoch) |
| P1 | `trainer.py` | `_save/_load_checkpoint` | 检查点保存和加载 |
| P1 | `trainer.py` | `_should_early_stop` | 早停逻辑 |
| P2 | `trainer.py` | `_setup_distributed` | DDP/FSDP/DeepSpeed 分布式训练 |
| P2 | `trainer.py` | `_log_metrics` | TensorBoard/WandB 日志 |

#### 关键技术点

1. **混合精度训练**:
   ```python
   scaler = torch.amp.GradScaler('cuda')
   with torch.amp.autocast('cuda', dtype=torch.float16):
       logits = model(src, tgt)
       loss = criterion(logits, labels)
   scaler.scale(loss).backward()
   scaler.step(optimizer)
   scaler.update()
   ```

2. **梯度累积**:
   ```python
   loss = loss / gradient_accumulation_steps
   loss.backward()
   if (step + 1) % gradient_accumulation_steps == 0:
       optimizer.step()
       optimizer.zero_grad()
   ```

3. **标签平滑** (smoothing=0.1):
   ```
   对于 target token: prob = 1 - smoothing = 0.9
   对于其他 token: prob = smoothing / (V - 1) ≈ 0.000003
   ```

#### 验收标准
- [ ] `pytest tests/test_training.py` 全部通过
- [ ] 能在小数据集上完成完整训练流程
- [ ] Loss 能正常下降
- [ ] 检查点能正确保存和加载
- [ ] 早停机制正常工作

---

### Person D — 评估与推理工程师

**负责文件**: `src/easytranslate/evaluation/` + `scripts/evaluate.py` + `scripts/translate.py`  
**预计工作量**: ~800 行代码  
**截止日期建议**: 第 2-3 周 (依赖 Person B)

#### 核心任务

| 优先级 | 文件 | 任务 | 要点 |
|--------|------|------|------|
| P0 | `metrics.py` | `compute_bleu` | SacreBLEU (tokenize="zh" 对中文分词) |
| P0 | `metrics.py` | `compute_comet` | COMET 神经网络评估指标 |
| P0 | `metrics.py` | `compute_chrf` / `compute_ter` | chrF++ 和 TER 指标 |
| P0 | `decoding.py` | `greedy_decode` | 贪心解码 (逐 token argmax) |
| P0 | `decoding.py` | `beam_search_decode` | 束搜索 (最关键的解码算法!) |
| P1 | `decoding.py` | `sample_decode` | 采样解码 (temperature + top-k + top-p) |
| P1 | `evaluator.py` | `Evaluator` | 统一评估接口 |
| P1 | `scripts/evaluate.py` | 评估入口脚本 | 加载模型 → 评估 → 输出结果 |
| P2 | `scripts/translate.py` | 交互翻译 | CLI 交互 + 文件翻译 |
| P2 | `scripts/translate.py` | `launch_web_ui` | Gradio Web 界面 |

#### 关键技术点

1. **Beam Search** (翻译最核心的算法):
   ```
   维护 beam_size 个候选序列
   每步扩展所有候选 → 选 top-k → 继续
   完成后按 score / length^penalty 排序
   ```

2. **SacreBLEU** (标准化评估):
   ```python
   import sacrebleu
   bleu = sacrebleu.corpus_bleu(hypotheses, [references], tokenize="zh")
   ```

3. **COMET** (最准确的评估指标):
   ```python
   from comet import download_model, load_from_checkpoint
   model = load_from_checkpoint(download_model("Unbabel/wmt22-comet-da"))
   data = [{"src": s, "mt": h, "ref": r} for s, h, r in zip(sources, hyps, refs)]
   output = model.predict(data)
   ```

#### 验收标准
- [ ] `pytest tests/test_evaluation.py` 全部通过
- [ ] BLEU 计算结果与 sacrebleu CLI 一致
- [ ] Beam search 翻译质量 ≥ Greedy
- [ ] 能完成端到端评估流程
- [ ] 交互翻译模式正常工作

---

### Person E — 实验与报告

**负责文件**: `src/easytranslate/utils/` + `scripts/run_experiments.py` + `scripts/visualize.py` + 实验报告  
**预计工作量**: ~600 行代码 + 实验报告  
**截止日期建议**: 第 1 周 (utils) + 第 3-4 周 (实验)

#### 核心任务

| 优先级 | 文件 | 任务 | 要点 |
|--------|------|------|------|
| P0 | `utils/config.py` | 配置管理 | OmegaConf 加载 + CLI 覆盖 |
| P0 | `utils/seed.py` | 随机种子 | 全局种子设置 (可复现) |
| P0 | `utils/logging.py` | 日志系统 | Rich 美化 + 文件日志 |
| P1 | `scripts/run_experiments.py` | 实验运行器 | 自动化运行 6 组消融实验 |
| P1 | `scripts/visualize.py` | 训练曲线 | Loss/BLEU/LR 曲线绘制 |
| P1 | `scripts/visualize.py` | 实验对比 | 柱状图 + 表格对比 |
| P2 | `scripts/visualize.py` | 注意力可视化 | Cross-attention 热力图 |
| P2 | `scripts/visualize.py` | 翻译样例 | 好坏案例展示 |
| P0 | — | 实验报告 | 撰写完整实验报告 |

#### 实验设计 (6 组消融实验)

| 实验编号 | 实验名称 | 关键变量 | 目的 |
|----------|----------|----------|------|
| Exp1 | Baseline Transformer | 标准 6 层, d=512, 正弦位置编码 | 基线 |
| Exp2 | + RoPE | 替换正弦编码为 RoPE | 验证旋转位置编码效果 |
| Exp3 | + Flash Attention | 使用 Flash Attention 2 | 验证训练加速效果 |
| Exp4 | Full (RoPE + Flash) | RoPE + Flash Attention | 完整前沿方案 |
| Exp5 | NLLB + LoRA | 预训练 NLLB-600M + LoRA r=16 | 预训练微调效果 |
| Exp6 | NLLB Full FT | 预训练 NLLB-600M 全量微调 | 微调上界 |

#### 报告结构建议
1. **摘要**: 项目目标、方法、主要结果
2. **引言**: 机器翻译背景、Transformer 发展、研究动机
3. **相关工作**: Transformer、NLLB、LoRA、Flash Attention、RoPE
4. **方法**:
   - 模型架构 (附架构图)
   - 训练策略 (混合精度、标签平滑、调度器)
   - 预训练微调方案 (NLLB + LoRA)
5. **实验**:
   - 数据集 (WMT19 zh-en)
   - 实验设置 (超参数表)
   - 实验结果 (BLEU/COMET/chrF 表格)
   - 消融分析 (各组件贡献)
   - 训练效率对比 (Flash Attention 加速比)
6. **分析**:
   - 注意力可视化
   - 翻译案例分析
   - 错误分析
7. **结论与展望**

#### 验收标准
- [ ] utils 模块功能正常
- [ ] 6 组实验能自动化运行
- [ ] 生成完整的可视化图表
- [ ] 实验报告完整,图表清晰

---

## 🗓️ 时间线

```
第 1 周: Person A (数据) + Person B (模型) + Person E (utils) 并行开发

第 2 周: Person C (训练, 依赖 A+B) + Person D (评估, 依赖 B) 开始
         Person A, B 完善和联调

第 3 周: 全体联调 → scripts/train.py 整合
         Person E 开始跑实验

第 4 周: Person E 完成实验 + 报告
         全体 Review + 优化
```

## 🔌 模块接口约定

### 分词器接口 (Person A 定义, Person B/C/D 使用)
```python
tokenizer.encode(text: str) -> list[int]
tokenizer.decode(ids: list[int]) -> str
tokenizer.vocab_size -> int
tokenizer.pad_id -> int
tokenizer.bos_id -> int
tokenizer.eos_id -> int
```

### 模型接口 (Person B 定义, Person C/D 使用)
```python
# 训练
logits = model(src_ids, tgt_input_ids, src_padding_mask, tgt_padding_mask)
# logits: [B, T, vocab_size]

# 推理
encoder_output = model.encode(src_ids, src_padding_mask)
next_logits = model.decode_step(tgt_input_ids, encoder_output, src_padding_mask)
```

### 数据 Batch 格式 (Person A 定义, Person C 使用)
```python
batch = {
    "src_ids": Tensor[B, S],
    "tgt_input_ids": Tensor[B, T],
    "labels": Tensor[B, T],
    "src_padding_mask": BoolTensor[B, S],
    "tgt_padding_mask": BoolTensor[B, T],
}
```

### 评估接口 (Person D 定义, Person C/E 使用)
```python
evaluator = Evaluator(model, tokenizer, config)
results = evaluator.evaluate(dataloader)
# results: {"bleu": 25.6, "comet": 0.82, "chrf": 45.3, "ter": 55.2}
```

---

## ⚙️ 开发环境

```bash
# 1. 创建虚拟环境
conda create -n easytranslate python=3.11
conda activate easytranslate

# 2. 安装 PyTorch (CUDA 12.1)
pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121

# 3. 安装项目依赖
pip install -r requirements.txt

# 4. 安装项目 (开发模式)
pip install -e .

# 5. 运行测试
pytest tests/ -v
```

---

## ✅ 最终交付清单

- [ ] 代码: 所有 `TODO: Person X` 均已实现
- [ ] 测试: `pytest tests/` 全部通过
- [ ] 训练: 至少完成 Exp1 (基线) 和 Exp5 (NLLB+LoRA) 两组实验
- [ ] 评估: 在 WMT19 测试集上报告 BLEU / COMET 分数
- [ ] 报告: 完整实验报告 (含图表)
- [ ] 演示: 交互翻译 Demo 可运行