updates
This commit is contained in:
@@ -1,373 +0,0 @@
|
|||||||
# Minimum Inference Code for FLUX
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import datetime
|
|
||||||
import math
|
|
||||||
import os
|
|
||||||
import random
|
|
||||||
from typing import Callable, List, Optional, Tuple
|
|
||||||
import einops
|
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from safetensors.torch import safe_open, load_file
|
|
||||||
from tqdm import tqdm
|
|
||||||
from PIL import Image
|
|
||||||
import accelerate
|
|
||||||
|
|
||||||
from .library import device_utils
|
|
||||||
from .library.device_utils import init_ipex, get_preferred_device
|
|
||||||
|
|
||||||
init_ipex()
|
|
||||||
|
|
||||||
|
|
||||||
from .library.utils import setup_logging
|
|
||||||
|
|
||||||
setup_logging()
|
|
||||||
import logging
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
import networks.lora_flux as lora_flux
|
|
||||||
from .library import flux_models, flux_utils, strategy_flux
|
|
||||||
|
|
||||||
|
|
||||||
def time_shift(mu: float, sigma: float, t: torch.Tensor):
|
|
||||||
return math.exp(mu) / (math.exp(mu) + (1 / t - 1) ** sigma)
|
|
||||||
|
|
||||||
|
|
||||||
def get_lin_function(x1: float = 256, y1: float = 0.5, x2: float = 4096, y2: float = 1.15) -> Callable[[float], float]:
|
|
||||||
m = (y2 - y1) / (x2 - x1)
|
|
||||||
b = y1 - m * x1
|
|
||||||
return lambda x: m * x + b
|
|
||||||
|
|
||||||
|
|
||||||
def get_schedule(
|
|
||||||
num_steps: int,
|
|
||||||
image_seq_len: int,
|
|
||||||
base_shift: float = 0.5,
|
|
||||||
max_shift: float = 1.15,
|
|
||||||
shift: bool = True,
|
|
||||||
) -> list[float]:
|
|
||||||
# extra step for zero
|
|
||||||
timesteps = torch.linspace(1, 0, num_steps + 1)
|
|
||||||
|
|
||||||
# shifting the schedule to favor high timesteps for higher signal images
|
|
||||||
if shift:
|
|
||||||
# eastimate mu based on linear estimation between two points
|
|
||||||
mu = get_lin_function(y1=base_shift, y2=max_shift)(image_seq_len)
|
|
||||||
timesteps = time_shift(mu, 1.0, timesteps)
|
|
||||||
|
|
||||||
return timesteps.tolist()
|
|
||||||
|
|
||||||
|
|
||||||
def denoise(
|
|
||||||
model: flux_models.Flux,
|
|
||||||
img: torch.Tensor,
|
|
||||||
img_ids: torch.Tensor,
|
|
||||||
txt: torch.Tensor,
|
|
||||||
txt_ids: torch.Tensor,
|
|
||||||
vec: torch.Tensor,
|
|
||||||
timesteps: list[float],
|
|
||||||
guidance: float = 4.0,
|
|
||||||
):
|
|
||||||
# this is ignored for schnell
|
|
||||||
guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype)
|
|
||||||
for t_curr, t_prev in zip(tqdm(timesteps[:-1]), timesteps[1:]):
|
|
||||||
t_vec = torch.full((img.shape[0],), t_curr, dtype=img.dtype, device=img.device)
|
|
||||||
pred = model(img=img, img_ids=img_ids, txt=txt, txt_ids=txt_ids, y=vec, timesteps=t_vec, guidance=guidance_vec)
|
|
||||||
|
|
||||||
img = img + (t_prev - t_curr) * pred
|
|
||||||
|
|
||||||
return img
|
|
||||||
|
|
||||||
|
|
||||||
def do_sample(
|
|
||||||
accelerator: Optional[accelerate.Accelerator],
|
|
||||||
model: flux_models.Flux,
|
|
||||||
img: torch.Tensor,
|
|
||||||
img_ids: torch.Tensor,
|
|
||||||
l_pooled: torch.Tensor,
|
|
||||||
t5_out: torch.Tensor,
|
|
||||||
txt_ids: torch.Tensor,
|
|
||||||
num_steps: int,
|
|
||||||
guidance: float,
|
|
||||||
is_schnell: bool,
|
|
||||||
device: torch.device,
|
|
||||||
flux_dtype: torch.dtype,
|
|
||||||
):
|
|
||||||
timesteps = get_schedule(num_steps, img.shape[1], shift=not is_schnell)
|
|
||||||
|
|
||||||
# denoise initial noise
|
|
||||||
if accelerator:
|
|
||||||
with accelerator.autocast(), torch.no_grad():
|
|
||||||
x = denoise(model, img, img_ids, t5_out, txt_ids, l_pooled, timesteps=timesteps, guidance=guidance)
|
|
||||||
else:
|
|
||||||
with torch.autocast(device_type=device.type, dtype=flux_dtype), torch.no_grad():
|
|
||||||
x = denoise(model, img, img_ids, t5_out, txt_ids, l_pooled, timesteps=timesteps, guidance=guidance)
|
|
||||||
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
def generate_image(
|
|
||||||
model,
|
|
||||||
clip_l,
|
|
||||||
t5xxl,
|
|
||||||
tokenize_strategy,
|
|
||||||
encoding_strategy,
|
|
||||||
ae,
|
|
||||||
prompt: str,
|
|
||||||
seed: Optional[int],
|
|
||||||
image_width: int,
|
|
||||||
image_height: int,
|
|
||||||
steps: Optional[int],
|
|
||||||
guidance: float,
|
|
||||||
):
|
|
||||||
seed = seed if seed is not None else random.randint(0, 2**32 - 1)
|
|
||||||
logger.info(f"Seed: {seed}")
|
|
||||||
|
|
||||||
# make first noise with packed shape
|
|
||||||
# original: b,16,2*h//16,2*w//16, packed: b,h//16*w//16,16*2*2
|
|
||||||
packed_latent_height, packed_latent_width = math.ceil(image_height / 16), math.ceil(image_width / 16)
|
|
||||||
noise = torch.randn(
|
|
||||||
1,
|
|
||||||
packed_latent_height * packed_latent_width,
|
|
||||||
16 * 2 * 2,
|
|
||||||
device=device,
|
|
||||||
dtype=dtype,
|
|
||||||
generator=torch.Generator(device=device).manual_seed(seed),
|
|
||||||
)
|
|
||||||
|
|
||||||
# prepare img and img ids
|
|
||||||
|
|
||||||
# this is needed only for img2img
|
|
||||||
# img = rearrange(img, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2)
|
|
||||||
# if img.shape[0] == 1 and bs > 1:
|
|
||||||
# img = repeat(img, "1 ... -> bs ...", bs=bs)
|
|
||||||
|
|
||||||
# txt2img only needs img_ids
|
|
||||||
img_ids = flux_utils.prepare_img_ids(1, packed_latent_height, packed_latent_width)
|
|
||||||
|
|
||||||
# prepare embeddings
|
|
||||||
logger.info("Encoding prompts...")
|
|
||||||
tokens_and_masks = tokenize_strategy.tokenize(prompt)
|
|
||||||
clip_l = clip_l.to(device)
|
|
||||||
t5xxl = t5xxl.to(device)
|
|
||||||
with torch.no_grad():
|
|
||||||
if is_fp8(clip_l_dtype) or is_fp8(t5xxl_dtype):
|
|
||||||
clip_l.to(clip_l_dtype)
|
|
||||||
t5xxl.to(t5xxl_dtype)
|
|
||||||
with accelerator.autocast():
|
|
||||||
_, t5_out, txt_ids = encoding_strategy.encode_tokens(
|
|
||||||
tokenize_strategy, [clip_l, t5xxl], tokens_and_masks, args.apply_t5_attn_mask
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
with torch.autocast(device_type=device.type, dtype=clip_l_dtype):
|
|
||||||
l_pooled, _, _ = encoding_strategy.encode_tokens(tokenize_strategy, [clip_l, None], tokens_and_masks)
|
|
||||||
with torch.autocast(device_type=device.type, dtype=t5xxl_dtype):
|
|
||||||
_, t5_out, txt_ids = encoding_strategy.encode_tokens(
|
|
||||||
tokenize_strategy, [None, t5xxl], tokens_and_masks, args.apply_t5_attn_mask
|
|
||||||
)
|
|
||||||
|
|
||||||
# NaN check
|
|
||||||
if torch.isnan(l_pooled).any():
|
|
||||||
raise ValueError("NaN in l_pooled")
|
|
||||||
if torch.isnan(t5_out).any():
|
|
||||||
raise ValueError("NaN in t5_out")
|
|
||||||
|
|
||||||
if args.offload:
|
|
||||||
clip_l = clip_l.cpu()
|
|
||||||
t5xxl = t5xxl.cpu()
|
|
||||||
# del clip_l, t5xxl
|
|
||||||
device_utils.clean_memory()
|
|
||||||
|
|
||||||
# generate image
|
|
||||||
logger.info("Generating image...")
|
|
||||||
model = model.to(device)
|
|
||||||
if steps is None:
|
|
||||||
steps = 4 if is_schnell else 50
|
|
||||||
|
|
||||||
img_ids = img_ids.to(device)
|
|
||||||
x = do_sample(accelerator, model, noise, img_ids, l_pooled, t5_out, txt_ids, steps, guidance, is_schnell, device, flux_dtype)
|
|
||||||
if args.offload:
|
|
||||||
model = model.cpu()
|
|
||||||
# del model
|
|
||||||
device_utils.clean_memory()
|
|
||||||
|
|
||||||
# unpack
|
|
||||||
x = x.float()
|
|
||||||
x = einops.rearrange(x, "b (h w) (c ph pw) -> b c (h ph) (w pw)", h=packed_latent_height, w=packed_latent_width, ph=2, pw=2)
|
|
||||||
|
|
||||||
# decode
|
|
||||||
logger.info("Decoding image...")
|
|
||||||
ae = ae.to(device)
|
|
||||||
with torch.no_grad():
|
|
||||||
if is_fp8(ae_dtype):
|
|
||||||
with accelerator.autocast():
|
|
||||||
x = ae.decode(x)
|
|
||||||
else:
|
|
||||||
with torch.autocast(device_type=device.type, dtype=ae_dtype):
|
|
||||||
x = ae.decode(x)
|
|
||||||
if args.offload:
|
|
||||||
ae = ae.cpu()
|
|
||||||
|
|
||||||
x = x.clamp(-1, 1)
|
|
||||||
x = x.permute(0, 2, 3, 1)
|
|
||||||
#img = Image.fromarray((127.5 * (x + 1.0)).float().cpu().numpy().astype(np.uint8)[0])
|
|
||||||
|
|
||||||
# # save image
|
|
||||||
# output_dir = args.output_dir
|
|
||||||
# os.makedirs(output_dir, exist_ok=True)
|
|
||||||
# output_path = os.path.join(output_dir, f"{datetime.datetime.now().strftime('%Y%m%d_%H%M%S')}.png")
|
|
||||||
# img.save(output_path)
|
|
||||||
|
|
||||||
# logger.info(f"Saved image to {output_path}")
|
|
||||||
return x
|
|
||||||
|
|
||||||
|
|
||||||
# if __name__ == "__main__":
|
|
||||||
# target_height = 768 # 1024
|
|
||||||
# target_width = 1360 # 1024
|
|
||||||
|
|
||||||
# # steps = 50 # 28 # 50
|
|
||||||
# # guidance_scale = 5
|
|
||||||
# # seed = 1 # None # 1
|
|
||||||
|
|
||||||
# device = get_preferred_device()
|
|
||||||
|
|
||||||
# parser = argparse.ArgumentParser()
|
|
||||||
# parser.add_argument("--ckpt_path", type=str, required=True)
|
|
||||||
# parser.add_argument("--clip_l", type=str, required=False)
|
|
||||||
# parser.add_argument("--t5xxl", type=str, required=False)
|
|
||||||
# parser.add_argument("--ae", type=str, required=False)
|
|
||||||
# parser.add_argument("--apply_t5_attn_mask", action="store_true")
|
|
||||||
# parser.add_argument("--prompt", type=str, default="A photo of a cat")
|
|
||||||
# parser.add_argument("--output_dir", type=str, default=".")
|
|
||||||
# parser.add_argument("--dtype", type=str, default="bfloat16", help="base dtype")
|
|
||||||
# parser.add_argument("--clip_l_dtype", type=str, default=None, help="dtype for clip_l")
|
|
||||||
# parser.add_argument("--ae_dtype", type=str, default=None, help="dtype for ae")
|
|
||||||
# parser.add_argument("--t5xxl_dtype", type=str, default=None, help="dtype for t5xxl")
|
|
||||||
# parser.add_argument("--flux_dtype", type=str, default=None, help="dtype for flux")
|
|
||||||
# parser.add_argument("--seed", type=int, default=None)
|
|
||||||
# parser.add_argument("--steps", type=int, default=None, help="Number of steps. Default is 4 for schnell, 50 for dev")
|
|
||||||
# parser.add_argument("--guidance", type=float, default=3.5)
|
|
||||||
# parser.add_argument("--offload", action="store_true", help="Offload to CPU")
|
|
||||||
# parser.add_argument(
|
|
||||||
# "--lora_weights",
|
|
||||||
# type=str,
|
|
||||||
# nargs="*",
|
|
||||||
# default=[],
|
|
||||||
# help="LoRA weights, only supports networks.lora_flux, each argument is a `path;multiplier` (semi-colon separated)",
|
|
||||||
# )
|
|
||||||
# parser.add_argument("--merge_lora_weights", action="store_true", help="Merge LoRA weights to model")
|
|
||||||
# parser.add_argument("--width", type=int, default=target_width)
|
|
||||||
# parser.add_argument("--height", type=int, default=target_height)
|
|
||||||
# parser.add_argument("--interactive", action="store_true")
|
|
||||||
# args = parser.parse_args()
|
|
||||||
|
|
||||||
# seed = args.seed
|
|
||||||
# steps = args.steps
|
|
||||||
# guidance_scale = args.guidance
|
|
||||||
|
|
||||||
# name = "schnell" if "schnell" in args.ckpt_path else "dev" # TODO change this to a more robust way
|
|
||||||
# is_schnell = name == "schnell"
|
|
||||||
|
|
||||||
# def str_to_dtype(s: Optional[str], default_dtype: Optional[torch.dtype] = None) -> torch.dtype:
|
|
||||||
# if s is None:
|
|
||||||
# return default_dtype
|
|
||||||
# if s in ["bf16", "bfloat16"]:
|
|
||||||
# return torch.bfloat16
|
|
||||||
# elif s in ["fp16", "float16"]:
|
|
||||||
# return torch.float16
|
|
||||||
# elif s in ["fp32", "float32"]:
|
|
||||||
# return torch.float32
|
|
||||||
# elif s in ["fp8_e4m3fn", "e4m3fn", "float8_e4m3fn"]:
|
|
||||||
# return torch.float8_e4m3fn
|
|
||||||
# elif s in ["fp8_e4m3fnuz", "e4m3fnuz", "float8_e4m3fnuz"]:
|
|
||||||
# return torch.float8_e4m3fnuz
|
|
||||||
# elif s in ["fp8_e5m2", "e5m2", "float8_e5m2"]:
|
|
||||||
# return torch.float8_e5m2
|
|
||||||
# elif s in ["fp8_e5m2fnuz", "e5m2fnuz", "float8_e5m2fnuz"]:
|
|
||||||
# return torch.float8_e5m2fnuz
|
|
||||||
# elif s in ["fp8", "float8"]:
|
|
||||||
# return torch.float8_e4m3fn # default fp8
|
|
||||||
# else:
|
|
||||||
# raise ValueError(f"Unsupported dtype: {s}")
|
|
||||||
|
|
||||||
# def is_fp8(dt):
|
|
||||||
# return dt in [torch.float8_e4m3fn, torch.float8_e4m3fnuz, torch.float8_e5m2, torch.float8_e5m2fnuz]
|
|
||||||
|
|
||||||
# dtype = str_to_dtype(args.dtype)
|
|
||||||
# clip_l_dtype = str_to_dtype(args.clip_l_dtype, dtype)
|
|
||||||
# t5xxl_dtype = str_to_dtype(args.t5xxl_dtype, dtype)
|
|
||||||
# ae_dtype = str_to_dtype(args.ae_dtype, dtype)
|
|
||||||
# flux_dtype = str_to_dtype(args.flux_dtype, dtype)
|
|
||||||
|
|
||||||
# logger.info(f"Dtypes for clip_l, t5xxl, ae, flux: {clip_l_dtype}, {t5xxl_dtype}, {ae_dtype}, {flux_dtype}")
|
|
||||||
|
|
||||||
# loading_device = "cpu" if args.offload else device
|
|
||||||
|
|
||||||
# use_fp8 = [is_fp8(d) for d in [dtype, clip_l_dtype, t5xxl_dtype, ae_dtype, flux_dtype]]
|
|
||||||
# if any(use_fp8):
|
|
||||||
# accelerator = accelerate.Accelerator(mixed_precision="bf16")
|
|
||||||
# else:
|
|
||||||
# accelerator = None
|
|
||||||
|
|
||||||
# # load clip_l
|
|
||||||
# logger.info(f"Loading clip_l from {args.clip_l}...")
|
|
||||||
# clip_l = flux_utils.load_clip_l(args.clip_l, clip_l_dtype, loading_device)
|
|
||||||
# clip_l.eval()
|
|
||||||
|
|
||||||
# logger.info(f"Loading t5xxl from {args.t5xxl}...")
|
|
||||||
# t5xxl = flux_utils.load_t5xxl(args.t5xxl, t5xxl_dtype, loading_device)
|
|
||||||
# t5xxl.eval()
|
|
||||||
|
|
||||||
# if is_fp8(clip_l_dtype):
|
|
||||||
# clip_l = accelerator.prepare(clip_l)
|
|
||||||
# if is_fp8(t5xxl_dtype):
|
|
||||||
# t5xxl = accelerator.prepare(t5xxl)
|
|
||||||
|
|
||||||
# t5xxl_max_length = 256 if is_schnell else 512
|
|
||||||
# tokenize_strategy = strategy_flux.FluxTokenizeStrategy(t5xxl_max_length)
|
|
||||||
# encoding_strategy = strategy_flux.FluxTextEncodingStrategy()
|
|
||||||
|
|
||||||
# # DiT
|
|
||||||
# model = flux_utils.load_flow_model(name, args.ckpt_path, flux_dtype, loading_device)
|
|
||||||
# model.eval()
|
|
||||||
# logger.info(f"Casting model to {flux_dtype}")
|
|
||||||
# model.to(flux_dtype) # make sure model is dtype
|
|
||||||
# if is_fp8(flux_dtype):
|
|
||||||
# model = accelerator.prepare(model)
|
|
||||||
|
|
||||||
# # AE
|
|
||||||
# ae = flux_utils.load_ae(name, args.ae, ae_dtype, loading_device)
|
|
||||||
# ae.eval()
|
|
||||||
# if is_fp8(ae_dtype):
|
|
||||||
# ae = accelerator.prepare(ae)
|
|
||||||
|
|
||||||
# # LoRA
|
|
||||||
# lora_models: List[lora_flux.LoRANetwork] = []
|
|
||||||
# for weights_file in args.lora_weights:
|
|
||||||
# if ";" in weights_file:
|
|
||||||
# weights_file, multiplier = weights_file.split(";")
|
|
||||||
# multiplier = float(multiplier)
|
|
||||||
# else:
|
|
||||||
# multiplier = 1.0
|
|
||||||
|
|
||||||
# lora_model, weights_sd = lora_flux.create_network_from_weights(
|
|
||||||
# multiplier, weights_file, ae, [clip_l, t5xxl], model, None, True
|
|
||||||
# )
|
|
||||||
# if args.merge_lora_weights:
|
|
||||||
# lora_model.merge_to([clip_l, t5xxl], model, weights_sd)
|
|
||||||
# else:
|
|
||||||
# lora_model.apply_to([clip_l, t5xxl], model)
|
|
||||||
# info = lora_model.load_state_dict(weights_sd, strict=True)
|
|
||||||
# logger.info(f"Loaded LoRA weights from {weights_file}: {info}")
|
|
||||||
# lora_model.eval()
|
|
||||||
# lora_model.to(device)
|
|
||||||
|
|
||||||
# lora_models.append(lora_model)
|
|
||||||
|
|
||||||
# if not args.interactive:
|
|
||||||
# generate_image(model, clip_l, t5xxl, ae, args.prompt, args.seed, args.width, args.height, args.steps, args.guidance)
|
|
||||||
|
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
{
|
||||||
|
"hf_token": "your_token_here"
|
||||||
|
}
|
||||||
@@ -25,7 +25,7 @@ setup_logging()
|
|||||||
import logging
|
import logging
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
from comfy.utils import ProgressBar
|
||||||
|
|
||||||
def sample_images(
|
def sample_images(
|
||||||
accelerator: Accelerator,
|
accelerator: Accelerator,
|
||||||
@@ -39,26 +39,10 @@ def sample_images(
|
|||||||
validation_settings=None,
|
validation_settings=None,
|
||||||
prompt_replacement=None,
|
prompt_replacement=None,
|
||||||
):
|
):
|
||||||
# if steps == 0:
|
|
||||||
# if not args.sample_at_first:
|
|
||||||
# return
|
|
||||||
# else:
|
|
||||||
# if args.sample_every_n_steps is None and args.sample_every_n_epochs is None:
|
|
||||||
# return
|
|
||||||
# if args.sample_every_n_epochs is not None:
|
|
||||||
# # sample_every_n_steps は無視する
|
|
||||||
# if epoch is None or epoch % args.sample_every_n_epochs != 0:
|
|
||||||
# return
|
|
||||||
# else:
|
|
||||||
# if steps % args.sample_every_n_steps != 0 or epoch is not None: # steps is not divisible or end of epoch
|
|
||||||
# return
|
|
||||||
|
|
||||||
logger.info("")
|
logger.info("")
|
||||||
logger.info(f"generating sample images at step: {steps}")
|
logger.info(f"generating sample images at step: {steps}")
|
||||||
#if not os.path.isfile(args.sample_prompts):
|
|
||||||
# logger.error(f"No prompt file / プロンプトファイルがありません: {args.sample_prompts}")
|
|
||||||
# return
|
|
||||||
|
|
||||||
#distributed_state = PartialState() # for multi gpu distributed inference. this is a singleton, so it's safe to use it here
|
#distributed_state = PartialState() # for multi gpu distributed inference. this is a singleton, so it's safe to use it here
|
||||||
|
|
||||||
# unwrap unet and text_encoder(s)
|
# unwrap unet and text_encoder(s)
|
||||||
@@ -300,10 +284,12 @@ def denoise(
|
|||||||
print("IMAGE DTYPE: ", img.dtype)
|
print("IMAGE DTYPE: ", img.dtype)
|
||||||
guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype)
|
guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype)
|
||||||
print("GUIDANCE VECTOR: ", guidance_vec)
|
print("GUIDANCE VECTOR: ", guidance_vec)
|
||||||
|
comfy_pbar = ProgressBar(total=len(timesteps))
|
||||||
for t_curr, t_prev in zip(tqdm(timesteps[:-1]), timesteps[1:]):
|
for t_curr, t_prev in zip(tqdm(timesteps[:-1]), timesteps[1:]):
|
||||||
t_vec = torch.full((img.shape[0],), t_curr, dtype=img.dtype, device=img.device)
|
t_vec = torch.full((img.shape[0],), t_curr, dtype=img.dtype, device=img.device)
|
||||||
pred = model(img=img, img_ids=img_ids, txt=txt, txt_ids=txt_ids, y=vec, timesteps=t_vec, guidance=guidance_vec)
|
pred = model(img=img, img_ids=img_ids, txt=txt, txt_ids=txt_ids, y=vec, timesteps=t_vec, guidance=guidance_vec)
|
||||||
|
|
||||||
img = img + (t_prev - t_curr) * pred
|
img = img + (t_prev - t_curr) * pred
|
||||||
|
comfy_pbar.update(1)
|
||||||
|
|
||||||
return img
|
return img
|
||||||
@@ -6,8 +6,9 @@ import folder_paths
|
|||||||
import comfy.model_management as mm
|
import comfy.model_management as mm
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import toml
|
import toml
|
||||||
|
import json
|
||||||
import time
|
import time
|
||||||
|
from pathlib import Path
|
||||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
from .flux_train_network_comfy import FluxNetworkTrainer
|
from .flux_train_network_comfy import FluxNetworkTrainer
|
||||||
@@ -69,6 +70,7 @@ class TrainDatasetConfig:
|
|||||||
"max_bucket_resos": ("STRING",{"default": "1024, 768, 512"}),
|
"max_bucket_resos": ("STRING",{"default": "1024, 768, 512"}),
|
||||||
"color_aug": ("BOOLEAN",{"default": False, "tooltip": "enable weak color augmentation"}),
|
"color_aug": ("BOOLEAN",{"default": False, "tooltip": "enable weak color augmentation"}),
|
||||||
"flip_aug": ("BOOLEAN",{"default": False, "tooltip": "enable horizontal flip augmentation"}),
|
"flip_aug": ("BOOLEAN",{"default": False, "tooltip": "enable horizontal flip augmentation"}),
|
||||||
|
"dataset_repeats": ("INT", {"default": 1, "min": 1, "tooltip": "number of times to repeat dataset for an epoch"}),
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -77,7 +79,7 @@ class TrainDatasetConfig:
|
|||||||
FUNCTION = "create_config"
|
FUNCTION = "create_config"
|
||||||
CATEGORY = "FluxTrainer"
|
CATEGORY = "FluxTrainer"
|
||||||
|
|
||||||
def create_config(self, dataset_path, class_tokens, width, height, batch_size, enable_bucket, color_aug, flip_aug,
|
def create_config(self, dataset_path, class_tokens, width, height, batch_size, dataset_repeats, enable_bucket, color_aug, flip_aug,
|
||||||
bucket_no_upscale, min_bucket_reso, max_bucket_resos):
|
bucket_no_upscale, min_bucket_reso, max_bucket_resos):
|
||||||
|
|
||||||
|
|
||||||
@@ -89,7 +91,7 @@ class TrainDatasetConfig:
|
|||||||
"datasets": [
|
"datasets": [
|
||||||
{
|
{
|
||||||
"resolution": (width, height),
|
"resolution": (width, height),
|
||||||
"batch_size": batch_size,
|
"batch_size": batch_size,
|
||||||
"keep_tokens": 2,
|
"keep_tokens": 2,
|
||||||
"enable_bucket": enable_bucket,
|
"enable_bucket": enable_bucket,
|
||||||
"bucket_no_upscale": bucket_no_upscale,
|
"bucket_no_upscale": bucket_no_upscale,
|
||||||
@@ -106,15 +108,18 @@ class TrainDatasetConfig:
|
|||||||
}
|
}
|
||||||
]
|
]
|
||||||
}
|
}
|
||||||
|
dataset_settings = {
|
||||||
return (toml.dumps(dataset),)
|
"repeats": dataset_repeats,
|
||||||
|
"dataset": toml.dumps(dataset)
|
||||||
|
}
|
||||||
|
return (dataset_settings,)
|
||||||
|
|
||||||
class InitFluxTraining:
|
class InitFluxTraining:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": {
|
return {"required": {
|
||||||
"flux_models": ("TRAIN_FLUX_MODELS",),
|
"flux_models": ("TRAIN_FLUX_MODELS",),
|
||||||
"dataset": ("TOML_DATASET",),
|
"dataset_settings": ("TOML_DATASET",),
|
||||||
"output_name": ("STRING", {"default": "flux_lora", "multiline": False}),
|
"output_name": ("STRING", {"default": "flux_lora", "multiline": False}),
|
||||||
"output_dir": ("STRING", {"default": "flux_trainer_output", "multiline": False}),
|
"output_dir": ("STRING", {"default": "flux_trainer_output", "multiline": False}),
|
||||||
"network_dim": ("INT", {"default": 4, "min": 1, "max": 256, "step": 1, "tooltip": "network dim"}),
|
"network_dim": ("INT", {"default": 4, "min": 1, "max": 256, "step": 1, "tooltip": "network dim"}),
|
||||||
@@ -152,8 +157,11 @@ class InitFluxTraining:
|
|||||||
FUNCTION = "init_training"
|
FUNCTION = "init_training"
|
||||||
CATEGORY = "FluxTrainer"
|
CATEGORY = "FluxTrainer"
|
||||||
|
|
||||||
def init_training(self, flux_models, dataset, sample_prompts, output_name, optimizer_type, attention_mode, training_dtype, save_dtype, **kwargs,):
|
def init_training(self, flux_models, dataset_settings, sample_prompts, output_name, optimizer_type, attention_mode, training_dtype, save_dtype, **kwargs,):
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
|
dataset = dataset_settings["dataset"]
|
||||||
|
dataset_repeats = dataset_settings["repeats"]
|
||||||
|
|
||||||
parser = setup_parser()
|
parser = setup_parser()
|
||||||
args, _ = parser.parse_known_args()
|
args, _ = parser.parse_known_args()
|
||||||
@@ -193,6 +201,7 @@ class InitFluxTraining:
|
|||||||
config_dict = {
|
config_dict = {
|
||||||
"sample_prompts": prompts,
|
"sample_prompts": prompts,
|
||||||
"save_precision": save_dtype,
|
"save_precision": save_dtype,
|
||||||
|
"dataset_repeats": dataset_repeats,
|
||||||
"mixed_precision": "bf16",
|
"mixed_precision": "bf16",
|
||||||
"num_cpu_threads_per_process": 1,
|
"num_cpu_threads_per_process": 1,
|
||||||
"pretrained_model_name_or_path": flux_models["transformer"],
|
"pretrained_model_name_or_path": flux_models["transformer"],
|
||||||
@@ -309,16 +318,16 @@ class FluxTrainSave:
|
|||||||
with torch.inference_mode(False):
|
with torch.inference_mode(False):
|
||||||
trainer = network_trainer["network_trainer"]
|
trainer = network_trainer["network_trainer"]
|
||||||
|
|
||||||
ckpt_name = train_util.get_epoch_ckpt_name(trainer.args, "." + trainer.args.save_model_as, trainer.current_epoch.value + 1)
|
ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, trainer.global_step)
|
||||||
trainer.save_model(ckpt_name, trainer.accelerator.unwrap_model(trainer.network), trainer.global_step, trainer.current_epoch.value + 1)
|
trainer.save_model(ckpt_name, trainer.accelerator.unwrap_model(trainer.network), trainer.global_step, trainer.current_epoch.value + 1)
|
||||||
|
|
||||||
remove_epoch_no = train_util.get_remove_epoch_no(trainer.args, trainer.current_epoch.value + 1)
|
remove_step_no = train_util.get_remove_step_no(trainer.args, trainer.global_step)
|
||||||
if remove_epoch_no is not None:
|
if remove_step_no is not None:
|
||||||
remove_ckpt_name = train_util.get_epoch_ckpt_name(trainer.args, "." + trainer.args.save_model_as, remove_epoch_no)
|
remove_ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, remove_step_no)
|
||||||
trainer.remove_model(remove_ckpt_name)
|
trainer.remove_model(remove_ckpt_name)
|
||||||
|
|
||||||
if save_state:
|
if save_state:
|
||||||
train_util.save_and_remove_state_on_epoch_end(trainer.args, trainer.accelerator, trainer.current_epoch.value + 1)
|
train_util.save_and_remove_state_stepwise(trainer.args, trainer.accelerator, trainer.global_step)
|
||||||
|
|
||||||
lora_path = os.path.join(trainer.args.output_dir, "output", ckpt_name)
|
lora_path = os.path.join(trainer.args.output_dir, "output", ckpt_name)
|
||||||
|
|
||||||
@@ -333,8 +342,8 @@ class FluxTrainEnd:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("STRING",)
|
RETURN_TYPES = ("STRING", "STRING",)
|
||||||
RETURN_NAMES = ("lora_path",)
|
RETURN_NAMES = ("lora_path", "metadata",)
|
||||||
FUNCTION = "endtrain"
|
FUNCTION = "endtrain"
|
||||||
CATEGORY = "FluxTrainer"
|
CATEGORY = "FluxTrainer"
|
||||||
|
|
||||||
@@ -359,11 +368,14 @@ class FluxTrainEnd:
|
|||||||
|
|
||||||
final_output_lora_path = os.path.join(network_trainer.args.output_dir, "output", network_trainer.args.output_name)
|
final_output_lora_path = os.path.join(network_trainer.args.output_dir, "output", network_trainer.args.output_name)
|
||||||
|
|
||||||
|
# metadata
|
||||||
|
metadata = json.dumps(network_trainer.metadata, indent=2)
|
||||||
|
|
||||||
training_loop = None
|
training_loop = None
|
||||||
network_trainer = None
|
network_trainer = None
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
return (final_output_lora_path,)
|
return (final_output_lora_path, metadata)
|
||||||
|
|
||||||
class FluxTrainValidationSettings:
|
class FluxTrainValidationSettings:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -467,9 +479,7 @@ class VisualizeLoss:
|
|||||||
|
|
||||||
# Convert the PIL Image to a torch tensor
|
# Convert the PIL Image to a torch tensor
|
||||||
image_tensor = transforms.ToTensor()(image)
|
image_tensor = transforms.ToTensor()(image)
|
||||||
print(image_tensor.shape)
|
|
||||||
image_tensor = image_tensor.unsqueeze(0).permute(0, 2, 3, 1).cpu().float()
|
image_tensor = image_tensor.unsqueeze(0).permute(0, 2, 3, 1).cpu().float()
|
||||||
print(image_tensor.shape)
|
|
||||||
|
|
||||||
return image_tensor,
|
return image_tensor,
|
||||||
|
|
||||||
@@ -668,11 +678,13 @@ class FluxKohyaInferenceSampler:
|
|||||||
):
|
):
|
||||||
# this is ignored for schnell
|
# this is ignored for schnell
|
||||||
guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype)
|
guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype)
|
||||||
|
comfy_pbar = comfy.utils.ProgressBar(total=len(timesteps))
|
||||||
for t_curr, t_prev in zip(tqdm(timesteps[:-1]), timesteps[1:]):
|
for t_curr, t_prev in zip(tqdm(timesteps[:-1]), timesteps[1:]):
|
||||||
t_vec = torch.full((img.shape[0],), t_curr, dtype=img.dtype, device=img.device)
|
t_vec = torch.full((img.shape[0],), t_curr, dtype=img.dtype, device=img.device)
|
||||||
pred = model(img=img, img_ids=img_ids, txt=txt, txt_ids=txt_ids, y=vec, timesteps=t_vec, guidance=guidance_vec)
|
pred = model(img=img, img_ids=img_ids, txt=txt, txt_ids=txt_ids, y=vec, timesteps=t_vec, guidance=guidance_vec)
|
||||||
|
|
||||||
img = img + (t_prev - t_curr) * pred
|
img = img + (t_prev - t_curr) * pred
|
||||||
|
comfy_pbar.update(1)
|
||||||
|
|
||||||
return img
|
return img
|
||||||
def do_sample(
|
def do_sample(
|
||||||
@@ -729,6 +741,92 @@ class FluxKohyaInferenceSampler:
|
|||||||
|
|
||||||
return ((0.5 * (x + 1.0)).cpu().float(),)
|
return ((0.5 * (x + 1.0)).cpu().float(),)
|
||||||
|
|
||||||
|
class UploadToHuggingFace:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"network_trainer": ("NETWORKTRAINER",),
|
||||||
|
"source_path": ("STRING", {"default": ""}),
|
||||||
|
"repo_id": ("STRING",{"default": ""}),
|
||||||
|
"path_in_repo": ("STRING",{"default": "model"}),
|
||||||
|
"revision": ("STRING", {"default": "main"}),
|
||||||
|
"private": ("BOOLEAN", {"default": True, "tooltip": "If creating a new repo, leave it private"}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"token": ("STRING", {"default": "","tooltip":"DO NOT LEAVE IN THE NODE or it might save in metadata, can also use the hf_token.json"}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("STRING",)
|
||||||
|
RETURN_NAMES = ("status",)
|
||||||
|
FUNCTION = "upload"
|
||||||
|
CATEGORY = "FluxTrainer"
|
||||||
|
|
||||||
|
def upload(self, source_path, network_trainer, repo_id, path_in_repo, private, revision,token):
|
||||||
|
from huggingface_hub import HfApi
|
||||||
|
|
||||||
|
with open(os.path.join(script_directory, "hf_token.json"), "r") as file:
|
||||||
|
token_data = json.load(file)
|
||||||
|
token = token_data["hf_token"]
|
||||||
|
|
||||||
|
# Save metadata to a JSON file
|
||||||
|
metadata = network_trainer["network_trainer"].metadata
|
||||||
|
metadata_file_path = Path(source_path) / "metadata.json"
|
||||||
|
with open(metadata_file_path, 'w') as f:
|
||||||
|
json.dump(metadata, f)
|
||||||
|
|
||||||
|
repo_type = "model"
|
||||||
|
api = HfApi(token=token)
|
||||||
|
|
||||||
|
try:
|
||||||
|
api.repo_info(repo_id=repo_id, revision=revision, repo_type=repo_type)
|
||||||
|
repo_exists = True
|
||||||
|
except:
|
||||||
|
repo_exists = False
|
||||||
|
|
||||||
|
if not repo_exists(repo_id=repo_id, repo_type=repo_type, token=token):
|
||||||
|
try:
|
||||||
|
api.create_repo(repo_id=repo_id, repo_type=repo_type, private=private)
|
||||||
|
except Exception as e: # Checked for RepositoryNotFoundError, but other exceptions could be problematic
|
||||||
|
logger.error("===========================================")
|
||||||
|
logger.error(f"failed to create HuggingFace repo: {e}")
|
||||||
|
logger.error("===========================================")
|
||||||
|
|
||||||
|
is_folder = (type(source_path) == str and os.path.isdir(source_path)) or (isinstance(source_path, Path) and source_path.is_dir())
|
||||||
|
|
||||||
|
try:
|
||||||
|
if is_folder:
|
||||||
|
api.upload_folder(
|
||||||
|
repo_id=repo_id,
|
||||||
|
repo_type=repo_type,
|
||||||
|
folder_path=source_path,
|
||||||
|
path_in_repo=path_in_repo,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
api.upload_file(
|
||||||
|
repo_id=repo_id,
|
||||||
|
repo_type=repo_type,
|
||||||
|
path_or_fileobj=source_path,
|
||||||
|
path_in_repo=path_in_repo,
|
||||||
|
)
|
||||||
|
# Upload the metadata file separately if it's not a folder upload
|
||||||
|
if not is_folder:
|
||||||
|
api.upload_file(
|
||||||
|
repo_id=repo_id,
|
||||||
|
repo_type=repo_type,
|
||||||
|
path_or_fileobj=str(metadata_file_path),
|
||||||
|
path_in_repo=path_in_repo + '/metadata.json',
|
||||||
|
)
|
||||||
|
status = "Uploaded to HuggingFace succesfully"
|
||||||
|
except Exception as e: # RuntimeErrorを確認済みだが他にあると困るので
|
||||||
|
logger.error("===========================================")
|
||||||
|
logger.error(f"failed to upload to HuggingFace / HuggingFaceへのアップロードに失敗しました : {e}")
|
||||||
|
logger.error("===========================================")
|
||||||
|
status = f"Failed to upload to HuggingFace {e}"
|
||||||
|
|
||||||
|
return (status,)
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"InitFluxTraining": InitFluxTraining,
|
"InitFluxTraining": InitFluxTraining,
|
||||||
"FluxTrainModelSelect": FluxTrainModelSelect,
|
"FluxTrainModelSelect": FluxTrainModelSelect,
|
||||||
@@ -739,7 +837,8 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"FluxTrainValidationSettings": FluxTrainValidationSettings,
|
"FluxTrainValidationSettings": FluxTrainValidationSettings,
|
||||||
"FluxTrainEnd": FluxTrainEnd,
|
"FluxTrainEnd": FluxTrainEnd,
|
||||||
"FluxTrainSave": FluxTrainSave,
|
"FluxTrainSave": FluxTrainSave,
|
||||||
"FluxKohyaInferenceSampler": FluxKohyaInferenceSampler
|
"FluxKohyaInferenceSampler": FluxKohyaInferenceSampler,
|
||||||
|
"UploadToHuggingFace": UploadToHuggingFace
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"InitFluxTraining": "Init Flux Training",
|
"InitFluxTraining": "Init Flux Training",
|
||||||
@@ -751,5 +850,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
|||||||
"FluxTrainValidationSettings": "Flux Train Validation Settings",
|
"FluxTrainValidationSettings": "Flux Train Validation Settings",
|
||||||
"FluxTrainEnd": "Flux Train End",
|
"FluxTrainEnd": "Flux Train End",
|
||||||
"FluxTrainSave": "Flux Train Save",
|
"FluxTrainSave": "Flux Train Save",
|
||||||
"FluxKohyaInferenceSampler": "Flux Kohya Inference Sampler"
|
"FluxKohyaInferenceSampler": "Flux Kohya Inference Sampler",
|
||||||
|
"UploadToHuggingFace": "Upload To HuggingFace"
|
||||||
}
|
}
|
||||||
|
|||||||
+1
-2
@@ -4,11 +4,10 @@ diffusers>=0.25.0
|
|||||||
ftfy>=6.1.1
|
ftfy>=6.1.1
|
||||||
opencv-python>=4.7.0.68
|
opencv-python>=4.7.0.68
|
||||||
einops>=0.7.0
|
einops>=0.7.0
|
||||||
pytorch-lightning>=1.9.0
|
#pytorch-lightning>=1.9.0
|
||||||
bitsandbytes>=0.43.3
|
bitsandbytes>=0.43.3
|
||||||
prodigyopt>=1.0
|
prodigyopt>=1.0
|
||||||
lion-pytorch>=0.0.6
|
lion-pytorch>=0.0.6
|
||||||
tensorboard
|
|
||||||
safetensors>=0.4.2
|
safetensors>=0.4.2
|
||||||
altair>=4.2.2
|
altair>=4.2.2
|
||||||
toml>=0.10.2
|
toml>=0.10.2
|
||||||
|
|||||||
+8
-8
@@ -913,7 +913,7 @@ class NetworkTrainer:
|
|||||||
# if initial_epoch or initial_step is specified, steps_from_state is ignored even when resuming
|
# if initial_epoch or initial_step is specified, steps_from_state is ignored even when resuming
|
||||||
if steps_from_state is not None:
|
if steps_from_state is not None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"steps from the state is ignored because initial_step is specified / initial_stepが指定されているため、stateからのステップ数は無視されます"
|
"steps from the state is ignored because initial_step is specified"
|
||||||
)
|
)
|
||||||
if args.initial_step is not None:
|
if args.initial_step is not None:
|
||||||
initial_step = args.initial_step
|
initial_step = args.initial_step
|
||||||
@@ -931,7 +931,7 @@ class NetworkTrainer:
|
|||||||
if initial_step > 0:
|
if initial_step > 0:
|
||||||
assert (
|
assert (
|
||||||
args.max_train_steps > initial_step
|
args.max_train_steps > initial_step
|
||||||
), f"max_train_steps should be greater than initial step / max_train_stepsは初期ステップより大きい必要があります: {args.max_train_steps} vs {initial_step}"
|
), f"max_train_steps should be greater than initial step: {args.max_train_steps} vs {initial_step}"
|
||||||
|
|
||||||
epoch_to_start = 0
|
epoch_to_start = 0
|
||||||
if initial_step > 0:
|
if initial_step > 0:
|
||||||
@@ -939,9 +939,9 @@ class NetworkTrainer:
|
|||||||
# if skip_until_initial_step is specified, load data and discard it to ensure the same data is used
|
# if skip_until_initial_step is specified, load data and discard it to ensure the same data is used
|
||||||
if not args.resume:
|
if not args.resume:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"initial_step is specified but not resuming. lr scheduler will be started from the beginning / initial_stepが指定されていますがresumeしていないため、lr schedulerは最初から始まります"
|
f"initial_step is specified but not resuming. lr scheduler will be started from the beginning"
|
||||||
)
|
)
|
||||||
logger.info(f"skipping {initial_step} steps / {initial_step}ステップをスキップします")
|
logger.info(f"skipping {initial_step} steps")
|
||||||
initial_step *= args.gradient_accumulation_steps
|
initial_step *= args.gradient_accumulation_steps
|
||||||
|
|
||||||
# set epoch to start to make initial_step less than len(train_dataloader)
|
# set epoch to start to make initial_step less than len(train_dataloader)
|
||||||
@@ -1074,11 +1074,11 @@ class NetworkTrainer:
|
|||||||
latents = batch["latents"].to(accelerator.device).to(dtype=weight_dtype)
|
latents = batch["latents"].to(accelerator.device).to(dtype=weight_dtype)
|
||||||
else:
|
else:
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
# latentに変換
|
# encode latents
|
||||||
latents = self.encode_images_to_latents(args, accelerator, vae, batch["images"].to(vae_dtype))
|
latents = self.encode_images_to_latents(args, accelerator, vae, batch["images"].to(vae_dtype))
|
||||||
latents = latents.to(dtype=weight_dtype)
|
latents = latents.to(dtype=weight_dtype)
|
||||||
|
|
||||||
# NaNが含まれていれば警告を表示し0に置き換える
|
# NaN check
|
||||||
if torch.any(torch.isnan(latents)):
|
if torch.any(torch.isnan(latents)):
|
||||||
accelerator.print("NaN found in latents, replacing with zeros")
|
accelerator.print("NaN found in latents, replacing with zeros")
|
||||||
latents = torch.nan_to_num(latents, 0, out=latents)
|
latents = torch.nan_to_num(latents, 0, out=latents)
|
||||||
@@ -1145,13 +1145,13 @@ class NetworkTrainer:
|
|||||||
loss = apply_masked_loss(loss, batch)
|
loss = apply_masked_loss(loss, batch)
|
||||||
loss = loss.mean([1, 2, 3])
|
loss = loss.mean([1, 2, 3])
|
||||||
|
|
||||||
loss_weights = batch["loss_weights"] # 各sampleごとのweight
|
loss_weights = batch["loss_weights"] # weight for each sample
|
||||||
loss = loss * loss_weights
|
loss = loss * loss_weights
|
||||||
|
|
||||||
# min snr gamma, scale v pred loss like noise pred, v pred like loss, debiased estimation etc.
|
# min snr gamma, scale v pred loss like noise pred, v pred like loss, debiased estimation etc.
|
||||||
loss = self.post_process_loss(loss, args, timesteps, noise_scheduler)
|
loss = self.post_process_loss(loss, args, timesteps, noise_scheduler)
|
||||||
|
|
||||||
loss = loss.mean() # 平均なのでbatch_sizeで割る必要なし
|
loss = loss.mean() # No need to divide by batch_size since it's an average
|
||||||
|
|
||||||
accelerator.backward(loss)
|
accelerator.backward(loss)
|
||||||
if accelerator.sync_gradients:
|
if accelerator.sync_gradients:
|
||||||
|
|||||||
Reference in New Issue
Block a user