FST_code / src /lmr /data /dataset_loader.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
1.94 kB
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