()
| 975 | # --------------------------------------------------------------------------- |
| 976 | |
| 977 | def main() -> None: |
| 978 | global WORKSPACE_DIR, ORCHESTRATION_STATE, WARMUP_RUNS, TIMED_RUNS |
| 979 | |
| 980 | parser = argparse.ArgumentParser( |
| 981 | description="AutoKernel End-to-End Verifier", |
| 982 | formatter_class=argparse.RawDescriptionHelpFormatter, |
| 983 | epilog=__doc__, |
| 984 | ) |
| 985 | |
| 986 | # Model loading |
| 987 | model_group = parser.add_mutually_exclusive_group(required=True) |
| 988 | model_group.add_argument( |
| 989 | "--model", type=str, |
| 990 | help="Path to a Python file containing the model class" |
| 991 | ) |
| 992 | model_group.add_argument( |
| 993 | "--module", type=str, |
| 994 | help="Python module name (e.g. 'transformers')" |
| 995 | ) |
| 996 | |
| 997 | parser.add_argument( |
| 998 | "--class-name", type=str, required=True, |
| 999 | help="Name of the model class to instantiate" |
| 1000 | ) |
| 1001 | parser.add_argument( |
| 1002 | "--pretrained", type=str, default=None, |
| 1003 | help="Pretrained model name/path (for HuggingFace models)" |
| 1004 | ) |
| 1005 | parser.add_argument( |
| 1006 | "--input-shape", type=str, default="1,2048", |
| 1007 | help="Comma-separated input shape, e.g. '1,2048' (default: 1,2048)" |
| 1008 | ) |
| 1009 | parser.add_argument( |
| 1010 | "--dtype", type=str, default="float16", |
| 1011 | help="Data type: float16, bfloat16, float32 (default: float16)" |
| 1012 | ) |
| 1013 | |
| 1014 | # Benchmark tuning |
| 1015 | parser.add_argument( |
| 1016 | "--warmup", type=int, default=WARMUP_RUNS, |
| 1017 | help=f"Number of warmup iterations (default: {WARMUP_RUNS})" |
| 1018 | ) |
| 1019 | parser.add_argument( |
| 1020 | "--timed", type=int, default=TIMED_RUNS, |
| 1021 | help=f"Number of timed iterations (default: {TIMED_RUNS})" |
| 1022 | ) |
| 1023 | |
| 1024 | # Tolerance overrides |
| 1025 | parser.add_argument("--atol", type=float, default=None, help="Override absolute tolerance") |
| 1026 | parser.add_argument("--rtol", type=float, default=None, help="Override relative tolerance") |
| 1027 | |
| 1028 | # Modes |
| 1029 | parser.add_argument( |
| 1030 | "--diagnose", action="store_true", |
| 1031 | help="On failure, test each kernel replacement individually to find the culprit" |
| 1032 | ) |
| 1033 | parser.add_argument( |
| 1034 | "--json", type=str, default=None, |
no test coverage detected