Merge pull request #3841 from unslothai/fix-vllm-pdl-blackwell

Fix vLLM PDL bug on Blackwell GPUs (B200/B100)
This commit is contained in:
Daniel Han 2026-01-05 04:37:58 -08:00 committed by GitHub
commit d9d26699b5
2 changed files with 144 additions and 0 deletions

View file

@ -126,6 +126,7 @@ from .import_fixes import (
fix_xformers_performance_issue,
fix_vllm_aimv2_issue,
fix_vllm_guided_decoding_params,
fix_vllm_pdl_blackwell,
ignore_logger_messages,
patch_ipykernel_hf_xet,
patch_trackio,
@ -138,6 +139,7 @@ from .import_fixes import (
fix_xformers_performance_issue()
fix_vllm_aimv2_issue()
fix_vllm_guided_decoding_params()
fix_vllm_pdl_blackwell()
ignore_logger_messages()
patch_ipykernel_hf_xet()
patch_trackio()
@ -149,6 +151,7 @@ fix_executorch()
del fix_xformers_performance_issue
del fix_vllm_aimv2_issue
del fix_vllm_guided_decoding_params
del fix_vllm_pdl_blackwell
del ignore_logger_messages
del patch_ipykernel_hf_xet
del patch_trackio

View file

@ -556,3 +556,144 @@ def fix_huggingface_hub():
huggingface_hub.is_offline_mode = (
lambda: huggingface_hub.constants.HF_HUB_OFFLINE
)
def fix_vllm_pdl_blackwell():
"""
Fix vLLM PDL (Programmatic Dependent Launch) bug on Blackwell GPUs (SM100).
The issue: vLLM's LoRA Triton kernels use tl.extra.cuda.gdc_wait() for PDL
optimization on SM90+ GPUs. This fails on SM100 (B200/B100) during CUDA graph
capture because Triton's pipeliner can't handle gdc_wait in complex kernels.
See: https://github.com/vllm-project/vllm/issues/30872
"""
if importlib.util.find_spec("vllm") is None:
return
# Check if any CUDA GPU is SM100 (Blackwell)
try:
import torch
if not torch.cuda.is_available():
return
# Scan all GPUs for SM100 - fix applies globally via env var and monkey-patch
has_sm100 = False
sm100_gpu_name = None
for i in range(torch.cuda.device_count()):
major, minor = torch.cuda.get_device_capability(i)
if major == 10:
has_sm100 = True
sm100_gpu_name = torch.cuda.get_device_name(i)
break
if not has_sm100:
return
except Exception:
return
# Helper to check if module spec exists
def _spec_exists(name):
try:
return importlib.util.find_spec(name) is not None
except (ModuleNotFoundError, ValueError):
return False
# Check if vLLM has the PDL-related modules before doing internet check
has_utils = _spec_exists("vllm.lora.ops.triton_ops.utils")
has_expand_op = _spec_exists("vllm.lora.ops.triton_ops.lora_expand_op")
has_shrink_op = _spec_exists("vllm.lora.ops.triton_ops.lora_shrink_op")
if not has_utils and not has_expand_op and not has_shrink_op:
# Old vLLM version without PDL support - nothing to patch
return
# Check if GitHub issue is closed (fix merged upstream)
issue_closed = False
try:
import socket
import urllib.request
import json as json_module
# Quick internet connectivity check (0.5s timeout)
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
sock.settimeout(0.5)
try:
sock.connect(("api.github.com", 443))
has_internet = True
except (socket.timeout, OSError):
has_internet = False
finally:
sock.close()
if has_internet:
api_url = "https://api.github.com/repos/vllm-project/vllm/issues/30872"
req = urllib.request.Request(
api_url,
headers = {
"User-Agent": "Unsloth-PDL-Fix",
"Accept": "application/vnd.github.v3+json",
},
)
with urllib.request.urlopen(req, timeout = 3) as response:
data = json_module.loads(response.read().decode())
issue_closed = data.get("state") == "closed"
except Exception:
# If we can't check, assume issue is still open (apply fix to be safe)
pass
if issue_closed:
logger.info(
f"Unsloth: SM100 ({sm100_gpu_name}) detected but PDL issue #30872 "
f"is closed - skipping PDL fix"
)
return
# Apply the PDL fix
os.environ["TRITON_DISABLE_PDL"] = "1"
def fake_supports_pdl(*args, **kwargs):
return False
patched = []
# First, patch the source module (utils.py) where supports_pdl is defined.
# This is critical because supports_pdl uses @lru_cache - we must clear the
# cache to prevent stale cached results from the original function.
try:
utils_module = importlib.import_module("vllm.lora.ops.triton_ops.utils")
if hasattr(utils_module, "supports_pdl"):
original_fn = utils_module.supports_pdl
if hasattr(original_fn, "cache_clear"):
original_fn.cache_clear()
utils_module.supports_pdl = fake_supports_pdl
patched.append("utils")
except (ImportError, ModuleNotFoundError, AttributeError):
pass
# Also patch the consumer modules that import supports_pdl from utils.
# This ensures the patched function is used even if the module was already
# imported before this fix runs.
consumer_modules = {
"lora_expand_op": "vllm.lora.ops.triton_ops.lora_expand_op",
"lora_shrink_op": "vllm.lora.ops.triton_ops.lora_shrink_op",
"fused_moe_lora_op": "vllm.lora.ops.triton_ops.fused_moe_lora_op",
}
for name, path in consumer_modules.items():
try:
module = importlib.import_module(path)
if hasattr(module, "supports_pdl"):
module.supports_pdl = fake_supports_pdl
patched.append(name)
except (ImportError, ModuleNotFoundError, AttributeError):
pass
if patched:
logger.info(
f"Unsloth: Applied PDL fix for SM100 ({sm100_gpu_name}) - "
f"patched: {', '.join(patched)}"
)
else:
# Just set the env var - vLLM might be an older version without supports_pdl
logger.info(f"Unsloth: Set TRITON_DISABLE_PDL=1 for SM100 ({sm100_gpu_name})")