Skip to content

Commit be99753

Browse files
committed
guard triton version and sm version better
Signed-off-by: Hao Wu <skyw@nvidia.com>
1 parent f6bc122 commit be99753

4 files changed

Lines changed: 24 additions & 6 deletions

File tree

emerging_optimizers/orthogonalized_optimizers/muon.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,8 +16,10 @@
1616
from typing import Callable
1717

1818
import torch
19+
from absl import logging
1920
from torch.optim.optimizer import ParamsT
2021

22+
from emerging_optimizers import triton_kernels
2123
from emerging_optimizers.orthogonalized_optimizers.muon_utils import newton_schulz
2224
from emerging_optimizers.orthogonalized_optimizers.orthogonalized_optimizer import OrthogonalizedOptimizer, _args_doc
2325

@@ -80,6 +82,18 @@ def __init__(
8082
if num_ns_steps < 1:
8183
raise ValueError(f"num_ns_steps must be at least 1, got {num_ns_steps}")
8284

85+
if torch.cuda.is_available():
86+
sm_version = torch.cuda.get_device_capability()
87+
else:
88+
sm_version = (0, 0)
89+
if not triton_kernels.HAS_TRITON_340: # type: ignore[attr-defined]
90+
logging.error("Triton 3.4.0 or higher is required for use_syrk to be True.")
91+
use_syrk = False
92+
elif sm_version not in ((8, 0), (9, 0), (10, 0), (10, 3)):
93+
logging.error(
94+
f"Correctness of Triton kernel on SM {sm_version} cannot be guaranteed. Setting use_syrk to False."
95+
)
96+
use_syrk = False
8397
orthogonalize_fn = partial(
8498
newton_schulz, steps=num_ns_steps, coefficient_type=coefficient_type, use_syrk=use_syrk
8599
)

emerging_optimizers/orthogonalized_optimizers/muon_utils.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -278,6 +278,9 @@ def newton_schulz_step_tsyrk(
278278
Returns:
279279
The orthogonalization of X.
280280
"""
281+
assert triton_kernels.HAS_TRITON_340, ( # type: ignore[attr-defined]
282+
"Triton version doesn't support tensor descriptor API. Minimum required version is 3.4.0."
283+
)
281284
A = triton_kernels.tsyrk_ex(X) # type: ignore[attr-defined]
282285
if tp_group is not None:
283286
torch.distributed.all_reduce(A, op=torch.distributed.ReduceOp.SUM, group=tp_group)

emerging_optimizers/triton_kernels/syrk.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -20,12 +20,13 @@
2020

2121
try:
2222
from triton.tools.tensor_descriptor import TensorDescriptor
23+
24+
HAS_TRITON_340 = True
2325
except ImportError:
24-
raise ImportError(
25-
f"Triton version ({triton.__version__}) doesn't support tensor descriptor API. Minimum required version is 3.4.0."
26-
)
26+
HAS_TRITON_340 = False
27+
2728

28-
__all__ = ["ssyrk", "tsyrk_ex"]
29+
__all__ = ["ssyrk", "tsyrk_ex", "HAS_TRITON_340"]
2930

3031

3132
@triton.jit

tests/test_muon_utils.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
from emerging_optimizers.orthogonalized_optimizers import muon, muon_utils
2323

2424

25-
_SM_VERSION = torch.cuda.get_device_capability() if torch.cuda.is_available() else None
25+
_SM_VERSION = torch.cuda.get_device_capability() if torch.cuda.is_available() else (0, 0)
2626

2727

2828
def newton_schulz_ref(x: torch.Tensor, coefficient_sets: list[tuple[float, float, float]]) -> torch.Tensor:
@@ -212,7 +212,7 @@ def test_qkv_split_shapes_validation(self):
212212

213213

214214
@absltest.skipIf(
215-
_SM_VERSION is None or _SM_VERSION not in ((8, 0), (9, 0), (10, 0), (11, 0)),
215+
_SM_VERSION not in ((8, 0), (9, 0), (10, 0), (10, 3)),
216216
f"Correctness of Triton kernel on SM {_SM_VERSION} cannot be guaranteed.",
217217
)
218218
class TestNewtonSchulzStepWithTsyrk(parameterized.TestCase):

0 commit comments

Comments
 (0)