From 33e865120c1ea08c39e9fb7b68e2ebd09b37ce58 Mon Sep 17 00:00:00 2001 From: Can Date: Tue, 4 Feb 2025 22:20:42 +0300 Subject: [PATCH] Refactor imports and improve code formatting in dataset and trainer modules --- src/f5_tts/model/dataset.py | 1 - src/f5_tts/model/trainer.py | 10 +++++----- 2 files changed, 5 insertions(+), 6 deletions(-) diff --git a/src/f5_tts/model/dataset.py b/src/f5_tts/model/dataset.py index 17c9f8b..75eeddd 100644 --- a/src/f5_tts/model/dataset.py +++ b/src/f5_tts/model/dataset.py @@ -1,5 +1,4 @@ import json -import random from importlib.resources import files import torch diff --git a/src/f5_tts/model/trainer.py b/src/f5_tts/model/trainer.py index b62a767..26970a3 100644 --- a/src/f5_tts/model/trainer.py +++ b/src/f5_tts/model/trainer.py @@ -279,11 +279,11 @@ class Trainer: self.accelerator.even_batches = False sampler = SequentialSampler(train_dataset) batch_sampler = DynamicBatchSampler( - sampler, - self.batch_size, - max_samples=self.max_samples, + sampler, + self.batch_size, + max_samples=self.max_samples, random_seed=resumable_with_seed, # This enables reproducible shuffling - drop_last=False + drop_last=False, ) train_dataloader = DataLoader( train_dataset, @@ -334,7 +334,7 @@ class Trainer: current_dataloader = train_dataloader # Set epoch for the batch sampler if it exists - if hasattr(train_dataloader, 'batch_sampler') and hasattr(train_dataloader.batch_sampler, 'set_epoch'): + if hasattr(train_dataloader, "batch_sampler") and hasattr(train_dataloader.batch_sampler, "set_epoch"): train_dataloader.batch_sampler.set_epoch(epoch) progress_bar = tqdm(