Update fix sample size && Fix vram memory bug in wan2.1 lora training (#227)

This commit is contained in:
Bubbliiiing
2025-06-24 10:09:23 +08:00
committed by GitHub
parent ba1da31fb5
commit 86028a9ce6
14 changed files with 464 additions and 182 deletions
+4 -4
View File
@@ -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()
+5 -5
View File
@@ -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()
+6 -5
View File
@@ -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}")
+58 -18
View File
@@ -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()
+84 -30
View File
@@ -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}")
+64 -19
View File
@@ -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()
+65 -20
View File
@@ -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()
+68 -22
View File
@@ -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}")
+68 -22
View File
@@ -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}")
+25 -29
View File
@@ -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
+6 -3
View File
@@ -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)
+6 -3
View File
@@ -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)
+2 -1
View File
@@ -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)
+3 -1
View File
@@ -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