| """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 | |