From c8be524730312249af8ea4a528a26b1b63d9888f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 18 Aug 2024 02:10:18 +0300 Subject: [PATCH] updates --- flux_minimal_inference_comfy.py | 373 -------------------------------- hf_token.json | 3 + library/flux_train_utils.py | 22 +- nodes.py | 138 ++++++++++-- requirements.txt | 3 +- train_network.py | 16 +- 6 files changed, 135 insertions(+), 420 deletions(-) delete mode 100644 flux_minimal_inference_comfy.py create mode 100644 hf_token.json diff --git a/flux_minimal_inference_comfy.py b/flux_minimal_inference_comfy.py deleted file mode 100644 index cbbf4c3..0000000 --- a/flux_minimal_inference_comfy.py +++ /dev/null @@ -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) - \ No newline at end of file diff --git a/hf_token.json b/hf_token.json new file mode 100644 index 0000000..ab7d5c7 --- /dev/null +++ b/hf_token.json @@ -0,0 +1,3 @@ +{ + "hf_token": "your_token_here" +} \ No newline at end of file diff --git a/library/flux_train_utils.py b/library/flux_train_utils.py index cf3c8db..5e0eafc 100644 --- a/library/flux_train_utils.py +++ b/library/flux_train_utils.py @@ -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 \ No newline at end of file diff --git a/nodes.py b/nodes.py index 6dbf9ae..cb85ec8 100644 --- a/nodes.py +++ b/nodes.py @@ -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" } diff --git a/requirements.txt b/requirements.txt index 60c7a01..dc4313e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -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 diff --git a/train_network.py b/train_network.py index ff509d6..c712349 100644 --- a/train_network.py +++ b/train_network.py @@ -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: