From f3a31815ca93d8206b5594f488f58275a92765d5 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 1 Nov 2023 19:31:06 -0500 Subject: [PATCH] Hopefully decreased RAM usage for large batches of cond_hint images by applying scaling to only sub_idx cond hints, when possible --- control/control.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/control/control.py b/control/control.py index 18fff63..c851fd2 100644 --- a/control/control.py +++ b/control/control.py @@ -204,10 +204,11 @@ class ControlNetAdvanced(ControlNet): if self.cond_hint is not None: del self.cond_hint self.cond_hint = None - self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device) - # if self.cond_hint length matches real latent count, need to subdivide it - if self.cond_hint.size(0) == self.full_latent_length: - self.cond_hint = self.cond_hint[self.sub_idxs] + # if self.cond_hint_original length matches real latent count, need to subdivide it + if self.cond_hint_original.size(0) == self.full_latent_length: + self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device) + else: + self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(self.control_model.dtype).to(self.device) 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) @@ -217,6 +218,7 @@ class ControlNetAdvanced(ControlNet): if self.mask_cond_hint is not None: del self.mask_cond_hint self.mask_cond_hint = None + # TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM # resize mask and match batch count self.mask_cond_hint = prepare_mask_batch(self.mask_cond_hint_original, x_noisy.shape, multiplier=8) actual_latent_length = x_noisy.shape[0] // batched_number