MCPcopy Create free account
hub / github.com/ace-step/ACE-Step / train_dataloader

Method train_dataloader

trainer.py:445–456  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

443 return [optimizer], [{"scheduler": lr_scheduler, "interval": "step"}]
444
445 def train_dataloader(self):
446 self.train_dataset = Text2MusicDataset(
447 train=True,
448 train_dataset_path=self.hparams.dataset_path,
449 )
450 return DataLoader(
451 self.train_dataset,
452 shuffle=True,
453 num_workers=self.hparams.num_workers,
454 pin_memory=True,
455 collate_fn=self.train_dataset.collate_fn,
456 )
457
458 def get_sd3_sigmas(self, timesteps, device, n_dim=4, dtype=torch.float32):
459 sigmas = self.scheduler.sigmas.to(device=device, dtype=dtype)

Callers

nothing calls this directly

Calls 1

Text2MusicDatasetClass · 0.90

Tested by

no test coverage detected