Fix bugs in the training code (#356)

This commit is contained in:
Bubbliiiing
2025-10-17 16:41:47 +08:00
committed by GitHub
parent 8c34acc600
commit 900b181ac6
30 changed files with 578 additions and 162 deletions
+6 -5
View File
@@ -72,7 +72,7 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
from videox_fun.pipeline import (CogVideoXFunPipeline,
CogVideoXFunControlPipeline,
CogVideoXFunInpaintPipeline)
from videox_fun.pipeline.pipeline_CogVideoXFuninpaint import (
from videox_fun.pipeline.pipeline_cogvideox_fun_inpaint import (
add_noise_to_reference_video, get_3d_rotary_pos_embed,
get_resize_crop_region_for_grid)
from videox_fun.utils.discrete_sampler import DiscreteSampling
@@ -1244,9 +1244,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1366,7 +1367,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.auto_tile_batch_size and args.training_with_video_token_length:
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]:
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
mask = torch.tile(mask, (4, 1, 1, 1, 1))
+5 -5
View File
@@ -70,7 +70,7 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
from videox_fun.pipeline import (CogVideoXFunPipeline,
CogVideoXFunControlPipeline,
CogVideoXFunInpaintPipeline)
from videox_fun.pipeline.pipeline_CogVideoXFuninpaint import (
from videox_fun.pipeline.pipeline_cogvideox_fun_inpaint import (
add_noise_to_reference_video, get_3d_rotary_pos_embed,
get_resize_crop_region_for_grid)
from videox_fun.utils.discrete_sampler import DiscreteSampling
@@ -1179,10 +1179,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("validation_paths")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+6 -3
View File
@@ -71,7 +71,7 @@ from videox_fun.models import (AutoencoderKLCogVideoX,
from videox_fun.pipeline import (CogVideoXFunPipeline,
CogVideoXFunControlPipeline,
CogVideoXFunInpaintPipeline)
from videox_fun.pipeline.pipeline_CogVideoXFuninpaint import (
from videox_fun.pipeline.pipeline_cogvideox_fun_inpaint import (
add_noise_to_reference_video, get_3d_rotary_pos_embed,
get_resize_crop_region_for_grid)
from videox_fun.utils.discrete_sampler import DiscreteSampling
@@ -1180,7 +1180,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1361,7 +1364,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.auto_tile_batch_size and args.training_with_video_token_length:
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]:
mask_pixel_values = torch.tile(mask_pixel_values, (4, 1, 1, 1, 1))
mask = torch.tile(mask, (4, 1, 1, 1, 1))
+4 -1
View File
@@ -1049,7 +1049,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+4 -4
View File
@@ -1362,10 +1362,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+4 -2
View File
@@ -1297,8 +1297,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+5 -5
View File
@@ -1228,10 +1228,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
@@ -1472,7 +1472,7 @@ def main():
masked_loss = masked_loss * weighting
final_loss = masked_loss.mean()
return final_loss
weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas)
loss = custom_mse_loss(noise_pred.float(), target.float(), weighting.float())
loss = loss.mean()
+35 -14
View File
@@ -76,11 +76,11 @@ from videox_fun.data.dataset_image import ImageEditDataset
from videox_fun.models import (AutoencoderKLQwenImage, Qwen2VLProcessor,
Qwen2_5_VLForConditionalGeneration,
Qwen2Tokenizer, QwenImageTransformer2DModel)
from videox_fun.pipeline import QwenImageEditPipeline
from videox_fun.pipeline import QwenImageEditPipeline, QwenImageEditPlusPipeline
from videox_fun.pipeline.pipeline_qwenimage_edit import PREFERRED_QWENIMAGE_RESOLUTIONS, calculate_dimensions
from videox_fun.pipeline.pipeline_qwenimage_edit_plus import CONDITION_IMAGE_SIZE, VAE_IMAGE_SIZE
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid, get_image
if is_wandb_available():
import wandb
@@ -150,13 +150,22 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato
subfolder="scheduler"
)
transformer3d = transformer3d.to("cpu")
pipeline = QwenImageEditPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
scheduler=scheduler,
)
if args.train_mode == "qwen_image_edit":
pipeline = QwenImageEditPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
scheduler=scheduler,
)
else:
pipeline = QwenImageEditPlusPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
scheduler=scheduler,
)
pipeline = pipeline.to(accelerator.device)
if args.seed is None:
@@ -166,12 +175,17 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, accelerato
for i in range(len(args.validation_prompts)):
with torch.no_grad():
if args.train_mode == "qwen_image_edit":
image = get_image(args.validation_image_paths[i])
else:
image = [get_image(args.validation_image_paths[i])]
sample = pipeline(
args.validation_prompts[i],
negative_prompt = "bad detailed",
height = args.image_sample_size,
width = args.image_sample_size,
generator = generator
generator = generator,
image = image
).images
os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True)
image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif"))
@@ -246,6 +260,13 @@ def parse_args():
nargs="+",
help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--validation_image_paths",
type=str,
default=None,
nargs="+",
help=("A set of images evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--output_dir",
type=str,
@@ -1250,10 +1271,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+35 -12
View File
@@ -77,7 +77,7 @@ from videox_fun.models import (AutoencoderKLQwenImage, AutoencoderKLWan,
Qwen2_5_VLForConditionalGeneration,
Qwen2Tokenizer, Qwen2VLProcessor,
QwenImageTransformer2DModel)
from videox_fun.pipeline import QwenImageEditPipeline, QwenImagePipeline
from videox_fun.pipeline import QwenImageEditPipeline, QwenImageEditPlusPipeline
from videox_fun.pipeline.pipeline_qwenimage_edit import (
PREFERRED_QWENIMAGE_RESOLUTIONS, calculate_dimensions)
from videox_fun.pipeline.pipeline_qwenimage_edit_plus import (
@@ -85,7 +85,7 @@ from videox_fun.pipeline.pipeline_qwenimage_edit_plus import (
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.lora_utils import (create_network, merge_lora,
unmerge_lora)
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid, get_image
if is_wandb_available():
import wandb
@@ -155,13 +155,22 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a
subfolder="scheduler"
)
transformer3d = transformer3d.to("cpu")
pipeline = QwenImageEditPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
scheduler=scheduler,
)
if args.train_mode == "qwen_image_edit":
pipeline = QwenImageEditPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
scheduler=scheduler,
)
else:
pipeline = QwenImageEditPlusPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
scheduler=scheduler,
)
pipeline = pipeline.to(accelerator.device)
pipeline = merge_lora(
@@ -175,12 +184,17 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, args, a
for i in range(len(args.validation_prompts)):
with torch.no_grad():
if args.train_mode == "qwen_image_edit":
image = get_image(args.validation_image_paths[i])
else:
image = [get_image(args.validation_image_paths[i])]
sample = pipeline(
args.validation_prompts[i],
negative_prompt = "bad detailed",
height = args.image_sample_size,
width = args.image_sample_size,
generator = generator
generator = generator,
image = image
).images
os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True)
image = sample[0].save(os.path.join(args.output_dir, f"sample/sample-{global_step}-image-{i}.gif"))
@@ -255,6 +269,13 @@ def parse_args():
nargs="+",
help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--validation_image_paths",
type=str,
default=None,
nargs="+",
help=("A set of images evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--output_dir",
type=str,
@@ -1202,8 +1223,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+4 -2
View File
@@ -1173,8 +1173,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+4 -4
View File
@@ -1420,10 +1420,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+4 -2
View File
@@ -1355,8 +1355,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+4 -2
View File
@@ -1054,8 +1054,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("backprop_step_list", None)
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Train!
+4 -4
View File
@@ -1417,10 +1417,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+11 -4
View File
@@ -230,6 +230,13 @@ def parse_args():
nargs="+",
help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--validation_paths",
type=str,
default=None,
nargs="+",
help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--output_dir",
type=str,
@@ -1415,10 +1422,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+11 -2
View File
@@ -234,6 +234,13 @@ def parse_args():
nargs="+",
help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--validation_paths",
type=str,
default=None,
nargs="+",
help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--output_dir",
type=str,
@@ -1360,8 +1367,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+4 -2
View File
@@ -1356,8 +1356,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+4 -2
View File
@@ -1067,8 +1067,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("backprop_step_list", None)
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Train!
+18 -6
View File
@@ -146,7 +146,8 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
with torch.no_grad():
with torch.autocast("cuda", dtype=weight_dtype):
video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size])
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size])
control_video, _, _, _ = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size])
sample = pipeline(
args.validation_prompts[i],
num_frames = video_length,
@@ -155,7 +156,11 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
width = args.video_sample_size,
generator = generator,
control_video = input_video,
video = inpaint_video,
mask_video = inpaint_video_mask,
control_video = control_video,
subject_ref_images = None,
vace_context_scale = 1,
).videos
os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True)
save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif"))
@@ -231,6 +236,13 @@ def parse_args():
nargs="+",
help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--validation_paths",
type=str,
default=None,
nargs="+",
help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--output_dir",
type=str,
@@ -1403,10 +1415,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+50 -13
View File
@@ -73,7 +73,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset,
get_random_mask)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel,
Wan2_2Transformer3DModel)
from videox_fun.pipeline import WanPipeline, WanI2VPipeline
from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
@@ -165,29 +165,66 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, ac
try:
logger.info("Running validation... ")
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
if args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
transformer3d_2_val = None
else:
if args.boundary_type == "low":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
if args.train_mode != "normal":
pipeline = WanI2VPipeline(
pipeline = Wan2_2I2VPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
else:
pipeline = WanPipeline(
pipeline = Wan2_2Pipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
pipeline = pipeline.to(accelerator.device)
@@ -1415,10 +1452,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+46 -7
View File
@@ -163,11 +163,46 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config,
try:
logger.info("Running validation... ")
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
if args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
transformer3d_2_val = None
else:
if args.boundary_type == "low":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
@@ -178,6 +213,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config,
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
else:
@@ -186,6 +222,7 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config,
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
pipeline = pipeline.to(accelerator.device)
@@ -1362,8 +1399,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+53 -7
View File
@@ -75,7 +75,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset,
get_random_mask)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel,
Wan2_2Transformer3DModel)
from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline
from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.utils import get_image_to_video_latent, save_videos_grid
@@ -157,20 +157,66 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, ac
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
if args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
transformer3d_2_val = None
else:
if args.boundary_type == "low":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
if args.train_mode != "normal":
pipeline = WanFunInpaintPipeline(
pipeline = Wan2_2I2VPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
else:
pipeline = WanFunPipeline(
pipeline = Wan2_2Pipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
pipeline = pipeline.to(accelerator.device)
@@ -1423,10 +1469,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+54 -12
View File
@@ -75,7 +75,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset,
process_pose_params)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel,
Wan2_2Transformer3DModel)
from videox_fun.pipeline import WanFunControlPipeline
from videox_fun.pipeline import Wan2_2FunControlPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.lora_utils import (create_network, merge_lora,
unmerge_lora)
@@ -152,20 +152,55 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, args, config, ac
try:
logger.info("Running validation... ")
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
if args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
transformer3d_2_val = None
else:
if args.boundary_type == "low":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
pipeline = WanFunControlPipeline(
pipeline = Wan2_2FunControlPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
pipeline = pipeline.to(accelerator.device)
@@ -265,6 +300,13 @@ def parse_args():
nargs="+",
help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--validation_paths",
type=str,
default=None,
nargs="+",
help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--output_dir",
type=str,
@@ -1484,10 +1526,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("trainable_modules")
tracker_config.pop("trainable_modules_low_learning_rate")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+1
View File
@@ -37,6 +37,7 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_control.py \
--training_with_video_token_length \
--enable_bucket \
--uniform_sampling \
--boundary_type="low" \
--train_mode="control_ref" \
--control_ref_image="random" \
--add_inpaint_info \
+54 -9
View File
@@ -75,7 +75,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset,
process_pose_params)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, CLIPModel, WanT5EncoderModel,
Wan2_2Transformer3DModel)
from videox_fun.pipeline import WanFunControlPipeline
from videox_fun.pipeline import Wan2_2FunControlPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.lora_utils import (create_network, merge_lora,
unmerge_lora)
@@ -152,20 +152,56 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config,
try:
logger.info("Running validation... ")
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
if args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
transformer3d_2_val = None
else:
if args.boundary_type == "low":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
pipeline = WanFunControlPipeline(
pipeline = Wan2_2FunControlPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
pipeline = pipeline.to(accelerator.device)
@@ -269,6 +305,13 @@ def parse_args():
nargs="+",
help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--validation_paths",
type=str,
default=None,
nargs="+",
help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--output_dir",
type=str,
@@ -1431,8 +1474,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
+49 -10
View File
@@ -71,7 +71,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoDataset,
get_random_mask)
from videox_fun.models import (AutoencoderKLWan, AutoencoderKLWan3_8, WanT5EncoderModel,
Wan2_2Transformer3DModel)
from videox_fun.pipeline import WanFunInpaintPipeline, WanFunPipeline
from videox_fun.pipeline import Wan2_2Pipeline, Wan2_2I2VPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
from videox_fun.utils.lora_utils import (create_network, merge_lora,
unmerge_lora)
@@ -146,29 +146,66 @@ def log_validation(vae, text_encoder, tokenizer, transformer3d, network, config,
try:
logger.info("Running validation... ")
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
if args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
transformer3d_2_val = None
else:
if args.boundary_type == "low":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
if args.train_mode != "normal":
pipeline = WanFunInpaintPipeline(
pipeline = Wan2_2I2VPipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
else:
pipeline = WanFunPipeline(
pipeline = Wan2_2Pipeline(
vae=accelerator.unwrap_model(vae).to(weight_dtype),
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
)
pipeline = pipeline.to(accelerator.device)
@@ -1363,8 +1400,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("fix_sample_size")
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Function for unwrapping if model was compiled with `torch.compile`.
-1
View File
@@ -38,5 +38,4 @@ accelerate launch --mixed_precision="bf16" scripts/wan2.2_fun/train_lora.py \
--train_mode="inpaint" \
--boundary_type="low" \
--lora_skip_name="ffn" \
--boundary_type="low" \
--low_vram
+4 -2
View File
@@ -1230,8 +1230,10 @@ def main():
# The trackers initializes automatically on the main process.
if accelerator.is_main_process:
tracker_config = dict(vars(args))
tracker_config.pop("validation_prompts")
tracker_config.pop("backprop_step_list", None)
keys_to_pop = [k for k, v in tracker_config.items() if isinstance(v, list)]
for k in keys_to_pop:
tracker_config.pop(k)
print(f"Removed tracker_config['{k}']")
accelerator.init_trackers(args.tracker_project_name, tracker_config)
# Train!
+67 -14
View File
@@ -74,7 +74,7 @@ from videox_fun.data.dataset_image_video import (ImageVideoControlDataset,
padding_image,
process_pose_file,
process_pose_params)
from videox_fun.models import (AutoencoderKLWan, CLIPModel,
from videox_fun.models import (AutoencoderKLWan, CLIPModel, AutoencoderKLWan3_8,
VaceWanTransformer3DModel, WanT5EncoderModel)
from videox_fun.pipeline import Wan2_2VaceFunPipeline
from videox_fun.utils.discrete_sampler import DiscreteSampling
@@ -117,11 +117,46 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
try:
logger.info("Running validation... ")
transformer3d_val = VaceWanTransformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, config['transformer_additional_kwargs'].get('transformer_subpath', 'transformer')),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
if args.boundary_type == "full":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
transformer3d_2_val = None
else:
if args.boundary_type == "low":
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
else:
sub_path = config['transformer_additional_kwargs'].get('transformer_low_noise_model_subpath', 'transformer')
transformer3d_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
sub_path = config['transformer_additional_kwargs'].get('transformer_high_noise_model_subpath', 'transformer')
transformer3d_2_val = Wan2_2Transformer3DModel.from_pretrained(
os.path.join(args.pretrained_model_name_or_path, sub_path),
transformer_additional_kwargs=OmegaConf.to_container(config['transformer_additional_kwargs']),
).to(weight_dtype)
transformer3d_2_val.load_state_dict(accelerator.unwrap_model(transformer3d).state_dict())
scheduler = FlowMatchEulerDiscreteScheduler(
**filter_kwargs(FlowMatchEulerDiscreteScheduler, OmegaConf.to_container(config['scheduler_kwargs']))
)
@@ -131,6 +166,7 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
text_encoder=accelerator.unwrap_model(text_encoder),
tokenizer=tokenizer,
transformer=transformer3d_val,
transformer_2=transformer3d_2_val,
scheduler=scheduler,
clip_image_encoder=clip_image_encoder,
)
@@ -146,7 +182,8 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
with torch.no_grad():
with torch.autocast("cuda", dtype=weight_dtype):
video_length = int(args.video_sample_n_frames // vae.config.temporal_compression_ratio * vae.config.temporal_compression_ratio) + 1 if args.video_sample_n_frames != 1 else 1
input_video, input_video_mask, ref_image, clip_image = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size])
inpaint_video, inpaint_video_mask, clip_image = get_image_to_video_latent(None, None, video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size])
control_video, _, _, _ = get_video_to_video_latent(args.validation_paths[i], video_length=video_length, sample_size=[args.video_sample_size, args.video_sample_size])
sample = pipeline(
args.validation_prompts[i],
num_frames = video_length,
@@ -155,7 +192,11 @@ def log_validation(vae, text_encoder, tokenizer, clip_image_encoder, transformer
width = args.video_sample_size,
generator = generator,
control_video = input_video,
video = inpaint_video,
mask_video = inpaint_video_mask,
control_video = control_video,
subject_ref_images = None,
vace_context_scale = 1,
).videos
os.makedirs(os.path.join(args.output_dir, "sample"), exist_ok=True)
save_videos_grid(sample, os.path.join(args.output_dir, f"sample/sample-{global_step}-{i}.gif"))
@@ -231,6 +272,13 @@ def parse_args():
nargs="+",
help=("A set of prompts evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--validation_paths",
type=str,
default=None,
nargs="+",
help=("A set of control videos evaluated every `--validation_epochs` and logged to `--report_to`."),
)
parser.add_argument(
"--output_dir",
type=str,
@@ -780,7 +828,11 @@ def main():
)
text_encoder = text_encoder.eval()
# Get Vae
vae = AutoencoderKLWan.from_pretrained(
Chosen_AutoencoderKL = {
"AutoencoderKLWan": AutoencoderKLWan,
"AutoencoderKLWan3_8": AutoencoderKLWan3_8
}[config['vae_kwargs'].get('vae_type', 'AutoencoderKLWan')]
vae = Chosen_AutoencoderKL.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']),
)
@@ -1023,7 +1075,8 @@ def main():
# Get the training dataset
sample_n_frames_bucket_interval = vae.config.temporal_compression_ratio
spatial_compression_ratio = vae.config.spatial_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)
@@ -1186,7 +1239,7 @@ def main():
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()}
if args.fix_sample_size is not None:
fix_sample_size = [int(x / 16) * 16 for x in args.fix_sample_size]
fix_sample_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 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[
@@ -1196,10 +1249,10 @@ def main():
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]
random_sample_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 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]
closest_size = [int(x / spatial_compression_ratio / 2) * spatial_compression_ratio * 2 for x in closest_size]
for example in examples:
# To 0~1
@@ -1808,7 +1861,7 @@ def main():
vace_latents = vace_encode_frames(control_pixel_values, subject_ref_images, mask)
mask = torch.ones_like(mask)
mask_latents = vace_encode_masks(mask, subject_ref_images)
mask_latents = vace_encode_masks(mask, subject_ref_images, vae_stride=[4, spatial_compression_ratio, spatial_compression_ratio])
vace_context = torch.stack(vace_latent(vace_latents, mask_latents))
if subject_ref_images is not None:
+28 -6
View File
@@ -33,8 +33,8 @@ from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import (AdaLayerNormContinuous,
AdaLayerNormZero,
AdaLayerNormZeroSingle)
from diffusers.utils import (USE_PEFT_BACKEND, logging, scale_lora_layers,
unscale_lora_layers)
from diffusers.utils import (USE_PEFT_BACKEND, is_torch_version, logging,
scale_lora_layers, unscale_lora_layers)
from diffusers.utils.torch_utils import maybe_allow_in_graph
from ..dist import (FluxMultiGPUsAttnProcessor2_0, get_sequence_parallel_rank,
@@ -695,6 +695,14 @@ class FluxTransformer2DModel(
self.sp_world_size = 1
self.sp_world_rank = 0
def _set_gradient_checkpointing(self, *args, **kwargs):
if "value" in kwargs:
self.gradient_checkpointing = kwargs["value"]
elif "enable" in kwargs:
self.gradient_checkpointing = kwargs["enable"]
else:
raise ValueError("Invalid set gradient checkpointing")
def enable_multi_gpus_inference(self,):
self.sp_world_size = get_sequence_parallel_world_size()
self.sp_world_rank = get_sequence_parallel_rank()
@@ -868,13 +876,20 @@ class FluxTransformer2DModel(
for index_block, block in enumerate(self.transformer_blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
encoder_hidden_states, hidden_states = self._gradient_checkpointing_func(
block,
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
encoder_hidden_states,
temb,
image_rotary_emb,
joint_attention_kwargs,
**ckpt_kwargs,
)
else:
@@ -900,13 +915,20 @@ class FluxTransformer2DModel(
for index_block, block in enumerate(self.single_transformer_blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
encoder_hidden_states, hidden_states = self._gradient_checkpointing_func(
block,
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
encoder_hidden_states,
temb,
image_rotary_emb,
joint_attention_kwargs,
**ckpt_kwargs,
)
else: