diff --git a/controlnet/nodes.py b/controlnet/nodes.py index 14975d1..28235f6 100644 --- a/controlnet/nodes.py +++ b/controlnet/nodes.py @@ -41,9 +41,11 @@ class WanVideoControlnetLoader: model_path = folder_paths.get_full_path_or_raise("controlnet", model) sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True) - + num_layers = 8 if "blocks.7.scale_shift_table" in sd else 6 - out_proj_dim = 5120 if num_layers == 6 else 1536 + out_proj_dim = sd["controlnet_blocks.0.bias"].shape[0] + downscale_coef = 16 if out_proj_dim == 3072 else 8 + vae_channels = 48 if out_proj_dim == 3072 else 16 if not "control_encoder.0.0.weight" in sd: raise ValueError("Invalid ControlNet model") @@ -52,7 +54,7 @@ class WanVideoControlnetLoader: "added_kv_proj_dim": None, "attention_head_dim": 128, "cross_attn_norm": None, - "downscale_coef": 8, + "downscale_coef": downscale_coef, "eps": 1e-06, "ffn_dim": 8960, "freq_dim": 256, @@ -69,8 +71,9 @@ class WanVideoControlnetLoader: "qk_norm": "rms_norm_across_heads", "rope_max_seq_len": 1024, "text_dim": 4096, - "vae_channels": 16 + "vae_channels": vae_channels } + print(f"Loading WanControlnet with config: {controlnet_cfg}") from .wan_controlnet import WanControlnet diff --git a/controlnet/wan_controlnet.py b/controlnet/wan_controlnet.py index af4c71e..f7b90d2 100644 --- a/controlnet/wan_controlnet.py +++ b/controlnet/wan_controlnet.py @@ -24,6 +24,12 @@ def zero_module(module): logger = logging.get_logger(__name__) # pylint: disable=invalid-name +def zero_module(module): + for p in module.parameters(): + nn.init.zeros_(p) + return module + + class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): r""" A Controlnet Transformer model for video-like data used in the Wan model. @@ -110,7 +116,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel nn.GELU(approximate="tanh"), nn.GroupNorm(2, input_channels[0]), ), - ## Temporal compression with spatial awareness + ## Spatio-Temporal compression with spatial awareness nn.Sequential( nn.Conv3d(input_channels[0], input_channels[1], kernel_size=3, stride=(2, 1, 1), padding=1), nn.GELU(approximate="tanh"), @@ -196,15 +202,27 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel hidden_states = self.patch_embedding(hidden_states) hidden_states = hidden_states.flatten(2).transpose(1, 2) + # timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v) + if timestep.ndim == 2: + ts_seq_len = timestep.shape[1] + timestep = timestep.flatten() # batch_size * seq_len + else: + ts_seq_len = None + temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder( - timestep, encoder_hidden_states, encoder_hidden_states_image + timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len ) - timestep_proj = timestep_proj.unflatten(1, (6, -1)) + if ts_seq_len is not None: + # batch_size, seq_len, 6, inner_dim + timestep_proj = timestep_proj.unflatten(2, (6, -1)) + else: + # batch_size, 6, inner_dim + timestep_proj = timestep_proj.unflatten(1, (6, -1)) if encoder_hidden_states_image is not None: encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) - # 2. Transformer blocks + # 4. Transformer blocks controlnet_hidden_states = () if torch.is_grad_enabled() and self.gradient_checkpointing: for block, controlnet_block in zip(self.blocks, self.controlnet_blocks): @@ -246,13 +264,14 @@ if __name__ == "__main__": "text_dim": 4096, "downscale_coef": 8, "out_proj_dim": 12 * 128, + "vae_channels": 16 } controlnet = WanControlnet(**parameters) - hidden_states = torch.rand(1, 16, 21, 60, 90) - timestep = torch.randint(low=0, high=1000, size=(1,), dtype=torch.long) + hidden_states = torch.rand(1, 16, 13, 60, 90) + timestep = torch.tensor([1000]).repeat(17550).unsqueeze(0) #torch.randint(low=0, high=1000, size=(1,), dtype=torch.long) encoder_hidden_states = torch.rand(1, 512, 4096) - controlnet_states = torch.rand(1, 3, 81, 480, 720) + controlnet_states = torch.rand(1, 3, 49, 480, 720) controlnet_hidden_states = controlnet( hidden_states=hidden_states,