mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-10-08 01:27:05 +08:00
Add changes manually from 'master' so merge conflict resolution goes more smoothly
This commit is contained in:
parent
602c12b515
commit
ac5a3bde6a
@ -523,19 +523,28 @@ class ModelPatcher:
|
|||||||
else:
|
else:
|
||||||
set_func(out_weight, inplace_update=inplace_update, seed=string_to_seed(key))
|
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):
|
def load(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, full_load=False):
|
||||||
with self.use_ejected():
|
with self.use_ejected():
|
||||||
self.unpatch_hooks()
|
self.unpatch_hooks()
|
||||||
mem_counter = 0
|
mem_counter = 0
|
||||||
patch_counter = 0
|
patch_counter = 0
|
||||||
lowvram_counter = 0
|
lowvram_counter = 0
|
||||||
loading = []
|
loading = self._load_list()
|
||||||
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))
|
|
||||||
|
|
||||||
load_completely = []
|
load_completely = []
|
||||||
loading.sort(reverse=True)
|
loading.sort(reverse=True)
|
||||||
@ -578,8 +587,9 @@ class ModelPatcher:
|
|||||||
if m.comfy_cast_weights:
|
if m.comfy_cast_weights:
|
||||||
wipe_lowvram_weight(m)
|
wipe_lowvram_weight(m)
|
||||||
|
|
||||||
mem_counter += module_mem
|
if full_load or mem_counter + module_mem < lowvram_model_memory:
|
||||||
load_completely.append((module_mem, n, m, params))
|
mem_counter += module_mem
|
||||||
|
load_completely.append((module_mem, n, m, params))
|
||||||
|
|
||||||
load_completely.sort(reverse=True)
|
load_completely.sort(reverse=True)
|
||||||
for x in load_completely:
|
for x in load_completely:
|
||||||
@ -678,14 +688,7 @@ class ModelPatcher:
|
|||||||
with self.use_ejected():
|
with self.use_ejected():
|
||||||
memory_freed = 0
|
memory_freed = 0
|
||||||
patch_counter = 0
|
patch_counter = 0
|
||||||
unload_list = []
|
unload_list = self._load_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.sort()
|
unload_list.sort()
|
||||||
for unload in unload_list:
|
for unload in unload_list:
|
||||||
if memory_to_free < memory_freed:
|
if memory_to_free < memory_freed:
|
||||||
@ -693,32 +696,42 @@ class ModelPatcher:
|
|||||||
module_mem = unload[0]
|
module_mem = unload[0]
|
||||||
n = unload[1]
|
n = unload[1]
|
||||||
m = unload[2]
|
m = unload[2]
|
||||||
weight_key = "{}.weight".format(n)
|
params = unload[3]
|
||||||
bias_key = "{}.bias".format(n)
|
|
||||||
|
|
||||||
|
lowvram_possible = hasattr(m, "comfy_cast_weights")
|
||||||
if hasattr(m, "comfy_patched_weights") and m.comfy_patched_weights == True:
|
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)
|
bk = self.backup.get(key, None)
|
||||||
if bk is not None:
|
if bk is not None:
|
||||||
|
if not lowvram_possible:
|
||||||
|
move_weight = False
|
||||||
|
break
|
||||||
|
|
||||||
if bk.inplace_update:
|
if bk.inplace_update:
|
||||||
comfy.utils.copy_to_param(self.model, key, bk.weight)
|
comfy.utils.copy_to_param(self.model, key, bk.weight)
|
||||||
else:
|
else:
|
||||||
comfy.utils.set_attr_param(self.model, key, bk.weight)
|
comfy.utils.set_attr_param(self.model, key, bk.weight)
|
||||||
self.backup.pop(key)
|
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)
|
m.prev_comfy_cast_weights = m.comfy_cast_weights
|
||||||
if weight_key in self.patches:
|
m.comfy_cast_weights = True
|
||||||
m.weight_function = LowVramPatch(weight_key, self.patches)
|
m.comfy_patched_weights = False
|
||||||
patch_counter += 1
|
memory_freed += module_mem
|
||||||
if bias_key in self.patches:
|
logging.debug("freed {}".format(n))
|
||||||
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))
|
|
||||||
|
|
||||||
self.model.model_lowvram = True
|
self.model.model_lowvram = True
|
||||||
self.model.lowvram_patch_counter += patch_counter
|
self.model.lowvram_patch_counter += patch_counter
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user