From 7bebf61aad028e4d00832bc1d9bdccff9992ff1f Mon Sep 17 00:00:00 2001 From: Shmuel Ronen <80190186+ShmuelRonen@users.noreply.github.com> Date: Sat, 22 Mar 2025 16:12:09 +0200 Subject: [PATCH] LatentSync_Wrapper_1.5 --- scripts/__pycache__/inference.cpython-312.pyc | Bin 0 -> 7383 bytes scripts/inference.py | 162 ++++++ scripts/train_syncnet.py | 332 +++++++++++ scripts/train_unet.py | 516 ++++++++++++++++++ 4 files changed, 1010 insertions(+) create mode 100644 scripts/__pycache__/inference.cpython-312.pyc create mode 100644 scripts/inference.py create mode 100644 scripts/train_syncnet.py create mode 100644 scripts/train_unet.py diff --git a/scripts/__pycache__/inference.cpython-312.pyc b/scripts/__pycache__/inference.cpython-312.pyc new file mode 100644 index 0000000000000000000000000000000000000000..3b718d4c55e77ee84f8c24c463f35772da83cc31 GIT binary patch literal 7383 zcmb_BTWlNGl|ykz4&Rh0O4QS|O;wVsYSLun|!c4lZr z4CN^80+ANkm4$xDMqAhc3PgnjthxwLFR&l^>JwNDLwU)a%IFqrpzA+vVr+o)qkHZQ zFGqH$+HH3X>z;e=Ip>~x?s?7d*G{L6faf1>{QLON`~>k|Xo4S81^E0_9YHJ+EWsKQ z#HcnJMh!5pOVlNeqeg>;Yzwk-2hX6M~r&7de%C*mmq}S*AW~M4`PM%ais!|dO2@Aq~kHf+BhF;Uv;b) zHH?AShRz9eu2uJQ={zq;SO0=^4KGON<(k;W7dVHHYi65XfUX(j`JbWd$5zwFs(oAF z28{6o`U-})miQ2c^>M|Z!K3ZK=i%D9mKE?F^caX~!?YpP@)%9dhYX7K8%b`QIgw7q zLUoFBI4h;OR5Zt$^47VvA(OtE{0m?e`&oerUuT-!c3}I;i1`chW{I@hKgQphFGX zujt^wvf?fb>0`WJhcSd*4M*#bz`G&w>LGj;ho!LvnHqfU>Z zlzYp16UOS}j8Q@Y)rU22cLP7H!`1@bb$ZKo*4l=i0{T|Gn?8P@c#6KdP}9dcs>@a2 zIUcd2N>~M?00!~7@Hj(0N7`=8sgHS*wX!zWzCvoa8PjuQ)QS@m)&VZ;kipH-w=)mAA32t8IEA1DWs-YYg74 zOSJqyNz~(|{n^{J5>hYLEgr^@KIW@lF=NjY@d14WL$!(yyF+aYg3qyjt=mVQ9AI~{ zyH+Z9GviEs{m6?m|FS)+`*B2!L0zBm%+%*S@o(r-F$8Dr#j<8R(wkMhJMYSS^7hp*wp%T3 zc2AYJ0W|q0=GI3xG6^>>G)b2*t~MQ(V<1wE66#UZruF=f^wnonY*tc|<7^;npe+#ecBiqjoV87bNb8GpyAwH(dz!2M z+d2(~{%XqWSdu=j4zAd>8MM8`Zla$M4O4Xn;wrSk2I5xT2S0z?c$Jtk)L2=;D^eep z7bJw98ON5|*^zfVeS>YaQ?GthUM4*H1yHQOqx$gc%?2JRxD9?2jIs=#VX>CKNmP5K+n#n#sgJcfaRfu zm(s#?fRO?hQK!$n&5IJ9p3M2GRJ0~692e86oPEngKxu#c%bW3kmA*H3Tycs~Cc&{Z zdTC&UaoX~4EpIAkj-TTH+X$yatw_kG*y(735k)%1Bsqn;&Pfc-rC7zOLU>EXrcEVA zkQ93~!DncZNoEq9;!S0fv?MTmikA2!CrVsKR6OYQgh^y37+Olx=@h3p)CD>ni@}oO zR2e`sBvn*g;v@*n2$QtNZkT$E><_0m^C^}~Nt$t=wp2_=C&BD3Cj}ub1QKBUr`8V# zLKdMBqzDb@59-MUANtdb{=ioNLO{JNkKT{qMAb_Z70c8FFJ?F)M~O)$kqBobeim$% zv%pnN!z{NaDx|?+j3j|-@SaAqd@_*Cic;V@7dSX{C;<8l9y*|q(Jae^NX3>!FLE?; zI>pR#H~1*0*pS&`0@z$s>Qz6YW8e&Qf=i7{6S)>qVy^QEUYch3h{*QdV7T6x)SDF1 z%by+mBRWB|(sYI!|KmTu^V!-PKR>P*qnT{TAnbxe0gXvrh0Tq?syMYBD({|(3v8O^ zMX(p0O>vTdQt;zH!jp5IPQj}ulL(-9ORz2(PG%0|taOqQClMI3D>gL}wA0u$x$4PW z8a|zsv^BRjN1Q4`Km}5B7d{o^1W<;n+_!5y%d?C+M-HnvrZ|3lLefsCn5K9ZS#Cl_ zDsELGpJc{4Rc>fNsY{CrshL-e7E#!b_HukP3-^=cR1XnQr>2-x(+L4}nTI>P!SD&@ zI_&8i2ba|@hmJCtiuDys3_MHfAF2R@1BxjlKx`{s{SXlDl`~r)eAPIzFxN#@A&RDB ze1fBCg^WW^P^cs`$pIB8jG|W|(-{u_ve{%t4PM2K{2F2bt_5O>V-@Skx1${RAP}2W zM^d~s{7>6@^LJusC93jS_v;%QE z!N;V^OhDCAfB*qoLL38&gme;=Iv_>KQ6Z6^2$V{}Yp$bymJm=ktXKptlVHHx-MX>V z2p7=T7{w-VVmgsUM-yHGM93?krbQtkS3y+3^~5s$L$sL9WYPkhN$5ta_GnfB=Yh-O z1*CrE`-$dsU&)cG^?^stzS_j2^Jr+{6q0L*sd1nfIHjSI0cbsN1@vMez~?$b*r!s{ z3YDpRj}cLF34HbW^76>xu@mWJZ2Hn^dI~aGdP>Ceu_)vVxabt5hAOi6EQ3A9brn?3o6w91ao z+tS^4@4PE_Mr2#xI@!OmEi!oTVyWwBq3fs|JSN+YuamEW5?bi@RPaHXNBfv*BxAy0Um>X>4)qUeB_$cJ&i`(RX(CJTN+3Rm+6e?V%0I zDSP_XI_17sW!LaJ_1XqyTd>dDZ=GMh_<$M!hMl`VvV3Uy$nl|L`I>z6Vli+@raCs7 z{7d7D<4cLf#Jyu{M~h9vbM{B>`h}_asfC;KH|5U3b@zcyGZBcC+jo|C_LV!k$~%JP z?w)cmR1N^9tGsK^|M=_ym+WZ&yp?eJH|l+IQ_u2+_4@v`H_M*JCF`P9Zr!&$yvE$J zu6quBVI&#`o=`+X)2+#SgX{I(%l+F>?FFSxHrWxpcTEl+kZmumlS7+zPLpk;ac8Nq zztGsf7A`gp&sjG+3Ci_;{=Iz3(p|80FAsid32$`lxO?T!mAhki#+G~5r2BvKi_`KI zrr2@4K(&^w?uGXG_L4PJu!fc|d}{5hN&W3&2LsfPTs!4muNPfsWa`YL#+_wvSJ~|= zds~)XS$w7B4Hdkha%%l_n!6#Oe&pdnXiHorNOSbNUty|uELVok|y6wvB ziE^NO_QdB9S-x}U@5KY07Z%@IuXZe!s+f(!n&Ypkd?kqdp<)+=`=B{!>8|dOO ze_)E+3c_Z8av%2g&`YuVMfFgZV2as>;nAlP-`o%wf4?yiEntL*Q* zYrkXHdb;t}Jo_W3SZ ziS>x-;>(uxBc@Ac%LA+R(oy4shQOsm#s`N?04q@cP_m-wTtja!YBnmptD3HJ^;!ZJ zBWNW&%_I}>nVdv@jbdTB7?Vv%IbUyYHTluiQhPB}Crp4s3MFuVn}wzs8u;E`y){8I zy*bHw@EUGUa`juqszvEB>J|~@)7{88VBy(3o-_CMqH@Ro7V>SVL!qEdWf(!^1oRt$ zVojq@{pdhy#T3w24aLgxu~-&JM78W{CB7-4rg`*9Kde?5F|4*>sM;!BTNbMYCJYTX zBBwD%tG*s48s*@(0}vO6GMW&zN*vfCqEgAisx1akiwVCG2&gYm>j%0Y)UZH*;u;pl zv#5D^Lp_sX*3PCl7?!2ssOp+x6eU41q6WyMc1DU#B|`+Fa1|YZM6I~udG@Z=(*9%k zB?4{DAwud193Fo9I3~OcQ}i|>!nbu|(`YalKJO$9p5GAGhs2KG5RTsx0}qLV4~gN2 zMAt*&@Mq>JgTcHh841!cm;2P%Qg#OC-YPkR1!u76+%tQsOxj9hbAfD@{rlGjd_XpBb`WH%a1*Tcf0tvRU;qFB literal 0 HcmV?d00001 diff --git a/scripts/inference.py b/scripts/inference.py new file mode 100644 index 0000000..8ca2b62 --- /dev/null +++ b/scripts/inference.py @@ -0,0 +1,162 @@ +# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import argparse +import os +from omegaconf import OmegaConf +import torch +from diffusers import AutoencoderKL, DDIMScheduler +from latentsync.models.unet import UNet3DConditionModel +from latentsync.pipelines.lipsync_pipeline import LipsyncPipeline +from accelerate.utils import set_seed +from latentsync.whisper.audio2feature import Audio2Feature + + +def main(config, args): + if not os.path.exists(args.video_path): + raise RuntimeError(f"Video path '{args.video_path}' not found") + if not os.path.exists(args.audio_path): + raise RuntimeError(f"Audio path '{args.audio_path}' not found") + + # Check if the GPU supports float16 + is_fp16_supported = torch.cuda.is_available() and torch.cuda.get_device_capability()[0] > 7 + dtype = torch.float16 if is_fp16_supported else torch.float32 + + print(f"Input video path: {args.video_path}") + print(f"Input audio path: {args.audio_path}") + print(f"Loaded checkpoint path: {args.inference_ckpt_path}") + + # Use relative path for scheduler configuration + current_dir = os.path.dirname(os.path.abspath(__file__)) + scheduler_path = os.path.join(current_dir, "..", "configs", "scheduler") + + # Check if scheduler directory exists + if not os.path.exists(scheduler_path): + print(f"Creating scheduler directory at {scheduler_path}") + os.makedirs(scheduler_path, exist_ok=True) + + # Create scheduler config file if it doesn't exist + scheduler_config_file = os.path.join(scheduler_path, "scheduler_config.json") + config_file = os.path.join(scheduler_path, "config.json") + + if not os.path.exists(scheduler_config_file): + # Default scheduler config + scheduler_config = { + "_class_name": "DDIMScheduler", + "beta_end": 0.012, + "beta_schedule": "scaled_linear", + "beta_start": 0.00085, + "clip_sample": False, + "num_train_timesteps": 1000, + "set_alpha_to_one": False, + "steps_offset": 1, + "trained_betas": None, + "skip_prk_steps": True + } + + import json + with open(scheduler_config_file, 'w') as f: + json.dump(scheduler_config, f, indent=2) + + # Also create a copy as config.json for compatibility + with open(config_file, 'w') as f: + json.dump(scheduler_config, f, indent=2) + + print(f"Loading scheduler from: {scheduler_path}") + try: + scheduler = DDIMScheduler.from_pretrained(scheduler_path) + except Exception as e: + print(f"Error loading scheduler: {e}") + # Fallback to creating scheduler directly + scheduler = DDIMScheduler( + beta_start=0.00085, + beta_end=0.012, + beta_schedule="scaled_linear", + clip_sample=False, + set_alpha_to_one=False, + steps_offset=1, + skip_prk_steps=True + ) + + # Use relative paths for whisper models as well + if config.model.cross_attention_dim == 768: + whisper_model_path = os.path.join(current_dir, "..", "checkpoints", "whisper", "small.pt") + elif config.model.cross_attention_dim == 384: + whisper_model_path = os.path.join(current_dir, "..", "checkpoints", "whisper", "tiny.pt") + else: + raise NotImplementedError("cross_attention_dim must be 768 or 384") + + audio_encoder = Audio2Feature( + model_path=whisper_model_path, + device="cuda", + num_frames=config.data.num_frames, + audio_feat_length=config.data.audio_feat_length, + ) + + vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse", torch_dtype=dtype) + vae.config.scaling_factor = 0.18215 + vae.config.shift_factor = 0 + + denoising_unet, _ = UNet3DConditionModel.from_pretrained( + OmegaConf.to_container(config.model), + args.inference_ckpt_path, + device="cpu", + ) + + denoising_unet = denoising_unet.to(dtype=dtype) + + pipeline = LipsyncPipeline( + vae=vae, + audio_encoder=audio_encoder, + denoising_unet=denoising_unet, + scheduler=scheduler, + ).to("cuda") + + if args.seed != -1: + set_seed(args.seed) + else: + torch.seed() + + print(f"Initial seed: {torch.initial_seed()}") + + pipeline( + video_path=args.video_path, + audio_path=args.audio_path, + video_out_path=args.video_out_path, + video_mask_path=args.video_out_path.replace(".mp4", "_mask.mp4"), + num_frames=config.data.num_frames, + num_inference_steps=args.inference_steps, + guidance_scale=args.guidance_scale, + weight_dtype=dtype, + width=config.data.resolution, + height=config.data.resolution, + mask_image_path=config.data.mask_image_path, + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--unet_config_path", type=str, default="configs/unet.yaml") + parser.add_argument("--inference_ckpt_path", type=str, required=True) + parser.add_argument("--video_path", type=str, required=True) + parser.add_argument("--audio_path", type=str, required=True) + parser.add_argument("--video_out_path", type=str, required=True) + parser.add_argument("--inference_steps", type=int, default=20) + parser.add_argument("--guidance_scale", type=float, default=1.0) + parser.add_argument("--seed", type=int, default=1247) + args = parser.parse_args() + + config = OmegaConf.load(args.unet_config_path) + + main(config, args) \ No newline at end of file diff --git a/scripts/train_syncnet.py b/scripts/train_syncnet.py new file mode 100644 index 0000000..371d240 --- /dev/null +++ b/scripts/train_syncnet.py @@ -0,0 +1,332 @@ +# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from tqdm.auto import tqdm +import os, argparse, datetime, math +import logging +from omegaconf import OmegaConf +import shutil + +from latentsync.data.syncnet_dataset import SyncNetDataset +from latentsync.models.stable_syncnet import StableSyncNet +from latentsync.models.wav2lip_syncnet import Wav2LipSyncNet +from latentsync.utils.util import gather_loss, plot_loss_chart +from accelerate.utils import set_seed + +import torch +from diffusers import AutoencoderKL +from diffusers.utils.logging import get_logger +from einops import rearrange +import torch.distributed as dist +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.utils.data.distributed import DistributedSampler +from latentsync.utils.util import init_dist, cosine_loss + +logger = get_logger(__name__) + + +def main(config): + # Initialize distributed training + local_rank = init_dist() + global_rank = dist.get_rank() + num_processes = dist.get_world_size() + is_main_process = global_rank == 0 + + seed = config.run.seed + global_rank + set_seed(seed) + + # Logging folder + folder_name = "train" + datetime.datetime.now().strftime(f"-%Y_%m_%d-%H:%M:%S") + output_dir = os.path.join(config.data.train_output_dir, folder_name) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + + # Handle the output folder creation + if is_main_process: + os.makedirs(output_dir, exist_ok=True) + os.makedirs(f"{output_dir}/checkpoints", exist_ok=True) + os.makedirs(f"{output_dir}/loss_charts", exist_ok=True) + shutil.copy(config.config_path, output_dir) + + device = torch.device(local_rank) + + if config.data.latent_space: + vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse", torch_dtype=torch.float16) + vae.requires_grad_(False) + vae.to(device) + else: + vae = None + + # Dataset and Dataloader setup + train_dataset = SyncNetDataset(config.data.train_data_dir, config.data.train_fileslist, config) + val_dataset = SyncNetDataset(config.data.val_data_dir, config.data.val_fileslist, config) + + train_distributed_sampler = DistributedSampler( + train_dataset, + num_replicas=num_processes, + rank=global_rank, + shuffle=True, + seed=config.run.seed, + ) + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_size=config.data.batch_size, + shuffle=False, + sampler=train_distributed_sampler, + num_workers=config.data.num_workers, + pin_memory=False, + drop_last=True, + worker_init_fn=train_dataset.worker_init_fn, + ) + + num_samples_limit = 640 + + val_batch_size = min( + num_samples_limit // config.data.num_frames, config.data.batch_size + ) # limit batch size to avoid CUDA OOM + + val_dataloader = torch.utils.data.DataLoader( + val_dataset, + batch_size=val_batch_size, + shuffle=False, + num_workers=config.data.num_workers, + pin_memory=False, + drop_last=False, + worker_init_fn=val_dataset.worker_init_fn, + ) + + # Model + syncnet = StableSyncNet(OmegaConf.to_container(config.model)).to(device) + # syncnet = Wav2LipSyncNet().to(device) + + optimizer = torch.optim.AdamW( + list(filter(lambda p: p.requires_grad, syncnet.parameters())), lr=config.optimizer.lr + ) + + if config.ckpt.resume_ckpt_path != "": + if is_main_process: + logger.info(f"Load checkpoint from: {config.ckpt.resume_ckpt_path}") + ckpt = torch.load(config.ckpt.resume_ckpt_path, map_location=device, weights_only=True) + + syncnet.load_state_dict(ckpt["state_dict"]) + global_step = ckpt["global_step"] + train_step_list = ckpt["train_step_list"] + train_loss_list = ckpt["train_loss_list"] + val_step_list = ckpt["val_step_list"] + val_loss_list = ckpt["val_loss_list"] + else: + global_step = 0 + train_step_list = [] + train_loss_list = [] + val_step_list = [] + val_loss_list = [] + + # DDP wrapper + syncnet = DDP(syncnet, device_ids=[local_rank], output_device=local_rank) + + num_update_steps_per_epoch = math.ceil(len(train_dataloader)) + num_train_epochs = math.ceil(config.run.max_train_steps / num_update_steps_per_epoch) + + if is_main_process: + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {config.data.batch_size}") + logger.info(f" Total train batch size (w. parallel & distributed) = {config.data.batch_size * num_processes}") + logger.info(f" Total optimization steps = {config.run.max_train_steps}") + + first_epoch = global_step // num_update_steps_per_epoch + num_val_batches = config.data.num_val_samples // (num_processes * config.data.batch_size) + + # Only show the progress bar once on each machine. + progress_bar = tqdm( + range(0, config.run.max_train_steps), initial=global_step, desc="Steps", disable=not is_main_process + ) + + # Support mixed-precision training + scaler = torch.amp.GradScaler("cuda") if config.run.mixed_precision_training else None + + for epoch in range(first_epoch, num_train_epochs): + train_dataloader.sampler.set_epoch(epoch) + syncnet.train() + + for step, batch in enumerate(train_dataloader): + ### >>>> Training >>>> ### + + frames = batch["frames"].to(device, dtype=torch.float16) + audio_samples = batch["audio_samples"].to(device, dtype=torch.float16) + y = batch["y"].to(device, dtype=torch.float32) + + if config.data.latent_space: + max_batch_size = ( + num_samples_limit // config.data.num_frames + ) # due to the limited cuda memory, we split the input frames into parts + if frames.shape[0] > max_batch_size: + assert ( + frames.shape[0] % max_batch_size == 0 + ), f"max_batch_size {max_batch_size} should be divisible by batch_size {frames.shape[0]}" + frames_part_results = [] + for i in range(0, frames.shape[0], max_batch_size): + frames_part = frames[i : i + max_batch_size] + frames_part = rearrange(frames_part, "b f c h w -> (b f) c h w") + with torch.no_grad(): + frames_part = vae.encode(frames_part).latent_dist.sample() * 0.18215 + frames_part_results.append(frames_part) + frames = torch.cat(frames_part_results, dim=0) + else: + frames = rearrange(frames, "b f c h w -> (b f) c h w") + with torch.no_grad(): + frames = vae.encode(frames).latent_dist.sample() * 0.18215 + + frames = rearrange(frames, "(b f) c h w -> b (f c) h w", f=config.data.num_frames) + else: + frames = rearrange(frames, "b f c h w -> b (f c) h w") + + if config.data.lower_half: + height = frames.shape[2] + frames = frames[:, :, height // 2 :, :] + + # Mixed-precision training + with torch.autocast(device_type="cuda", dtype=torch.float16, enabled=config.run.mixed_precision_training): + vision_embeds, audio_embeds = syncnet(frames, audio_samples) + + loss = cosine_loss(vision_embeds.float(), audio_embeds.float(), y).mean() + + optimizer.zero_grad() + + # Backpropagate + if config.run.mixed_precision_training: + scaler.scale(loss).backward() + """ >>> gradient clipping >>> """ + scaler.unscale_(optimizer) + torch.nn.utils.clip_grad_norm_(syncnet.parameters(), config.optimizer.max_grad_norm) + """ <<< gradient clipping <<< """ + scaler.step(optimizer) + scaler.update() + else: + loss.backward() + """ >>> gradient clipping >>> """ + torch.nn.utils.clip_grad_norm_(syncnet.parameters(), config.optimizer.max_grad_norm) + """ <<< gradient clipping <<< """ + optimizer.step() + + progress_bar.update(1) + global_step += 1 + + global_average_loss = gather_loss(loss, device) + train_step_list.append(global_step) + train_loss_list.append(global_average_loss) + + if is_main_process and global_step % config.run.validation_steps == 0: + logger.info(f"Validation at step {global_step}") + val_loss = validation( + val_dataloader, + device, + syncnet, + cosine_loss, + config.data.latent_space, + config.data.lower_half, + vae, + num_val_batches, + ) + val_step_list.append(global_step) + val_loss_list.append(val_loss) + logger.info(f"Validation loss at step {global_step} is {val_loss:0.3f}") + + if is_main_process and global_step % config.ckpt.save_ckpt_steps == 0: + checkpoint_save_path = os.path.join(output_dir, f"checkpoints/checkpoint-{global_step}.pt") + torch.save( + { + "state_dict": syncnet.module.state_dict(), # to unwrap DDP + "global_step": global_step, + "train_step_list": train_step_list, + "train_loss_list": train_loss_list, + "val_step_list": val_step_list, + "val_loss_list": val_loss_list, + }, + checkpoint_save_path, + ) + logger.info(f"Saved checkpoint to {checkpoint_save_path}") + plot_loss_chart( + os.path.join(output_dir, f"loss_charts/loss_chart-{global_step}.png"), + ("Train loss", train_step_list, train_loss_list), + ("Val loss", val_step_list, val_loss_list), + ) + + progress_bar.set_postfix({"step_loss": global_average_loss}) + if global_step >= config.run.max_train_steps: + break + + progress_bar.close() + dist.destroy_process_group() + + +@torch.no_grad() +def validation(val_dataloader, device, syncnet, cosine_loss, latent_space, lower_half, vae, num_val_batches): + syncnet.eval() + + losses = [] + val_step = 0 + while True: + for step, batch in enumerate(val_dataloader): + ### >>>> Validation >>>> ### + + frames = batch["frames"].to(device, dtype=torch.float16) + audio_samples = batch["audio_samples"].to(device, dtype=torch.float16) + y = batch["y"].to(device, dtype=torch.float32) + + if latent_space: + num_frames = frames.shape[1] + frames = rearrange(frames, "b f c h w -> (b f) c h w") + frames = vae.encode(frames).latent_dist.sample() * 0.18215 + frames = rearrange(frames, "(b f) c h w -> b (f c) h w", f=num_frames) + else: + frames = rearrange(frames, "b f c h w -> b (f c) h w") + + if lower_half: + height = frames.shape[2] + frames = frames[:, :, height // 2 :, :] + + with torch.autocast(device_type="cuda", dtype=torch.float16): + vision_embeds, audio_embeds = syncnet(frames, audio_samples) + + loss = cosine_loss(vision_embeds.float(), audio_embeds.float(), y).mean() + + losses.append(loss.item()) + + val_step += 1 + if val_step > num_val_batches: + syncnet.train() + if len(losses) == 0: + raise RuntimeError("No validation data") + return sum(losses) / len(losses) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="Code to train the SyncNet") + parser.add_argument("--config_path", type=str, default="configs/syncnet/syncnet_16_pixel.yaml") + args = parser.parse_args() + + # Load a configuration file + config = OmegaConf.load(args.config_path) + config.config_path = args.config_path + + main(config) diff --git a/scripts/train_unet.py b/scripts/train_unet.py new file mode 100644 index 0000000..cf9f7bc --- /dev/null +++ b/scripts/train_unet.py @@ -0,0 +1,516 @@ +# Copyright (c) 2024 Bytedance Ltd. and/or its affiliates +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import math +import argparse +import shutil +import datetime +import logging +from omegaconf import OmegaConf + +from tqdm.auto import tqdm +from einops import rearrange + +import torch +import torch.nn.functional as F +import torch.nn as nn +import torch.distributed as dist +from torch.utils.data.distributed import DistributedSampler +from torch.nn.parallel import DistributedDataParallel as DDP + +import diffusers +from diffusers import AutoencoderKL, DDIMScheduler +from diffusers.utils.logging import get_logger +from diffusers.optimization import get_scheduler +from accelerate.utils import set_seed + +from latentsync.data.unet_dataset import UNetDataset +from latentsync.models.unet import UNet3DConditionModel +from latentsync.models.stable_syncnet import StableSyncNet +from latentsync.pipelines.lipsync_pipeline import LipsyncPipeline +from latentsync.utils.util import ( + init_dist, + cosine_loss, + one_step_sampling, +) +from latentsync.utils.util import plot_loss_chart +from latentsync.whisper.audio2feature import Audio2Feature +from latentsync.trepa.loss import TREPALoss +from eval.syncnet import SyncNetEval +from eval.syncnet_detect import SyncNetDetector +from eval.eval_sync_conf import syncnet_eval +import lpips + + +logger = get_logger(__name__) + + +def main(config): + # Initialize distributed training + local_rank = init_dist() + global_rank = dist.get_rank() + num_processes = dist.get_world_size() + is_main_process = global_rank == 0 + + seed = config.run.seed + global_rank + set_seed(seed) + + # Logging folder + folder_name = "train" + datetime.datetime.now().strftime(f"-%Y_%m_%d-%H:%M:%S") + output_dir = os.path.join(config.data.train_output_dir, folder_name) + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + + # Handle the output folder creation + if is_main_process: + diffusers.utils.logging.set_verbosity_info() + os.makedirs(output_dir, exist_ok=True) + os.makedirs(f"{output_dir}/checkpoints", exist_ok=True) + os.makedirs(f"{output_dir}/val_videos", exist_ok=True) + os.makedirs(f"{output_dir}/sync_conf_results", exist_ok=True) + shutil.copy(config.unet_config_path, output_dir) + shutil.copy(config.data.syncnet_config_path, output_dir) + + device = torch.device(local_rank) + + noise_scheduler = DDIMScheduler.from_pretrained("configs") + + vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse", torch_dtype=torch.float16) + vae.config.scaling_factor = 0.18215 + vae.config.shift_factor = 0 + + vae_scale_factor = 2 ** (len(vae.config.block_out_channels) - 1) + vae.requires_grad_(False) + vae.to(device) + + if config.run.pixel_space_supervise: + vae.enable_gradient_checkpointing() + + syncnet_eval_model = SyncNetEval(device=device) + syncnet_eval_model.loadParameters("checkpoints/auxiliary/syncnet_v2.model") + + syncnet_detector = SyncNetDetector(device=device, detect_results_dir="detect_results") + + if config.model.cross_attention_dim == 768: + whisper_model_path = "checkpoints/whisper/small.pt" + elif config.model.cross_attention_dim == 384: + whisper_model_path = "checkpoints/whisper/tiny.pt" + else: + raise NotImplementedError("cross_attention_dim must be 768 or 384") + + audio_encoder = Audio2Feature( + model_path=whisper_model_path, + device=device, + audio_embeds_cache_dir=config.data.audio_embeds_cache_dir, + num_frames=config.data.num_frames, + audio_feat_length=config.data.audio_feat_length, + ) + + denoising_unet, resume_global_step = UNet3DConditionModel.from_pretrained( + OmegaConf.to_container(config.model), + config.ckpt.resume_ckpt_path, + device=device, + ) + + if config.model.add_audio_layer and config.run.use_syncnet: + syncnet_config = OmegaConf.load(config.data.syncnet_config_path) + if syncnet_config.ckpt.inference_ckpt_path == "": + raise ValueError("SyncNet path is not provided") + syncnet = StableSyncNet(OmegaConf.to_container(syncnet_config.model), gradient_checkpointing=True).to( + device=device, dtype=torch.float16 + ) + syncnet_checkpoint = torch.load( + syncnet_config.ckpt.inference_ckpt_path, map_location=device, weights_only=True + ) + syncnet.load_state_dict(syncnet_checkpoint["state_dict"]) + syncnet.requires_grad_(False) + + del syncnet_checkpoint + torch.cuda.empty_cache() + + if config.model.use_motion_module: + denoising_unet.requires_grad_(False) + for name, param in denoising_unet.named_parameters(): + for trainable_module_name in config.run.trainable_modules: + if trainable_module_name in name: + param.requires_grad = True + break + trainable_params = list(filter(lambda p: p.requires_grad, denoising_unet.parameters())) + else: + denoising_unet.requires_grad_(True) + trainable_params = list(denoising_unet.parameters()) + + if config.optimizer.scale_lr: + config.optimizer.lr = config.optimizer.lr * num_processes + + optimizer = torch.optim.AdamW(trainable_params, lr=config.optimizer.lr) + + if is_main_process: + logger.info(f"trainable params number: {len(trainable_params)}") + logger.info(f"trainable params scale: {sum(p.numel() for p in trainable_params) / 1e6:.3f} M") + + # Enable gradient checkpointing + if config.run.enable_gradient_checkpointing: + denoising_unet.enable_gradient_checkpointing() + + # Get the training dataset + train_dataset = UNetDataset(config.data.train_data_dir, config) + distributed_sampler = DistributedSampler( + train_dataset, + num_replicas=num_processes, + rank=global_rank, + shuffle=True, + seed=config.run.seed, + ) + + # DataLoaders creation: + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_size=config.data.batch_size, + shuffle=False, + sampler=distributed_sampler, + num_workers=config.data.num_workers, + pin_memory=False, + drop_last=True, + worker_init_fn=train_dataset.worker_init_fn, + ) + + # Get the training iteration + if config.run.max_train_steps == -1: + assert config.run.max_train_epochs != -1 + config.run.max_train_steps = config.run.max_train_epochs * len(train_dataloader) + + # Scheduler + lr_scheduler = get_scheduler( + config.optimizer.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=config.optimizer.lr_warmup_steps, + num_training_steps=config.run.max_train_steps, + ) + + if config.run.perceptual_loss_weight != 0 and config.run.pixel_space_supervise: + lpips_loss_func = lpips.LPIPS(net="vgg").to(device) + + if config.run.trepa_loss_weight != 0 and config.run.pixel_space_supervise: + trepa_loss_func = TREPALoss(device=device, with_cp=True) + + # Validation pipeline + pipeline = LipsyncPipeline( + vae=vae, + audio_encoder=audio_encoder, + denoising_unet=denoising_unet, + scheduler=noise_scheduler, + ).to(device) + pipeline.set_progress_bar_config(disable=True) + + # DDP warpper + denoising_unet = DDP(denoising_unet, device_ids=[local_rank], output_device=local_rank) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader)) + # Afterwards we recalculate our number of training epochs + num_train_epochs = math.ceil(config.run.max_train_steps / num_update_steps_per_epoch) + + # Train! + total_batch_size = config.data.batch_size * num_processes + + if is_main_process: + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num Epochs = {num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {config.data.batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Total optimization steps = {config.run.max_train_steps}") + global_step = resume_global_step + first_epoch = resume_global_step // num_update_steps_per_epoch + + # Only show the progress bar once on each machine. + progress_bar = tqdm( + range(0, config.run.max_train_steps), + initial=resume_global_step, + desc="Steps", + disable=not is_main_process, + ) + + train_step_list = [] + val_step_list = [] + sync_conf_list = [] + + # Support mixed-precision training + scaler = torch.amp.GradScaler("cuda") if config.run.mixed_precision_training else None + + for epoch in range(first_epoch, num_train_epochs): + train_dataloader.sampler.set_epoch(epoch) + denoising_unet.train() + + for step, batch in enumerate(train_dataloader): + ### >>>> Training >>>> ### + + if config.model.add_audio_layer: + if batch["mel"] != []: + mel = batch["mel"].to(device, dtype=torch.float16) + + audio_embeds_list = [] + try: + for idx in range(len(batch["video_path"])): + video_path = batch["video_path"][idx] + start_idx = batch["start_idx"][idx] + + with torch.no_grad(): + audio_feat = audio_encoder.audio2feat(video_path) + audio_embeds = audio_encoder.crop_overlap_audio_window(audio_feat, start_idx) + audio_embeds_list.append(audio_embeds) + except Exception as e: + logger.info(f"{type(e).__name__} - {e} - {video_path}") + continue + audio_embeds = torch.stack(audio_embeds_list) # (B, 16, 50, 384) + audio_embeds = audio_embeds.to(device, dtype=torch.float16) + else: + audio_embeds = None + + # Convert videos to latent space + gt_pixel_values = batch["gt_pixel_values"].to(device, dtype=torch.float16) + masked_pixel_values = batch["masked_pixel_values"].to(device, dtype=torch.float16) + masks = batch["masks"].to(device, dtype=torch.float16) + ref_pixel_values = batch["ref_pixel_values"].to(device, dtype=torch.float16) + + gt_pixel_values = rearrange(gt_pixel_values, "b f c h w -> (b f) c h w") + masked_pixel_values = rearrange(masked_pixel_values, "b f c h w -> (b f) c h w") + masks = rearrange(masks, "b f c h w -> (b f) c h w") + ref_pixel_values = rearrange(ref_pixel_values, "b f c h w -> (b f) c h w") + + with torch.no_grad(): + gt_latents = vae.encode(gt_pixel_values).latent_dist.sample() + masked_latents = vae.encode(masked_pixel_values).latent_dist.sample() + ref_latents = vae.encode(ref_pixel_values).latent_dist.sample() + + masks = torch.nn.functional.interpolate(masks, size=config.data.resolution // vae_scale_factor) + + gt_latents = ( + rearrange(gt_latents, "(b f) c h w -> b c f h w", f=config.data.num_frames) - vae.config.shift_factor + ) * vae.config.scaling_factor + masked_latents = ( + rearrange(masked_latents, "(b f) c h w -> b c f h w", f=config.data.num_frames) + - vae.config.shift_factor + ) * vae.config.scaling_factor + ref_latents = ( + rearrange(ref_latents, "(b f) c h w -> b c f h w", f=config.data.num_frames) - vae.config.shift_factor + ) * vae.config.scaling_factor + masks = rearrange(masks, "(b f) c h w -> b c f h w", f=config.data.num_frames) + + # Sample noise that we'll add to the latents + if config.run.use_mixed_noise: + # Refer to the paper: https://arxiv.org/abs/2305.10474 + noise_shared_std_dev = (config.run.mixed_noise_alpha**2 / (1 + config.run.mixed_noise_alpha**2)) ** 0.5 + noise_shared = torch.randn_like(gt_latents) * noise_shared_std_dev + noise_shared = noise_shared[:, :, 0:1].repeat(1, 1, config.data.num_frames, 1, 1) + + noise_ind_std_dev = (1 / (1 + config.run.mixed_noise_alpha**2)) ** 0.5 + noise_ind = torch.randn_like(gt_latents) * noise_ind_std_dev + noise = noise_ind + noise_shared + else: + noise = torch.randn_like(gt_latents) + noise = noise[:, :, 0:1].repeat( + 1, 1, config.data.num_frames, 1, 1 + ) # Using the same noise for all frames, refer to the paper: https://arxiv.org/abs/2308.09716 + + bsz = gt_latents.shape[0] + + # Sample a random timestep for each video + timesteps = torch.randint(0, noise_scheduler.config.num_train_timesteps, (bsz,), device=gt_latents.device) + timesteps = timesteps.long() + + # Add noise to the latents according to the noise magnitude at each timestep + # (this is the forward diffusion process) + noisy_gt_latents = noise_scheduler.add_noise(gt_latents, noise, timesteps) + + # Get the target for loss depending on the prediction type + if noise_scheduler.config.prediction_type == "epsilon": + target = noise + elif noise_scheduler.config.prediction_type == "v_prediction": + raise NotImplementedError + else: + raise ValueError(f"Unknown prediction type {noise_scheduler.config.prediction_type}") + + denoising_unet_input = torch.cat([noisy_gt_latents, masks, masked_latents, ref_latents], dim=1) + + # Predict the noise and compute loss + # Mixed-precision training + with torch.autocast(device_type="cuda", dtype=torch.float16, enabled=config.run.mixed_precision_training): + pred_noise = denoising_unet(denoising_unet_input, timesteps, encoder_hidden_states=audio_embeds).sample + + if config.run.recon_loss_weight != 0: + recon_loss = F.mse_loss(pred_noise.float(), target.float(), reduction="mean") + else: + recon_loss = 0 + + pred_latents = one_step_sampling(noise_scheduler, pred_noise, timesteps, noisy_gt_latents) + + if config.run.pixel_space_supervise: + pred_pixel_values = vae.decode( + rearrange(pred_latents, "b c f h w -> (b f) c h w") / vae.config.scaling_factor + + vae.config.shift_factor + ).sample + + if config.run.perceptual_loss_weight != 0 and config.run.pixel_space_supervise: + pred_pixel_values_perceptual = pred_pixel_values[:, :, pred_pixel_values.shape[2] // 2 :, :] + gt_pixel_values_perceptual = gt_pixel_values[:, :, gt_pixel_values.shape[2] // 2 :, :] + lpips_loss = lpips_loss_func( + pred_pixel_values_perceptual.float(), gt_pixel_values_perceptual.float() + ).mean() + else: + lpips_loss = 0 + + if config.run.trepa_loss_weight != 0 and config.run.pixel_space_supervise: + trepa_pred_pixel_values = rearrange( + pred_pixel_values, "(b f) c h w -> b c f h w", f=config.data.num_frames + ) + trepa_gt_pixel_values = rearrange( + gt_pixel_values, "(b f) c h w -> b c f h w", f=config.data.num_frames + ) + trepa_loss = trepa_loss_func(trepa_pred_pixel_values, trepa_gt_pixel_values) + else: + trepa_loss = 0 + + if config.model.add_audio_layer and config.run.use_syncnet: + if config.run.pixel_space_supervise: + syncnet_input = rearrange( + pred_pixel_values, "(b f) c h w -> b (f c) h w", f=config.data.num_frames + ) + else: + syncnet_input = rearrange(pred_latents, "b c f h w -> b (f c) h w") + + if syncnet_config.data.lower_half: + height = syncnet_input.shape[2] + syncnet_input = syncnet_input[:, :, height // 2 :, :] + ones_tensor = torch.ones((config.data.batch_size, 1)).float().to(device=device) + vision_embeds, audio_embeds = syncnet(syncnet_input, mel) + sync_loss = cosine_loss(vision_embeds.float(), audio_embeds.float(), ones_tensor).mean() + else: + sync_loss = 0 + + loss = ( + recon_loss * config.run.recon_loss_weight + + sync_loss * config.run.sync_loss_weight + + lpips_loss * config.run.perceptual_loss_weight + + trepa_loss * config.run.trepa_loss_weight + ) + + train_step_list.append(global_step) + + optimizer.zero_grad() + + # Backpropagate + if config.run.mixed_precision_training: + scaler.scale(loss).backward() + """ >>> gradient clipping >>> """ + scaler.unscale_(optimizer) + torch.nn.utils.clip_grad_norm_(trainable_params, config.optimizer.max_grad_norm) + """ <<< gradient clipping <<< """ + scaler.step(optimizer) + scaler.update() + else: + loss.backward() + """ >>> gradient clipping >>> """ + torch.nn.utils.clip_grad_norm_(trainable_params, config.optimizer.max_grad_norm) + """ <<< gradient clipping <<< """ + optimizer.step() + + # Check the grad of attn blocks for debugging + # print(denoising_unet.module.up_blocks[3].attentions[2].transformer_blocks[0].attn2.to_q.weight.grad) + + lr_scheduler.step() + progress_bar.update(1) + global_step += 1 + + ### <<<< Training <<<< ### + + # Save checkpoint and conduct validation + if is_main_process and (global_step % config.ckpt.save_ckpt_steps == 0): + model_save_path = os.path.join(output_dir, f"checkpoints/checkpoint-{global_step}.pt") + state_dict = { + "global_step": global_step, + "state_dict": denoising_unet.module.state_dict(), + } + try: + torch.save(state_dict, model_save_path) + logger.info(f"Saved checkpoint to {model_save_path}") + except Exception as e: + logger.error(f"Error saving model: {e}") + + # Validation + logger.info("Running validation... ") + + validation_video_out_path = os.path.join(output_dir, f"val_videos/val_video_{global_step}.mp4") + validation_video_mask_path = os.path.join(output_dir, f"val_videos/val_video_mask.mp4") + + with torch.autocast(device_type="cuda", dtype=torch.float16): + pipeline( + config.data.val_video_path, + config.data.val_audio_path, + validation_video_out_path, + validation_video_mask_path, + num_frames=config.data.num_frames, + num_inference_steps=config.run.inference_steps, + guidance_scale=config.run.guidance_scale, + weight_dtype=torch.float16, + width=config.data.resolution, + height=config.data.resolution, + mask=config.data.mask, + mask_image_path=config.data.mask_image_path, + ) + + logger.info(f"Saved validation video output to {validation_video_out_path}") + + val_step_list.append(global_step) + + if config.model.add_audio_layer and os.path.exists(validation_video_out_path): + try: + _, conf = syncnet_eval(syncnet_eval_model, syncnet_detector, validation_video_out_path, "temp") + except Exception as e: + logger.info(e) + conf = 0 + sync_conf_list.append(conf) + plot_loss_chart( + os.path.join(output_dir, f"sync_conf_results/sync_conf_chart-{global_step}.png"), + ("Sync confidence", val_step_list, sync_conf_list), + ) + + logs = {"step_loss": loss.item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + + if global_step >= config.run.max_train_steps: + break + + progress_bar.close() + dist.destroy_process_group() + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + + # Config file path + parser.add_argument("--unet_config_path", type=str, default="configs/unet.yaml") + + args = parser.parse_args() + config = OmegaConf.load(args.unet_config_path) + config.unet_config_path = args.unet_config_path + + main(config)