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