fred-dev/comfy_ui_ali
0
1from __future__ import annotations2import uuid3import comfy.model_management4import comfy.conds5import comfy.utils6import comfy.hooks7import comfy.patcher_extension8from typing import TYPE_CHECKING9if TYPE_CHECKING:10 from comfy.model_patcher import ModelPatcher11 from comfy.model_base import BaseModel12 from comfy.controlnet import ControlBase13 14def prepare_mask(noise_mask, shape, device):15 return comfy.utils.reshape_mask(noise_mask, shape).to(device)16 17def get_models_from_cond(cond, model_type):18 models = []19 for c in cond:20 if model_type in c:21 if isinstance(c[model_type], list):22 models += c[model_type]23 else:24 models += [c[model_type]]25 return models26 27def get_hooks_from_cond(cond, full_hooks: comfy.hooks.HookGroup):28 # get hooks from conds, and collect cnets so they can be checked for extra_hooks29 cnets: list[ControlBase] = []30 for c in cond:31 if 'hooks' in c:32 for hook in c['hooks'].hooks:33 full_hooks.add(hook)34 if 'control' in c:35 cnets.append(c['control'])36 37 def get_extra_hooks_from_cnet(cnet: ControlBase, _list: list):38 if cnet.extra_hooks is not None:39 _list.append(cnet.extra_hooks)40 if cnet.previous_controlnet is None:41 return _list42 return get_extra_hooks_from_cnet(cnet.previous_controlnet, _list)43 44 hooks_list = []45 cnets = set(cnets)46 for base_cnet in cnets:47 get_extra_hooks_from_cnet(base_cnet, hooks_list)48 extra_hooks = comfy.hooks.HookGroup.combine_all_hooks(hooks_list)49 if extra_hooks is not None:50 for hook in extra_hooks.hooks:51 full_hooks.add(hook)52 53 return full_hooks54 55def convert_cond(cond):56 out = []57 for c in cond:58 temp = c[1].copy()59 model_conds = temp.get("model_conds", {})60 if c[0] is not None:61 temp["cross_attn"] = c[0]62 temp["model_conds"] = model_conds63 temp["uuid"] = uuid.uuid4()64 out.append(temp)65 return out66 67def get_additional_models(conds, dtype):68 """loads additional models in conditioning"""69 cnets: list[ControlBase] = []70 gligen = []71 add_models = []72 73 for k in conds:74 cnets += get_models_from_cond(conds[k], "control")75 gligen += get_models_from_cond(conds[k], "gligen")76 add_models += get_models_from_cond(conds[k], "additional_models")77 78 control_nets = set(cnets)79 80 inference_memory = 081 control_models = []82 for m in control_nets:83 control_models += m.get_models()84 inference_memory += m.inference_memory_requirements(dtype)85 86 gligen = [x[1] for x in gligen]87 models = control_models + gligen + add_models88 89 return models, inference_memory90 91def get_additional_models_from_model_options(model_options: dict[str]=None):92 """loads additional models from registered AddModels hooks"""93 models = []94 if model_options is not None and "registered_hooks" in model_options:95 registered: comfy.hooks.HookGroup = model_options["registered_hooks"]96 for hook in registered.get_type(comfy.hooks.EnumHookType.AdditionalModels):97 hook: comfy.hooks.AdditionalModelsHook98 models.extend(hook.models)99 return models100 101def cleanup_additional_models(models):102 """cleanup additional models that were loaded"""103 for m in models:104 if hasattr(m, 'cleanup'):105 m.cleanup()106 107 108def prepare_sampling(model: ModelPatcher, noise_shape, conds, model_options=None):109 real_model: BaseModel = None110 models, inference_memory = get_additional_models(conds, model.model_dtype())111 models += get_additional_models_from_model_options(model_options)112 models += model.get_nested_additional_models() # TODO: does this require inference_memory update?113 memory_required = model.memory_required([noise_shape[0] * 2] + list(noise_shape[1:])) + inference_memory114 minimum_memory_required = model.memory_required([noise_shape[0]] + list(noise_shape[1:])) + inference_memory115 comfy.model_management.load_models_gpu([model] + models, memory_required=memory_required, minimum_memory_required=minimum_memory_required)116 real_model = model.model117 118 return real_model, conds, models119 120def cleanup_models(conds, models):121 cleanup_additional_models(models)122 123 control_cleanup = []124 for k in conds:125 control_cleanup += get_models_from_cond(conds[k], "control")126 127 cleanup_additional_models(set(control_cleanup))128 129def prepare_model_patcher(model: 'ModelPatcher', conds, model_options: dict):130 '''131 Registers hooks from conds.132 '''133 # check for hooks in conds - if not registered, see if can be applied134 hooks = comfy.hooks.HookGroup()135 for k in conds:136 get_hooks_from_cond(conds[k], hooks)137 # add wrappers and callbacks from ModelPatcher to transformer_options138 model_options["transformer_options"]["wrappers"] = comfy.patcher_extension.copy_nested_dicts(model.wrappers)139 model_options["transformer_options"]["callbacks"] = comfy.patcher_extension.copy_nested_dicts(model.callbacks)140 # begin registering hooks141 registered = comfy.hooks.HookGroup()142 target_dict = comfy.hooks.create_target_dict(comfy.hooks.EnumWeightTarget.Model)143 # handle all TransformerOptionsHooks144 for hook in hooks.get_type(comfy.hooks.EnumHookType.TransformerOptions):145 hook: comfy.hooks.TransformerOptionsHook146 hook.add_hook_patches(model, model_options, target_dict, registered)147 # handle all AddModelsHooks148 for hook in hooks.get_type(comfy.hooks.EnumHookType.AdditionalModels):149 hook: comfy.hooks.AdditionalModelsHook150 hook.add_hook_patches(model, model_options, target_dict, registered)151 # handle all WeightHooks by registering on ModelPatcher152 model.register_all_hook_patches(hooks, target_dict, model_options, registered)153 # add registered_hooks onto model_options for further reference154 if len(registered) > 0:155 model_options["registered_hooks"] = registered156 # merge original wrappers and callbacks with hooked wrappers and callbacks157 to_load_options: dict[str] = model_options.setdefault("to_load_options", {})158 for wc_name in ["wrappers", "callbacks"]:159 comfy.patcher_extension.merge_nested_dicts(to_load_options.setdefault(wc_name, {}), model_options["transformer_options"][wc_name],160 copy_dict1=False)161 return to_load_options162 