Allow loading 5B controlnet

inference doesn't work yet
This commit is contained in:
kijai
2025-08-03 12:17:10 +03:00
parent 309a93c221
commit 26c1911e41
2 changed files with 33 additions and 11 deletions
+7 -4
View File
@@ -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
+26 -7
View File
@@ -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,