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
7 changes: 5 additions & 2 deletions python/cudnn/discrete_grouped_gemm/discrete_kernel_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@

import cutlass
import cutlass.cute as cute
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw
import cutlass.cute.testing as testing
from cutlass.cute.nvgpu import cpasync, tcgen05
from cutlass.cutlass_dsl import T, dsl_user_op
Expand Down Expand Up @@ -301,7 +302,8 @@ def atomic_max_float32(
) -> Float32:
value_int = llvm.bitcast(T.i32(), value.ir_value(loc=loc, ip=ip), loc=loc, ip=ip)

old_value_int = nvvm.atomicrmw(
old_value_int = nvvm_atomicrmw(
T.i32(),
op=cutlass._mlir.dialects.nvvm.AtomicOpKind.MAX,
ptr=ptr,
a=value_int,
Expand All @@ -320,7 +322,8 @@ def atomic_add_float32(
ip=None,
) -> Float32:
"""Atomic FP32 addition in global memory (used for dprob gradient accumulation)."""
old_value = nvvm.atomicrmw(
old_value = nvvm_atomicrmw(
T.f32(),
op=AtomicOpKind.FADD,
ptr=ptr,
a=value.ir_value(loc=loc, ip=ip),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@

import cutlass
import cutlass.cute as cute
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw
from cutlass.cutlass_dsl import (
Boolean,
Int32,
Expand Down Expand Up @@ -59,7 +60,8 @@ def atomic_add_i32(
ip=None,
) -> Int32:
"""Perform an atomic add on an int32 value in global memory."""
old_value = nvvm.atomicrmw(
old_value = nvvm_atomicrmw(
T.i32(),
op=AtomicOpKind.ADD,
ptr=ptr,
a=value.ir_value(loc=loc, ip=ip),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@

import cutlass
import cutlass.cute as cute
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw
from cutlass.cute.nvgpu import cpasync, tcgen05
from cutlass._mlir.dialects import math, nvvm, llvm
from cutlass.cutlass_dsl import T
Expand Down Expand Up @@ -1324,7 +1325,8 @@ def kernel(
# Global atomic max (accumulates across all tiles for final tensor amax)
# Since we compute absolute values, all values are non-negative
_value_int = llvm.bitcast(T.i32(), block_amax.ir_value(), loc=None, ip=None)
_old_value_int = nvvm.atomicrmw(
_old_value_int = nvvm_atomicrmw(
T.i32(),
op=nvvm.AtomicOpKind.MAX,
ptr=mAmax.iterator.llvm_ptr,
a=_value_int,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@

import cutlass
import cutlass.cute as cute
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.utils as utils
import cutlass.pipeline as pipeline
Expand All @@ -54,7 +55,8 @@ def atomic_add_float32(
loc=None,
ip=None,
) -> Float32:
old_value = nvvm.atomicrmw(
old_value = nvvm_atomicrmw(
T.f32(),
AtomicOpKind.FADD,
ptr,
value.ir_value(loc=loc, ip=ip),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@

import cutlass
import cutlass.cute as cute
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw
from cutlass.cute.nvgpu import cpasync, tcgen05
import cutlass.utils as utils
import cutlass.pipeline as pipeline
Expand Down Expand Up @@ -1704,7 +1705,8 @@ def kernel(
loc=None,
ip=None,
)
_old_value_int = cutlass._mlir.dialects.nvvm.atomicrmw(
_old_value_int = nvvm_atomicrmw(
cutlass.cutlass_dsl.T.i32(),
op=cutlass._mlir.dialects.nvvm.AtomicOpKind.MAX,
ptr=mAmax_tensor.iterator.llvm_ptr,
a=_value_int,
Expand Down
7 changes: 5 additions & 2 deletions python/cudnn/grouped_gemm/moe_kernel_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,7 @@

import cutlass
import cutlass.cute as cute
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw
import cutlass.cute.testing as testing
from cutlass.cute.nvgpu import cpasync, tcgen05
from cutlass.cutlass_dsl import T, dsl_user_op
Expand Down Expand Up @@ -283,7 +284,8 @@ def atomic_max_float32(
) -> Float32:
value_int = llvm.bitcast(T.i32(), value.ir_value(loc=loc, ip=ip), loc=loc, ip=ip)

old_value_int = nvvm.atomicrmw(
old_value_int = nvvm_atomicrmw(
T.i32(),
op=cutlass._mlir.dialects.nvvm.AtomicOpKind.MAX,
ptr=ptr,
a=value_int,
Expand All @@ -302,7 +304,8 @@ def atomic_add_float32(
ip=None,
) -> Float32:
"""Atomic FP32 addition in global memory (used for dprob gradient accumulation)."""
old_value = nvvm.atomicrmw(
old_value = nvvm_atomicrmw(
T.f32(),
op=AtomicOpKind.FADD,
ptr=ptr,
a=value.ir_value(loc=loc, ip=ip),
Expand Down
4 changes: 3 additions & 1 deletion python/cudnn/grouped_gemm/moe_persistent_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@

import cutlass
import cutlass.cute as cute
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw
from cutlass.cutlass_dsl import (
Boolean,
Int32,
Expand Down Expand Up @@ -59,7 +60,8 @@ def atomic_add_i32(
ip=None,
) -> Int32:
"""Perform an atomic add on an int32 value in global memory."""
old_value = nvvm.atomicrmw(
old_value = nvvm_atomicrmw(
T.i32(),
op=AtomicOpKind.ADD,
ptr=ptr,
a=value.ir_value(loc=loc, ip=ip),
Expand Down
10 changes: 7 additions & 3 deletions python/cudnn/grouped_gemm/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@
from cutlass.cutlass_dsl import T
from cutlass.cute.typing import Float32, Int32
import cutlass.cute as cute
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw
import cutlass
import torch
import cutlass.pipeline as pipeline
Expand All @@ -75,7 +76,8 @@ def atomic_add_i32(
ip=None,
) -> Int32:
"""Perform an atomic add on an int32 value in global memory."""
old_value = nvvm.atomicrmw(
old_value = nvvm_atomicrmw(
T.i32(),
op=AtomicOpKind.ADD,
ptr=ptr,
a=value.ir_value(loc=loc, ip=ip),
Expand Down Expand Up @@ -229,7 +231,8 @@ def atomic_max_float32(
"""
value_int = llvm.bitcast(T.i32(), value.ir_value(loc=loc, ip=ip), loc=loc, ip=ip)

old_value_int = nvvm.atomicrmw(
old_value_int = nvvm_atomicrmw(
T.i32(),
op=cutlass._mlir.dialects.nvvm.AtomicOpKind.MAX,
ptr=ptr,
a=value_int,
Expand All @@ -253,7 +256,8 @@ def atomic_add_float32(
:param value: The float32 value to add
:return: The old value at the memory location
"""
old_value = nvvm.atomicrmw(
old_value = nvvm_atomicrmw(
T.f32(),
op=cutlass._mlir.dialects.nvvm.AtomicOpKind.FADD,
ptr=ptr,
a=value.ir_value(loc=loc, ip=ip),
Expand Down
20 changes: 20 additions & 0 deletions python/cudnn/nvvm_compat.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,20 @@
# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: MIT

"""Signature-compat wrappers for nvvm dialect builders that changed across
nvidia-cutlass-dsl releases."""

import inspect

from cutlass._mlir.dialects import nvvm

# nvidia-cutlass-dsl <= 4.5.x generates atomicrmw(res, op, ptr, a, ...) with an
# explicit result type; 4.6.0+ infers the result type and dropped the parameter.
_ATOMICRMW_TAKES_RES = "res" in inspect.signature(nvvm.atomicrmw).parameters


def atomicrmw(res, op, ptr, a, *, loc=None, ip=None):
"""nvvm.atomicrmw that works on both cutlass-dsl 4.5.x and 4.6.0+."""
if _ATOMICRMW_TAKES_RES:
return nvvm.atomicrmw(res=res, op=op, ptr=ptr, a=a, loc=loc, ip=ip)
return nvvm.atomicrmw(op=op, ptr=ptr, a=a, loc=loc, ip=ip)
3 changes: 2 additions & 1 deletion python/cudnn/sdpa/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
from cutlass._mlir.dialects import llvm, nvvm # noqa: PLC2701
from cutlass.cute.runtime import from_dlpack
from cutlass.cutlass_dsl import T, dsl_user_op
from ..nvvm_compat import atomicrmw as nvvm_atomicrmw

ARCH_SM90 = 90
ARCH_SM100 = 100
Expand Down Expand Up @@ -455,7 +456,7 @@ def fadd_reduce(x: cute.TensorSSA, init_val: float | Float32 | None = None, arch
@dsl_user_op
def atomic_add_fp32(a: float | Float32, gmem_ptr: cute.Pointer, *, loc=None, ip=None) -> None:
"""Wrapper of atomic add for fp32."""
nvvm.atomicrmw(op=nvvm.AtomicOpKind.FADD, ptr=gmem_ptr.llvm_ptr, a=Float32(a).ir_value(), loc=loc, ip=ip)
nvvm_atomicrmw(T.f32(), op=nvvm.AtomicOpKind.FADD, ptr=gmem_ptr.llvm_ptr, a=Float32(a).ir_value(), loc=loc, ip=ip)


@dsl_user_op
Expand Down