Support control-lora

https: //huggingface.co/spacepxl/Wan2.1-control-loras/blob/main/wan2.1-1.3b-control-lora-tile-v0.1_comfy.safetensors
Co-Authored-By: spacepxl <143970342+spacepxl@users.noreply.github.com>
This commit is contained in:
kijai
2025-03-10 21:31:06 +02:00
co-authored by spacepxl
parent 9e83063dcf
commit ffaf1afe4e
2 changed files with 109 additions and 5 deletions
+98 -3
View File
@@ -25,6 +25,7 @@ from comfy.utils import load_torch_file, save_torch_file, ProgressBar, common_up
import comfy.model_base
import comfy.latent_formats
from comfy.clip_vision import clip_preprocess, ClipVisionModel
from comfy.sd import load_lora_for_models
script_directory = os.path.dirname(os.path.abspath(__file__))
@@ -432,9 +433,10 @@ class WanVideoModelLoader:
comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
patcher.is_patched = False
if lora is not None:
from comfy.sd import load_lora_for_models
for l in lora:
log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}")
lora_path = l["path"]
@@ -444,10 +446,41 @@ class WanVideoModelLoader:
if l["blocks"]:
lora_sd = filter_state_dict_by_blocks(lora_sd, l["blocks"])
#spacepxl's control LoRA patch
# for key in lora_sd.keys():
# print(key)
if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd:
log.info("Control-LoRA detected, patching model...")
in_cls = transformer.patch_embedding.__class__ # nn.Conv3d
old_in_dim = transformer.in_dim # 16
new_in_dim = lora_sd["diffusion_model.patch_embedding.lora_A.weight"].shape[1]
assert new_in_dim == 32
new_in = in_cls(
new_in_dim,
transformer.patch_embedding.out_channels,
transformer.patch_embedding.kernel_size,
transformer.patch_embedding.stride,
transformer.patch_embedding.padding,
).to(device=device, dtype=torch.bfloat16)
new_in.weight.zero_()
new_in.bias.zero_()
new_in.weight[:, :old_in_dim].copy_(transformer.patch_embedding.weight)
new_in.bias.copy_(transformer.patch_embedding.bias)
transformer.patch_embedding = new_in
transformer.expanded_patch_embedding = new_in
transformer.register_to_config(in_dim=new_in_dim)
patcher, _ = load_lora_for_models(patcher, None, lora_sd, lora_strength, 0)
del lora_sd
patcher.patch_model(device)
patcher.is_patched = True
del sd
gc.collect()
@@ -992,6 +1025,40 @@ class WanVideoEmptyEmbeds:
}
return (embeds,)
class WanVideoControlEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"latents": ("LATENT", {"tooltip": "Encoded latents to use as control signals"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the control signal"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the control signal"}),
},
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", )
RETURN_NAMES = ("image_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, latents, start_percent, end_percent):
samples = latents["samples"].squeeze(0)
C, T, H, W = samples.shape
num_frames = (T - 1) * 4 + 1
seq_len = math.ceil((H * W) / 4 * ((num_frames - 1) // 4 + 1))
embeds = {
"max_seq_len": seq_len,
"target_shape": samples.shape,
"num_frames": num_frames,
"control_images": samples,
"start_percent": start_percent,
"end_percent": end_percent,
}
return (embeds,)
#region Sampler
@@ -1176,6 +1243,17 @@ class WanVideoSampler:
device=torch.device("cpu"),
generator=seed_g)
control_latents = image_embeds.get("control_images", None)
if control_latents is not None:
image_cond = control_latents.to(device)
control_start_percent = image_embeds.get("start_percent", 0.0)
control_end_percent = image_embeds.get("end_percent", 1.0)
if not patcher.is_patched:
print("Patching model for control")
patcher.patch_model(device)
patcher.is_patched = True
latent_video_length = noise.shape[1]
if context_options is not None:
@@ -1373,6 +1451,17 @@ class WanVideoSampler:
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, teacache_state=None):
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True):
current_step_percentage = idx / len(timesteps)
if not control_start_percent <= current_step_percentage <= control_end_percent:
image_cond = None
control_enabled = False
if patcher.is_patched:
patcher.unpatch_model(device)
patcher.is_patched = False
else:
control_enabled = True
base_params = {
'clip_fea': clip_fea,
'seq_len': seq_len,
@@ -1381,6 +1470,7 @@ class WanVideoSampler:
't': timestep,
'current_step': idx,
'y': image_cond,
'control_enabled': control_enabled,
}
# Get conditional prediction
@@ -1589,7 +1679,10 @@ class WanVideoSampler:
partial_img_emb = None
if image_cond is not None:
partial_img_emb = image_cond[:, c, :, :]
partial_img_emb[:, 0, :, :] = image_cond[:, 0, :, :].to(intermediate_device)
partial_image_cond = image_cond[:, 0, :, :].to(intermediate_device)
if min(c) > len(c) // 2:
partial_image_cond *= 0.1
partial_img_emb[:, 0, :, :] = partial_image_cond
partial_latent_model_input = latent_model_input[:, c, :, :]
@@ -1801,7 +1894,7 @@ class WanVideoEncode:
mask = torch.nn.functional.interpolate(
mask.unsqueeze(1), # Add channel dim for interpolate
size=(target_h, target_w),
mode='nearest'
mode='bilinear'
).squeeze(1) # Remove channel dim
# Add batch & channel dims for final output
@@ -1913,6 +2006,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoVRAMManagement": WanVideoVRAMManagement,
"WanVideoTextEmbedBridge": WanVideoTextEmbedBridge,
"WanVideoFlowEdit": WanVideoFlowEdit,
"WanVideoControlEmbeds": WanVideoControlEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
@@ -1937,4 +2031,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoVRAMManagement": "WanVideo VRAM Management",
"WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge",
"WanVideoFlowEdit": "WanVideo FlowEdit",
"WanVideoControlEmbeds": "WanVideo Control Embeds",
}
+11 -2
View File
@@ -520,6 +520,10 @@ class WanModel(ModelMixin, ConfigMixin):
# embeddings
self.patch_embedding = nn.Conv3d(
in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.original_patch_embedding = self.patch_embedding
self.expanded_patch_embedding = self.patch_embedding
self.text_embedding = nn.Sequential(
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
nn.Linear(dim, dim))
@@ -591,7 +595,8 @@ class WanModel(ModelMixin, ConfigMixin):
device=torch.device('cuda'),
freqs=None,
current_step=0,
pred_id=None
pred_id=None,
control_enabled=False,
):
r"""
Forward pass through the diffusion model
@@ -625,7 +630,11 @@ class WanModel(ModelMixin, ConfigMixin):
x = torch.cat([x, y], dim=0)
# embeddings
x = [self.patch_embedding(x.unsqueeze(0))]
if control_enabled:
x = [self.expanded_patch_embedding(x.unsqueeze(0))]
else:
x = [self.original_patch_embedding(x.unsqueeze(0))]
grid_sizes = torch.stack(
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
x = [u.flatten(2).transpose(1, 2) for u in x]