(test_case, op, test_dtype, i, orig, decomp, ref, args, kwargs)
| 166 | |
| 167 | |
| 168 | def op_assert_ref(test_case, op, test_dtype, i, orig, decomp, ref, args, kwargs): |
| 169 | assert orig.dtype == decomp.dtype, f"{i} Operation: {op}" |
| 170 | if orig.numel() == 0 or decomp.numel() == 0: |
| 171 | assert orig.numel() == decomp.numel() |
| 172 | return |
| 173 | assert orig.shape == decomp.shape, f"{i} Operation: {op}" |
| 174 | tol_table = { |
| 175 | (torch.bfloat16, torch.ops.aten.native_layer_norm.default): 1e-5, |
| 176 | (torch.float16, torch.ops.aten.native_layer_norm.default): 1e-5, |
| 177 | (torch.float16, torch.ops.aten.native_layer_norm_backward.default): 1e-3, |
| 178 | (torch.bfloat16, torch.ops.aten.native_layer_norm_backward.default): 2e-2, |
| 179 | (torch.bfloat16, torch.ops.aten.native_batch_norm.default): 1e-5, |
| 180 | (torch.float16, torch.ops.aten.native_batch_norm.default): 1e-5, |
| 181 | (torch.bfloat16, torch.ops.aten._native_batch_norm_legit.default): 1e-5, |
| 182 | (torch.bfloat16, torch.ops.aten._native_batch_norm_legit.no_stats): 1e-5, |
| 183 | (torch.float16, torch.ops.aten._native_batch_norm_legit.default): 1e-5, |
| 184 | (torch.float16, torch.ops.aten._native_batch_norm_legit.no_stats): 1e-5, |
| 185 | (torch.bfloat16, torch.ops.aten.linalg_vector_norm.default): 1e-4, |
| 186 | (torch.float16, torch.ops.aten.linalg_vector_norm.default): 1e-4, |
| 187 | (torch.bfloat16, torch.ops.aten.var_mean.correction): 5e-7, |
| 188 | (torch.float16, torch.ops.aten.var_mean.correction): 5e-7, |
| 189 | (torch.bfloat16, torch.ops.aten.var_mean.dim): 5e-7, |
| 190 | (torch.float16, torch.ops.aten.var_mean.dim): 5e-7, |
| 191 | (torch.float16, torch.ops.aten.nll_loss_forward.default): 1e-2, |
| 192 | (torch.bfloat16, torch.ops.aten.nll_loss_forward.default): 1e-1, |
| 193 | (torch.float16, torch.ops.aten.nll_loss2d_forward.default): 1e-2, |
| 194 | (torch.bfloat16, torch.ops.aten.nll_loss2d_forward.default): 2e-1, |
| 195 | # see https://github.com/pytorch/pytorch/pull/96264 |
| 196 | (torch.float16, torch.ops.aten.mv.default): 1e-5, |
| 197 | } |
| 198 | if ref.is_floating_point(): |
| 199 | orig_diff = (orig - ref).abs().max() |
| 200 | decomp_diff = (decomp - ref).abs().max() |
| 201 | atol = tol_table.get((test_dtype, op), 1e-7) |
| 202 | if decomp_diff > orig_diff + atol: |
| 203 | raise RuntimeError( |
| 204 | f"Difference from float64 is larger with decomposition {op.__name__}" |
| 205 | f" than original on output {i}. Original max diff: {orig_diff}, Decomp max diff: {decomp_diff}\n" |
| 206 | f"atol = {atol}\n" |
| 207 | f"args = {args}\n" |
| 208 | f"kwargs = {kwargs}" |
| 209 | ) |
| 210 | else: |
| 211 | test_case.assertEqual( |
| 212 | orig, decomp, msg=f"{op.__name__}\nargs = {args}\nkwargs = {kwargs}" |
| 213 | ) |
| 214 | |
| 215 | |
| 216 | def op_assert_equal(test_case, op, test_dtype, orig, decomp, args, kwargs): |
no test coverage detected
searching dependent graphs…