From 1f32cef3573471529777b3558eeebc20d7fa7c94 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Fri, 17 Jul 2026 16:41:31 +0200 Subject: [PATCH 1/2] Reject negative PyTorch one-hot class counts --- .../backend_support/_pytorch_one_hot_scalar_contract.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/src/pyrecest/backend_support/_pytorch_one_hot_scalar_contract.py b/src/pyrecest/backend_support/_pytorch_one_hot_scalar_contract.py index 95db51d96..8f0d94e07 100644 --- a/src/pyrecest/backend_support/_pytorch_one_hot_scalar_contract.py +++ b/src/pyrecest/backend_support/_pytorch_one_hot_scalar_contract.py @@ -6,6 +6,7 @@ _INTEGER_MESSAGE = "num_classes must be an integer" +_NONNEGATIVE_MESSAGE = "num_classes must be non-negative" def _is_boolean_take_axis(axis, torch_module) -> bool: @@ -20,7 +21,7 @@ def _is_boolean_take_axis(axis, torch_module) -> bool: def _normalize_num_classes(num_classes, torch_module) -> int: - """Return a non-boolean integer ``num_classes`` value.""" + """Return a non-negative, non-boolean integer ``num_classes`` value.""" if isinstance(num_classes, bool) or type(num_classes).__name__ == "bool_": raise TypeError(f"{_INTEGER_MESSAGE}, not boolean") if torch_module.is_tensor(num_classes): @@ -28,9 +29,12 @@ def _normalize_num_classes(num_classes, torch_module) -> int: raise TypeError(_INTEGER_MESSAGE) num_classes = num_classes.item() try: - return _operator_index(num_classes) + normalized = _operator_index(num_classes) except TypeError as exc: raise TypeError(_INTEGER_MESSAGE) from exc + if normalized < 0: + raise ValueError(_NONNEGATIVE_MESSAGE) + return normalized def _patch_pytorch_take_axis_contract(pytorch_backend, torch_module) -> None: From bf9eb53cc02f7740a2863ff2702a2970f96a91c9 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Fri, 17 Jul 2026 16:41:46 +0200 Subject: [PATCH 2/2] Test negative PyTorch one-hot class counts --- .../test_pytorch_one_hot_scalar_contract.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/tests/backend_support/test_pytorch_one_hot_scalar_contract.py b/tests/backend_support/test_pytorch_one_hot_scalar_contract.py index c8e821642..0f45ace0e 100644 --- a/tests/backend_support/test_pytorch_one_hot_scalar_contract.py +++ b/tests/backend_support/test_pytorch_one_hot_scalar_contract.py @@ -38,6 +38,14 @@ def _scalar_contract_code(target_module): else: raise AssertionError("boolean num_classes was accepted") +for bad_num_classes in (-1, torch.tensor(-1)): + try: + target.one_hot([0], bad_num_classes) + except ValueError as exc: + assert "num_classes must be non-negative" in str(exc) + else: + raise AssertionError("negative num_classes was accepted") + device = torch.device("cuda") if torch.cuda.is_available() else torch.device("meta") device_result = target.one_hot(torch.tensor(2, device=device), 4) assert device_result.device.type == device.type