()
| 61 | |
| 62 | |
| 63 | def main(): |
| 64 | parser = argparse.ArgumentParser(description="MegEngine Classification Dump .mge") |
| 65 | parser.add_argument( |
| 66 | "-a", |
| 67 | "--arch", |
| 68 | default="resnet18", |
| 69 | help="model architecture (default: resnet18)", |
| 70 | ) |
| 71 | parser.add_argument( |
| 72 | "-s", |
| 73 | "--shape", |
| 74 | type=int, |
| 75 | nargs='+', |
| 76 | default="1 3 224 224", |
| 77 | help="input shape (default: 1 3 224 224)" |
| 78 | ) |
| 79 | parser.add_argument( |
| 80 | "-o", |
| 81 | "--output", |
| 82 | type=str, |
| 83 | default="model.mge", |
| 84 | help="output filename" |
| 85 | ) |
| 86 | |
| 87 | args = parser.parse_args() |
| 88 | if 'resnet' in args.arch: |
| 89 | model = getattr(resnet_model, args.arch)(pretrained=True) |
| 90 | elif 'shufflenet' in args.arch: |
| 91 | model = getattr(snet_model, args.arch)(pretrained=True) |
| 92 | else: |
| 93 | print('unavailable arch {}'.format(args.arch)) |
| 94 | sys.exit() |
| 95 | print(model) |
| 96 | dump_static_graph(model, args.output, tuple(args.shape)) |
| 97 | |
| 98 | |
| 99 | if __name__ == "__main__": |
no test coverage detected