xiexinyuan341's picture
Implement data pipeline module
a2d6c00 verified
Raw
History Blame
5.43 kB
"""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