controlnet device fix
This commit is contained in:
@@ -2907,12 +2907,13 @@ class WanVideoSampler:
|
||||
|
||||
if controlnet_latents is not None:
|
||||
if (controlnet_start <= current_step_percentage < controlnet_end):
|
||||
self.controlnet.to(device)
|
||||
controlnet_states = self.controlnet(
|
||||
hidden_states=latent_model_input.unsqueeze(0).to(self.controlnet.dtype),
|
||||
hidden_states=latent_model_input.unsqueeze(0).to(device, self.controlnet.dtype),
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=positive_embeds[0].unsqueeze(0).to(self.controlnet.dtype),
|
||||
encoder_hidden_states=positive_embeds[0].unsqueeze(0).to(device, self.controlnet.dtype),
|
||||
attention_kwargs=None,
|
||||
controlnet_states=controlnet_latents.to(self.controlnet.dtype).to(self.controlnet.device),
|
||||
controlnet_states=controlnet_latents.to(device, self.controlnet.dtype),
|
||||
return_dict=False,
|
||||
)[0]
|
||||
if isinstance(controlnet_states, (tuple, list)):
|
||||
|
||||
@@ -1451,7 +1451,7 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
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"]
|
||||
x += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
|
||||
|
||||
if b <= self.blocks_to_swap and self.blocks_to_swap >= 0:
|
||||
block.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
Reference in New Issue
Block a user