From 879336654521fa656db2f969d83a12ce70999eef Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Sun, 8 Mar 2026 13:07:52 +0000 Subject: [PATCH] Fix SM-gating for CUTLASS/TMA compile options in GRPO trainer Change `get_device_capability()[0] >= 9` to `== 9` for `triton.enable_persistent_tma_matmul`, `cuda.cutlass_epilogue_fusion_enabled`, and `cuda.cutlass_tma_only` in the GRPO torch_compile_options. These options enable SM90-specific CUTLASS TMA kernels which are not valid on SM100+ (Blackwell B200/B100). The `>= 9` check incorrectly enables them on any GPU with SM >= 9, but they should only run on SM 9.x (Hopper). This matches the existing pattern in `fp8.py:test_has_fbgemm()` which documents that SM100 GPUs fail with CUTLASS SM90 kernels. --- unsloth/models/rl.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index e4f34c908e..7faa61fd59 100755 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -1322,12 +1322,12 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): if DEVICE_TYPE == "cuda": # CUDA-specific options (added to base options) cuda_options = """ - "triton.enable_persistent_tma_matmul": torch.cuda.get_device_capability()[0] >= 9,""" + "triton.enable_persistent_tma_matmul": torch.cuda.get_device_capability()[0] == 9,""" # cutlass options were added in PyTorch 2.8.0 if torch_version >= Version("2.8.0"): cuda_options += """ - "cuda.cutlass_epilogue_fusion_enabled": torch.cuda.get_device_capability()[0] >= 9, - "cuda.cutlass_tma_only": torch.cuda.get_device_capability()[0] >= 9,""" + "cuda.cutlass_epilogue_fusion_enabled": torch.cuda.get_device_capability()[0] == 9, + "cuda.cutlass_tma_only": torch.cuda.get_device_capability()[0] == 9,""" cuda_options += """ "cuda.compile_opt_level" : "-O2", "cuda.enable_cuda_lto" : True,