rsil-benchmark / tests /test_data.py
kiruluta's picture
Upload folder using huggingface_hub
f048438 verified
Raw History Blame Contribute Delete
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)