This commit is contained in:
kijai
2024-08-18 02:10:18 +03:00
parent b2353741fc
commit c8be524730
6 changed files with 135 additions and 420 deletions
-373
View File
@@ -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)
+3
View File
@@ -0,0 +1,3 @@
{
"hf_token": "your_token_here"
}
+3 -17
View File
@@ -25,7 +25,7 @@ setup_logging()
import logging
logger = logging.getLogger(__name__)
from comfy.utils import ProgressBar
def sample_images(
accelerator: Accelerator,
@@ -39,25 +39,9 @@ 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
@@ -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
+118 -18
View File
@@ -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):
@@ -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,9 +157,12 @@ 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
View File
@@ -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
View File
@@ -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: