Allex21/LT
0
1# training with captions2# XXX dropped option: hypernetwork training3 4import argparse5import math6import os7from multiprocessing import Value8import toml9 10from tqdm import tqdm11 12import torch13from library import deepspeed_utils14from library.device_utils import init_ipex, clean_memory_on_device15 16init_ipex()17 18from accelerate.utils import set_seed19from diffusers import DDPMScheduler20 21from library.utils import setup_logging, add_logging_arguments22 23setup_logging()24import logging25 26logger = logging.getLogger(__name__)27 28import library.train_util as train_util29import library.config_util as config_util30from library.config_util import (31 ConfigSanitizer,32 BlueprintGenerator,33)34import library.custom_train_functions as custom_train_functions35from library.custom_train_functions import (36 apply_snr_weight,37 get_weighted_text_embeddings,38 prepare_scheduler_for_custom_training,39 scale_v_prediction_loss_like_noise_prediction,40 apply_debiased_estimation,41)42 43 44def train(args):45 train_util.verify_training_args(args)46 train_util.prepare_dataset_args(args, True)47 deepspeed_utils.prepare_deepspeed_args(args)48 setup_logging(args, reset=True)49 50 cache_latents = args.cache_latents51 52 if args.seed is not None:53 set_seed(args.seed) # 乱数系列を初期化する54 55 tokenizer = train_util.load_tokenizer(args)56 57 # データセットを準備する58 if args.dataset_class is None:59 blueprint_generator = BlueprintGenerator(ConfigSanitizer(False, True, False, True))60 if args.dataset_config is not None:61 logger.info(f"Load dataset config from {args.dataset_config}")62 user_config = config_util.load_user_config(args.dataset_config)63 ignored = ["train_data_dir", "in_json"]64 if any(getattr(args, attr) is not None for attr in ignored):65 logger.warning(66 "ignore following options because config file is found: {0} / 設定ファイルが利用されるため以下のオプションは無視されます: {0}".format(67 ", ".join(ignored)68 )69 )70 else:71 user_config = {72 "datasets": [73 {74 "subsets": [75 {76 "image_dir": args.train_data_dir,77 "metadata_file": args.in_json,78 }79 ]80 }81 ]82 }83 84 blueprint = blueprint_generator.generate(user_config, args, tokenizer=tokenizer)85 train_dataset_group = config_util.generate_dataset_group_by_blueprint(blueprint.dataset_group)86 else:87 train_dataset_group = train_util.load_arbitrary_dataset(args, tokenizer)88 89 current_epoch = Value("i", 0)90 current_step = Value("i", 0)91 ds_for_collator = train_dataset_group if args.max_data_loader_n_workers == 0 else None92 collator = train_util.collator_class(current_epoch, current_step, ds_for_collator)93 94 train_dataset_group.verify_bucket_reso_steps(64)95 96 if args.debug_dataset:97 train_util.debug_dataset(train_dataset_group)98 return99 if len(train_dataset_group) == 0:100 logger.error(101 "No data found. Please verify the metadata file and train_data_dir option. / 画像がありません。メタデータおよびtrain_data_dirオプションを確認してください。"102 )103 return104 105 if cache_latents:106 assert (107 train_dataset_group.is_latent_cacheable()108 ), "when caching latents, either color_aug or random_crop cannot be used / latentをキャッシュするときはcolor_augとrandom_cropは使えません"109 110 # acceleratorを準備する111 logger.info("prepare accelerator")112 accelerator = train_util.prepare_accelerator(args)113 114 # mixed precisionに対応した型を用意しておき適宜castする115 weight_dtype, save_dtype = train_util.prepare_dtype(args)116 vae_dtype = torch.float32 if args.no_half_vae else weight_dtype117 118 # モデルを読み込む119 text_encoder, vae, unet, load_stable_diffusion_format = train_util.load_target_model(args, weight_dtype, accelerator)120 121 # verify load/save model formats122 if load_stable_diffusion_format:123 src_stable_diffusion_ckpt = args.pretrained_model_name_or_path124 src_diffusers_model_path = None125 else:126 src_stable_diffusion_ckpt = None127 src_diffusers_model_path = args.pretrained_model_name_or_path128 129 if args.save_model_as is None:130 save_stable_diffusion_format = load_stable_diffusion_format131 use_safetensors = args.use_safetensors132 else:133 save_stable_diffusion_format = args.save_model_as.lower() == "ckpt" or args.save_model_as.lower() == "safetensors"134 use_safetensors = args.use_safetensors or ("safetensors" in args.save_model_as.lower())135 136 # Diffusers版のxformers使用フラグを設定する関数137 def set_diffusers_xformers_flag(model, valid):138 # model.set_use_memory_efficient_attention_xformers(valid) # 次のリリースでなくなりそう139 # pipeが自動で再帰的にset_use_memory_efficient_attention_xformersを探すんだって(;´Д`)140 # U-Netだけ使う時にはどうすればいいのか……仕方ないからコピって使うか141 # 0.10.2でなんか巻き戻って個別に指定するようになった(;^ω^)142 143 # Recursively walk through all the children.144 # Any children which exposes the set_use_memory_efficient_attention_xformers method145 # gets the message146 def fn_recursive_set_mem_eff(module: torch.nn.Module):147 if hasattr(module, "set_use_memory_efficient_attention_xformers"):148 module.set_use_memory_efficient_attention_xformers(valid)149 150 for child in module.children():151 fn_recursive_set_mem_eff(child)152 153 fn_recursive_set_mem_eff(model)154 155 # モデルに xformers とか memory efficient attention を組み込む156 if args.diffusers_xformers:157 accelerator.print("Use xformers by Diffusers")158 set_diffusers_xformers_flag(unet, True)159 else:160 # Windows版のxformersはfloatで学習できないのでxformersを使わない設定も可能にしておく必要がある161 accelerator.print("Disable Diffusers' xformers")162 set_diffusers_xformers_flag(unet, False)163 train_util.replace_unet_modules(unet, args.mem_eff_attn, args.xformers, args.sdpa)164 165 # 学習を準備する166 if cache_latents:167 vae.to(accelerator.device, dtype=vae_dtype)168 vae.requires_grad_(False)169 vae.eval()170 with torch.no_grad():171 train_dataset_group.cache_latents(vae, args.vae_batch_size, args.cache_latents_to_disk, accelerator.is_main_process)172 vae.to("cpu")173 clean_memory_on_device(accelerator.device)174 175 accelerator.wait_for_everyone()176 177 # 学習を準備する:モデルを適切な状態にする178 training_models = []179 if args.gradient_checkpointing:180 unet.enable_gradient_checkpointing()181 training_models.append(unet)182 183 if args.train_text_encoder:184 accelerator.print("enable text encoder training")185 if args.gradient_checkpointing:186 text_encoder.gradient_checkpointing_enable()187 training_models.append(text_encoder)188 else:189 text_encoder.to(accelerator.device, dtype=weight_dtype)190 text_encoder.requires_grad_(False) # text encoderは学習しない191 if args.gradient_checkpointing:192 text_encoder.gradient_checkpointing_enable()193 text_encoder.train() # required for gradient_checkpointing194 else:195 text_encoder.eval()196 197 if not cache_latents:198 vae.requires_grad_(False)199 vae.eval()200 vae.to(accelerator.device, dtype=vae_dtype)201 202 for m in training_models:203 m.requires_grad_(True)204 205 trainable_params = []206 if args.learning_rate_te is None or not args.train_text_encoder:207 for m in training_models:208 trainable_params.extend(m.parameters())209 else:210 trainable_params = [211 {"params": list(unet.parameters()), "lr": args.learning_rate},212 {"params": list(text_encoder.parameters()), "lr": args.learning_rate_te},213 ]214 215 # 学習に必要なクラスを準備する216 accelerator.print("prepare optimizer, data loader etc.")217 _, _, optimizer = train_util.get_optimizer(args, trainable_params=trainable_params)218 219 # dataloaderを準備する220 # DataLoaderのプロセス数:0 は persistent_workers が使えないので注意221 n_workers = min(args.max_data_loader_n_workers, os.cpu_count()) # cpu_count or max_data_loader_n_workers222 train_dataloader = torch.utils.data.DataLoader(223 train_dataset_group,224 batch_size=1,225 shuffle=True,226 collate_fn=collator,227 num_workers=n_workers,228 persistent_workers=args.persistent_data_loader_workers,229 )230 231 # 学習ステップ数を計算する232 if args.max_train_epochs is not None:233 args.max_train_steps = args.max_train_epochs * math.ceil(234 len(train_dataloader) / accelerator.num_processes / args.gradient_accumulation_steps235 )236 accelerator.print(237 f"override steps. steps for {args.max_train_epochs} epochs is / 指定エポックまでのステップ数: {args.max_train_steps}"238 )239 240 # データセット側にも学習ステップを送信241 train_dataset_group.set_max_train_steps(args.max_train_steps)242 243 # lr schedulerを用意する244 lr_scheduler = train_util.get_scheduler_fix(args, optimizer, accelerator.num_processes)245 246 # 実験的機能:勾配も含めたfp16学習を行う モデル全体をfp16にする247 if args.full_fp16:248 assert (249 args.mixed_precision == "fp16"250 ), "full_fp16 requires mixed precision='fp16' / full_fp16を使う場合はmixed_precision='fp16'を指定してください。"251 accelerator.print("enable full fp16 training.")252 unet.to(weight_dtype)253 text_encoder.to(weight_dtype)254 255 if args.deepspeed:256 if args.train_text_encoder:257 ds_model = deepspeed_utils.prepare_deepspeed_model(args, unet=unet, text_encoder=text_encoder)258 else:259 ds_model = deepspeed_utils.prepare_deepspeed_model(args, unet=unet)260 ds_model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(261 ds_model, optimizer, train_dataloader, lr_scheduler262 )263 training_models = [ds_model]264 else:265 # acceleratorがなんかよろしくやってくれるらしい266 if args.train_text_encoder:267 unet, text_encoder, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(268 unet, text_encoder, optimizer, train_dataloader, lr_scheduler269 )270 else:271 unet, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(unet, optimizer, train_dataloader, lr_scheduler)272 273 # 実験的機能:勾配も含めたfp16学習を行う PyTorchにパッチを当ててfp16でのgrad scaleを有効にする274 if args.full_fp16:275 train_util.patch_accelerator_for_fp16_training(accelerator)276 277 # resumeする278 train_util.resume_from_local_or_hf_if_specified(accelerator, args)279 280 # epoch数を計算する281 num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)282 num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)283 if (args.save_n_epoch_ratio is not None) and (args.save_n_epoch_ratio > 0):284 args.save_every_n_epochs = math.floor(num_train_epochs / args.save_n_epoch_ratio) or 1285 286 # 学習する287 total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps288 accelerator.print("running training / 学習開始")289 accelerator.print(f" num examples / サンプル数: {train_dataset_group.num_train_images}")290 accelerator.print(f" num batches per epoch / 1epochのバッチ数: {len(train_dataloader)}")291 accelerator.print(f" num epochs / epoch数: {num_train_epochs}")292 accelerator.print(f" batch size per device / バッチサイズ: {args.train_batch_size}")293 accelerator.print(294 f" total train batch size (with parallel & distributed & accumulation) / 総バッチサイズ(並列学習、勾配合計含む): {total_batch_size}"295 )296 accelerator.print(f" gradient accumulation steps / 勾配を合計するステップ数 = {args.gradient_accumulation_steps}")297 accelerator.print(f" total optimization steps / 学習ステップ数: {args.max_train_steps}")298 299 progress_bar = tqdm(range(args.max_train_steps), smoothing=0, disable=not accelerator.is_local_main_process, desc="steps")300 global_step = 0301 302 noise_scheduler = DDPMScheduler(303 beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", num_train_timesteps=1000, clip_sample=False304 )305 prepare_scheduler_for_custom_training(noise_scheduler, accelerator.device)306 if args.zero_terminal_snr:307 custom_train_functions.fix_noise_scheduler_betas_for_zero_terminal_snr(noise_scheduler)308 309 if accelerator.is_main_process:310 init_kwargs = {}311 if args.wandb_run_name:312 init_kwargs["wandb"] = {"name": args.wandb_run_name}313 if args.log_tracker_config is not None:314 init_kwargs = toml.load(args.log_tracker_config)315 accelerator.init_trackers(316 "finetuning" if args.log_tracker_name is None else args.log_tracker_name,317 config=train_util.get_sanitized_config_or_none(args),318 init_kwargs=init_kwargs,319 )320 321 # For --sample_at_first322 train_util.sample_images(accelerator, args, 0, global_step, accelerator.device, vae, tokenizer, text_encoder, unet)323 324 loss_recorder = train_util.LossRecorder()325 for epoch in range(num_train_epochs):326 accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}")327 current_epoch.value = epoch + 1328 329 for m in training_models:330 m.train()331 332 for step, batch in enumerate(train_dataloader):333 current_step.value = global_step334 with accelerator.accumulate(*training_models):335 with torch.no_grad():336 if "latents" in batch and batch["latents"] is not None:337 latents = batch["latents"].to(accelerator.device).to(dtype=weight_dtype)338 else:339 # latentに変換340 latents = vae.encode(batch["images"].to(dtype=vae_dtype)).latent_dist.sample().to(weight_dtype)341 latents = latents * 0.18215342 b_size = latents.shape[0]343 344 with torch.set_grad_enabled(args.train_text_encoder):345 # Get the text embedding for conditioning346 if args.weighted_captions:347 encoder_hidden_states = get_weighted_text_embeddings(348 tokenizer,349 text_encoder,350 batch["captions"],351 accelerator.device,352 args.max_token_length // 75 if args.max_token_length else 1,353 clip_skip=args.clip_skip,354 )355 else:356 input_ids = batch["input_ids"].to(accelerator.device)357 encoder_hidden_states = train_util.get_hidden_states(358 args, input_ids, tokenizer, text_encoder, None if not args.full_fp16 else weight_dtype359 )360 361 # Sample noise, sample a random timestep for each image, and add noise to the latents,362 # with noise offset and/or multires noise if specified363 noise, noisy_latents, timesteps, huber_c = train_util.get_noise_noisy_latents_and_timesteps(364 args, noise_scheduler, latents365 )366 367 # Predict the noise residual368 with accelerator.autocast():369 noise_pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample370 371 if args.v_parameterization:372 # v-parameterization training373 target = noise_scheduler.get_velocity(latents, noise, timesteps)374 else:375 target = noise376 377 if args.min_snr_gamma or args.scale_v_pred_loss_like_noise_pred or args.debiased_estimation_loss:378 # do not mean over batch dimension for snr weight or scale v-pred loss379 loss = train_util.conditional_loss(380 noise_pred.float(), target.float(), reduction="none", loss_type=args.loss_type, huber_c=huber_c381 )382 loss = loss.mean([1, 2, 3])383 384 if args.min_snr_gamma:385 loss = apply_snr_weight(loss, timesteps, noise_scheduler, args.min_snr_gamma, args.v_parameterization)386 if args.scale_v_pred_loss_like_noise_pred:387 loss = scale_v_prediction_loss_like_noise_prediction(loss, timesteps, noise_scheduler)388 if args.debiased_estimation_loss:389 loss = apply_debiased_estimation(loss, timesteps, noise_scheduler, args.v_parameterization)390 391 loss = loss.mean() # mean over batch dimension392 else:393 loss = train_util.conditional_loss(394 noise_pred.float(), target.float(), reduction="mean", loss_type=args.loss_type, huber_c=huber_c395 )396 397 accelerator.backward(loss)398 if accelerator.sync_gradients and args.max_grad_norm != 0.0:399 params_to_clip = []400 for m in training_models:401 params_to_clip.extend(m.parameters())402 accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm)403 404 optimizer.step()405 lr_scheduler.step()406 optimizer.zero_grad(set_to_none=True)407 408 # Checks if the accelerator has performed an optimization step behind the scenes409 if accelerator.sync_gradients:410 progress_bar.update(1)411 global_step += 1412 413 train_util.sample_images(414 accelerator, args, None, global_step, accelerator.device, vae, tokenizer, text_encoder, unet415 )416 417 # 指定ステップごとにモデルを保存418 if args.save_every_n_steps is not None and global_step % args.save_every_n_steps == 0:419 accelerator.wait_for_everyone()420 if accelerator.is_main_process:421 src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path422 train_util.save_sd_model_on_epoch_end_or_stepwise(423 args,424 False,425 accelerator,426 src_path,427 save_stable_diffusion_format,428 use_safetensors,429 save_dtype,430 epoch,431 num_train_epochs,432 global_step,433 accelerator.unwrap_model(text_encoder),434 accelerator.unwrap_model(unet),435 vae,436 )437 438 current_loss = loss.detach().item() # 平均なのでbatch sizeは関係ないはず439 if args.logging_dir is not None:440 logs = {"loss": current_loss}441 train_util.append_lr_to_logs(logs, lr_scheduler, args.optimizer_type, including_unet=True)442 accelerator.log(logs, step=global_step)443 444 loss_recorder.add(epoch=epoch, step=step, loss=current_loss)445 avr_loss: float = loss_recorder.moving_average446 logs = {"avr_loss": avr_loss} # , "lr": lr_scheduler.get_last_lr()[0]}447 progress_bar.set_postfix(**logs)448 449 if global_step >= args.max_train_steps:450 break451 452 if args.logging_dir is not None:453 logs = {"loss/epoch": loss_recorder.moving_average}454 accelerator.log(logs, step=epoch + 1)455 456 accelerator.wait_for_everyone()457 458 if args.save_every_n_epochs is not None:459 if accelerator.is_main_process:460 src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path461 train_util.save_sd_model_on_epoch_end_or_stepwise(462 args,463 True,464 accelerator,465 src_path,466 save_stable_diffusion_format,467 use_safetensors,468 save_dtype,469 epoch,470 num_train_epochs,471 global_step,472 accelerator.unwrap_model(text_encoder),473 accelerator.unwrap_model(unet),474 vae,475 )476 477 train_util.sample_images(accelerator, args, epoch + 1, global_step, accelerator.device, vae, tokenizer, text_encoder, unet)478 479 is_main_process = accelerator.is_main_process480 if is_main_process:481 unet = accelerator.unwrap_model(unet)482 text_encoder = accelerator.unwrap_model(text_encoder)483 484 accelerator.end_training()485 486 if is_main_process and (args.save_state or args.save_state_on_train_end):487 train_util.save_state_on_train_end(args, accelerator)488 489 del accelerator # この後メモリを使うのでこれは消す490 491 if is_main_process:492 src_path = src_stable_diffusion_ckpt if save_stable_diffusion_format else src_diffusers_model_path493 train_util.save_sd_model_on_train_end(494 args, src_path, save_stable_diffusion_format, use_safetensors, save_dtype, epoch, global_step, text_encoder, unet, vae495 )496 logger.info("model saved.")497 498 499def setup_parser() -> argparse.ArgumentParser:500 parser = argparse.ArgumentParser()501 502 add_logging_arguments(parser)503 train_util.add_sd_models_arguments(parser)504 train_util.add_dataset_arguments(parser, False, True, True)505 train_util.add_training_arguments(parser, False)506 deepspeed_utils.add_deepspeed_arguments(parser)507 train_util.add_sd_saving_arguments(parser)508 train_util.add_optimizer_arguments(parser)509 config_util.add_config_arguments(parser)510 custom_train_functions.add_custom_train_arguments(parser)511 512 parser.add_argument(513 "--diffusers_xformers", action="store_true", help="use xformers by diffusers / Diffusersでxformersを使用する"514 )515 parser.add_argument("--train_text_encoder", action="store_true", help="train text encoder / text encoderも学習する")516 parser.add_argument(517 "--learning_rate_te",518 type=float,519 default=None,520 help="learning rate for text encoder, default is same as unet / Text Encoderの学習率、デフォルトはunetと同じ",521 )522 parser.add_argument(523 "--no_half_vae",524 action="store_true",525 help="do not use fp16/bf16 VAE in mixed precision (use float VAE) / mixed precisionでも fp16/bf16 VAEを使わずfloat VAEを使う",526 )527 528 return parser529 530 531if __name__ == "__main__":532 parser = setup_parser()533 534 args = parser.parse_args()535 train_util.verify_command_line_training_args(args)536 args = train_util.read_config_from_file(args, parser)537 538 train(args)539 