Download tests/test_data.py from kiruluta/rsil-benchmark: direct link, hf CLI and curl.
- Browser
- Download file 1.8 kB
-
https://huggingface.co/kiruluta/rsil-benchmark/resolve/main/tests/test_data.py
- Command line
-
hf download hf://kiruluta/rsil-benchmark/tests/test_data.py
-
curl -L -o test_data.py https://huggingface.co/kiruluta/rsil-benchmark/resolve/main/tests/test_data.py
1.8 kB
| import torch | |
| from rsil.data import PackedTokenStream, DeterministicMLMCollator | |
| class TinyTok: | |
| mask_token_id=1; sep_token_id=2 | |
| def __len__(self): return 64 | |
| def __call__(self,text,add_special_tokens=False,truncation=False): return {"input_ids":[3+(ord(c)%40) for c in text if not c.isspace()]} | |
| def get_special_tokens_mask(self, ids, already_has_special_tokens=True): return [1 if i in (1,2) else 0 for i in ids] | |
| def test_packed_real_text_pipeline_has_no_strings_and_fixed_length(): | |
| src=[{"text":"alpha beta gamma"},{"text":"delta epsilon zeta"},{"text":"eta theta iota"}] | |
| blocks=list(PackedTokenStream(src,TinyTok(),8)) | |
| assert blocks and all(b["input_ids"].shape==(8,) for b in blocks) | |
| assert all(torch.is_tensor(b["input_ids"]) for b in blocks) | |
| batch=DeterministicMLMCollator(TinyTok(),seed=7)(blocks[:2]) | |
| assert batch["input_ids"].dtype==torch.long and batch["labels"].shape==batch["input_ids"].shape | |
| assert (batch["labels"]!=-100).any() | |
| class _BackendEncoding: | |
| def __init__(self, ids): self.ids = ids | |
| class _LongSafeBackend: | |
| def encode(self, text, add_special_tokens=False): | |
| # Deliberately return >512 tokens from one source paragraph. | |
| return _BackendEncoding(list(range(600))) | |
| class _BackendOnlyTokenizer: | |
| sep_token_id = None | |
| backend_tokenizer = _LongSafeBackend() | |
| def __call__(self, *args, **kwargs): | |
| raise AssertionError("wrapper tokenizer must not be called for fast backend path") | |
| def test_long_source_uses_backend_and_packs_exact_blocks(): | |
| from rsil.data import PackedTokenStream | |
| src = [{"text": "long source"}] | |
| blocks = list(PackedTokenStream(src, _BackendOnlyTokenizer(), seq_len=256)) | |
| assert len(blocks) == 2 | |
| assert all(tuple(b["input_ids"].shape) == (256,) for b in blocks) | |