from __future__ import annotations import argparse from pathlib import Path REPO = Path(__file__).resolve().parents[2] SUITE = REPO / "validation/networks/pimcomp_models" MODELS = { "vgg8": SUITE / "vgg8/vgg8-mnist-reconstructed.onnx", "resnet18": SUITE / "resnet18/resnet18-v1-7.onnx", "resnet34": SUITE / "resnet34/resnet34-v1-7.onnx", "googlenet": SUITE / "googlenet/googlenet-12-pimsim-nn.onnx", "yolo11n": SUITE / "yolo11n/yolo11n-pimsim-nn.onnx", } FUNCTIONAL_MODELS = { **MODELS, "yolo11n": REPO / "validation/networks/yolo11n/depth_51/yolo11n_depth_51.onnx", } MODEL_NAMES = tuple(MODELS) DEFAULT_MODELS = MODEL_NAMES ABLATION_DEFAULT_MODELS = ("vgg8", "resnet18", "resnet34", "googlenet") def add_models_argument( parser: argparse.ArgumentParser, default: tuple[str, ...] = DEFAULT_MODELS, ) -> None: help_text = "Models to run (default: " + ", ".join(default) + ")." if "yolo11n" not in default: help_text += " Select yolo11n explicitly when needed." parser.add_argument( "--models", nargs="+", choices=MODEL_NAMES, default=list(default), metavar="MODEL", help=help_text, )