fix lora
This commit is contained in:
+2
-2
@@ -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"]
|
||||
+360
-370
@@ -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
|
||||
+449
-329
@@ -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",
|
||||
}
|
||||
Reference in New Issue
Block a user