Implement data pipeline module
#2
by xiexinyuan341 - opened
- .gitignore +1 -1
- requirements.txt +1 -1
- setup.py +2 -0
- src/easytranslate/data/README.md +126 -0
- src/easytranslate/data/__init__.py +32 -0
- src/easytranslate/data/collator.py +114 -0
- src/easytranslate/data/dataset.py +157 -0
- src/easytranslate/data/preprocessing.py +125 -0
- src/easytranslate/data/tokenizer.py +197 -0
- tests/test_data.py +94 -23
.gitignore
CHANGED
|
@@ -21,7 +21,7 @@ env/
|
|
| 21 |
*.swo
|
| 22 |
|
| 23 |
# Data
|
| 24 |
-
data/
|
| 25 |
*.tsv
|
| 26 |
*.csv
|
| 27 |
|
|
|
|
| 21 |
*.swo
|
| 22 |
|
| 23 |
# Data
|
| 24 |
+
/data/
|
| 25 |
*.tsv
|
| 26 |
*.csv
|
| 27 |
|
requirements.txt
CHANGED
|
@@ -20,7 +20,7 @@ rouge-score>=0.1.2
|
|
| 20 |
nltk>=3.8.0
|
| 21 |
|
| 22 |
# Utilities
|
| 23 |
-
numpy>=1.24.0
|
| 24 |
pandas>=2.0.0
|
| 25 |
tqdm>=4.66.0
|
| 26 |
pyyaml>=6.0.0
|
|
|
|
| 20 |
nltk>=3.8.0
|
| 21 |
|
| 22 |
# Utilities
|
| 23 |
+
numpy>=1.24.0,<2.0
|
| 24 |
pandas>=2.0.0
|
| 25 |
tqdm>=4.66.0
|
| 26 |
pyyaml>=6.0.0
|
setup.py
CHANGED
|
@@ -11,9 +11,11 @@ setup(
|
|
| 11 |
"torch>=2.1.0",
|
| 12 |
"transformers>=4.36.0",
|
| 13 |
"datasets>=2.16.0",
|
|
|
|
| 14 |
"sentencepiece>=0.1.99",
|
| 15 |
"accelerate>=0.25.0",
|
| 16 |
"sacrebleu>=2.4.0",
|
|
|
|
| 17 |
"omegaconf>=2.3.0",
|
| 18 |
"rich>=13.0.0",
|
| 19 |
"tqdm>=4.66.0",
|
|
|
|
| 11 |
"torch>=2.1.0",
|
| 12 |
"transformers>=4.36.0",
|
| 13 |
"datasets>=2.16.0",
|
| 14 |
+
"tokenizers>=0.15.0",
|
| 15 |
"sentencepiece>=0.1.99",
|
| 16 |
"accelerate>=0.25.0",
|
| 17 |
"sacrebleu>=2.4.0",
|
| 18 |
+
"numpy>=1.24.0,<2.0",
|
| 19 |
"omegaconf>=2.3.0",
|
| 20 |
"rich>=13.0.0",
|
| 21 |
"tqdm>=4.66.0",
|
src/easytranslate/data/README.md
ADDED
|
@@ -0,0 +1,126 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# EasyTranslate Data 模块使用说明
|
| 2 |
+
|
| 3 |
+
本目录负责英中翻译任务的数据加载、清洗、分词、样本构造和动态批处理。推荐主数据集使用 WMT19 zh-en,调试或小样本实验可使用 OPUS-100 en-zh。
|
| 4 |
+
|
| 5 |
+
## 文件职责
|
| 6 |
+
|
| 7 |
+
- `dataset.py`: 加载 WMT/OPUS/custom 数据,并提供 `TranslationDataset`。
|
| 8 |
+
- `tokenizer.py`: 训练 BPE tokenizer,并用 `TokenizerWrapper` 统一 tokenizer 接口。
|
| 9 |
+
- `preprocessing.py`: 文本标准化、长度过滤、去重。
|
| 10 |
+
- `collator.py`: batch padding、attention mask、动态 token batch。
|
| 11 |
+
|
| 12 |
+
## 推荐数据集
|
| 13 |
+
|
| 14 |
+
正式实验建议使用:
|
| 15 |
+
|
| 16 |
+
```python
|
| 17 |
+
from easytranslate.data import load_wmt_dataset
|
| 18 |
+
|
| 19 |
+
raw = load_wmt_dataset(year="19", language_pair="zh-en")
|
| 20 |
+
```
|
| 21 |
+
|
| 22 |
+
快速调试建议使用:
|
| 23 |
+
|
| 24 |
+
```python
|
| 25 |
+
from easytranslate.data import load_opus_dataset
|
| 26 |
+
|
| 27 |
+
raw = load_opus_dataset(subset="en-zh")
|
| 28 |
+
```
|
| 29 |
+
|
| 30 |
+
两个加载函数都会把样本统一成:
|
| 31 |
+
|
| 32 |
+
```python
|
| 33 |
+
{"src": "English sentence", "tgt": "中文句子"}
|
| 34 |
+
```
|
| 35 |
+
|
| 36 |
+
## 从文本到 DataLoader
|
| 37 |
+
|
| 38 |
+
```python
|
| 39 |
+
from torch.utils.data import DataLoader
|
| 40 |
+
|
| 41 |
+
from easytranslate.data import (
|
| 42 |
+
DynamicBatchSampler,
|
| 43 |
+
TranslationCollator,
|
| 44 |
+
TranslationDataset,
|
| 45 |
+
preprocess_pipeline,
|
| 46 |
+
train_bpe_tokenizer,
|
| 47 |
+
)
|
| 48 |
+
|
| 49 |
+
src_texts = raw["train"]["src"]
|
| 50 |
+
tgt_texts = raw["train"]["tgt"]
|
| 51 |
+
|
| 52 |
+
src_texts, tgt_texts = preprocess_pipeline(
|
| 53 |
+
src_texts,
|
| 54 |
+
tgt_texts,
|
| 55 |
+
lowercase_src=False,
|
| 56 |
+
remove_punctuation=False,
|
| 57 |
+
max_src_len=256,
|
| 58 |
+
max_tgt_len=256,
|
| 59 |
+
length_ratio_threshold=3.0,
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
tokenizer = train_bpe_tokenizer(
|
| 63 |
+
list(src_texts) + list(tgt_texts),
|
| 64 |
+
vocab_size=32000,
|
| 65 |
+
min_frequency=2,
|
| 66 |
+
save_path="outputs/tokenizer/bpe.json",
|
| 67 |
+
)
|
| 68 |
+
|
| 69 |
+
train_dataset = TranslationDataset(
|
| 70 |
+
src_texts,
|
| 71 |
+
tgt_texts,
|
| 72 |
+
tokenizer=tokenizer,
|
| 73 |
+
max_src_len=256,
|
| 74 |
+
max_tgt_len=256,
|
| 75 |
+
)
|
| 76 |
+
|
| 77 |
+
lengths = [
|
| 78 |
+
(len(tokenizer.encode(src, add_special_tokens=True)), len(tokenizer.encode(tgt, add_special_tokens=True)))
|
| 79 |
+
for src, tgt in zip(src_texts, tgt_texts)
|
| 80 |
+
]
|
| 81 |
+
|
| 82 |
+
batch_sampler = DynamicBatchSampler(lengths, max_tokens_per_batch=8192)
|
| 83 |
+
collator = TranslationCollator(pad_token_id=tokenizer.pad_token_id)
|
| 84 |
+
|
| 85 |
+
loader = DataLoader(
|
| 86 |
+
train_dataset,
|
| 87 |
+
batch_sampler=batch_sampler,
|
| 88 |
+
collate_fn=collator,
|
| 89 |
+
num_workers=4,
|
| 90 |
+
pin_memory=True,
|
| 91 |
+
)
|
| 92 |
+
```
|
| 93 |
+
|
| 94 |
+
## Batch 字段
|
| 95 |
+
|
| 96 |
+
`TranslationCollator` 输出:
|
| 97 |
+
|
| 98 |
+
- `src_ids`: `[B, S]`
|
| 99 |
+
- `tgt_input_ids`: `[B, T]`,以 `<s>` 开头,用于 teacher forcing
|
| 100 |
+
- `labels`: `[B, T]`,以 `</s>` 结尾,padding 为 `-100`
|
| 101 |
+
- `src_padding_mask`: `[B, S]`,padding 位置为 `True`
|
| 102 |
+
- `tgt_padding_mask`: `[B, T]`,padding 位置为 `True`
|
| 103 |
+
- `src_attention_mask` / `tgt_attention_mask`: 有效 token 为 `1`
|
| 104 |
+
- `src_lens` / `tgt_lens`: 原始长度
|
| 105 |
+
|
| 106 |
+
## 自定义平行语料
|
| 107 |
+
|
| 108 |
+
```python
|
| 109 |
+
from easytranslate.data import load_custom_dataset
|
| 110 |
+
|
| 111 |
+
data = load_custom_dataset(
|
| 112 |
+
train_src="data/train.en",
|
| 113 |
+
train_tgt="data/train.zh",
|
| 114 |
+
val_src="data/val.en",
|
| 115 |
+
val_tgt="data/val.zh",
|
| 116 |
+
test_src="data/test.en",
|
| 117 |
+
test_tgt="data/test.zh",
|
| 118 |
+
)
|
| 119 |
+
```
|
| 120 |
+
|
| 121 |
+
## 处理原则
|
| 122 |
+
|
| 123 |
+
- 保留英文大小写和中英文标点,默认不做 lowercase、不去标点。
|
| 124 |
+
- 清洗只做 Unicode NFKC、控制字符删除、空白合并。
|
| 125 |
+
- tokenizer 只用训练集训练,不使用验证集或测试集。
|
| 126 |
+
- 从零训练 Transformer 时建议使用共享 bilingual BPE;微调 NLLB 时直接使用 NLLB tokenizer。
|
src/easytranslate/data/__init__.py
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Data utilities for EasyTranslate."""
|
| 2 |
+
|
| 3 |
+
from easytranslate.data.collator import DynamicBatchSampler, TranslationCollator
|
| 4 |
+
from easytranslate.data.dataset import (
|
| 5 |
+
TranslationDataset,
|
| 6 |
+
load_custom_dataset,
|
| 7 |
+
load_opus_dataset,
|
| 8 |
+
load_wmt_dataset,
|
| 9 |
+
)
|
| 10 |
+
from easytranslate.data.preprocessing import (
|
| 11 |
+
clean_text,
|
| 12 |
+
deduplicate_pairs,
|
| 13 |
+
filter_by_length,
|
| 14 |
+
preprocess_pipeline,
|
| 15 |
+
)
|
| 16 |
+
from easytranslate.data.tokenizer import TokenizerWrapper, build_tokenizer, train_bpe_tokenizer
|
| 17 |
+
|
| 18 |
+
__all__ = [
|
| 19 |
+
"DynamicBatchSampler",
|
| 20 |
+
"TokenizerWrapper",
|
| 21 |
+
"TranslationCollator",
|
| 22 |
+
"TranslationDataset",
|
| 23 |
+
"build_tokenizer",
|
| 24 |
+
"clean_text",
|
| 25 |
+
"deduplicate_pairs",
|
| 26 |
+
"filter_by_length",
|
| 27 |
+
"load_custom_dataset",
|
| 28 |
+
"load_opus_dataset",
|
| 29 |
+
"load_wmt_dataset",
|
| 30 |
+
"preprocess_pipeline",
|
| 31 |
+
"train_bpe_tokenizer",
|
| 32 |
+
]
|
src/easytranslate/data/collator.py
ADDED
|
@@ -0,0 +1,114 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Batch collation and dynamic batching for translation training."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import random
|
| 6 |
+
from typing import Iterator, Sequence
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
from torch.nn.utils.rnn import pad_sequence
|
| 10 |
+
from torch.utils.data import Sampler
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class TranslationCollator:
|
| 14 |
+
"""Pad variable-length translation examples into one batch."""
|
| 15 |
+
|
| 16 |
+
def __init__(self, pad_token_id: int = 0, label_pad_token_id: int = -100):
|
| 17 |
+
self.pad_token_id = pad_token_id
|
| 18 |
+
self.label_pad_token_id = label_pad_token_id
|
| 19 |
+
|
| 20 |
+
def __call__(self, batch: Sequence[dict[str, torch.Tensor]]) -> dict[str, torch.Tensor]:
|
| 21 |
+
src_ids = pad_sequence(
|
| 22 |
+
[item["src_ids"] for item in batch],
|
| 23 |
+
batch_first=True,
|
| 24 |
+
padding_value=self.pad_token_id,
|
| 25 |
+
)
|
| 26 |
+
tgt_input_ids = pad_sequence(
|
| 27 |
+
[item["tgt_input_ids"] for item in batch],
|
| 28 |
+
batch_first=True,
|
| 29 |
+
padding_value=self.pad_token_id,
|
| 30 |
+
)
|
| 31 |
+
labels = pad_sequence(
|
| 32 |
+
[item["labels"] for item in batch],
|
| 33 |
+
batch_first=True,
|
| 34 |
+
padding_value=self.label_pad_token_id,
|
| 35 |
+
)
|
| 36 |
+
|
| 37 |
+
src_padding_mask = src_ids.eq(self.pad_token_id)
|
| 38 |
+
tgt_padding_mask = tgt_input_ids.eq(self.pad_token_id)
|
| 39 |
+
|
| 40 |
+
return {
|
| 41 |
+
"src_ids": src_ids,
|
| 42 |
+
"tgt_input_ids": tgt_input_ids,
|
| 43 |
+
"labels": labels,
|
| 44 |
+
"src_padding_mask": src_padding_mask,
|
| 45 |
+
"tgt_padding_mask": tgt_padding_mask,
|
| 46 |
+
"src_attention_mask": (~src_padding_mask).long(),
|
| 47 |
+
"tgt_attention_mask": (~tgt_padding_mask).long(),
|
| 48 |
+
"src_lens": torch.stack([item["src_len"] for item in batch]),
|
| 49 |
+
"tgt_lens": torch.stack([item["tgt_len"] for item in batch]),
|
| 50 |
+
}
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
class DynamicBatchSampler(Sampler[list[int]]):
|
| 54 |
+
"""Create batches constrained by an approximate max token budget."""
|
| 55 |
+
|
| 56 |
+
def __init__(
|
| 57 |
+
self,
|
| 58 |
+
lengths: Sequence[int | tuple[int, int]],
|
| 59 |
+
max_tokens_per_batch: int = 8192,
|
| 60 |
+
shuffle: bool = True,
|
| 61 |
+
drop_last: bool = False,
|
| 62 |
+
):
|
| 63 |
+
self.lengths = [max(length) if isinstance(length, tuple) else int(length) for length in lengths]
|
| 64 |
+
self.max_tokens_per_batch = max_tokens_per_batch
|
| 65 |
+
self.shuffle = shuffle
|
| 66 |
+
self.drop_last = drop_last
|
| 67 |
+
|
| 68 |
+
def __iter__(self) -> Iterator[list[int]]:
|
| 69 |
+
indices = list(range(len(self.lengths)))
|
| 70 |
+
if self.shuffle:
|
| 71 |
+
random.shuffle(indices)
|
| 72 |
+
|
| 73 |
+
indices.sort(key=lambda idx: self.lengths[idx])
|
| 74 |
+
batches: list[list[int]] = []
|
| 75 |
+
batch: list[int] = []
|
| 76 |
+
max_len = 0
|
| 77 |
+
|
| 78 |
+
for idx in indices:
|
| 79 |
+
candidate_max_len = max(max_len, self.lengths[idx])
|
| 80 |
+
candidate_tokens = candidate_max_len * (len(batch) + 1)
|
| 81 |
+
|
| 82 |
+
if batch and candidate_tokens > self.max_tokens_per_batch:
|
| 83 |
+
batches.append(batch)
|
| 84 |
+
batch = []
|
| 85 |
+
max_len = 0
|
| 86 |
+
|
| 87 |
+
batch.append(idx)
|
| 88 |
+
max_len = max(max_len, self.lengths[idx])
|
| 89 |
+
|
| 90 |
+
if batch and not self.drop_last:
|
| 91 |
+
batches.append(batch)
|
| 92 |
+
|
| 93 |
+
if self.shuffle:
|
| 94 |
+
random.shuffle(batches)
|
| 95 |
+
|
| 96 |
+
yield from batches
|
| 97 |
+
|
| 98 |
+
def __len__(self) -> int:
|
| 99 |
+
count = 0
|
| 100 |
+
batch_size = 0
|
| 101 |
+
max_len = 0
|
| 102 |
+
|
| 103 |
+
for length in sorted(self.lengths):
|
| 104 |
+
candidate_max_len = max(max_len, length)
|
| 105 |
+
if batch_size and candidate_max_len * (batch_size + 1) > self.max_tokens_per_batch:
|
| 106 |
+
count += 1
|
| 107 |
+
batch_size = 0
|
| 108 |
+
max_len = 0
|
| 109 |
+
batch_size += 1
|
| 110 |
+
max_len = max(max_len, length)
|
| 111 |
+
|
| 112 |
+
if batch_size and not self.drop_last:
|
| 113 |
+
count += 1
|
| 114 |
+
return count
|
src/easytranslate/data/dataset.py
ADDED
|
@@ -0,0 +1,157 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Dataset loading and PyTorch dataset classes for English-Chinese translation."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
from typing import Mapping, Sequence
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
from torch.utils.data import Dataset
|
| 10 |
+
|
| 11 |
+
from easytranslate.data.preprocessing import preprocess_pipeline
|
| 12 |
+
from easytranslate.data.tokenizer import TokenizerWrapper
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
class TranslationDataset(Dataset):
|
| 16 |
+
"""PyTorch dataset that builds seq2seq inputs for teacher forcing."""
|
| 17 |
+
|
| 18 |
+
def __init__(
|
| 19 |
+
self,
|
| 20 |
+
src_texts: Sequence[str],
|
| 21 |
+
tgt_texts: Sequence[str],
|
| 22 |
+
tokenizer: TokenizerWrapper | None = None,
|
| 23 |
+
src_tokenizer: TokenizerWrapper | None = None,
|
| 24 |
+
tgt_tokenizer: TokenizerWrapper | None = None,
|
| 25 |
+
max_src_len: int = 256,
|
| 26 |
+
max_tgt_len: int = 256,
|
| 27 |
+
):
|
| 28 |
+
if len(src_texts) != len(tgt_texts):
|
| 29 |
+
raise ValueError("src_texts and tgt_texts must have the same length")
|
| 30 |
+
|
| 31 |
+
if tokenizer is not None:
|
| 32 |
+
src_tokenizer = src_tokenizer or tokenizer
|
| 33 |
+
tgt_tokenizer = tgt_tokenizer or tokenizer
|
| 34 |
+
if src_tokenizer is None or tgt_tokenizer is None:
|
| 35 |
+
raise ValueError("Provide tokenizer or both src_tokenizer and tgt_tokenizer")
|
| 36 |
+
|
| 37 |
+
self.src_texts = list(src_texts)
|
| 38 |
+
self.tgt_texts = list(tgt_texts)
|
| 39 |
+
self.src_tokenizer = src_tokenizer
|
| 40 |
+
self.tgt_tokenizer = tgt_tokenizer
|
| 41 |
+
self.max_src_len = max_src_len
|
| 42 |
+
self.max_tgt_len = max_tgt_len
|
| 43 |
+
|
| 44 |
+
def __len__(self) -> int:
|
| 45 |
+
return len(self.src_texts)
|
| 46 |
+
|
| 47 |
+
def __getitem__(self, index: int) -> dict[str, torch.Tensor]:
|
| 48 |
+
src_ids = self.src_tokenizer.encode(
|
| 49 |
+
self.src_texts[index],
|
| 50 |
+
add_special_tokens=True,
|
| 51 |
+
max_length=self.max_src_len,
|
| 52 |
+
)
|
| 53 |
+
|
| 54 |
+
target_core = self.tgt_tokenizer.encode(
|
| 55 |
+
self.tgt_texts[index],
|
| 56 |
+
add_special_tokens=False,
|
| 57 |
+
max_length=max(1, self.max_tgt_len - 1),
|
| 58 |
+
)
|
| 59 |
+
tgt_input_ids = [self.tgt_tokenizer.bos_token_id] + target_core
|
| 60 |
+
labels = target_core + [self.tgt_tokenizer.eos_token_id]
|
| 61 |
+
|
| 62 |
+
return {
|
| 63 |
+
"src_ids": torch.tensor(src_ids, dtype=torch.long),
|
| 64 |
+
"tgt_input_ids": torch.tensor(tgt_input_ids, dtype=torch.long),
|
| 65 |
+
"labels": torch.tensor(labels, dtype=torch.long),
|
| 66 |
+
"src_len": torch.tensor(len(src_ids), dtype=torch.long),
|
| 67 |
+
"tgt_len": torch.tensor(len(labels), dtype=torch.long),
|
| 68 |
+
}
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def _translation_to_columns(dataset, src_lang: str, tgt_lang: str):
|
| 72 |
+
def convert(example):
|
| 73 |
+
translation = example["translation"]
|
| 74 |
+
return {"src": translation[src_lang], "tgt": translation[tgt_lang]}
|
| 75 |
+
|
| 76 |
+
return dataset.map(convert, remove_columns=dataset.column_names)
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def load_wmt_dataset(
|
| 80 |
+
year: str = "19",
|
| 81 |
+
language_pair: str = "zh-en",
|
| 82 |
+
src_lang: str = "en",
|
| 83 |
+
tgt_lang: str = "zh",
|
| 84 |
+
split: str | None = None,
|
| 85 |
+
cache_dir: str | None = None,
|
| 86 |
+
):
|
| 87 |
+
"""Load WMT zh-en and normalize rows to {'src', 'tgt'}."""
|
| 88 |
+
from datasets import load_dataset
|
| 89 |
+
|
| 90 |
+
dataset_name = f"wmt/wmt{year}"
|
| 91 |
+
try:
|
| 92 |
+
dataset = load_dataset(dataset_name, language_pair, split=split, cache_dir=cache_dir)
|
| 93 |
+
except Exception:
|
| 94 |
+
dataset = load_dataset(f"wmt{year}", language_pair, split=split, cache_dir=cache_dir)
|
| 95 |
+
|
| 96 |
+
if split is not None:
|
| 97 |
+
return _translation_to_columns(dataset, src_lang, tgt_lang)
|
| 98 |
+
|
| 99 |
+
return dataset.map(
|
| 100 |
+
lambda example: {"src": example["translation"][src_lang], "tgt": example["translation"][tgt_lang]},
|
| 101 |
+
remove_columns=next(iter(dataset.values())).column_names,
|
| 102 |
+
)
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def load_opus_dataset(
|
| 106 |
+
subset: str = "en-zh",
|
| 107 |
+
src_lang: str = "en",
|
| 108 |
+
tgt_lang: str = "zh",
|
| 109 |
+
split: str | None = None,
|
| 110 |
+
cache_dir: str | None = None,
|
| 111 |
+
):
|
| 112 |
+
"""Load OPUS-100 and normalize rows to {'src', 'tgt'}."""
|
| 113 |
+
from datasets import load_dataset
|
| 114 |
+
|
| 115 |
+
dataset = load_dataset("Helsinki-NLP/opus-100", subset, split=split, cache_dir=cache_dir)
|
| 116 |
+
if split is not None:
|
| 117 |
+
return _translation_to_columns(dataset, src_lang, tgt_lang)
|
| 118 |
+
|
| 119 |
+
return dataset.map(
|
| 120 |
+
lambda example: {"src": example["translation"][src_lang], "tgt": example["translation"][tgt_lang]},
|
| 121 |
+
remove_columns=next(iter(dataset.values())).column_names,
|
| 122 |
+
)
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def _read_lines(path: str | Path) -> list[str]:
|
| 126 |
+
with Path(path).open("r", encoding="utf-8") as f:
|
| 127 |
+
return [line.rstrip("\n") for line in f]
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def load_custom_dataset(
|
| 131 |
+
train_src: str | Path,
|
| 132 |
+
train_tgt: str | Path,
|
| 133 |
+
val_src: str | Path | None = None,
|
| 134 |
+
val_tgt: str | Path | None = None,
|
| 135 |
+
test_src: str | Path | None = None,
|
| 136 |
+
test_tgt: str | Path | None = None,
|
| 137 |
+
preprocessing_config: Mapping | None = None,
|
| 138 |
+
) -> dict[str, dict[str, list[str]]]:
|
| 139 |
+
"""Load parallel text files and return split dictionaries."""
|
| 140 |
+
preprocessing_config = dict(preprocessing_config or {})
|
| 141 |
+
|
| 142 |
+
def load_split(src_path: str | Path, tgt_path: str | Path) -> dict[str, list[str]]:
|
| 143 |
+
src_texts, tgt_texts = preprocess_pipeline(
|
| 144 |
+
_read_lines(src_path),
|
| 145 |
+
_read_lines(tgt_path),
|
| 146 |
+
**preprocessing_config,
|
| 147 |
+
)
|
| 148 |
+
return {"src": src_texts, "tgt": tgt_texts}
|
| 149 |
+
|
| 150 |
+
result = {"train": load_split(train_src, train_tgt)}
|
| 151 |
+
|
| 152 |
+
if val_src and val_tgt:
|
| 153 |
+
result["validation"] = load_split(val_src, val_tgt)
|
| 154 |
+
if test_src and test_tgt:
|
| 155 |
+
result["test"] = load_split(test_src, test_tgt)
|
| 156 |
+
|
| 157 |
+
return result
|
src/easytranslate/data/preprocessing.py
ADDED
|
@@ -0,0 +1,125 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Text cleaning and filtering utilities for translation corpora."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import re
|
| 6 |
+
import unicodedata
|
| 7 |
+
from typing import Callable, Iterable, Sequence
|
| 8 |
+
|
| 9 |
+
CONTROL_OR_ZERO_WIDTH_RE = re.compile(r"[\u0000-\u001f\u007f-\u009f\u200b\u200c\u200d\ufeff]")
|
| 10 |
+
SPACE_RE = re.compile(r"\s+")
|
| 11 |
+
PUNCT_RE = re.compile(r"[^\w\s\u4e00-\u9fff]", flags=re.UNICODE)
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def clean_text(text: str, lowercase: bool = False, remove_punctuation: bool = False) -> str:
|
| 15 |
+
"""Normalize a single sentence without changing its meaning aggressively."""
|
| 16 |
+
if text is None:
|
| 17 |
+
return ""
|
| 18 |
+
|
| 19 |
+
text = unicodedata.normalize("NFKC", str(text))
|
| 20 |
+
text = CONTROL_OR_ZERO_WIDTH_RE.sub("", text)
|
| 21 |
+
text = SPACE_RE.sub(" ", text).strip()
|
| 22 |
+
|
| 23 |
+
if lowercase:
|
| 24 |
+
text = text.lower()
|
| 25 |
+
if remove_punctuation:
|
| 26 |
+
text = PUNCT_RE.sub("", text)
|
| 27 |
+
text = SPACE_RE.sub(" ", text).strip()
|
| 28 |
+
|
| 29 |
+
return text
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def _default_length(text: str) -> int:
|
| 33 |
+
"""Use whitespace tokens for Latin text and character count for CJK-heavy text."""
|
| 34 |
+
cjk_chars = sum(1 for ch in text if "\u4e00" <= ch <= "\u9fff")
|
| 35 |
+
if cjk_chars >= max(1, len(text) // 3):
|
| 36 |
+
return len(text.replace(" ", ""))
|
| 37 |
+
return len(text.split())
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def filter_by_length(
|
| 41 |
+
src: str,
|
| 42 |
+
tgt: str,
|
| 43 |
+
min_src_len: int = 1,
|
| 44 |
+
min_tgt_len: int = 1,
|
| 45 |
+
max_src_len: int = 256,
|
| 46 |
+
max_tgt_len: int = 256,
|
| 47 |
+
length_ratio_threshold: float = 3.0,
|
| 48 |
+
length_fn: Callable[[str], int] | None = None,
|
| 49 |
+
) -> bool:
|
| 50 |
+
"""Return True when a sentence pair passes basic length and ratio checks."""
|
| 51 |
+
length_fn = length_fn or _default_length
|
| 52 |
+
src_len = length_fn(src)
|
| 53 |
+
tgt_len = length_fn(tgt)
|
| 54 |
+
|
| 55 |
+
if src_len < min_src_len or tgt_len < min_tgt_len:
|
| 56 |
+
return False
|
| 57 |
+
if src_len > max_src_len or tgt_len > max_tgt_len:
|
| 58 |
+
return False
|
| 59 |
+
|
| 60 |
+
shorter = max(1, min(src_len, tgt_len))
|
| 61 |
+
longer = max(src_len, tgt_len)
|
| 62 |
+
return longer / shorter <= length_ratio_threshold
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
def deduplicate_pairs(pairs: Iterable[tuple[str, str]]) -> list[tuple[str, str]]:
|
| 66 |
+
"""Deduplicate by exact cleaned source-target pair while preserving order."""
|
| 67 |
+
seen: set[tuple[str, str]] = set()
|
| 68 |
+
result: list[tuple[str, str]] = []
|
| 69 |
+
|
| 70 |
+
for src, tgt in pairs:
|
| 71 |
+
key = (src, tgt)
|
| 72 |
+
if key in seen:
|
| 73 |
+
continue
|
| 74 |
+
seen.add(key)
|
| 75 |
+
result.append(key)
|
| 76 |
+
|
| 77 |
+
return result
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def preprocess_pipeline(
|
| 81 |
+
src_texts: Sequence[str],
|
| 82 |
+
tgt_texts: Sequence[str],
|
| 83 |
+
lowercase_src: bool = False,
|
| 84 |
+
lowercase_tgt: bool = False,
|
| 85 |
+
remove_punctuation: bool = False,
|
| 86 |
+
max_src_len: int = 256,
|
| 87 |
+
max_tgt_len: int = 256,
|
| 88 |
+
min_src_len: int = 1,
|
| 89 |
+
min_tgt_len: int = 1,
|
| 90 |
+
filter_by_length_enabled: bool = True,
|
| 91 |
+
length_ratio_threshold: float = 3.0,
|
| 92 |
+
deduplicate: bool = True,
|
| 93 |
+
) -> tuple[list[str], list[str]]:
|
| 94 |
+
"""Clean, filter, and optionally deduplicate parallel source-target texts."""
|
| 95 |
+
if len(src_texts) != len(tgt_texts):
|
| 96 |
+
raise ValueError("src_texts and tgt_texts must have the same length")
|
| 97 |
+
|
| 98 |
+
pairs: list[tuple[str, str]] = []
|
| 99 |
+
for raw_src, raw_tgt in zip(src_texts, tgt_texts):
|
| 100 |
+
src = clean_text(raw_src, lowercase=lowercase_src, remove_punctuation=remove_punctuation)
|
| 101 |
+
tgt = clean_text(raw_tgt, lowercase=lowercase_tgt, remove_punctuation=remove_punctuation)
|
| 102 |
+
|
| 103 |
+
if not src or not tgt:
|
| 104 |
+
continue
|
| 105 |
+
if filter_by_length_enabled and not filter_by_length(
|
| 106 |
+
src,
|
| 107 |
+
tgt,
|
| 108 |
+
min_src_len=min_src_len,
|
| 109 |
+
min_tgt_len=min_tgt_len,
|
| 110 |
+
max_src_len=max_src_len,
|
| 111 |
+
max_tgt_len=max_tgt_len,
|
| 112 |
+
length_ratio_threshold=length_ratio_threshold,
|
| 113 |
+
):
|
| 114 |
+
continue
|
| 115 |
+
|
| 116 |
+
pairs.append((src, tgt))
|
| 117 |
+
|
| 118 |
+
if deduplicate:
|
| 119 |
+
pairs = deduplicate_pairs(pairs)
|
| 120 |
+
|
| 121 |
+
if not pairs:
|
| 122 |
+
return [], []
|
| 123 |
+
|
| 124 |
+
src_clean, tgt_clean = zip(*pairs)
|
| 125 |
+
return list(src_clean), list(tgt_clean)
|
src/easytranslate/data/tokenizer.py
ADDED
|
@@ -0,0 +1,197 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Tokenizer wrappers and BPE training helpers."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
from typing import Iterable, Mapping, Sequence
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
class TokenizerWrapper:
|
| 10 |
+
"""Small adapter that gives HF tokenizers and tokenizers.Tokenizer one API."""
|
| 11 |
+
|
| 12 |
+
def __init__(
|
| 13 |
+
self,
|
| 14 |
+
tokenizer,
|
| 15 |
+
pad_token: str = "<pad>",
|
| 16 |
+
unk_token: str = "<unk>",
|
| 17 |
+
bos_token: str = "<s>",
|
| 18 |
+
eos_token: str = "</s>",
|
| 19 |
+
):
|
| 20 |
+
self.tokenizer = tokenizer
|
| 21 |
+
self.pad_token = pad_token
|
| 22 |
+
self.unk_token = unk_token
|
| 23 |
+
self.bos_token = bos_token
|
| 24 |
+
self.eos_token = eos_token
|
| 25 |
+
|
| 26 |
+
@property
|
| 27 |
+
def pad_token_id(self) -> int:
|
| 28 |
+
return self.token_to_id(self.pad_token)
|
| 29 |
+
|
| 30 |
+
@property
|
| 31 |
+
def unk_token_id(self) -> int:
|
| 32 |
+
return self.token_to_id(self.unk_token)
|
| 33 |
+
|
| 34 |
+
@property
|
| 35 |
+
def bos_token_id(self) -> int:
|
| 36 |
+
return self.token_to_id(self.bos_token)
|
| 37 |
+
|
| 38 |
+
@property
|
| 39 |
+
def eos_token_id(self) -> int:
|
| 40 |
+
return self.token_to_id(self.eos_token)
|
| 41 |
+
|
| 42 |
+
@property
|
| 43 |
+
def vocab_size(self) -> int:
|
| 44 |
+
if hasattr(self.tokenizer, "get_vocab_size"):
|
| 45 |
+
return int(self.tokenizer.get_vocab_size())
|
| 46 |
+
return int(len(self.tokenizer))
|
| 47 |
+
|
| 48 |
+
def token_to_id(self, token: str) -> int:
|
| 49 |
+
if hasattr(self.tokenizer, "token_to_id"):
|
| 50 |
+
idx = self.tokenizer.token_to_id(token)
|
| 51 |
+
elif hasattr(self.tokenizer, "convert_tokens_to_ids"):
|
| 52 |
+
idx = self.tokenizer.convert_tokens_to_ids(token)
|
| 53 |
+
else:
|
| 54 |
+
raise TypeError("Unsupported tokenizer type")
|
| 55 |
+
|
| 56 |
+
if idx is None:
|
| 57 |
+
raise ValueError(f"Token {token!r} is not in the tokenizer vocabulary")
|
| 58 |
+
return int(idx)
|
| 59 |
+
|
| 60 |
+
def encode(self, text: str, add_special_tokens: bool = False, max_length: int | None = None) -> list[int]:
|
| 61 |
+
if hasattr(self.tokenizer, "encode") and self.tokenizer.__class__.__module__.startswith("tokenizers"):
|
| 62 |
+
ids = self.tokenizer.encode(text).ids
|
| 63 |
+
else:
|
| 64 |
+
ids = self.tokenizer.encode(text, add_special_tokens=add_special_tokens)
|
| 65 |
+
add_special_tokens = False
|
| 66 |
+
|
| 67 |
+
if add_special_tokens:
|
| 68 |
+
ids = [self.bos_token_id] + list(ids) + [self.eos_token_id]
|
| 69 |
+
|
| 70 |
+
if max_length is not None:
|
| 71 |
+
ids = list(ids)[:max_length]
|
| 72 |
+
|
| 73 |
+
return list(ids)
|
| 74 |
+
|
| 75 |
+
def decode(self, ids: Sequence[int], skip_special_tokens: bool = True) -> str:
|
| 76 |
+
if hasattr(self.tokenizer, "decode"):
|
| 77 |
+
try:
|
| 78 |
+
return self.tokenizer.decode(list(ids), skip_special_tokens=skip_special_tokens)
|
| 79 |
+
except TypeError:
|
| 80 |
+
return self.tokenizer.decode(list(ids))
|
| 81 |
+
raise TypeError("Unsupported tokenizer type")
|
| 82 |
+
|
| 83 |
+
def save(self, path: str | Path) -> None:
|
| 84 |
+
path = Path(path)
|
| 85 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 86 |
+
|
| 87 |
+
if hasattr(self.tokenizer, "save"):
|
| 88 |
+
self.tokenizer.save(str(path))
|
| 89 |
+
return
|
| 90 |
+
if hasattr(self.tokenizer, "save_pretrained"):
|
| 91 |
+
self.tokenizer.save_pretrained(str(path))
|
| 92 |
+
return
|
| 93 |
+
raise TypeError("Unsupported tokenizer type")
|
| 94 |
+
|
| 95 |
+
|
| 96 |
+
def _special_tokens(config: Mapping | None = None) -> dict[str, str]:
|
| 97 |
+
tokens = {
|
| 98 |
+
"pad": "<pad>",
|
| 99 |
+
"unk": "<unk>",
|
| 100 |
+
"bos": "<s>",
|
| 101 |
+
"eos": "</s>",
|
| 102 |
+
}
|
| 103 |
+
if config:
|
| 104 |
+
tokens.update(dict(config))
|
| 105 |
+
return tokens
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
def train_bpe_tokenizer(
|
| 109 |
+
texts: Iterable[str],
|
| 110 |
+
vocab_size: int = 32000,
|
| 111 |
+
min_frequency: int = 2,
|
| 112 |
+
special_tokens: Mapping[str, str] | None = None,
|
| 113 |
+
save_path: str | Path | None = None,
|
| 114 |
+
) -> TokenizerWrapper:
|
| 115 |
+
"""Train a byte-level BPE tokenizer on source and target training text."""
|
| 116 |
+
from tokenizers import Tokenizer
|
| 117 |
+
from tokenizers.decoders import ByteLevel as ByteLevelDecoder
|
| 118 |
+
from tokenizers.models import BPE
|
| 119 |
+
from tokenizers.normalizers import NFKC, Sequence as NormalizerSequence
|
| 120 |
+
from tokenizers.pre_tokenizers import ByteLevel
|
| 121 |
+
from tokenizers.trainers import BpeTrainer
|
| 122 |
+
|
| 123 |
+
tokens = _special_tokens(special_tokens)
|
| 124 |
+
ordered_specials = [tokens["pad"], tokens["unk"], tokens["bos"], tokens["eos"]]
|
| 125 |
+
|
| 126 |
+
tokenizer = Tokenizer(BPE(unk_token=tokens["unk"]))
|
| 127 |
+
tokenizer.normalizer = NormalizerSequence([NFKC()])
|
| 128 |
+
tokenizer.pre_tokenizer = ByteLevel(add_prefix_space=False)
|
| 129 |
+
tokenizer.decoder = ByteLevelDecoder()
|
| 130 |
+
|
| 131 |
+
trainer = BpeTrainer(
|
| 132 |
+
vocab_size=vocab_size,
|
| 133 |
+
min_frequency=min_frequency,
|
| 134 |
+
special_tokens=ordered_specials,
|
| 135 |
+
show_progress=True,
|
| 136 |
+
)
|
| 137 |
+
tokenizer.train_from_iterator((text for text in texts if text), trainer=trainer)
|
| 138 |
+
|
| 139 |
+
wrapper = TokenizerWrapper(
|
| 140 |
+
tokenizer,
|
| 141 |
+
pad_token=tokens["pad"],
|
| 142 |
+
unk_token=tokens["unk"],
|
| 143 |
+
bos_token=tokens["bos"],
|
| 144 |
+
eos_token=tokens["eos"],
|
| 145 |
+
)
|
| 146 |
+
|
| 147 |
+
if save_path is not None:
|
| 148 |
+
wrapper.save(save_path)
|
| 149 |
+
|
| 150 |
+
return wrapper
|
| 151 |
+
|
| 152 |
+
|
| 153 |
+
def build_tokenizer(config: Mapping, train_texts: Iterable[str] | None = None) -> TokenizerWrapper:
|
| 154 |
+
"""Build a tokenizer from project config."""
|
| 155 |
+
tokenizer_type = config.get("type", "bpe")
|
| 156 |
+
tokens = _special_tokens(config.get("special_tokens"))
|
| 157 |
+
|
| 158 |
+
if tokenizer_type == "pretrained":
|
| 159 |
+
from transformers import AutoTokenizer
|
| 160 |
+
|
| 161 |
+
model_name = config.get("model_name") or config.get("pretrained_model_name")
|
| 162 |
+
if not model_name:
|
| 163 |
+
raise ValueError("pretrained tokenizer requires config['model_name']")
|
| 164 |
+
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
| 165 |
+
return TokenizerWrapper(
|
| 166 |
+
tokenizer,
|
| 167 |
+
pad_token=tokenizer.pad_token or tokens["pad"],
|
| 168 |
+
unk_token=tokenizer.unk_token or tokens["unk"],
|
| 169 |
+
bos_token=tokenizer.bos_token or tokens["bos"],
|
| 170 |
+
eos_token=tokenizer.eos_token or tokens["eos"],
|
| 171 |
+
)
|
| 172 |
+
|
| 173 |
+
if tokenizer_type in {"bpe", "sentencepiece"}:
|
| 174 |
+
tokenizer_path = config.get("path") or config.get("tokenizer_path")
|
| 175 |
+
if tokenizer_path and Path(tokenizer_path).exists():
|
| 176 |
+
from tokenizers import Tokenizer
|
| 177 |
+
|
| 178 |
+
return TokenizerWrapper(
|
| 179 |
+
Tokenizer.from_file(str(tokenizer_path)),
|
| 180 |
+
pad_token=tokens["pad"],
|
| 181 |
+
unk_token=tokens["unk"],
|
| 182 |
+
bos_token=tokens["bos"],
|
| 183 |
+
eos_token=tokens["eos"],
|
| 184 |
+
)
|
| 185 |
+
|
| 186 |
+
if train_texts is None:
|
| 187 |
+
raise ValueError("BPE tokenizer requires train_texts when no tokenizer path is provided")
|
| 188 |
+
|
| 189 |
+
return train_bpe_tokenizer(
|
| 190 |
+
train_texts,
|
| 191 |
+
vocab_size=int(config.get("vocab_size", 32000)),
|
| 192 |
+
min_frequency=int(config.get("min_frequency", 2)),
|
| 193 |
+
special_tokens=tokens,
|
| 194 |
+
save_path=tokenizer_path,
|
| 195 |
+
)
|
| 196 |
+
|
| 197 |
+
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
|
tests/test_data.py
CHANGED
|
@@ -5,43 +5,83 @@
|
|
| 5 |
import pytest
|
| 6 |
import torch
|
| 7 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 8 |
|
| 9 |
class TestTranslationDataset:
|
| 10 |
"""测试 TranslationDataset 类。"""
|
| 11 |
|
| 12 |
-
def test_dataset_length(self):
|
| 13 |
"""测试数据集长度。"""
|
| 14 |
-
|
| 15 |
-
|
| 16 |
|
| 17 |
-
def test_getitem_returns_correct_keys(self):
|
| 18 |
"""测试 __getitem__ 返回正确的字段。"""
|
| 19 |
-
|
| 20 |
-
|
|
|
|
| 21 |
|
| 22 |
-
def test_getitem_tensor_types(self):
|
| 23 |
"""测试返回的 tensor 类型正确。"""
|
| 24 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 25 |
|
| 26 |
def test_src_tgt_mismatch_raises(self):
|
| 27 |
"""测试源目标数量不匹配时抛出异常。"""
|
| 28 |
-
|
|
|
|
| 29 |
|
| 30 |
|
| 31 |
class TestTokenizer:
|
| 32 |
"""测试分词器。"""
|
| 33 |
|
| 34 |
-
def test_bpe_train_and_encode(self):
|
| 35 |
"""测试 BPE 训练和编码。"""
|
| 36 |
-
|
|
|
|
|
|
|
|
|
|
| 37 |
|
| 38 |
-
def test_encode_decode_roundtrip(self):
|
| 39 |
"""测试编码-解码往返一致性。"""
|
| 40 |
-
|
|
|
|
|
|
|
|
|
|
| 41 |
|
| 42 |
-
def test_special_tokens(self):
|
| 43 |
"""测试特殊 token 正确。"""
|
| 44 |
-
|
|
|
|
|
|
|
|
|
|
| 45 |
|
| 46 |
|
| 47 |
class TestPreprocessing:
|
|
@@ -49,24 +89,55 @@ class TestPreprocessing:
|
|
| 49 |
|
| 50 |
def test_clean_text_unicode(self):
|
| 51 |
"""测试 Unicode 标准化。"""
|
| 52 |
-
|
| 53 |
|
| 54 |
def test_filter_by_length(self):
|
| 55 |
"""测试按长度过滤。"""
|
| 56 |
-
|
|
|
|
| 57 |
|
| 58 |
def test_deduplicate(self):
|
| 59 |
"""测试去重。"""
|
| 60 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 61 |
|
| 62 |
|
| 63 |
class TestCollator:
|
| 64 |
"""测试数据整理器。"""
|
| 65 |
|
| 66 |
-
def test_padding(self):
|
| 67 |
"""测试 padding 正确。"""
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 71 |
"""测试 attention mask 正确。"""
|
| 72 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
import pytest
|
| 6 |
import torch
|
| 7 |
|
| 8 |
+
from easytranslate.data import (
|
| 9 |
+
DynamicBatchSampler,
|
| 10 |
+
TokenizerWrapper,
|
| 11 |
+
TranslationCollator,
|
| 12 |
+
TranslationDataset,
|
| 13 |
+
clean_text,
|
| 14 |
+
deduplicate_pairs,
|
| 15 |
+
filter_by_length,
|
| 16 |
+
preprocess_pipeline,
|
| 17 |
+
train_bpe_tokenizer,
|
| 18 |
+
)
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
@pytest.fixture()
|
| 22 |
+
def tiny_tokenizer() -> TokenizerWrapper:
|
| 23 |
+
texts = [
|
| 24 |
+
"Hello world",
|
| 25 |
+
"Machine translation is useful",
|
| 26 |
+
"I love natural language processing",
|
| 27 |
+
"你好 世界",
|
| 28 |
+
"机器翻译 很 有用",
|
| 29 |
+
"我 喜欢 自然语言处理",
|
| 30 |
+
]
|
| 31 |
+
return train_bpe_tokenizer(texts, vocab_size=80, min_frequency=1)
|
| 32 |
+
|
| 33 |
|
| 34 |
class TestTranslationDataset:
|
| 35 |
"""测试 TranslationDataset 类。"""
|
| 36 |
|
| 37 |
+
def test_dataset_length(self, tiny_tokenizer):
|
| 38 |
"""测试数据集长度。"""
|
| 39 |
+
dataset = TranslationDataset(["Hello world"], ["你好 世界"], tokenizer=tiny_tokenizer)
|
| 40 |
+
assert len(dataset) == 1
|
| 41 |
|
| 42 |
+
def test_getitem_returns_correct_keys(self, tiny_tokenizer):
|
| 43 |
"""测试 __getitem__ 返回正确的字段。"""
|
| 44 |
+
dataset = TranslationDataset(["Hello world"], ["你好 世界"], tokenizer=tiny_tokenizer)
|
| 45 |
+
item = dataset[0]
|
| 46 |
+
assert set(item) == {"src_ids", "tgt_input_ids", "labels", "src_len", "tgt_len"}
|
| 47 |
|
| 48 |
+
def test_getitem_tensor_types(self, tiny_tokenizer):
|
| 49 |
"""测试返回的 tensor 类型正确。"""
|
| 50 |
+
dataset = TranslationDataset(["Hello world"], ["你好 世界"], tokenizer=tiny_tokenizer)
|
| 51 |
+
item = dataset[0]
|
| 52 |
+
assert item["src_ids"].dtype == torch.long
|
| 53 |
+
assert item["tgt_input_ids"].dtype == torch.long
|
| 54 |
+
assert item["labels"].dtype == torch.long
|
| 55 |
|
| 56 |
def test_src_tgt_mismatch_raises(self):
|
| 57 |
"""测试源目标数量不匹配时抛出异常。"""
|
| 58 |
+
with pytest.raises(ValueError):
|
| 59 |
+
TranslationDataset(["a", "b"], ["甲"], tokenizer=object())
|
| 60 |
|
| 61 |
|
| 62 |
class TestTokenizer:
|
| 63 |
"""测试分词器。"""
|
| 64 |
|
| 65 |
+
def test_bpe_train_and_encode(self, tiny_tokenizer):
|
| 66 |
"""测试 BPE 训练和编码。"""
|
| 67 |
+
ids = tiny_tokenizer.encode("Hello world", add_special_tokens=True)
|
| 68 |
+
assert len(ids) >= 3
|
| 69 |
+
assert ids[0] == tiny_tokenizer.bos_token_id
|
| 70 |
+
assert ids[-1] == tiny_tokenizer.eos_token_id
|
| 71 |
|
| 72 |
+
def test_encode_decode_roundtrip(self, tiny_tokenizer):
|
| 73 |
"""测试编码-解码往返一致性。"""
|
| 74 |
+
ids = tiny_tokenizer.encode("Hello world")
|
| 75 |
+
decoded = tiny_tokenizer.decode(ids)
|
| 76 |
+
assert "Hello" in decoded
|
| 77 |
+
assert "world" in decoded
|
| 78 |
|
| 79 |
+
def test_special_tokens(self, tiny_tokenizer):
|
| 80 |
"""测试特殊 token 正确。"""
|
| 81 |
+
assert tiny_tokenizer.pad_token_id == 0
|
| 82 |
+
assert tiny_tokenizer.unk_token_id == 1
|
| 83 |
+
assert tiny_tokenizer.bos_token_id == 2
|
| 84 |
+
assert tiny_tokenizer.eos_token_id == 3
|
| 85 |
|
| 86 |
|
| 87 |
class TestPreprocessing:
|
|
|
|
| 89 |
|
| 90 |
def test_clean_text_unicode(self):
|
| 91 |
"""测试 Unicode 标准化。"""
|
| 92 |
+
assert clean_text("ABC\u200b 123") == "ABC 123"
|
| 93 |
|
| 94 |
def test_filter_by_length(self):
|
| 95 |
"""测试按长度过滤。"""
|
| 96 |
+
assert filter_by_length("hello world", "你好世界", max_src_len=10, max_tgt_len=10)
|
| 97 |
+
assert not filter_by_length("hello " * 300, "你好", max_src_len=256)
|
| 98 |
|
| 99 |
def test_deduplicate(self):
|
| 100 |
"""测试去重。"""
|
| 101 |
+
pairs = deduplicate_pairs([("a", "甲"), ("a", "甲"), ("b", "乙")])
|
| 102 |
+
assert pairs == [("a", "甲"), ("b", "乙")]
|
| 103 |
+
|
| 104 |
+
def test_preprocess_pipeline(self):
|
| 105 |
+
"""测试完整预处理流水线。"""
|
| 106 |
+
src, tgt = preprocess_pipeline([" Hello world ", "", "Hello world"], [" 你好 世界 ", "空", "你好 世界"])
|
| 107 |
+
assert src == ["Hello world"]
|
| 108 |
+
assert tgt == ["你好 世界"]
|
| 109 |
|
| 110 |
|
| 111 |
class TestCollator:
|
| 112 |
"""测试数据整理器。"""
|
| 113 |
|
| 114 |
+
def test_padding(self, tiny_tokenizer):
|
| 115 |
"""测试 padding 正确。"""
|
| 116 |
+
dataset = TranslationDataset(
|
| 117 |
+
["Hello world", "Machine translation is useful"],
|
| 118 |
+
["你好 世界", "机器翻译 很 有用"],
|
| 119 |
+
tokenizer=tiny_tokenizer,
|
| 120 |
+
)
|
| 121 |
+
batch = TranslationCollator(pad_token_id=tiny_tokenizer.pad_token_id)([dataset[0], dataset[1]])
|
| 122 |
+
assert batch["src_ids"].ndim == 2
|
| 123 |
+
assert batch["tgt_input_ids"].ndim == 2
|
| 124 |
+
assert batch["labels"].shape == batch["tgt_input_ids"].shape
|
| 125 |
+
|
| 126 |
+
def test_attention_mask(self, tiny_tokenizer):
|
| 127 |
"""测试 attention mask 正确。"""
|
| 128 |
+
dataset = TranslationDataset(
|
| 129 |
+
["Hello", "Machine translation is useful"],
|
| 130 |
+
["你好", "机器翻译 很 有用"],
|
| 131 |
+
tokenizer=tiny_tokenizer,
|
| 132 |
+
)
|
| 133 |
+
batch = TranslationCollator(pad_token_id=tiny_tokenizer.pad_token_id)([dataset[0], dataset[1]])
|
| 134 |
+
assert batch["src_padding_mask"].dtype == torch.bool
|
| 135 |
+
assert torch.equal(batch["src_attention_mask"], (~batch["src_padding_mask"]).long())
|
| 136 |
+
|
| 137 |
+
def test_dynamic_batch_sampler(self):
|
| 138 |
+
"""测试动态 batch 不超过 token 预算。"""
|
| 139 |
+
sampler = DynamicBatchSampler([5, 6, 20, 21], max_tokens_per_batch=24, shuffle=False)
|
| 140 |
+
batches = list(sampler)
|
| 141 |
+
assert batches
|
| 142 |
+
for batch in batches:
|
| 143 |
+
assert max([5, 6, 20, 21][idx] for idx in batch) * len(batch) <= 24
|