mirror of
https://github.com/storytold/FineTrainers-Conditioning.git
synced 2026-10-09 00:09:45 +00:00
+10
-4
@@ -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})
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user