This repository has been archived on 2026-08-19. You can view files and clone it. You cannot open issues or pull requests or push a commit.
Files
ROCm-Research-Archive/_TestScripts/Scripts and Tests/bc250_softmax_patch_v11.py
T
2026-08-20 00:45:43 +02:00

540 lines
20 KiB
Python

"""
BC-250 gfx1010 Comprehensive Monkey-Patch v11
1. Softmax: manual impl for dim > threshold (VGPR overflow fix)
2. SDPA: manual impl for large sequences
3. GGUF: GPU dequant with weight cache (eliminates per-step dequant)
4. Mmap: pre-clones non-GGUF tensor data before GPU transfer
5. GPU: pre-warms context and caching allocator
6. CLIP: forces text encoder to CPU (memory constraint)
7. VAE: GPU fp16 decode with persistent caching (shared memory APU)
8. Threads: all CPU cores for intra-op parallelism
9. Preload: background warmup prompt on server start
NOTE: mlockall REMOVED — on APU with shared memory, pinning 10GB of mmap'd
GGUF files leaves no room for GPU GTT allocations → OOM kill.
The kernel page cache handles this correctly without mlockall.
v11 changes vs v10:
- GPU dequant instead of CPU (GGUF dequant ops are pure PyTorch, run on GPU)
- Weight cache: dequanted fp16 weights cached per-layer, reused across steps
- mlockall() to pin process memory in RAM (no zram/swap penalty)
- Removed duplicate preload section
- Clean rewrite
BC-250 APU / gfx1010: XNACK disabled, no GPU page fault handling.
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
logger = logging.getLogger(__name__)
# === THREAD CONFIGURATION ===
_NUM_THREADS = int(os.environ.get("BC250_NUM_THREADS", str(os.cpu_count() or 12)))
torch.set_num_threads(_NUM_THREADS)
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 (XNACK workaround).
Note: GGMLTensor.clone() returns self, so GGUF weights are unaffected.
They're handled by GGMLTensor.to() which preserves metadata."""
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: shape={list(input.shape)}, dim={dim}")
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: shape={list(self.shape)}, dim={dim}")
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: Q={list(query.shape)}, S={S}")
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 WEIGHT CACHE + GPU DEQUANT ===
_gguf_patched = False
_weight_cache = {}
_weight_cache_bytes = 0
_WEIGHT_CACHE_MB = int(os.environ.get("BC250_WEIGHT_CACHE_MB", "0"))
def _try_patch_gguf():
"""Patch GGMLLayer.cast_bias_weight: GPU dequant + weight caching."""
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
_original_cast = getattr(GGMLLayer, 'cast_bias_weight', None)
cache_budget = _WEIGHT_CACHE_MB * 1024 * 1024
def _bc250_cast_bias_weight(s, input=None, dtype=None, device=None, bias_dtype=None):
"""GPU dequant with optional weight caching.
With --highvram, GGUF weights are already on GPU. Dequant happens
via PyTorch tensor ops on GPU (parallel) instead of CPU (sequential).
If weight cache is enabled (BC250_WEIGHT_CACHE_MB > 0), dequanted
weights are cached per-layer to eliminate dequant on steps 2+.
"""
global _weight_cache_bytes
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)
# Check weight cache
if cache_budget > 0:
cache_key = id(s)
cached = _weight_cache.get(cache_key)
if cached is not None:
return cached
# Bias
bias = None
if s.bias is not None:
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)
# Weight: .to(device) moves GGMLTensor to GPU, get_weight dequants on GPU
weight = s.get_weight(s.weight.to(device), dtype)
weight = comfy.ops.cast_to(weight, dtype, device, non_blocking=non_blocking, copy=False)
# Cache if within budget
if cache_budget > 0:
entry_bytes = weight.nelement() * weight.element_size()
if bias is not None:
entry_bytes += bias.nelement() * bias.element_size()
if _weight_cache_bytes + entry_bytes <= cache_budget:
_weight_cache[cache_key] = (weight, bias)
_weight_cache_bytes += entry_bytes
return weight, bias
GGMLLayer.cast_bias_weight = _bc250_cast_bias_weight
_gguf_patched = True
cache_str = f", weight cache={_WEIGHT_CACHE_MB}MB" if cache_budget > 0 else ""
logger.warning(f"[BC-250] GGUF patched: GPU dequant{cache_str}")
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")
return True
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
# === VAE GPU FP16 WITH CACHING ===
# BC-250 = APU with shared memory. GPU VRAM = CPU RAM = same physical pool.
# No OOM risk from "using VRAM" — it's all the same 16GB.
# GPU fp16 VAE is ~10x faster than CPU float32.
_vae_patched = False
_vae_cached = False
def _try_patch_vae_gpu():
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
def _ensure_vae_on_gpu_f16(self):
"""Move VAE to GPU fp16 once, keep it cached. Shared memory = no OOM risk."""
global _vae_cached
gpu = torch.device("cuda")
try:
p = next(self.first_stage_model.parameters())
already_ready = _vae_cached and p.device.type == 'cuda' and p.dtype == torch.float16
except StopIteration:
already_ready = False
if not already_ready:
logger.warning("[BC-250] Loading VAE to GPU fp16 (shared memory, will stay cached)")
# Bypass _bc250_safe_apply (mmap pre-clone) — VAE is safetensors, not GGUF
old_apply = torch.nn.Module._apply
torch.nn.Module._apply = _original_module_apply
try:
self.first_stage_model.half().cuda()
finally:
torch.nn.Module._apply = old_apply
self.first_stage_model.eval()
_vae_cached = True
self.disable_offload = True
def _bc250_vae_decode(self, samples_in, vae_options={}):
t0 = _time.time()
self.throw_exception_if_invalid()
if self.latent_dim == 2 and samples_in.ndim == 5:
samples_in = samples_in[:, :, 0]
# Free GPU memory from UNET before loading VAE
mm = sys.modules.get('comfy.model_management')
if mm:
mm.unload_all_models()
gc.collect()
torch.cuda.empty_cache()
_ensure_vae_on_gpu_f16(self)
pixel_samples = None
with torch.no_grad():
for x in range(samples_in.shape[0]):
sample = samples_in[x:x+1].to(torch.float16).cuda()
decoded = self.first_stage_model.decode(sample, **vae_options)
if decoded.ndim == 5:
decoded = decoded[:, :, 0]
out = self.process_output(decoded.float().cpu())
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
pixel_samples = pixel_samples.movedim(1, -1)
elapsed = _time.time() - t0
logger.warning(f"[BC-250] VAE decode (GPU fp16): {elapsed:.1f}s")
return pixel_samples
def _bc250_vae_encode(self, pixel_samples):
t0 = _time.time()
self.throw_exception_if_invalid()
mm = sys.modules.get('comfy.model_management')
if mm:
mm.unload_all_models()
gc.collect()
torch.cuda.empty_cache()
_ensure_vae_on_gpu_f16(self)
with torch.no_grad():
pixels_in = self.process_input(pixel_samples).to(torch.float16).cuda()
result = self.first_stage_model.encode(pixels_in).float().cpu()
logger.warning(f"[BC-250] VAE encode (GPU fp16): {_time.time() - t0:.1f}s")
return result
VAE.decode = _bc250_vae_decode
VAE.encode = _bc250_vae_encode
_vae_patched = True
logger.warning("[BC-250] VAE patched: GPU fp16 (shared memory = zero OOM risk)")
return True
class _SDModuleWatcher:
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_gpu():
self.done = True
return mod
# === GPU MEMORY CLEANUP HOOK ===
_load_patched = False
def _try_patch_load_models():
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):
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")
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():
import urllib.request
import urllib.error
url = f"http://127.0.0.1:{_PRELOAD_PORT}"
logger.warning("[BC-250] Preload: waiting for 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, skip")
return
logger.warning("[BC-250] Preload: server ready, submitting warmup...")
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 queued (id={prompt_id})")
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 cached. Ready.")
return
except Exception:
pass
logger.warning("[BC-250] Preload: warmup timed out (5min)")
except Exception as e:
logger.warning(f"[BC-250] Preload 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 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")
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")
# 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 patch installed (threshold={SAFE_SOFTMAX_THRESHOLD})")
# Deferred patches via import hooks
sys.meta_path.insert(0, _GGUFImportWatcher())
sys.meta_path.insert(0, _ModelMgmtWatcher())
sys.meta_path.insert(0, _SDModuleWatcher())
# Try immediate patches
_try_patch_gguf()
_try_patch_text_encoder_device()
_try_patch_vae_gpu()
_try_patch_load_models()
_prewarm_gpu()
_start_preload_thread()
install()