()
| 65 | |
| 66 | |
| 67 | def main(): |
| 68 | args = parse_args() |
| 69 | |
| 70 | if args.model_name not in MODEL_NAME_TO_MODEL: |
| 71 | raise RuntimeError( |
| 72 | f"Model {args.model_name} is not a valid name. " |
| 73 | f"Available models are {list(MODEL_NAME_TO_MODEL.keys())}." |
| 74 | ) |
| 75 | |
| 76 | ( |
| 77 | model, |
| 78 | example_args, |
| 79 | example_kwargs, |
| 80 | dynamic_shapes, |
| 81 | ) = EagerModelFactory.create_model(*MODEL_NAME_TO_MODEL[args.model_name]) |
| 82 | model = model.eval() |
| 83 | exported_programs = torch.export.export( |
| 84 | model, |
| 85 | args=example_args, |
| 86 | kwargs=example_kwargs, |
| 87 | dynamic_shapes=dynamic_shapes, |
| 88 | ) |
| 89 | |
| 90 | partitioner = CudaPartitioner( |
| 91 | [CudaBackend.generate_method_name_compile_spec(args.model_name)] |
| 92 | ) |
| 93 | |
| 94 | et_prog = to_edge_transform_and_lower( |
| 95 | exported_programs, |
| 96 | partitioner=[partitioner], |
| 97 | compile_config=_EDGE_COMPILE_CONFIG, |
| 98 | generate_etrecord=args.generate_etrecord, |
| 99 | ) |
| 100 | exec_program = et_prog.to_executorch() |
| 101 | save_pte_program(exec_program, args.model_name, args.output_dir) |
| 102 | if args.generate_etrecord: |
| 103 | exec_program.get_etrecord().save(f"{args.model_name}_etrecord.bin") |
| 104 | |
| 105 | |
| 106 | if __name__ == "__main__": |
no test coverage detected