Load model according to CLI args. Returns (model, description_string).
(args: argparse.Namespace)
| 291 | |
| 292 | |
| 293 | def load_model(args: argparse.Namespace) -> Tuple[nn.Module, str]: |
| 294 | """Load model according to CLI args. Returns (model, description_string).""" |
| 295 | if args.model: |
| 296 | model = _load_model_from_file(args.model, args.class_name) |
| 297 | desc = f"{args.class_name} from {args.model}" |
| 298 | elif args.module: |
| 299 | model = _load_model_from_module(args.module, args.class_name, args.pretrained) |
| 300 | desc = f"{args.class_name} from {args.module}" |
| 301 | if args.pretrained: |
| 302 | desc += f" (pretrained: {args.pretrained})" |
| 303 | else: |
| 304 | raise ValueError("Must specify either --model <file> or --module <package>") |
| 305 | |
| 306 | return model, desc |
| 307 | |
| 308 | |
| 309 | # --------------------------------------------------------------------------- |
no test coverage detected