import hydra from hydra.utils import get_method from pathlib import Path from lmr.data.disk_dataset import DiskDataset from lmr.data.proportional_dataset import ProportionalDataset from lmr.config import initialize_config def initialize_dataset(dataset_config, dataset_dir): get_method(dataset_config.init_fn)(dataset_config, dataset_dir) def get_dataset_splits(dataset_config, max_seq_len, dataset_dir, split_names=('train', 'validation', 'test')): output_dir = dataset_dir / dataset_config.dataset_name # print(output_dir) # Initialize dataset if folder missing if not output_dir.exists() or dataset_config.component_whitelist is not None: initialize_dataset(dataset_config, dataset_dir) splits = {} # Normalize dataset name for checks ds_name = str(dataset_config.dataset_name).lower() for split_name in split_names: # Default: only use sliding window for training if enabled in config stride_fraction = 0.5 if getattr(dataset_config, "use_sliding_window", False) and split_name == "train" else None # FORCE: tinygsm should never use sliding window if ds_name == "tinygsm": stride_fraction = -1 if dataset_config.sampling_type == 'proportional': components = {} for component_name in dataset_config.proportions.keys(): file_path = output_dir / component_name / f"{split_name}.bin" component_dataset = DiskDataset(file_path, max_seq_len, stride_fraction=stride_fraction, allow_cycling=True) components[component_name] = component_dataset splits[split_name] = ProportionalDataset(components, dataset_config.proportions) else: file_path = output_dir / f"{split_name}.bin" splits[split_name] = DiskDataset(file_path, max_seq_len, stride_fraction=stride_fraction, allow_cycling=False) return splits