Fix uni3c on fp16 and add strength parameters

This commit is contained in:
kijai
2025-05-29 13:25:52 +03:00
parent 87ae18e203
commit 129f368380
3 changed files with 27 additions and 16 deletions
+3
View File
@@ -2738,6 +2738,9 @@ class WanVideoSampler:
"render_latent": uni3c_embeds["render_latent"],
"render_mask": uni3c_embeds["render_mask"],
"camera_embedding": uni3c_embeds["camera_embedding"],
"controlnet_weight": uni3c_embeds["controlnet_weight"],
"start": uni3c_embeds["start"],
"end": uni3c_embeds["end"],
}
#feta
+9 -5
View File
@@ -137,12 +137,12 @@ class WanVideoUni3C_embeds:
return {"required": {
"controlnet": ("WANVIDEOCONTROLNET",),
"render_latent": ("LATENT",),
# "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
# "vace_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply VACE"}),
# "vace_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply VACE"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply the controlnet"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply the controlnet"}),
},
"optional": {
"render_mask": ("MASK",),
"render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}),
},
}
@@ -151,7 +151,7 @@ class WanVideoUni3C_embeds:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, controlnet, render_latent, render_mask=None):
def process(self, controlnet, render_latent, strength, start_percent, end_percent, render_mask=None):
device = mm.get_torch_device()
@@ -163,6 +163,7 @@ class WanVideoUni3C_embeds:
width = latents.shape[4] * 8
if render_mask is not None:
raise NotImplementedError("render_mask is not implemented at this time")
mask = torch.nn.functional.interpolate(
render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(nframe, height, width),
@@ -218,6 +219,9 @@ class WanVideoUni3C_embeds:
uni3c_embeds = {
"controlnet": controlnet,
"controlnet_weight": strength,
"start": start_percent,
"end": end_percent,
"render_latent": latents.to(device),
"render_mask": latent_mask,
"camera_embedding": None
+15 -11
View File
@@ -1423,17 +1423,21 @@ class WanModel(ModelMixin, ConfigMixin):
kwargs['vace_context_scale'] = vace_scale_list
#uni3c controlnet
pdc_controlnet_states = None
if pcd_data is not None:
self.controlnet.to(self.main_device)
pdc_controlnet_states = self.controlnet(
render_latent=render_latent.to(self.main_device),
render_mask=pcd_data["render_mask"],
camera_embedding=pcd_data["camera_embedding"],
temb=e.to(self.main_device),
device=self.offload_device)
self.controlnet.to(self.offload_device)
if (pcd_data["start"] <= current_step_percentage <= pcd_data["end"]) or \
(pcd_data["end"] > 0 and current_step == 0 and current_step_percentage >= pcd_data["start"]):
self.controlnet.to(self.main_device)
pdc_controlnet_states = self.controlnet(
render_latent=render_latent.to(self.main_device, self.controlnet.dtype),
render_mask=pcd_data["render_mask"],
camera_embedding=pcd_data["camera_embedding"],
temb=e.to(self.main_device),
device=self.offload_device)
self.controlnet.to(self.offload_device)
for b, block in enumerate(self.blocks):
#skip layer guidance
if self.slg_blocks is not None:
if b in self.slg_blocks and is_uncond:
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
@@ -1443,9 +1447,9 @@ class WanModel(ModelMixin, ConfigMixin):
x = block(x, **kwargs)
#uni3c controlnet
if pcd_data is not None:
if b < len(pdc_controlnet_states):
x += pdc_controlnet_states[b].to(x.device)
if pdc_controlnet_states is not None and b < len(pdc_controlnet_states):
x += pdc_controlnet_states[b].to(x) * pcd_data["controlnet_weight"]
#controlnet
if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])):
x += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]] * controlnet["controlnet_weight"]