Fix scheduler bugs (#177)

* fix scheduler bugs.

* fix
This commit is contained in:
Sayak Paul
2025-01-03 18:52:40 +05:30
committed by GitHub
parent b8352abf70
commit 071c140879
2 changed files with 11 additions and 5 deletions
+10 -4
View File
@@ -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})
+1 -1
View File
@@ -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.