File size: 1,935 Bytes
3b2d368
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
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