| """Dataset loading and PyTorch dataset classes for English-Chinese translation.""" |
|
|
| from __future__ import annotations |
|
|
| from pathlib import Path |
| from typing import Mapping, Sequence |
|
|
| import torch |
| from torch.utils.data import Dataset |
|
|
| from easytranslate.data.preprocessing import preprocess_pipeline |
| from easytranslate.data.tokenizer import TokenizerWrapper |
|
|
|
|
| class TranslationDataset(Dataset): |
| """PyTorch dataset that builds seq2seq inputs for teacher forcing.""" |
|
|
| def __init__( |
| self, |
| src_texts: Sequence[str], |
| tgt_texts: Sequence[str], |
| tokenizer: TokenizerWrapper | None = None, |
| src_tokenizer: TokenizerWrapper | None = None, |
| tgt_tokenizer: TokenizerWrapper | None = None, |
| max_src_len: int = 256, |
| max_tgt_len: int = 256, |
| ): |
| if len(src_texts) != len(tgt_texts): |
| raise ValueError("src_texts and tgt_texts must have the same length") |
|
|
| if tokenizer is not None: |
| src_tokenizer = src_tokenizer or tokenizer |
| tgt_tokenizer = tgt_tokenizer or tokenizer |
| if src_tokenizer is None or tgt_tokenizer is None: |
| raise ValueError("Provide tokenizer or both src_tokenizer and tgt_tokenizer") |
|
|
| self.src_texts = list(src_texts) |
| self.tgt_texts = list(tgt_texts) |
| self.src_tokenizer = src_tokenizer |
| self.tgt_tokenizer = tgt_tokenizer |
| self.max_src_len = max_src_len |
| self.max_tgt_len = max_tgt_len |
|
|
| def __len__(self) -> int: |
| return len(self.src_texts) |
|
|
| def __getitem__(self, index: int) -> dict[str, torch.Tensor]: |
| src_ids = self.src_tokenizer.encode( |
| self.src_texts[index], |
| add_special_tokens=True, |
| max_length=self.max_src_len, |
| ) |
|
|
| target_core = self.tgt_tokenizer.encode( |
| self.tgt_texts[index], |
| add_special_tokens=False, |
| max_length=max(1, self.max_tgt_len - 1), |
| ) |
| tgt_input_ids = [self.tgt_tokenizer.bos_token_id] + target_core |
| labels = target_core + [self.tgt_tokenizer.eos_token_id] |
|
|
| return { |
| "src_ids": torch.tensor(src_ids, dtype=torch.long), |
| "tgt_input_ids": torch.tensor(tgt_input_ids, dtype=torch.long), |
| "labels": torch.tensor(labels, dtype=torch.long), |
| "src_len": torch.tensor(len(src_ids), dtype=torch.long), |
| "tgt_len": torch.tensor(len(labels), dtype=torch.long), |
| } |
|
|
|
|
| def _translation_to_columns(dataset, src_lang: str, tgt_lang: str): |
| def convert(example): |
| translation = example["translation"] |
| return {"src": translation[src_lang], "tgt": translation[tgt_lang]} |
|
|
| return dataset.map(convert, remove_columns=dataset.column_names) |
|
|
|
|
| def load_wmt_dataset( |
| year: str = "19", |
| language_pair: str = "zh-en", |
| src_lang: str = "en", |
| tgt_lang: str = "zh", |
| split: str | None = None, |
| cache_dir: str | None = None, |
| ): |
| """Load WMT zh-en and normalize rows to {'src', 'tgt'}.""" |
| from datasets import load_dataset |
|
|
| dataset_name = f"wmt/wmt{year}" |
| try: |
| dataset = load_dataset(dataset_name, language_pair, split=split, cache_dir=cache_dir) |
| except Exception: |
| dataset = load_dataset(f"wmt{year}", language_pair, split=split, cache_dir=cache_dir) |
|
|
| if split is not None: |
| return _translation_to_columns(dataset, src_lang, tgt_lang) |
|
|
| return dataset.map( |
| lambda example: {"src": example["translation"][src_lang], "tgt": example["translation"][tgt_lang]}, |
| remove_columns=next(iter(dataset.values())).column_names, |
| ) |
|
|
|
|
| def load_opus_dataset( |
| subset: str = "en-zh", |
| src_lang: str = "en", |
| tgt_lang: str = "zh", |
| split: str | None = None, |
| cache_dir: str | None = None, |
| ): |
| """Load OPUS-100 and normalize rows to {'src', 'tgt'}.""" |
| from datasets import load_dataset |
|
|
| dataset = load_dataset("Helsinki-NLP/opus-100", subset, split=split, cache_dir=cache_dir) |
| if split is not None: |
| return _translation_to_columns(dataset, src_lang, tgt_lang) |
|
|
| return dataset.map( |
| lambda example: {"src": example["translation"][src_lang], "tgt": example["translation"][tgt_lang]}, |
| remove_columns=next(iter(dataset.values())).column_names, |
| ) |
|
|
|
|
| def _read_lines(path: str | Path) -> list[str]: |
| with Path(path).open("r", encoding="utf-8") as f: |
| return [line.rstrip("\n") for line in f] |
|
|
|
|
| def load_custom_dataset( |
| train_src: str | Path, |
| train_tgt: str | Path, |
| val_src: str | Path | None = None, |
| val_tgt: str | Path | None = None, |
| test_src: str | Path | None = None, |
| test_tgt: str | Path | None = None, |
| preprocessing_config: Mapping | None = None, |
| ) -> dict[str, dict[str, list[str]]]: |
| """Load parallel text files and return split dictionaries.""" |
| preprocessing_config = dict(preprocessing_config or {}) |
|
|
| def load_split(src_path: str | Path, tgt_path: str | Path) -> dict[str, list[str]]: |
| src_texts, tgt_texts = preprocess_pipeline( |
| _read_lines(src_path), |
| _read_lines(tgt_path), |
| **preprocessing_config, |
| ) |
| return {"src": src_texts, "tgt": tgt_texts} |
|
|
| result = {"train": load_split(train_src, train_tgt)} |
|
|
| if val_src and val_tgt: |
| result["validation"] = load_split(val_src, val_tgt) |
| if test_src and test_tgt: |
| result["test"] = load_split(test_src, test_tgt) |
|
|
| return result |
|
|