| 206 | |
| 207 | |
| 208 | def check_correctness(gold_outputs, outputs, rtol=1e-3, atol=1e-3): |
| 209 | if len(gold_outputs) != len(outputs): |
| 210 | print("Number of outputs {} is not equal to expected number {}".format( |
| 211 | len(outputs), len(gold_outputs))) |
| 212 | return False |
| 213 | |
| 214 | out_num = len(gold_outputs) |
| 215 | ret = True |
| 216 | for i in range(out_num): |
| 217 | if not np.allclose(gold_outputs[i], outputs[i], rtol, atol): |
| 218 | print("\nOutput {} is incorrect ...".format(i)) |
| 219 | print("Expected value: \n{}".format(gold_outputs[i])) |
| 220 | print("......") |
| 221 | print("Actual value: \n{}\n".format(outputs[i])) |
| 222 | ret = False |
| 223 | |
| 224 | return ret |
| 225 | |
| 226 | |
| 227 | def tune_input_shape(model, input_data): |