feature:add head size detect and patch some ops
All checks were successful
Docker Build and Push / docker (push) Successful in 1m5s
All checks were successful
Docker Build and Push / docker (push) Successful in 1m5s
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
151
patch_ops.py
Executable file
151
patch_ops.py
Executable file
@@ -0,0 +1,151 @@
|
||||
#!/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启动
|
||||
Reference in New Issue
Block a user