mirror of
https://github.com/Comfy-Org/ComfyUI.git
synced 2026-10-01 18:38:01 -05:00
The V2 additions the KJNodes completion needed on the core side: - _model_transforms: the closed, core-owned transform vocabulary behind ModelRef.patch — 29 named transforms, declaratively parameterized, validated host-side, immutable and stacking. No function ever crosses the boundary; a pack cannot register one. - structured-vs-live split: value()/from_value() only on structured data refs (LATENT, AUDIO, TRACKS...). MODEL/CLIP/VAE/asset refs are handles in every execution mode — in-process identity resolution no longer hands a live model to node code. - preview overrides (tiny-VAE, LTX factors), triton VAE seam, memory attention, and profiling surfaces backing the corresponding closed brokers in the overlay. - torch_compile/model_patcher/model_management: compiled-view aliasing recognized by the model manager (no double-counted weights); shared state-dict loading path so the native loader and the V2 broker cannot drift.
94 lines
3.5 KiB
Python
94 lines
3.5 KiB
Python
from __future__ import annotations
|
|
import torch
|
|
|
|
import comfy.utils
|
|
from comfy.patcher_extension import CallbacksMP, WrappersMP
|
|
from typing import TYPE_CHECKING, Callable, Optional
|
|
if TYPE_CHECKING:
|
|
from comfy.model_patcher import ModelPatcher
|
|
from comfy.patcher_extension import WrapperExecutor
|
|
|
|
|
|
COMPILE_KEY = "torch.compile"
|
|
TORCH_COMPILE_KWARGS = "torch_compile_kwargs"
|
|
TORCH_COMPILE_KEYS = "torch_compile_keys"
|
|
|
|
|
|
def apply_torch_compile_factory(compiled_module_dict: dict[str, Callable]) -> Callable:
|
|
'''
|
|
Create a wrapper that will refer to the compiled_diffusion_model.
|
|
'''
|
|
def apply_torch_compile_wrapper(executor: WrapperExecutor, *args, **kwargs):
|
|
try:
|
|
orig_modules = {}
|
|
for key, value in compiled_module_dict.items():
|
|
orig_modules[key] = comfy.utils.get_attr(executor.class_obj, key)
|
|
comfy.utils.set_attr(executor.class_obj, key, value)
|
|
return executor(*args, **kwargs)
|
|
finally:
|
|
for key, value in orig_modules.items():
|
|
comfy.utils.set_attr(executor.class_obj, key, value)
|
|
return apply_torch_compile_wrapper
|
|
|
|
|
|
def set_torch_compile_wrapper(model: ModelPatcher, backend: str, options: Optional[dict[str,str]]=None,
|
|
mode: Optional[str]=None, fullgraph=False, dynamic: Optional[bool]=None,
|
|
keys: list[str]=["diffusion_model"], *args, **kwargs):
|
|
'''
|
|
Perform torch.compile that will be applied at sample time for either the whole model or specific params of the BaseModel instance.
|
|
|
|
When keys is None, it will default to using ["diffusion_model"], compiling the whole diffusion_model.
|
|
When a list of keys is provided, it will perform torch.compile on only the selected modules.
|
|
'''
|
|
# clear out any other torch.compile wrappers
|
|
model.remove_wrappers_with_key(WrappersMP.APPLY_MODEL, COMPILE_KEY)
|
|
model.remove_callbacks_with_key(
|
|
CallbacksMP.ON_DEEPCLONE_MULTIGPU, COMPILE_KEY)
|
|
# if no keys, default to 'diffusion_model'
|
|
if not keys:
|
|
keys = ["diffusion_model"]
|
|
# create kwargs dict that can be referenced later
|
|
compile_kwargs = {
|
|
"backend": backend,
|
|
"options": options,
|
|
"mode": mode,
|
|
"fullgraph": fullgraph,
|
|
"dynamic": dynamic,
|
|
}
|
|
# get a dict of compiled keys
|
|
compiled_modules = {}
|
|
for key in keys:
|
|
compiled_modules[key] = torch.compile(
|
|
model=model.get_model_object(key),
|
|
**compile_kwargs,
|
|
)
|
|
# add torch.compile wrapper
|
|
wrapper_func = apply_torch_compile_factory(
|
|
compiled_module_dict=compiled_modules,
|
|
)
|
|
# store wrapper to run on BaseModel's apply_model function
|
|
model.add_wrapper_with_key(WrappersMP.APPLY_MODEL, COMPILE_KEY, wrapper_func)
|
|
# keep compile kwargs for reference
|
|
model.model_options[TORCH_COMPILE_KWARGS] = compile_kwargs
|
|
model.model_options[TORCH_COMPILE_KEYS] = list(keys)
|
|
|
|
def compile_multigpu_clone(_source, clone):
|
|
refresh_torch_compile_wrapper(clone)
|
|
|
|
model.add_callback_with_key(
|
|
CallbacksMP.ON_DEEPCLONE_MULTIGPU,
|
|
COMPILE_KEY,
|
|
compile_multigpu_clone,
|
|
)
|
|
|
|
|
|
def refresh_torch_compile_wrapper(model: ModelPatcher) -> None:
|
|
compile_kwargs = model.model_options.get(TORCH_COMPILE_KWARGS)
|
|
if compile_kwargs is None:
|
|
return
|
|
set_torch_compile_wrapper(
|
|
model,
|
|
keys=model.model_options.get(TORCH_COMPILE_KEYS),
|
|
**compile_kwargs,
|
|
)
|