(platform, model_file,
input_file, mace_out_file,
input_names, input_shapes, input_data_formats,
output_names, output_shapes, output_data_formats,
validation_threshold, input_data_types, log_file)
| 248 | |
| 249 | |
| 250 | def validate_pytorch_model(platform, model_file, |
| 251 | input_file, mace_out_file, |
| 252 | input_names, input_shapes, input_data_formats, |
| 253 | output_names, output_shapes, output_data_formats, |
| 254 | validation_threshold, input_data_types, log_file): |
| 255 | import torch |
| 256 | loaded_model = torch.jit.load(model_file) |
| 257 | pytorch_inputs = [] |
| 258 | for i in range(len(input_names)): |
| 259 | input_value = load_data( |
| 260 | util.formatted_file_name(input_file, input_names[i]), |
| 261 | input_data_types[i]) |
| 262 | input_value = input_value.reshape(input_shapes[i]) |
| 263 | if input_data_formats[i] == DataFormat.NHWC and \ |
| 264 | len(input_shapes[i]) == 4: |
| 265 | input_value = input_value.transpose((0, 3, 1, 2)) |
| 266 | input_value = torch.from_numpy(input_value) |
| 267 | pytorch_inputs.append(input_value) |
| 268 | with torch.no_grad(): |
| 269 | pytorch_outputs = loaded_model(*pytorch_inputs) |
| 270 | |
| 271 | if isinstance(pytorch_outputs, torch.Tensor): |
| 272 | pytorch_outputs = [pytorch_outputs] |
| 273 | else: |
| 274 | if not isinstance(pytorch_outputs, (list, tuple)): |
| 275 | print('return type {} unsupported'.format(type(pytorch_outputs))) |
| 276 | sys.exit(1) |
| 277 | for i in range(len(output_names)): |
| 278 | value = pytorch_outputs[i].numpy() |
| 279 | output_file_name = util.formatted_file_name( |
| 280 | mace_out_file, output_names[i]) |
| 281 | mace_out_value = load_data(output_file_name) |
| 282 | mace_out_value, real_output_shape, real_output_data_format = \ |
| 283 | get_real_out_value_shape_df(platform, |
| 284 | mace_out_value, |
| 285 | output_shapes[i], |
| 286 | output_data_formats[i]) |
| 287 | compare_output(output_names[i], mace_out_value, |
| 288 | value, validation_threshold, log_file, |
| 289 | real_output_shape, real_output_data_format) |
| 290 | |
| 291 | |
| 292 | def validate_caffe_model(platform, model_file, input_file, |
no test coverage detected