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