(args)
| 39 | |
| 40 | |
| 41 | def convert(args): |
| 42 | # parameters from arguments |
| 43 | model_name = args.model |
| 44 | pretrained_model = args.pretrained_model |
| 45 | if not os.path.exists(pretrained_model): |
| 46 | print("pretrained_model doesn't exist!") |
| 47 | sys.exit(-1) |
| 48 | image_shape = [int(m) for m in args.image_shape.split(",")] |
| 49 | |
| 50 | assert model_name in model_list, "{} is not in lists: {}".format(args.model, |
| 51 | model_list) |
| 52 | |
| 53 | image = fluid.layers.data(name='image', shape=image_shape, dtype='float32') |
| 54 | |
| 55 | # model definition |
| 56 | model = models.__dict__[model_name]() |
| 57 | if args.task_mode == 'retrieval': |
| 58 | out = model.net(input=image, embedding_size=args.embedding_size) |
| 59 | else: |
| 60 | out = model.net(input=image) |
| 61 | place = fluid.CPUPlace() |
| 62 | exe = fluid.Executor(place) |
| 63 | exe.run(fluid.default_startup_program()) |
| 64 | |
| 65 | def if_exist(var): |
| 66 | return os.path.exists(os.path.join(pretrained_model, var.name)) |
| 67 | fluid.io.load_vars(exe, pretrained_model, predicate=if_exist) |
| 68 | |
| 69 | fluid.io.save_inference_model( |
| 70 | dirname = args.binary_model, |
| 71 | feeded_var_names = ['image'], |
| 72 | target_vars = [out['embedding']] if args.task_mode == 'retrieval' else [out], |
| 73 | executor = exe, |
| 74 | main_program = None, |
| 75 | model_filename = 'model', |
| 76 | params_filename = 'params') |
| 77 | |
| 78 | print('input_name: {}'.format('image')) |
| 79 | print('output_name: {}'.format(out['embedding'].name)) if args.task_mode == 'retrieval' else ('output_name: {}'.format(out.name)) |
| 80 | print("convert done.") |
| 81 | |
| 82 | |
| 83 | def main(): |
no test coverage detected