From 7fa9bb3dbc04048d050f65337f59ae319682e87b Mon Sep 17 00:00:00 2001 From: smthemex <138738845+smthemex@users.noreply.github.com> Date: Mon, 10 Aug 2026 12:05:51 +0800 Subject: [PATCH] fix lora --- __init__.py | 4 +- block_swap.py | 730 ++++++++++++++++++++-------------------- uniblockswap_node.py | 778 +++++++++++++++++++++++++------------------ 3 files changed, 811 insertions(+), 701 deletions(-) diff --git a/__init__.py b/__init__.py index 8f76d5b..e58fd39 100644 --- a/__init__.py +++ b/__init__.py @@ -1,3 +1,3 @@ -from .uniblockswap_node import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS - +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 index 31a9b09..c0b2ff7 100644 --- a/block_swap.py +++ b/block_swap.py @@ -1,371 +1,361 @@ -""" -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 _backup_ggml_refs(module): - """Preserve the ORIGINAL mmap-backed GGMLTensor objects for every GGML - parameter in `module`. - - Why a full-reference backup (not just .data): tensor_type / tensor_shape / - patches live on the GGMLTensor *object*, and a .to(...) round trip creates - a fresh tensor that loses the mmap mapping. We must keep the original object - alive so we can point the parameter back at it later. - """ - if getattr(module, "_ggml_mmap_backup", None) is not None: - return - backup = {} - for name, param in module.named_parameters(recurse=True): - t = param.data - if hasattr(t, "tensor_type"): # a GGMLTensor - backup[name] = t # keep the object alive, mmap intact - module._ggml_mmap_backup = backup - - -def _restore_ggml_refs(module): - """Point params back at the original mmap GGMLTensors and drop any GPU - copies. This is a *pointer assignment* (p.data = orig), so NO anonymous - heap allocation happens -- unlike module.to(offload_device), which would - reallocate the dequantized weights as non-reclaimable RAM. - - If a block was never GPU-loaded (no backup), fall back to .to(cpu) which is - a no-op for an already-mmap'd CPU tensor. - """ - backup = getattr(module, "_ggml_mmap_backup", None) - if not backup: - module.to(module.offload_device if hasattr(module, "offload_device") else "cpu") - return - params = dict(module.named_parameters(recurse=True)) - for name, orig in backup.items(): - p = params.get(name) - if p is not None: - p.data = orig - # free the GPU copy of the now-unreferenced tensor - if torch.cuda.is_available(): - gc.collect() - torch.cuda.empty_cache() - - -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: - prev_mod = self._modules[str(prev)] - # FREE previous block GPU memory - if _has_ggml_params(prev_mod): - # GGUF: restore the original mmap-backed GGMLTensor by - # pointer assignment. This drops the GPU copy WITHOUT - # reallocating the weights as anonymous CPU RAM (which - # .to(offload_device) would do after a .to(cuda) round - # trip, blowing RAM from 40G to 60G). - _restore_ggml_refs(prev_mod) - else: - # Safetensor: set to meta (vbar restores automatically) - _free_to_meta(prev_mod) - for m in prev_mod.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 - cur_mod = self._modules[str(idx)] - if _has_ggml_params(cur_mod): - # Snapshot the mmap reference so we can later restore it. We do NOT - # call cur_mod.to(compute_device) here: GGUF weights are dequantized - # per-layer on demand inside GGMLLayer.cast_bias_weight() when each - # op runs (self.weight.to(input.device)). Pre-moving the whole block - # to GPU would force a full dequantization of every layer at once, - # spiking VRAM and -- on the next swap -- a GPU->"CPU" round trip, - # both of which defeat the mmap model's whole point. - _backup_ggml_refs(cur_mod) - # 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: - blk = self._modules[str(i)] - if _has_ggml_params(blk): - # Restore the original mmap-backed GGMLTensor (pointer - # assignment, no anonymous RAM). If a block was never - # GPU-loaded the backup is empty and the helper safely - # falls back to a no-op .to(cpu). - _restore_ggml_refs(blk) - else: - _free_to_meta(blk) - for m in blk.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): - # Support slicing: blocks[start:end] - if isinstance(idx, slice): - start, stop, step = idx.indices(self.total_count) - return [self[i] for i in range(start, stop, step)] - - 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: the swap blocks already live in the mmap file-backed mapping - # on CPU. No copy is needed now; we just record the original references - # so a later offload can restore them (pointer assignment, no anon RAM). - # Safetensor blocks stay on GPU (original behavior). - for i in range(total - n, total): - blk = swl._modules[str(i)] - if _has_ggml_params(blk): - _backup_ggml_refs(blk) - - 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): - # Record mmap references; the block stays file-backed on CPU. - _backup_ggml_refs(blk) - 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 - +""" +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 _backup_ggml_refs(module): + """Preserve the ORIGINAL mmap-backed GGMLTensor objects for every GGML + parameter in `module`. + + Why a full-reference backup (not just .data): tensor_type / tensor_shape / + patches live on the GGMLTensor *object*, and a .to(...) round trip creates + a fresh tensor that loses the mmap mapping. We must keep the original object + alive so we can point the parameter back at it later. + """ + if getattr(module, "_ggml_mmap_backup", None) is not None: + return + backup = {} + for name, param in module.named_parameters(recurse=True): + t = param.data + if hasattr(t, "tensor_type"): # a GGMLTensor + backup[name] = t # keep the object alive, mmap intact + module._ggml_mmap_backup = backup + + +def _restore_ggml_refs(module): + """Point params back at the original mmap GGMLTensors and drop any GPU + copies. This is a *pointer assignment* (p.data = orig), so NO anonymous + heap allocation happens -- unlike module.to(offload_device), which would + reallocate the dequantized weights as non-reclaimable RAM. + + If a block was never GPU-loaded (no backup), fall back to .to(cpu) which is + a no-op for an already-mmap'd CPU tensor. + """ + backup = getattr(module, "_ggml_mmap_backup", None) + if not backup: + module.to(module.offload_device if hasattr(module, "offload_device") else "cpu") + return + params = dict(module.named_parameters(recurse=True)) + for name, orig in backup.items(): + p = params.get(name) + if p is not None: + p.data = orig + # free the GPU copy of the now-unreferenced tensor + if torch.cuda.is_available(): + gc.collect() + torch.cuda.empty_cache() + + +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: + prev_mod = self._modules[str(prev)] + # FREE previous block GPU memory + if _has_ggml_params(prev_mod): + # GGUF: restore the original mmap-backed GGMLTensor by + # pointer assignment. This drops the GPU copy WITHOUT + # reallocating the weights as anonymous CPU RAM (which + # .to(offload_device) would do after a .to(cuda) round + # trip, blowing RAM from 40G to 60G). + _restore_ggml_refs(prev_mod) + else: + # Safetensor: set to meta (vbar restores automatically) + _free_to_meta(prev_mod) + # NOTE: vbar state (_v/_prefetch/_v_signature) is intentionally + # left untouched - vbar is READ-ONLY here and restores the block + # on its own fault mechanism. + except Exception: + pass + # LOAD current block if GGUF + cur_mod = self._modules[str(idx)] + if _has_ggml_params(cur_mod): + # Snapshot the mmap reference so we can later restore it. We do NOT + # call cur_mod.to(compute_device) here: GGUF weights are dequantized + # per-layer on demand inside GGMLLayer.cast_bias_weight() when each + # op runs (self.weight.to(input.device)). Pre-moving the whole block + # to GPU would force a full dequantization of every layer at once, + # spiking VRAM and -- on the next swap -- a GPU->"CPU" round trip, + # both of which defeat the mmap model's whole point. + _backup_ggml_refs(cur_mod) + # 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: + blk = self._modules[str(i)] + if _has_ggml_params(blk): + # Restore the original mmap-backed GGMLTensor (pointer + # assignment, no anonymous RAM). If a block was never + # GPU-loaded the backup is empty and the helper safely + # falls back to a no-op .to(cpu). + _restore_ggml_refs(blk) + else: + _free_to_meta(blk) + # vbar state left untouched (read-only by design) + 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): + # Support slicing: blocks[start:end] + if isinstance(idx, slice): + start, stop, step = idx.indices(self.total_count) + return [self[i] for i in range(start, stop, step)] + + 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: the swap blocks already live in the mmap file-backed mapping + # on CPU. No copy is needed now; we just record the original references + # so a later offload can restore them (pointer assignment, no anon RAM). + # Safetensor blocks stay on GPU (original behavior). + for i in range(total - n, total): + blk = swl._modules[str(i)] + if _has_ggml_params(blk): + _backup_ggml_refs(blk) + + 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): + # Record mmap references; the block stays file-backed on CPU. + _backup_ggml_refs(blk) + 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/uniblockswap_node.py b/uniblockswap_node.py index eeaae3e..887aeb4 100644 --- a/uniblockswap_node.py +++ b/uniblockswap_node.py @@ -1,330 +1,450 @@ -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", +import logging +import torch +import comfy.model_management as mm +import comfy.patcher_extension +import gc +import uuid +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 ('ggml_weight', 'ggml_weight_data'): + if hasattr(m, attr): + try: + delattr(m, attr) + except Exception: + pass + except Exception: + pass + + +def _ensure_lora_functions(patcher, swl): + """Attach LoRA LowVramPatch functions to every swap block module. + + This reuses the LoRA patches ComfyUI already attached (patcher.patches) - + no weight data is re-read and nothing is re-wrapped. ComfyUI's cast path + applies module.weight_function (incl. vbar fault restore, ops.py post_cast), + so LoRA stays effective every time a swap block is loaded into CUDA. + """ + if getattr(patcher, "mmap_released", False): + return # GGUF: keep original dequant/cast path behavior + import comfy.model_patcher as mp + patches = getattr(patcher, "patches", None) + if not patches: + return + full_path = None + for path, mod in patcher.model.named_modules(): + if mod is swl: + full_path = path + break + if full_path is None: + return + for i in range(swl.non_swap_count, swl.total_count): + try: + blk = swl._modules.get(str(i)) + if blk is None: + continue + prefix = f"{full_path}.{i}" + for mname, m in blk.named_modules(): + base = prefix if not mname else f"{prefix}.{mname}" + for pname in ("weight", "bias"): + key = f"{base}.{pname}" + if key not in patches: + continue + try: + _, set_func, convert_func = mp.get_key_weight( + patcher.model, key) + except Exception: + continue + fn = mp.LowVramPatch(key, patches, convert_func, set_func) + attr_name = pname + "_function" + cur = list(getattr(m, attr_name, None) or []) + if not any(getattr(f, "key", None) == key for f in cur): + cur.append(fn) + setattr(m, attr_name, cur) + except Exception: + continue + + +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." + + # NOTE: no IS_CHANGED on purpose. apply_swap() is a one-time install step + # (wraps the shared model object in SwappableModuleList + attaches LoRA + # weight_functions). Forcing re-execution every run would re-wrap the + # already-wrapped model (nested swap) and break swap + LoRA state. + # Per-inference VRAM cleanup of TE/VAE is handled by the separate + # UniBlockSwapCacheControl node, which re-runs every inference. + 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,) + + # Re-entry guard: model.clone() shares the SAME underlying model object + # (no deepcopy). If this node re-runs on an already-swapped model (e.g. + # num_blocks changed -> input changed -> node re-executed), fully restore + # the original structure first, otherwise install_block_swap would wrap + # the SwappableModuleList again (nested) and break swap/LoRA state. + prev_cleanup = getattr(diffusion_model, "_uniblockswap_cleanup", None) + if prev_cleanup is not None: + try: + prev_cleanup() + except Exception: + logger.warning("UniBlockSwap: failed to restore previous swap before reinstall", exc_info=True) + + 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 swl in _dit_all_swls: + _ensure_lora_functions(p, swl) + 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, + ) + + # Swap blocks: never write LoRA-patched weights back into parameters + # here. LoRA is applied at cast time via weight_function + # (_ensure_lora_functions), otherwise it would double-apply. + _orig_patch = patcher.patch_weight_to_device + def _skip_swap_patch(key, *args, **kwargs): + if _is_dit_swap_key(key): + return + return _orig_patch(key, *args, **kwargs) + patcher.patch_weight_to_device = _skip_swap_patch + + # CRITICAL: filter swap blocks from _load_list so load() never + # iterates over them (m.to(device_to) would load all swap blocks to + # GPU at once). Blocks are loaded one-at-a-time by swap + vbar fault. + _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 + + # Attach LoRA weight_functions to swap blocks (idempotent). + for swl in _dit_all_swls: + _ensure_lora_functions(patcher, swl) + + 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." + + # NOTE: no IS_CHANGED on purpose (see UniBlockSwap comment). Per-inference + # VRAM cleanup of the text encoder is handled by UniBlockSwapCacheControl. + 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,) + + # Re-entry guard: clip.clone() shares the same underlying model object. + # Restore any previously installed TE swap structure before reinstalling + # to avoid nesting SwappableModuleList / double-wrapping forward. + prev_cleanup = getattr(new_clip.patcher.model, "_uniblockswap_te_cleanup", None) + if prev_cleanup is not None: + try: + prev_cleanup() + except Exception: + logger.warning("UniBlockSwapTE: failed to restore previous swap before reinstall", exc_info=True) + + 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) + for mgr in mgr_list: + _ensure_lora_functions(p, mgr) + 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 + + for mgr in mgr_list: + _ensure_lora_functions(new_clip.patcher, mgr) + + 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,) + + +class UniBlockSwapCacheControl: + """Per-inference VRAM cleanup for the other models (TE/VAE), passthrough. + + UniBlockSwap / UniBlockSwapTE no longer force re-execution via IS_CHANGED + (re-running would re-wrap the shared model object and break swap/LoRA). + Instead this node re-runs every inference but only performs the cheap + clear_comfyui_cache_except() side effect, leaving swap installation intact. + + Usage: UniBlockSwap -> UniBlockSwapCacheControl -> KSampler + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": {"model": ("MODEL",)}, + } + + RETURN_TYPES = ("MODEL",) + RETURN_NAMES = ("model",) + FUNCTION = "clear_cache" + CATEGORY = "model/loaders" + DESCRIPTION = ("每次推理清除除传入 model 外的其他模型(TE/VAE 等)的显存," + "并原样透传 model。放在 UniBlockSwap 之后、KSampler 之前。" + "替代原先用 IS_CHANGED 强制 UniBlockSwap 重跑的做法," + "避免重复安装 swap 导致 LoRA 失效。") + + @classmethod + def IS_CHANGED(cls, **kwargs): + # Re-run every inference, but this node only clears cache - it does not + # reinstall swap, so it has no destructive side effects. + return uuid.uuid4().hex + + def clear_cache(self, model): + clear_comfyui_cache_except(model) + return (model,) + + +NODE_CLASS_MAPPINGS = { + "UniBlockSwap": UniBlockSwap, + "UniBlockSwapTE": UniBlockSwapTE, + "UniBlockSwapCacheControl": UniBlockSwapCacheControl, +} +NODE_DISPLAY_NAME_MAPPINGS = { + "UniBlockSwap": "UniBlockSwap", + "UniBlockSwapTE": "UniBlockSwap TE", + "UniBlockSwapCacheControl": "UniBlockSwap Cache Control", } \ No newline at end of file