sdfjliom's picture
Implement data pipeline module (#2)
6b73a07
|
Raw
History Blame
3.33 kB
# EasyTranslate Data 模块使用说明
本目录负责英中翻译任务的数据加载、清洗、分词、样本构造和动态批处理。推荐主数据集使用 WMT19 zh-en,调试或小样本实验可使用 OPUS-100 en-zh。
## 文件职责
- `dataset.py`: 加载 WMT/OPUS/custom 数据,并提供 `TranslationDataset`
- `tokenizer.py`: 训练 BPE tokenizer,并用 `TokenizerWrapper` 统一 tokenizer 接口。
- `preprocessing.py`: 文本标准化、长度过滤、去重。
- `collator.py`: batch padding、attention mask、动态 token batch。
## 推荐数据集
正式实验建议使用:
```python
from easytranslate.data import load_wmt_dataset
raw = load_wmt_dataset(year="19", language_pair="zh-en")
```
快速调试建议使用:
```python
from easytranslate.data import load_opus_dataset
raw = load_opus_dataset(subset="en-zh")
```
两个加载函数都会把样本统一成:
```python
{"src": "English sentence", "tgt": "中文句子"}
```
## 从文本到 DataLoader
```python
from torch.utils.data import DataLoader
from easytranslate.data import (
DynamicBatchSampler,
TranslationCollator,
TranslationDataset,
preprocess_pipeline,
train_bpe_tokenizer,
)
src_texts = raw["train"]["src"]
tgt_texts = raw["train"]["tgt"]
src_texts, tgt_texts = preprocess_pipeline(
src_texts,
tgt_texts,
lowercase_src=False,
remove_punctuation=False,
max_src_len=256,
max_tgt_len=256,
length_ratio_threshold=3.0,
)
tokenizer = train_bpe_tokenizer(
list(src_texts) + list(tgt_texts),
vocab_size=32000,
min_frequency=2,
save_path="outputs/tokenizer/bpe.json",
)
train_dataset = TranslationDataset(
src_texts,
tgt_texts,
tokenizer=tokenizer,
max_src_len=256,
max_tgt_len=256,
)
lengths = [
(len(tokenizer.encode(src, add_special_tokens=True)), len(tokenizer.encode(tgt, add_special_tokens=True)))
for src, tgt in zip(src_texts, tgt_texts)
]
batch_sampler = DynamicBatchSampler(lengths, max_tokens_per_batch=8192)
collator = TranslationCollator(pad_token_id=tokenizer.pad_token_id)
loader = DataLoader(
train_dataset,
batch_sampler=batch_sampler,
collate_fn=collator,
num_workers=4,
pin_memory=True,
)
```
## Batch 字段
`TranslationCollator` 输出:
- `src_ids`: `[B, S]`
- `tgt_input_ids`: `[B, T]`,以 `<s>` 开头,用于 teacher forcing
- `labels`: `[B, T]`,以 `</s>` 结尾,padding 为 `-100`
- `src_padding_mask`: `[B, S]`,padding 位置为 `True`
- `tgt_padding_mask`: `[B, T]`,padding 位置为 `True`
- `src_attention_mask` / `tgt_attention_mask`: 有效 token 为 `1`
- `src_lens` / `tgt_lens`: 原始长度
## 自定义平行语料
```python
from easytranslate.data import load_custom_dataset
data = load_custom_dataset(
train_src="data/train.en",
train_tgt="data/train.zh",
val_src="data/val.en",
val_tgt="data/val.zh",
test_src="data/test.en",
test_tgt="data/test.zh",
)
```
## 处理原则
- 保留英文大小写和中英文标点,默认不做 lowercase、不去标点。
- 清洗只做 Unicode NFKC、控制字符删除、空白合并。
- tokenizer 只用训练集训练,不使用验证集或测试集。
- 从零训练 Transformer 时建议使用共享 bilingual BPE;微调 NLLB 时直接使用 NLLB tokenizer。