MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedTTS2 / create_dataloaders

Function create_dataloaders

bin/finetune_example/posttrain_dataloader.py:288–319  ·  view source on GitHub ↗
(
    train_datasets: str,
    validation_datasets: str,
    batch_size: int,
    device,  # for debug
    infinite_train: bool = False,
    num_workers: int = 0,
)

Source from the content-addressed store, hash-verified

286
287
288def create_dataloaders(
289 train_datasets: str,
290 validation_datasets: str,
291 batch_size: int,
292 device, # for debug
293 infinite_train: bool = False,
294 num_workers: int = 0,
295):
296
297 trainset = TokenizedDataset(dataset_dir=train_datasets, device=device)
298 valset = TokenizedDataset(dataset_dir=validation_datasets, device=device)
299 trainloader = DataLoader(
300 trainset,
301 # batch_sampler=trainsampler,
302 batch_size=batch_size,
303 collate_fn=collate_fn,
304 num_workers=12,
305 pin_memory=True,
306 shuffle=True,
307 )
308
309 valloader = DataLoader(
310 valset,
311 # batch_sampler=valsampler,
312 batch_size=batch_size,
313 collate_fn=collate_fn,
314 num_workers=num_workers,
315 pin_memory=True,
316 shuffle=False,
317 )
318
319 return trainloader, valloader

Callers 1

trainFunction · 0.90

Calls 1

TokenizedDatasetClass · 0.85

Tested by

no test coverage detected