Files
mineru-rocm/docker/scripts/patch_vllm_platform.py
T
2026-06-23 13:58:15 +08:00

310 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""vllm 平台检测补丁
问题 1:amdsmi 不可用时 platform 回退到 torch.version.hip
问题 2:rocm.py 中 logger.warning_once() 导致循环导入(默认不再改写;仅保留为显式开关)
"""
import os
import re
import sys
import sysconfig
import traceback
VLLM_DIR = '/opt/vllm/vllm'
VLLM_SRC_DIR = '/opt/vllm'
def patch6_init_platform_fallback():
"""补丁 6:platforms/__init__.py —— torch.version.hip 兜底
vLLM 的 is_rocm 检测依赖 import amdsmi;amdsmi 未装时 is_rocm=False,
current_platform 落到 UnspecifiedPlatform,device_type='' →
RuntimeError: Device string must not be empty。
本补丁在 RocmPlatform 的 return 语句前注入 torch.version.hip 兜底。
采用行结构 + return 定位,不依赖精确字符串匹配
(vllm main 分支重构频繁,固定字符串匹配已多次失效)。
"""
f = os.path.join(VLLM_DIR, 'platforms', '__init__.py')
lines = open(f).read().splitlines(keepends=True)
# 兼容单/双引号两种写法(vllm 不同版本可能不同)
return_idx = None
for i, ln in enumerate(lines):
if 'RocmPlatform' in ln and 'return' in ln and 'is_rocm' in ln:
return_idx = i
break
if return_idx is None:
print('Patch 6: RocmPlatform return statement not found; skipping.')
return
if 'torch.version.hip is not None' in ''.join(lines):
print('Patch 6: already applied (torch.version.hip fallback present).')
return
indent = ' ' * (len(lines[return_idx]) - len(lines[return_idx].lstrip()))
inject = (
f"{indent}# amdsmi fallback: also check torch.version.hip\n"
f"{indent}if not is_rocm:\n"
f"{indent} try:\n"
f"{indent} import torch as _torch\n"
f"{indent} if _torch.version.hip is not None:\n"
f"{indent} is_rocm = True\n"
f"{indent} except Exception:\n"
f"{indent} pass\n"
)
lines.insert(return_idx, inject)
open(f, 'w').write(''.join(lines))
print('Patch 6: __init__.py torch.version.hip fallback applied.')
def patch7_rocm_break_import_cycle():
"""补丁 7:platforms/rocm.py —— logger.warning_once → sys.stderr.write
默认启用安全的“整段 except 块”替换,避免 logger.warning_once() 在平台检测
期间触发循环导入,导致 current_platform 落到 UnspecifiedPlatform。
如需禁用,设置 VLLM_PATCH_ROCM_WARNING_ONCE=0。
"""
if os.environ.get('VLLM_PATCH_ROCM_WARNING_ONCE') == '0':
print('Patch 7: skipped (VLLM_PATCH_ROCM_WARNING_ONCE=0).')
return
f = os.path.join(VLLM_DIR, 'platforms', 'rocm.py')
c = open(f).read()
if 'amdsmi unavailable, using torch.cuda fallback' in c:
print('Patch 7: already applied.')
return
pattern = re.compile(
r'(\n\s*except Exception as e:\n'
r'\s*logger\.debug\("Failed to get GCN arch via amdsmi: %s", e\)\n'
r'\s*logger\.warning_once\(\n'
r'(?:\s*"[^"]*"\n)+'
r'\s*\)\n)',
re.MULTILINE,
)
def repl(match):
indent = re.search(r'\n(\s*)except Exception as e:', match.group(1)).group(1)
body_indent = indent + ' '
return (
f'\n{indent}except Exception as e:\n'
f'{body_indent}import sys as _sys\n'
f'{body_indent}_sys.stderr.write('
'"vLLM ROCm: amdsmi unavailable, using torch.cuda fallback for GPU detection\\n")\n'
)
c2, n = pattern.subn(repl, c, count=1)
if c2 != c:
open(f, 'w').write(c2)
print('Patch 7: rocm.py logger.warning_once circular import patch applied.')
else:
print('Patch 7: already applied or pattern not found.')
def patch9_registry_model_impl_compat():
"""补丁 9:registry.py —— vllm main 与 transformers v4 兼容
问题:vllm main 的 ModelRegistry 大量访问 model_config.model_impl,
该属性是 vllm ModelConfig 的字段(默认 "auto"),但 inspect 链路传入的
有时是 transformers config 对象(如 Qwen2VLConfig),v4 没有此属性 →
AttributeError → "Model architectures [...] failed to be inspected"。
同时 _try_resolve_transformers 末尾调用 model_config._get_transformers_backend_cls(),
v4 的 config 也没有该方法。
根因:mineru[core] 锁定 transformers<5.0.0(v4),vllm main 期望 v5。
本补丁把所有 model_config.model_impl 访问改成 getattr 兜底(缺属性时当 "auto",
走 fallback 分支匹配 vllm 注册表),并给 _get_transformers_backend_cls 加兜底。
采用逐处 getattr 替换,不依赖方法定位(比方法注入更可靠)。
"""
f = os.path.join(VLLM_DIR, 'model_executor', 'models', 'registry.py')
if not os.path.exists(f):
print('Patch 9: registry.py not found; skipping.')
return
c = open(f).read()
if '# mineru-rocm: model_impl v4 compat' in c:
print('Patch 9: already applied (model_impl v4 compat present).')
return
# 把 model_config.model_impl 访问替换为 getattr 兜底
before = c.count('model_config.model_impl')
c2 = c.replace(
'model_config.model_impl',
'getattr(model_config, "model_impl", "auto")'
)
replaced = before - c2.count('model_config.model_impl')
# _get_transformers_backend_cls 兜底:v4 无此方法
c2 = c2.replace(
'return model_config._get_transformers_backend_cls()',
'return getattr(model_config, "_get_transformers_backend_cls", lambda: None)()'
)
# 写入幂等标记(注释,便于重入检测)
c2 = '# mineru-rocm: model_impl v4 compat\n' + c2
if c2 != c:
open(f, 'w').write(c2)
print(f'Patch 9: registry.py model_impl v4 compat applied '
f'({replaced} access(es) wrapped).')
else:
print('Patch 9: no changes applied (pattern not found).')
def ensure_vllm_dist_info():
"""为 /opt/vllm 源码导入创建最小 dist-info。
vLLM 平台检测会通过 importlib.metadata 查询 vllm 分发元数据/entry points。
如果只靠 PYTHONPATH 导入源码,没有 dist-info,就会出现:
"The vLLM package was not found...",并可能得到 UnspecifiedPlatform。
注意:仅当 rocm_platform_plugin 函数实际存在时才注册 entry_point,
否则 vLLM 会因 AttributeError 回退到 UnspecifiedPlatform。
"""
try:
import vllm
except Exception:
return
version = getattr(vllm, '__version__', '0.1.dev1') or '0.1.dev1'
# PEP440 合法性校验:vllm 在 _version.py 缺失时会回退到 'dev',
# 这不是合法 PEP440,会触发下游 packaging.version.parse 抛 InvalidVersion
# (例如 mineru.backend.vlm.utils:set_default_gpu_memory_utilization)。
try:
from packaging.version import Version as _PEP440Version
_PEP440Version(version)
except Exception:
print(f'WARNING: vllm.__version__={version!r} is not PEP440-compliant; '
f'falling back to "0.11.0" for dist-info.')
version = '0.11.0'
try:
vllm.__version__ = version
except Exception:
pass
site_packages = sysconfig.get_paths().get('purelib')
if not site_packages:
return
dist_info = os.path.join(site_packages, f'vllm-{version}.dist-info')
os.makedirs(dist_info, exist_ok=True)
metadata = os.path.join(dist_info, 'METADATA')
if not os.path.exists(metadata):
with open(metadata, 'w') as fp:
fp.write(
'Metadata-Version: 2.1\n'
'Name: vllm\n'
f'Version: {version}\n'
'Summary: vLLM source tree mounted at /opt/vllm\n'
)
with open(os.path.join(dist_info, 'top_level.txt'), 'w') as fp:
fp.write('vllm\n')
with open(os.path.join(dist_info, 'INSTALLER'), 'w') as fp:
fp.write('mineru-rocm-runtime\n')
# 仅当 rocm_platform_plugin 实际存在时才注册,避免 vLLM 加载失败
entry_points = ''
try:
from vllm.platforms.rocm import rocm_platform_plugin # noqa: F401
entry_points += (
'[vllm.platform_plugins]\n'
'rocm = vllm.platforms.rocm:rocm_platform_plugin\n'
)
print('rocm_platform_plugin found, registering entry point.')
except (ImportError, AttributeError, SystemExit, KeyboardInterrupt):
print('WARNING: rocm_platform_plugin not found in vllm.platforms.rocm; '
'skipping entry_point registration.')
except BaseException as e:
print(f'WARNING: rocm_platform_plugin check failed ({type(e).__name__}: {e}); '
'skipping entry_point registration.')
with open(os.path.join(dist_info, 'entry_points.txt'), 'w') as fp:
fp.write(entry_points)
with open(os.path.join(dist_info, 'RECORD'), 'w') as fp:
fp.write('')
print(f'vllm dist-info ensured: {dist_info}')
def ensure_vllm_installed():
"""确保 vLLM Python 包在当前虚拟环境中可导入。
某些 MinerU 依赖安装流程后,site-packages 中可能缺少 vLLM 分发元数据,但
/opt/vllm 源码仍在。此时直接把 /opt/vllm 加入 PYTHONPATH/sys.path 再导入,
避免 editable 安装触发 pyproject 元数据校验失败。
"""
try:
import vllm # noqa: F401
return
except ModuleNotFoundError as exc:
if exc.name != 'vllm':
raise
if not os.path.isdir(VLLM_SRC_DIR):
raise ModuleNotFoundError(
'vllm is not installed and /opt/vllm source directory is missing'
)
print('vllm package missing; trying PYTHONPATH fallback with /opt/vllm...')
current = os.environ.get('PYTHONPATH', '')
paths = [p for p in current.split(':') if p]
if VLLM_SRC_DIR not in paths:
os.environ['PYTHONPATH'] = f"{VLLM_SRC_DIR}:{current}" if current else VLLM_SRC_DIR
if VLLM_SRC_DIR not in sys.path:
sys.path.insert(0, VLLM_SRC_DIR)
import vllm # noqa: F401
print('vllm import recovered via PYTHONPATH fallback.')
def install_sitecustomize_force_rocm():
"""已弃用:sitecustomize 时机问题(current_platform 是 lazy init,
sitecustomize 触发提前 resolve 时补丁 6 尚未应用)导致兜底无效。
保留空壳仅为兼容旧调用,实际不做任何事。
平台检测统一由补丁 6(rocm_platform_plugin torch.version.hip 兜底)解决。
"""
# 清理历史遗留的 sitecustomize 片段(幂等)
site_packages = sysconfig.get_paths().get('purelib')
if site_packages:
sc = os.path.join(site_packages, 'sitecustomize.py')
marker = '# mineru-rocm: force rocm platform'
if os.path.exists(sc) and marker in open(sc).read():
# 重写文件,移除我们的片段
c = open(sc).read()
# 片段从 marker 行开始到文件末尾
idx = c.find(marker)
# 回退到 marker 前的换行
while idx > 0 and c[idx - 1] == '\n':
idx -= 1
c = c[:idx].rstrip() + '\n'
open(sc, 'w').write(c)
print('Patch 8: removed legacy sitecustomize force-rocm snippet.')
def main():
patch6_init_platform_fallback()
patch7_rocm_break_import_cycle()
patch9_registry_model_impl_compat()
install_sitecustomize_force_rocm()
try:
ensure_vllm_installed()
ensure_vllm_dist_info()
import vllm
from vllm.platforms import current_platform
platform_name = type(current_platform).__name__
print(f'vllm runtime import OK: {vllm.__version__}, platform={platform_name}')
if platform_name == 'UnspecifiedPlatform':
print(
'WARNING: vLLM platform detection returned UnspecifiedPlatform. '
'This is expected on containers without GPU access (router/gradio). '
'On worker containers, check: amdsmi, /dev/kfd, /dev/dri, '
'and vllm.platform_plugins metadata.'
)
except Exception:
print('WARNING: vllm runtime import failed after platform patches (non-fatal):')
traceback.print_exc()
print('vllm platform patches done.')
if __name__ == '__main__':
main()