sdfjliom xiexinyuan341 commited on
Commit
6b73a07
·
1 Parent(s): 4b16d58

Implement data pipeline module (#2)

Browse files

- Implement data pipeline module (a2d6c00f06d350e44f367ae2a14caee908f11454)


Co-authored-by: xinyuanxie <xiexinyuan341@users.noreply.huggingface.co>

.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
- # TODO: 创建 mock tokenizer,构建小型数据集,验证 len() 正确
15
- pass
16
 
17
- def test_getitem_returns_correct_keys(self):
18
  """测试 __getitem__ 返回正确的字段。"""
19
- # TODO: 验证返回 dict 包含 src_ids, tgt_input_ids, labels, src_len, tgt_len
20
- pass
 
21
 
22
- def test_getitem_tensor_types(self):
23
  """测试返回的 tensor 类型正确。"""
24
- pass
 
 
 
 
25
 
26
  def test_src_tgt_mismatch_raises(self):
27
  """测试源目标数量不匹配时抛出异常。"""
28
- pass
 
29
 
30
 
31
  class TestTokenizer:
32
  """测试分词器。"""
33
 
34
- def test_bpe_train_and_encode(self):
35
  """测试 BPE 训练和编码。"""
36
- pass
 
 
 
37
 
38
- def test_encode_decode_roundtrip(self):
39
  """测试编码-解码往返一致性。"""
40
- pass
 
 
 
41
 
42
- def test_special_tokens(self):
43
  """测试特殊 token 正确。"""
44
- pass
 
 
 
45
 
46
 
47
  class TestPreprocessing:
@@ -49,24 +89,55 @@ class TestPreprocessing:
49
 
50
  def test_clean_text_unicode(self):
51
  """测试 Unicode 标准化。"""
52
- pass
53
 
54
  def test_filter_by_length(self):
55
  """测试按长度过滤。"""
56
- pass
 
57
 
58
  def test_deduplicate(self):
59
  """测试去重。"""
60
- pass
 
 
 
 
 
 
 
61
 
62
 
63
  class TestCollator:
64
  """测试数据整理器。"""
65
 
66
- def test_padding(self):
67
  """测试 padding 正确。"""
68
- pass
69
-
70
- def test_attention_mask(self):
 
 
 
 
 
 
 
 
71
  """测试 attention mask 正确。"""
72
- pass
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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