Files
mineru-rocm/docker/scripts/patch_vllm_platform.py
T
2026-06-12 10:56:36 +08:00

75 lines
2.7 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 traceback
VLLM_DIR = '/opt/vllm/vllm'
def patch6_init_platform_fallback():
"""补丁 6:platforms/__init__.py —— torch.version.hip 兜底"""
f = os.path.join(VLLM_DIR, 'platforms', '__init__.py')
c = open(f).read()
old = " return 'vllm.platforms.rocm.RocmPlatform' if is_rocm else None"
new = (" # amdsmi fallback: also check torch.version.hip\n"
" if not is_rocm:\n"
" try:\n"
" import torch\n"
" if torch.version.hip is not None:\n"
" is_rocm = True\n"
" except Exception:\n"
" pass\n"
" return 'vllm.platforms.rocm.RocmPlatform' if is_rocm else None")
c2 = c.replace(old, new)
if c2 != c:
open(f, 'w').write(c2)
print('Patch 6: __init__.py platform fallback applied.')
else:
print('Patch 6: already applied or pattern not found.')
def patch7_rocm_break_import_cycle():
"""补丁 7:platforms/rocm.py —— logger.warning_once → sys.stderr.write
这个补丁会直接改写 vLLM 源码中的函数调用。部分 vLLM 版本里
logger.warning_once(...) 参数不一定兼容 sys.stderr.write(...),改写后可能
导致 import vllm / import vllm.platforms.rocm 失败。MinerU 会把这种导入失败
包装成 "Please install vllm",因此默认禁用,只在显式设置环境变量时启用。
"""
if os.environ.get('VLLM_PATCH_ROCM_WARNING_ONCE') != '1':
print('Patch 7: skipped (set VLLM_PATCH_ROCM_WARNING_ONCE=1 to enable).')
return
f = os.path.join(VLLM_DIR, 'platforms', 'rocm.py')
c = open(f).read()
old = 'logger.warning_once('
new = 'import sys as _sys\n _sys.stderr.write('
c2 = c.replace(old, new)
if c2 != c:
open(f, 'w').write(c2)
print('Patch 7: rocm.py circular import broken.')
else:
print('Patch 7: already applied or pattern not found.')
def main():
patch6_init_platform_fallback()
patch7_rocm_break_import_cycle()
try:
import vllm
from vllm.platforms import current_platform
print(f'vllm runtime import OK: {vllm.__version__}, platform={type(current_platform).__name__}')
except Exception:
print('ERROR: vllm runtime import failed after platform patches:')
traceback.print_exc()
raise
print('vllm platform patches done.')
if __name__ == '__main__':
main()