diff --git a/library/flux_train_utils.py b/library/flux_train_utils.py index 0aa2c09..a2cff8d 100644 --- a/library/flux_train_utils.py +++ b/library/flux_train_utils.py @@ -36,6 +36,7 @@ def sample_images( ae, text_encoders, sample_prompts_te_outputs, + validation_settings=None, prompt_replacement=None, ): # if steps == 0: @@ -53,7 +54,7 @@ def sample_images( # return 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 @@ -111,6 +112,7 @@ def sample_images( steps, sample_prompts_te_outputs, prompt_replacement, + validation_settings ) image_tensor_list.append(image_tensor) @@ -134,14 +136,22 @@ def sample_image_inference( steps, sample_prompts_te_outputs, prompt_replacement, + validation_settings=None ): assert isinstance(prompt_dict, dict) # negative_prompt = prompt_dict.get("negative_prompt") - sample_steps = prompt_dict.get("sample_steps", 20) - width = prompt_dict.get("width", 512) - height = prompt_dict.get("height", 512) - scale = prompt_dict.get("scale", 3.5) - seed = prompt_dict.get("seed") + if validation_settings is not None: + sample_steps = validation_settings["steps"] + width = validation_settings["width"] + height = validation_settings["height"] + scale = validation_settings["guidance_scale"] + seed = validation_settings["seed"] + else: + sample_steps = prompt_dict.get("sample_steps", 20) + width = prompt_dict.get("width", 512) + height = prompt_dict.get("height", 512) + scale = prompt_dict.get("scale", 3.5) + seed = prompt_dict.get("seed") # controlnet_image = prompt_dict.get("controlnet_image") prompt: str = prompt_dict.get("prompt", "") # sampler_name: str = prompt_dict.get("sample_sampler", args.sample_sampler) diff --git a/library/strategy_base.py b/library/strategy_base.py index 7c59097..b5cb6be 100644 --- a/library/strategy_base.py +++ b/library/strategy_base.py @@ -107,8 +107,8 @@ class TextEncodingStrategy: @classmethod def set_strategy(cls, strategy): - if cls._strategy is not None: - raise RuntimeError(f"Internal error. {cls.__name__} strategy is already set") + #if cls._strategy is not None: + # raise RuntimeError(f"Internal error. {cls.__name__} strategy is already set") cls._strategy = strategy @classmethod @@ -139,8 +139,8 @@ class TextEncoderOutputsCachingStrategy: @classmethod def set_strategy(cls, strategy): - if cls._strategy is not None: - raise RuntimeError(f"Internal error. {cls.__name__} strategy is already set") + #if cls._strategy is not None: + # raise RuntimeError(f"Internal error. {cls.__name__} strategy is already set") cls._strategy = strategy @classmethod @@ -186,8 +186,8 @@ class LatentsCachingStrategy: @classmethod def set_strategy(cls, strategy): - if cls._strategy is not None: - raise RuntimeError(f"Internal error. {cls.__name__} strategy is already set") + #if cls._strategy is not None: + # raise RuntimeError(f"Internal error. {cls.__name__} strategy is already set") cls._strategy = strategy @classmethod @@ -260,7 +260,7 @@ class LatentsCachingStrategy: """ Default implementation for cache_batch_latents. Image loading, VAE, flipping, alpha mask handling are common. """ - from library import train_util # import here to avoid circular import + from . import train_util # import here to avoid circular import img_tensor, alpha_masks, original_sizes, crop_ltrbs = train_util.load_images_and_masks_for_caching( image_infos, alpha_mask, random_crop diff --git a/nodes.py b/nodes.py index fa1b635..d110ba9 100644 --- a/nodes.py +++ b/nodes.py @@ -1,5 +1,6 @@ import os import torch +from torchvision import transforms import math import copy import folder_paths @@ -15,7 +16,6 @@ import torch from accelerate import Accelerator accelerator = Accelerator(mixed_precision='bf16', cpu=False) from .library.device_utils import init_ipex, clean_memory_on_device -from .library.train_util import sample_images_common init_ipex() from .library import flux_models, flux_train_utils, flux_utils, sd3_train_utils, strategy_base, strategy_flux, train_util @@ -35,11 +35,11 @@ class FluxNetworkTrainer(NetworkTrainer): if args.cache_text_encoder_outputs: assert ( train_dataset_group.is_text_encoder_output_cacheable() - ), "when caching Text Encoder output, either caption_dropout_rate, shuffle_caption, token_warmup_step or caption_tag_dropout_rate cannot be used / Text Encoderの出力をキャッシュするときはcaption_dropout_rate, shuffle_caption, token_warmup_step, caption_tag_dropout_rateは使えません" + ), "when caching Text Encoder output, either caption_dropout_rate, shuffle_caption, token_warmup_step or caption_tag_dropout_rate cannot be used" assert ( args.network_train_unet_only or not args.cache_text_encoder_outputs - ), "network for Text Encoder cannot be trained with caching Text Encoder outputs / Text Encoderの出力をキャッシュしながらText Encoderのネットワークを学習することはできません" + ), "network for Text Encoder cannot be trained with caching Text Encoder outputs" train_dataset_group.verify_bucket_reso_steps(32) # TODO check this @@ -129,7 +129,7 @@ class FluxNetworkTrainer(NetworkTrainer): ): if args.cache_text_encoder_outputs: if not args.lowram: - # メモリ消費を減らす + # reduce memory consumption logger.info("move vae and unet to cpu to save memory") org_vae_device = vae.device org_unet_device = unet.device @@ -144,7 +144,7 @@ class FluxNetworkTrainer(NetworkTrainer): with accelerator.autocast(): dataset.new_cache_text_encoder_outputs(text_encoders, accelerator.is_main_process) - # cache sample prompts + # cache sample prompts self.sample_prompts_te_outputs = None if args.sample_prompts is not None: logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}") @@ -152,8 +152,6 @@ class FluxNetworkTrainer(NetworkTrainer): tokenize_strategy: strategy_flux.FluxTokenizeStrategy = strategy_base.TokenizeStrategy.get_strategy() text_encoding_strategy: strategy_flux.FluxTextEncodingStrategy = strategy_base.TextEncodingStrategy.get_strategy() - #prompts = sd3_train_utils.load_prompts(args.sample_prompts) - prompts = [] for line in args.sample_prompts: line = line.strip() @@ -201,12 +199,12 @@ class FluxNetworkTrainer(NetworkTrainer): text_encoders[0].to(accelerator.device, dtype=weight_dtype) text_encoders[1].to(accelerator.device, dtype=weight_dtype) - def sample_images(self, accelerator, args, epoch, global_step, device, ae, tokenizer, text_encoder, flux): + def sample_images(self, accelerator, args, epoch, global_step, ae, text_encoder, flux, validation_settings): if not args.split_mode: - flux_train_utils.sample_images( - accelerator, args, epoch, global_step, flux, ae, text_encoder, self.sample_prompts_te_outputs + image_tensors = flux_train_utils.sample_images( + accelerator, args, epoch, global_step, flux, ae, text_encoder, self.sample_prompts_te_outputs, validation_settings ) - return + return image_tensors class FluxUpperLowerWrapper(torch.nn.Module): def __init__(self, flux_upper: flux_models.FluxUpper, flux_lower: flux_models.FluxLower, device: torch.device): @@ -228,7 +226,7 @@ class FluxNetworkTrainer(NetworkTrainer): wrapper = FluxUpperLowerWrapper(self.flux_upper, flux, accelerator.device) clean_memory_on_device(accelerator.device) flux_train_utils.sample_images( - accelerator, args, epoch, global_step, wrapper, ae, text_encoder, self.sample_prompts_te_outputs + accelerator, args, epoch, global_step, flux, ae, text_encoder, self.sample_prompts_te_outputs, validation_settings ) clean_memory_on_device(accelerator.device) @@ -462,7 +460,7 @@ class FluxTrainModelSelect: RETURN_TYPES = ("TRAIN_FLUX_MODELS",) RETURN_NAMES = ("flux_models",) FUNCTION = "loadmodel" - CATEGORY = "TrainFlux" + CATEGORY = "FluxTrainer" def loadmodel(self, transformer, vae, clip_l, t5): @@ -501,7 +499,7 @@ class TrainDatasetConfig: RETURN_TYPES = ("TOML_DATASET",) RETURN_NAMES = ("dataset",) FUNCTION = "create_config" - CATEGORY = "TrainFlux" + CATEGORY = "FluxTrainer" def create_config(self, dataset_path, class_tokens, width, height, batch_size, enable_bucket, color_aug, flip_aug, bucket_no_upscale, min_bucket_reso, max_bucket_reso): @@ -564,16 +562,17 @@ class InitFluxTraining: "model_prediction_type": (["raw", "additive", "sigma_scaled"], {"tooltip": "How to interpret and process the model prediction: raw (use as is), additive (add to noisy input), sigma_scaled (apply sigma scaling)."}), "discrete_flow_shift": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "for the Euler Discrete Scheduler, default is 3.0"}), "highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}), - "sample_prompts": ("STRING", {"multiline": True, "default": "sample prompts", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}), + "attention_mode": (["sdpa", "xformers", "disabled"], {"default": "default", "tooltip": "memory efficient attention mode"}), + "sample_prompts": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}), }, } - RETURN_TYPES = ("NETWORKTRAINER",) - RETURN_NAMES = ("network_trainer",) + RETURN_TYPES = ("NETWORKTRAINER", "INT", "STRING", ) + RETURN_NAMES = ("network_trainer", "epochs_count", "output_path",) FUNCTION = "init_training" - CATEGORY = "TrainFlux" + CATEGORY = "FluxTrainer" - def init_training(self, flux_models, dataset, sample_prompts, output_name, optimizer_type, **kwargs,): + def init_training(self, flux_models, dataset, sample_prompts, output_name, optimizer_type, attention_mode, **kwargs,): mm.soft_empty_cache() parser = setup_parser() @@ -619,7 +618,6 @@ class InitFluxTraining: "t5xxl": flux_models["t5"], "ae": flux_models["vae"], "save_model_as": "safetensors", - "sdpa": True, "persistent_data_loader_workers": False, "max_data_loader_n_workers": 0, "seed": 42, @@ -633,6 +631,12 @@ class InitFluxTraining: "loss_type": "l2", "optimizer_type": optimizer_type, } + attention_settings = { + "sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True}, + "xformers": {"mem_eff_attn": True, "xformers": True, "spda": False} + } + config_dict.update(attention_settings.get(attention_mode, {})) + if optimizer_type == "adafactor": config_dict["optimizer_args"] = [ "relative_step=False", @@ -650,37 +654,41 @@ class InitFluxTraining: final_output_lora_path = os.path.join(output_dir, "output", output_name) + epochs_count = network_trainer.num_train_epochs + trainer = { "network_trainer": network_trainer, "training_loop": training_loop, } - return (trainer, ) + return (trainer, epochs_count, final_output_lora_path) -class TrainLoop: +class FluxTrainLoop: @classmethod def INPUT_TYPES(s): return {"required": { "network_trainer": ("NETWORKTRAINER",), - "epochs": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1}), + "steps": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1}), "end": ("BOOLEAN", {"default": False, "tooltip": "whether to end training"}), }, } - RETURN_TYPES = ("NETWORKTRAINER", "IMAGE", "LOSSRECORDER",) - RETURN_NAMES = ("network_trainer", "validation_images", "loss_recorder") - FUNCTION = "loadmodel" - CATEGORY = "TrainFlux" + RETURN_TYPES = ("NETWORKTRAINER",) + RETURN_NAMES = ("network_trainer",) + FUNCTION = "train" + CATEGORY = "FluxTrainer" - def loadmodel(self, network_trainer, epochs, end): + def train(self, network_trainer, steps, end): with torch.inference_mode(False): training_loop = network_trainer["training_loop"] network_trainer = network_trainer["network_trainer"] - - print(network_trainer.num_train_epochs) - pbar = comfy.utils.ProgressBar(epochs) - for epoch in range(epochs): - global_step, current_epoch = training_loop( - epoch=epoch, + initial_global_step = network_trainer.global_step + + target_global_step = network_trainer.global_step + steps + pbar = comfy.utils.ProgressBar(steps) + while network_trainer.global_step < target_global_step: + epoch = training_loop( + break_at_steps=target_global_step, + epoch=network_trainer.current_epoch.value, num_train_epochs=network_trainer.num_train_epochs, accelerator=network_trainer.accelerator, network=network_trainer.network, @@ -698,38 +706,39 @@ class TrainLoop: lr_scheduler=network_trainer.lr_scheduler, loss_recorder=network_trainer.loss_recorder ) - pbar.update(1) - print("GLOBAL STEP: ", global_step) - print("CURRENT EPOCH: ", current_epoch.value) + pbar.update(network_trainer.global_step - initial_global_step) + + # Also break if the global steps have reached the max train steps + if network_trainer.global_step >= network_trainer.args.max_train_steps: + break - with torch.inference_mode(True): - image_tensors = flux_train_utils.sample_images( - accelerator, - network_trainer.args, - epoch, - global_step, - network_trainer.unet, - network_trainer.vae, - network_trainer.text_encoder, - network_trainer.sample_prompts_te_outputs - ) - print(image_tensors.min(), image_tensors.max()) + # with torch.inference_mode(True): + # image_tensors = network_trainer.sample_images( + # network_trainer.accelerator, + # network_trainer.args, + # epoch, + # network_trainer.global_step, + # network_trainer.vae, + # network_trainer.text_encoder, + # network_trainer.unet + # ) + # print(image_tensors.min(), image_tensors.max(), image_tensors.shape) if end: network_trainer.metadata["ss_epoch"] = str(network_trainer.num_train_epochs) network_trainer.metadata["ss_training_finished_at"] = str(time.time()) - network = accelerator.unwrap_model(network) + network = network_trainer.accelerator.unwrap_model(network_trainer.network) - accelerator.end_training() + network_trainer.accelerator.end_training() - train_util.save_state_on_train_end(network_trainer.args, accelerator) + train_util.save_state_on_train_end(network_trainer.args, network_trainer.accelerator) ckpt_name = train_util.get_last_ckpt_name(network_trainer.args, "." + network_trainer.args.save_model_as) - network_trainer.save_model(ckpt_name, network, global_step, network_trainer.num_train_epochs, force_sync_upload=True) + network_trainer.save_model(ckpt_name, network, network_trainer.global_step, network_trainer.num_train_epochs, force_sync_upload=True) logger.info("model saved.") else: ckpt_name = train_util.get_epoch_ckpt_name(network_trainer.args, "." + network_trainer.args.save_model_as, epoch + 1) - network_trainer.save_model(ckpt_name, accelerator.unwrap_model(network_trainer.network), global_step, epoch + 1) + network_trainer.save_model(ckpt_name, accelerator.unwrap_model(network_trainer.network), network_trainer.global_step, epoch + 1) remove_epoch_no = train_util.get_remove_epoch_no(network_trainer.args, epoch + 1) if remove_epoch_no is not None: @@ -743,7 +752,116 @@ class TrainLoop: "network_trainer": network_trainer, "training_loop": training_loop, } - return (trainer, (0.5 * (image_tensors + 1.0)).cpu().float(), network_trainer.loss_recorder.loss_list) + return (trainer, ) + +class FluxTrainValidationSettings: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "steps": ("INT", {"default": 20, "min": 1, "max": 256, "step": 1, "tooltip": "sampling steps"}), + "width": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8, "tooltip": "image width"}), + "height": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8, "tooltip": "image height"}), + "guidance_scale": ("FLOAT", {"default": 3.5, "min": 1.0, "max": 32.0, "step": 0.05, "tooltip": "guidance scale"}), + "seed": ("INT", {"default": 42,"min": 0, "max": 0xffffffffffffffff, "step": 1}), + }, + } + + RETURN_TYPES = ("VALSETTINGS", ) + RETURN_NAMES = ("validation_settings", ) + FUNCTION = "set" + CATEGORY = "FluxTrainer" + + def set(self, **kwargs): + validation_settings = kwargs + print(validation_settings) + + return (validation_settings,) + +class FluxTrainValidate: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "network_trainer": ("NETWORKTRAINER",), + }, + "optional": { + "validation_settings": ("VALSETTINGS",), + } + } + + RETURN_TYPES = ("NETWORKTRAINER", "IMAGE",) + RETURN_NAMES = ("network_trainer", "validation_images",) + FUNCTION = "validate" + CATEGORY = "FluxTrainer" + + def validate(self, network_trainer, validation_settings=None): + training_loop = network_trainer["training_loop"] + network_trainer = network_trainer["network_trainer"] + + with torch.inference_mode(True): + image_tensors = network_trainer.sample_images( + network_trainer.accelerator, + network_trainer.args, + network_trainer.current_epoch.value, + network_trainer.global_step, + network_trainer.vae, + network_trainer.text_encoder, + network_trainer.unet, + validation_settings + ) + + trainer = { + "network_trainer": network_trainer, + "training_loop": training_loop, + } + return (trainer, (0.5 * (image_tensors + 1.0)).cpu().float(),) + +class VisualizeLoss: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "network_trainer": ("NETWORKTRAINER",), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("plot",) + FUNCTION = "draw" + CATEGORY = "FluxTrainer" + + def draw(self, network_trainer): + import matplotlib.pyplot as plt + import io + from PIL import Image + + # Example list of loss values + loss_values = network_trainer["network_trainer"].loss_recorder.loss_list + + # Create a plot + fig, ax = plt.subplots() + ax.plot(loss_values, label='Training Loss') + ax.set_xlabel('Epoch') + ax.set_ylabel('Loss') + ax.set_title('Training Loss Over Time') + ax.legend() + ax.grid(True) + + # Save the plot to a BytesIO object + buf = io.BytesIO() + plt.savefig(buf, format='png') + plt.close(fig) + buf.seek(0) + + # Convert the BytesIO object to a PIL Image + image = Image.open(buf).convert('RGB') + + # 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, @@ -751,11 +869,17 @@ NODE_CLASS_MAPPINGS = { "InitFluxTraining": InitFluxTraining, "FluxTrainModelSelect": FluxTrainModelSelect, "TrainDatasetConfig": TrainDatasetConfig, - "TrainLoop": TrainLoop + "FluxTrainLoop": FluxTrainLoop, + "VisualizeLoss": VisualizeLoss, + "FluxTrainValidate": FluxTrainValidate, + "FluxTrainValidationSettings": FluxTrainValidationSettings } NODE_DISPLAY_NAME_MAPPINGS = { "InitFluxTraining": "Init Flux Training", "FluxTrainModelSelect": "FluxTrain ModelSelect", "TrainDatasetConfig": "Train Dataset Config", - "TrainLoop": "Train Loop" + "FluxTrainLoop": "Flux Train Loop", + "VisualizeLoss": "Visualize Loss", + "FluxTrainValidate": "Flux Train Validate", + "FluxTrainValidationSettings": "Flux Train Validation Settings" } diff --git a/train_network.py b/train_network.py index 1194285..8a9745c 100644 --- a/train_network.py +++ b/train_network.py @@ -102,9 +102,9 @@ class NetworkTrainer: def load_target_model(self, args, weight_dtype, accelerator): text_encoder, vae, unet, _ = train_util.load_target_model(args, weight_dtype, accelerator) - # モデルに xformers とか memory efficient attention を組み込む + # Incorporate xformers or memory efficient attention into the model train_util.replace_unet_modules(unet, args.mem_eff_attn, args.xformers, args.sdpa) - if torch.__version__ >= "2.0.0": # PyTorch 2.0.0 以上対応のxformersなら以下が使える + if torch.__version__ >= "2.0.0": # If you have xformers compatible with PyTorch 2.0.0 or higher, you can use the following vae.set_use_memory_efficient_attention_xformers(args.xformers) return model_util.get_model_version_str_for_sd1_sd2(args.v2, args.v_parameterization), text_encoder, vae, unet @@ -262,7 +262,9 @@ class NetworkTrainer: latents_caching_strategy = self.get_latents_caching_strategy(args) strategy_base.LatentsCachingStrategy.set_strategy(latents_caching_strategy) - # データセットを準備する + pbar = ProgressBar(5) + + # Prepare the dataset if args.dataset_class is None: blueprint_generator = BlueprintGenerator(ConfigSanitizer(True, True, args.masked_loss, True)) if use_user_config: @@ -271,7 +273,7 @@ class NetworkTrainer: ignored = ["train_data_dir", "reg_data_dir", "in_json"] if any(getattr(args, attr) is not None for attr in ignored): logger.warning( - "ignoring the following options because config file is found: {0} / 設定ファイルが利用されるため以下のオプションは無視されます: {0}".format( + "ignoring the following options because config file is found: {0}".format( ", ".join(ignored) ) ) @@ -329,21 +331,24 @@ class NetworkTrainer: self.assert_extra_args(args, train_dataset_group) - # acceleratorを準備する + # prepare accelerator logger.info("preparing accelerator") accelerator = train_util.prepare_accelerator(args) - # mixed precisionに対応した型を用意しておき適宜castする + + # Prepare a type that supports mixed precision and cast it as appropriate. weight_dtype, save_dtype = train_util.prepare_dtype(args) vae_dtype = torch.float32 if args.no_half_vae else weight_dtype - # モデルを読み込む + # Load the model model_version, text_encoder, vae, unet = self.load_target_model(args, weight_dtype, accelerator) # text_encoder is List[CLIPTextModel] or CLIPTextModel text_encoders = text_encoder if isinstance(text_encoder, list) else [text_encoder] - # 差分追加学習のためにモデルを読み込む + pbar.update(1) + + # Load the model for incremental learning sys.path.append(os.path.dirname(__file__)) accelerator.print("import network module:", args.network_module) network_module = importlib.import_module(args.network_module) @@ -365,7 +370,8 @@ class NetworkTrainer: accelerator.print(f"all weights merged: {', '.join(args.base_weights)}") - # 学習を準備する + + # cache latents if cache_latents: vae.to(accelerator.device, dtype=vae_dtype) vae.requires_grad_(False) @@ -376,7 +382,6 @@ class NetworkTrainer: vae.to("cpu") clean_memory_on_device(accelerator.device) - # 必要ならテキストエンコーダーの出力をキャッシュする: Text Encoderはcpuまたはgpuへ移される # cache text encoder outputs if needed: Text Encoder is moved to cpu or gpu text_encoding_strategy = self.get_text_encoding_strategy(args) strategy_base.TextEncodingStrategy.set_strategy(text_encoding_strategy) @@ -386,6 +391,8 @@ class NetworkTrainer: strategy_base.TextEncoderOutputsCachingStrategy.set_strategy(text_encoder_outputs_caching_strategy) self.cache_text_encoder_outputs_if_needed(args, accelerator, unet, vae, text_encoders, train_dataset_group, weight_dtype) + pbar.update(1) + # prepare network net_kwargs = {} if args.network_args is not None: @@ -439,10 +446,10 @@ class NetworkTrainer: del t_enc network.enable_gradient_checkpointing() # may have no effect - # 学習に必要なクラスを準備する + # Prepare classes necessary for learning accelerator.print("prepare optimizer, data loader etc.") - # 後方互換性を確保するよ + # Ensure backward compatibility try: results = network.prepare_optimizer_params(args.text_encoder_lr, args.unet_lr, args.learning_rate) if type(results) is tuple: @@ -488,22 +495,22 @@ class NetworkTrainer: persistent_workers=args.persistent_data_loader_workers, ) - # 学習ステップ数を計算する - if args.max_train_epochs is not None: - args.max_train_steps = args.max_train_epochs * math.ceil( - len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps - ) - accelerator.print( - f"override steps. steps for {args.max_train_epochs} epochs is / 指定エポックまでのステップ数: {args.max_train_steps}" - ) + # # Calculate the number of learning steps + # if args.max_train_epochs is not None: + # args.max_train_steps = args.max_train_epochs * math.ceil( + # len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps + # ) + # accelerator.print( + # f"override steps. steps for {args.max_train_epochs} epochs is {args.max_train_steps}" + # ) - # データセット側にも学習ステップを送信 + # Send learning steps to the dataset side as well train_dataset_group.set_max_train_steps(args.max_train_steps) - # lr schedulerを用意する + # lr scheduler init lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes) - # 実験的機能:勾配も含めたfp16/bf16学習を行う モデル全体をfp16/bf16にする + # Experimental function: performs fp16/bf16 learning including gradients, sets the entire model to fp16/bf16 if args.full_fp16: assert ( args.mixed_precision == "fp16" @@ -520,10 +527,10 @@ class NetworkTrainer: unet_weight_dtype = te_weight_dtype = weight_dtype # Experimental Feature: Put base model into fp8 to save vram if args.fp8_base: - assert torch.__version__ >= "2.1.0", "fp8_base requires torch>=2.1.0 / fp8を使う場合はtorch>=2.1.0が必要です。" + assert torch.__version__ >= "2.1.0", "fp8_base requires torch>=2.1.0" assert ( args.mixed_precision != "no" - ), "fp8_base requires mixed precision='fp16' or 'bf16' / fp8を使う場合はmixed_precision='fp16'または'bf16'が必要です。" + ), "fp8_base requires mixed precision='fp16' or 'bf16'" accelerator.print("enable fp8 training.") unet_weight_dtype = torch.float8_e4m3fn te_weight_dtype = torch.float8_e4m3fn @@ -597,15 +604,17 @@ class NetworkTrainer: accelerator.unwrap_model(network).prepare_grad_etc(text_encoder, unet) - if not cache_latents: # キャッシュしない場合はVAEを使うのでVAEを準備する + if not cache_latents: # If you do not cache, VAE will be used, so enable VAE preparation. vae.requires_grad_(False) vae.eval() vae.to(accelerator.device, dtype=vae_dtype) - # 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする + # Experimental feature: Perform fp16 learning including gradients Apply a patch to PyTorch to enable grad scale in fp16 if args.full_fp16: train_util.patch_accelerator_for_fp16_training(accelerator) + pbar.update(1) + # before resuming make hook for saving/loading to save/load the network weights only def save_model_hook(models, weights, output_dir): # pop weights of other models than network to save only network weights @@ -651,10 +660,12 @@ class NetworkTrainer: accelerator.register_save_state_pre_hook(save_model_hook) accelerator.register_load_state_pre_hook(load_model_hook) - # resumeする + # resume from local or huggingface train_util.resume_from_local_or_hf_if_specified(accelerator, args) - # epoch数を計算する + pbar.update(1) + + # Calculate the number of epochs num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) if (args.save_n_epoch_ratio is not None) and (args.save_n_epoch_ratio > 0): @@ -664,17 +675,16 @@ class NetworkTrainer: # TODO: find a way to handle total batch size when there are multiple datasets total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps - accelerator.print("running training / 学習開始") - accelerator.print(f" num train images * repeats / 学習画像の数×繰り返し回数: {train_dataset_group.num_train_images}") - accelerator.print(f" num reg images / 正則化画像の数: {train_dataset_group.num_reg_images}") - accelerator.print(f" num batches per epoch / 1epochのバッチ数: {len(train_dataloader)}") - accelerator.print(f" num epochs / epoch数: {num_train_epochs}") + accelerator.print("running training") + accelerator.print(f" num train images * repeats: {train_dataset_group.num_train_images}") + accelerator.print(f" num reg images: {train_dataset_group.num_reg_images}") + accelerator.print(f" num batches per epoch: {len(train_dataloader)}") + accelerator.print(f" num epochs: {num_train_epochs}") accelerator.print( - f" batch size per device / バッチサイズ: {', '.join([str(d.batch_size) for d in train_dataset_group.datasets])}" + f" batch size per device: {', '.join([str(d.batch_size) for d in train_dataset_group.datasets])}" ) - # accelerator.print(f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}") - accelerator.print(f" gradient accumulation steps / 勾配を合計するステップ数 = {args.gradient_accumulation_steps}") - accelerator.print(f" total optimization steps / 学習ステップ数: {args.max_train_steps}") + accelerator.print(f" gradient accumulation steps: {args.gradient_accumulation_steps}") + accelerator.print(f" total optimization steps: {args.max_train_steps}") # TODO refactor metadata creation and move to util metadata = { @@ -813,9 +823,9 @@ class NetworkTrainer: # merge tag frequency: for ds_dir_name, ds_freq_for_dir in dataset.tag_frequency.items(): - # あるディレクトリが複数のdatasetで使用されている場合、一度だけ数える - # もともと繰り返し回数を指定しているので、キャプション内でのタグの出現回数と、それが学習で何度使われるかは一致しない - # なので、ここで複数datasetの回数を合算してもあまり意味はない + # If a directory is used by multiple datasets, count only once + # Since the number of repetitions is originally specified, the number of times a tag appears in the caption does not match the number of times it is used in training. + # Therefore, it is not very meaningful to add up the number of times for multiple datasets here. if ds_dir_name in tag_frequency: continue tag_frequency[ds_dir_name] = ds_freq_for_dir @@ -827,7 +837,7 @@ class NetworkTrainer: # conserving backward compatibility when using train_dataset_dir and reg_dataset_dir assert ( len(train_dataset_group.datasets) == 1 - ), f"There should be a single dataset but {len(train_dataset_group.datasets)} found. This seems to be a bug. / データセットは1個だけ存在するはずですが、実際には{len(train_dataset_group.datasets)}個でした。プログラムのバグかもしれません。" + ), f"There should be a single dataset but {len(train_dataset_group.datasets)} found. This seems to be a bug." dataset = train_dataset_group.datasets[0] @@ -938,7 +948,7 @@ class NetworkTrainer: epoch_to_start = initial_step // math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) initial_step = 0 # do not skip - global_step = 0 + noise_scheduler = self.get_noise_scheduler(args, accelerator.device) @@ -956,6 +966,8 @@ class NetworkTrainer: loss_recorder = train_util.LossRecorder() del train_dataset_group + pbar.update(1) + # callback for step start if hasattr(accelerator.unwrap_model(network), "on_step_start"): on_step_start = accelerator.unwrap_model(network).on_step_start @@ -995,12 +1007,13 @@ class NetworkTrainer: # For --sample_at_first #self.sample_images(accelerator, args, 0, global_step, accelerator.device, vae, tokenizers, text_encoder, unet) + self.global_step = 0 # training loop if initial_step > 0: # only if skip_until_initial_step is specified for skip_epoch in range(epoch_to_start): # skip epochs logger.info(f"skipping epoch {skip_epoch+1} because initial_step (multiplied) is {initial_step}") initial_step -= len(train_dataloader) - global_step = initial_step + self.global_step = initial_step # log device and dtype for each model logger.info(f"unet dtype: {unet_weight_dtype}, device: {unet.device}") @@ -1009,15 +1022,13 @@ class NetworkTrainer: clean_memory_on_device(accelerator.device) - - def training_loop(epoch, num_train_epochs, accelerator, network, text_encoder, + progress_bar = tqdm( + range(args.max_train_steps - initial_step), smoothing=0, disable=False, desc="steps" + ) + def training_loop(break_at_steps, epoch, num_train_epochs, accelerator, network, text_encoder, unet, vae, tokenizers, args, train_dataloader, initial_step, global_step, current_epoch, metadata, optimizer, lr_scheduler, loss_recorder): - progress_bar = tqdm( - range(args.max_train_steps - initial_step), smoothing=0, disable=not accelerator.is_local_main_process, desc="steps" - ) - #comfy_progress_bar = ProgressBar(args.max_train_steps - initial_step) - #for epoch in range(epoch_to_start, num_train_epochs): + accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}") current_epoch.value = epoch + 1 @@ -1161,14 +1172,14 @@ class NetworkTrainer: ) accelerator.log(logs, step=global_step) - if global_step >= args.max_train_steps: + if global_step >= break_at_steps: break if args.logging_dir is not None: logs = {"loss/epoch": loss_recorder.moving_average} accelerator.log(logs, step=epoch + 1) - - return global_step, current_epoch + self.global_step = global_step + return current_epoch.value self.epoch_to_start = epoch_to_start self.num_train_epochs = num_train_epochs @@ -1181,7 +1192,6 @@ class NetworkTrainer: self.args = args self.train_dataloader = train_dataloader self.initial_step = initial_step - self.global_step = global_step self.current_epoch = current_epoch self.metadata = metadata self.optimizer = optimizer