diff --git a/python/triton/language/semantic.py b/python/triton/language/semantic.py index 52192cf3ee..42bf7024f4 100644 --- a/python/triton/language/semantic.py +++ b/python/triton/language/semantic.py @@ -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)) + 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') @@ -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) + 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') diff --git a/third_party/ascend/lib/DiscreteMaskAccessConversion/CMakeLists.txt b/third_party/ascend/lib/DiscreteMaskAccessConversion/CMakeLists.txt index 0c7f145684..5317125f4a 100644 --- a/third_party/ascend/lib/DiscreteMaskAccessConversion/CMakeLists.txt +++ b/third_party/ascend/lib/DiscreteMaskAccessConversion/CMakeLists.txt @@ -11,4 +11,5 @@ add_triton_library(DiscreteMaskAccessConversion MLIRTransforms MLIRSupport TritonIR + TritonToLinalg ) diff --git a/third_party/ascend/lib/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.cpp b/third_party/ascend/lib/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.cpp index 9a4158b35f..e31c9f44f4 100644 --- a/third_party/ascend/lib/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.cpp +++ b/third_party/ascend/lib/DiscreteMaskAccessConversion/DiscreteMaskAccessConversionPass.cpp @@ -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" @@ -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( + 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(patterns.getContext()); diff --git a/third_party/ascend/lib/TritonToLinalg/LoadStoreConverter.cpp b/third_party/ascend/lib/TritonToLinalg/LoadStoreConverter.cpp index 446694912f..114521e4f9 100644 --- a/third_party/ascend/lib/TritonToLinalg/LoadStoreConverter.cpp +++ b/third_party/ascend/lib/TritonToLinalg/LoadStoreConverter.cpp @@ -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(); @@ -1179,7 +1181,29 @@ AtomicMaxMinCanonicalizer::matchAndRewrite(triton::AtomicRMWOp op, if (auto andOp = originalMask.getDefiningOp()) // LHS is convention in semantic interpreter originalMask = andOp.getLhs(); - else if (auto cmpOp = originalMask.getDefiningOp()) { + else if (auto xorOp = originalMask.getDefiningOp()) { + // 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(); + 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(); + 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( + op->getLoc(), + DenseElementsAttr::get(cast(op.getMask().getType()), true)); + } else if (auto cmpOp = originalMask.getDefiningOp()) { if (cmpOp.getPredicate() != mlir::arith::CmpFPredicate::OGE || !matchPattern(cmpOp.getRhs(), /*positive float zero matcher*/ m_PosZeroFloat())) @@ -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(); } diff --git a/third_party/ascend/unittest/pytest_ut/test_float_indirect_atomic_min.py b/third_party/ascend/unittest/pytest_ut/test_float_indirect_atomic_min.py new file mode 100644 index 0000000000..1000605879 --- /dev/null +++ b/third_party/ascend/unittest/pytest_ut/test_float_indirect_atomic_min.py @@ -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)