(args)
| 28 | |
| 29 | |
| 30 | def main(args): |
| 31 | # User defined keyword arguments |
| 32 | kwargs = {"order": "NCHW"} |
| 33 | kwargs.update(dict(args.kwargs)) |
| 34 | |
| 35 | model = ModelHelper(name=args.benchmark_name) |
| 36 | |
| 37 | op_type = args.operator # assumes a brew type op name |
| 38 | input_name = args.input_name |
| 39 | output_name = args.output_name |
| 40 | |
| 41 | iters = int(args.iters) |
| 42 | for i in range(iters): |
| 43 | input_blob_name = input_name + (str(i) if i > 0 and args.chain else '') |
| 44 | output_blob_name = output_name + str(i + 1) |
| 45 | add_op = getattr(brew, op_type) |
| 46 | add_op(model, input_blob_name, output_blob_name, **kwargs) |
| 47 | if args.chain: |
| 48 | input_name, output_name = output_name, input_name |
| 49 | |
| 50 | workspace.RunNetOnce(model.param_init_net) |
| 51 | extra_init_net_ops = [] |
| 52 | |
| 53 | def make_blob_on_context(blob_name, blob_data, context): |
| 54 | if context.upper() != "CPU": |
| 55 | blob_name_modified = "{}_CPU".format(blob_name) |
| 56 | else: # CPU case is simple |
| 57 | blob_name_modified = blob_name |
| 58 | |
| 59 | fill_op = core.CreateOperator( |
| 60 | "GivenTensorFill", [], [blob_name_modified], |
| 61 | arg=[ |
| 62 | utils.MakeArgument("shape", blob_data.shape), |
| 63 | utils.MakeArgument("values", blob_data) |
| 64 | ] |
| 65 | ) |
| 66 | extra_init_net_ops.append(fill_op) |
| 67 | |
| 68 | # We need to create CPU blobs and add some copy operations in |
| 69 | # the init_net |
| 70 | if context.upper() == "OPENGL": |
| 71 | copy_op = core.CreateOperator("CopyToOpenGL", [blob_name_modified], |
| 72 | [blob_name]) |
| 73 | extra_init_net_ops.append(copy_op) |
| 74 | |
| 75 | for unparsed_blob in args.blob: |
| 76 | name, unparsed_dims = unparsed_blob.split('=') |
| 77 | dims = [int(d) for d in unparsed_dims.split(',')] |
| 78 | np_input = np.random.rand(*dims).astype(np.float32) |
| 79 | make_blob_on_context(name, np_input, args.context) |
| 80 | |
| 81 | init_net, predict_net = mobile_exporter.Export( |
| 82 | workspace, model.net, model.params |
| 83 | ) |
| 84 | init_net.op.extend(extra_init_net_ops) |
| 85 | |
| 86 | # Handle manual rewrite |
| 87 | if args.context.upper() == "OPENGL": |
no test coverage detected
searching dependent graphs…