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
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
from comfy.utils import ProgressBar
|
||||
|
||||
def sample_images(
|
||||
accelerator: Accelerator,
|
||||
@@ -39,26 +39,10 @@ def sample_images(
|
||||
validation_settings=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(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
|
||||
|
||||
# unwrap unet and text_encoder(s)
|
||||
@@ -300,10 +284,12 @@ def denoise(
|
||||
print("IMAGE DTYPE: ", img.dtype)
|
||||
guidance_vec = torch.full((img.shape[0],), guidance, device=img.device, dtype=img.dtype)
|
||||
print("GUIDANCE VECTOR: ", guidance_vec)
|
||||
comfy_pbar = ProgressBar(total=len(timesteps))
|
||||
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
|
||||
comfy_pbar.update(1)
|
||||
|
||||
return img
|
||||
@@ -6,8 +6,9 @@ import folder_paths
|
||||
import comfy.model_management as mm
|
||||
import comfy.utils
|
||||
import toml
|
||||
import json
|
||||
import time
|
||||
|
||||
from pathlib import Path
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
from .flux_train_network_comfy import FluxNetworkTrainer
|
||||
@@ -69,6 +70,7 @@ class TrainDatasetConfig:
|
||||
"max_bucket_resos": ("STRING",{"default": "1024, 768, 512"}),
|
||||
"color_aug": ("BOOLEAN",{"default": False, "tooltip": "enable weak color 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"
|
||||
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):
|
||||
|
||||
|
||||
@@ -89,7 +91,7 @@ class TrainDatasetConfig:
|
||||
"datasets": [
|
||||
{
|
||||
"resolution": (width, height),
|
||||
"batch_size": batch_size,
|
||||
"batch_size": batch_size,
|
||||
"keep_tokens": 2,
|
||||
"enable_bucket": enable_bucket,
|
||||
"bucket_no_upscale": bucket_no_upscale,
|
||||
@@ -106,15 +108,18 @@ class TrainDatasetConfig:
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
return (toml.dumps(dataset),)
|
||||
dataset_settings = {
|
||||
"repeats": dataset_repeats,
|
||||
"dataset": toml.dumps(dataset)
|
||||
}
|
||||
return (dataset_settings,)
|
||||
|
||||
class InitFluxTraining:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"flux_models": ("TRAIN_FLUX_MODELS",),
|
||||
"dataset": ("TOML_DATASET",),
|
||||
"dataset_settings": ("TOML_DATASET",),
|
||||
"output_name": ("STRING", {"default": "flux_lora", "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"}),
|
||||
@@ -152,8 +157,11 @@ class InitFluxTraining:
|
||||
FUNCTION = "init_training"
|
||||
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()
|
||||
|
||||
dataset = dataset_settings["dataset"]
|
||||
dataset_repeats = dataset_settings["repeats"]
|
||||
|
||||
parser = setup_parser()
|
||||
args, _ = parser.parse_known_args()
|
||||
@@ -193,6 +201,7 @@ class InitFluxTraining:
|
||||
config_dict = {
|
||||
"sample_prompts": prompts,
|
||||
"save_precision": save_dtype,
|
||||
"dataset_repeats": dataset_repeats,
|
||||
"mixed_precision": "bf16",
|
||||
"num_cpu_threads_per_process": 1,
|
||||
"pretrained_model_name_or_path": flux_models["transformer"],
|
||||
@@ -309,16 +318,16 @@ class FluxTrainSave:
|
||||
with torch.inference_mode(False):
|
||||
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)
|
||||
|
||||
remove_epoch_no = train_util.get_remove_epoch_no(trainer.args, trainer.current_epoch.value + 1)
|
||||
if remove_epoch_no is not None:
|
||||
remove_ckpt_name = train_util.get_epoch_ckpt_name(trainer.args, "." + trainer.args.save_model_as, remove_epoch_no)
|
||||
remove_step_no = train_util.get_remove_step_no(trainer.args, trainer.global_step)
|
||||
if remove_step_no is not None:
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
@@ -333,8 +342,8 @@ class FluxTrainEnd:
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lora_path",)
|
||||
RETURN_TYPES = ("STRING", "STRING",)
|
||||
RETURN_NAMES = ("lora_path", "metadata",)
|
||||
FUNCTION = "endtrain"
|
||||
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)
|
||||
|
||||
# metadata
|
||||
metadata = json.dumps(network_trainer.metadata, indent=2)
|
||||
|
||||
training_loop = None
|
||||
network_trainer = None
|
||||
mm.soft_empty_cache()
|
||||
|
||||
return (final_output_lora_path,)
|
||||
return (final_output_lora_path, metadata)
|
||||
|
||||
class FluxTrainValidationSettings:
|
||||
@classmethod
|
||||
@@ -467,9 +479,7 @@ class VisualizeLoss:
|
||||
|
||||
# Convert the PIL Image to a torch tensor
|
||||
image_tensor = transforms.ToTensor()(image)
|
||||
print(image_tensor.shape)
|
||||
image_tensor = image_tensor.unsqueeze(0).permute(0, 2, 3, 1).cpu().float()
|
||||
print(image_tensor.shape)
|
||||
|
||||
return image_tensor,
|
||||
|
||||
@@ -668,11 +678,13 @@ class FluxKohyaInferenceSampler:
|
||||
):
|
||||
# this is ignored for schnell
|
||||
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:]):
|
||||
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
|
||||
comfy_pbar.update(1)
|
||||
|
||||
return img
|
||||
def do_sample(
|
||||
@@ -729,6 +741,92 @@ class FluxKohyaInferenceSampler:
|
||||
|
||||
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 = {
|
||||
"InitFluxTraining": InitFluxTraining,
|
||||
"FluxTrainModelSelect": FluxTrainModelSelect,
|
||||
@@ -739,7 +837,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"FluxTrainValidationSettings": FluxTrainValidationSettings,
|
||||
"FluxTrainEnd": FluxTrainEnd,
|
||||
"FluxTrainSave": FluxTrainSave,
|
||||
"FluxKohyaInferenceSampler": FluxKohyaInferenceSampler
|
||||
"FluxKohyaInferenceSampler": FluxKohyaInferenceSampler,
|
||||
"UploadToHuggingFace": UploadToHuggingFace
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"InitFluxTraining": "Init Flux Training",
|
||||
@@ -751,5 +850,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FluxTrainValidationSettings": "Flux Train Validation Settings",
|
||||
"FluxTrainEnd": "Flux Train End",
|
||||
"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
|
||||
opencv-python>=4.7.0.68
|
||||
einops>=0.7.0
|
||||
pytorch-lightning>=1.9.0
|
||||
#pytorch-lightning>=1.9.0
|
||||
bitsandbytes>=0.43.3
|
||||
prodigyopt>=1.0
|
||||
lion-pytorch>=0.0.6
|
||||
tensorboard
|
||||
safetensors>=0.4.2
|
||||
altair>=4.2.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 steps_from_state is not None:
|
||||
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:
|
||||
initial_step = args.initial_step
|
||||
@@ -931,7 +931,7 @@ class NetworkTrainer:
|
||||
if initial_step > 0:
|
||||
assert (
|
||||
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
|
||||
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 not args.resume:
|
||||
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
|
||||
|
||||
# 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)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
# latentに変換
|
||||
# encode latents
|
||||
latents = self.encode_images_to_latents(args, accelerator, vae, batch["images"].to(vae_dtype))
|
||||
latents = latents.to(dtype=weight_dtype)
|
||||
|
||||
# NaNが含まれていれば警告を表示し0に置き換える
|
||||
# NaN check
|
||||
if torch.any(torch.isnan(latents)):
|
||||
accelerator.print("NaN found in latents, replacing with zeros")
|
||||
latents = torch.nan_to_num(latents, 0, out=latents)
|
||||
@@ -1145,13 +1145,13 @@ class NetworkTrainer:
|
||||
loss = apply_masked_loss(loss, batch)
|
||||
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
|
||||
|
||||
# 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 = loss.mean() # 平均なのでbatch_sizeで割る必要なし
|
||||
loss = loss.mean() # No need to divide by batch_size since it's an average
|
||||
|
||||
accelerator.backward(loss)
|
||||
if accelerator.sync_gradients:
|
||||
|
||||
Reference in New Issue
Block a user