Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions thunder/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
5 changes: 5 additions & 0 deletions thunder/tests/opinfos.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down
Loading