Properly set video length in SparseCtrl motion module, account for potential batched_number changes

This commit is contained in:
Jedrzej Kosinski
2023-12-19 10:59:45 -06:00
parent 3bf77b1f82
commit 2757a4f17e
2 changed files with 19 additions and 9 deletions
+9 -5
View File
@@ -245,14 +245,18 @@ class SparseCtrlAdvanced(ControlNetAdvanced):
if self.manual_cast_dtype is not None:
dtype = self.manual_cast_dtype
output_dtype = x_noisy.dtype
# set actual input length on motion model
actual_length = x_noisy.size(0)//batched_number
full_length = actual_length if self.sub_idxs is None else self.full_latent_length
self.control_model.set_actual_length(actual_length=actual_length, full_length=full_length)
# prepare cond_hint, if needed
if self.sub_idxs is not None or self.cond_hint is None:
dim_mult = 1 if self.control_model.use_simplified_conditioning_embedding else 8
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[0] != self.cond_hint.shape[0] or x_noisy.shape[2]*dim_mult != self.cond_hint.shape[2] or x_noisy.shape[3]*dim_mult != self.cond_hint.shape[3]:
# clear out cond_hint and conditioning_mask
if self.cond_hint is not None:
del self.cond_hint
self.cond_hint = None
# first, figure out which cond idxs are relevant, and where they fit in
full_length = x_noisy.size(0)//batched_number if self.sub_idxs is None else self.full_latent_length
cond_idxs = self.sparse_settings.sparse_method.get_indeces(hint_length=self.cond_hint_original.size(0), full_length=full_length)
range_idxs = list(range(full_length)) if self.sub_idxs is None else self.sub_idxs
@@ -285,9 +289,9 @@ class SparseCtrlAdvanced(ControlNetAdvanced):
self.cond_hint = torch.cat([self.cond_hint, cond_mask], dim=1)
del sub_cond_hint
del cond_mask
# make cond_hint match x_noisy batch
if x_noisy.shape[0] != self.cond_hint.shape[0]:
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
# make cond_hint match x_noisy batch
if x_noisy.shape[0] != self.cond_hint.shape[0]:
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
# prepare mask_cond_hint
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype)
+10 -4
View File
@@ -43,6 +43,10 @@ class SparseControlNet(ControlNetCLDM):
)
self.motion_holder: MotionWrapperHolder = None
def set_actual_length(self, actual_length: int, full_length: int):
if self.motion_holder is not None:
self.motion_holder.motion_wrapper.set_video_length(video_length=actual_length, full_length=full_length)
def forward(self, x: Tensor, hint: Tensor, timesteps, context, y=None, **kwargs):
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype)
emb = self.time_embed(t_emb)
@@ -256,10 +260,12 @@ class SparseCtrlMotionWrapper(nn.Module):
def set_video_length(self, video_length: int, full_length: int):
self.AD_video_length = video_length
for block in self.down_blocks:
block.set_video_length(video_length, full_length)
for block in self.up_blocks:
block.set_video_length(video_length, full_length)
if self.down_blocks is not None:
for block in self.down_blocks:
block.set_video_length(video_length, full_length)
if self.up_blocks is not None:
for block in self.up_blocks:
block.set_video_length(video_length, full_length)
if self.mid_block is not None:
self.mid_block.set_video_length(video_length, full_length)