Skip to content
Merged
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
50 changes: 44 additions & 6 deletions python/triton/language/semantic.py
Original file line number Diff line number Diff line change
Expand Up @@ -1362,9 +1362,28 @@ def atomic_max(self, ptr: TensorTy, val: TensorTy, mask: TensorTy, sem: str, sco
return self.tensor(
self.builder.create_atomic_rmw(ir.ATOMIC_OP.UMAX, ptr.handle, val.handle, mask.handle, sem, scope),
val.type)
# Design for NPU
return self.tensor(
self.builder.create_atomic_rmw(ir.ATOMIC_OP.MAX, ptr.handle, val.handle, mask.handle, sem, scope), val.type)
# for float
# return atomic_smax(i_ptr, i_val) if val >= 0
# return atomic_umin(i_ptr, i_val) if val < 0
if sca_ty not in {tl.float32, tl.float64}:
raise TypeError(f"atomic_max not supported for dtype {sca_ty}")

i_type = tl.int32 if sca_ty == tl.float32 else tl.int64
i_val = self.bitcast(val, i_type)
i_ptr = self.bitcast(ptr, tl.pointer_type(i_type, 1))
ui_type = tl.uint32 if sca_ty == tl.float32 else tl.uint64
ui_val = self.bitcast(val, ui_type)
ui_ptr = self.bitcast(ptr, tl.pointer_type(ui_type, 1))
Comment thread
CHNJZ marked this conversation as resolved.
neg = self._signbit(val)
pos = self.not_(neg)
pos_ret = self.tensor(
self.builder.create_atomic_rmw(ir.ATOMIC_OP.MAX, i_ptr.handle, i_val.handle,
self.and_(mask, pos).handle, sem, scope), i_val.type)
neg_ret = self.tensor(
self.builder.create_atomic_rmw(ir.ATOMIC_OP.UMIN, ui_ptr.handle, ui_val.handle,
self.and_(mask, neg).handle, sem, scope), ui_val.type)
ret = self.where(pos, pos_ret, neg_ret)
return self.bitcast(ret, sca_ty)

def atomic_min(self, ptr: TensorTy, val: TensorTy, mask: TensorTy, sem: str, scope: str) -> TensorTy:
ptr, val, mask = self.atom_red_typechecking_impl(ptr, val, mask, 'min')
Expand All @@ -1381,9 +1400,28 @@ def atomic_min(self, ptr: TensorTy, val: TensorTy, mask: TensorTy, sem: str, sco
return self.tensor(
self.builder.create_atomic_rmw(ir.ATOMIC_OP.UMIN, ptr.handle, val.handle, mask.handle, sem, scope),
val.type)
# Design for NPU
return self.tensor(
self.builder.create_atomic_rmw(ir.ATOMIC_OP.MIN, ptr.handle, val.handle, mask.handle, sem, scope), val.type)
# for float
# return atomic_smin(i_ptr, i_val) if val >= 0
# return atomic_umax(i_ptr, i_val) if val < 0
if sca_ty not in {tl.float32, tl.float64}:
raise TypeError(f"atomic_min not supported for dtype {sca_ty}")

i_type = tl.int32 if sca_ty == tl.float32 else tl.int64
i_val = self.bitcast(val, i_type)
i_ptr = self.bitcast(ptr, tl.pointer_type(i_type, 1))
ui_type = tl.uint32 if sca_ty == tl.float32 else tl.uint64
ui_val = self.bitcast(val, ui_type)
ui_ptr = self.bitcast(ptr, tl.pointer_type(ui_type, 1))
neg = self._signbit(val)
pos = self.not_(neg)
pos_ret = self.tensor(
self.builder.create_atomic_rmw(ir.ATOMIC_OP.MIN, i_ptr.handle, i_val.handle,
self.and_(mask, pos).handle, sem, scope), i_val.type)
neg_ret = self.tensor(
self.builder.create_atomic_rmw(ir.ATOMIC_OP.UMAX, ui_ptr.handle, ui_val.handle,
self.and_(mask, neg).handle, sem, scope), ui_ptr.type)
Comment thread
CHNJZ marked this conversation as resolved.
Comment thread
CHNJZ marked this conversation as resolved.
Comment thread
CHNJZ marked this conversation as resolved.
Comment thread
CHNJZ marked this conversation as resolved.
ret = self.where(pos, pos_ret, neg_ret)
return self.bitcast(ret, sca_ty)

def atomic_add(self, ptr: TensorTy, val: TensorTy, mask: TensorTy, sem: str, scope: str) -> TensorTy:
ptr, val, mask = self.atom_red_typechecking_impl(ptr, val, mask, 'add')
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,4 +11,5 @@ add_triton_library(DiscreteMaskAccessConversion
MLIRTransforms
MLIRSupport
TritonIR
TritonToLinalg
)
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
#include "Utils/Utils.h"
#include "ascend/include/DiscreteMaskAccessConversion/Passes.h"

#include "ascend/include/TritonToLinalg/LoadStoreConverter.h"
#include "ascend/include/TritonToLinalg/MaskAnalysis.h"
#include "ascend/include/TritonToStructured/MemOpConverter.h"
#include "bishengir/Dialect/HIVM/IR/HIVM.h"
Expand Down Expand Up @@ -465,6 +466,20 @@ void DiscreteMaskAccessConversionPass::runOnOperation() {
enableSyncBlockLockFlag = !tileNonOverlap;
auto moduleOp = getOperation();

// Restore floating-point atomic max/min expanded by semantic.py before
// discrete-mask rewriting changes the atomic value into arith.select.
// Run this in a separate greedy-rewrite phase so that
// DiscreteMaskAtomicConversion cannot consume the expanded form first.
RewritePatternSet atomicMaxMinPatterns(&getContext());
atomicMaxMinPatterns.add<LoadStoreConverter::AtomicMaxMinCanonicalizer>(
atomicMaxMinPatterns.getContext());
if (failed(
applyPatternsGreedily(moduleOp, std::move(atomicMaxMinPatterns)))) {
moduleOp->emitError("failed to canonicalize floating-point atomic max/min");
signalPassFailure();
return;
}

RewritePatternSet patterns(&getContext());
patterns.add<DiscreteMaskLoadConversion, DiscreteMaskStoreConversion,
DiscreteMaskAtomicConversion>(patterns.getContext());
Expand Down
31 changes: 30 additions & 1 deletion third_party/ascend/lib/TritonToLinalg/LoadStoreConverter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1158,6 +1158,8 @@ AtomicMaxMinCanonicalizer::matchAndRewrite(triton::AtomicRMWOp op,
// if the return value of op is used, we can't simply erase it
if (op.getResult().use_empty()) {
rewriter.eraseOp(op);
if (ptrBitcastOp->use_empty())
rewriter.eraseOp(ptrBitcastOp);
return success();
}
return failure();
Expand All @@ -1179,7 +1181,29 @@ AtomicMaxMinCanonicalizer::matchAndRewrite(triton::AtomicRMWOp op,
if (auto andOp = originalMask.getDefiningOp<arith::AndIOp>())
// LHS is convention in semantic interpreter
originalMask = andOp.getLhs();
else if (auto cmpOp = originalMask.getDefiningOp<arith::CmpFOp>()) {
else if (auto xorOp = originalMask.getDefiningOp<arith::XOrIOp>()) {
// Current f32 atomic_min uses !signbit as the positive mask:
// shrui(value_bits, 31) -> cmpi ne 0 -> xori true.
if ((rmwOp != triton::RMWOp::MIN && rmwOp != triton::RMWOp::MAX) ||
(!elementType.isF32() && !elementType.isF64()) ||
!matchPattern(xorOp.getRhs(), m_One())) {
return failure();
}

auto cmpOp = xorOp.getLhs().getDefiningOp<arith::CmpIOp>();
if (!cmpOp || cmpOp.getPredicate() != arith::CmpIPredicate::ne ||
!matchPattern(cmpOp.getRhs(), m_Zero()))
return op->emitError("Illegal mask for atomicrmwOp of float type");

auto shiftOp = cmpOp.getLhs().getDefiningOp<arith::ShRUIOp>();
if (!shiftOp || shiftOp.getLhs() != valueBitcastOp.getResult())
return op->emitError("Illegal mask for atomicrmwOp of float type");

// Restore the implicit all-true mask.
originalMask = rewriter.create<arith::ConstantOp>(
op->getLoc(),
DenseElementsAttr::get(cast<ShapedType>(op.getMask().getType()), true));
} else if (auto cmpOp = originalMask.getDefiningOp<arith::CmpFOp>()) {
if (cmpOp.getPredicate() != mlir::arith::CmpFPredicate::OGE ||
!matchPattern(cmpOp.getRhs(),
/*positive float zero matcher*/ m_PosZeroFloat()))
Expand Down Expand Up @@ -1235,6 +1259,11 @@ AtomicMaxMinCanonicalizer::matchAndRewrite(triton::AtomicRMWOp op,
rewriter.eraseOp(op);
}

// The restored atomic uses ptrBitcastOp.getSrc(), i.e. the original GM
// pointer. Remove the pointer bitcast once the paired integer atomic is gone.
if (ptrBitcastOp->use_empty())
rewriter.eraseOp(ptrBitcastOp);

return success();
}

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
# Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in
# all copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
# THE SOFTWARE.

import torch
import torch_npu
import triton
import triton.language as tl


@triton.jit
def float_indirect_atomic_min_kernel(
value_ptr,
index_ptr,
output_ptr,
n_elements,
BLOCK_SIZE: tl.constexpr,
):
offsets = tl.program_id(0) * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements

values = tl.load(value_ptr + offsets, mask=mask, other=0.0)
indices = tl.load(index_ptr + offsets, mask=mask, other=0)
tl.atomic_min(output_ptr + indices, values, mask=mask)


def test_float_indirect_atomic_min():
n_elements = 257
block_size = 128

values_cpu = torch.linspace(-4.0, 4.0, n_elements, dtype=torch.float32)
indices_cpu = torch.arange(n_elements - 1, -1, -1, dtype=torch.int64)

values = values_cpu.npu()
indices = indices_cpu.npu()
output = torch.full((n_elements, ), float("inf"), dtype=torch.float32, device="npu")

grid = (triton.cdiv(n_elements, block_size), )
float_indirect_atomic_min_kernel[grid](
values,
indices,
output,
n_elements,
BLOCK_SIZE=block_size,
)

expected = torch.full((n_elements, ), float("inf"), dtype=torch.float32)
expected[indices_cpu] = values_cpu

torch.testing.assert_close(output.cpu(), expected, rtol=0, atol=0)
Loading