Allow loading 5B controlnet
inference doesn't work yet
This commit is contained in:
+7
-4
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user