427 lines
15 KiB
Python
427 lines
15 KiB
Python
"""
|
|
BC-250 gfx1010 Comprehensive Monkey-Patch v6
|
|
1. Replaces torch.softmax with manual implementation (VGPR overflow fix)
|
|
2. Replaces SDPA with manual implementation
|
|
3. Patches GGUF cast_bias_weight to dequant on CPU (avoids GPU page-fault hangs)
|
|
4. Pre-clones mmap'd tensor data before GPU transfer (XNACK workaround)
|
|
5. Pre-warms GPU context and caching allocator
|
|
6. Forces text encoder to CPU (memory constraint)
|
|
7. Forces VAE decode on CPU (prevents GPU page-fault hang on safetensors mmap)
|
|
|
|
v6 changes: Removed NO_VRAM (made sampling impossibly slow).
|
|
Instead, VAE is forced to decode on CPU. UNet uses normal lowvram path.
|
|
Previous LOWVRAM run: 4/4 steps in 21s (5.5s/step). NO_VRAM: stuck at 0/4 for 10+ min.
|
|
|
|
BC-250 APU / gfx1010: XNACK disabled, no GPU page fault handling.
|
|
GPU copy shader hangs on non-resident pages (mmap'd or swapped).
|
|
Place in ComfyUI root and import as first line of main.py.
|
|
"""
|
|
import torch
|
|
import torch.nn.functional as F
|
|
import os
|
|
import sys
|
|
import gc
|
|
import logging
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
SAFE_SOFTMAX_THRESHOLD = int(os.environ.get("BC250_SOFTMAX_THRESHOLD", "512"))
|
|
|
|
_original_softmax = torch.nn.functional.softmax
|
|
_original_tensor_softmax = torch.Tensor.softmax
|
|
_original_sdpa = torch.nn.functional.scaled_dot_product_attention
|
|
|
|
# === MMAP PRE-CLONE PATCH ===
|
|
_original_module_apply = torch.nn.Module._apply
|
|
|
|
def _bc250_safe_apply(self, fn, recurse=True):
|
|
"""Pre-clone mmap'd CPU tensor data before GPU transfer to avoid XNACK hangs."""
|
|
for key, param in self._parameters.items():
|
|
if param is not None and param.device.type == 'cpu':
|
|
param.data = param.data.clone()
|
|
for key, buf in self._buffers.items():
|
|
if buf is not None and buf.device.type == 'cpu':
|
|
self._buffers[key] = buf.clone()
|
|
return _original_module_apply(self, fn, recurse)
|
|
|
|
# === SOFTMAX PATCH ===
|
|
|
|
def _safe_softmax_impl(input, dim=-1):
|
|
x_max = input.max(dim=dim, keepdim=True).values
|
|
exp_x = torch.exp(input - x_max)
|
|
return exp_x / exp_x.sum(dim=dim, keepdim=True)
|
|
|
|
def patched_softmax(input, dim=None, _stacklevel=3, dtype=None):
|
|
if dim is None:
|
|
dim = -1
|
|
if dtype is not None:
|
|
input = input.to(dtype)
|
|
if input.is_cuda and input.shape[dim] > SAFE_SOFTMAX_THRESHOLD:
|
|
return _safe_softmax_impl(input, dim)
|
|
return _original_softmax(input, dim=dim)
|
|
|
|
def patched_tensor_softmax(self, dim=-1, dtype=None):
|
|
if dtype is not None:
|
|
self = self.to(dtype)
|
|
if self.is_cuda and self.shape[dim] > SAFE_SOFTMAX_THRESHOLD:
|
|
return _safe_softmax_impl(self, dim)
|
|
return _original_tensor_softmax(self, dim=dim)
|
|
|
|
def _safe_sdpa(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None):
|
|
L, S = query.size(-2), key.size(-2)
|
|
if scale is None:
|
|
scale = query.size(-1) ** -0.5
|
|
attn_weight = torch.matmul(query, key.transpose(-2, -1)) * scale
|
|
if is_causal:
|
|
causal_mask = torch.triu(torch.ones(L, S, dtype=torch.bool, device=query.device), diagonal=1)
|
|
attn_weight = attn_weight.masked_fill(causal_mask, float('-inf'))
|
|
if attn_mask is not None:
|
|
if attn_mask.dtype == torch.bool:
|
|
attn_weight = attn_weight.masked_fill(~attn_mask, float('-inf'))
|
|
else:
|
|
attn_weight = attn_weight + attn_mask
|
|
attn_weight = _safe_softmax_impl(attn_weight, dim=-1)
|
|
if dropout_p > 0.0:
|
|
attn_weight = torch.nn.functional.dropout(attn_weight, p=dropout_p)
|
|
return torch.matmul(attn_weight, value)
|
|
|
|
def patched_sdpa(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, scale=None):
|
|
S = key.size(-2)
|
|
if query.is_cuda and S > SAFE_SOFTMAX_THRESHOLD:
|
|
return _safe_sdpa(query, key, value, attn_mask=attn_mask,
|
|
dropout_p=dropout_p, is_causal=is_causal, scale=scale)
|
|
return _original_sdpa(query, key, value, attn_mask=attn_mask,
|
|
dropout_p=dropout_p, is_causal=is_causal, scale=scale)
|
|
|
|
# === GGUF CPU-DEQUANT PATCH (cast_bias_weight override) ===
|
|
_gguf_patched = False
|
|
|
|
def _try_patch_gguf():
|
|
"""Patch GGMLLayer.cast_bias_weight to dequant on CPU, send floats to GPU."""
|
|
global _gguf_patched
|
|
if _gguf_patched:
|
|
return True
|
|
|
|
ops_mod = None
|
|
dequant_mod = None
|
|
for name, mod in sys.modules.items():
|
|
if mod is None:
|
|
continue
|
|
if name.endswith('.ops') and 'GGUF' in name:
|
|
ops_mod = mod
|
|
if name.endswith('.dequant') and 'GGUF' in name:
|
|
dequant_mod = mod
|
|
|
|
if ops_mod is None or dequant_mod is None:
|
|
return False
|
|
|
|
GGMLLayer = getattr(ops_mod, 'GGMLLayer', None)
|
|
is_quantized_fn = getattr(dequant_mod, 'is_quantized', None)
|
|
if GGMLLayer is None or is_quantized_fn is None:
|
|
return False
|
|
|
|
def _bc250_cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None):
|
|
"""Dequant on CPU, only send float results to GPU."""
|
|
import comfy.model_management
|
|
import comfy.ops
|
|
|
|
if input is not None:
|
|
if dtype is None:
|
|
dtype = getattr(input, "dtype", torch.float32)
|
|
if bias_dtype is None:
|
|
bias_dtype = dtype
|
|
if device is None:
|
|
device = input.device
|
|
|
|
non_blocking = comfy.model_management.device_supports_non_blocking(device)
|
|
|
|
bias = None
|
|
if s.bias is not None:
|
|
if is_quantized_fn(s.bias):
|
|
bias = s.get_weight(s.bias, bias_dtype)
|
|
else:
|
|
bias = s.get_weight(s.bias.to(device), bias_dtype)
|
|
bias = comfy.ops.cast_to(bias, bias_dtype, device, non_blocking=non_blocking, copy=False)
|
|
|
|
if is_quantized_fn(s.weight):
|
|
weight = s.get_weight(s.weight, dtype)
|
|
else:
|
|
weight = s.get_weight(s.weight.to(device), dtype)
|
|
weight = comfy.ops.cast_to(weight, dtype, device, non_blocking=non_blocking, copy=False)
|
|
return weight, bias
|
|
|
|
GGMLLayer.cast_bias_weight = _bc250_cast_bias_weight
|
|
|
|
_gguf_patched = True
|
|
logger.warning("[BC-250] GGUF cast_bias_weight patched (CPU dequant, float-only GPU transfer)")
|
|
return True
|
|
|
|
# === IMPORT HOOK for deferred GGUF patching ===
|
|
|
|
class _GGUFImportWatcher:
|
|
def __init__(self):
|
|
self.done = False
|
|
|
|
def find_module(self, fullname, path=None):
|
|
if self.done:
|
|
return None
|
|
if 'GGUF' in fullname and ('dequant' in fullname or 'ops' in fullname):
|
|
return self
|
|
return None
|
|
|
|
def load_module(self, fullname):
|
|
if self in sys.meta_path:
|
|
sys.meta_path.remove(self)
|
|
try:
|
|
import importlib
|
|
mod = importlib.import_module(fullname)
|
|
finally:
|
|
if self not in sys.meta_path:
|
|
sys.meta_path.insert(0, self)
|
|
|
|
if _try_patch_gguf():
|
|
self.done = True
|
|
return mod
|
|
|
|
# === TEXT ENCODER CPU PATCH ===
|
|
_te_patched = False
|
|
|
|
def _try_patch_text_encoder_device():
|
|
global _te_patched
|
|
if _te_patched:
|
|
return True
|
|
mm = sys.modules.get('comfy.model_management')
|
|
if mm is None:
|
|
return False
|
|
mm.text_encoder_device = lambda: torch.device("cpu")
|
|
mm.text_encoder_offload_device = lambda: torch.device("cpu")
|
|
|
|
_te_patched = True
|
|
logger.warning("[BC-250] Text encoder forced to CPU (memory constraint)")
|
|
return True
|
|
|
|
# === VAE CPU-ONLY PATCH ===
|
|
# Force VAE to decode on CPU. VAE is only 320MB — fast enough on CPU for small images.
|
|
# Avoids GPU page-fault hangs from safetensors mmap'd weights on BC-250 (XNACK disabled).
|
|
|
|
_vae_patched = False
|
|
|
|
def _try_patch_vae_cpu():
|
|
"""Patch comfy.sd.VAE to decode and encode on CPU only."""
|
|
global _vae_patched
|
|
if _vae_patched:
|
|
return True
|
|
|
|
sd_mod = sys.modules.get('comfy.sd')
|
|
if sd_mod is None:
|
|
return False
|
|
|
|
VAE = getattr(sd_mod, 'VAE', None)
|
|
if VAE is None:
|
|
return False
|
|
|
|
_original_vae_decode = VAE.decode
|
|
_original_vae_encode = getattr(VAE, 'encode', None)
|
|
|
|
def _bc250_vae_decode(self, samples_in, vae_options={}):
|
|
"""Force VAE decode on CPU — bypass load_models_gpu entirely.
|
|
|
|
Root cause: load_models_gpu tries to unload UNet (5032MB in GPU managed memory)
|
|
before loading VAE. Unloading reads GPU pages that may be swapped → XNACK hang.
|
|
Solution: skip load_models_gpu, run VAE inference directly on CPU.
|
|
"""
|
|
import comfy.model_management as mm
|
|
|
|
logger.warning("[BC-250] VAE decode: CPU-only bypass (skipping load_models_gpu)")
|
|
torch.cuda.empty_cache()
|
|
gc.collect()
|
|
|
|
# Temporarily no-op load_models_gpu to prevent UNet unload hang
|
|
_orig_lmg = mm.load_models_gpu
|
|
mm.load_models_gpu = lambda *a, **kw: None
|
|
|
|
# Save and override device to CPU
|
|
orig_device = getattr(self, 'device', None)
|
|
orig_output_device = getattr(self, 'output_device', None)
|
|
self.device = torch.device("cpu")
|
|
self.output_device = torch.device("cpu")
|
|
|
|
try:
|
|
# Ensure VAE model weights are on CPU
|
|
if hasattr(self, 'first_stage_model'):
|
|
self.first_stage_model.to(torch.device("cpu"))
|
|
self.first_stage_model.eval()
|
|
|
|
# Run the original decode (which now skips load_models_gpu)
|
|
result = _original_vae_decode(self, samples_in, vae_options)
|
|
if isinstance(result, torch.Tensor):
|
|
result = result.to(device=torch.device("cpu"))
|
|
return result
|
|
finally:
|
|
# Restore everything
|
|
mm.load_models_gpu = _orig_lmg
|
|
if orig_device is not None:
|
|
self.device = orig_device
|
|
if orig_output_device is not None:
|
|
self.output_device = orig_output_device
|
|
|
|
VAE.decode = _bc250_vae_decode
|
|
|
|
if _original_vae_encode is not None:
|
|
def _bc250_vae_encode(self, pixel_samples):
|
|
"""Force VAE encode on CPU — same bypass as decode."""
|
|
import comfy.model_management as mm
|
|
logger.warning("[BC-250] VAE encode: CPU-only bypass")
|
|
torch.cuda.empty_cache()
|
|
gc.collect()
|
|
_orig_lmg = mm.load_models_gpu
|
|
mm.load_models_gpu = lambda *a, **kw: None
|
|
orig_device = getattr(self, 'device', None)
|
|
orig_output_device = getattr(self, 'output_device', None)
|
|
self.device = torch.device("cpu")
|
|
self.output_device = torch.device("cpu")
|
|
try:
|
|
if hasattr(self, 'first_stage_model'):
|
|
self.first_stage_model.to(torch.device("cpu"))
|
|
self.first_stage_model.eval()
|
|
pixel_samples = pixel_samples.to(device=torch.device("cpu"), dtype=torch.float32)
|
|
result = _original_vae_encode(self, pixel_samples)
|
|
if isinstance(result, torch.Tensor):
|
|
result = result.to(device=torch.device("cpu"))
|
|
return result
|
|
finally:
|
|
mm.load_models_gpu = _orig_lmg
|
|
if orig_device is not None:
|
|
self.device = orig_device
|
|
if orig_output_device is not None:
|
|
self.output_device = orig_output_device
|
|
|
|
VAE.encode = _bc250_vae_encode
|
|
|
|
_vae_patched = True
|
|
logger.warning("[BC-250] VAE forced to CPU decode/encode (prevents mmap GPU hangs)")
|
|
return True
|
|
|
|
class _SDModuleWatcher:
|
|
"""Patches comfy.sd.VAE after it's imported."""
|
|
def __init__(self):
|
|
self.done = False
|
|
def find_module(self, fullname, path=None):
|
|
if self.done:
|
|
return None
|
|
if fullname == 'comfy.sd':
|
|
return self
|
|
return None
|
|
def load_module(self, fullname):
|
|
if self in sys.meta_path:
|
|
sys.meta_path.remove(self)
|
|
try:
|
|
import importlib
|
|
mod = importlib.import_module(fullname)
|
|
finally:
|
|
if self not in sys.meta_path:
|
|
sys.meta_path.insert(0, self)
|
|
if _try_patch_vae_cpu():
|
|
self.done = True
|
|
return mod
|
|
|
|
class _ModelMgmtWatcher:
|
|
def __init__(self):
|
|
self.done = False
|
|
def find_module(self, fullname, path=None):
|
|
if self.done:
|
|
return None
|
|
if fullname == 'comfy.model_management':
|
|
return self
|
|
return None
|
|
def load_module(self, fullname):
|
|
if self in sys.meta_path:
|
|
sys.meta_path.remove(self)
|
|
try:
|
|
import importlib
|
|
mod = importlib.import_module(fullname)
|
|
finally:
|
|
if self not in sys.meta_path:
|
|
sys.meta_path.insert(0, self)
|
|
if _try_patch_text_encoder_device():
|
|
self.done = True
|
|
return mod
|
|
|
|
# === GPU MEMORY CLEANUP HOOK ===
|
|
# Patch model_management.load_models_gpu to clean up before loading
|
|
|
|
_load_patched = False
|
|
|
|
def _try_patch_load_models():
|
|
"""Add GPU memory cleanup before model loading."""
|
|
global _load_patched
|
|
if _load_patched:
|
|
return True
|
|
|
|
mm = sys.modules.get('comfy.model_management')
|
|
if mm is None:
|
|
return False
|
|
|
|
_original_load = getattr(mm, 'load_models_gpu', None)
|
|
if _original_load is None:
|
|
return False
|
|
|
|
def _bc250_load_models_gpu(models, *args, **kwargs):
|
|
"""Clean GPU cache before loading models to prevent memory pressure hangs."""
|
|
gc.collect()
|
|
torch.cuda.empty_cache()
|
|
return _original_load(models, *args, **kwargs)
|
|
|
|
mm.load_models_gpu = _bc250_load_models_gpu
|
|
_load_patched = True
|
|
logger.warning("[BC-250] GPU memory cleanup hook installed (load_models_gpu)")
|
|
return True
|
|
|
|
# === INSTALL ===
|
|
|
|
def _prewarm_gpu():
|
|
try:
|
|
if not torch.cuda.is_available():
|
|
return
|
|
dummy = torch.zeros(1, device='cuda')
|
|
_ = dummy + 1
|
|
torch.cuda.synchronize()
|
|
del dummy
|
|
torch.cuda.empty_cache()
|
|
logger.warning("[BC-250] GPU pre-warmed (context + allocator ready)")
|
|
except Exception as e:
|
|
logger.warning(f"[BC-250] GPU pre-warm failed: {e}")
|
|
|
|
|
|
def install():
|
|
# Mmap pre-clone patch
|
|
torch.nn.Module._apply = _bc250_safe_apply
|
|
logger.warning("[BC-250] Mmap pre-clone patch installed (XNACK workaround)")
|
|
|
|
# Softmax patches
|
|
torch.nn.functional.softmax = patched_softmax
|
|
torch.Tensor.softmax = patched_tensor_softmax
|
|
torch.nn.functional.scaled_dot_product_attention = patched_sdpa
|
|
logger.warning(f"[BC-250] Softmax monkey-patch installed (threshold={SAFE_SOFTMAX_THRESHOLD})")
|
|
|
|
# GGUF deferred cast_bias_weight patch
|
|
sys.meta_path.insert(0, _GGUFImportWatcher())
|
|
logger.warning("[BC-250] GGUF CPU-dequant hook registered (cast_bias_weight)")
|
|
|
|
# Text encoder CPU patch
|
|
sys.meta_path.insert(0, _ModelMgmtWatcher())
|
|
|
|
# VAE CPU-only patch
|
|
sys.meta_path.insert(0, _SDModuleWatcher())
|
|
|
|
# Try immediate patches if modules already loaded
|
|
_try_patch_gguf()
|
|
_try_patch_text_encoder_device()
|
|
_try_patch_vae_cpu()
|
|
_try_patch_load_models()
|
|
|
|
# Pre-warm GPU
|
|
_prewarm_gpu()
|
|
|
|
install()
|