MCPcopy Create free account
hub / github.com/pytorch/pytorch / op_assert_ref

Function op_assert_ref

test/test_decomp.py:168–213  ·  view source on GitHub ↗
(test_case, op, test_dtype, i, orig, decomp, ref, args, kwargs)

Source from the content-addressed store, hash-verified

166
167
168def 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
216def op_assert_equal(test_case, op, test_dtype, orig, decomp, args, kwargs):

Callers 1

__torch_dispatch__Method · 0.85

Calls 5

maxMethod · 0.80
numelMethod · 0.45
absMethod · 0.45
getMethod · 0.45
assertEqualMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…