Fix bugs in the training code (#356)
This commit is contained in:
@@ -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))
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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!
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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!
|
||||
|
||||
@@ -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
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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 \
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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!
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user