# lmr/data/streaming.py import torch from torch.utils.data import IterableDataset from datasets import load_dataset import math import os def _get_ddp_info(): import torch.distributed as dist if dist.is_available() and dist.is_initialized(): return dist.get_world_size(), dist.get_rank() return 1, 0 class StreamingLMIterableDataset(IterableDataset): """ Wrap HF streaming dataset into an IterableDataset that yields fixed-length chunks (seq_len). Yields dicts: {"input_ids": LongTensor(seq_len), "labels": LongTensor(seq_len)} Args: hf_repo_or_base: repository / base name (HF_BASE) - passed through to load_dataset calls data_files: str or pattern for the dataset split split: 'train' / 'validation' / 'test' tokenizer: tokenizer instance with .encode(text) -> List[int] seq_len: tokens per example (model context length) max_tokens: optional total token cap (useful to bound epoch length) sample_key: which field contains text, default 'text' use_shard: if True, attempt to shard HF streams across DDP ranks """ def __init__(self, hf_base, data_files_pattern, split, tokenizer, seq_len=2048, max_tokens=None, sample_key="text", use_shard=True, streaming_kwargs=None): self.hf_base = hf_base self.data_files_pattern = data_files_pattern self.split = split self.tokenizer = tokenizer self.seq_len = seq_len self.max_tokens = max_tokens self.sample_key = sample_key self.use_shard = use_shard self.streaming_kwargs = streaming_kwargs or {} def _make_stream(self): # note: data_files can be a glob or explicit file list ds = load_dataset( self.hf_base, data_files={self.split: self.data_files_pattern}, split=self.split, streaming=True, **self.streaming_kwargs ) # DDP sharding via datasets.shard world_size, rank = _get_ddp_info() if self.use_shard and world_size > 1: try: ds = ds.shard(num_shards=world_size, index=rank) except Exception: # fallback: some streaming backends might not support shard pass return ds def __iter__(self): ds_stream = self._make_stream() buffer = [] total_tokens = 0 for example in ds_stream: # accept dict-like, access text field text = None if isinstance(example, dict): text = example.get(self.sample_key) or example.get("content") or example.get("text") else: text = str(example) if text is None: continue tok = self.tokenizer.encode(text) if not tok: continue buffer.extend(tok) # emit chunks of seq_len while len(buffer) >= self.seq_len: chunk = buffer[:self.seq_len] del buffer[:self.seq_len] total_tokens += self.seq_len if self.max_tokens and total_tokens > self.max_tokens: return input_ids = torch.tensor(chunk, dtype=torch.long) yield {"input_ids": input_ids, "labels": input_ids.clone()} # optional: discard remainder or pad (we choose to discard remainder) return