import os
import re
import sys
import argparse
import logging
import subprocess
from pathlib import Path
KEYS = [
"OP_CATEGORY",
"OP_NAME",
"HOSTNAME",
"MODE",
"DIR",
"OPTYPE",
"ACLNNTYPE",
"DEPENDENCIES",
"COMPUTE_UNIT",
"TILING_DIR",
"DISABLE_IN_OPP",
]
OP_CATEGORY_SET = {""}
logger = logging.getLogger()
logging.basicConfig(level=logging.INFO, stream=sys.stdout)
def args_parse():
parser = argparse.ArgumentParser()
parser.add_argument(
"--ops",
nargs="?",
required=True,
help="Operators that need to find dependency.",
)
parser.add_argument("-p", "--path", nargs="?", required=True, help="Build path.")
return parser.parse_args()
def set_dict_value(dict_value, key, value):
if key not in dict_value:
dict_value[key] = []
dict_value[key].append(value)
def check_pytorch_extension_op(cmake_file: Path) -> bool:
if not cmake_file.exists():
return False
try:
add_sources_pattern = re.compile(r"add_sources\s*\([^)]*\)")
comment_pattern = re.compile(r"#[^\n]*")
content = cmake_file.read_text(encoding="utf-8")
content_no_comment = comment_pattern.sub("", content)
return add_sources_pattern.search(content_no_comment) is not None
except (IOError, UnicodeDecodeError) as e:
logging.warning("Failed to read CMakeLists.txt: %s, error: %s", cmake_file, e)
return False
class OpDependenciesParser:
def __init__(self, build_path):
self.all_ops_dependency = {}
self.all_ops_reverse_dependency = {}
self.all_ops = ["add_example", "add_example_aicpu"]
self.all_category_ops = {}
self.parse_dependency(build_path)
self.parse_pytorch_extension_ops(build_path)
self.framework_only_ops = self.parse_common_framework_ops(build_path)
pass
def find_all_dependency(self, op, result_dependencies, all_dependencies, src_op):
if op not in self.all_ops:
if op in self.framework_only_ops:
logging.warning(
"%s is not in the dependency graph (framework-only plugin without an op directory); "
"treat as no sub-dependencies.",
op,
)
if op not in result_dependencies:
result_dependencies.append(op)
return
logging.error("%s is not exists, please check.", op)
raise RuntimeError(f"{op} is not exists, please check.")
if op in result_dependencies:
return
result_dependencies.append(op)
for sub_op in all_dependencies.get(op, []):
self.find_all_dependency(
sub_op, result_dependencies, all_dependencies, src_op
)
def parse_line(self, line):
last_key = None
op_type = None
op_category = None
common_name = None
for value in line.strip().split(";"):
if value in KEYS:
last_key = value
continue
if last_key == "OP_CATEGORY":
op_category = value
common_name = op_category + ".common"
if last_key == "OP_NAME":
op_type = value
if op_type == "common":
op_type = common_name
set_dict_value(self.all_category_ops, op_category, op_type)
self.all_ops.append(op_type)
if last_key == "DEPENDENCIES":
set_dict_value(self.all_ops_dependency, op_type, value)
set_dict_value(self.all_ops_reverse_dependency, value, op_type)
def parse_dependency(self, build_path):
build_path = os.path.abspath(build_path)
file_path = os.path.join(build_path, "tmp", "ops_config.txt")
if not os.path.exists(file_path):
logging.error("%s config file is not exists.", file_path)
raise RuntimeError(f"{file_path} config file is not exists.")
with open(file_path, "r", encoding="utf-8") as file:
for line in file:
self.parse_line(line)
for op_category, ops in self.all_category_ops.items():
common_name = op_category + ".common"
self.all_ops.append(common_name)
if common_name in ops:
self.all_ops_reverse_dependency[common_name] = []
self.all_ops_reverse_dependency[common_name].extend(ops)
for op in ops:
set_dict_value(self.all_ops_dependency, op, common_name)
def parse_pytorch_extension_ops(self, build_path):
experimental_path = Path(build_path).resolve().parent / "experimental"
if not experimental_path.exists():
logging.warning("experimental directory not found: %s", experimental_path)
return
for op_class_dir in experimental_path.iterdir():
if not op_class_dir.is_dir():
continue
for op_dir in op_class_dir.iterdir():
if not op_dir.is_dir():
continue
cmake_file = op_dir / "CMakeLists.txt"
if check_pytorch_extension_op(cmake_file):
set_dict_value(
self.all_category_ops, op_class_dir.name, op_dir.name
)
self.all_ops.append(op_dir.name)
def parse_common_framework_ops(self, build_path):
plugin_dir = Path(build_path).resolve().parent / "common" / "src" / "framework"
framework_only_ops = set()
if not plugin_dir.exists():
return framework_only_ops
for plugin_file in plugin_dir.glob("*_onnx_plugin.cpp"):
op_name = plugin_file.name[: -len("_onnx_plugin.cpp")]
if op_name not in self.all_ops:
framework_only_ops.add(op_name)
return framework_only_ops
def get_dependencies_by_ops(self, ops):
result_ops = []
reverse_ops = []
for op in ops:
if op not in reverse_ops:
self.find_all_dependency(
op, reverse_ops, self.all_ops_reverse_dependency, op
)
if op not in result_ops:
self.find_all_dependency(op, result_ops, self.all_ops_dependency, op)
return (result_ops, reverse_ops)
def get_category_list(self):
return self.all_category_ops.keys()
def find_category(ops_list, all_category_ops):
result = []
for value in ops_list.split(";"):
keys = [key for key, val in all_category_ops.items() if value in val]
if keys:
result.append(keys[0])
return result
def main():
args = args_parse()
parser = OpDependenciesParser(args.path)
(op_dependencies, reverse_op_dependencies) = parser.get_dependencies_by_ops(
args.ops.split(";")
)
op_dependencies = ";".join(op_dependencies)
reverse_op_dependencies = ";".join(reverse_op_dependencies)
category_set = set(find_category(op_dependencies, parser.all_category_ops))
enable_asc_build = "FALSE"
if category_set == OP_CATEGORY_SET:
enable_asc_build = "TRUE"
logging.info(
"op_dependencies:%s, reverse_op_dependencies:%s",
op_dependencies,
reverse_op_dependencies,
)
subprocess.run(
[
"cmake",
"-DASCEND_COMPILE_OPS=" + op_dependencies,
"-DENABLE_ASC_BUILD=" + enable_asc_build,
"..",
],
cwd=args.path,
)
if __name__ == "__main__":
main()