(results, expected, check_shape=True)
| 70 | """ |
| 71 | |
| 72 | def check_results(results, expected, check_shape=True): |
| 73 | if not isinstance(results, (tuple, list)): |
| 74 | results = (results,) |
| 75 | for r, e in zip(results, expected): |
| 76 | if not isinstance(r, (tensor, VarNode)): |
| 77 | r = tensor(r) |
| 78 | if check_shape: |
| 79 | r_shape = r.numpy().shape |
| 80 | e_shape = e.shape if isinstance(e, np.ndarray) else () |
| 81 | assert r_shape == e_shape |
| 82 | compare_fn(r, e) |
| 83 | |
| 84 | def get_param(cases, idx): |
| 85 | case = cases[idx] |