From 66383c7da47fdb0955149c21d076939dabd78ea3 Mon Sep 17 00:00:00 2001 From: smthemex Date: Mon, 13 Jul 2026 15:44:43 +0800 Subject: [PATCH] init --- README.md | 2 +- __init__.py | 3 + block_swap.py | 301 +++++++++++++++++++++++++++++++++++++++ pyproject.toml | 15 ++ uniblockswap_node.py | 330 +++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 650 insertions(+), 1 deletion(-) create mode 100644 __init__.py create mode 100644 block_swap.py create mode 100644 pyproject.toml create mode 100644 uniblockswap_node.py diff --git a/README.md b/README.md index 44676a3..63fab9a 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..8f76d5b --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .uniblockswap_node import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/block_swap.py b/block_swap.py new file mode 100644 index 0000000..746aee5 --- /dev/null +++ b/block_swap.py @@ -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 \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..8ef15a7 --- /dev/null +++ b/pyproject.toml @@ -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 = [] diff --git a/uniblockswap_node.py b/uniblockswap_node.py new file mode 100644 index 0000000..eeaae3e --- /dev/null +++ b/uniblockswap_node.py @@ -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", +} \ No newline at end of file