diff --git a/video-generation/opensora/model_variants/wanx_diffusers_src/transformer_wan_refextractor_2d_controlnet_last_prefix.py b/video-generation/opensora/model_variants/wanx_diffusers_src/transformer_wan_refextractor_2d_controlnet_last_prefix.py index 2f2e9e7..9c08f5f 100644 --- a/video-generation/opensora/model_variants/wanx_diffusers_src/transformer_wan_refextractor_2d_controlnet_last_prefix.py +++ b/video-generation/opensora/model_variants/wanx_diffusers_src/transformer_wan_refextractor_2d_controlnet_last_prefix.py @@ -550,7 +550,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOrigi block_out_channels=(16, 16, 16, 16), norm_num_groups=4, layers_per_block=1, - spatial_compression_ratio=16 + spatial_compression_ratio=cfg.get("conditioning_embedding_spatial_compression_ratio", 16) ) self.controlnet = nn.Module() @@ -666,6 +666,15 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOrigi controlnet_cond_prefix = self.input_hint_block(image_pose) controlnet_cond = torch.concat([controlnet_cond_prefix,controlnet_cond],dim=2) + target_size = (post_patch_num_frames, post_patch_height, post_patch_width) + if tuple(controlnet_cond.shape[2:]) != target_size: + controlnet_cond = F.interpolate( + controlnet_cond.float(), + size=target_size, + mode="trilinear", + align_corners=False, + ).to(dtype=hidden_states.dtype) + controlnet_tokens = controlnet_cond.flatten(2).transpose(1, 2) diff --git a/video-generation/opensora/train/diffusion.py b/video-generation/opensora/train/diffusion.py index 110ea67..ef80ee9 100644 --- a/video-generation/opensora/train/diffusion.py +++ b/video-generation/opensora/train/diffusion.py @@ -769,6 +769,12 @@ class FlowMatching(): final_mask = loss_mask.float() * ocr_mask.float() + if model_pred.shape != target.shape: + raise RuntimeError( + f"model_pred shape {tuple(model_pred.shape)} does not match target shape {tuple(target.shape)}. " + "For Wan-style spatial patching, train image height/width must make latent H/W divisible by the transformer patch size." + ) + # auto broadcast loss = (weighting.float() * (model_pred.float() - target.float()) ** 2 * final_mask) diff --git a/video-generation/opensora/train/wanx_train/train_wanx_refextractor_mask2_controlnet2.py b/video-generation/opensora/train/wanx_train/train_wanx_refextractor_mask2_controlnet2.py index 8e80d67..3028720 100644 --- a/video-generation/opensora/train/wanx_train/train_wanx_refextractor_mask2_controlnet2.py +++ b/video-generation/opensora/train/wanx_train/train_wanx_refextractor_mask2_controlnet2.py @@ -200,9 +200,11 @@ def custom_models(args, weight_dtype): vae = get_vae(args.vae_name, args.vae_path, weight_dtype) vae.eval() - # vae.enable_tiling() + # Wan2.2 TI2V VAE has patch_size=2; diffusers tiled encode bypasses the RGB->patch latent path. try: - vae.enable_tiling() + inner_vae = getattr(vae, "vae", vae) + if getattr(getattr(inner_vae, "config", None), "patch_size", None) is None: + vae.enable_tiling() vae.enable_slicing() except Exception as e: pass @@ -550,11 +552,18 @@ def main(args): if name == special_key: pretrained_weight = params current_weight = model.refextractor.state_dict()[special_key] - if pretrained_weight.shape[1] == 16 and current_weight.shape[1] == 32: + if (pretrained_weight.ndim == current_weight.ndim + and pretrained_weight.shape[0] == current_weight.shape[0] + and pretrained_weight.shape[2:] == current_weight.shape[2:] + and pretrained_weight.shape[1] < current_weight.shape[1]): + extra_channels = current_weight.shape[1] - pretrained_weight.shape[1] if accelerator.is_main_process: - accelerator.print(f"Special handling for '{special_key}' with torch.cat: extending {pretrained_weight.shape[1]} to {current_weight.shape[1]}.") + accelerator.print( + f"Special handling for '{special_key}': extending " + f"{pretrained_weight.shape[1]} to {current_weight.shape[1]} input channels." + ) new_channel_shape = list(pretrained_weight.shape) - new_channel_shape[1] = 16 + new_channel_shape[1] = extra_channels zero_channels = pretrained_weight.new_zeros(new_channel_shape) adjusted_weight = torch.cat([pretrained_weight, zero_channels], dim=1) refextractor_state_dict[special_key] = adjusted_weight diff --git a/video-generation/opensora/utils/bucket.py b/video-generation/opensora/utils/bucket.py index 1444721..d6c15e4 100644 --- a/video-generation/opensora/utils/bucket.py +++ b/video-generation/opensora/utils/bucket.py @@ -1,4 +1,5 @@ from collections import OrderedDict +import os import numpy as np from .aspect import ASPECT_RATIOS, get_closest_ratio @@ -235,6 +236,18 @@ valid_bucket_configs = { "960_full": bucket_config_960_full, } + +def _align_down_to_multiple(value, multiple): + value = int(value) + multiple = int(multiple) + if multiple <= 1: + return value + return max(multiple, (value // multiple) * multiple) + +def _align_hw_from_env(height, width): + multiple = int(os.environ.get("OPENSORA_SPATIAL_ALIGN_MULTIPLE", "1")) + return _align_down_to_multiple(height, multiple), _align_down_to_multiple(width, multiple) + def find_approximate_hw(hw, hw_dict, approx=0.8): for k, v in hw_dict.items(): if hw >= v * approx: @@ -360,6 +373,7 @@ class Bucket: assert len(bucket_id) == 3 T = self.t_criteria[bucket_id[0]][bucket_id[1]] H, W = self.ar_criteria[bucket_id[0]][bucket_id[1]][bucket_id[2]] + H, W = _align_hw_from_env(H, W) return T, H, W def get_prob(self, bucket_id):