| """ |
| 单元测试 — 数据模块 (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 |
|
|