All checks were successful
Docker Build and Push / docker (push) Successful in 1m5s
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
152 lines
4.8 KiB
Python
Executable File
152 lines
4.8 KiB
Python
Executable File
#!/usr/bin/env python3
|
||
"""
|
||
Patch ixformer functions - 自动扫描并复制自定义的ops文件并修改__init__.py
|
||
|
||
支持批量处理patched_ops目录中的所有.py文件,无需每次修改代码。
|
||
只需将新的ops文件放入patched_ops目录即可。
|
||
"""
|
||
import os
|
||
import shutil
|
||
from typing import List
|
||
|
||
# 目标路径
|
||
TARGET_DIR = "/usr/local/corex/lib64/python3/dist-packages/ixformer/functions/"
|
||
TARGET_INIT_FILE = os.path.join(TARGET_DIR, "__init__.py")
|
||
SOURCE_DIR = "/opt/patched_ops/"
|
||
|
||
|
||
def get_ops_files(source_dir: str) -> List[str]:
|
||
"""
|
||
扫描源目录中的所有.py文件(排除__init__.py)
|
||
|
||
Args:
|
||
source_dir: 源目录路径
|
||
|
||
Returns:
|
||
.py文件名列表
|
||
"""
|
||
if not os.path.exists(source_dir):
|
||
print(f"[patch_ops] Source directory not found: {source_dir}")
|
||
return []
|
||
|
||
ops_files = []
|
||
for filename in os.listdir(source_dir):
|
||
if filename.endswith('.py') and filename != '__init__.py':
|
||
ops_files.append(filename)
|
||
|
||
return ops_files
|
||
|
||
|
||
def generate_import_statement(filename: str) -> str:
|
||
"""
|
||
根据文件名生成import语句
|
||
|
||
Args:
|
||
filename: .py文件名(不含扩展名)
|
||
|
||
Returns:
|
||
import语句
|
||
"""
|
||
module_name = filename.replace('.py', '')
|
||
return f"from .{module_name} import *"
|
||
|
||
|
||
def patch_ixformer_ops() -> bool:
|
||
"""
|
||
Patch ixformer functions - 批量复制ops文件并修改__init__.py
|
||
|
||
Returns:
|
||
是否成功patch
|
||
"""
|
||
try:
|
||
print(f"[patch_ops] Starting ixformer ops patch...")
|
||
|
||
# 检查目标目录是否存在
|
||
if not os.path.exists(TARGET_DIR):
|
||
print(f"[patch_ops] Target directory not found: {TARGET_DIR}")
|
||
return False
|
||
|
||
# 扫描源目录中的所有ops文件
|
||
ops_files = get_ops_files(SOURCE_DIR)
|
||
|
||
if not ops_files:
|
||
print(f"[patch_ops] No ops files found in {SOURCE_DIR}")
|
||
return False
|
||
|
||
print(f"[patch_ops] Found {len(ops_files)} ops file(s): {', '.join(ops_files)}")
|
||
|
||
# 备份原__init__.py文件
|
||
if os.path.exists(TARGET_INIT_FILE):
|
||
backup_file = TARGET_INIT_FILE + ".backup"
|
||
if not os.path.exists(backup_file):
|
||
shutil.copy2(TARGET_INIT_FILE, backup_file)
|
||
print(f"[patch_ops] Backed up {TARGET_INIT_FILE} to {backup_file}")
|
||
else:
|
||
print(f"[patch_ops] Backup already exists: {backup_file}")
|
||
|
||
# 读取现有的__init__.py内容
|
||
existing_imports = set()
|
||
if os.path.exists(TARGET_INIT_FILE):
|
||
with open(TARGET_INIT_FILE, 'r') as f:
|
||
init_content = f.read()
|
||
# 提取现有的from .xxx import *语句
|
||
for line in init_content.split('\n'):
|
||
line = line.strip()
|
||
if line.startswith('from .') and 'import *' in line:
|
||
existing_imports.add(line)
|
||
|
||
# 批量复制文件并收集需要添加的import语句
|
||
new_imports = []
|
||
for filename in ops_files:
|
||
source_file = os.path.join(SOURCE_DIR, filename)
|
||
target_file = os.path.join(TARGET_DIR, filename)
|
||
|
||
# 复制文件
|
||
shutil.copy2(source_file, target_file)
|
||
print(f"[patch_ops] Copied {filename} to {TARGET_DIR}")
|
||
|
||
# 生成import语句
|
||
import_stmt = generate_import_statement(filename)
|
||
if import_stmt not in existing_imports:
|
||
new_imports.append(import_stmt)
|
||
|
||
# 修改__init__.py文件,添加新的import语句
|
||
if new_imports:
|
||
with open(TARGET_INIT_FILE, 'a') as f:
|
||
f.write(f"\n# Auto-patched ixformer ops\n")
|
||
for import_stmt in new_imports:
|
||
f.write(f"{import_stmt}\n")
|
||
print(f"[patch_ops] Added {len(new_imports)} import statement(s) to {TARGET_INIT_FILE}")
|
||
else:
|
||
print(f"[patch_ops] All imports already exist in {TARGET_INIT_FILE}, skipping modification")
|
||
|
||
print(f"[patch_ops] Successfully patched ixformer ops")
|
||
return True
|
||
|
||
except Exception as e:
|
||
print(f"[patch_ops] Error patching ixformer ops: {e}")
|
||
return False
|
||
|
||
|
||
def main():
|
||
"""主函数"""
|
||
try:
|
||
success = patch_ixformer_ops()
|
||
if success:
|
||
print(f"[patch_ops] Patch completed successfully")
|
||
else:
|
||
print(f"[patch_ops] Patch failed, but this will not prevent vLLM from starting")
|
||
except Exception as e:
|
||
print(f"[patch_ops] Unexpected error: {e}")
|
||
print(f"[patch_ops] This will not prevent vLLM from starting")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
import sys
|
||
try:
|
||
main()
|
||
except Exception as e:
|
||
print(f"[patch_ops] Unexpected error: {e}")
|
||
print(f"[patch_ops] This will not prevent vLLM from starting")
|
||
sys.exit(0) # 总是成功退出,不阻止vLLM启动
|