GaussianNLLLoss no_batch_dim docs and testing (#69783) Summary: Pull Request resolved: https://github.com/pytorch/pytorch/pull/69783 Test Plan: Imported from OSS Reviewed By: jbschlosser Differential Revision: D33200486 Pulled By: george-qi fbshipit-source-id: a2bc2b366772682825f879dae4ac29c1f4d6a5f1
diff --git a/torch/nn/modules/loss.py b/torch/nn/modules/loss.py index e7ef57c..46c3e53 100644 --- a/torch/nn/modules/loss.py +++ b/torch/nn/modules/loss.py
@@ -323,11 +323,11 @@ losses. Default: ``'mean'``. Shape: - - Input: :math:`(N, *)` where :math:`*` means any number of additional + - Input: :math:`(N, *)` or :math:`(*)` where :math:`*` means any number of additional dimensions - - Target: :math:`(N, *)`, same shape as the input, or same shape as the input + - Target: :math:`(N, *)` or :math:`(*)`, same shape as the input, or same shape as the input but with one dimension equal to 1 (to allow for broadcasting) - - Var: :math:`(N, *)`, same shape as the input, or same shape as the input but + - Var: :math:`(N, *)` or :math:`(*)`, same shape as the input, or same shape as the input but with one dimension equal to 1, or same shape as the input but with one fewer dimension (to allow for broadcasting) - Output: scalar if :attr:`reduction` is ``'mean'`` (default) or
diff --git a/torch/testing/_internal/common_modules.py b/torch/testing/_internal/common_modules.py index 0f73f56..d3168fb 100644 --- a/torch/testing/_internal/common_modules.py +++ b/torch/testing/_internal/common_modules.py
@@ -257,6 +257,31 @@ return module_inputs +def module_inputs_torch_nn_GaussianNLLLoss(module_info, device, dtype, requires_grad, **kwargs): + make_input = partial(make_tensor, device=device, dtype=dtype, requires_grad=requires_grad) + make_target = partial(make_tensor, device=device, dtype=dtype, requires_grad=False) + + cases: List[Tuple[str, dict]] = [ + ('', {}), + ('reduction_sum', {'reduction': 'sum'}), + ('reduction_mean', {'reduction': 'mean'}), + ('reduction_none', {'reduction': 'none'}), + ] + + module_inputs = [] + for desc, constructor_kwargs in cases: + module_inputs.append( + ModuleInput(constructor_input=FunctionInput(**constructor_kwargs), + forward_input=FunctionInput(make_input((3)), + make_target((3)), + make_input((1)).abs()), + desc=desc, + reference_fn=no_batch_dim_reference_fn) + ) + + return module_inputs + + def no_batch_dim_reference_fn(m, p, *args, **kwargs): """Reference function for modules supporting no batch dimensions. @@ -527,6 +552,8 @@ ]), ModuleInfo(torch.nn.NLLLoss, module_inputs_func=module_inputs_torch_nn_NLLLoss), + ModuleInfo(torch.nn.GaussianNLLLoss, + module_inputs_func=module_inputs_torch_nn_GaussianNLLLoss), ModuleInfo(torch.nn.Hardswish, module_inputs_func=module_inputs_torch_nn_Hardswish, supports_gradgrad=False),