From ac5a3bde6a532ccb172c6535fe78ead1047a35ac Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 24 Nov 2024 15:40:06 -0600 Subject: [PATCH] Add changes manually from 'master' so merge conflict resolution goes more smoothly --- comfy/model_patcher.py | 79 ++++++++++++++++++++++++------------------ 1 file changed, 46 insertions(+), 33 deletions(-) diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index c2c61d432..643fdfa4e 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -523,19 +523,28 @@ class ModelPatcher: else: set_func(out_weight, inplace_update=inplace_update, seed=string_to_seed(key)) + def _load_list(self): + loading = [] + for n, m in self.model.named_modules(): + params = [] + skip = False + for name, param in m.named_parameters(recurse=False): + params.append(name) + for name, param in m.named_parameters(recurse=True): + if name not in params: + skip = True # skip random weights in non leaf modules + break + if not skip and (hasattr(m, "comfy_cast_weights") or len(params) > 0): + loading.append((comfy.model_management.module_size(m), n, m, params)) + return loading + def load(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, full_load=False): with self.use_ejected(): self.unpatch_hooks() mem_counter = 0 patch_counter = 0 lowvram_counter = 0 - loading = [] - for n, m in self.model.named_modules(): - params = [] - for name, param in m.named_parameters(recurse=False): - params.append(name) - if hasattr(m, "comfy_cast_weights") or len(params) > 0: - loading.append((comfy.model_management.module_size(m), n, m, params)) + loading = self._load_list() load_completely = [] loading.sort(reverse=True) @@ -578,8 +587,9 @@ class ModelPatcher: if m.comfy_cast_weights: wipe_lowvram_weight(m) - mem_counter += module_mem - load_completely.append((module_mem, n, m, params)) + if full_load or mem_counter + module_mem < lowvram_model_memory: + mem_counter += module_mem + load_completely.append((module_mem, n, m, params)) load_completely.sort(reverse=True) for x in load_completely: @@ -678,14 +688,7 @@ class ModelPatcher: with self.use_ejected(): memory_freed = 0 patch_counter = 0 - unload_list = [] - - for n, m in self.model.named_modules(): - shift_lowvram = False - if hasattr(m, "comfy_cast_weights"): - module_mem = comfy.model_management.module_size(m) - unload_list.append((module_mem, n, m)) - + unload_list = self._load_list() unload_list.sort() for unload in unload_list: if memory_to_free < memory_freed: @@ -693,32 +696,42 @@ class ModelPatcher: module_mem = unload[0] n = unload[1] m = unload[2] - weight_key = "{}.weight".format(n) - bias_key = "{}.bias".format(n) + params = unload[3] + lowvram_possible = hasattr(m, "comfy_cast_weights") if hasattr(m, "comfy_patched_weights") and m.comfy_patched_weights == True: - for key in [weight_key, bias_key]: + move_weight = True + for param in params: + key = "{}.{}".format(n, param) bk = self.backup.get(key, None) if bk is not None: + if not lowvram_possible: + move_weight = False + break + if bk.inplace_update: comfy.utils.copy_to_param(self.model, key, bk.weight) else: comfy.utils.set_attr_param(self.model, key, bk.weight) self.backup.pop(key) + + weight_key = "{}.weight".format(n) + bias_key = "{}.bias".format(n) + if move_weight: + m.to(device_to) + if lowvram_possible: + if weight_key in self.patches: + m.weight_function = LowVramPatch(weight_key, self.patches) + patch_counter += 1 + if bias_key in self.patches: + m.bias_function = LowVramPatch(bias_key, self.patches) + patch_counter += 1 - m.to(device_to) - if weight_key in self.patches: - m.weight_function = LowVramPatch(weight_key, self.patches) - patch_counter += 1 - if bias_key in self.patches: - m.bias_function = LowVramPatch(bias_key, self.patches) - patch_counter += 1 - - m.prev_comfy_cast_weights = m.comfy_cast_weights - m.comfy_cast_weights = True - m.comfy_patched_weights = False - memory_freed += module_mem - logging.debug("freed {}".format(n)) + m.prev_comfy_cast_weights = m.comfy_cast_weights + m.comfy_cast_weights = True + m.comfy_patched_weights = False + memory_freed += module_mem + logging.debug("freed {}".format(n)) self.model.model_lowvram = True self.model.lowvram_patch_counter += patch_counter