()
| 48 | |
| 49 | |
| 50 | def main() -> None: |
| 51 | parser = argparse.ArgumentParser() |
| 52 | parser.add_argument( |
| 53 | "-m", |
| 54 | "--model_name", |
| 55 | required=True, |
| 56 | help=f"provide a model name. Valid ones: {list(MODEL_NAME_TO_MODEL.keys())}", |
| 57 | ) |
| 58 | |
| 59 | parser.add_argument( |
| 60 | "-o", |
| 61 | "--output_path", |
| 62 | required=False, |
| 63 | help=f"Provide an output path to save the generated etrecord. Defaults to {DEFAULT_OUTPUT_PATH}.", |
| 64 | ) |
| 65 | |
| 66 | args = parser.parse_args() |
| 67 | |
| 68 | if args.model_name not in MODEL_NAME_TO_MODEL: |
| 69 | raise RuntimeError( |
| 70 | f"Model {args.model_name} is not a valid name. " |
| 71 | f"Available models are {list(MODEL_NAME_TO_MODEL.keys())}." |
| 72 | ) |
| 73 | |
| 74 | model, example_inputs, _, _ = EagerModelFactory.create_model( |
| 75 | *MODEL_NAME_TO_MODEL[args.model_name] |
| 76 | ) |
| 77 | |
| 78 | gen_etrecord(model, example_inputs, args.output_path) |
| 79 | |
| 80 | |
| 81 | if __name__ == "__main__": |
no test coverage detected