From 071c1408796dcdb03b69a8bbaed06649020aed2a Mon Sep 17 00:00:00 2001 From: Sayak Paul Date: Fri, 3 Jan 2025 18:52:40 +0530 Subject: [PATCH] Fix scheduler bugs (#177) * fix scheduler bugs. * fix --- finetrainers/trainer.py | 14 ++++++++++---- finetrainers/utils/diffusion_utils.py | 2 +- 2 files changed, 11 insertions(+), 5 deletions(-) diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index eb1a6de..cf114b2 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -11,6 +11,7 @@ import diffusers import torch import torch.backends import transformers +import wandb from accelerate import Accelerator, DistributedType from accelerate.logging import get_logger from accelerate.utils import ( @@ -29,8 +30,6 @@ from huggingface_hub import create_repo, upload_folder from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict from tqdm import tqdm -import wandb - from .args import _INVERSE_DTYPE_MAP, Args, validate_args from .constants import ( FINETRAINERS_LOG_LEVEL, @@ -645,8 +644,14 @@ class Trainer: generator = generator.manual_seed(self.args.seed) self.state.generator = generator - scheduler_sigmas = get_scheduler_sigmas(self.scheduler).to(device=accelerator.device, dtype=torch.float32) - scheduler_alphas = get_scheduler_alphas(self.scheduler).to(device=accelerator.device, dtype=torch.float32) + scheduler_sigmas = get_scheduler_sigmas(self.scheduler) + scheduler_sigmas = ( + scheduler_sigmas.to(device=accelerator.device, dtype=torch.float32) if scheduler_sigmas else None + ) + scheduler_alphas = get_scheduler_alphas(self.scheduler) + scheduler_alphas = ( + scheduler_alphas.to(device=accelerator.device, dtype=torch.float32) if scheduler_alphas else None + ) for epoch in range(first_epoch, self.state.train_epochs): logger.debug(f"Starting epoch ({epoch + 1}/{self.state.train_epochs})") @@ -751,6 +756,7 @@ class Trainer: else: # Default to flow-matching noise addition noisy_latents = (1.0 - sigmas) * latent_conditions["latents"] + sigmas * noise + noisy_latents = noisy_latents.to(latent_conditions["latents"].dtype) latent_conditions.update({"noisy_latents": noisy_latents}) diff --git a/finetrainers/utils/diffusion_utils.py b/finetrainers/utils/diffusion_utils.py index 1d46fbc..94eead7 100644 --- a/finetrainers/utils/diffusion_utils.py +++ b/finetrainers/utils/diffusion_utils.py @@ -121,7 +121,7 @@ def prepare_loss_weights( flow_weighting_scheme: str = "none", ) -> torch.Tensor: if isinstance(scheduler, FlowMatchEulerDiscreteScheduler): - return compute_loss_weighting_for_sd3(sigmas, weighting_scheme=flow_weighting_scheme) + return compute_loss_weighting_for_sd3(sigmas=sigmas, weighting_scheme=flow_weighting_scheme) elif isinstance(scheduler, CogVideoXDDIMScheduler): # SNR is computed as (alphas / (1 - alphas)), but for some reason CogVideoX uses 1 / (1 - alphas). # TODO(aryan): Experiment if using alphas / (1 - alphas) gives better results.