UCAS-EasyTranslate / tests /test_data.py
xiexinyuan341's picture
Implement data pipeline module
a2d6c00 verified
Raw
History Blame
5.35 kB
"""
单元测试 — 数据模块 (Person A 实现后需通过)
"""
import pytest
import torch
from easytranslate.data import (
DynamicBatchSampler,
TokenizerWrapper,
TranslationCollator,
TranslationDataset,
clean_text,
deduplicate_pairs,
filter_by_length,
preprocess_pipeline,
train_bpe_tokenizer,
)
@pytest.fixture()
def tiny_tokenizer() -> TokenizerWrapper:
texts = [
"Hello world",
"Machine translation is useful",
"I love natural language processing",
"你好 世界",
"机器翻译 很 有用",
"我 喜欢 自然语言处理",
]
return train_bpe_tokenizer(texts, vocab_size=80, min_frequency=1)
class TestTranslationDataset:
"""测试 TranslationDataset 类。"""
def test_dataset_length(self, tiny_tokenizer):
"""测试数据集长度。"""
dataset = TranslationDataset(["Hello world"], ["你好 世界"], tokenizer=tiny_tokenizer)
assert len(dataset) == 1
def test_getitem_returns_correct_keys(self, tiny_tokenizer):
"""测试 __getitem__ 返回正确的字段。"""
dataset = TranslationDataset(["Hello world"], ["你好 世界"], tokenizer=tiny_tokenizer)
item = dataset[0]
assert set(item) == {"src_ids", "tgt_input_ids", "labels", "src_len", "tgt_len"}
def test_getitem_tensor_types(self, tiny_tokenizer):
"""测试返回的 tensor 类型正确。"""
dataset = TranslationDataset(["Hello world"], ["你好 世界"], tokenizer=tiny_tokenizer)
item = dataset[0]
assert item["src_ids"].dtype == torch.long
assert item["tgt_input_ids"].dtype == torch.long
assert item["labels"].dtype == torch.long
def test_src_tgt_mismatch_raises(self):
"""测试源目标数量不匹配时抛出异常。"""
with pytest.raises(ValueError):
TranslationDataset(["a", "b"], ["甲"], tokenizer=object())
class TestTokenizer:
"""测试分词器。"""
def test_bpe_train_and_encode(self, tiny_tokenizer):
"""测试 BPE 训练和编码。"""
ids = tiny_tokenizer.encode("Hello world", add_special_tokens=True)
assert len(ids) >= 3
assert ids[0] == tiny_tokenizer.bos_token_id
assert ids[-1] == tiny_tokenizer.eos_token_id
def test_encode_decode_roundtrip(self, tiny_tokenizer):
"""测试编码-解码往返一致性。"""
ids = tiny_tokenizer.encode("Hello world")
decoded = tiny_tokenizer.decode(ids)
assert "Hello" in decoded
assert "world" in decoded
def test_special_tokens(self, tiny_tokenizer):
"""测试特殊 token 正确。"""
assert tiny_tokenizer.pad_token_id == 0
assert tiny_tokenizer.unk_token_id == 1
assert tiny_tokenizer.bos_token_id == 2
assert tiny_tokenizer.eos_token_id == 3
class TestPreprocessing:
"""测试预处理。"""
def test_clean_text_unicode(self):
"""测试 Unicode 标准化。"""
assert clean_text("ABC\u200b 123") == "ABC 123"
def test_filter_by_length(self):
"""测试按长度过滤。"""
assert filter_by_length("hello world", "你好世界", max_src_len=10, max_tgt_len=10)
assert not filter_by_length("hello " * 300, "你好", max_src_len=256)
def test_deduplicate(self):
"""测试去重。"""
pairs = deduplicate_pairs([("a", "甲"), ("a", "甲"), ("b", "乙")])
assert pairs == [("a", "甲"), ("b", "乙")]
def test_preprocess_pipeline(self):
"""测试完整预处理流水线。"""
src, tgt = preprocess_pipeline([" Hello world ", "", "Hello world"], [" 你好 世界 ", "空", "你好 世界"])
assert src == ["Hello world"]
assert tgt == ["你好 世界"]
class TestCollator:
"""测试数据整理器。"""
def test_padding(self, tiny_tokenizer):
"""测试 padding 正确。"""
dataset = TranslationDataset(
["Hello world", "Machine translation is useful"],
["你好 世界", "机器翻译 很 有用"],
tokenizer=tiny_tokenizer,
)
batch = TranslationCollator(pad_token_id=tiny_tokenizer.pad_token_id)([dataset[0], dataset[1]])
assert batch["src_ids"].ndim == 2
assert batch["tgt_input_ids"].ndim == 2
assert batch["labels"].shape == batch["tgt_input_ids"].shape
def test_attention_mask(self, tiny_tokenizer):
"""测试 attention mask 正确。"""
dataset = TranslationDataset(
["Hello", "Machine translation is useful"],
["你好", "机器翻译 很 有用"],
tokenizer=tiny_tokenizer,
)
batch = TranslationCollator(pad_token_id=tiny_tokenizer.pad_token_id)([dataset[0], dataset[1]])
assert batch["src_padding_mask"].dtype == torch.bool
assert torch.equal(batch["src_attention_mask"], (~batch["src_padding_mask"]).long())
def test_dynamic_batch_sampler(self):
"""测试动态 batch 不超过 token 预算。"""
sampler = DynamicBatchSampler([5, 6, 20, 21], max_tokens_per_batch=24, shuffle=False)
batches = list(sampler)
assert batches
for batch in batches:
assert max([5, 6, 20, 21][idx] for idx in batch) * len(batch) <= 24