acmyu/KeyframesAI
0
1import logging2import math3import os4from typing import Any, Dict, List, Optional, Tuple, Union5from diffusers.models.controlnet import ControlNetConditioningEmbedding6import torch7from torch import nn8import torch.nn.functional as F9import torch.utils.checkpoint10import transformers11from accelerate import Accelerator12from accelerate.logging import get_logger13from accelerate.utils import ProjectConfiguration, set_seed14 15from tqdm.auto import tqdm16from src.configs.stage2_config import args17 18import diffusers19from diffusers import (20 AutoencoderKL,21 DDPMScheduler,22)23from diffusers.optimization import get_scheduler24from diffusers.utils import check_min_version, is_wandb_available25from src.dataset.stage2_dataset import InpaintDataset, InpaintCollate_fn26from transformers import CLIPVisionModelWithProjection27from transformers import Dinov2Model28from src.models.stage2_inpaint_unet_2d_condition import Stage2_InapintUNet2DConditionModel29 30 31# Will error if the minimal version of diffusers is not installed. Remove at your own risks.32check_min_version("0.18.0.dev0")33 34logger = get_logger(__name__)35 36 37class ImageProjModel_p(torch.nn.Module):38 """SD model with image prompt"""39 40 def __init__(self, in_dim, hidden_dim, out_dim, dropout = 0.):41 super().__init__()42 43 self.net = nn.Sequential(44 nn.Linear(in_dim, hidden_dim),45 nn.GELU(),46 nn.Dropout(dropout),47 nn.LayerNorm(hidden_dim),48 nn.Linear(hidden_dim, out_dim),49 nn.Dropout(dropout)50 )51 52 def forward(self, x): 53 return self.net(x)54 55class ImageProjModel_g(torch.nn.Module):56 """SD model with image prompt"""57 58 def __init__(self, in_dim, hidden_dim, out_dim, dropout = 0.):59 super().__init__()60 61 self.net = nn.Sequential(62 nn.Linear(in_dim, hidden_dim),63 nn.GELU(),64 nn.Dropout(dropout),65 nn.LayerNorm(hidden_dim),66 nn.Linear(hidden_dim, out_dim),67 nn.Dropout(dropout)68 )69 70 def forward(self, x): # b, 257,128071 return self.net(x)72 73 74class SDModel(torch.nn.Module):75 """SD model with image prompt"""76 def __init__(self, unet) -> None:77 super().__init__()78 self.image_proj_model_p = ImageProjModel_p(in_dim=1536, hidden_dim=768, out_dim=1024)79 80 self.unet = unet81 self.pose_proj = ControlNetConditioningEmbedding(82 conditioning_embedding_channels=320,83 block_out_channels=(16, 32, 96, 256),84 conditioning_channels=3)85 86 87 def forward(self, noisy_latents, timesteps, simg_f_p, timg_f_g, pose_f):88 89 extra_image_embeddings_p = self.image_proj_model_p(simg_f_p)90 extra_image_embeddings_g = timg_f_g91 92 print(extra_image_embeddings_p.size())93 print(extra_image_embeddings_g.size())94 95 encoder_image_hidden_states = torch.cat([extra_image_embeddings_p ,extra_image_embeddings_g], dim=1)96 pose_cond = self.pose_proj(pose_f)97 98 pred_noise = self.unet(noisy_latents, timesteps, class_labels=timg_f_g, encoder_hidden_states=encoder_image_hidden_states,my_pose_cond=pose_cond).sample99 return pred_noise100 101 102 103 104def load_training_checkpoint2(model, load_dir, tag=None, **kwargs):105 """Utility function for checkpointing model + optimizer dictionaries106 The main purpose for this is to be able to resume training from that instant again107 """108 """109 checkpoint_state_dict= torch.load(load_dir, map_location="cpu")110 111 112 print(checkpoint_state_dict.keys())113 epoch = 0114 last_global_step = 0115 116 epoch = checkpoint_state_dict["epoch"]117 last_global_step = checkpoint_state_dict["last_global_step"]118 # TODO optimizer lr, and loss state119 120 121 weight_dict = checkpoint_state_dict["module"]122 new_weight_dict = {f"module.{key}": value for key, value in weight_dict.items()}123 model.load_state_dict(new_weight_dict)124 del checkpoint_state_dict125 126 return model, epoch, last_global_step127 """128 129 image_proj_model_p_dict = {}130 pose_proj_dict = {}131 unet_dict = {}132 model_sd = torch.load(load_dir, map_location="cpu")["module"]133 134 for k in model_sd.keys():135 if k.startswith("pose_proj"):136 137 pose_proj_dict[k.replace("pose_proj.", "")] = model_sd[k]138 139 elif k.startswith("image_proj_model_p"):140 image_proj_model_p_dict[k.replace("image_proj_model_p.", "")] = model_sd[k]141 142 elif k.startswith("unet"):143 unet_dict[k.replace("unet.", "")] = model_sd[k]144 145 else:146 print(k)147 148 model.pose_proj.load_state_dict(pose_proj_dict)149 model.image_proj_model_p.load_state_dict(image_proj_model_p_dict)150 model.unet.load_state_dict(unet_dict)151 152 return model, 0, 0153 154 155def load_training_checkpoint(model, load_dir, tag=None, **kwargs):156 model_sd = torch.load(load_dir, map_location="cpu")["module"]157 158 image_proj_model_dict = {}159 pose_proj_dict = {}160 unet_dict = {}161 for k in model_sd.keys():162 if k.startswith("pose_proj"):163 pose_proj_dict[k.replace("pose_proj.", "")] = model_sd[k]164 165 elif k.startswith("image_proj_model"):166 image_proj_model_dict[k.replace("image_proj_model.", "")] = model_sd[k]167 168 169 elif k.startswith("unet"):170 unet_dict[k.replace("unet.", "")] = model_sd[k]171 else:172 print(k)173 174 model.pose_proj.load_state_dict(pose_proj_dict)175 model.image_proj_model_p.load_state_dict(image_proj_model_dict)176 model.unet.load_state_dict(unet_dict)177 178 return model, 0, 0179 180 181def checkpoint_model(checkpoint_folder, ckpt_id, model, epoch, last_global_step, **kwargs):182 """Utility function for checkpointing model + optimizer dictionaries183 The main purpose for this is to be able to resume training from that instant again184 """185 checkpoint_state_dict = {186 "epoch": epoch,187 "last_global_step": last_global_step,188 }189 # Add extra kwargs too190 checkpoint_state_dict.update(kwargs)191 192 success = model.save_checkpoint(checkpoint_folder, ckpt_id, checkpoint_state_dict)193 status_msg = f"checkpointing: checkpoint_folder={checkpoint_folder}, ckpt_id={ckpt_id}"194 if success:195 logging.info(f"Success {status_msg}")196 else:197 logging.warning(f"Failure {status_msg}")198 return199 200 201 202def main():203 logging_dir = 'outputs/logging'204 205 accelerator = Accelerator(206 log_with=args.report_to,207 project_dir=logging_dir,208 mixed_precision=args.mixed_precision,209 gradient_accumulation_steps=args.gradient_accumulation_steps210 )211 212 # Make one log on every process with the configuration for debugging.213 logging.basicConfig(214 format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",215 datefmt="%m/%d/%Y %H:%M:%S",216 level=logging.INFO, )217 logger.info(accelerator.state, main_process_only=False)218 if accelerator.is_local_main_process:219 transformers.utils.logging.set_verbosity_warning()220 diffusers.utils.logging.set_verbosity_info()221 else:222 transformers.utils.logging.set_verbosity_error()223 diffusers.utils.logging.set_verbosity_error()224 225 # If passed along, set the training seed now.226 set_seed(42)227 228 # Handle the repository creation229 if accelerator.is_main_process:230 os.makedirs('outputs', exist_ok=True)231 232 233 234 # Load scheduler235 noise_scheduler = DDPMScheduler.from_pretrained("stabilityai/stable-diffusion-2-1-base", subfolder="scheduler")236 237 # Load model238 image_encoder_p = Dinov2Model.from_pretrained('facebook/dinov2-giant')239 image_encoder_g = CLIPVisionModelWithProjection.from_pretrained('laion/CLIP-ViT-H-14-laion2B-s32B-b79K')#("openai/clip-vit-base-patch32")240 241 vae = AutoencoderKL.from_pretrained("stabilityai/stable-diffusion-2-1-base", subfolder="vae")242 243 unet = Stage2_InapintUNet2DConditionModel.from_pretrained("stabilityai/stable-diffusion-2-1-base", torch_dtype=torch.float16,subfolder="unet",in_channels=9, low_cpu_mem_usage=False, ignore_mismatched_sizes=True)244 """245 unet = Stage2_InapintUNet2DConditionModel.from_pretrained("stabilityai/stable-diffusion-2-1-base", subfolder="unet",246 in_channels=9, class_embed_type="projection" ,projection_class_embeddings_input_dim=1024,247 low_cpu_mem_usage=False, ignore_mismatched_sizes=True)248 """249 image_encoder_p.requires_grad_(False)250 image_encoder_g.requires_grad_(False)251 vae.requires_grad_(False)252 253 sd_model = SDModel(unet=unet)254 sd_model.train()255 256 257 if args.gradient_checkpointing:258 sd_model.enable_gradient_checkpointing()259 260 261 # Enable TF32 for faster training on Ampere GPUs,262 # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices263 if args.allow_tf32:264 torch.backends.cuda.matmul.allow_tf32 = True265 266 learning_rate = 1e-4267 train_batch_size = 1268 269 270 # Optimizer creation271 params_to_optimize = sd_model.parameters()272 optimizer = torch.optim.AdamW(273 params_to_optimize,274 lr=learning_rate,275 betas=(args.adam_beta1, args.adam_beta2),276 weight_decay=args.adam_weight_decay,277 eps=args.adam_epsilon,278 )279 280 dataset = InpaintDataset(281 [{282 "source_image": "sm.png",283 "target_image": "target.png",284 }], 285 'imgs/', 286 size=(args.img_width, args.img_height), # w h287 imgp_drop_rate=0.1,288 imgg_drop_rate=0.1,289 )290 291 """292 dataset = InpaintDataset(293 args.json_path,294 args.image_root_path,295 size=(args.img_width, args.img_height), # w h296 imgp_drop_rate=0.1,297 imgg_drop_rate=0.1,298 )299 """300 301 train_sampler = torch.utils.data.distributed.DistributedSampler(302 dataset, num_replicas=accelerator.num_processes, rank=accelerator.process_index, shuffle=True)303 304 train_dataloader = torch.utils.data.DataLoader(305 dataset,306 sampler=train_sampler,307 collate_fn=InpaintCollate_fn,308 batch_size=train_batch_size,309 num_workers=2,)310 311 312 # Scheduler and math around the number of training steps.313 overrode_max_train_steps = False314 num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)315 if args.max_train_steps is None:316 args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch317 overrode_max_train_steps = True318 319 lr_scheduler = get_scheduler(320 args.lr_scheduler,321 optimizer=optimizer,322 num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes,323 num_training_steps=args.max_train_steps * accelerator.num_processes,324 num_cycles=args.lr_num_cycles,325 power=args.lr_power,326 )327 328 # Prepare everything with our `accelerator`.329 sd_model, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(sd_model, optimizer, train_dataloader, lr_scheduler)330 331 # For mixed precision training we cast the text_encoder and vae weights to half-precision332 # as these models are only used for inference, keeping weights in full precision is not required.333 weight_dtype = torch.float32334 """335 if accelerator.mixed_precision == "fp16":336 weight_dtype = torch.float16337 elif accelerator.mixed_precision == "bf16":338 weight_dtype = torch.bfloat16339 """340 341 # Move vae, unet and text_encoder to device and cast to weight_dtype342 vae.to(accelerator.device, dtype=weight_dtype)343 unet.to(accelerator.device, dtype=weight_dtype)344 image_encoder_p.to(accelerator.device, dtype=weight_dtype)345 image_encoder_g.to(accelerator.device, dtype=weight_dtype)346 347 # We need to recalculate our total training steps as the size of the training dataloader may have changed.348 num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps)349 if overrode_max_train_steps:350 args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch351 # Afterwards we recalculate our number of training epochs352 args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)353 354 355 356 # Train!357 total_batch_size = (358 train_batch_size359 * accelerator.num_processes360 * args.gradient_accumulation_steps361 )362 363 logger.info("***** Running training *****")364 logger.info(f" Num batches each epoch = {len(train_dataloader)}")365 logger.info(f" Num Epochs = {args.num_train_epochs}")366 logger.info(f" Instantaneous batch size per device = {train_batch_size}")367 logger.info(368 f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}"369 )370 logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")371 logger.info(f" Total optimization steps = {args.max_train_steps}")372 373 374 if args.resume_from_checkpoint:375 # New Code #376 # Loads the DeepSpeed checkpoint from the specified path377 prior_model, last_epoch, last_global_step = load_training_checkpoint(378 sd_model,379 args.resume_from_checkpoint,380 **{"load_optimizer_states": True, "load_lr_scheduler_states": True},381 )382 accelerator.print(f"Resumed from checkpoint: {args.resume_from_checkpoint}, global step: {last_global_step}")383 starting_epoch = last_epoch384 global_steps = last_global_step385 sd_model = sd_model386 else:387 global_steps = 0388 starting_epoch = 0389 sd_model = sd_model390 391 progress_bar = tqdm(range(global_steps, args.max_train_steps), initial=global_steps, desc="Steps",392 # Only show the progress bar once on each machine.393 disable=not accelerator.is_local_main_process, )394 395 bsz = train_batch_size396 397 398 for epoch in range(starting_epoch, args.num_train_epochs):399 for step, batch in enumerate(train_dataloader):400 with accelerator.accumulate(sd_model):401 with torch.no_grad():402 # Convert images to latent space403 latents = vae.encode(batch["source_target_image"].to(dtype=weight_dtype)).latent_dist.sample()404 latents = latents * vae.config.scaling_factor405 406 # Get the masked image latents407 masked_latents = vae.encode(batch["vae_source_mask_image"].to(dtype=weight_dtype)).latent_dist.sample()408 masked_latents = masked_latents * vae.config.scaling_factor409 410 # mask411 mask1 = torch.ones((bsz, 1, int(args.img_height / 8), int(args.img_width / 8))).to(accelerator.device, dtype=weight_dtype)412 mask0 = torch.zeros((bsz, 1, int(args.img_height / 8), int(args.img_width / 8))).to(accelerator.device, dtype=weight_dtype)413 mask = torch.cat([mask1, mask0], dim=3)414 # Get the image embedding for conditioning415 cond_image_feature_p = image_encoder_p(batch["source_image"].to(accelerator.device, dtype=weight_dtype))416 cond_image_feature_p = (cond_image_feature_p.last_hidden_state)417 418 419 cond_image_feature_g = image_encoder_g(batch["target_image"].to(accelerator.device, dtype=weight_dtype), ).image_embeds420 cond_image_feature_g =cond_image_feature_g.unsqueeze(1)421 422 # Sample noise that we'll add to the latents423 noise = torch.randn_like(latents)424 if args.noise_offset:425 # https://www.crosslabs.org//blog/diffusion-with-offset-noise426 noise += args.noise_offset * torch.randn(427 (latents.shape[0], latents.shape[1], 1, 1), device=latents.device428 )429 430 # Sample a random timestep for each image431 timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (train_batch_size,),device=latents.device, )432 timesteps = timesteps.long()433 434 # Add noise to the latents according to the noise magnitude at each timestep (this is the forward diffusion process)435 noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)436 437 noisy_latents = torch.cat([noisy_latents, mask, masked_latents], dim=1)438 # Get the text embedding for conditioning439 440 441 cond_pose = batch["source_target_pose"].to(dtype=weight_dtype)442 443 print(noisy_latents.size())444 print(cond_image_feature_p.size())445 print(cond_image_feature_g.size())446 print(cond_pose.size())447 448 # Predict the noise residual449 model_pred = sd_model(noisy_latents, timesteps, cond_image_feature_p,cond_image_feature_g, cond_pose, )450 451 # Get the target for loss depending on the prediction type452 if noise_scheduler.config.prediction_type == "epsilon":453 target = noise454 elif noise_scheduler.config.prediction_type == "v_prediction":455 target = noise_scheduler.get_velocity(latents, noise, timesteps)456 else:457 raise ValueError(458 f"Unknown prediction type {noise_scheduler.config.prediction_type}"459 )460 461 loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean")462 463 accelerator.backward(loss)464 if accelerator.sync_gradients:465 params_to_clip = sd_model.parameters()466 accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm)467 optimizer.step()468 lr_scheduler.step()469 optimizer.zero_grad(set_to_none=args.set_grads_to_none)470 471 # Checks if the accelerator has performed an optimization step behind the scenes472 if accelerator.sync_gradients:473 progress_bar.update(1)474 global_steps += 1475 476 if global_steps % args.checkpointing_steps == 0:477 """478 checkpoint_model(479 args.output_dir, global_steps, sd_model, epoch, global_steps480 )481 """482 checkpoint_state_dict = {483 "epoch": epoch,484 "module": sd_model.state_dict(),485 }486 print(list(sd_model.state_dict().keys())[:20])487 torch.save(checkpoint_state_dict, "fine_tuned_pcdms.pt")488 489 logs = {"loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]}490 print(logs)491 progress_bar.set_postfix(**logs)492 493 if global_steps >= args.max_train_steps:494 break495 496 # Create the pipeline using the trained modules and save it.497 accelerator.wait_for_everyone()498 accelerator.end_training()499 500 501if __name__ == "__main__":502 503 main()504 505"""506python train2.py \507 --pretrained_model_name_or_path="stabilityai/stable-diffusion-2-1-base" \508 --output_dir="out/" \509 --img_height=512 \510 --img_width=512 \511 --learning_rate=1e-4 \512 --train_batch_size=8 \513 --max_train_steps=1000000 \514 --mixed_precision="fp16" \515 --checkpointing_steps=1 \516 --noise_offset=0.1 \517 --lr_warmup_steps 5000 \518 --seed 42 \519 --resume_from_checkpoint s2_512.pt520 521 522"""