Files
mineru-rocm/docker/scripts/patch_vllm_platform.py
T

75 lines
2.7 KiB
Python
Raw Normal View History

2026-06-04 16:16:30 +08:00
#!/usr/bin/env python3
"""vllm 平台检测补丁
问题 1:amdsmi 不可用时 platform 回退到 torch.version.hip
2026-06-12 10:56:36 +08:00
问题 2:rocm.py 中 logger.warning_once() 导致循环导入(默认不再改写;仅保留为显式开关)
2026-06-04 16:16:30 +08:00
"""
import os
2026-06-12 10:56:36 +08:00
import traceback
2026-06-04 16:16:30 +08:00
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():
2026-06-12 10:56:36 +08:00
"""补丁 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
2026-06-04 16:16:30 +08:00
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()
2026-06-12 10:56:36 +08:00
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
2026-06-04 16:16:30 +08:00
print('vllm platform patches done.')
if __name__ == '__main__':
main()