MCPcopy Create free account
hub / github.com/PaddlePaddle/Research / convert

Function convert

CV/landmark/inference/convert_binary_model.py:41–80  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

39
40
41def 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
83def main():

Callers 3

mainFunction · 0.70
__type_castMethod · 0.50
__type_castMethod · 0.50

Calls 2

netMethod · 0.45
runMethod · 0.45

Tested by

no test coverage detected