Merge pull request #2 from smthemex/Pr

init
This commit is contained in:
smthemex
2026-07-13 15:45:05 +08:00
committed by GitHub
5 changed files with 650 additions and 1 deletions
+1 -1
View File
@@ -1,7 +1,7 @@
# ComfyUI_UniBlockSwap
A universal swap node that supports ComfyUI native workflow, allowing 4_6G users to experience Klein9B or other large models
# Coming soon
# Update
* Make it for ' low Vram and normal Ram' users to esay running ComfyUI origin workflows.(Support allmot all of comfyUI origin workflows)
* Support text encoder or diffusion models, is enable text encoder will need more Ram
+3
View File
@@ -0,0 +1,3 @@
from .uniblockswap_node import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+301
View File
@@ -0,0 +1,301 @@
"""
UniBlockSwap - Universal single-block swap for ComfyUI.
Safetensor blocks: freed to meta on swap, restored by vbar automatically.
GGUF blocks: freed to CPU on swap, moved to GPU when accessed.
"""
import gc
import logging
import torch
import torch.nn as nn
logger = logging.getLogger(__name__)
CONTAINER_NAMES = (
"blocks", "transformer_blocks", "double_blocks", "single_blocks",
"input_blocks", "output_blocks", "middle_block", "layers",
"double_stream_layers", "single_stream_layers",
"block",
)
def find_blocks(model):
for name in CONTAINER_NAMES:
c = getattr(model, name, None)
if isinstance(c, (nn.ModuleList, list)) and len(c) > 0 and hasattr(c[0], "forward"):
return name, c
return None, None
def _has_ggml_params(module):
"""Check if module has GGMLTensor parameters (quantized GGUF weights)."""
for p in module.parameters():
if hasattr(p, 'tensor_type'):
return True
return False
def _free_to_meta(module):
"""Free param data to meta tensor - NO CPU copy created.
The module structure is preserved. next load() restores from backup."""
for param in module.parameters(recurse=False):
param.data = torch.empty(0, device='meta')
class SwappableModuleList(nn.ModuleList):
def __init__(self, modules, compute_device, offload_device,
non_swap_count=0):
super().__init__(modules)
self.compute_device = compute_device
self.offload_device = offload_device
self.non_swap_count = non_swap_count
self.total_count = len(modules)
self._loaded_swap_idx = -1
self.container_name = ''
def _load_swap(self, local_idx):
idx = local_idx + self.non_swap_count
if local_idx == self._loaded_swap_idx:
return
if self._loaded_swap_idx >= 0:
prev = self._loaded_swap_idx + self.non_swap_count
try:
# FREE previous block GPU memory
if _has_ggml_params(self._modules[str(prev)]):
# GGUF: move quantized data to CPU (preserves GGMLTensor attributes)
self._modules[str(prev)].to(self.offload_device)
else:
# Safetensor: set to meta (vbar restores automatically)
_free_to_meta(self._modules[str(prev)])
for m in self._modules[str(prev)].modules():
for attr in ('_v', '_prefetch', '_v_signature'):
if hasattr(m, attr):
try:
delattr(m, attr)
except Exception:
pass
except Exception:
pass
# LOAD current block if GGUF
if _has_ggml_params(self._modules[str(idx)]):
self._modules[str(idx)].to(self.compute_device)
# else: safetensor - vbar handles restoration
self._loaded_swap_idx = local_idx
def offload_swap_blocks(self):
for i in range(self.non_swap_count, self.total_count):
try:
if _has_ggml_params(self._modules[str(i)]):
self._modules[str(i)].to(self.offload_device)
else:
_free_to_meta(self._modules[str(i)])
for m in self._modules[str(i)].modules():
for attr in ('_v', '_prefetch', '_v_signature'):
if hasattr(m, attr):
try:
delattr(m, attr)
except Exception:
pass
except Exception:
pass
self._loaded_swap_idx = -1
def _apply(self, fn, recurse=True):
"""Apply fn to non-swap blocks only.
CRITICAL: Prevents model.to(device_to) from moving swap block
GGMLTensors to GPU, which would cause a VRAM spike (12GB).
Safetensor swap blocks are already meta (no-op), so this only
affects GGUF paths.
nn.ModuleList._apply(recurse=False) applies fn to all _modules
entries INCLUDING swap blocks. We skip that and handle only
non_swap_count blocks manually.
"""
for i in range(self.non_swap_count):
try:
child = self._modules.get(str(i))
if child is not None:
child._apply(fn, recurse)
except Exception:
pass
return self
def __getattr__(self, name):
try:
idx = int(name)
if 0 <= idx < self.total_count:
return self.__getitem__(idx)
except (ValueError, TypeError):
pass
raise AttributeError(f"'{type(self).__name__}' has no attribute '{name}'")
def __getitem__(self, idx):
if idx >= self.non_swap_count:
self._load_swap(idx - self.non_swap_count)
return super().__getitem__(idx)
def __iter__(self):
for idx in range(self.total_count):
yield self.__getitem__(idx)
def install_block_swap(diffusion_model, compute_device, offload_device,
num_blocks=-1):
all_containers = []
for name in CONTAINER_NAMES:
c = getattr(diffusion_model, name, None)
if isinstance(c, (nn.ModuleList, list)) and len(c) > 0 and hasattr(c[0], "forward"):
all_containers.append((name, c))
if not all_containers:
return None, lambda: None, set()
first_swl = None
all_names = set()
for name, orig in all_containers:
total = len(orig)
n = num_blocks if num_blocks > 0 else total
n = max(1, min(n, total))
swl = SwappableModuleList(
orig, compute_device, offload_device,
non_swap_count=total - n,
)
swl.container_name = name
setattr(diffusion_model, name, swl)
all_names.add(name)
if first_swl is None:
first_swl = swl
logger.info("UniBlockSwap: '%s' = %d blocks, swapping %d",
name, total, n)
# For GGUF: offload swap blocks to CPU immediately.
# Safetensor blocks stay on GPU (original behavior).
for i in range(total - n, total):
blk = swl._modules[str(i)]
if _has_ggml_params(blk):
blk.to(offload_device)
orig_fwd = diffusion_model.forward
def wrapped(*args, **kwargs):
try:
return orig_fwd(*args, **kwargs)
finally:
if torch.cuda.is_available():
torch.cuda.synchronize(compute_device)
gc.collect()
torch.cuda.empty_cache()
diffusion_model.forward = wrapped
def cleanup():
diffusion_model.forward = orig_fwd
for name, orig in all_containers:
setattr(diffusion_model, name, orig)
all_swls = []
for name in CONTAINER_NAMES:
c = getattr(diffusion_model, name, None)
if hasattr(c, 'offload_swap_blocks'):
all_swls.append(c)
return first_swl, cleanup, all_names, all_swls
def find_te_containers(cond_stage_model):
results = []
seen_ids = set()
def _recurse(module, depth=0):
if depth > 20:
return
for name in CONTAINER_NAMES:
c = getattr(module, name, None)
if (isinstance(c, (nn.ModuleList, list)) and
len(c) > 0 and hasattr(c[0], "forward") and
id(c) not in seen_ids):
seen_ids.add(id(c))
results.append((name, c, module))
for child_name, child in module.named_children():
if isinstance(child, (nn.ModuleList, list)):
continue
_recurse(child, depth + 1)
_recurse(cond_stage_model)
return results
def install_te_block_swap(cond_stage_model, compute_device, offload_device,
num_blocks=-1):
containers = find_te_containers(cond_stage_model)
if not containers:
return [], lambda: None, set()
mgr_list = []
container_names = set()
parent_to_mgrs = {}
for name, orig, parent in containers:
total = len(orig)
n = num_blocks if num_blocks > 0 else total
n = max(1, min(n, total))
swl = SwappableModuleList(
orig, compute_device, offload_device,
non_swap_count=total - n,
)
swl.container_name = name
setattr(parent, name, swl)
mgr_list.append(swl)
container_names.add(name)
parent_id = id(parent)
if parent_id not in parent_to_mgrs:
parent_to_mgrs[parent_id] = (parent, parent.forward, [])
parent_to_mgrs[parent_id][2].append(swl)
logger.info("UniBlockSwapTE: '%s' (%s) = %d blocks, swapping %d",
name, type(parent).__name__, total, n)
for i in range(total - n, total):
blk = swl._modules[str(i)]
if _has_ggml_params(blk):
blk.to(offload_device)
else:
_free_to_meta(blk)
wrapped_parents = []
for parent_id, (parent, orig_fwd, parent_mgrs) in parent_to_mgrs.items():
def make_wrapped(_orig_fwd=orig_fwd, _mgrs=parent_mgrs, _cdevice=compute_device,
_root=cond_stage_model):
def wrapped(*args, **kwargs):
try:
return _orig_fwd(*args, **kwargs)
finally:
for m in _mgrs:
m.offload_swap_blocks()
backup_cleaner = getattr(_root, '_uniblockswap_backup_cleanup', None)
patcher = getattr(_root, '_patcher_ref', None)
if backup_cleaner is not None and patcher is not None:
backup_cleaner(patcher)
if torch.cuda.is_available():
torch.cuda.synchronize(_cdevice)
gc.collect()
torch.cuda.empty_cache()
return wrapped
parent.forward = make_wrapped()
wrapped_parents.append((parent, orig_fwd))
def cleanup():
for name, orig, parent in containers:
current = getattr(parent, name, None)
if hasattr(current, 'offload_swap_blocks'):
setattr(parent, name, orig)
for parent, orig_fwd in wrapped_parents:
parent.forward = orig_fwd
return mgr_list, cleanup, container_names
+15
View File
@@ -0,0 +1,15 @@
[project]
name = "uniblockswap"
description = "A universal swap node that supports ComfyUI native workflow, allowing 4_6G users to experience Klein9B or other large models"
version = "1.0.0"
license = {file = "LICENSE"}
[project.urls]
Repository = "https://github.com/smthemex/ComfyUI_UniBlockSwap"
# Used by Comfy Registry https://registry.comfy.org
[tool.comfy]
PublisherId = "smthemex"
DisplayName = "ComfyUI_UniBlockSwap"
Icon = ""
includes = []
+330
View File
@@ -0,0 +1,330 @@
import logging
import torch
import comfy.model_management as mm
import comfy.patcher_extension
import gc
from .block_swap import install_block_swap, install_te_block_swap, _free_to_meta, _has_ggml_params
logger = logging.getLogger(__name__)
def _get_diffusion_model(patcher):
if patcher is None:
return None
model_obj = getattr(patcher, "model", patcher)
diffusion = getattr(model_obj, "diffusion_model", None)
if diffusion is not None and isinstance(diffusion, torch.nn.Module):
return diffusion
if isinstance(model_obj, torch.nn.Module):
return model_obj
inner = getattr(patcher, "model", None)
if inner is not None and isinstance(inner, torch.nn.Module):
return inner
return None
def _get_cond_stage_model(clip_obj):
"""Extract the cond_stage_model from a CLIP wrapper."""
if clip_obj is None:
return None
cond_stage = getattr(clip_obj, "cond_stage_model", None)
if cond_stage is not None and isinstance(cond_stage, torch.nn.Module):
return cond_stage
return None
def _free_block_cleanup(swl):
"""Free swap block memory during ON_CLEANUP.
Safetensor: _free_to_meta (release to meta, vbar handles restore).
GGUF: to(offload_device) (quantized data to CPU, GGMLTensor preserved).
"""
for i in range(swl.non_swap_count, swl.total_count):
try:
blk = swl._modules.get(str(i))
if blk is None:
continue
if _has_ggml_params(blk):
blk.to(swl.offload_device)
else:
_free_to_meta(blk)
for m in blk.modules():
for attr in ('_v', '_prefetch', '_v_signature',
'ggml_weight', 'ggml_weight_data'):
if hasattr(m, attr):
try:
delattr(m, attr)
except Exception:
pass
except Exception:
pass
def clear_comfyui_cache_except(exclude_patcher=None):
"""Clear all models from GPU to CPU (unpatch), except exclude_patcher.
This frees VRAM used by TE/VAE/etc without touching the DIT model.
"""
cf_models = mm.loaded_models()
for pipe in cf_models:
if exclude_patcher is not None and pipe is exclude_patcher:
continue
try:
pipe.unpatch_model(device_to=torch.device("cpu"))
except Exception:
pass
mm.soft_empty_cache()
torch.cuda.empty_cache()
max_gpu_memory = torch.cuda.max_memory_allocated()
print(f"After Max GPU memory allocated: {max_gpu_memory / 1000 ** 3:.2f} GB")
class UniBlockSwap:
"""Swap blocks one-at-a-time between GPU/CPU to reduce VRAM.
Supports both safetensor and GGUF models.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"model": ("MODEL",)},
"optional": {
"num_blocks": ("INT", {
"default": -1, "min": -1, "max": 10000, "step": 1,
"tooltip": "Blocks from end to swap. -1 = all, 0 = disable",
}),
},
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("model",)
FUNCTION = "apply_swap"
CATEGORY = "model/loaders"
DESCRIPTION = "Swap blocks one-at-a-time between GPU/CPU to reduce VRAM."
def apply_swap(self, model, num_blocks=-1):
if num_blocks == 0:
return (model,)
patcher = model.clone()
if hasattr(model, 'backup'):
model.backup.clear()
patcher.backup = {}
clear_comfyui_cache_except(patcher)
diffusion_model = _get_diffusion_model(patcher)
if diffusion_model is None:
logger.warning("UniBlockSwap: no diffusion model found")
return (patcher,)
compute = mm.get_torch_device()
offload = mm.unet_offload_device()
logger.info("UniBlockSwap: %s, compute=%s, offload=%s",
type(diffusion_model).__name__, compute, offload)
mgr, cleanup, _dit_swap_names, _dit_all_swls = install_block_swap(
diffusion_model, compute, offload,
num_blocks=num_blocks,
)
if mgr is None:
return (patcher,)
def _is_dit_swap_key(key):
parts = key.split(".")
for i, part in enumerate(parts):
if part in _dit_swap_names and i + 1 < len(parts):
next_part = parts[i + 1]
if next_part.lstrip("-").isdigit():
return True
return False
def _on_load(p, device_to, lowvram, force, full):
try:
mgr.offload_swap_blocks()
for key in list(p.backup.keys()):
if _is_dit_swap_key(key):
p.backup.pop(key, None)
except Exception:
pass
mm.soft_empty_cache()
gc.collect()
patcher.add_callback_with_key(
comfy.patcher_extension.CallbacksMP.ON_LOAD,
"UniBlockSwap", _on_load,
)
# Detect if this patcher is a GGUFModelPatcher (which handles GGMLTensor weights).
_is_gguf = hasattr(patcher, 'mmap_released')
_orig_patch = patcher.patch_weight_to_device
def _skip_swap_patch(key, *args, **kwargs):
if _is_dit_swap_key(key):
if _is_gguf:
# GGUF: completely skip. _load_swap manages GPU loading.
return
# Safetensor: call original, delete backup.
result = _orig_patch(key, *args, **kwargs)
if key in patcher.backup:
patcher.backup.pop(key, None)
return result
return _orig_patch(key, *args, **kwargs)
patcher.patch_weight_to_device = _skip_swap_patch
# CRITICAL: _load_list filter for GGUF to prevent load() from
# iterating over swap blocks and calling m.to(device_to) on each,
# which would load all GGUF swap blocks to GPU at once (12GB spike).
if _is_gguf:
_orig_load_list = patcher._load_list
def _filtered_load_list(*args, **kwargs):
raw = _orig_load_list(*args, **kwargs)
return [item for item in raw if not _is_dit_swap_key(item[-3])]
patcher._load_list = _filtered_load_list
def _on_dit_cleanup(p):
try:
for swl in _dit_all_swls:
_free_block_cleanup(swl)
for key in list(p.backup.keys()):
if _is_dit_swap_key(key):
p.backup.pop(key, None)
for _ in range(3):
gc.collect()
except Exception:
pass
patcher.add_callback_with_key(
comfy.patcher_extension.CallbacksMP.ON_CLEANUP,
"UniBlockSwap", _on_dit_cleanup,
)
patcher.model._uniblockswap_cleanup = cleanup
return (patcher,)
class UniBlockSwapTE:
"""Swap text encoder blocks one-at-a-time between GPU/CPU to save VRAM."""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {"clip": ("CLIP",)},
"optional": {
"num_blocks": ("INT", {
"default": -1, "min": -1, "max": 10000, "step": 1,
"tooltip": "Blocks from end to swap. -1 = all, 0 = disable",
}),
},
}
RETURN_TYPES = ("CLIP",)
RETURN_NAMES = ("clip",)
FUNCTION = "apply_swap"
CATEGORY = "model/loaders"
DESCRIPTION = "Swap text encoder blocks one-at-a-time between GPU/CPU to reduce VRAM."
def apply_swap(self, clip, num_blocks=-1):
if num_blocks == 0:
return (clip,)
new_clip = clip.clone()
cond_stage = _get_cond_stage_model(new_clip)
if cond_stage is None:
logger.warning("UniBlockSwapTE: no cond_stage_model found")
return (new_clip,)
new_clip.patcher.backup = {}
mm.soft_empty_cache()
torch.cuda.empty_cache()
gc.collect()
compute = new_clip.patcher.load_device
offload = new_clip.patcher.offload_device
logger.info("UniBlockSwapTE: %s, compute=%s, offload=%s",
type(cond_stage).__name__, compute, offload)
mgr_list, cleanup, container_names = install_te_block_swap(
cond_stage, compute, offload,
num_blocks=num_blocks,
)
if not mgr_list:
logger.info("UniBlockSwapTE: no block containers found in %s",
type(cond_stage).__name__)
return (new_clip,)
def _is_swap_key(key):
for mgr in mgr_list:
cname = getattr(mgr, 'container_name', '')
if not cname:
continue
parts = key.split(".")
for i, part in enumerate(parts):
if part == cname and i + 1 < len(parts):
next_part = parts[i + 1]
if next_part.lstrip("-").isdigit():
return True
return False
def _purge_swap_from_backup(p):
if len(p.backup) == 0:
return
try:
keys_to_del = [k for k in p.backup if _is_swap_key(k)]
for k in keys_to_del:
p.backup.pop(k, None)
except Exception:
pass
def _on_load(p, device_to, lowvram, force, full):
_purge_swap_from_backup(p)
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
new_clip.patcher.add_callback_with_key(
comfy.patcher_extension.CallbacksMP.ON_LOAD,
"UniBlockSwapTE", _on_load,
)
_orig_patch = new_clip.patcher.patch_weight_to_device
def _skip_swap_patch(key, *args, **kwargs):
if _is_swap_key(key):
return
return _orig_patch(key, *args, **kwargs)
new_clip.patcher.patch_weight_to_device = _skip_swap_patch
_orig_load_list = new_clip.patcher._load_list
def _filtered_load_list(*args, **kwargs):
raw = _orig_load_list(*args, **kwargs)
return [item for item in raw if not _is_swap_key(item[-3])]
new_clip.patcher._load_list = _filtered_load_list
new_clip.patcher.model._uniblockswap_te_cleanup = cleanup
def _on_cleanup(p):
try:
for mgr in mgr_list:
_free_block_cleanup(mgr)
_purge_swap_from_backup(p)
for _ in range(3):
gc.collect()
except Exception:
pass
new_clip.patcher.add_callback_with_key(
comfy.patcher_extension.CallbacksMP.ON_CLEANUP,
"UniBlockSwapTE", _on_cleanup,
)
return (new_clip,)
NODE_CLASS_MAPPINGS = {
"UniBlockSwap": UniBlockSwap,
"UniBlockSwapTE": UniBlockSwapTE,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"UniBlockSwap": "UniBlockSwap",
"UniBlockSwapTE": "UniBlockSwap TE",
}