Fix uni3c on fp16 and add strength parameters
This commit is contained in:
@@ -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
@@ -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
@@ -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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user