""" BC-250 gfx1010 Comprehensive Monkey-Patch v10 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. VAE decode on CPU float32 (bypasses GPU managed memory issues) 8. Sets torch threads to all CPU cores (faster CPU ops + VAE decode) 9. Caches VAE model on CPU (avoids reload each generation) 10. Startup preloading: submits warmup prompt to preload all models on boot v10 changes: - Background warmup thread submits 64x64 @ 1 step prompt after server starts - All models (CLIP, UNET, VAE) preloaded before user interaction - Models configurable via BC250_PRELOAD_* env vars 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 import threading import json import time as _time import threading import json import time as _time logger = logging.getLogger(__name__) # === THREAD CONFIGURATION === # BC-250 has 12 threads (6C/12T Zen2). Use all for CPU-heavy work (VAE, CLIP, dequant). _NUM_THREADS = int(os.environ.get("BC250_NUM_THREADS", str(os.cpu_count() or 12))) torch.set_num_threads(_NUM_THREADS) # Note: set_num_interop_threads must be called before any parallel op, skip to avoid deadlock logger.warning(f"[BC-250] Torch threads: intra-op={_NUM_THREADS}") SAFE_SOFTMAX_THRESHOLD = int(os.environ.get("BC250_SOFTMAX_THRESHOLD", "4096")) _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: if not getattr(patched_softmax, '_logged', False): logger.warning(f"[BC-250] Manual F.softmax triggered: shape={list(input.shape)}, dim={dim}, threshold={SAFE_SOFTMAX_THRESHOLD}") patched_softmax._logged = True 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: if not getattr(patched_tensor_softmax, '_logged', False): logger.warning(f"[BC-250] Manual softmax triggered: shape={list(self.shape)}, dim={dim}, threshold={SAFE_SOFTMAX_THRESHOLD}") patched_tensor_softmax._logged = True 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: if not getattr(patched_sdpa, '_logged', False): logger.warning(f"[BC-250] Manual SDPA triggered: Q={list(query.shape)}, K={list(key.shape)}, S={S}, threshold={SAFE_SOFTMAX_THRESHOLD}") patched_sdpa._logged = True 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, send float results to GPU. Cannot use .to(device) on quantized GGUF tensors from mmap'd files (GPU copy shader hangs on non-resident pages, XNACK disabled). Dequant to float on CPU, then transfer dequantized float 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)") 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 FLOAT32 DECODE PATCH === # Decode VAE on CPU using float32 (not fp16). fp16 on CPU is emulated (10x slower). # Cannot use GPU because UNet managed memory blocks new GPU allocations (XNACK disabled). # 320MB VAE at float32 = 640MB RAM. For 256x256: ~2-3 min on 12-thread CPU. _vae_patched = False _vae_cached = False # Track whether VAE is already loaded to CPU float32 def _try_patch_vae_cpu(): """Patch comfy.sd.VAE to decode on CPU with float32, with persistent caching.""" 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_encode = getattr(VAE, 'encode', None) def _ensure_vae_on_cpu_f32(self): """Move VAE to CPU float32 once, then keep it cached.""" global _vae_cached if not _vae_cached or next(self.first_stage_model.parameters()).dtype != torch.float32: logger.warning("[BC-250] Loading VAE to CPU float32 (will stay cached)") self.first_stage_model.to(torch.float32).to(torch.device("cpu")) self.first_stage_model.eval() _vae_cached = True # Prevent ComfyUI model_management from offloading the VAE self.disable_offload = True def _bc250_vae_decode(self, samples_in, vae_options={}): """CPU float32 VAE decode — bypasses GPU managed memory entirely. The UNet (5032MB managed memory) blocks new GPU allocations when its pages are swapped by the OS (XNACK disabled on gfx1010). float32 on CPU is ~5x faster than fp16 (which requires emulation). VAE stays cached on CPU after first load — no re-conversion needed. """ import time t0 = time.time() logger.warning("[BC-250] VAE decode: CPU float32 (cached)") self.throw_exception_if_invalid() if self.latent_dim == 2 and samples_in.ndim == 5: samples_in = samples_in[:, :, 0] cpu = torch.device("cpu") _ensure_vae_on_cpu_f32(self) pixel_samples = None with torch.no_grad(): for x in range(samples_in.shape[0]): sample = samples_in[x:x+1].to(torch.float32) decoded = self.first_stage_model.decode(sample, **vae_options) # Squeeze temporal dim for 3D video autoencoders (single image) if decoded.ndim == 5: decoded = decoded[:, :, 0] out = self.process_output(decoded.float()) if pixel_samples is None: pixel_samples = torch.empty( (samples_in.shape[0],) + tuple(out.shape[1:]), device=cpu ) pixel_samples[x:x+1] = out del decoded, sample # NCHW → NHWC (same as original ComfyUI VAE.decode line 977) pixel_samples = pixel_samples.movedim(1, -1) elapsed = time.time() - t0 logger.warning(f"[BC-250] VAE decode complete in {elapsed:.1f}s") return pixel_samples VAE.decode = _bc250_vae_decode if _original_vae_encode is not None: def _bc250_vae_encode(self, pixel_samples): """CPU float32 VAE encode (cached).""" import time t0 = time.time() logger.warning("[BC-250] VAE encode: CPU float32 (cached)") self.throw_exception_if_invalid() _ensure_vae_on_cpu_f32(self) with torch.no_grad(): pixels_in = self.process_input(pixel_samples).to(torch.float32) result = self.first_stage_model.encode(pixels_in).float() elapsed = time.time() - t0 logger.warning(f"[BC-250] VAE encode complete in {elapsed:.1f}s") return result VAE.encode = _bc250_vae_encode _vae_patched = True logger.warning("[BC-250] VAE patched: CPU float32 decode/encode with caching (bypass GPU managed memory)") 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 # === STARTUP PRELOAD === _PRELOAD_CLIP = os.environ.get("BC250_PRELOAD_CLIP", "Qwen_3_4b-Q8_0.gguf") _PRELOAD_UNET = os.environ.get("BC250_PRELOAD_UNET", "z_image_turbo-Q5_K_S.gguf") _PRELOAD_VAE = os.environ.get("BC250_PRELOAD_VAE", "ae.safetensors") _PRELOAD_PORT = int(os.environ.get("BC250_PRELOAD_PORT", "8188")) _PRELOAD_ENABLED = os.environ.get("BC250_PRELOAD", "1") == "1" def _preload_models(): """Background thread: wait for ComfyUI server, then submit a warmup prompt.""" import urllib.request import urllib.error url = f"http://127.0.0.1:{_PRELOAD_PORT}" # Wait for server to be ready (max 120s) logger.warning("[BC-250] Preload: waiting for ComfyUI server...") for _ in range(240): try: urllib.request.urlopen(f"{url}/api/system_stats", timeout=2) break except (urllib.error.URLError, OSError, ConnectionRefusedError): _time.sleep(0.5) else: logger.warning("[BC-250] Preload: server not ready after 120s, skipping") return logger.warning("[BC-250] Preload: server ready, submitting warmup prompt...") warmup = { "1": {"class_type": "CLIPLoaderGGUF", "inputs": {"clip_name": _PRELOAD_CLIP, "type": "lumina2"}}, "2": {"class_type": "CLIPTextEncode", "inputs": {"text": "warmup", "clip": ["1", 0]}}, "3": {"class_type": "CLIPTextEncode", "inputs": {"text": "", "clip": ["1", 0]}}, "4": {"class_type": "UnetLoaderGGUF", "inputs": {"unet_name": _PRELOAD_UNET}}, "5": {"class_type": "EmptyLatentImage", "inputs": {"width": 64, "height": 64, "batch_size": 1}}, "6": {"class_type": "KSampler", "inputs": { "seed": 1, "steps": 1, "cfg": 1.0, "sampler_name": "euler", "scheduler": "normal", "denoise": 1.0, "model": ["4", 0], "positive": ["2", 0], "negative": ["3", 0], "latent_image": ["5", 0] }}, "7": {"class_type": "VAELoader", "inputs": {"vae_name": _PRELOAD_VAE}}, "8": {"class_type": "VAEDecode", "inputs": {"samples": ["6", 0], "vae": ["7", 0]}}, "9": {"class_type": "SaveImage", "inputs": {"filename_prefix": "_warmup", "images": ["8", 0]}} } payload = json.dumps({"prompt": warmup}).encode("utf-8") req = urllib.request.Request( f"{url}/api/prompt", data=payload, headers={"Content-Type": "application/json"}, method="POST" ) try: resp = urllib.request.urlopen(req, timeout=10) data = json.loads(resp.read()) prompt_id = data.get("prompt_id", "unknown") logger.warning(f"[BC-250] Preload: warmup prompt queued (id={prompt_id})") # Wait for completion (max 5min) for _ in range(300): _time.sleep(1) try: hist_resp = urllib.request.urlopen(f"{url}/api/history/{prompt_id}", timeout=5) hist = json.loads(hist_resp.read()) if prompt_id in hist: logger.warning("[BC-250] Preload: all models loaded and cached. Ready for user prompts.") return except Exception: pass logger.warning("[BC-250] Preload: warmup timed out after 5min") except Exception as e: logger.warning(f"[BC-250] Preload: warmup failed: {e}") def _start_preload_thread(): if not _PRELOAD_ENABLED: logger.warning("[BC-250] Preload: disabled (BC250_PRELOAD=0)") return t = threading.Thread(target=_preload_models, daemon=True, name="BC250-Preload") t.start() logger.warning("[BC-250] Preload: background warmup thread started") # === STARTUP PRELOAD === _PRELOAD_CLIP = os.environ.get("BC250_PRELOAD_CLIP", "Qwen_3_4b-Q8_0.gguf") _PRELOAD_UNET = os.environ.get("BC250_PRELOAD_UNET", "z_image_turbo-Q5_K_S.gguf") _PRELOAD_VAE = os.environ.get("BC250_PRELOAD_VAE", "ae.safetensors") _PRELOAD_PORT = int(os.environ.get("BC250_PRELOAD_PORT", "8188")) _PRELOAD_ENABLED = os.environ.get("BC250_PRELOAD", "1") == "1" def _preload_models(): """Background thread: wait for ComfyUI server, then submit a warmup prompt.""" import urllib.request import urllib.error url = f"http://127.0.0.1:{_PRELOAD_PORT}" # Wait for server to be ready (max 120s) logger.warning("[BC-250] Preload: waiting for ComfyUI server...") for _ in range(240): try: urllib.request.urlopen(f"{url}/api/system_stats", timeout=2) break except (urllib.error.URLError, OSError, ConnectionRefusedError): _time.sleep(0.5) else: logger.warning("[BC-250] Preload: server not ready after 120s, skipping") return logger.warning("[BC-250] Preload: server ready, submitting warmup prompt...") warmup = { "1": {"class_type": "CLIPLoaderGGUF", "inputs": {"clip_name": _PRELOAD_CLIP, "type": "lumina2"}}, "2": {"class_type": "CLIPTextEncode", "inputs": {"text": "warmup", "clip": ["1", 0]}}, "3": {"class_type": "CLIPTextEncode", "inputs": {"text": "", "clip": ["1", 0]}}, "4": {"class_type": "UnetLoaderGGUF", "inputs": {"unet_name": _PRELOAD_UNET}}, "5": {"class_type": "EmptyLatentImage", "inputs": {"width": 64, "height": 64, "batch_size": 1}}, "6": {"class_type": "KSampler", "inputs": { "seed": 1, "steps": 1, "cfg": 1.0, "sampler_name": "euler", "scheduler": "normal", "denoise": 1.0, "model": ["4", 0], "positive": ["2", 0], "negative": ["3", 0], "latent_image": ["5", 0] }}, "7": {"class_type": "VAELoader", "inputs": {"vae_name": _PRELOAD_VAE}}, "8": {"class_type": "VAEDecode", "inputs": {"samples": ["6", 0], "vae": ["7", 0]}}, "9": {"class_type": "SaveImage", "inputs": {"filename_prefix": "_warmup", "images": ["8", 0]}} } payload = json.dumps({"prompt": warmup}).encode("utf-8") req = urllib.request.Request( f"{url}/api/prompt", data=payload, headers={"Content-Type": "application/json"}, method="POST" ) try: resp = urllib.request.urlopen(req, timeout=10) data = json.loads(resp.read()) prompt_id = data.get("prompt_id", "unknown") logger.warning(f"[BC-250] Preload: warmup prompt queued (id={prompt_id})") # Wait for completion (max 5min) for _ in range(300): _time.sleep(1) try: hist_resp = urllib.request.urlopen(f"{url}/api/history/{prompt_id}", timeout=5) hist = json.loads(hist_resp.read()) if prompt_id in hist: logger.warning("[BC-250] Preload: all models loaded and cached. Ready for user prompts.") return except Exception: pass logger.warning("[BC-250] Preload: warmup timed out after 5min") except Exception as e: logger.warning(f"[BC-250] Preload: warmup failed: {e}") def _start_preload_thread(): if not _PRELOAD_ENABLED: logger.warning("[BC-250] Preload: disabled (BC250_PRELOAD=0)") return t = threading.Thread(target=_preload_models, daemon=True, name="BC250-Preload") t.start() logger.warning("[BC-250] Preload: background warmup thread started") # === 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() # Start background preload thread _start_preload_thread() install()