Update fix sample size && Fix vram memory bug in wan2.1 lora training (#227)
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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}")
|
||||
|
||||
|
||||
@@ -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}")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Vendored
+6
-3
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,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
|
||||
|
||||
Reference in New Issue
Block a user