xiexinyuan341's picture
Implement data pipeline module
a2d6c00 verified
|
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。

推荐数据集

正式实验建议使用:

from easytranslate.data import load_wmt_dataset

raw = load_wmt_dataset(year="19", language_pair="zh-en")

快速调试建议使用:

from easytranslate.data import load_opus_dataset

raw = load_opus_dataset(subset="en-zh")

两个加载函数都会把样本统一成:

{"src": "English sentence", "tgt": "中文句子"}

从文本到 DataLoader

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: 原始长度

自定义平行语料

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。