from torch_dataset_loader import JsonlShardDataset, StreamConfig, collate_dynamic_padding
loader = create_loader(args.shards_index, args.batch_size, args.num_workers, shuffle_buffer=10000)

def create_loader(shards_index: str, batch_size: int, num_workers: int, shuffle_buffer: int):
    ds = JsonlShardDataset(StreamConfig(
        shards_index=shards_index,
        shuffle_shards=True,
        shuffle_records_buffer=shuffle_buffer,
        seed=42,
        repeat=1,
    ))
    loader = torch.utils.data.DataLoader(
        ds,
        batch_size=batch_size,
        num_workers=num_workers,
        collate_fn=collate_dynamic_padding,
        pin_memory=True,
        prefetch_factor=2 if num_workers > 0 else None,
        persistent_workers=True if num_workers > 0 else False,
    )
    return loader
