CoolFace
Apppublic

fred-dev/comfy_ui_ali

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
sampler_helpers.py162 linesDownload Raw Back to comfy
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