CoolFace
Apppublic

acmyu/KeyframesAI

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
train2.py522 linesDownload Raw Back to root
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"""