blob: 7d1a2ff050ad12555b9bb3f65b1bec9141bdcbba [file] [edit]
#!/usr/bin/env python3
# Copyright 2024 The ChromiumOS Authors
# Use of this source code is governed by a BSD-style license that can be
# found in the LICENSE file.
"""A script to generate accel_test.conf for stable_delegate_test_suite."""
import argparse
import logging
from pathlib import Path
import subprocess
import sys
from typing import Dict, List, Optional, TextIO, Tuple
_ACCEL_CONF_HEADER = """# Copyright 2024 The ChromiumOS Authors
# Use of this source code is governed by a BSD-style license that can be
# found in the LICENSE file.
# Auto-generated by stable_delegate_accel_test_conf.py
"""
def list_test(
stable_delegate_test_suite: Path,
gtest_filter: str,
) -> List[str]:
stable_delegate_test_suite_cmd = [
str(stable_delegate_test_suite.absolute()),
"--gtest_list_tests",
"--gtest_filter=%s" % gtest_filter,
]
logging.debug("Running: %s", " ".join(stable_delegate_test_suite_cmd))
result = subprocess.check_output(stable_delegate_test_suite_cmd, text=True)
prefix = ""
test_list: List[str] = []
for line in result.splitlines():
# ignoring comments for test parameters
item = line.split("#", 1)[0].strip()
if item.endswith("."):
prefix = item
else:
test_list.append(prefix + item)
return test_list
def run_tests(
stable_delegate_test_suite: Path,
stable_delegate_settings_file: Path,
allow_fp16_precision_for_fp32: bool,
gtest_filter: str,
) -> Tuple[Dict[str, List[Tuple[str, bool]]], List[str]]:
test_results: Dict[str, List[Tuple[str, bool]]] = {}
passed_count = 0
failed_count = 0
crashed_tests = []
for test in list_test(stable_delegate_test_suite, gtest_filter):
stable_delegate_test_suite_cmd = [
str(stable_delegate_test_suite.absolute()),
"--stable_delegate_settings_file=%s"
% (stable_delegate_settings_file.absolute(),),
"--allow_fp16_precision_for_fp32=%s"
% str(allow_fp16_precision_for_fp32).lower(),
"--gtest_filter=%s" % test,
]
result = subprocess.run(
stable_delegate_test_suite_cmd,
stdout=subprocess.DEVNULL,
stderr=subprocess.PIPE,
check=False,
)
if result.returncode not in [0, 1]:
logging.debug("%s: CRASH", test)
crashed_tests.append(test)
continue
testsuite = test.replace(".", "/").split("/")[0]
if testsuite not in test_results:
test_results[testsuite] = []
if result.returncode == 0:
logging.debug("%s: PASS", test)
test_results[testsuite].append((test, True))
passed_count += 1
else:
logging.debug("%s: FAIL", test)
test_results[testsuite].append((test, False))
failed_count += 1
logging.debug("Finished running stable_delegate_test_suite")
logging.debug(
"pass: %d, fail %d, crash, %d",
passed_count,
failed_count,
len(crashed_tests),
)
return test_results, crashed_tests
def generate_accel_test_conf(
stable_delegate_test_suite: Path,
stable_delegate_settings_file: Path,
allow_fp16_precision_for_fp32: bool,
output_file: TextIO,
crash_output_file: TextIO,
gtest_filter: str,
):
assert stable_delegate_test_suite.is_file()
assert stable_delegate_settings_file.is_file()
test_results, crashed_tests = run_tests(
stable_delegate_test_suite,
stable_delegate_settings_file,
allow_fp16_precision_for_fp32,
gtest_filter,
)
passed_subtests: List[str] = []
passed_testsuites: List[str] = []
tests: List[str] = []
for testsuite, subtest_results in test_results.items():
passed: List[str] = []
failed: List[str] = []
for subtest_result in subtest_results:
test, status = subtest_result
if status:
passed.append(test)
else:
failed.append("-" + test)
if len(failed) == 0:
passed_testsuites.append(f"{testsuite}/.*")
elif len(passed) > len(failed):
tests += failed
tests.append(f"{testsuite}/.*")
else:
passed_subtests += passed
output_file.write(_ACCEL_CONF_HEADER)
output_file.write("\n# Seperate passed tests in testsuite mostly failed\n")
for test in passed_subtests:
output_file.write("%s\n" % test)
output_file.write("\n# Passed testsuites\n")
for test in passed_testsuites:
output_file.write("%s\n" % test)
output_file.write("\n# Testsuites with failures\n")
for test in tests:
output_file.write("%s\n" % test)
output_file.write("\n# Wildcard for all the other failed tests\n")
output_file.write("-.*\n")
crash_output_file.write("-%s" % ":".join(crashed_tests))
logging.debug("Finished generating accel_test_conf")
def setup_argument_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
formatter_class=argparse.ArgumentDefaultsHelpFormatter
)
parser.add_argument(
"--stable_delegate_settings_file",
required=True,
type=Path,
help="The path to the delegate settings JSON file.",
)
parser.add_argument(
"--allow_fp16_precision_for_fp32",
dest="allow_fp16_precision_for_fp32",
action="store_true",
help="Allow FP16 precision for FP32 operators in DTS.",
)
parser.add_argument(
"--no-allow_fp16_precision_for_fp32",
dest="allow_fp16_precision_for_fp32",
action="store_false",
help="Don't allow FP16 precision for FP32 operators in DTS.",
)
parser.set_defaults(allow_fp16_precision_for_fp32=True)
parser.add_argument(
"--output",
default="accel_test.conf",
type=argparse.FileType("w"),
help="output path for generated acceleration test config",
)
parser.add_argument(
"--crash_output",
type=argparse.FileType("w"),
default="crash_filter.txt",
help="output path for generated gtest filter for crashed tests",
)
parser.add_argument(
"--gtest_filter",
type=str,
default="*",
help="gtest filter for quick verifying subset of tests",
)
parser.add_argument(
"--stable_delegate_test_suite",
default="/usr/local/bin/stable_delegate_test_suite",
type=Path,
help="path to the stable_delegate_test_suite binary",
)
parser.add_argument(
"--debug",
action="store_true",
help="enable debug logging",
)
return parser
def main(argv: Optional[List[str]] = None) -> Optional[int]:
parser = setup_argument_parser()
args = parser.parse_args(argv)
log_level = logging.DEBUG if args.debug else logging.INFO
log_format = "%(asctime)s - %(levelname)s - %(funcName)s: %(message)s"
logging.basicConfig(level=log_level, format=log_format)
generate_accel_test_conf(
args.stable_delegate_test_suite,
args.stable_delegate_settings_file,
args.allow_fp16_precision_for_fp32,
args.output,
args.crash_output,
args.gtest_filter,
)
if __name__ == "__main__":
sys.exit(main(sys.argv[1:]))