diff --git a/scripts/cogvideox_fun/train.py b/scripts/cogvideox_fun/train.py index cb2ef53..2cc8d4a 100755 --- a/scripts/cogvideox_fun/train.py +++ b/scripts/cogvideox_fun/train.py @@ -1738,10 +1738,10 @@ def main(): if args.use_ema: ema_transformer3d.copy_to(transformer3d.parameters()) - if args.use_deepspeed or accelerator.is_main_process: - save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - accelerator.save_state(save_path) - logger.info(f"Saved state to {save_path}") + if args.use_deepspeed or accelerator.is_main_process: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") accelerator.end_training() diff --git a/scripts/cogvideox_fun/train_control.py b/scripts/cogvideox_fun/train_control.py index fdcdae2..7e1cfeb 100755 --- a/scripts/cogvideox_fun/train_control.py +++ b/scripts/cogvideox_fun/train_control.py @@ -944,7 +944,7 @@ def main(): video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, video_repeat=args.video_repeat, image_sample_size=args.image_sample_size, - enable_bucket=args.enable_bucket, enable_inpaint=False, + enable_bucket=args.enable_bucket, ) if args.enable_bucket: @@ -1629,10 +1629,10 @@ def main(): if args.use_ema: ema_transformer3d.copy_to(transformer3d.parameters()) - if args.use_deepspeed or accelerator.is_main_process: - save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - accelerator.save_state(save_path) - logger.info(f"Saved state to {save_path}") + if args.use_deepspeed or accelerator.is_main_process: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") accelerator.end_training() diff --git a/scripts/cogvideox_fun/train_lora.py b/scripts/cogvideox_fun/train_lora.py index a0c6d0d..5761b08 100755 --- a/scripts/cogvideox_fun/train_lora.py +++ b/scripts/cogvideox_fun/train_lora.py @@ -1697,11 +1697,12 @@ def main(): # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() - if accelerator.is_main_process: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - if args.save_state: + if args.use_deepspeed or accelerator.is_main_process: + if not args.save_state: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + else: + accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") diff --git a/scripts/wan2.1/train.py b/scripts/wan2.1/train.py index c5aa7e6..4d51476 100755 --- a/scripts/wan2.1/train.py +++ b/scripts/wan2.1/train.py @@ -606,7 +606,12 @@ def parse_args(): "--image_sample_size", type=int, default=512, - help="Sample size of the video.", + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." ) parser.add_argument( "--video_sample_stride", @@ -768,6 +773,9 @@ def main(): print(f"Using DeepSpeed Zero stage: {zero_stage}") args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True elif fsdp_plugin is not None: from torch.distributed.fsdp import ShardingStrategy zero_stage = 0 @@ -782,6 +790,9 @@ def main(): print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True else: zero_stage = 0 fsdp_stage = 0 @@ -875,6 +886,13 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ) + vae.eval() + # Get Clip Image Encoder + if args.train_mode != "normal": + clip_image_encoder = CLIPModel.from_pretrained( + 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() # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -882,13 +900,6 @@ def main(): transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), ).to(weight_dtype) - if args.train_mode != "normal": - # Get Clip Image Encoder - clip_image_encoder = CLIPModel.from_pretrained( - 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() - # Freeze vae and text_encoder and set transformer3d to trainable vae.requires_grad_(False) text_encoder.requires_grad_(False) @@ -1120,6 +1131,13 @@ def main(): # Get the training dataset sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + 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 + + # Get the dataset train_dataset = ImageVideoDataset( args.train_data_meta, args.train_data_dir, video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, @@ -1204,9 +1222,9 @@ def main(): aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} - closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) - closest_size = [int(x / 16) * 16 for x in closest_size] - if args.random_ratio_crop: + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: if rng is None: random_sample_size = aspect_ratio_random_crop_sample_size[ np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) @@ -1216,9 +1234,24 @@ def main(): rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) ] random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] for example in examples: - if args.random_ratio_crop: + if args.fix_sample_size is not None: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: # To 0~1 pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() pixel_values = pixel_values / 255. @@ -1341,6 +1374,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) + if fsdp_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) + if args.use_ema: ema_transformer3d.to(accelerator.device) @@ -1365,6 +1404,7 @@ def main(): tracker_config.pop("validation_prompts") tracker_config.pop("trainable_modules") tracker_config.pop("trainable_modules_low_learning_rate") + tracker_config.pop("fix_sample_size") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1465,7 +1505,7 @@ def main(): pixel_values = batch["pixel_values"].to(weight_dtype) # Increase the batch size when the length of the latent sequence of the current sample is small - if args.training_with_video_token_length and not zero_stage == 3: + if args.training_with_video_token_length and zero_stage != 3: if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1)) if args.enable_text_encoder_in_dataloader: @@ -1486,7 +1526,7 @@ def main(): mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype) mask = batch["mask"].to(weight_dtype) # Increase the batch size when the length of the latent sequence of the current sample is small - if args.training_with_video_token_length and not zero_stage == 3: + if args.training_with_video_token_length and zero_stage != 3: if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1)) mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1)) @@ -1856,10 +1896,10 @@ def main(): if args.use_ema: ema_transformer3d.copy_to(transformer3d.parameters()) - if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: - save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - accelerator.save_state(save_path) - logger.info(f"Saved state to {save_path}") + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") accelerator.end_training() diff --git a/scripts/wan2.1/train_lora.py b/scripts/wan2.1/train_lora.py index 7c021dc..de056e5 100755 --- a/scripts/wan2.1/train_lora.py +++ b/scripts/wan2.1/train_lora.py @@ -625,7 +625,12 @@ def parse_args(): "--image_sample_size", type=int, default=512, - help="Sample size of the video.", + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." ) parser.add_argument( "--video_sample_stride", @@ -653,12 +658,6 @@ def parse_args(): "The config of the model in training." ), ) - parser.add_argument( - "--image_repeat_in_forward", - type=int, - default=0, - help="Num of repeat image in forward.", - ) parser.add_argument( "--transformer_path", type=str, @@ -676,7 +675,7 @@ def parse_args(): parser.add_argument( '--tokenizer_max_length', type=int, - default=226, + default=512, help='Max length of tokenizer' ) parser.add_argument( @@ -694,7 +693,7 @@ def parse_args(): default="normal", help=( 'The format of training data. Support `"normal"`' - ' (default), `"inpaint"`.' + ' (default), `"i2v"`.' ), ) parser.add_argument( @@ -716,6 +715,12 @@ 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( + "--lora_skip_name", + type=str, + default=None, + help=("The module is not trained in loras. "), + ) args = parser.parse_args() env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) @@ -767,6 +772,9 @@ def main(): print(f"Using DeepSpeed Zero stage: {zero_stage}") args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True elif fsdp_plugin is not None: from torch.distributed.fsdp import ShardingStrategy zero_stage = 0 @@ -781,6 +789,9 @@ def main(): print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True else: zero_stage = 0 fsdp_stage = 0 @@ -868,11 +879,19 @@ def main(): low_cpu_mem_usage=True, torch_dtype=weight_dtype, ) + text_encoder = text_encoder.eval() # Get Vae vae = AutoencoderKLWan.from_pretrained( os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ) + vae.eval() + # Get Clip Image Encoder + if args.train_mode != "normal": + clip_image_encoder = CLIPModel.from_pretrained( + 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() # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -880,13 +899,6 @@ def main(): transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), ).to(weight_dtype) - if args.train_mode != "normal": - # Get Clip Image Encoder - clip_image_encoder = CLIPModel.from_pretrained( - 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() - # Freeze vae and text_encoder and set transformer3d to trainable vae.requires_grad_(False) text_encoder.requires_grad_(False) @@ -902,7 +914,7 @@ def main(): text_encoder, transformer3d, neuron_dropout=None, - add_lora_in_attn_temporal=True, + skip_name=args.lora_skip_name, ) network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True) @@ -1055,12 +1067,20 @@ def main(): # Get the training dataset sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + 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 + + # Get the dataset train_dataset = ImageVideoDataset( args.train_data_meta, args.train_data_dir, video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, video_repeat=args.video_repeat, image_sample_size=args.image_sample_size, - enable_bucket=args.enable_bucket, enable_inpaint=True if args.train_mode != "normal" else False, + enable_bucket=args.enable_bucket, + enable_inpaint=True if args.train_mode != "normal" else False, ) if args.enable_bucket: @@ -1086,6 +1106,7 @@ def main(): } return length_to_frame_num + def collate_fn(examples): # Get token length target_token_length = args.video_sample_n_frames * args.token_sample_size * args.token_sample_size @@ -1138,9 +1159,9 @@ def main(): aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} - closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) - closest_size = [int(x / 16) * 16 for x in closest_size] - if args.random_ratio_crop: + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: if rng is None: random_sample_size = aspect_ratio_random_crop_sample_size[ np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) @@ -1150,9 +1171,24 @@ def main(): rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) ] random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] for example in examples: - if args.random_ratio_crop: + if args.fix_sample_size is not None: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: # To 0~1 pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() pixel_values = pixel_values / 255. @@ -1282,6 +1318,18 @@ def main(): network, optimizer, train_dataloader, lr_scheduler ) + if zero_stage == 3: + 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: + 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 vae.to(accelerator.device, dtype=weight_dtype) transformer3d.to(accelerator.device, dtype=weight_dtype) @@ -1302,6 +1350,7 @@ def main(): if accelerator.is_main_process: tracker_config = dict(vars(args)) tracker_config.pop("validation_prompts") + tracker_config.pop("fix_sample_size") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1354,7 +1403,7 @@ def main(): first_epoch = global_step // num_update_steps_per_epoch print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") - if zero_stage != 3: + if zero_stage != 3 and not args.use_fsdp: from safetensors.torch import load_file state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device)) m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False) @@ -1463,7 +1512,7 @@ def main(): pixel_values = batch["pixel_values"].to(weight_dtype) # Increase the batch size when the length of the latent sequence of the current sample is small - if args.training_with_video_token_length: + if args.training_with_video_token_length and zero_stage != 3: if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: pixel_values = torch.tile(pixel_values, (4, 1, 1, 1, 1)) if args.enable_text_encoder_in_dataloader: @@ -1484,7 +1533,7 @@ def main(): mask_pixel_values = batch["mask_pixel_values"].to(weight_dtype) mask = batch["mask"].to(weight_dtype) # Increase the batch size when the length of the latent sequence of the current sample is small - if args.training_with_video_token_length: + if args.training_with_video_token_length and zero_stage != 3: if args.video_sample_n_frames * args.token_sample_size * args.token_sample_size // 16 >= pixel_values.size()[1] * pixel_values.size()[3] * pixel_values.size()[4]: clip_pixel_values = torch.tile(clip_pixel_values, (4, 1, 1, 1)) mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1)) @@ -1563,6 +1612,8 @@ def main(): if args.low_vram: torch.cuda.empty_cache() vae.to(accelerator.device) + if args.train_mode != "normal": + clip_image_encoder.to(accelerator.device) if not args.enable_text_encoder_in_dataloader: text_encoder.to("cpu") @@ -1619,6 +1670,8 @@ def main(): if args.low_vram: vae.to('cpu') + if args.train_mode != "normal": + clip_image_encoder.to('cpu') torch.cuda.empty_cache() if not args.enable_text_encoder_in_dataloader: text_encoder.to(accelerator.device) @@ -1813,11 +1866,12 @@ def main(): # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() - if accelerator.is_main_process: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - if args.save_state: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + if not args.save_state: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + else: + accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") diff --git a/scripts/wan2.1_fun/train.py b/scripts/wan2.1_fun/train.py index cf89f6e..bf19ea3 100755 --- a/scripts/wan2.1_fun/train.py +++ b/scripts/wan2.1_fun/train.py @@ -576,7 +576,12 @@ def parse_args(): "--image_sample_size", type=int, default=512, - help="Sample size of the video.", + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." ) parser.add_argument( "--video_sample_stride", @@ -738,6 +743,9 @@ def main(): print(f"Using DeepSpeed Zero stage: {zero_stage}") args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True elif fsdp_plugin is not None: from torch.distributed.fsdp import ShardingStrategy zero_stage = 0 @@ -752,6 +760,9 @@ def main(): print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True else: zero_stage = 0 fsdp_stage = 0 @@ -845,6 +856,13 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ) + vae.eval() + # Get Clip Image Encoder + if args.train_mode != "normal": + clip_image_encoder = CLIPModel.from_pretrained( + 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() # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -852,13 +870,6 @@ def main(): transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), ).to(weight_dtype) - if args.train_mode != "normal": - # Get Clip Image Encoder - clip_image_encoder = CLIPModel.from_pretrained( - 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() - # Freeze vae and text_encoder and set transformer3d to trainable vae.requires_grad_(False) text_encoder.requires_grad_(False) @@ -1090,6 +1101,13 @@ def main(): # Get the training dataset sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + 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 + + # Get the dataset train_dataset = ImageVideoDataset( args.train_data_meta, args.train_data_dir, video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, @@ -1210,16 +1228,36 @@ def main(): aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} - closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) - closest_size = [int(x / 16) * 16 for x in closest_size] - if args.random_ratio_crop: - random_sample_size = aspect_ratio_random_crop_sample_size[ - np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) - ] + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] for example in examples: - if args.random_ratio_crop: + if args.fix_sample_size is not None: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: # To 0~1 pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() pixel_values = pixel_values / 255. @@ -1344,6 +1382,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) + if fsdp_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) + if args.use_ema: ema_transformer3d.to(accelerator.device) @@ -1368,6 +1412,7 @@ def main(): tracker_config.pop("validation_prompts") tracker_config.pop("trainable_modules") tracker_config.pop("trainable_modules_low_learning_rate") + tracker_config.pop("fix_sample_size") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1865,10 +1910,10 @@ def main(): if args.use_ema: ema_transformer3d.copy_to(transformer3d.parameters()) - if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: - save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - accelerator.save_state(save_path) - logger.info(f"Saved state to {save_path}") + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") accelerator.end_training() diff --git a/scripts/wan2.1_fun/train_control.py b/scripts/wan2.1_fun/train_control.py index ef471de..d45fd27 100755 --- a/scripts/wan2.1_fun/train_control.py +++ b/scripts/wan2.1_fun/train_control.py @@ -492,7 +492,12 @@ def parse_args(): "--image_sample_size", type=int, default=512, - help="Sample size of the video.", + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." ) parser.add_argument( "--video_sample_stride", @@ -670,6 +675,9 @@ def main(): print(f"Using DeepSpeed Zero stage: {zero_stage}") args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True elif fsdp_plugin is not None: from torch.distributed.fsdp import ShardingStrategy zero_stage = 0 @@ -684,6 +692,9 @@ def main(): print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True else: zero_stage = 0 fsdp_stage = 0 @@ -777,6 +788,13 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ) + vae.eval() + # Get Clip Image Encoder + if args.train_mode != "normal": + clip_image_encoder = CLIPModel.from_pretrained( + 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() # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -784,13 +802,6 @@ def main(): transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), ).to(weight_dtype) - if args.train_mode != "normal": - # Get Clip Image Encoder - clip_image_encoder = CLIPModel.from_pretrained( - 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() - # Freeze vae and text_encoder and set transformer3d to trainable vae.requires_grad_(False) text_encoder.requires_grad_(False) @@ -1020,13 +1031,19 @@ def main(): # Get the training dataset sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + 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 + # Get the dataset train_dataset = ImageVideoControlDataset( args.train_data_meta, args.train_data_dir, video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, video_repeat=args.video_repeat, image_sample_size=args.image_sample_size, - enable_bucket=args.enable_bucket, enable_inpaint=False, + enable_bucket=args.enable_bucket, enable_camera_info=args.train_mode == "control_camera_ref" ) @@ -1170,13 +1187,21 @@ def main(): aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} - closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) - closest_size = [int(x / 16) * 16 for x in closest_size] - if args.random_ratio_crop: - random_sample_size = aspect_ratio_random_crop_sample_size[ - np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) - ] + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] for example in examples: # To 0~1 @@ -1186,7 +1211,20 @@ def main(): control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() control_pixel_values = control_pixel_values / 255. - if args.random_ratio_crop: + if args.fix_sample_size is not None: + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + + transform_no_normalize = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + ]) + elif args.random_ratio_crop: # Get adapt hw for resize b, c, h, w = pixel_values.size() th, tw = random_sample_size @@ -1347,6 +1385,12 @@ def main(): transformer3d, optimizer, train_dataloader, lr_scheduler ) + if fsdp_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) + if args.use_ema: ema_transformer3d.to(accelerator.device) @@ -1371,6 +1415,7 @@ def main(): tracker_config.pop("validation_prompts") tracker_config.pop("trainable_modules") tracker_config.pop("trainable_modules_low_learning_rate") + tracker_config.pop("fix_sample_size") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1902,10 +1947,10 @@ def main(): if args.use_ema: ema_transformer3d.copy_to(transformer3d.parameters()) - if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: - save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - accelerator.save_state(save_path) - logger.info(f"Saved state to {save_path}") + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") accelerator.end_training() diff --git a/scripts/wan2.1_fun/train_control_lora.py b/scripts/wan2.1_fun/train_control_lora.py index b36c39b..862a2d8 100755 --- a/scripts/wan2.1_fun/train_control_lora.py +++ b/scripts/wan2.1_fun/train_control_lora.py @@ -510,7 +510,12 @@ def parse_args(): "--image_sample_size", type=int, default=512, - help="Sample size of the video.", + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." ) parser.add_argument( "--video_sample_stride", @@ -668,6 +673,9 @@ def main(): print(f"Using DeepSpeed Zero stage: {zero_stage}") args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True elif fsdp_plugin is not None: from torch.distributed.fsdp import ShardingStrategy zero_stage = 0 @@ -682,6 +690,9 @@ def main(): print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True else: zero_stage = 0 fsdp_stage = 0 @@ -775,6 +786,13 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ) + vae.eval() + # Get Clip Image Encoder + if args.train_mode != "normal": + clip_image_encoder = CLIPModel.from_pretrained( + 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() # Get Transformer transformer3d = WanTransformer3DModel.from_pretrained( @@ -782,13 +800,6 @@ def main(): transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']), ).to(weight_dtype) - if args.train_mode != "normal": - # Get Clip Image Encoder - clip_image_encoder = CLIPModel.from_pretrained( - 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() - # Freeze vae and text_encoder and set transformer3d to trainable vae.requires_grad_(False) text_encoder.requires_grad_(False) @@ -955,13 +966,19 @@ def main(): # Get the training dataset sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + 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 + # Get the dataset train_dataset = ImageVideoControlDataset( args.train_data_meta, args.train_data_dir, video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, video_repeat=args.video_repeat, image_sample_size=args.image_sample_size, - enable_bucket=args.enable_bucket, enable_inpaint=False, + enable_bucket=args.enable_bucket, enable_camera_info=args.train_mode == "control_camera_ref" ) @@ -1105,13 +1122,21 @@ def main(): aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} - closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) - closest_size = [int(x / 16) * 16 for x in closest_size] - if args.random_ratio_crop: - random_sample_size = aspect_ratio_random_crop_sample_size[ - np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) - ] + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] for example in examples: # To 0~1 @@ -1121,7 +1146,20 @@ def main(): control_pixel_values = torch.from_numpy(example["control_pixel_values"]).permute(0, 3, 1, 2).contiguous() control_pixel_values = control_pixel_values / 255. - if args.random_ratio_crop: + if args.fix_sample_size is not None: + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + + transform_no_normalize = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + ]) + elif args.random_ratio_crop: # Get adapt hw for resize b, c, h, w = pixel_values.size() th, tw = random_sample_size @@ -1295,6 +1333,12 @@ def main(): shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) transformer3d = shard_fn(transformer3d) + if fsdp_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 vae.to(accelerator.device, dtype=weight_dtype) transformer3d.to(accelerator.device, dtype=weight_dtype) @@ -1315,6 +1359,7 @@ def main(): if accelerator.is_main_process: tracker_config = dict(vars(args)) tracker_config.pop("validation_prompts") + tracker_config.pop("fix_sample_size") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1367,7 +1412,7 @@ def main(): first_epoch = global_step // num_update_steps_per_epoch print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") - if zero_stage != 3: + if zero_stage != 3 and not args.use_fsdp: from safetensors.torch import load_file state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device)) m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False) @@ -1870,11 +1915,12 @@ def main(): # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() - if accelerator.is_main_process: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - if args.save_state: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + if not args.save_state: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + else: + accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") diff --git a/scripts/wan2.1_fun/train_lora.py b/scripts/wan2.1_fun/train_lora.py index 5d392fe..512c4ee 100755 --- a/scripts/wan2.1_fun/train_lora.py +++ b/scripts/wan2.1_fun/train_lora.py @@ -588,7 +588,12 @@ def parse_args(): "--image_sample_size", type=int, default=512, - help="Sample size of the video.", + help="Sample size of the image.", + ) + parser.add_argument( + "--fix_sample_size", + nargs=2, type=int, default=None, + help="Fix Sample size [height, width] when using bucket and collate_fn." ) parser.add_argument( "--video_sample_stride", @@ -730,6 +735,9 @@ def main(): print(f"Using DeepSpeed Zero stage: {zero_stage}") args.use_deepspeed = True + if zero_stage == 3: + print(f"Auto set save_state to True because zero_stage == 3") + args.save_state = True elif fsdp_plugin is not None: from torch.distributed.fsdp import ShardingStrategy zero_stage = 0 @@ -744,6 +752,9 @@ def main(): print(f"Using FSDP stage: {fsdp_stage}") args.use_fsdp = True + if fsdp_stage == 3: + print(f"Auto set save_state to True because fsdp_stage == 3") + args.save_state = True else: zero_stage = 0 fsdp_stage = 0 @@ -837,20 +848,20 @@ def main(): os.path.join(args.pretrained_model_name_or_path, config['vae_kwargs'].get('vae_subpath', 'vae')), additional_kwargs=OmegaConf.to_container(config['vae_kwargs']), ) - + vae.eval() + # Get Clip Image Encoder + if args.train_mode != "normal": + clip_image_encoder = CLIPModel.from_pretrained( + 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() + # Get Transformer transformer3d = 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) - if args.train_mode != "normal": - # Get Clip Image Encoder - clip_image_encoder = CLIPModel.from_pretrained( - 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() - # Freeze vae and text_encoder and set transformer3d to trainable vae.requires_grad_(False) text_encoder.requires_grad_(False) @@ -1019,6 +1030,13 @@ def main(): # Get the training dataset sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio + 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 + + # Get the dataset train_dataset = ImageVideoDataset( args.train_data_meta, args.train_data_dir, video_sample_size=args.video_sample_size, video_sample_stride=args.video_sample_stride, video_sample_n_frames=args.video_sample_n_frames, @@ -1139,16 +1157,36 @@ def main(): aspect_ratio_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_512[key]] for key in ASPECT_RATIO_512.keys()} aspect_ratio_random_crop_sample_size = {key : [x / 512 * args.video_sample_size / random_downsample_ratio for x in ASPECT_RATIO_RANDOM_CROP_512[key]] for key in ASPECT_RATIO_RANDOM_CROP_512.keys()} - closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) - closest_size = [int(x / 16) * 16 for x in closest_size] - if args.random_ratio_crop: - random_sample_size = aspect_ratio_random_crop_sample_size[ - np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) - ] + if args.fix_sample_size is not None: + fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size] + elif args.random_ratio_crop: + if rng is None: + random_sample_size = aspect_ratio_random_crop_sample_size[ + np.random.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] + else: + random_sample_size = aspect_ratio_random_crop_sample_size[ + rng.choice(list(aspect_ratio_random_crop_sample_size.keys()), p = ASPECT_RATIO_RANDOM_CROP_PROB) + ] random_sample_size = [int(x / 16) * 16 for x in random_sample_size] + else: + closest_size, closest_ratio = get_closest_ratio(h, w, ratios=aspect_ratio_sample_size) + closest_size = [int(x / 16) * 16 for x in closest_size] for example in examples: - if args.random_ratio_crop: + if args.fix_sample_size is not None: + # To 0~1 + pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() + pixel_values = pixel_values / 255. + + # Get adapt hw for resize + fix_sample_size = list(map(lambda x: int(x), fix_sample_size)) + transform = transforms.Compose([ + transforms.Resize(fix_sample_size, interpolation=transforms.InterpolationMode.BILINEAR), # Image.BICUBIC + transforms.CenterCrop(fix_sample_size), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True), + ]) + elif args.random_ratio_crop: # To 0~1 pixel_values = torch.from_numpy(example["pixel_values"]).permute(0, 3, 1, 2).contiguous() pixel_values = pixel_values / 255. @@ -1286,6 +1324,12 @@ def main(): shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype) transformer3d = shard_fn(transformer3d) + if fsdp_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 vae.to(accelerator.device, dtype=weight_dtype) transformer3d.to(accelerator.device, dtype=weight_dtype) @@ -1306,6 +1350,7 @@ def main(): if accelerator.is_main_process: tracker_config = dict(vars(args)) tracker_config.pop("validation_prompts") + tracker_config.pop("fix_sample_size") accelerator.init_trackers(args.tracker_project_name, tracker_config) # Function for unwrapping if model was compiled with `torch.compile`. @@ -1358,7 +1403,7 @@ def main(): first_epoch = global_step // num_update_steps_per_epoch print(f"Load pkl from {pkl_path}. Get first_epoch = {first_epoch}.") - if zero_stage != 3: + if zero_stage != 3 and not args.use_fsdp: from safetensors.torch import load_file state_dict = load_file(os.path.join(checkpoint_folder_path, "lora_diffusion_pytorch_model.safetensors"), device=str(accelerator.device)) m, u = accelerator.unwrap_model(network).load_state_dict(state_dict, strict=False) @@ -1827,11 +1872,12 @@ def main(): # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() - if accelerator.is_main_process: - safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") - accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") - save_model(safetensor_save_path, accelerator.unwrap_model(network)) - if args.save_state: + if args.use_deepspeed or args.use_fsdp or accelerator.is_main_process: + if not args.save_state: + safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors") + save_model(safetensor_save_path, accelerator.unwrap_model(network)) + else: + accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") accelerator.save_state(accelerator_save_path) logger.info(f"Saved state to {accelerator_save_path}") diff --git a/videox_fun/data/dataset_image_video.py b/videox_fun/data/dataset_image_video.py index 395cd05..8270aa7 100755 --- a/videox_fun/data/dataset_image_video.py +++ b/videox_fun/data/dataset_image_video.py @@ -334,17 +334,17 @@ def resize_frame(frame, target_short_side): class ImageVideoDataset(Dataset): def __init__( - self, - ann_path, data_root=None, - video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16, - image_sample_size=512, - video_repeat=0, - text_drop_ratio=0.1, - enable_bucket=False, - video_length_drop_start=0.0, - video_length_drop_end=1.0, - enable_inpaint=False, - ): + self, + ann_path, data_root=None, + video_sample_size=512, video_sample_stride=4, video_sample_n_frames=16, + image_sample_size=512, + video_repeat=0, + text_drop_ratio=0.1, + enable_bucket=False, + video_length_drop_start=0.0, + video_length_drop_end=1.0, + enable_inpaint=False, + ): # Loading annotations from files print(f"loading annotations from {ann_path} ...") if ann_path.endswith('.csv'): @@ -356,15 +356,18 @@ class ImageVideoDataset(Dataset): self.data_root = data_root # It's used to balance num of images and videos. - self.dataset = [] - for data in dataset: - if data.get('type', 'image') != 'video': - self.dataset.append(data) if video_repeat > 0: + self.dataset = [] + for data in dataset: + if data.get('type', 'image') != 'video': + self.dataset.append(data) + for _ in range(video_repeat): for data in dataset: if data.get('type', 'image') == 'video': self.dataset.append(data) + else: + self.dataset = dataset del dataset self.length = len(self.dataset) @@ -503,11 +506,6 @@ class ImageVideoDataset(Dataset): clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255 sample["clip_pixel_values"] = clip_pixel_values - ref_pixel_values = sample["pixel_values"][0].unsqueeze(0) - if (mask == 1).all(): - ref_pixel_values = torch.ones_like(ref_pixel_values) * -1 - sample["ref_pixel_values"] = ref_pixel_values - return sample class ImageVideoControlDataset(Dataset): @@ -535,15 +533,18 @@ class ImageVideoControlDataset(Dataset): self.data_root = data_root # It's used to balance num of images and videos. - self.dataset = [] - for data in dataset: - if data.get('type', 'image') != 'video': - self.dataset.append(data) if video_repeat > 0: + self.dataset = [] + for data in dataset: + if data.get('type', 'image') != 'video': + self.dataset.append(data) + for _ in range(video_repeat): for data in dataset: if data.get('type', 'image') == 'video': self.dataset.append(data) + else: + self.dataset = dataset del dataset self.length = len(self.dataset) @@ -767,9 +768,4 @@ class ImageVideoControlDataset(Dataset): clip_pixel_values = (clip_pixel_values * 0.5 + 0.5) * 255 sample["clip_pixel_values"] = clip_pixel_values - ref_pixel_values = sample["pixel_values"][0].unsqueeze(0) - if (mask == 1).all(): - ref_pixel_values = torch.ones_like(ref_pixel_values) * -1 - sample["ref_pixel_values"] = ref_pixel_values - return sample diff --git a/videox_fun/dist/__init__.py b/videox_fun/dist/__init__.py index ac1612e..267c0b5 100755 --- a/videox_fun/dist/__init__.py +++ b/videox_fun/dist/__init__.py @@ -29,13 +29,16 @@ if importlib.util.find_spec("pai_fuser") is not None: if ENABLE_KERNEL: import torch + import types from .wan_xfuser import rope_apply + def deepcopy_function(f): + return types.FunctionType(f.__code__, f.__globals__, name=f.__name__, argdefs=f.__defaults__,closure=f.__closure__) + + local_rope_apply_qk = deepcopy_function(wan_xfuser.rope_apply_qk) def adaptive_fast_usp_rope_apply_qk(q, k, grid_sizes, freqs): if torch.is_grad_enabled(): - q = rope_apply(q, grid_sizes, freqs) - k = rope_apply(k, grid_sizes, freqs) - return q, k + return local_rope_apply_qk(q, k, grid_sizes, freqs) else: return usp_fast_rope_apply_qk(q, k, grid_sizes, freqs) diff --git a/videox_fun/models/__init__.py b/videox_fun/models/__init__.py index 2878198..6dccf30 100755 --- a/videox_fun/models/__init__.py +++ b/videox_fun/models/__init__.py @@ -37,13 +37,16 @@ if importlib.util.find_spec("pai_fuser") is not None: from pai_fuser.core.rope import ENABLE_KERNEL, fast_rope_apply_qk if ENABLE_KERNEL: + import types from .wan_transformer3d import rope_apply + def deepcopy_function(f): + return types.FunctionType(f.__code__, f.__globals__, name=f.__name__, argdefs=f.__defaults__,closure=f.__closure__) + + local_rope_apply_qk = deepcopy_function(wan_transformer3d.rope_apply_qk) def adaptive_fast_rope_apply_qk(q, k, grid_sizes, freqs): if torch.is_grad_enabled(): - q = rope_apply(q, grid_sizes, freqs) - k = rope_apply(k, grid_sizes, freqs) - return q, k + return local_rope_apply_qk(q, k, grid_sizes, freqs) else: return fast_rope_apply_qk(q, k, grid_sizes, freqs) diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 17796a4..de9a987 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -912,6 +912,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): def enable_multi_gpus_inference(self,): self.sp_world_size = get_sequence_parallel_world_size() self.sp_world_rank = get_sequence_parallel_rank() + self.all_gather = get_sp_group().all_gather for block in self.blocks: block.self_attn.forward = types.MethodType( usp_attn_forward, block.self_attn) @@ -1132,7 +1133,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): x = block(x, **kwargs) if self.sp_world_size > 1: - x = get_sp_group().all_gather(x, dim=1) + x = self.all_gather(x, dim=1) if self.ref_conv is not None and full_ref is not None: full_ref_length = full_ref.size(1) diff --git a/videox_fun/utils/discrete_sampler.py b/videox_fun/utils/discrete_sampler.py index 149dbe7..84d23d2 100644 --- a/videox_fun/utils/discrete_sampler.py +++ b/videox_fun/utils/discrete_sampler.py @@ -3,7 +3,7 @@ import torch class DiscreteSampling: - def __init__(self, num_idx, uniform_sampling=False): + def __init__(self, num_idx, uniform_sampling=False, sp_size=1): self.num_idx = num_idx self.uniform_sampling = uniform_sampling self.is_distributed = torch.distributed.is_available() and torch.distributed.is_initialized() @@ -21,6 +21,8 @@ class DiscreteSampling: break assert self.group_num > 0 assert world_size % self.group_num == 0 + if self.group_num >= sp_size: + self.group_num = self.group_num // sp_size # the number of rank in one group self.group_width = world_size // self.group_num self.sigma_interval = self.num_idx // self.group_num