From 1ca41624472551bf994cb67eee93e5fc7a739cec Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Tue, 10 Feb 2026 17:23:23 +0800 Subject: [PATCH] Support validations in all models && Update Saving Code && Preparing Code (#452) --- scripts/cogvideox_fun/train.py | 323 +++++----- scripts/cogvideox_fun/train.sh | 4 +- scripts/cogvideox_fun/train_control.py | 235 ++++---- scripts/cogvideox_fun/train_control.sh | 4 +- scripts/cogvideox_fun/train_lora.py | 301 ++++------ scripts/cogvideox_fun/train_lora.sh | 4 +- scripts/fantasytalking/train.py | 247 ++++---- scripts/fantasytalking/train.sh | 2 +- scripts/flux/train.py | 26 +- scripts/flux/train.sh | 2 +- scripts/flux/train_lora.py | 59 +- scripts/flux/train_lora.sh | 2 +- scripts/flux2/train.py | 24 +- scripts/flux2/train.sh | 2 +- scripts/flux2/train_lora.py | 57 +- scripts/flux2/train_lora.sh | 2 +- scripts/flux2_fun/train_control.py | 24 +- scripts/flux2_fun/train_control_distill.py | 24 +- scripts/flux2_fun/train_control_distill.sh | 2 +- scripts/hunyuanvideo/train.py | 217 ++++--- scripts/hunyuanvideo/train.sh | 2 +- scripts/hunyuanvideo/train_lora.py | 241 ++++---- scripts/hunyuanvideo/train_lora.sh | 2 +- scripts/longcatvideo/train.py | 297 ++++----- scripts/longcatvideo/train.sh | 2 +- scripts/longcatvideo/train_lora.py | 279 ++++----- scripts/longcatvideo/train_lora.sh | 2 +- scripts/qwenimage/train.py | 26 +- scripts/qwenimage/train.sh | 2 +- scripts/qwenimage/train_edit.py | 28 +- scripts/qwenimage/train_edit.sh | 2 +- scripts/qwenimage/train_edit_lora.py | 59 +- scripts/qwenimage/train_edit_lora.sh | 2 +- scripts/qwenimage/train_lora.py | 57 +- scripts/qwenimage/train_lora.sh | 2 +- scripts/qwenimage_fun/train_control.py | 26 +- scripts/qwenimage_instantx/train_control.py | 26 +- scripts/turbodiffusion/train_distill.py | 285 ++++----- scripts/turbodiffusion/train_distill.sh | 2 +- scripts/wan2.1/train.py | 328 +++++----- scripts/wan2.1/train.sh | 4 +- scripts/wan2.1/train_distill.py | 284 ++++----- scripts/wan2.1/train_distill.sh | 2 +- scripts/wan2.1/train_distill_lora.py | 322 +++++----- scripts/wan2.1/train_distill_lora.sh | 2 +- scripts/wan2.1/train_lora.py | 310 +++++----- scripts/wan2.1/train_lora.sh | 4 +- scripts/wan2.1_fun/train.py | 307 +++++----- scripts/wan2.1_fun/train.sh | 4 +- scripts/wan2.1_fun/train_control.py | 222 ++++--- scripts/wan2.1_fun/train_control.sh | 2 +- scripts/wan2.1_fun/train_control_lora.py | 236 ++++---- scripts/wan2.1_fun/train_control_lora.sh | 2 +- scripts/wan2.1_fun/train_lora.py | 318 +++++----- scripts/wan2.1_fun/train_lora.sh | 4 +- scripts/wan2.1_vace/train.py | 231 ++++--- scripts/wan2.1_vace/train.sh | 2 +- scripts/wan2.2/train.py | 389 ++++++------ scripts/wan2.2/train.sh | 4 +- scripts/wan2.2/train_animate.py | 329 +++++----- scripts/wan2.2/train_animate.sh | 2 +- scripts/wan2.2/train_animate_lora.py | 324 +++++----- scripts/wan2.2/train_animate_lora.sh | 2 +- scripts/wan2.2/train_distill.py | 349 +++++------ scripts/wan2.2/train_distill.sh | 2 +- scripts/wan2.2/train_distill_lora.py | 393 ++++++------ scripts/wan2.2/train_distill_lora.sh | 2 +- scripts/wan2.2/train_lora.py | 381 +++++------- scripts/wan2.2/train_lora.sh | 4 +- scripts/wan2.2/train_s2v.py | 370 ++++++------ scripts/wan2.2/train_s2v.sh | 4 +- scripts/wan2.2/train_s2v_lora.py | 350 +++++------ scripts/wan2.2/train_s2v_lora.sh | 2 +- scripts/wan2.2_fun/train.py | 383 ++++++------ scripts/wan2.2_fun/train.sh | 2 +- scripts/wan2.2_fun/train_control.py | 280 +++++---- scripts/wan2.2_fun/train_control.sh | 2 +- scripts/wan2.2_fun/train_control_lora.py | 290 ++++----- scripts/wan2.2_fun/train_control_lora.sh | 2 +- scripts/wan2.2_fun/train_lora.py | 380 +++++------- scripts/wan2.2_fun/train_lora.sh | 2 +- scripts/wan2.2_vace_fun/train.py | 294 +++++---- scripts/wan2.2_vace_fun/train.sh | 2 +- scripts/z_image/train.py | 24 +- scripts/z_image/train_distill.py | 589 ++++++++++-------- scripts/z_image/train_distill_lora.py | 595 +++++++++++-------- scripts/z_image/train_lora.py | 57 +- scripts/z_image_fun/train_control.py | 24 +- scripts/z_image_fun/train_control_distill.py | 40 +- videox_fun/models/cogvideox_transformer3d.py | 3 +- videox_fun/utils/__init__.py | 11 +- videox_fun/utils/trigflow_sampler.py | 23 + 92 files changed, 5118 insertions(+), 6279 deletions(-) create mode 100644 videox_fun/utils/trigflow_sampler.py diff --git a/scripts/cogvideox_fun/train.py b/scripts/cogvideox_fun/train.py index b199440..50c9104 100755 --- a/scripts/cogvideox_fun/train.py +++ b/scripts/cogvideox_fun/train.py @@ -58,28 +58,30 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, - ASPECT_RATIO_RANDOM_CROP_512, - ASPECT_RATIO_RANDOM_CROP_PROB, - AspectRatioBatchImageVideoSampler, - RandomSampler, get_closest_ratio) + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, - ImageVideoDataset, - ImageVideoSampler, - get_random_mask) + ImageVideoDataset, + ImageVideoSampler, + get_random_mask) from videox_fun.models import (AutoencoderKLCogVideoX, - CogVideoXTransformer3DModel, T5EncoderModel, - T5Tokenizer) -from videox_fun.pipeline import (CogVideoXFunPipeline, - CogVideoXFunControlPipeline, - CogVideoXFunInpaintPipeline) + CogVideoXTransformer3DModel, T5EncoderModel, + T5Tokenizer) +from videox_fun.pipeline import (CogVideoXFunControlPipeline, + CogVideoXFunInpaintPipeline, + CogVideoXFunPipeline) from videox_fun.pipeline.pipeline_cogvideox_fun_inpaint import ( add_noise_to_reference_video, get_3d_rotary_pos_embed, get_resize_crop_region_for_grid) from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.lora_utils import create_network, merge_lora, unmerge_lora -from videox_fun.utils.utils import (get_image_to_video_latent, - get_video_to_video_latent, save_videos_grid) - +from videox_fun.utils.lora_utils import (create_network, merge_lora, + unmerge_lora) +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + get_video_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -160,116 +162,106 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = CogVideoXTransformer3DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer" - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + + if args.train_mode != "normal": + pipeline = CogVideoXFunInpaintPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + else: + pipeline = CogVideoXFunPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - if args.train_mode != "normal": - pipeline = CogVideoXFunInpaintPipeline.from_pretrained( - args.pretrained_model_name_or_path, - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - torch_dtype=weight_dtype - ) - else: - pipeline = CogVideoXFunPipeline.from_pretrained( - args.pretrained_model_name_or_path, - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - torch_dtype=weight_dtype - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -341,6 +333,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -905,7 +904,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -925,26 +924,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1771,25 +1750,24 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1797,25 +1775,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/cogvideox_fun/train.sh b/scripts/cogvideox_fun/train.sh index d3e7508..889b0a3 100755 --- a/scripts/cogvideox_fun/train.sh +++ b/scripts/cogvideox_fun/train.sh @@ -25,7 +25,7 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_cog" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -65,7 +65,7 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train.py \ # --lr_scheduler="constant_with_warmup" \ # --lr_warmup_steps=100 \ # --seed=42 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_cog" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/cogvideox_fun/train_control.py b/scripts/cogvideox_fun/train_control.py index 0bd223f..a94c2f0 100755 --- a/scripts/cogvideox_fun/train_control.py +++ b/scripts/cogvideox_fun/train_control.py @@ -35,7 +35,7 @@ from accelerate import Accelerator from accelerate.logging import get_logger from accelerate.state import AcceleratorState from accelerate.utils import ProjectConfiguration, set_seed -from diffusers import AutoencoderKL, DDPMScheduler +from diffusers import AutoencoderKL, DDIMScheduler, DDPMScheduler from diffusers.optimization import get_scheduler from diffusers.training_utils import EMAModel from diffusers.utils import check_min_version, deprecate, is_wandb_available @@ -56,26 +56,28 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, - ASPECT_RATIO_RANDOM_CROP_512, - ASPECT_RATIO_RANDOM_CROP_PROB, - AspectRatioBatchImageVideoSampler, - RandomSampler, get_closest_ratio) + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, - ImageVideoDataset, - ImageVideoSampler, - get_random_mask) + ImageVideoDataset, + ImageVideoSampler, + get_random_mask) from videox_fun.models import (AutoencoderKLCogVideoX, - CogVideoXTransformer3DModel, T5EncoderModel, - T5Tokenizer) -from videox_fun.pipeline import (CogVideoXFunPipeline, - CogVideoXFunControlPipeline, - CogVideoXFunInpaintPipeline) + CogVideoXTransformer3DModel, T5EncoderModel, + T5Tokenizer) +from videox_fun.pipeline import (CogVideoXFunControlPipeline, + CogVideoXFunInpaintPipeline, + CogVideoXFunPipeline) from videox_fun.pipeline.pipeline_cogvideox_fun_inpaint import ( add_noise_to_reference_video, get_3d_rotary_pos_embed, get_resize_crop_region_for_grid) from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import (get_image_to_video_latent, - get_video_to_video_latent, save_videos_grid) +from videox_fun.utils.utils import (calculate_dimensions, + get_image_to_video_latent, + get_video_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -156,60 +158,79 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = CogVideoXTransformer3DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer" - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") - pipeline = CogVideoXFunControlPipeline.from_pretrained( - args.pretrained_model_name_or_path, - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - torch_dtype=weight_dtype, - ) - pipeline = pipeline.to(accelerator.device) + pipeline = CogVideoXFunControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + for i in range(len(args.validation_prompts)): + import cv2 + cap = cv2.VideoCapture(args.validation_paths[i]) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - control_video = input_video, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + + input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "The video is not of a high quality, it has a low resolution. Watermark present in each frame. The background is solid. Strange body and strange trajectory. Distortion. ", + height = height, + width = width, + generator = generator, - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + control_video = input_video, + num_inference_steps = 25, + guidance_scale = 6.0, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -843,7 +864,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -863,26 +884,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1658,25 +1659,24 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1684,25 +1684,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/cogvideox_fun/train_control.sh b/scripts/cogvideox_fun/train_control.sh index 1ee000e..a253ff2 100755 --- a/scripts/cogvideox_fun/train_control.sh +++ b/scripts/cogvideox_fun/train_control.sh @@ -25,7 +25,7 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train_control.p --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=50 \ --seed=43 \ - --output_dir="output_dir" \ + --output_dir="output_dir_cog_control" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -64,7 +64,7 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train_control.p # --lr_scheduler="constant_with_warmup" \ # --lr_warmup_steps=50 \ # --seed=43 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_cog_control" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/cogvideox_fun/train_lora.py b/scripts/cogvideox_fun/train_lora.py index b802670..fc2dc38 100755 --- a/scripts/cogvideox_fun/train_lora.py +++ b/scripts/cogvideox_fun/train_lora.py @@ -78,7 +78,8 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -161,116 +162,108 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = CogVideoXTransformer3DModel.from_pretrained( - args.pretrained_model_name_or_path, subfolder="transformer", - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = DDIMScheduler.from_pretrained(args.pretrained_model_name_or_path, subfolder="scheduler") - if args.train_mode != "normal": - pipeline = CogVideoXFunInpaintPipeline.from_pretrained( - args.pretrained_model_name_or_path, - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - torch_dtype=weight_dtype, - ) - else: - pipeline = CogVideoXFunPipeline.from_pretrained( - args.pretrained_model_name_or_path, - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - torch_dtype=weight_dtype - ) + if args.train_mode != "normal": + pipeline = CogVideoXFunInpaintPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + else: + pipeline = CogVideoXFunPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 7, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -342,6 +335,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -916,7 +916,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -945,15 +945,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -966,34 +958,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1301,23 +1271,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=transformer3d.transformer_blocks) - transformer3d = shard_fn(transformer3d) # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) @@ -1841,19 +1800,18 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1861,19 +1819,18 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/cogvideox_fun/train_lora.sh b/scripts/cogvideox_fun/train_lora.sh index 19ec0ac..de54bda 100755 --- a/scripts/cogvideox_fun/train_lora.sh +++ b/scripts/cogvideox_fun/train_lora.sh @@ -23,7 +23,7 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_cog_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -65,7 +65,7 @@ accelerate launch --mixed_precision="bf16" scripts/cogvideox_fun/train_lora.py \ # --checkpointing_steps=50 \ # --learning_rate=1e-04 \ # --seed=42 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_cog_lora" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/fantasytalking/train.py b/scripts/fantasytalking/train.py index 6787b99..90d43be 100644 --- a/scripts/fantasytalking/train.py +++ b/scripts/fantasytalking/train.py @@ -70,13 +70,18 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_PROB, AspectRatioBatchImageVideoSampler, RandomSampler, get_closest_ratio) -from videox_fun.data.dataset_image_video import ImageVideoSampler, get_random_mask +from videox_fun.data.dataset_image_video import (ImageVideoSampler, + get_random_mask) from videox_fun.data.dataset_video import VideoSpeechDataset -from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, FantasyTalkingAudioEncoder, - FantasyTalkingTransformer3DModel) +from videox_fun.models import (AutoencoderKLWan, CLIPModel, + FantasyTalkingAudioEncoder, + FantasyTalkingTransformer3DModel, + WanT5EncoderModel) from videox_fun.pipeline import FantasyTalkingPipeline, WanFunPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -143,73 +148,86 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, audio_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) - transformer3d_val = FantasyTalkingTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + pipeline = FantasyTalkingPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + audio_encoder=audio_encoder, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = FantasyTalkingPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): + start_image = Image.open(args.validation_image_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, clip_image = get_image_to_video_latent(args.validation_image_paths[i], None, video_length=video_length, sample_size=[height, width]) + audio_path = args.validation_audio_paths[i] - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, clip_image = get_image_to_video_latent(args.validation_image_paths[i], None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - audio_path = args.validation_audio_paths[i] + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + video = input_video, + mask_video = input_video_mask, + clip_image = clip_image, + audio_path = audio_path, + shift = 5, + fps = 16 + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - video = input_video, - mask_video = input_video_mask, - clip_image = clip_image, - audio_path = audio_path, - shift = 5, - fps = 16 - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -895,7 +913,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -915,26 +933,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1350,8 +1348,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -1836,27 +1835,27 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + audio_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1864,27 +1863,27 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + audio_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/fantasytalking/train.sh b/scripts/fantasytalking/train.sh index b7eac02..c50c43e 100644 --- a/scripts/fantasytalking/train.sh +++ b/scripts/fantasytalking/train.sh @@ -24,7 +24,7 @@ accelerate launch --mixed_precision="bf16" scripts/fantasytalking/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_fantasytalking" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/flux/train.py b/scripts/flux/train.py index f6a5783..366c3f9 100644 --- a/scripts/flux/train.py +++ b/scripts/flux/train.py @@ -262,7 +262,7 @@ def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, tr text_encoder_2=text_encoder_2, tokenizer=tokenizer, tokenizer_2=tokenizer_2, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -926,7 +926,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -946,26 +946,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1302,7 +1282,7 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model diff --git a/scripts/flux/train.sh b/scripts/flux/train.sh index 5488c30..3096b8d 100644 --- a/scripts/flux/train.sh +++ b/scripts/flux/train.sh @@ -20,7 +20,7 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_flux" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/flux/train_lora.py b/scripts/flux/train_lora.py index f9b7811..a60c442 100644 --- a/scripts/flux/train_lora.py +++ b/scripts/flux/train_lora.py @@ -265,7 +265,7 @@ def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, tr text_encoder_2=text_encoder_2, tokenizer=tokenizer, tokenizer_2=tokenizer_2, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -929,7 +929,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -958,15 +958,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -979,34 +971,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1291,23 +1261,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(transformer3d.transformer_blocks) + list(transformer3d.single_transformer_blocks)) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1315,8 +1274,6 @@ def main(): from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.text_model.encoder.layers) text_encoder = shard_fn(text_encoder) - # shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder_2.encoder.block) - # text_encoder_2 = shard_fn(text_encoder_2) # Move text_encode and vae to gpu and cast to weight_dtype vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) diff --git a/scripts/flux/train_lora.sh b/scripts/flux/train_lora.sh index 5080571..46855fc 100644 --- a/scripts/flux/train_lora.sh +++ b/scripts/flux/train_lora.sh @@ -18,7 +18,7 @@ accelerate launch --mixed_precision="bf16" scripts/flux/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_flux_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/flux2/train.py b/scripts/flux2/train.py index 7c9a423..9411ed1 100644 --- a/scripts/flux2/train.py +++ b/scripts/flux2/train.py @@ -329,7 +329,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -1018,7 +1018,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1038,26 +1038,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): diff --git a/scripts/flux2/train.sh b/scripts/flux2/train.sh index 84271f6..abc3d31 100644 --- a/scripts/flux2/train.sh +++ b/scripts/flux2/train.sh @@ -20,7 +20,7 @@ accelerate launch --mixed_precision="bf16" scripts/flux2/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_flux2" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/flux2/train_lora.py b/scripts/flux2/train_lora.py index 4405278..ddfa389 100644 --- a/scripts/flux2/train_lora.py +++ b/scripts/flux2/train_lora.py @@ -332,7 +332,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -990,7 +990,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1019,15 +1019,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1040,34 +1032,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1348,23 +1318,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(transformer3d.transformer_blocks) + list(transformer3d.single_transformer_blocks)) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial diff --git a/scripts/flux2/train_lora.sh b/scripts/flux2/train_lora.sh index 5b1078d..0a4b5b4 100644 --- a/scripts/flux2/train_lora.sh +++ b/scripts/flux2/train_lora.sh @@ -18,7 +18,7 @@ accelerate launch --mixed_precision="bf16" scripts/flux2/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir_lora" \ + --output_dir="output_dir_flux2_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/flux2_fun/train_control.py b/scripts/flux2_fun/train_control.py index 7ea5e49..ca519ea 100644 --- a/scripts/flux2_fun/train_control.py +++ b/scripts/flux2_fun/train_control.py @@ -333,7 +333,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -1040,7 +1040,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1060,26 +1060,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): diff --git a/scripts/flux2_fun/train_control_distill.py b/scripts/flux2_fun/train_control_distill.py index 7c6b869..3d3d215 100644 --- a/scripts/flux2_fun/train_control_distill.py +++ b/scripts/flux2_fun/train_control_distill.py @@ -337,7 +337,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -1046,7 +1046,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1066,26 +1066,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): diff --git a/scripts/flux2_fun/train_control_distill.sh b/scripts/flux2_fun/train_control_distill.sh index 42b9898..83faf02 100644 --- a/scripts/flux2_fun/train_control_distill.sh +++ b/scripts/flux2_fun/train_control_distill.sh @@ -21,7 +21,7 @@ accelerate launch --mixed_precision="bf16" scripts/flux2_fun/train_control_disti --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir_flux2_control_CFG_Distill" \ + --output_dir="output_dir_flux2_control_distill" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/hunyuanvideo/train.py b/scripts/hunyuanvideo/train.py index a9626ca..9e4da08 100644 --- a/scripts/hunyuanvideo/train.py +++ b/scripts/hunyuanvideo/train.py @@ -167,67 +167,80 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = HunyuanVideoTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, 'transformer'), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - - if args.train_mode != "normal": - raise NotImplementedError("train_mode is not implemented") - else: - pipeline = HunyuanVideoPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - text_encoder_2=accelerator.unwrap_model(text_encoder_2), - tokenizer=tokenizer, - tokenizer_2=tokenizer_2, - transformer=transformer3d_val, - scheduler=scheduler, + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" ) - pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.train_mode != "normal": + raise NotImplementedError("train_mode is not implemented") + else: + pipeline = HunyuanVideoPipeline( + vae=vae, + text_encoder=text_encoder, + text_encoder_2=text_encoder_2, + tokenizer=tokenizer, + tokenizer_2=tokenizer_2, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": raise NotImplementedError("train_mode is not implemented") else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + text_encoder_2.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + text_encoder_2.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + DEFAULT_PROMPT_TEMPLATE = { "template": ( @@ -1026,7 +1039,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1046,26 +1059,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -2054,27 +2047,26 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2082,27 +2074,26 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - args, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/hunyuanvideo/train.sh b/scripts/hunyuanvideo/train.sh index 2ba2234..11a06a6 100644 --- a/scripts/hunyuanvideo/train.sh +++ b/scripts/hunyuanvideo/train.sh @@ -25,7 +25,7 @@ accelerate launch --mixed_precision="bf16" scripts/hunyuanvideo/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_hunyuanvideo" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/hunyuanvideo/train_lora.py b/scripts/hunyuanvideo/train_lora.py index 9d5d918..93066c7 100644 --- a/scripts/hunyuanvideo/train_lora.py +++ b/scripts/hunyuanvideo/train_lora.py @@ -171,70 +171,81 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, text_encoder_2, tokenizer, tokenizer_2, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = HunyuanVideoTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, 'transformer'), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - - if args.train_mode != "normal": - raise NotImplementedError("train_mode is not implemented") - else: - pipeline = HunyuanVideoPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - text_encoder_2=accelerator.unwrap_model(text_encoder_2), - tokenizer=tokenizer, - tokenizer_2=tokenizer_2, - transformer=transformer3d_val, - scheduler=scheduler, + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" ) - pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.train_mode != "normal": + raise NotImplementedError("train_mode is not implemented") + else: + pipeline = HunyuanVideoPipeline( + vae=vae, + text_encoder=text_encoder, + text_encoder_2=text_encoder_2, + tokenizer=tokenizer, + tokenizer_2=tokenizer_2, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": raise NotImplementedError("train_mode is not implemented") else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + text_encoder_2.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + text_encoder_2.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) DEFAULT_PROMPT_TEMPLATE = { "template": ( @@ -1034,7 +1045,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1063,15 +1074,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1084,34 +1087,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1429,22 +1410,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network + transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(transformer3d.transformer_blocks) + list(transformer3d.single_transformer_blocks)) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1534,6 +1505,16 @@ def main(): else: initial_global_step = 0 + # function for saving/removing + def save_model(ckpt_file, unwrapped_nw): + os.makedirs(args.output_dir, exist_ok=True) + accelerator.print(f"\nsaving checkpoint: {ckpt_file}") + if isinstance(unwrapped_nw, dict): + from safetensors.torch import save_file + save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"}) + return ckpt_file + unwrapped_nw.save_weights(ckpt_file, weight_dtype, None) + progress_bar = tqdm( range(0, args.max_train_steps), initial=initial_global_step, @@ -2041,21 +2022,20 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2063,21 +2043,20 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - text_encoder_2, - tokenizer, - tokenizer_2, - transformer3d, - network, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + text_encoder_2, + tokenizer, + tokenizer_2, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/hunyuanvideo/train_lora.sh b/scripts/hunyuanvideo/train_lora.sh index 6200750..ec68d56 100644 --- a/scripts/hunyuanvideo/train_lora.sh +++ b/scripts/hunyuanvideo/train_lora.sh @@ -23,7 +23,7 @@ accelerate launch --mixed_precision="bf16" scripts/hunyuanvideo/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir_lora" \ + --output_dir="output_dir_hunyuanvideo_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/longcatvideo/train.py b/scripts/longcatvideo/train.py index f63ceac..4440c86 100644 --- a/scripts/longcatvideo/train.py +++ b/scripts/longcatvideo/train.py @@ -47,7 +47,8 @@ from einops import rearrange from packaging import version from PIL import Image from torch.distributed.fsdp.fully_sharded_data_parallel import ( - FullOptimStateDictConfig, FullStateDictConfig, ShardedStateDictConfig, ShardedOptimStateDictConfig) + FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig, + ShardedStateDictConfig) from torch.utils.data import RandomSampler from torch.utils.tensorboard import SummaryWriter from torchvision import transforms @@ -62,20 +63,6 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None -from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, - ASPECT_RATIO_RANDOM_CROP_512, - ASPECT_RATIO_RANDOM_CROP_PROB, - AspectRatioBatchImageVideoSampler, - RandomSampler, get_closest_ratio) -from videox_fun.data.dataset_image_video import (ImageVideoDataset, - ImageVideoSampler, - get_random_mask) -from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, - LongCatVideoTransformer3DModel) -from videox_fun.pipeline import WanPipeline, WanI2VPipeline -from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid - from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512, ASPECT_RATIO_RANDOM_CROP_PROB, @@ -84,9 +71,11 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, from videox_fun.data.dataset_image_video import (ImageVideoDataset, ImageVideoSampler, get_random_mask) -from videox_fun.models import (AutoencoderKLLongCatVideo, CLIPModel, UMT5EncoderModel, - LongCatVideoTransformer3DModel) -from videox_fun.pipeline import WanI2VPipeline, WanPipeline +from videox_fun.models import (AutoencoderKLLongCatVideo, AutoencoderKLWan, + CLIPModel, LongCatVideoTransformer3DModel, + UMT5EncoderModel, WanT5EncoderModel) +from videox_fun.pipeline import (LongCatVideoPipeline, WanI2VPipeline, + WanPipeline) from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, @@ -177,117 +166,71 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = LongCatVideoTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, "dit"), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) - - if args.train_mode != "normal": - pipeline = WanI2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) + + pipeline = LongCatVideoPipeline( + vae=vae, + text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d_val, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) - else: - pipeline = WanPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.scale_factor_temporal * vae.config.scale_factor_temporal) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + for i in range(len(args.validation_prompts)): + sample = pipeline( + prompt = args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -974,7 +917,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -994,26 +937,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1068,8 +991,8 @@ def main(): if args.gradient_checkpointing: transformer3d.enable_gradient_checkpointing() elif args.selective_ac > 0: - from videox_fun.utils.ac_handle import apply_checkpointing, partial from videox_fun.models.wan_transformer3d import WanAttentionBlock + from videox_fun.utils.ac_handle import apply_checkpointing, partial apply_selective_ac = partial(apply_checkpointing, block=WanAttentionBlock) apply_selective_ac(transformer3d, p=args.selective_ac) @@ -1350,16 +1273,17 @@ def main(): new_examples['text'], max_length=args.tokenizer_max_length, padding="max_length", - add_special_tokens=True, truncation=True, + add_special_tokens=True, + return_attention_mask=True, return_tensors="pt" ) encoder_hidden_states = text_encoder( - prompt_ids.input_ids, attention_mask=prompt_ids.attention_mask.to(latents.device) - )[0] + prompt_ids.input_ids, attention_mask=prompt_ids.attention_mask + ).last_hidden_state encoder_hidden_states = encoder_hidden_states.unsqueeze(1) - new_examples['encoder_attention_mask'] = prompt_ids.attention_mask new_examples['encoder_hidden_states'] = encoder_hidden_states + new_examples['encoder_attention_mask'] = prompt_ids.attention_mask return new_examples @@ -1403,6 +1327,7 @@ def main(): if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block) text_encoder = shard_fn(text_encoder) @@ -1689,13 +1614,13 @@ def main(): max_length=args.tokenizer_max_length, truncation=True, add_special_tokens=True, + return_attention_mask=True, return_tensors="pt" ) - text_input_ids = prompt_ids.input_ids - prompt_attention_mask = prompt_ids.attention_mask + text_input_ids = prompt_ids.input_ids.to(latents.device) + prompt_attention_mask = prompt_ids.attention_mask.to(latents.device) - seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long() - prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0] + prompt_embeds = text_encoder(text_input_ids, attention_mask=prompt_attention_mask).last_hidden_state prompt_embeds = prompt_embeds.unsqueeze(1) if args.low_vram and not args.enable_text_encoder_in_dataloader: @@ -1850,26 +1775,24 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1877,26 +1800,24 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/longcatvideo/train.sh b/scripts/longcatvideo/train.sh index 1f8e227..37ce2da 100644 --- a/scripts/longcatvideo/train.sh +++ b/scripts/longcatvideo/train.sh @@ -25,7 +25,7 @@ accelerate launch --mixed_precision="bf16" scripts/longcatvideo/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir_longcat_full_finetune" \ + --output_dir="output_dir_longcat" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/longcatvideo/train_lora.py b/scripts/longcatvideo/train_lora.py index fc109df..f318573 100644 --- a/scripts/longcatvideo/train_lora.py +++ b/scripts/longcatvideo/train_lora.py @@ -68,9 +68,10 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, from videox_fun.data.dataset_image_video import (ImageVideoDataset, ImageVideoSampler, get_random_mask) -from videox_fun.models import (AutoencoderKLLongCatVideo, CLIPModel, UMT5EncoderModel, - LongCatVideoTransformer3DModel) -from videox_fun.pipeline import WanI2VPipeline, WanPipeline +from videox_fun.models import (AutoencoderKLLongCatVideo, CLIPModel, + LongCatVideoTransformer3DModel, + UMT5EncoderModel) +from videox_fun.pipeline import LongCatVideoPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, @@ -161,118 +162,73 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = LongCatVideoTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, "dit"), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( - args.pretrained_model_name_or_path, - subfolder="scheduler" - ) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="scheduler" + ) - if args.train_mode != "normal": - pipeline = WanI2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), + pipeline = LongCatVideoPipeline( + vae=vae, + text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d_val, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) - else: - pipeline = WanPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + for i in range(len(args.validation_prompts)): + sample = pipeline( + prompt = args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.scale_factor_temporal * vae.config.scale_factor_temporal) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -951,7 +907,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -980,7 +936,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -993,34 +949,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1298,16 +1232,17 @@ def main(): new_examples['text'], max_length=args.tokenizer_max_length, padding="max_length", - add_special_tokens=True, truncation=True, + add_special_tokens=True, + return_attention_mask=True, return_tensors="pt" ) encoder_hidden_states = text_encoder( - prompt_ids.input_ids, attention_mask=prompt_ids.attention_mask.to(latents.device) - )[0] + prompt_ids.input_ids, attention_mask=prompt_ids.attention_mask + ).last_hidden_state encoder_hidden_states = encoder_hidden_states.unsqueeze(1) - new_examples['encoder_attention_mask'] = prompt_ids.attention_mask new_examples['encoder_hidden_states'] = encoder_hidden_states + new_examples['encoder_attention_mask'] = prompt_ids.attention_mask return new_examples @@ -1349,26 +1284,16 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.encoder.block) text_encoder = shard_fn(text_encoder) @@ -1719,13 +1644,13 @@ def main(): max_length=args.tokenizer_max_length, truncation=True, add_special_tokens=True, + return_attention_mask=True, return_tensors="pt" ) - text_input_ids = prompt_ids.input_ids - prompt_attention_mask = prompt_ids.attention_mask + text_input_ids = prompt_ids.input_ids.to(latents.device) + prompt_attention_mask = prompt_ids.attention_mask.to(latents.device) - seq_lens = prompt_attention_mask.gt(0).sum(dim=1).long() - prompt_embeds = text_encoder(text_input_ids.to(latents.device), attention_mask=prompt_attention_mask.to(latents.device))[0] + prompt_embeds = text_encoder(text_input_ids, attention_mask=prompt_attention_mask).last_hidden_state prompt_embeds = prompt_embeds.unsqueeze(1) if args.low_vram and not args.enable_text_encoder_in_dataloader: @@ -1871,20 +1796,18 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1892,20 +1815,18 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/longcatvideo/train_lora.sh b/scripts/longcatvideo/train_lora.sh index 0423d30..655f856 100644 --- a/scripts/longcatvideo/train_lora.sh +++ b/scripts/longcatvideo/train_lora.sh @@ -23,7 +23,7 @@ accelerate launch --mixed_precision="bf16" scripts/longcatvideo/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir_longcat_lora" \ + --output_dir="output_dir_longcatvideo_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/qwenimage/train.py b/scripts/qwenimage/train.py index 14e82f6..4b5fba1 100644 --- a/scripts/qwenimage/train.py +++ b/scripts/qwenimage/train.py @@ -151,7 +151,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -827,7 +827,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -847,26 +847,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1220,7 +1200,7 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers) diff --git a/scripts/qwenimage/train.sh b/scripts/qwenimage/train.sh index 2e81f31..52932d4 100644 --- a/scripts/qwenimage/train.sh +++ b/scripts/qwenimage/train.sh @@ -20,7 +20,7 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_qwenimage" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/qwenimage/train_edit.py b/scripts/qwenimage/train_edit.py index 1c39914..2f02e7d 100644 --- a/scripts/qwenimage/train_edit.py +++ b/scripts/qwenimage/train_edit.py @@ -155,7 +155,7 @@ def log_validation(vae, text_encoder, tokenizer, processor, transformer3d, args, vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, processor=processor, scheduler=scheduler, ) @@ -164,7 +164,7 @@ def log_validation(vae, text_encoder, tokenizer, processor, transformer3d, args, vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, processor=processor, scheduler=scheduler, ) @@ -863,7 +863,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -883,26 +883,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1269,7 +1249,7 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.language_model.layers) diff --git a/scripts/qwenimage/train_edit.sh b/scripts/qwenimage/train_edit.sh index 50c6f51..943d248 100644 --- a/scripts/qwenimage/train_edit.sh +++ b/scripts/qwenimage/train_edit.sh @@ -20,7 +20,7 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train_edit.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_qwenimage_edit" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/qwenimage/train_edit_lora.py b/scripts/qwenimage/train_edit_lora.py index d54ca3c..d625d77 100644 --- a/scripts/qwenimage/train_edit_lora.py +++ b/scripts/qwenimage/train_edit_lora.py @@ -162,7 +162,7 @@ def log_validation(vae, text_encoder, tokenizer, processor, transformer3d, netwo vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, processor=processor, scheduler=scheduler, ) @@ -171,7 +171,7 @@ def log_validation(vae, text_encoder, tokenizer, processor, transformer3d, netwo vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, processor=processor, scheduler=scheduler, ) @@ -870,7 +870,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -899,15 +899,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -920,34 +912,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1262,23 +1232,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=transformer3d.transformer_blocks) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial diff --git a/scripts/qwenimage/train_edit_lora.sh b/scripts/qwenimage/train_edit_lora.sh index ec17fe7..d73a3ba 100644 --- a/scripts/qwenimage/train_edit_lora.sh +++ b/scripts/qwenimage/train_edit_lora.sh @@ -18,7 +18,7 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train_edit_lora.py --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir_lora" \ + --output_dir="output_dir_qwenimage_edit_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/qwenimage/train_lora.py b/scripts/qwenimage/train_lora.py index 18ba850..d9fa0ef 100644 --- a/scripts/qwenimage/train_lora.py +++ b/scripts/qwenimage/train_lora.py @@ -150,7 +150,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -826,7 +826,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -855,15 +855,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -876,34 +868,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1205,23 +1175,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=transformer3d.transformer_blocks) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial diff --git a/scripts/qwenimage/train_lora.sh b/scripts/qwenimage/train_lora.sh index 7a35068..a578fed 100644 --- a/scripts/qwenimage/train_lora.sh +++ b/scripts/qwenimage/train_lora.sh @@ -18,7 +18,7 @@ accelerate launch --mixed_precision="bf16" scripts/qwenimage/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir_lora" \ + --output_dir="output_dir_qwenimage_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/qwenimage_fun/train_control.py b/scripts/qwenimage_fun/train_control.py index 55b2c5c..5619806 100644 --- a/scripts/qwenimage_fun/train_control.py +++ b/scripts/qwenimage_fun/train_control.py @@ -151,7 +151,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -857,7 +857,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -877,26 +877,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1273,7 +1253,7 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import shard_model diff --git a/scripts/qwenimage_instantx/train_control.py b/scripts/qwenimage_instantx/train_control.py index 8fa883b..8ee0c8e 100644 --- a/scripts/qwenimage_instantx/train_control.py +++ b/scripts/qwenimage_instantx/train_control.py @@ -151,7 +151,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, cn_transformer, vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, controlnet=cn_transformer, scheduler=scheduler, ) @@ -865,7 +865,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -885,26 +885,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1315,7 +1295,7 @@ def main(): cn_transformer, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import shard_model diff --git a/scripts/turbodiffusion/train_distill.py b/scripts/turbodiffusion/train_distill.py index ebe1b9e..862ac50 100644 --- a/scripts/turbodiffusion/train_distill.py +++ b/scripts/turbodiffusion/train_distill.py @@ -81,7 +81,9 @@ from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanTransformer3DModel) from videox_fun.pipeline import WanI2VPipeline, WanPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -169,120 +171,109 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = WanI2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) - else: - pipeline = WanPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + + if args.train_mode != "normal": + pipeline = WanI2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + else: + pipeline = WanPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - num_inference_steps = len(args.denoising_step_list), - guidance_scale = 1.0, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = len(args.denoising_step_indices_list), + guidance_scale = 1.0, + ).videos - video = input_video, - mask_video = input_video_mask, - num_inference_steps = len(args.denoising_step_list), - guidance_scale = 1.0, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, - num_inference_steps = len(args.denoising_step_list), - guidance_scale = 1.0, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = len(args.denoising_step_indices_list), + guidance_scale = 1.0, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -354,6 +345,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--negative_prompt", type=str, @@ -912,6 +910,8 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')), ) clip_image_encoder = clip_image_encoder.eval() + else: + clip_image_encoder = None # Get Transformer generator_transformer3d = TurboWanTransformer3DModel.from_pretrained( @@ -987,7 +987,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1007,26 +1007,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1536,13 +1516,14 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler= accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - if fsdp_stage != 0: + + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) real_score_transformer3d = shard_fn(real_score_transformer3d) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -2268,20 +2249,19 @@ def main(): accelerator_fake_score_transformer3d.save_state(fake_score_save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - generator_transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + generator_transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"denoising_loss": denoising_loss.detach().item(), "dmd_loss": dmd_loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2289,20 +2269,19 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - generator_transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + generator_transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/turbodiffusion/train_distill.sh b/scripts/turbodiffusion/train_distill.sh index 889ec62..3941d3b 100644 --- a/scripts/turbodiffusion/train_distill.sh +++ b/scripts/turbodiffusion/train_distill.sh @@ -28,7 +28,7 @@ accelerate launch --mixed_precision="bf16" scripts/turbodiffusion/train_distill. --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir_distill_turbodiffusion" \ + --output_dir="output_dir_turbodiffusion_distill" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index 9e0a954..bc889b7 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -48,7 +48,8 @@ from omegaconf import OmegaConf from packaging import version from PIL import Image from torch.distributed.fsdp.fully_sharded_data_parallel import ( - FullOptimStateDictConfig, FullStateDictConfig, ShardedStateDictConfig, ShardedOptimStateDictConfig) + FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig, + ShardedStateDictConfig) from torch.utils.data import RandomSampler from torch.utils.tensorboard import SummaryWriter from torchvision import transforms @@ -64,18 +65,20 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, - ASPECT_RATIO_RANDOM_CROP_512, - ASPECT_RATIO_RANDOM_CROP_PROB, - AspectRatioBatchImageVideoSampler, - RandomSampler, get_closest_ratio) + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) from videox_fun.data.dataset_image_video import (ImageVideoDataset, - ImageVideoSampler, - get_random_mask) + ImageVideoSampler, + get_random_mask) from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, - WanTransformer3DModel) -from videox_fun.pipeline import WanPipeline, WanI2VPipeline + WanTransformer3DModel) +from videox_fun.pipeline import WanI2VPipeline, WanPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -163,116 +166,109 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = WanI2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) - else: - pipeline = WanPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + + if args.train_mode != "normal": + pipeline = WanI2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + else: + pipeline = WanPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -344,6 +340,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -902,6 +905,8 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')), ) clip_image_encoder = clip_image_encoder.eval() + else: + clip_image_encoder = None # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -971,7 +976,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -991,26 +996,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1066,8 +1051,8 @@ def main(): if args.gradient_checkpointing: transformer3d.enable_gradient_checkpointing() elif args.selective_ac > 0: - from videox_fun.utils.ac_handle import apply_checkpointing, partial from videox_fun.models.wan_transformer3d import WanAttentionBlock + from videox_fun.utils.ac_handle import apply_checkpointing, partial apply_selective_ac = partial(apply_checkpointing, block=WanAttentionBlock) apply_selective_ac(transformer3d, p=args.selective_ac) @@ -1398,8 +1383,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -1866,27 +1852,26 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1894,27 +1879,26 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1/train.sh b/scripts/wan2.1/train.sh index 0afb828..72bfcf8 100755 --- a/scripts/wan2.1/train.sh +++ b/scripts/wan2.1/train.sh @@ -26,7 +26,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.1" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -69,7 +69,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train.py \ # --lr_scheduler="constant_with_warmup" \ # --lr_warmup_steps=100 \ # --seed=42 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_wan2.1" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1/train_distill.py b/scripts/wan2.1/train_distill.py index 8d5c13f..390eef6 100644 --- a/scripts/wan2.1/train_distill.py +++ b/scripts/wan2.1/train_distill.py @@ -80,7 +80,9 @@ from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, WanTransformer3DModel) from videox_fun.pipeline import WanI2VPipeline, WanPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -168,120 +170,109 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = WanI2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) - else: - pipeline = WanPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + + if args.train_mode != "normal": + pipeline = WanI2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + else: + pipeline = WanPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - num_inference_steps = len(args.denoising_step_list), - guidance_scale = 1.0, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = len(args.denoising_step_indices_list), + guidance_scale = 1.0, + ).videos - video = input_video, - mask_video = input_video_mask, - num_inference_steps = len(args.denoising_step_list), - guidance_scale = 1.0, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, - num_inference_steps = len(args.denoising_step_list), - guidance_scale = 1.0, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = len(args.denoising_step_indices_list), + guidance_scale = 1.0, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -353,6 +344,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--negative_prompt", type=str, @@ -911,6 +909,8 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')), ) clip_image_encoder = clip_image_encoder.eval() + else: + clip_image_encoder = None # Get Transformer generator_transformer3d = WanTransformer3DModel.from_pretrained( @@ -986,7 +986,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1006,26 +1006,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1535,13 +1515,13 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler= accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) real_score_transformer3d = shard_fn(real_score_transformer3d) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -2267,20 +2247,19 @@ def main(): accelerator_fake_score_transformer3d.save_state(fake_score_save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - generator_transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + generator_transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"denoising_loss": denoising_loss.detach().item(), "dmd_loss": dmd_loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2288,20 +2267,19 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - generator_transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + generator_transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1/train_distill.sh b/scripts/wan2.1/train_distill.sh index f1b5dc2..dbe4524 100644 --- a/scripts/wan2.1/train_distill.sh +++ b/scripts/wan2.1/train_distill.sh @@ -27,7 +27,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_distill.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir_distill" \ + --output_dir="output_dir_wan2.1_distill" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1/train_distill_lora.py b/scripts/wan2.1/train_distill_lora.py index a62d69e..428c99c 100644 --- a/scripts/wan2.1/train_distill_lora.py +++ b/scripts/wan2.1/train_distill_lora.py @@ -83,7 +83,9 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -169,119 +171,113 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) - if args.train_mode != "normal": - pipeline = WanI2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, - ) - else: - pipeline = WanPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.train_mode != "normal": + pipeline = WanI2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + else: + pipeline = WanPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = len(args.denoising_step_indices_list), + guidance_scale = 1.0, + ).videos - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = len(args.denoising_step_indices_list), + guidance_scale = 1.0, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -353,6 +349,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--negative_prompt", type=str, @@ -933,6 +936,8 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')), ) clip_image_encoder = clip_image_encoder.eval() + else: + clip_image_encoder = None # Get Transformer generator_transformer3d = WanTransformer3DModel.from_pretrained( @@ -1028,7 +1033,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1057,15 +1062,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1078,34 +1075,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1562,7 +1537,7 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - elif fsdp_stage != 0: + else: generator_transformer3d.network = network generator_transformer3d = generator_transformer3d.to(dtype=weight_dtype) generator_transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( @@ -1573,21 +1548,6 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - fake_score_network, critic_optimizer, fake_score_lr_scheduler= accelerator_fake_score_transformer3d.prepare( - fake_score_network, critic_optimizer, fake_score_lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - generator_transformer3d = shard_fn(generator_transformer3d) - fake_score_transformer3d = shard_fn(fake_score_transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -2362,21 +2322,20 @@ def main(): accelerator_fake_score_transformer3d.save_state(fake_score_save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - generator_transformer3d, - network, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + generator_transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"denoising_loss": denoising_loss.detach().item(), "dmd_loss": dmd_loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2384,21 +2343,20 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - generator_transformer3d, - network, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + generator_transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1/train_distill_lora.sh b/scripts/wan2.1/train_distill_lora.sh index 2d55cb1..7801835 100644 --- a/scripts/wan2.1/train_distill_lora.sh +++ b/scripts/wan2.1/train_distill_lora.sh @@ -25,7 +25,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_distill_lora.py --learning_rate=1e-05 \ --learning_rate_critic=1e-06 \ --seed=42 \ - --output_dir="output_dir_distill" \ + --output_dir="output_dir_wan2.1_distill_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index 4c796e6..a3e50ef 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -75,7 +75,9 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -161,119 +163,113 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) - if args.train_mode != "normal": - pipeline = WanI2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, - ) - else: - pipeline = WanPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.train_mode != "normal": + pipeline = WanI2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + else: + pipeline = WanPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -345,6 +341,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -906,6 +909,8 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')), ) clip_image_encoder = clip_image_encoder.eval() + else: + clip_image_encoder = None # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -971,7 +976,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1000,7 +1005,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1013,34 +1018,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1370,23 +1353,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1910,21 +1882,20 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1932,21 +1903,20 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1/train_lora.sh b/scripts/wan2.1/train_lora.sh index be0ece9..348d030 100755 --- a/scripts/wan2.1/train_lora.sh +++ b/scripts/wan2.1/train_lora.sh @@ -24,7 +24,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.1_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -67,7 +67,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1/train_lora.py \ # --checkpointing_steps=50 \ # --learning_rate=1e-04 \ # --seed=42 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_wan2.1_lora" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index 7c2a4fa..01551ec 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -77,7 +77,9 @@ from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, WanTransformer3DModel) from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -146,116 +148,109 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = WanFunInpaintPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) ) - else: - pipeline = WanFunPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + + if args.train_mode != "normal": + pipeline = WanFunInpaintPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + else: + pipeline = WanFunPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -314,6 +309,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -866,6 +868,8 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')), ) clip_image_encoder = clip_image_encoder.eval() + else: + clip_image_encoder = None # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -935,7 +939,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -955,26 +959,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1395,8 +1379,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -1869,27 +1854,26 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1897,27 +1881,26 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1_fun/train.sh b/scripts/wan2.1_fun/train.sh index 385f033..b0bc2b3 100755 --- a/scripts/wan2.1_fun/train.sh +++ b/scripts/wan2.1_fun/train.sh @@ -26,7 +26,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.1_fun" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -69,7 +69,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train.py \ # --lr_scheduler="constant_with_warmup" \ # --lr_warmup_steps=100 \ # --seed=42 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_wan2.1_fun" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index 95df6d1..efa4761 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -77,7 +77,8 @@ from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, WanTransformer3DModel) from videox_fun.pipeline import WanFunControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -112,64 +113,80 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + pipeline = WanFunControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = WanFunControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + for i in range(len(args.validation_prompts)): + import cv2 + cap = cv2.VideoCapture(args.validation_paths[i]) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + + input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - control_video = input_video, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + control_video = input_video, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -868,7 +885,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -888,26 +905,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1399,8 +1396,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -1904,27 +1902,26 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1932,27 +1929,26 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1_fun/train_control.sh b/scripts/wan2.1_fun/train_control.sh index 4cf7fbd..8aec313 100755 --- a/scripts/wan2.1_fun/train_control.sh +++ b/scripts/wan2.1_fun/train_control.sh @@ -26,7 +26,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.1_fun_control" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index 26e0d42..cad7f2a 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -80,7 +80,8 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -113,70 +114,84 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + pipeline = WanFunControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = WanFunControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + for i in range(len(args.validation_prompts)): + import cv2 + cap = cv2.VideoCapture(args.validation_paths[i]) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + + input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + control_video = input_video, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - control_video = input_video, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -874,7 +889,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -903,15 +918,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -924,34 +931,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1390,23 +1375,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1967,21 +1941,20 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1989,21 +1962,20 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1_fun/train_control_lora.sh b/scripts/wan2.1_fun/train_control_lora.sh index a5df6dd..ad745dc 100755 --- a/scripts/wan2.1_fun/train_control_lora.sh +++ b/scripts/wan2.1_fun/train_control_lora.sh @@ -24,7 +24,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_control_lora --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.1_fun_control_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index 9fbedd4..422cab0 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -76,7 +76,9 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -142,119 +144,113 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = WanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) - if args.train_mode != "normal": - pipeline = WanFunInpaintPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, - ) - else: - pipeline = WanFunPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.train_mode != "normal": + pipeline = WanFunInpaintPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + else: + pipeline = WanFunPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): + for i in range(len(args.validation_prompts)): if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -313,6 +309,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -868,6 +871,8 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['image_encoder_kwargs'].get('image_encoder_subpath', 'image_encoder')), ) clip_image_encoder = clip_image_encoder.eval() + else: + clip_image_encoder = None # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -933,7 +938,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -962,15 +967,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -983,34 +980,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1378,23 +1353,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1924,21 +1888,20 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1946,21 +1909,20 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1_fun/train_lora.sh b/scripts/wan2.1_fun/train_lora.sh index 4d0aa0c..970f803 100755 --- a/scripts/wan2.1_fun/train_lora.sh +++ b/scripts/wan2.1_fun/train_lora.sh @@ -24,7 +24,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.1_fun_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -68,7 +68,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_fun/train_lora.py \ # --checkpointing_steps=50 \ # --learning_rate=1e-04 \ # --seed=42 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_wan2.1_fun_lora" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.1_vace/train.py b/scripts/wan2.1_vace/train.py index 520f2da..aacd0aa 100644 --- a/scripts/wan2.1_vace/train.py +++ b/scripts/wan2.1_vace/train.py @@ -78,7 +78,8 @@ from videox_fun.models import (AutoencoderKLWan, CLIPModel, VaceWanTransformer3DModel, WanT5EncoderModel) from videox_fun.pipeline import WanVacePipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -111,71 +112,86 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) - transformer3d_val = VaceWanTransformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + pipeline = WanVacePipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = WanVacePipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + for i in range(len(args.validation_prompts)): + import cv2 + cap = cv2.VideoCapture(args.validation_paths[i]) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - control_video, _, _, _ = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[height, width]) + control_video, _, _, _ = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - video = inpaint_video, - mask_video = inpaint_video_mask, - control_video = control_video, - subject_ref_images = None, - vace_context_scale = 1, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = control_video, + subject_ref_images = None, + vace_context_scale = 1, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -854,7 +870,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -874,26 +890,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1395,8 +1391,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -1999,27 +1996,25 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2027,27 +2022,25 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.1_vace/train.sh b/scripts/wan2.1_vace/train.sh index 867964f..95f7ad2 100644 --- a/scripts/wan2.1_vace/train.sh +++ b/scripts/wan2.1_vace/train.sh @@ -25,7 +25,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.1_vace/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.1_vace" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2/train.py b/scripts/wan2.2/train.py index e42adf9..526cc3a 100644 --- a/scripts/wan2.2/train.py +++ b/scripts/wan2.2/train.py @@ -48,7 +48,8 @@ from omegaconf import OmegaConf from packaging import version from PIL import Image from torch.distributed.fsdp.fully_sharded_data_parallel import ( - FullOptimStateDictConfig, FullStateDictConfig, ShardedStateDictConfig, ShardedOptimStateDictConfig) + FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig, + ShardedStateDictConfig) from torch.utils.data import RandomSampler from torch.utils.tensorboard import SummaryWriter from torchvision import transforms @@ -64,18 +65,20 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, - ASPECT_RATIO_RANDOM_CROP_512, - ASPECT_RATIO_RANDOM_CROP_PROB, - AspectRatioBatchImageVideoSampler, - RandomSampler, get_closest_ratio) + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) from videox_fun.data.dataset_image_video import (ImageVideoDataset, - ImageVideoSampler, - get_random_mask) -from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel, - Wan2_2Transformer3DModel) -from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline + ImageVideoSampler, + get_random_mask) +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, + Wan2_2Transformer3DModel, WanT5EncoderModel) +from videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -163,152 +166,132 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = Wan2_2I2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - else: - pipeline = Wan2_2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + if args.train_mode != "normal": + pipeline = Wan2_2I2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + else: + pipeline = Wan2_2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - return images + for i in range(len(args.validation_prompts)): + if args.train_mode != "normal": + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + else: + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -380,6 +363,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1009,7 +999,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1029,26 +1019,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1432,8 +1402,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -1921,26 +1892,25 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1948,26 +1918,25 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2/train.sh b/scripts/wan2.2/train.sh index 1224fb6..78b707e 100644 --- a/scripts/wan2.2/train.sh +++ b/scripts/wan2.2/train.sh @@ -26,7 +26,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -73,7 +73,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train.py \ # --lr_scheduler="constant_with_warmup" \ # --lr_warmup_steps=100 \ # --seed=42 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_wan2.2" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2/train_animate.py b/scripts/wan2.2/train_animate.py index a8123c1..817030c 100644 --- a/scripts/wan2.2/train_animate.py +++ b/scripts/wan2.2/train_animate.py @@ -62,22 +62,22 @@ project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dir for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None +from videox_fun.data import (ImageVideoDataset, ImageVideoSampler, + VideoAnimateDataset, get_random_mask, + process_pose_file, process_pose_params) from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, ASPECT_RATIO_RANDOM_CROP_512, ASPECT_RATIO_RANDOM_CROP_PROB, AspectRatioBatchImageVideoSampler, RandomSampler, get_closest_ratio) -from videox_fun.data import (VideoAnimateDataset, - ImageVideoDataset, - ImageVideoSampler, - get_random_mask, - process_pose_file, - process_pose_params) -from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, - Wan2_2Transformer3DModel_Animate) -from videox_fun.pipeline import Wan2_2FunControlPipeline +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, + CLIPModel, Wan2_2Transformer3DModel_Animate, + WanT5EncoderModel) +from videox_fun.pipeline import Wan2_2AnimatePipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image, + get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -157,100 +157,132 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - transformer3d_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + pipeline = Wan2_2AnimatePipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + pipeline = pipeline.to(accelerator.device) - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - pipeline = Wan2_2FunControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + images = [] + for i in range(len(args.validation_prompts)): + src_root_path = args.validation_paths[i] + src_pose_path = os.path.join(src_root_path, "src_pose.mp4") + src_face_path = os.path.join(src_root_path, "src_face.mp4") + src_ref_path = os.path.join(src_root_path, "src_ref.png") + src_bg_path = os.path.join(src_root_path, "src_bg.mp4") + src_mask_path = os.path.join(src_root_path, "src_mask.mp4") - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + import cv2 + cap = cv2.VideoCapture(src_pose_path) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - control_video = input_video, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + fps = 16 - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + pose_video, _, _, _ = get_video_to_video_latent(src_pose_path, video_length=video_length, sample_size=[height, width], fps=fps, ref_image=None) + face_video, _, _, _ = get_video_to_video_latent(src_face_path, video_length=video_length, sample_size=[512, 512], fps=fps, ref_image=None) + ref_image = get_image(src_ref_path) - return images + if os.path.exists(src_bg_path): + bg_video, _, _, _ = get_video_to_video_latent(src_bg_path, video_length=video_length, sample_size=[height, width], fps=fps, ref_image=None) + mask_video, _, _, _ = get_video_to_video_latent(src_mask_path, video_length=video_length, sample_size=[height, width], fps=fps, ref_image=None) + mask_video = mask_video[:, :1] + replace_flag = True + else: + bg_video = None + mask_video = None + replace_flag = False + + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + guidance_scale = 4.5, + num_inference_steps = 25, + + pose_video = pose_video, + face_video = face_video, + ref_image = ref_image, + bg_video = bg_video, + mask_video = mask_video, + replace_flag = replace_flag, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -450,9 +482,6 @@ def parse_args(): " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." ), ) - parser.add_argument( - "--report_model_info", action="store_true", help="Whether or not to report more info about model (such as norm, grad)." - ) parser.add_argument( "--mixed_precision", type=str, @@ -577,12 +606,6 @@ def parse_args(): default=512, help="Sample size of the video.", ) - parser.add_argument( - "--image_sample_size", - type=int, - default=512, - help="Sample size of the image.", - ) parser.add_argument( "--fix_sample_size", nargs=2, type=int, default=None, @@ -940,7 +963,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -960,26 +983,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1118,7 +1121,6 @@ def main(): if args.fix_sample_size is not None and args.enable_bucket: args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size) - args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) args.training_with_video_token_length = False args.random_hw_adapt = False @@ -1423,8 +1425,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -1951,26 +1954,26 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1978,26 +1981,26 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2/train_animate.sh b/scripts/wan2.2/train_animate.sh index ac27520..5f3ff8c 100644 --- a/scripts/wan2.2/train_animate.sh +++ b/scripts/wan2.2/train_animate.sh @@ -25,7 +25,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2_animate" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2/train_animate_lora.py b/scripts/wan2.2/train_animate_lora.py index 29b5c96..b9573af 100644 --- a/scripts/wan2.2/train_animate_lora.py +++ b/scripts/wan2.2/train_animate_lora.py @@ -77,13 +77,14 @@ from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, Wan2_2Transformer3DModel, Wan2_2Transformer3DModel_Animate, WanT5EncoderModel) -from videox_fun.pipeline import (Wan2_2FunControlPipeline, Wan2_2I2VPipeline, - Wan2_2Pipeline) +from videox_fun.pipeline import Wan2_2AnimatePipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image, + get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -163,104 +164,134 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, config, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - transformer3d_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + pipeline = Wan2_2AnimatePipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + clip_image_encoder=clip_image_encoder, + ) + pipeline = pipeline.to(accelerator.device) - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel_Animate.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - pipeline = Wan2_2FunControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + images = [] + for i in range(len(args.validation_prompts)): + src_root_path = args.validation_paths[i] + src_pose_path = os.path.join(src_root_path, "src_pose.mp4") + src_face_path = os.path.join(src_root_path, "src_face.mp4") + src_ref_path = os.path.join(src_root_path, "src_ref.png") + src_bg_path = os.path.join(src_root_path, "src_bg.mp4") + src_mask_path = os.path.join(src_root_path, "src_mask.mp4") - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + import cv2 + cap = cv2.VideoCapture(src_pose_path) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + fps = 16 - control_video = input_video, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + pose_video, _, _, _ = get_video_to_video_latent(src_pose_path, video_length=video_length, sample_size=[height, width], fps=fps, ref_image=None) + face_video, _, _, _ = get_video_to_video_latent(src_face_path, video_length=video_length, sample_size=[512, 512], fps=fps, ref_image=None) + ref_image = get_image(src_ref_path) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if os.path.exists(src_bg_path): + bg_video, _, _, _ = get_video_to_video_latent(src_bg_path, video_length=video_length, sample_size=[height, width], fps=fps, ref_image=None) + mask_video, _, _, _ = get_video_to_video_latent(src_mask_path, video_length=video_length, sample_size=[height, width], fps=fps, ref_image=None) + mask_video = mask_video[:, :1] + replace_flag = True + else: + bg_video = None + mask_video = None + replace_flag = False - return images + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + guidance_scale = 4.5, + num_inference_steps = 25, + + pose_video = pose_video, + face_video = face_video, + ref_image = ref_image, + bg_video = bg_video, + mask_video = mask_video, + replace_flag = replace_flag, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -604,12 +635,6 @@ def parse_args(): default=512, help="Sample size of the video.", ) - parser.add_argument( - "--image_sample_size", - type=int, - default=512, - help="Sample size of the image.", - ) parser.add_argument( "--fix_sample_size", nargs=2, type=int, default=None, @@ -949,7 +974,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -978,15 +1003,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -999,34 +1016,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1112,7 +1107,6 @@ def main(): if args.fix_sample_size is not None and args.enable_bucket: args.video_sample_size = max(max(args.fix_sample_size), args.video_sample_size) - args.image_sample_size = max(max(args.fix_sample_size), args.image_sample_size) args.training_with_video_token_length = False args.random_hw_adapt = False @@ -1417,23 +1411,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1962,40 +1945,41 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + clip_image_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2/train_animate_lora.sh b/scripts/wan2.2/train_animate_lora.sh index e38a4b1..19cd349 100644 --- a/scripts/wan2.2/train_animate_lora.sh +++ b/scripts/wan2.2/train_animate_lora.sh @@ -23,7 +23,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_animate_lora.py --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_animate_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2/train_distill.py b/scripts/wan2.2/train_distill.py index 9863483..7ca7368 100644 --- a/scripts/wan2.2/train_distill.py +++ b/scripts/wan2.2/train_distill.py @@ -76,11 +76,14 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, from videox_fun.data.dataset_image_video import (ImageVideoDataset, ImageVideoSampler, TextDataset, get_random_mask) -from videox_fun.models import (AutoencoderKLWan, CLIPModel, WanT5EncoderModel, AutoencoderKLWan3_8, - Wan2_2Transformer3DModel) -from videox_fun.pipeline import WanI2VPipeline, WanPipeline +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, + CLIPModel, Wan2_2Transformer3DModel, + WanT5EncoderModel) +from videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -168,152 +171,133 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = Wan2_2I2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - else: - pipeline = Wan2_2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + if args.train_mode != "normal": + pipeline = Wan2_2I2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + else: + pipeline = Wan2_2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - return images + for i in range(len(args.validation_prompts)): + if args.train_mode != "normal": + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 4, + guidance_scale = 1.0, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + else: + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 4, + guidance_scale = 1.0, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -385,6 +369,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--negative_prompt", type=str, @@ -1026,7 +1017,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1046,26 +1037,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1575,13 +1546,13 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler= accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) real_score_transformer3d = shard_fn(real_score_transformer3d) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -2367,19 +2338,18 @@ def main(): accelerator_fake_score_transformer3d.save_state(fake_score_save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - generator_transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + generator_transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"denoising_loss": denoising_loss.detach().item(), "dmd_loss": dmd_loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2387,19 +2357,18 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - generator_transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + generator_transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2/train_distill.sh b/scripts/wan2.2/train_distill.sh index d9af9aa..8c91e4e 100644 --- a/scripts/wan2.2/train_distill.sh +++ b/scripts/wan2.2/train_distill.sh @@ -27,7 +27,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_distill.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir_distill" \ + --output_dir="output_dir_wan2.2_distill" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2/train_distill_lora.py b/scripts/wan2.2/train_distill_lora.py index 812b009..c2babaa 100644 --- a/scripts/wan2.2/train_distill_lora.py +++ b/scripts/wan2.2/train_distill_lora.py @@ -79,12 +79,14 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset, from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, Wan2_2Transformer3DModel, WanT5EncoderModel) -from videox_fun.pipeline import WanI2VPipeline, WanPipeline +from videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -170,158 +172,136 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = Wan2_2I2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - else: - pipeline = Wan2_2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + if args.train_mode != "normal": + pipeline = Wan2_2I2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + else: + pipeline = Wan2_2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - return images + for i in range(len(args.validation_prompts)): + if args.train_mode != "normal": + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 4, + guidance_scale = 1.0, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + else: + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 4, + guidance_scale = 1.0, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -393,6 +373,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--negative_prompt", type=str, @@ -1073,7 +1060,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1102,15 +1089,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1123,34 +1102,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1607,7 +1564,7 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - elif fsdp_stage != 0: + else: generator_transformer3d.network = network generator_transformer3d = generator_transformer3d.to(dtype=weight_dtype) generator_transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( @@ -1618,21 +1575,6 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - fake_score_network, critic_optimizer, fake_score_lr_scheduler= accelerator_fake_score_transformer3d.prepare( - fake_score_network, critic_optimizer, fake_score_lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - generator_transformer3d = shard_fn(generator_transformer3d) - fake_score_transformer3d = shard_fn(fake_score_transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1640,6 +1582,11 @@ def main(): from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) real_score_transformer3d = shard_fn(real_score_transformer3d) + if fsdp_stage != 0 or zero_stage != 0: + from functools import partial + + from videox_fun.dist import set_multi_gpus_devices, shard_model + shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) # Move text_encode and vae to gpu and cast to weight_dtype @@ -2467,20 +2414,19 @@ def main(): accelerator_fake_score_transformer3d.save_state(fake_score_save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - generator_transformer3d, - network, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + generator_transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"denoising_loss": denoising_loss.detach().item(), "dmd_loss": dmd_loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2488,20 +2434,19 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - generator_transformer3d, - network, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + generator_transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2/train_distill_lora.sh b/scripts/wan2.2/train_distill_lora.sh index 241daa3..afba111 100644 --- a/scripts/wan2.2/train_distill_lora.sh +++ b/scripts/wan2.2/train_distill_lora.sh @@ -25,7 +25,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_distill_lora.py --learning_rate=1e-05 \ --learning_rate_critic=1e-06 \ --seed=42 \ - --output_dir="output_dir_distill" \ + --output_dir="output_dir_wan2.2_distill_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2/train_lora.py b/scripts/wan2.2/train_lora.py index 0bb8778..982a25b 100755 --- a/scripts/wan2.2/train_lora.py +++ b/scripts/wan2.2/train_lora.py @@ -75,7 +75,9 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -161,155 +163,136 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = Wan2_2I2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - else: - pipeline = Wan2_2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + if args.train_mode != "normal": + pipeline = Wan2_2I2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + else: + pipeline = Wan2_2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): + if args.train_mode != "normal": + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + else: + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -381,6 +364,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -1015,7 +1005,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1044,15 +1034,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1065,34 +1047,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1424,23 +1384,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1986,20 +1935,19 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2007,20 +1955,19 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2/train_lora.sh b/scripts/wan2.2/train_lora.sh index 112b669..73f8a67 100755 --- a/scripts/wan2.2/train_lora.sh +++ b/scripts/wan2.2/train_lora.sh @@ -24,7 +24,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -72,7 +72,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_lora.py \ # --checkpointing_steps=50 \ # --learning_rate=1e-04 \ # --seed=42 \ -# --output_dir="output_dir" \ +# --output_dir="output_dir_wan2.2_lora" \ # --gradient_checkpointing \ # --mixed_precision="bf16" \ # --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2/train_s2v.py b/scripts/wan2.2/train_s2v.py index 84b3183..5b0ee4c 100644 --- a/scripts/wan2.2/train_s2v.py +++ b/scripts/wan2.2/train_s2v.py @@ -21,6 +21,7 @@ import logging import math import os import pickle +import random import shutil import sys @@ -32,7 +33,6 @@ import torch.nn.functional as F import torch.utils.checkpoint import torchvision.transforms.functional as TF import transformers -import random from accelerate import Accelerator, FullyShardedDataParallelPlugin from accelerate.logging import get_logger from accelerate.state import AcceleratorState @@ -49,7 +49,8 @@ from omegaconf import OmegaConf from packaging import version from PIL import Image from torch.distributed.fsdp.fully_sharded_data_parallel import ( - FullOptimStateDictConfig, FullStateDictConfig, ShardedStateDictConfig, ShardedOptimStateDictConfig) + FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig, + ShardedStateDictConfig) from torch.utils.data import RandomSampler from torch.utils.tensorboard import SummaryWriter from torchvision import transforms @@ -65,19 +66,23 @@ for project_root in project_roots: sys.path.insert(0, project_root) if project_root not in sys.path else None from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, - ASPECT_RATIO_RANDOM_CROP_512, - ASPECT_RATIO_RANDOM_CROP_PROB, - AspectRatioBatchImageVideoSampler, - RandomSampler, get_closest_ratio) + ASPECT_RATIO_RANDOM_CROP_512, + ASPECT_RATIO_RANDOM_CROP_PROB, + AspectRatioBatchImageVideoSampler, + RandomSampler, get_closest_ratio) from videox_fun.data.dataset_image_video import (ImageVideoDataset, - ImageVideoSampler, - get_random_mask) + ImageVideoSampler, + get_random_mask) from videox_fun.data.dataset_video import VideoSpeechControlDataset -from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel, WanAudioEncoder, - Wan2_2Transformer3DModel_S2V) -from videox_fun.pipeline import Wan2_2S2VPipeline, Wan2_2I2VPipeline +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, + Wan2_2Transformer3DModel_S2V, WanAudioEncoder, + WanT5EncoderModel) +from videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2S2VPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + get_video_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -131,144 +136,108 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - pipeline = Wan2_2S2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel_S2V.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel_S2V.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + pipeline = Wan2_2S2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + audio_encoder=audio_encoder, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - return images + for i in range(len(args.validation_prompts)): + start_image = Image.open(args.validation_image_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + + pose_video, _, _, _ = get_video_to_video_latent(None, video_length=video_length, sample_size=(height, width), ref_image=None) + ref_image = get_image_latent(args.validation_image_paths[i], sample_size=(height, width)) + + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + + audio_path = args.validation_audio_paths[i], + pose_video = pose_video, + ref_image = ref_image, + init_first_frame = False, + num_inference_steps = 25, + guidance_scale = 4.5, + fps = 16, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -340,6 +309,20 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_image_paths", + type=str, + default=None, + nargs="+", + help=("A set of images evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--validation_audio_paths", + type=str, + default=None, + nargs="+", + help=("A set of audios evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -680,7 +663,7 @@ def parse_args(): parser.add_argument( "--boundary_type", type=str, - default="low", + default="full", help=( 'The format of training data. Support `"low"` and `"high"`' ), @@ -976,7 +959,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -996,26 +979,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1453,8 +1416,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -2006,26 +1970,26 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + audio_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2033,26 +1997,26 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + audio_encoder, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2/train_s2v.sh b/scripts/wan2.2/train_s2v.sh index a8d2d95..83c9797 100644 --- a/scripts/wan2.2/train_s2v.sh +++ b/scripts/wan2.2/train_s2v.sh @@ -24,7 +24,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_s2v.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2_s2v" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ @@ -35,5 +35,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_s2v.py \ --training_with_video_token_length \ --enable_bucket \ --uniform_sampling \ + --boundary_type="full" \ + --control_ref_image="random" \ --low_vram \ --trainable_modules "." diff --git a/scripts/wan2.2/train_s2v_lora.py b/scripts/wan2.2/train_s2v_lora.py index 4f0af45..f0d5a6e 100644 --- a/scripts/wan2.2/train_s2v_lora.py +++ b/scripts/wan2.2/train_s2v_lora.py @@ -89,7 +89,8 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -145,144 +146,110 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, audio_encoder, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel_S2V.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - pipeline = Wan2_2S2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + pipeline = Wan2_2S2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + audio_encoder=audio_encoder, + ) + pipeline = pipeline.to(accelerator.device) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - return images + for i in range(len(args.validation_prompts)): + start_image = Image.open(args.validation_image_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.video_sample_size * args.video_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + + pose_video, _, _, _ = get_video_to_video_latent(None, video_length=video_length, sample_size=(height, width), ref_image=None) + ref_image = get_image_latent(args.validation_image_paths[i], sample_size=(height, width)) + + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + + audio_path = args.validation_audio_paths[i], + pose_video = pose_video, + ref_image = ref_image, + init_first_frame = False, + num_inference_steps = 25, + guidance_scale = 4.5, + fps = 16, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def linear_decay(initial_value, final_value, total_steps, current_step): if current_step >= total_steps: @@ -354,6 +321,20 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_image_paths", + type=str, + default=None, + nargs="+", + help=("A set of images evaluated every `--validation_epochs` and logged to `--report_to`."), + ) + parser.add_argument( + "--validation_audio_paths", + type=str, + default=None, + nargs="+", + help=("A set of audios evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -711,7 +692,7 @@ def parse_args(): parser.add_argument( "--boundary_type", type=str, - default="low", + default="full", help=( 'The format of training data. Support `"low"` and `"high"`' ), @@ -999,7 +980,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1028,15 +1009,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1049,34 +1022,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1461,23 +1412,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -2034,19 +1974,20 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + audio_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2054,19 +1995,20 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + audio_encoder, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2/train_s2v_lora.sh b/scripts/wan2.2/train_s2v_lora.sh index f52c617..e9a07b8 100644 --- a/scripts/wan2.2/train_s2v_lora.sh +++ b/scripts/wan2.2/train_s2v_lora.sh @@ -23,7 +23,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2/train_s2v_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2_s2v_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2_fun/train.py b/scripts/wan2.2_fun/train.py index 42d311d..68d30b5 100644 --- a/scripts/wan2.2_fun/train.py +++ b/scripts/wan2.2_fun/train.py @@ -73,11 +73,14 @@ from videox_fun.data.bucket_sampler import (ASPECT_RATIO_512, from videox_fun.data.dataset_image_video import (ImageVideoDataset, ImageVideoSampler, get_random_mask) -from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, - Wan2_2Transformer3DModel) -from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, + CLIPModel, Wan2_2Transformer3DModel, + WanT5EncoderModel) +from videox_fun.pipeline import Wan2_2I2VPipeline, Wan2_2Pipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -146,161 +149,131 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = Wan2_2I2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - else: - pipeline = Wan2_2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + if args.train_mode != "normal": + pipeline = Wan2_2I2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + else: + pipeline = Wan2_2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - return images + for i in range(len(args.validation_prompts)): + if args.train_mode != "normal": + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + else: + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -359,6 +332,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -988,7 +968,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1008,26 +988,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1449,8 +1409,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -1926,26 +1887,25 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1953,26 +1913,25 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2_fun/train.sh b/scripts/wan2.2_fun/train.sh index dc1efb1..e48d7d2 100644 --- a/scripts/wan2.2_fun/train.sh +++ b/scripts/wan2.2_fun/train.sh @@ -26,7 +26,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2_fun" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2_fun/train_control.py b/scripts/wan2.2_fun/train_control.py index 14acef5..2a9a4d3 100644 --- a/scripts/wan2.2_fun/train_control.py +++ b/scripts/wan2.2_fun/train_control.py @@ -73,11 +73,13 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, get_random_mask, process_pose_file, process_pose_params) -from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel, - Wan2_2Transformer3DModel) +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, + CLIPModel, Wan2_2Transformer3DModel, + WanT5EncoderModel) from videox_fun.pipeline import Wan2_2FunControlPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -148,98 +150,107 @@ logger = get_logger(__name__, log_level="INFO") def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - pipeline = Wan2_2FunControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + pipeline = Wan2_2FunControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + for i in range(len(args.validation_prompts)): + import cv2 + cap = cv2.VideoCapture(args.validation_paths[i]) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - control_video = input_video, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[height, width]) + input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + control_video = input_video, + video = inpaint_video, + mask_video = inpaint_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -955,7 +966,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -975,26 +986,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1504,8 +1495,9 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial + from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) @@ -2071,26 +2063,25 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2098,26 +2089,25 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2_fun/train_control.sh b/scripts/wan2.2_fun/train_control.sh index 0ae5eab..8fd71a6 100644 --- a/scripts/wan2.2_fun/train_control.sh +++ b/scripts/wan2.2_fun/train_control.sh @@ -23,7 +23,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control.py \ --num_train_epochs=100 \ --checkpointing_steps=50 \ --learning_rate=2e-05 \ - --lr_scheduler="constant_with_warmup" \ + --output_dir="output_dir_wan2.2_fun_control" \ --lr_warmup_steps=100 \ --seed=42 \ --output_dir="output_dir" \ diff --git a/scripts/wan2.2_fun/train_control_lora.py b/scripts/wan2.2_fun/train_control_lora.py index 4f9018e..968a09f 100644 --- a/scripts/wan2.2_fun/train_control_lora.py +++ b/scripts/wan2.2_fun/train_control_lora.py @@ -81,7 +81,8 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -150,105 +151,111 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + pipeline = Wan2_2FunControlPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - pipeline = Wan2_2FunControlPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) + for i in range(len(args.validation_prompts)): + import cv2 + cap = cv2.VideoCapture(args.validation_paths[i]) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[height, width]) + input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + control_video = input_video, + video = inpaint_video, + mask_video = inpaint_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - control_video = input_video, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -963,7 +970,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -992,15 +999,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1013,34 +1012,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1497,23 +1474,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -2136,20 +2102,19 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2157,20 +2122,19 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + config, + args, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2_fun/train_control_lora.sh b/scripts/wan2.2_fun/train_control_lora.sh index ebafdef..91370d6 100644 --- a/scripts/wan2.2_fun/train_control_lora.sh +++ b/scripts/wan2.2_fun/train_control_lora.sh @@ -24,7 +24,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control_lora --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2_fun_control_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2_fun/train_lora.py b/scripts/wan2.2_fun/train_lora.py index be42be3..a229cac 100644 --- a/scripts/wan2.2_fun/train_lora.py +++ b/scripts/wan2.2_fun/train_lora.py @@ -76,7 +76,9 @@ from videox_fun.utils.discrete_sampler import DiscreteSampling from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, create_network, merge_lora, unmerge_lora) -from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, + save_videos_grid) if is_wandb_available(): import wandb @@ -143,155 +145,135 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config, args, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) - - if args.train_mode != "normal": - pipeline = Wan2_2I2VPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - else: - pipeline = Wan2_2Pipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - ) - pipeline = pipeline.to(accelerator.device) - - pipeline = merge_lora( - pipeline, None, 1, accelerator.device, state_dict=accelerator.unwrap_model(network).state_dict(), transformer_only=True - ) - - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) - - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - if args.train_mode != "normal": - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - video_length = 1 - input_video, input_video_mask, _ = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - guidance_scale = 6.0, - generator = generator, - - video = input_video, - mask_video = input_video_mask, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) else: - with torch.autocast("cuda", dtype=weight_dtype): - sample = pipeline( - args.validation_prompts[i], - num_frames = args.video_sample_n_frames, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = Wan2_2Transformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - sample = pipeline( - args.validation_prompts[i], - num_frames = 1, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif")) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + if args.train_mode != "normal": + pipeline = Wan2_2I2VPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + else: + pipeline = Wan2_2Pipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") + + for i in range(len(args.validation_prompts)): + if args.train_mode != "normal": + start_image = Image.open(args.validation_paths[i]) + width, height = start_image.width, start_image.height + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + input_video, input_video_mask, _ = get_image_to_video_latent(args.validation_paths[i], None, video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, + + video = input_video, + mask_video = input_video_mask, + num_inference_steps = 25, + guidance_scale = 4.5, + boundary = config['transformer_additional_kwargs'].get('boundary', 0.900) + ).videos + + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + else: + sample = pipeline( + args.validation_prompts[i], + num_frames = args.video_sample_n_frames, + negative_prompt = "bad detailed", + height = args.video_sample_size, + width = args.video_sample_size, + generator = generator, + num_inference_steps = 25, + guidance_scale = 4.5, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) + + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + transformer3d.to(accelerator.device, dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -350,6 +332,13 @@ def parse_args(): nargs="+", help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."), ) + parser.add_argument( + "--validation_paths", + type=str, + default=None, + nargs="+", + help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."), + ) parser.add_argument( "--output_dir", type=str, @@ -978,7 +967,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1007,15 +996,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -1028,34 +1009,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1424,23 +1383,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial @@ -1974,20 +1922,19 @@ def main(): accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -1995,20 +1942,19 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - log_validation( - vae, - text_encoder, - tokenizer, - transformer3d, - network, - config, - args, - accelerator, - weight_dtype, - global_step, - ) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + network, + args, + config, + accelerator, + weight_dtype, + global_step, + ) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2_fun/train_lora.sh b/scripts/wan2.2_fun/train_lora.sh index 3ca2d05..b2b0522 100644 --- a/scripts/wan2.2_fun/train_lora.sh +++ b/scripts/wan2.2_fun/train_lora.sh @@ -24,7 +24,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_lora.py \ --checkpointing_steps=50 \ --learning_rate=1e-04 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2_fun_lora" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/wan2.2_vace_fun/train.py b/scripts/wan2.2_vace_fun/train.py index aeed55e..62fd968 100644 --- a/scripts/wan2.2_vace_fun/train.py +++ b/scripts/wan2.2_vace_fun/train.py @@ -74,11 +74,13 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset, padding_image, process_pose_file, process_pose_params) -from videox_fun.models import (AutoencoderKLWan, CLIPModel, AutoencoderKLWan3_8, - VaceWanTransformer3DModel, WanT5EncoderModel) +from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, + CLIPModel, VaceWanTransformer3DModel, + WanT5EncoderModel) from videox_fun.pipeline import Wan2_2VaceFunPipeline from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import (get_image_to_video_latent, +from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, + get_image_to_video_latent, get_video_to_video_latent, save_videos_grid) @@ -111,107 +113,108 @@ check_min_version("0.18.0.dev0") logger = get_logger(__name__, log_level="INFO") -def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer3d, args, config, accelerator, weight_dtype, global_step): +def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, accelerator, weight_dtype, global_step): try: - logger.info("Running validation... ") - - if args.boundary_type == "full": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - transformer3d_2_val = None - else: - if args.boundary_type == "low": - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') - - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + is_deepspeed = type(transformer3d).__name__ == 'DeepSpeedEngine' + if is_deepspeed: + origin_config = transformer3d.config + transformer3d.config = accelerator.unwrap_model(transformer3d).config + with torch.no_grad(), torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): + logger.info("Running validation... ") + scheduler = FlowMatchEulerDiscreteScheduler( + **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) + ) + if args.boundary_type == "full": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + transformer3d_2 = None else: - sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + if args.boundary_type == "low": + transformer3d_1 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') + transformer3d_2 = VaceWanTransformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) + + else: + sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer') + transformer3d_1 = VaceWanTransformer3DModel.from_pretrained( + os.path.join(args.pretrained_model_name_or_path, sub_path), + transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), + ).to(weight_dtype) - transformer3d_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) + transformer3d_2 = accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d + + pipeline = Wan2_2VaceFunPipeline( + vae=vae, + text_encoder=text_encoder, + tokenizer=tokenizer, + transformer=transformer3d_1, + transformer_2=transformer3d_2, + scheduler=scheduler, + ) + pipeline = pipeline.to(accelerator.device) - sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer') - transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained( - os.path.join(args.pretrained_model_name_or_path, sub_path), - transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), - ).to(weight_dtype) - transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict()) - - scheduler = FlowMatchEulerDiscreteScheduler( - **filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs'])) - ) + if args.seed is None: + generator = None + else: + rank_seed = args.seed + accelerator.process_index + generator = torch.Generator(device=accelerator.device).manual_seed(rank_seed) + logger.info(f"Rank {accelerator.process_index} using seed: {rank_seed}") - pipeline = Wan2_2VaceFunPipeline( - vae=accelerator.unwrap_model(vae).to(weight_dtype), - text_encoder=accelerator.unwrap_model(text_encoder), - tokenizer=tokenizer, - transformer=transformer3d_val, - transformer_2=transformer3d_2_val, - scheduler=scheduler, - clip_image_encoder=clip_image_encoder, - ) - pipeline = pipeline.to(accelerator.device) + for i in range(len(args.validation_prompts)): + import cv2 + cap = cv2.VideoCapture(args.validation_paths[i]) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + cap.release() - if args.seed is None: - generator = None - else: - generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) + width, height = calculate_dimensions(args.image_sample_size * args.image_sample_size, width / height) + + video_length = int((args.video_sample_n_frames - 1) // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 + inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[height, width]) + control_video, _, _, _ = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[height, width]) + sample = pipeline( + args.validation_prompts[i], + num_frames = video_length, + negative_prompt = "bad detailed", + height = height, + width = width, + generator = generator, - images = [] - for i in range(len(args.validation_prompts)): - with torch.no_grad(): - with torch.autocast("cuda", dtype=weight_dtype): - video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1 - inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - control_video, _, _, _ = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size]) - sample = pipeline( - args.validation_prompts[i], - num_frames = video_length, - negative_prompt = "bad detailed", - height = args.video_sample_size, - width = args.video_sample_size, - generator = generator, + video = inpaint_video, + mask_video = inpaint_video_mask, + control_video = control_video, + subject_ref_images = None, + vace_context_scale = 1, + ).videos + os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) + save_videos_grid( + sample, + os.path.join( + args.output_dir, + f"sample/sample-{global_step}-rank{accelerator.process_index}-image-{i}.gif" + ) + ) - video = inpaint_video, - mask_video = inpaint_video_mask, - control_video = control_video, - subject_ref_images = None, - vace_context_scale = 1, - ).videos - os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True) - save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif")) - - del pipeline - del transformer3d_val - gc.collect() - torch.cuda.empty_cache() - torch.cuda.ipc_collect() - - return images + del pipeline + gc.collect() + torch.cuda.empty_cache() + torch.cuda.ipc_collect() + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if is_deepspeed: + transformer3d.config = origin_config except Exception as e: gc.collect() torch.cuda.empty_cache() torch.cuda.ipc_collect() - print(f"Eval error with info {e}") - return None + print(f"Eval error on rank {accelerator.process_index} with info {e}") + vae.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) + if not args.enable_text_encoder_in_dataloader: + text_encoder.to(accelerator.device if not args.low_vram else "cpu", dtype=weight_dtype) def parse_args(): parser = argparse.ArgumentParser(description="Simple example of a training script.") @@ -906,7 +909,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -926,26 +929,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1448,16 +1431,13 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) text_encoder = shard_fn(text_encoder) - # shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) - # transformer3d = shard_fn(transformer3d) - if args.use_ema: ema_transformer3d.to(accelerator.device) @@ -2070,27 +2050,25 @@ def main(): accelerator.save_state(save_path) logger.info(f"Saved state to {save_path}") - if accelerator.is_main_process: - if args.validation_prompts is not None and global_step % args.validation_steps == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and global_step % args.validation_steps == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -2098,27 +2076,25 @@ def main(): if global_step >= args.max_train_steps: break - if accelerator.is_main_process: - if args.validation_prompts is not None and epoch % args.validation_epochs == 0: - if args.use_ema: - # Store the UNet parameters temporarily and load the EMA parameters to perform inference. - ema_transformer3d.store(transformer3d.parameters()) - ema_transformer3d.copy_to(transformer3d.parameters()) - log_validation( - vae, - text_encoder, - tokenizer, - clip_image_encoder, - transformer3d, - args, - config, - accelerator, - weight_dtype, - global_step, - ) - if args.use_ema: - # Switch back to the original transformer3d parameters. - ema_transformer3d.restore(transformer3d.parameters()) + if args.validation_prompts is not None and epoch % args.validation_epochs == 0: + if args.use_ema: + # Store the UNet parameters temporarily and load the EMA parameters to perform inference. + ema_transformer3d.store(transformer3d.parameters()) + ema_transformer3d.copy_to(transformer3d.parameters()) + log_validation( + vae, + text_encoder, + tokenizer, + transformer3d, + args, + config, + accelerator, + weight_dtype, + global_step, + ) + if args.use_ema: + # Switch back to the original transformer3d parameters. + ema_transformer3d.restore(transformer3d.parameters()) # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() diff --git a/scripts/wan2.2_vace_fun/train.sh b/scripts/wan2.2_vace_fun/train.sh index 26c59c7..fd4bcae 100644 --- a/scripts/wan2.2_vace_fun/train.sh +++ b/scripts/wan2.2_vace_fun/train.sh @@ -26,7 +26,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_vace_fun/train.py \ --lr_scheduler="constant_with_warmup" \ --lr_warmup_steps=100 \ --seed=42 \ - --output_dir="output_dir" \ + --output_dir="output_dir_wan2.2_vace_fun" \ --gradient_checkpointing \ --mixed_precision="bf16" \ --adam_weight_decay=3e-2 \ diff --git a/scripts/z_image/train.py b/scripts/z_image/train.py index 9f89dce..adeefdf 100644 --- a/scripts/z_image/train.py +++ b/scripts/z_image/train.py @@ -203,7 +203,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -869,7 +869,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -889,26 +889,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): diff --git a/scripts/z_image/train_distill.py b/scripts/z_image/train_distill.py index 65b17ed..1efa10f 100644 --- a/scripts/z_image/train_distill.py +++ b/scripts/z_image/train_distill.py @@ -25,6 +25,7 @@ import pickle import random import shutil import sys +from functools import partial from typing import (Any, Callable, Dict, List, NamedTuple, Optional, Tuple, Union) @@ -55,7 +56,7 @@ from PIL import Image from torch.distributed.fsdp.fully_sharded_data_parallel import ( FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig, ShardedStateDictConfig) -from torch.utils.data import Dataset, RandomSampler, BatchSampler +from torch.utils.data import BatchSampler, Dataset, RandomSampler from torch.utils.tensorboard import SummaryWriter from torchvision import transforms from tqdm.auto import tqdm @@ -84,10 +85,12 @@ from videox_fun.models import (AutoencoderKL, AutoProcessor, AutoTokenizer, Qwen3ForCausalLM, QwenImageTransformer2DModel, ZImageTransformer2DModel) from videox_fun.pipeline import ZImagePipeline -from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, - get_image_to_video_latent, - save_videos_grid) +from videox_fun.utils import (DiscreteSampling, RectifiedFlow_TrigFlowWrapper, + calculate_dimensions, + convert_peft_lora_to_kohya_lora, create_network, + get_image_latent, get_image_to_video_latent, + merge_lora, sample_trigflow_timesteps, + save_videos_grid, unmerge_lora) if is_wandb_available(): import wandb @@ -208,7 +211,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -603,12 +606,6 @@ def parse_args(): default=[], help='Enter a list of trainable modules with lower learning rate' ) - parser.add_argument( - '--tokenizer_max_length', - type=int, - default=512, - help='Max length of tokenizer' - ) parser.add_argument( "--use_deepspeed", action="store_true", help="Whether or not to use deepspeed." ) @@ -618,47 +615,6 @@ def parse_args(): parser.add_argument( "--low_vram", action="store_true", help="Whether enable low_vram mode." ) - parser.add_argument( - "--prompt_template_encode", - type=str, - default="<|im_start|>system\nDescribe the image by detailing the color, shape, size, texture, quantity, text, spatial relationships of the objects and background:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n", - help=( - 'The prompt template for text encoder.' - ), - ) - parser.add_argument( - "--prompt_template_encode_start_idx", - type=int, - default=34, - help=( - 'The start idx for prompt template.' - ), - ) - parser.add_argument( - "--train_mode", - type=str, - default="normal", - help=( - 'The format of training data. Support `"normal"`' - ' (default), `"i2v"`.' - ), - ) - parser.add_argument( - "--abnormal_norm_clip_start", - type=int, - default=1000, - help=( - 'When do we start doing additional processing on abnormal gradients. ' - ), - ) - parser.add_argument( - "--initial_grad_norm_ratio", - type=int, - default=5, - help=( - 'The initial gradient is relative to the multiple of the max_grad_norm. ' - ), - ) parser.add_argument( "--weighting_scheme", type=str, @@ -678,11 +634,17 @@ def parse_args(): default=1.29, help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", ) + parser.add_argument( - "--guidance_scale", + "--use_trigflow", + action="store_true", + help="whether to use trigflow in training.", + ) + parser.add_argument( + "--sigma_max", type=float, - default=3.5, - help="the FLUX.1 dev variant is a guidance distilled model", + default=80.0, + help="The max value of sigma in trigflow.", ) parser.add_argument( "--gen_update_interval", @@ -834,6 +796,7 @@ def main(): weight_dtype = torch.bfloat16 args.mixed_precision = accelerator.mixed_precision + args.denoising_step_indices_list = [int(i) for i in args.denoising_step_indices_list] # Load scheduler, tokenizer and models. noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( args.pretrained_model_name_or_path, @@ -956,7 +919,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -976,26 +939,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1252,16 +1195,12 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - if fsdp_stage != 0: - from functools import partial - + if fsdp_stage != 0 or zero_stage != 0: from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(real_score_transformer3d.layers)) real_score_transformer3d = shard_fn(real_score_transformer3d) - if fsdp_stage != 0: - from functools import partial - + if fsdp_stage != 0 or zero_stage != 0: from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(text_encoder.model.layers)) text_encoder = shard_fn(text_encoder) @@ -1361,9 +1300,170 @@ def main(): vae_stream_1 = None vae_stream_2 = None - # Calculate the index we need】 - idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) + # RectifiedFlow Mode denoising_step_list = noise_scheduler.timesteps[args.train_sampling_steps - torch.tensor(args.denoising_step_indices_list)] + idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) + + # TrigFlow Mode + scaling = RectifiedFlow_TrigFlowWrapper(1, args.train_sampling_steps) + sample_trigflow_timesteps_D = partial( + sample_trigflow_timesteps, + P_mean=0.0, + P_std=1.6 + ) + + def denoise(model, xt, timestep, prompt_embeds, noise_scheduler=None, trigflow_scaling=None, multiply_c_in=True): + """ + Unified denoise function supporting both TrigFlow and Rectified Flow + + Args: + model: Diffusion model + xt: Noised input (B, C, T, H, W) or (B, C, H, W) + timestep: Timesteps (B,) or (B, 1) + prompt_embeds: Text condition embeddings + noise_scheduler: Noise scheduler (required for Rectified Flow) + trigflow_scaling: TrigFlow scaling function (required for TrigFlow) + multiply_c_in: Whether to multiply c_in with input (TrigFlow only) + + Returns: + x0_pred: Predicted clean data + flow_pred: Predicted velocity/flow field + """ + use_trigflow = getattr(args, 'use_trigflow', False) + original_dtype = xt.dtype + device = xt.device + + if use_trigflow: + # TrigFlow path + if trigflow_scaling is None: + raise ValueError("trigflow_scaling is required when using trigflow") + + ndim = xt.ndim + trigflow_t = timestep + + if trigflow_t.ndim == 1: + trigflow_t = trigflow_t.view(-1, 1) + + if ndim == 4: + trigflow_t_expanded = trigflow_t.view(-1, 1, 1, 1) + elif ndim == 5: + trigflow_t_expanded = trigflow_t.view(-1, 1, 1, 1, 1) + else: + raise ValueError(f"Expected 4D or 5D input, got {ndim}D tensor.") + + # Get TrigFlow preconditioning coefficients + c_skip, c_out, c_in, c_noise = trigflow_scaling(trigflow_t_expanded) + + # Precondition input + if multiply_c_in: + model_input = (xt * c_in).to(xt.dtype) + else: + model_input = xt.to(xt.dtype) + + timestep_normalized = c_noise.squeeze(1).squeeze(1).squeeze(1).squeeze(1) + + # Model inference + model_output = model( + x=model_input, + cap_feats=prompt_embeds, + t=(1000 - timestep_normalized) / 1000, + )[0] + + flow_pred = -model_output.double() + + # EDM-style x0 reconstruction + x0_pred = c_skip * xt + c_out * flow_pred + + else: + # Rectified Flow path + if noise_scheduler is None: + raise ValueError("scheduler is required for Rectified Flow") + + xt_double = xt.double() + timestep = timestep.to(device).double() + + timesteps = noise_scheduler.timesteps.to(device).double() + sigmas = noise_scheduler.sigmas.to(device).double() + timestep_id = torch.argmin((timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) + sigma_t = sigmas[timestep_id] + + ndim = xt.ndim + if ndim == 4: + sigma_t_expanded = sigma_t.view(-1, 1, 1, 1) + elif ndim == 5: + sigma_t_expanded = sigma_t.view(-1, 1, 1, 1, 1) + else: + raise ValueError(f"Expected 4D or 5D input, got {ndim}D tensor.") + + model_output = model( + x=xt, + cap_feats=prompt_embeds, + t=(1000 - timestep) / 1000, + )[0] + + flow_pred = -model_output.double() + + x0_pred = xt_double - sigma_t_expanded * flow_pred + + return x0_pred.to(original_dtype), flow_pred + + def add_noise(x0, noise, timesteps, noise_scheduler=None): + """ + Unified add noise function supporting both TrigFlow and Rectified Flow + + Args: + x0: Clean data + noise: Gaussian noise + timesteps: Timesteps + + Returns: + xt: Noised data + """ + use_trigflow = getattr(args, 'use_trigflow', False) + + if use_trigflow: + # TrigFlow path: xt = cos(t) * x0 + sin(t) * noise + trigflow_t = timesteps + ndim = x0.ndim + + if trigflow_t.ndim == 1: + trigflow_t = trigflow_t.view(-1, 1) + + if ndim == 4: + trigflow_t_expanded = trigflow_t.view(-1, 1, 1, 1) + elif ndim == 5: + trigflow_t_expanded = trigflow_t.view(-1, 1, 1, 1, 1) + else: + raise ValueError(f"Expected 4D or 5D input, got {ndim}D tensor.") + + cos_t = torch.cos(trigflow_t_expanded) + sin_t = torch.sin(trigflow_t_expanded) + + return cos_t * x0 + sin_t * noise + + else: + # Rectified Flow path: xt = (1 - sigma) * x0 + sigma * noise + if noise_scheduler is None: + raise ValueError("noise_scheduler are required for Rectified Flow") + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + + step_indices = [ + torch.argmin(torch.abs(schedule_timesteps - t)).item() + for t in timesteps + ] + step_indices = torch.tensor(step_indices, device=accelerator.device) + sigma = sigmas[step_indices].flatten() + + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + sigmas = get_sigmas(timesteps, n_dim=x0.ndim, dtype=x0.dtype) + return (1.0 - sigmas) * x0 + sigmas * noise for epoch in range(first_epoch, args.num_train_epochs): train_dmd_loss = 0.0 @@ -1420,12 +1520,14 @@ def main(): else: with torch.no_grad(): prompt_embeds = encode_prompt( - batch['text'], device=accelerator.device, + batch['text'], + device=accelerator.device, text_encoder=text_encoder, tokenizer=tokenizer, ) neg_prompt_embeds = encode_prompt( - ["低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感。构图混乱。文字模糊,扭曲。"], device=accelerator.device, + ["亮度过高,过曝,严重的色彩失真,低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感。构图混乱。文字模糊,扭曲。"], + device=accelerator.device, text_encoder=text_encoder, tokenizer=tokenizer, ) @@ -1436,148 +1538,105 @@ def main(): if args.low_vram: real_score_transformer3d = real_score_transformer3d.to(accelerator.device) + if getattr(args, 'use_trigflow', False): + # Create discrete denoising steps + t_max = torch.arctan(torch.tensor(args.sigma_max)) + denoising_step_list = torch.linspace(t_max.item(), 0.0, args.train_sampling_steps) + denoising_step_list = denoising_step_list[args.train_sampling_steps - torch.tensor(args.denoising_step_indices_list)] + else: + image_seq_len = int(target_shape[-1] // 2 * target_shape[-2] // 2) + mu = calculate_shift( + image_seq_len, + noise_scheduler.config.get("base_image_seq_len", 256), + noise_scheduler.config.get("max_image_seq_len", 4096), + noise_scheduler.config.get("base_shift", 0.5), + noise_scheduler.config.get("max_shift", 1.15), + ) + noise_scheduler.sigma_min = 0.0 + noise_scheduler.set_timesteps(args.train_sampling_steps, device=accelerator.device, mu=mu) + denoising_step_list = noise_scheduler.timesteps[args.train_sampling_steps - torch.tensor(args.denoising_step_indices_list)] + + # ==================== Generator Update (DMD) ==================== with accelerator.accumulate(generator_transformer3d): - def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): - sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) - schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) - timesteps = timesteps.to(accelerator.device) - - step_indices = [ - torch.argmin(torch.abs(schedule_timesteps - t)).item() - for t in timesteps - ] - step_indices = torch.tensor(step_indices, device=accelerator.device) - sigma = sigmas[step_indices].flatten() - - while len(sigma.shape) < n_dim: - sigma = sigma.unsqueeze(-1) - return sigma - - def add_noise(latents, noise, timesteps): - sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) - return (1.0 - sigmas) * latents + sigmas * noise - def generate_and_sync_list(num_denoising_steps, device): indices = torch.randint(low=0, high=num_denoising_steps, size=(1,), generator=torch_rng, device=device) if dist.is_initialized(): dist.broadcast(indices, src=0) return indices.tolist() - def convert_flow_pred_to_x0( - scheduler, - flow_pred: torch.Tensor, - xt: torch.Tensor, - timestep: torch.Tensor - ) -> torch.Tensor: - """ - Convert flow matching's prediction to x0 prediction. - Supports both 4D [B, C, H, W] and 5D [B, C, F, H, W] inputs. - """ - original_dtype = flow_pred.dtype - device = flow_pred.device - - flow_pred = flow_pred.double() - xt = xt.double() - timesteps = scheduler.timesteps.to(device).double() - sigmas = scheduler.sigmas.to(device).double() - timestep = timestep.to(device).double() - - timestep_id = torch.argmin((timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) - sigma_t = sigmas[timestep_id] - - ndim = flow_pred.ndim - if ndim == 4: - sigma_t = sigma_t.view(-1, 1, 1, 1) - elif ndim == 5: - sigma_t = sigma_t.view(-1, 1, 1, 1, 1) - else: - raise ValueError(f"Expected 4D or 5D input, got {ndim}D tensor.") - - x0_pred = xt - sigma_t * flow_pred - return x0_pred.to(original_dtype) - - # --- Main Training Logic --- bsz, channel, num_frames, height, width = target_shape + if step % args.gen_update_interval == 0: generator_noise = torch.randn(target_shape, device=accelerator.device, generator=torch_rng, dtype=weight_dtype) num_denoising_steps = len(denoising_step_list) final_step_index = generate_and_sync_list(num_denoising_steps, device=generator_noise.device)[0] - - # Precompute seq_len once (same for all steps) - - for index, current_timestep in enumerate(denoising_step_list): + + # Multi-step denoising (backward simulation) + for index in range(num_denoising_steps): is_final_step = (index == final_step_index) - timestep = torch.full( - generator_noise.shape[:1], - current_timestep, - device=generator_noise.device, - dtype=torch.int64 - ) + current_t = denoising_step_list[index].expand(bsz).to(accelerator.device) with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): context_manager = torch.no_grad() if not is_final_step else contextlib.nullcontext() - + with context_manager: - generator_pred = generator_transformer3d( - x=generator_noise, - cap_feats=prompt_embeds, - t=(1000 - timestep) / 1000, - )[0] - generator_pred = -generator_pred - generator_pred = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=generator_pred, + generator_pred, _ = denoise( + model=generator_transformer3d, xt=generator_noise, - timestep=timestep + timestep=current_t, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling, + multiply_c_in=False if index == 0 else True, ) if is_final_step: break - next_timestep = denoising_step_list[index + 1] * torch.ones( - generator_noise.shape[:1], dtype=torch.long, device=generator_noise.device - ) - generator_noise = add_noise( - generator_pred, - torch.randn(generator_pred.shape, dtype=generator_pred.dtype, device=generator_pred.device, generator=torch_rng), - next_timestep - ) + # Add noise for next step + if index < num_denoising_steps - 1: + next_t = denoising_step_list[index + 1].expand(bsz).to(accelerator.device) + generator_noise = add_noise( + generator_pred, + torch.randn(generator_pred.shape, dtype=generator_pred.dtype, device=generator_pred.device, generator=torch_rng), + next_t, + noise_scheduler=noise_scheduler + ) + + if getattr(args, 'use_trigflow', False): + # Sample timesteps for discriminator (D distribution) + generator_timestep = sample_trigflow_timesteps_D(bsz, device=accelerator.device) + else: + indices = idx_sampling(bsz, generator=torch_rng, device=accelerator.device).long().cpu() + generator_timestep = noise_scheduler.timesteps[indices].to(device=accelerator.device) - indices = idx_sampling(bsz, generator=torch_rng, device=accelerator.device).long().cpu() - generator_timestep = noise_scheduler.timesteps[indices].to(device=accelerator.device) + # Add noise to generated samples generator_denoised_input = add_noise( generator_pred, torch.randn(generator_pred.shape, dtype=generator_pred.dtype, device=generator_pred.device, generator=torch_rng), - generator_timestep + generator_timestep, + noise_scheduler=noise_scheduler ).detach().to(accelerator.device, dtype=weight_dtype) # Compute fake score with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device), torch.no_grad(): - fake_score_main_cond = fake_score_transformer3d( - x=generator_denoised_input, - cap_feats=prompt_embeds, - t=(1000 - generator_timestep) / 1000 - )[0] - fake_score_main_cond = -fake_score_main_cond - fake_score_main_cond = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=fake_score_main_cond, + fake_score_main_cond, _ = denoise( + model=fake_score_transformer3d, xt=generator_denoised_input, - timestep=generator_timestep + timestep=generator_timestep, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling ) if args.fake_guidance_scale != 0.0: - fake_score_main_uncond = fake_score_transformer3d( - x=generator_denoised_input, - cap_feats=neg_prompt_embeds, - t=(1000 - generator_timestep) / 1000 - )[0] - fake_score_main_uncond = -fake_score_main_uncond - fake_score_main_uncond = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=fake_score_main_uncond, + fake_score_main_uncond, _ = denoise( + model=fake_score_transformer3d, xt=generator_denoised_input, - timestep=generator_timestep + timestep=generator_timestep, + prompt_embeds=neg_prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling ) fake_score_main = fake_score_main_uncond + ( fake_score_main_cond - fake_score_main_uncond @@ -1585,32 +1644,24 @@ def main(): else: fake_score_main = fake_score_main_cond - # Compute real score - real_score_main_cond = real_score_transformer3d( - x=generator_denoised_input, - cap_feats=prompt_embeds, - t=(1000 - generator_timestep) / 1000 - )[0] - real_score_main_cond = -real_score_main_cond - real_score_main_cond = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=real_score_main_cond, + # Compute real score (teacher) + real_score_main_cond, _ = denoise( + model=real_score_transformer3d, xt=generator_denoised_input, - timestep=generator_timestep + timestep=generator_timestep, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling ) if args.real_guidance_scale != 0.0: - real_score_main_uncond = real_score_transformer3d( - x=generator_denoised_input, - cap_feats=neg_prompt_embeds, - t=(1000 - generator_timestep) / 1000 - )[0] - real_score_main_uncond = -real_score_main_uncond - real_score_main_uncond = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=real_score_main_uncond, + real_score_main_uncond, _ = denoise( + model=real_score_transformer3d, xt=generator_denoised_input, - timestep=generator_timestep + timestep=generator_timestep, + prompt_embeds=neg_prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling ) real_score_main = real_score_main_uncond + ( @@ -1622,7 +1673,7 @@ def main(): # DMD loss fake_to_real_grad = fake_score_main - real_score_main generator_to_real_norm = generator_pred - real_score_main - normalizer = torch.abs(generator_to_real_norm).mean(dim=[1, 2, 3, 4], keepdim=True) + normalizer = torch.abs(generator_to_real_norm).mean(dim=[1, 2, 3, 4], keepdim=True).clip(min=1e-5) fake_to_real_grad = fake_to_real_grad / normalizer fake_to_real_grad = torch.nan_to_num(fake_to_real_grad) @@ -1657,61 +1708,58 @@ def main(): num_denoising_steps = len(denoising_step_list) final_step_index = generate_and_sync_list(num_denoising_steps, device=fake_score_critic_noise.device)[0] - for index, current_timestep in enumerate(denoising_step_list): + for index in range(num_denoising_steps): is_final_step = (index == final_step_index) - timestep = torch.full( - fake_score_critic_noise.shape[:1], - current_timestep, - device=fake_score_critic_noise.device, - dtype=torch.int64 - ) - + current_t = denoising_step_list[index].expand(bsz).to(accelerator.device) + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): - fake_score_denoised_pred = generator_transformer3d( - x=fake_score_critic_noise, - cap_feats=prompt_embeds, - t=(1000 - timestep) / 1000 - )[0] - fake_score_denoised_pred = -fake_score_denoised_pred - fake_score_denoised_pred = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=fake_score_denoised_pred, + fake_score_denoised_pred, _ = denoise( + model=generator_transformer3d, xt=fake_score_critic_noise, - timestep=timestep + timestep=current_t, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling, + multiply_c_in=False if index == 0 else True, ) if is_final_step: break - - next_timestep = denoising_step_list[index + 1] * torch.ones( - fake_score_critic_noise.shape[:1], - dtype=torch.long, - device=fake_score_critic_noise.device - ) - - fake_score_critic_noise = add_noise( - fake_score_denoised_pred, - torch.randn(fake_score_denoised_pred.shape, dtype=fake_score_denoised_pred.dtype, device=fake_score_denoised_pred.device, generator=torch_rng), - next_timestep - ) - indices = idx_sampling(bsz, generator=torch_rng, device=accelerator.device).long().cpu() - critic_timestep = noise_scheduler.timesteps[indices].to(device=accelerator.device) + if index < num_denoising_steps - 1: + next_t = denoising_step_list[index + 1].expand(bsz).to(accelerator.device) + fake_score_critic_noise = add_noise( + fake_score_denoised_pred, + torch.randn(fake_score_denoised_pred.shape, dtype=fake_score_denoised_pred.dtype, device=fake_score_denoised_pred.device, generator=torch_rng), + next_t, + noise_scheduler=noise_scheduler + ) + + # Sample timesteps for critic + if getattr(args, 'use_trigflow', False): + # Sample timesteps for discriminator (D distribution) + critic_timestep = sample_trigflow_timesteps_D(bsz, device=accelerator.device) + else: + indices = idx_sampling(bsz, generator=torch_rng, device=accelerator.device).long().cpu() + critic_timestep = noise_scheduler.timesteps[indices].to(device=accelerator.device) critic_noise = torch.randn(fake_score_denoised_pred.shape, dtype=fake_score_denoised_pred.dtype, device=fake_score_denoised_pred.device, generator=torch_rng) fake_score_denoised_input = add_noise( fake_score_denoised_pred, critic_noise, - critic_timestep + critic_timestep, + noise_scheduler=noise_scheduler ) with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): - fake_score_denoised_output = fake_score_transformer3d( - x=fake_score_denoised_input, - cap_feats=prompt_embeds, - t=(1000 - critic_timestep) / 1000 - )[0] - fake_score_denoised_output = -fake_score_denoised_output + fake_score_pred, _ = denoise( + model=fake_score_transformer3d, + xt=fake_score_denoised_input, + timestep=critic_timestep, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling + ) def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): noise_pred = noise_pred.float() @@ -1725,10 +1773,37 @@ def main(): final_loss = masked_loss.mean() return final_loss - denoising_loss = custom_mse_loss(fake_score_denoised_output, critic_noise - fake_score_denoised_pred) - avg_denoising_loss = accelerator.gather(denoising_loss.repeat(args.train_batch_size)).mean() - train_denoising_loss += avg_denoising_loss.item() / args.gradient_accumulation_steps + # Compute weighting based on sin(t) (following rCM) + if getattr(args, 'use_trigflow', False): + ndim = fake_score_denoised_input.ndim + if critic_timestep.ndim == 1: + critic_timestep_view = critic_timestep.view(-1, 1) + else: + critic_timestep_view = critic_timestep + + if ndim == 4: + critic_t_expanded = critic_timestep_view.view(-1, 1, 1, 1) + elif ndim == 5: + critic_t_expanded = critic_timestep_view.view(-1, 1, 1, 1, 1) + + sin_t = torch.sin(critic_t_expanded) + weighting = 1.0 / (sin_t ** 2 + 1e-8) + else: + weighting = None + denoising_loss = custom_mse_loss( + fake_score_pred, + fake_score_denoised_pred, + weighting=weighting + ) + + avg_denoising_loss = accelerator_fake_score_transformer3d.gather(denoising_loss.repeat(args.train_batch_size)).mean() + train_denoising_loss += avg_denoising_loss.item() / args.gradient_accumulation_steps + + if args.low_vram: + generator_transformer3d = generator_transformer3d.to("cpu") + torch.cuda.empty_cache() + accelerator_fake_score_transformer3d.backward(denoising_loss) if accelerator_fake_score_transformer3d.sync_gradients: accelerator_fake_score_transformer3d.clip_grad_norm_(fake_trainable_params, args.max_grad_norm) diff --git a/scripts/z_image/train_distill_lora.py b/scripts/z_image/train_distill_lora.py index 5ffd384..acb2e9e 100644 --- a/scripts/z_image/train_distill_lora.py +++ b/scripts/z_image/train_distill_lora.py @@ -25,6 +25,7 @@ import pickle import random import shutil import sys +from functools import partial from typing import (Any, Callable, Dict, List, NamedTuple, Optional, Tuple, Union) @@ -55,7 +56,7 @@ from PIL import Image from torch.distributed.fsdp.fully_sharded_data_parallel import ( FullOptimStateDictConfig, FullStateDictConfig, ShardedOptimStateDictConfig, ShardedStateDictConfig) -from torch.utils.data import Dataset, RandomSampler, BatchSampler +from torch.utils.data import BatchSampler, Dataset, RandomSampler from torch.utils.tensorboard import SummaryWriter from torchvision import transforms from tqdm.auto import tqdm @@ -84,13 +85,12 @@ from videox_fun.models import (AutoencoderKL, AutoProcessor, AutoTokenizer, Qwen3ForCausalLM, QwenImageTransformer2DModel, ZImageTransformer2DModel) from videox_fun.pipeline import ZImagePipeline -from videox_fun.utils.discrete_sampler import DiscreteSampling -from videox_fun.utils.lora_utils import (convert_peft_lora_to_kohya_lora, - create_network, merge_lora, - unmerge_lora) -from videox_fun.utils.utils import (calculate_dimensions, get_image_latent, - get_image_to_video_latent, - save_videos_grid) +from videox_fun.utils import (DiscreteSampling, RectifiedFlow_TrigFlowWrapper, + calculate_dimensions, + convert_peft_lora_to_kohya_lora, create_network, + get_image_latent, get_image_to_video_latent, + merge_lora, sample_trigflow_timesteps, + save_videos_grid, unmerge_lora) if is_wandb_available(): import wandb @@ -102,36 +102,6 @@ def filter_kwargs(cls, kwargs): filtered_kwargs = {k: v for k, v in kwargs.items() if k in valid_params} return filtered_kwargs -def linear_decay(initial_value, final_value, total_steps, current_step): - if current_step >= total_steps: - return final_value - current_step = max(0, current_step) - step_size = (final_value - initial_value) / total_steps - current_value = initial_value + step_size * current_step - return current_value - -def generate_timestep_with_lognorm(low, high, shape, device="cpu", generator=None): - u = torch.normal(mean=0.0, std=1.0, size=shape, device=device, generator=generator) - t = 1 / (1 + torch.exp(-u)) * (high - low) + low - return torch.clip(t.to(torch.int32), low, high - 1) - -def compute_empirical_mu(image_seq_len: int, num_steps: int) -> float: - a1, b1 = 8.73809524e-05, 1.89833333 - a2, b2 = 0.00016927, 0.45666666 - - if image_seq_len > 4300: - mu = a2 * image_seq_len + b2 - return float(mu) - - m_200 = a2 * image_seq_len + b2 - m_10 = a1 * image_seq_len + b1 - - a = (m_200 - m_10) / 190.0 - b = m_200 - 200.0 * a - mu = a * num_steps + b - - return float(mu) - def calculate_shift( image_seq_len, base_seq_len: int = 256, @@ -211,7 +181,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -659,6 +629,17 @@ def parse_args(): help=("The module is trained in loras. "), ) + parser.add_argument( + "--use_trigflow", + action="store_true", + help="whether to use trigflow in training.", + ) + parser.add_argument( + "--sigma_max", + type=float, + default=80.0, + help="The max value of sigma in trigflow.", + ) parser.add_argument( "--gen_update_interval", type=int, @@ -810,6 +791,7 @@ def main(): args.mixed_precision = accelerator.mixed_precision # Load scheduler, tokenizer and models. + args.denoising_step_indices_list = [int(i) for i in args.denoising_step_indices_list] noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( args.pretrained_model_name_or_path, subfolder="scheduler" @@ -948,7 +930,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -977,15 +959,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -998,34 +972,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1228,7 +1180,7 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - elif fsdp_stage != 0: + else: generator_transformer3d.network = network generator_transformer3d = generator_transformer3d.to(dtype=weight_dtype) generator_transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( @@ -1239,24 +1191,13 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler = accelerator_fake_score_transformer3d.prepare( fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - fake_score_network, critic_optimizer, fake_score_lr_scheduler= accelerator_fake_score_transformer3d.prepare( - fake_score_network, critic_optimizer, fake_score_lr_scheduler - ) - - if fsdp_stage != 0: - from functools import partial + if fsdp_stage != 0 or zero_stage != 0: from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(real_score_transformer3d.layers)) real_score_transformer3d = shard_fn(real_score_transformer3d) - if fsdp_stage != 0: - from functools import partial - + if fsdp_stage != 0 or zero_stage != 0: from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(text_encoder.model.layers)) text_encoder = shard_fn(text_encoder) @@ -1368,10 +1309,170 @@ def main(): vae_stream_1 = None vae_stream_2 = None - # Calculate the index we need】 - idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) - args.denoising_step_indices_list = [int(tmp) for tmp in args.denoising_step_indices_list] + # RectifiedFlow Mode denoising_step_list = noise_scheduler.timesteps[args.train_sampling_steps - torch.tensor(args.denoising_step_indices_list)] + idx_sampling = DiscreteSampling(args.train_sampling_steps, uniform_sampling=args.uniform_sampling) + + # TrigFlow Mode + scaling = RectifiedFlow_TrigFlowWrapper(1, args.train_sampling_steps) + sample_trigflow_timesteps_D = partial( + sample_trigflow_timesteps, + P_mean=0.0, + P_std=1.6 + ) + + def denoise(model, xt, timestep, prompt_embeds, noise_scheduler=None, trigflow_scaling=None, multiply_c_in=True): + """ + Unified denoise function supporting both TrigFlow and Rectified Flow + + Args: + model: Diffusion model + xt: Noised input (B, C, T, H, W) or (B, C, H, W) + timestep: Timesteps (B,) or (B, 1) + prompt_embeds: Text condition embeddings + noise_scheduler: Noise scheduler (required for Rectified Flow) + trigflow_scaling: TrigFlow scaling function (required for TrigFlow) + multiply_c_in: Whether to multiply c_in with input (TrigFlow only) + + Returns: + x0_pred: Predicted clean data + flow_pred: Predicted velocity/flow field + """ + use_trigflow = getattr(args, 'use_trigflow', False) + original_dtype = xt.dtype + device = xt.device + + if use_trigflow: + # TrigFlow path + if trigflow_scaling is None: + raise ValueError("trigflow_scaling is required when using trigflow") + + ndim = xt.ndim + trigflow_t = timestep + + if trigflow_t.ndim == 1: + trigflow_t = trigflow_t.view(-1, 1) + + if ndim == 4: + trigflow_t_expanded = trigflow_t.view(-1, 1, 1, 1) + elif ndim == 5: + trigflow_t_expanded = trigflow_t.view(-1, 1, 1, 1, 1) + else: + raise ValueError(f"Expected 4D or 5D input, got {ndim}D tensor.") + + # Get TrigFlow preconditioning coefficients + c_skip, c_out, c_in, c_noise = trigflow_scaling(trigflow_t_expanded) + + # Precondition input + if multiply_c_in: + model_input = (xt * c_in).to(xt.dtype) + else: + model_input = xt.to(xt.dtype) + + timestep_normalized = c_noise.squeeze(1).squeeze(1).squeeze(1).squeeze(1) + + # Model inference + model_output = model( + x=model_input, + cap_feats=prompt_embeds, + t=(1000 - timestep_normalized) / 1000, + )[0] + + flow_pred = -model_output.double() + + # EDM-style x0 reconstruction + x0_pred = c_skip * xt + c_out * flow_pred + + else: + # Rectified Flow path + if noise_scheduler is None: + raise ValueError("scheduler is required for Rectified Flow") + + xt_double = xt.double() + timestep = timestep.to(device).double() + + timesteps = noise_scheduler.timesteps.to(device).double() + sigmas = noise_scheduler.sigmas.to(device).double() + timestep_id = torch.argmin((timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) + sigma_t = sigmas[timestep_id] + + ndim = xt.ndim + if ndim == 4: + sigma_t_expanded = sigma_t.view(-1, 1, 1, 1) + elif ndim == 5: + sigma_t_expanded = sigma_t.view(-1, 1, 1, 1, 1) + else: + raise ValueError(f"Expected 4D or 5D input, got {ndim}D tensor.") + + model_output = model( + x=xt, + cap_feats=prompt_embeds, + t=(1000 - timestep) / 1000, + )[0] + + flow_pred = -model_output.double() + + x0_pred = xt_double - sigma_t_expanded * flow_pred + + return x0_pred.to(original_dtype), flow_pred + + def add_noise(x0, noise, timesteps, noise_scheduler=None): + """ + Unified add noise function supporting both TrigFlow and Rectified Flow + + Args: + x0: Clean data + noise: Gaussian noise + timesteps: Timesteps + + Returns: + xt: Noised data + """ + use_trigflow = getattr(args, 'use_trigflow', False) + + if use_trigflow: + # TrigFlow path: xt = cos(t) * x0 + sin(t) * noise + trigflow_t = timesteps + ndim = x0.ndim + + if trigflow_t.ndim == 1: + trigflow_t = trigflow_t.view(-1, 1) + + if ndim == 4: + trigflow_t_expanded = trigflow_t.view(-1, 1, 1, 1) + elif ndim == 5: + trigflow_t_expanded = trigflow_t.view(-1, 1, 1, 1, 1) + else: + raise ValueError(f"Expected 4D or 5D input, got {ndim}D tensor.") + + cos_t = torch.cos(trigflow_t_expanded) + sin_t = torch.sin(trigflow_t_expanded) + + return cos_t * x0 + sin_t * noise + + else: + # Rectified Flow path: xt = (1 - sigma) * x0 + sigma * noise + if noise_scheduler is None: + raise ValueError("noise_scheduler are required for Rectified Flow") + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + + step_indices = [ + torch.argmin(torch.abs(schedule_timesteps - t)).item() + for t in timesteps + ] + step_indices = torch.tensor(step_indices, device=accelerator.device) + sigma = sigmas[step_indices].flatten() + + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + sigmas = get_sigmas(timesteps, n_dim=x0.ndim, dtype=x0.dtype) + return (1.0 - sigmas) * x0 + sigmas * noise for epoch in range(first_epoch, args.num_train_epochs): train_dmd_loss = 0.0 @@ -1428,12 +1529,14 @@ def main(): else: with torch.no_grad(): prompt_embeds = encode_prompt( - batch['text'], device=accelerator.device, + batch['text'], + device=accelerator.device, text_encoder=text_encoder, tokenizer=tokenizer, ) neg_prompt_embeds = encode_prompt( - ["亮度过高,过曝,严重的色彩失真,低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感。构图混乱。文字模糊,扭曲。"], device=accelerator.device, + ["亮度过高,过曝,严重的色彩失真,低分辨率,低画质,肢体畸形,手指畸形,画面过饱和,蜡像感,人脸无细节,过度光滑,画面具有AI感。构图混乱。文字模糊,扭曲。"], + device=accelerator.device, text_encoder=text_encoder, tokenizer=tokenizer, ) @@ -1444,147 +1547,105 @@ def main(): if args.low_vram: real_score_transformer3d = real_score_transformer3d.to(accelerator.device) + if getattr(args, 'use_trigflow', False): + # Create discrete denoising steps + t_max = torch.arctan(torch.tensor(args.sigma_max)) + denoising_step_list = torch.linspace(t_max.item(), 0.0, args.train_sampling_steps) + denoising_step_list = denoising_step_list[args.train_sampling_steps - torch.tensor(args.denoising_step_indices_list)] + else: + image_seq_len = int(target_shape[-1] // 2 * target_shape[-2] // 2) + mu = calculate_shift( + image_seq_len, + noise_scheduler.config.get("base_image_seq_len", 256), + noise_scheduler.config.get("max_image_seq_len", 4096), + noise_scheduler.config.get("base_shift", 0.5), + noise_scheduler.config.get("max_shift", 1.15), + ) + noise_scheduler.sigma_min = 0.0 + noise_scheduler.set_timesteps(args.train_sampling_steps, device=accelerator.device, mu=mu) + denoising_step_list = noise_scheduler.timesteps[args.train_sampling_steps - torch.tensor(args.denoising_step_indices_list)] + + # ==================== Generator Update (DMD) ==================== with accelerator.accumulate(generator_transformer3d): - def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): - sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) - schedule_timesteps = noise_scheduler.timesteps.to(accelerator.device) - timesteps = timesteps.to(accelerator.device) - - step_indices = [ - torch.argmin(torch.abs(schedule_timesteps - t)).item() - for t in timesteps - ] - step_indices = torch.tensor(step_indices, device=accelerator.device) - sigma = sigmas[step_indices].flatten() - - while len(sigma.shape) < n_dim: - sigma = sigma.unsqueeze(-1) - return sigma - - def add_noise(latents, noise, timesteps): - sigmas = get_sigmas(timesteps, n_dim=latents.ndim, dtype=latents.dtype) - return (1.0 - sigmas) * latents + sigmas * noise - def generate_and_sync_list(num_denoising_steps, device): indices = torch.randint(low=0, high=num_denoising_steps, size=(1,), generator=torch_rng, device=device) if dist.is_initialized(): dist.broadcast(indices, src=0) return indices.tolist() - def convert_flow_pred_to_x0( - scheduler, - flow_pred: torch.Tensor, - xt: torch.Tensor, - timestep: torch.Tensor - ) -> torch.Tensor: - """ - Convert flow matching's prediction to x0 prediction. - Supports both 4D [B, C, H, W] and 5D [B, C, F, H, W] inputs. - """ - original_dtype = flow_pred.dtype - device = flow_pred.device - - flow_pred = flow_pred.double() - xt = xt.double() - timesteps = scheduler.timesteps.to(device).double() - sigmas = scheduler.sigmas.to(device).double() - timestep = timestep.to(device).double() - - timestep_id = torch.argmin((timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) - sigma_t = sigmas[timestep_id] - - ndim = flow_pred.ndim - if ndim == 4: - sigma_t = sigma_t.view(-1, 1, 1, 1) - elif ndim == 5: - sigma_t = sigma_t.view(-1, 1, 1, 1, 1) - else: - raise ValueError(f"Expected 4D or 5D input, got {ndim}D tensor.") - - x0_pred = xt - sigma_t * flow_pred - return x0_pred.to(original_dtype) - - # --- Main Training Logic --- bsz, channel, num_frames, height, width = target_shape + if step % args.gen_update_interval == 0: generator_noise = torch.randn(target_shape, device=accelerator.device, generator=torch_rng, dtype=weight_dtype) num_denoising_steps = len(denoising_step_list) final_step_index = generate_and_sync_list(num_denoising_steps, device=generator_noise.device)[0] - - # Precompute seq_len once (same for all steps) - for index, current_timestep in enumerate(denoising_step_list): + + # Multi-step denoising (backward simulation) + for index in range(num_denoising_steps): is_final_step = (index == final_step_index) - timestep = torch.full( - generator_noise.shape[:1], - current_timestep, - device=generator_noise.device, - dtype=torch.int64 - ) + current_t = denoising_step_list[index].expand(bsz).to(accelerator.device) with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): context_manager = torch.no_grad() if not is_final_step else contextlib.nullcontext() - + with context_manager: - generator_pred = generator_transformer3d( - x=generator_noise, - cap_feats=prompt_embeds, - t=(1000 - timestep) / 1000, - )[0] - generator_pred = -generator_pred - generator_pred = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=generator_pred, + generator_pred, _ = denoise( + model=generator_transformer3d, xt=generator_noise, - timestep=timestep + timestep=current_t, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling, + multiply_c_in=False if index == 0 else True, ) if is_final_step: break - next_timestep = denoising_step_list[index + 1] * torch.ones( - generator_noise.shape[:1], dtype=torch.long, device=generator_noise.device - ) - generator_noise = add_noise( - generator_pred, - torch.randn(generator_pred.shape, dtype=generator_pred.dtype, device=generator_pred.device, generator=torch_rng), - next_timestep - ) + # Add noise for next step + if index < num_denoising_steps - 1: + next_t = denoising_step_list[index + 1].expand(bsz).to(accelerator.device) + generator_noise = add_noise( + generator_pred, + torch.randn(generator_pred.shape, dtype=generator_pred.dtype, device=generator_pred.device, generator=torch_rng), + next_t, + noise_scheduler=noise_scheduler + ) + + if getattr(args, 'use_trigflow', False): + # Sample timesteps for discriminator (D distribution) + generator_timestep = sample_trigflow_timesteps_D(bsz, device=accelerator.device) + else: + indices = idx_sampling(bsz, generator=torch_rng, device=accelerator.device).long().cpu() + generator_timestep = noise_scheduler.timesteps[indices].to(device=accelerator.device) - indices = idx_sampling(bsz, generator=torch_rng, device=accelerator.device).long().cpu() - generator_timestep = noise_scheduler.timesteps[indices].to(device=accelerator.device) + # Add noise to generated samples generator_denoised_input = add_noise( generator_pred, torch.randn(generator_pred.shape, dtype=generator_pred.dtype, device=generator_pred.device, generator=torch_rng), - generator_timestep + generator_timestep, + noise_scheduler=noise_scheduler ).detach().to(accelerator.device, dtype=weight_dtype) # Compute fake score with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device), torch.no_grad(): - fake_score_main_cond = fake_score_transformer3d( - x=generator_denoised_input, - cap_feats=prompt_embeds, - t=(1000 - generator_timestep) / 1000 - )[0] - fake_score_main_cond = -fake_score_main_cond - fake_score_main_cond = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=fake_score_main_cond, + fake_score_main_cond, _ = denoise( + model=fake_score_transformer3d, xt=generator_denoised_input, - timestep=generator_timestep + timestep=generator_timestep, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling ) if args.fake_guidance_scale != 0.0: - fake_score_main_uncond = fake_score_transformer3d( - x=generator_denoised_input, - cap_feats=neg_prompt_embeds, - t=(1000 - generator_timestep) / 1000 - )[0] - fake_score_main_uncond = -fake_score_main_uncond - fake_score_main_uncond = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=fake_score_main_uncond, + fake_score_main_uncond, _ = denoise( + model=fake_score_transformer3d, xt=generator_denoised_input, - timestep=generator_timestep + timestep=generator_timestep, + prompt_embeds=neg_prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling ) fake_score_main = fake_score_main_uncond + ( fake_score_main_cond - fake_score_main_uncond @@ -1592,32 +1653,24 @@ def main(): else: fake_score_main = fake_score_main_cond - # Compute real score - real_score_main_cond = real_score_transformer3d( - x=generator_denoised_input, - cap_feats=prompt_embeds, - t=(1000 - generator_timestep) / 1000 - )[0] - real_score_main_cond = -real_score_main_cond - real_score_main_cond = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=real_score_main_cond, + # Compute real score (teacher) + real_score_main_cond, _ = denoise( + model=real_score_transformer3d, xt=generator_denoised_input, - timestep=generator_timestep + timestep=generator_timestep, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling ) if args.real_guidance_scale != 0.0: - real_score_main_uncond = real_score_transformer3d( - x=generator_denoised_input, - cap_feats=neg_prompt_embeds, - t=(1000 - generator_timestep) / 1000 - )[0] - real_score_main_uncond = -real_score_main_uncond - real_score_main_uncond = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=real_score_main_uncond, + real_score_main_uncond, _ = denoise( + model=real_score_transformer3d, xt=generator_denoised_input, - timestep=generator_timestep + timestep=generator_timestep, + prompt_embeds=neg_prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling ) real_score_main = real_score_main_uncond + ( @@ -1629,7 +1682,7 @@ def main(): # DMD loss fake_to_real_grad = fake_score_main - real_score_main generator_to_real_norm = generator_pred - real_score_main - normalizer = torch.abs(generator_to_real_norm).mean(dim=[1, 2, 3, 4], keepdim=True) + normalizer = torch.abs(generator_to_real_norm).mean(dim=[1, 2, 3, 4], keepdim=True).clip(min=1e-5) fake_to_real_grad = fake_to_real_grad / normalizer fake_to_real_grad = torch.nan_to_num(fake_to_real_grad) @@ -1664,61 +1717,58 @@ def main(): num_denoising_steps = len(denoising_step_list) final_step_index = generate_and_sync_list(num_denoising_steps, device=fake_score_critic_noise.device)[0] - for index, current_timestep in enumerate(denoising_step_list): + for index in range(num_denoising_steps): is_final_step = (index == final_step_index) - timestep = torch.full( - fake_score_critic_noise.shape[:1], - current_timestep, - device=fake_score_critic_noise.device, - dtype=torch.int64 - ) - + current_t = denoising_step_list[index].expand(bsz).to(accelerator.device) + with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): - fake_score_denoised_pred = generator_transformer3d( - x=fake_score_critic_noise, - cap_feats=prompt_embeds, - t=(1000 - timestep) / 1000 - )[0] - fake_score_denoised_pred = -fake_score_denoised_pred - fake_score_denoised_pred = convert_flow_pred_to_x0( - scheduler=noise_scheduler, - flow_pred=fake_score_denoised_pred, + fake_score_denoised_pred, _ = denoise( + model=generator_transformer3d, xt=fake_score_critic_noise, - timestep=timestep + timestep=current_t, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling, + multiply_c_in=False if index == 0 else True, ) if is_final_step: break - - next_timestep = denoising_step_list[index + 1] * torch.ones( - fake_score_critic_noise.shape[:1], - dtype=torch.long, - device=fake_score_critic_noise.device - ) - - fake_score_critic_noise = add_noise( - fake_score_denoised_pred, - torch.randn(fake_score_denoised_pred.shape, dtype=fake_score_denoised_pred.dtype, device=fake_score_denoised_pred.device, generator=torch_rng), - next_timestep - ) - indices = idx_sampling(bsz, generator=torch_rng, device=accelerator.device).long().cpu() - critic_timestep = noise_scheduler.timesteps[indices].to(device=accelerator.device) + if index < num_denoising_steps - 1: + next_t = denoising_step_list[index + 1].expand(bsz).to(accelerator.device) + fake_score_critic_noise = add_noise( + fake_score_denoised_pred, + torch.randn(fake_score_denoised_pred.shape, dtype=fake_score_denoised_pred.dtype, device=fake_score_denoised_pred.device, generator=torch_rng), + next_t, + noise_scheduler=noise_scheduler + ) + + # Sample timesteps for critic + if getattr(args, 'use_trigflow', False): + # Sample timesteps for discriminator (D distribution) + critic_timestep = sample_trigflow_timesteps_D(bsz, device=accelerator.device) + else: + indices = idx_sampling(bsz, generator=torch_rng, device=accelerator.device).long().cpu() + critic_timestep = noise_scheduler.timesteps[indices].to(device=accelerator.device) critic_noise = torch.randn(fake_score_denoised_pred.shape, dtype=fake_score_denoised_pred.dtype, device=fake_score_denoised_pred.device, generator=torch_rng) fake_score_denoised_input = add_noise( fake_score_denoised_pred, critic_noise, - critic_timestep + critic_timestep, + noise_scheduler=noise_scheduler ) with torch.cuda.amp.autocast(dtype=weight_dtype), torch.cuda.device(device=accelerator.device): - fake_score_denoised_output = fake_score_transformer3d( - x=fake_score_denoised_input, - cap_feats=prompt_embeds, - t=(1000 - critic_timestep) / 1000 - )[0] - fake_score_denoised_output = -fake_score_denoised_output + fake_score_pred, _ = denoise( + model=fake_score_transformer3d, + xt=fake_score_denoised_input, + timestep=critic_timestep, + prompt_embeds=prompt_embeds, + noise_scheduler=noise_scheduler, + trigflow_scaling=scaling + ) def custom_mse_loss(noise_pred, target, weighting=None, threshold=50): noise_pred = noise_pred.float() @@ -1732,7 +1782,30 @@ def main(): final_loss = masked_loss.mean() return final_loss - denoising_loss = custom_mse_loss(fake_score_denoised_output, critic_noise - fake_score_denoised_pred) + # Compute weighting based on sin(t) (following rCM) + if getattr(args, 'use_trigflow', False): + ndim = fake_score_denoised_input.ndim + if critic_timestep.ndim == 1: + critic_timestep_view = critic_timestep.view(-1, 1) + else: + critic_timestep_view = critic_timestep + + if ndim == 4: + critic_t_expanded = critic_timestep_view.view(-1, 1, 1, 1) + elif ndim == 5: + critic_t_expanded = critic_timestep_view.view(-1, 1, 1, 1, 1) + + sin_t = torch.sin(critic_t_expanded) + weighting = 1.0 / (sin_t ** 2 + 1e-8) + else: + weighting = None + + denoising_loss = custom_mse_loss( + fake_score_pred, + fake_score_denoised_pred, + weighting=weighting + ) + avg_denoising_loss = accelerator_fake_score_transformer3d.gather(denoising_loss.repeat(args.train_batch_size)).mean() train_denoising_loss += avg_denoising_loss.item() / args.gradient_accumulation_steps diff --git a/scripts/z_image/train_lora.py b/scripts/z_image/train_lora.py index 7d92de8..d53692e 100644 --- a/scripts/z_image/train_lora.py +++ b/scripts/z_image/train_lora.py @@ -206,7 +206,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -850,7 +850,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -879,15 +879,7 @@ def main(): batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: + else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) @@ -900,34 +892,12 @@ def main(): safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) else: - network_state_dict = accelerate_state_dict + network_state_dict = {} + for key in accelerate_state_dict: + if "network" in key: + network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype) save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"}) - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - else: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - if accelerator.is_main_process: - safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors") - if args.use_peft_lora: - network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1])) - save_model(safetensor_save_path, network_state_dict) - - network_state_dict_kohya = convert_peft_lora_to_kohya_lora(network_state_dict) - safetensor_kohya_format_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model_compatible_with_comfyui.safetensors") - save_model(safetensor_kohya_format_save_path, network_state_dict_kohya) - else: - save_model(safetensor_save_path, accelerator.unwrap_model(models[-1])) - if not args.use_deepspeed: for _ in range(len(weights)): weights.pop() @@ -1207,23 +1177,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - elif fsdp_stage != 0: + else: transformer3d.network = network transformer3d = transformer3d.to(dtype=weight_dtype) transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( transformer3d, optimizer, train_dataloader, lr_scheduler ) - else: - network, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( - network, optimizer, train_dataloader, lr_scheduler - ) - - if zero_stage != 0 and not args.use_peft_lora: - from functools import partial - - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(transformer3d.layers)) - transformer3d = shard_fn(transformer3d) if fsdp_stage != 0 or zero_stage != 0: from functools import partial diff --git a/scripts/z_image_fun/train_control.py b/scripts/z_image_fun/train_control.py index 2bd7048..e4d5b28 100644 --- a/scripts/z_image_fun/train_control.py +++ b/scripts/z_image_fun/train_control.py @@ -207,7 +207,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -928,7 +928,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -948,26 +948,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): diff --git a/scripts/z_image_fun/train_control_distill.py b/scripts/z_image_fun/train_control_distill.py index dc91a51..8e452d4 100644 --- a/scripts/z_image_fun/train_control_distill.py +++ b/scripts/z_image_fun/train_control_distill.py @@ -213,7 +213,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato vae=vae, text_encoder=text_encoder, tokenizer=tokenizer, - transformer=transformer3d, + transformer=accelerator.unwrap_model(transformer3d) if type(transformer3d).__name__ == 'DistributedDataParallel' else transformer3d, scheduler=scheduler, ) pipeline = pipeline.to(accelerator.device) @@ -987,7 +987,7 @@ def main(): # `accelerate` 0.16.0 will have better support for customized saving if version.parse(accelerate.__version__) >= version.parse("0.16.0"): # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage == 3: def save_model_hook(models, weights, output_dir): accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) if accelerator.is_main_process: @@ -1007,26 +1007,6 @@ def main(): loaded_number, _ = pickle.load(file) batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") - - elif zero_stage == 3: - # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format - def save_model_hook(models, weights, output_dir): - accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True) - if accelerator.is_main_process: - from safetensors.torch import save_file - safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors") - save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"}) - - with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file: - pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file) - - def load_model_hook(models, input_dir): - pkl_path = os.path.join(input_dir, "sampler_pos_start.pkl") - if os.path.exists(pkl_path): - with open(pkl_path, 'rb') as file: - loaded_number, _ = pickle.load(file) - batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0) - print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.") else: # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format def save_model_hook(models, weights, output_dir): @@ -1420,14 +1400,14 @@ def main(): fake_score_transformer3d, critic_optimizer, fake_score_lr_scheduler ) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=list(real_score_transformer3d.layers)) real_score_transformer3d = shard_fn(real_score_transformer3d) - if fsdp_stage != 0: + if fsdp_stage != 0 or zero_stage != 0: from functools import partial from videox_fun.dist import set_multi_gpus_devices, shard_model @@ -1662,6 +1642,18 @@ def main(): if args.low_vram: real_score_transformer3d = real_score_transformer3d.to(accelerator.device) + image_seq_len = int(target_shape[-1] // 2 * target_shape[-2] // 2) + mu = calculate_shift( + image_seq_len, + noise_scheduler.config.get("base_image_seq_len", 256), + noise_scheduler.config.get("max_image_seq_len", 4096), + noise_scheduler.config.get("base_shift", 0.5), + noise_scheduler.config.get("max_shift", 1.15), + ) + noise_scheduler.sigma_min = 0.0 + noise_scheduler.set_timesteps(args.train_sampling_steps, device=accelerator.device, mu=mu) + denoising_step_list = noise_scheduler.timesteps[args.train_sampling_steps - torch.tensor(args.denoising_step_indices_list)] + with accelerator.accumulate(generator_transformer3d): def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): sigmas = noise_scheduler.sigmas.to(device=accelerator.device, dtype=dtype) diff --git a/videox_fun/models/cogvideox_transformer3d.py b/videox_fun/models/cogvideox_transformer3d.py index 0b15b93..beb2971 100755 --- a/videox_fun/models/cogvideox_transformer3d.py +++ b/videox_fun/models/cogvideox_transformer3d.py @@ -182,8 +182,9 @@ class CogVideoXPatchEmbed(nn.Module): post_time_compression_frames, self.spatial_interpolation_scale, self.temporal_interpolation_scale, + output_type="pt", ) - pos_embedding = torch.from_numpy(pos_embedding).flatten(0, 1) + pos_embedding = pos_embedding.flatten(0, 1) joint_pos_embedding = torch.zeros( 1, self.max_text_seq_length + num_patches, self.embed_dim, requires_grad=False ) diff --git a/videox_fun/utils/__init__.py b/videox_fun/utils/__init__.py index 7628f0c..e6956eb 100755 --- a/videox_fun/utils/__init__.py +++ b/videox_fun/utils/__init__.py @@ -11,10 +11,13 @@ from .fp8_optimization import (autocast_model_forward, from .group_offload import (register_auto_device_hook, safe_enable_group_offload, safe_remove_group_offloading) -from .lora_utils import merge_lora, unmerge_lora -from .utils import (filter_kwargs, get_autocast_dtype, get_image_latent, - get_image_to_video_latent, get_video_to_video_latent, - save_videos_grid) +from .lora_utils import (convert_peft_lora_to_kohya_lora, create_network, + merge_lora, unmerge_lora) +from .trigflow_sampler import (RectifiedFlow_TrigFlowWrapper, + sample_trigflow_timesteps) +from .utils import (calculate_dimensions, filter_kwargs, get_autocast_dtype, + get_image_latent, get_image_to_video_latent, + get_video_to_video_latent, save_videos_grid) # The pai_fuser is an internally developed acceleration package, which can be used on PAI. if importlib.util.find_spec("paifuser") is not None: diff --git a/videox_fun/utils/trigflow_sampler.py b/videox_fun/utils/trigflow_sampler.py new file mode 100644 index 0000000..4d7fbc1 --- /dev/null +++ b/videox_fun/utils/trigflow_sampler.py @@ -0,0 +1,23 @@ +import torch + +# Copied from https://github.com/NVlabs/rcm/blob/main/rcm/utils/denoiser_scaling.py +class RectifiedFlow_TrigFlowWrapper: + def __init__(self, sigma_data: float = 1.0, t_scaling_factor: float = 1.0): + assert abs(sigma_data - 1.0) < 1e-6, "sigma_data must be 1.0 for RectifiedFlowScaling" + self.t_scaling_factor = t_scaling_factor + + def __call__(self, trigflow_t: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + trigflow_t = trigflow_t.to(torch.float64) + c_skip = 1 / (torch.cos(trigflow_t) + torch.sin(trigflow_t)) + c_out = -1 * torch.sin(trigflow_t) / (torch.cos(trigflow_t) + torch.sin(trigflow_t)) + c_in = 1 / (torch.cos(trigflow_t) + torch.sin(trigflow_t)) + c_noise = (torch.sin(trigflow_t) / (torch.cos(trigflow_t) + torch.sin(trigflow_t))) * self.t_scaling_factor + return c_skip, c_out, c_in, c_noise + +# Sample timesteps +def sample_trigflow_timesteps(batch_size, device, P_mean=0.0, P_std=1.6): + """Sample timesteps for training""" + sigma = torch.randn(batch_size, device=device) + sigma = (sigma * P_std + P_mean).exp() + timesteps = torch.arctan(sigma) + return timesteps \ No newline at end of file