From 74d1cb5433ac54645b33c57e2f4aedf9939f7593 Mon Sep 17 00:00:00 2001 From: Bhimraj Yadav Date: Thu, 24 Sep 2026 08:30:40 +0000 Subject: [PATCH] Skip degenerate group_norm samples on torch 2.14+ torch 2.14 rejects group_norm inputs with an empty channel or spatial extent ("Expected number of channels to be greater than 0", "Expected HxW to be greater than 0"); 2.13 and earlier accepted them. The sample generator builds shapes from (0, 2) in every dim, so the torch reference failed on inputs it now refuses. An empty batch is still accepted, so only channels and the spatial dims are filtered. The boundary is 2.14, not 2.13, hence the new constant. --- thunder/constants.py | 1 + thunder/tests/opinfos.py | 5 +++++ 2 files changed, 6 insertions(+) diff --git a/thunder/constants.py b/thunder/constants.py index 5a49f674a0..29d6e8c292 100644 --- a/thunder/constants.py +++ b/thunder/constants.py @@ -5,3 +5,4 @@ # PyTorch version constants # `use_base_version` so that dev builds like "2.13.0a0+gitabc123" compare as "2.13.0" _TORCH_GREATER_EQUAL_2_13 = compare_version("torch", operator.ge, "2.13.0", use_base_version=True) +_TORCH_GREATER_EQUAL_2_14 = compare_version("torch", operator.ge, "2.14.0", use_base_version=True) diff --git a/thunder/tests/opinfos.py b/thunder/tests/opinfos.py index f6997926de..a39cc5fa2c 100644 --- a/thunder/tests/opinfos.py +++ b/thunder/tests/opinfos.py @@ -23,6 +23,7 @@ import thunder.core.prims as prims from thunder.core.pytree import tree_map from thunder.core.symbol import Symbol +from thunder.constants import _TORCH_GREATER_EQUAL_2_14 import thunder.executors as executors from thunder.tests.framework import _all_devicetypes, custom_comparator, IS_WINDOWS from thunder.tests.make_tensor import make_tensor, make_tensor_like @@ -8574,6 +8575,10 @@ def group_norm_sample_generator(op, device, dtype, requires_grad, **kwargs): if torch.device(device).type == "cuda" and ndim >= 3 and num_channels == 0: continue + # torch 2.14 rejects empty channel or spatial dims (an empty batch is still fine) + if _TORCH_GREATER_EQUAL_2_14 and (num_channels == 0 or math.prod(inner_dims) == 0): + continue + a = make(shape) for weight, bias in itertools.product((False, True), repeat=2):