Skip to content

Commit bce5578

Browse files
authored
Adding support for atomics lowering (#5)
Adding the lowering path from `tt.atomic_rmw` -> `tts.unstructured_atomic_rmw` -> `linalg.generic + memref.atomic_rmw` for scenarios when tensor of masks is required for the atomic ops.
1 parent e0c5133 commit bce5578

4 files changed

Lines changed: 229 additions & 4 deletions

File tree

include/triton-shared/Dialect/TritonStructured/IR/TritonStructuredDialect.td

Lines changed: 28 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010

1111
include "mlir/IR/OpBase.td"
1212
include "triton/Dialect/Triton/IR/TritonTypes.td"
13+
include "triton/Dialect/Triton/IR/TritonAttrDefs.td"
1314
include "mlir/Interfaces/SideEffectInterfaces.td"
1415

1516
def Triton_Structured_Dialect : Dialect {
@@ -297,6 +298,33 @@ def TTS_ScatterOp : TTS_Op<"scatter", [
297298
}];
298299
}
299300

301+
def TTS_UnstructuredAtomicRMWOp : TTS_Op<"unstructured_atomic_rmw", [
302+
MemoryEffects<[MemRead, MemWrite]>,
303+
OptionalTypesMatchWith<"mask type matches offset type", "offset", "mask",
304+
"triton::getI1SameShape($_self)">
305+
]> {
306+
let summary = "atomic read-modify-write with a tensor of per-element offsets";
307+
308+
let arguments = (
309+
ins
310+
TT_Ptr:$ptr,
311+
TT_IntLike:$offset,
312+
TT_Type:$value,
313+
Optional<TT_BoolLike>:$mask,
314+
TT_AtomicRMWAttr:$atomic_rmw_op,
315+
TT_MemSemanticAttr:$sem,
316+
TT_MemSyncScopeAttr:$scope
317+
);
318+
319+
let results = (outs TT_Type:$result);
320+
321+
let assemblyFormat = [{
322+
$atomic_rmw_op `,` $sem `,` $scope `,` $ptr `[` $offset `]` `,` $value
323+
(`mask` `=` $mask^)?
324+
attr-dict `:` `(` type($ptr) `,` type($offset) `,` type($value) `)` `->` type($result)
325+
}];
326+
}
327+
300328
def TTS_LoadOp : TTS_Op<"load", [
301329
MemoryEffects<[MemRead]>,
302330
AttrSizedOperandSegments

lib/Conversion/TritonToUnstructured/TritonToUnstructuredPass.cpp

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -395,8 +395,8 @@ class TritonToUnstructuredPass
395395
})
396396
.Case<tts::MakeGatherScatterTensorPtrOp>(
397397
[&](Operation *op) { return success(); })
398-
.Case<triton::LoadOp, triton::StoreOp, tts::MakeTensorPtrOp>(
399-
[&](Operation *op) {
398+
.Case<triton::LoadOp, triton::StoreOp, triton::AtomicRMWOp,
399+
tts::MakeTensorPtrOp>([&](Operation *op) {
400400
// Special case:
401401
// We do not want to create "unstructured tensor pointer"
402402
// into tts.make_tptr if the base pointer is directly from
@@ -508,6 +508,17 @@ class TritonToUnstructuredPass
508508
store->erase();
509509
return success();
510510
})
511+
.Case<triton::AtomicRMWOp>([&](triton::AtomicRMWOp atomicOp) {
512+
auto offsetInfo = offsetMap.at(atomicOp.getPtr());
513+
auto newOp = tts::UnstructuredAtomicRMWOp::create(
514+
b, loc, atomicOp.getType(), offsetInfo.ptr,
515+
offsetInfo.offset, atomicOp.getVal(), atomicOp.getMask(),
516+
atomicOp.getAtomicRmwOp(), atomicOp.getSem(),
517+
atomicOp.getScope());
518+
atomicOp->replaceAllUsesWith(newOp->getResults());
519+
atomicOp->erase();
520+
return success();
521+
})
511522
.Case<tts::MakeTensorPtrOp>([&](auto makeTensorPtr) {
512523
// For block pointers, the base could come from a sequence of
513524
// `tt.addptr`. Accumulate the target offset with the offset

lib/Conversion/UnstructuredToMemref/UnstructuredToMemrefPass.cpp

Lines changed: 114 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -376,6 +376,116 @@ struct ScatterConverter : public OpConversionPattern<tts::ScatterOp> {
376376
}
377377
};
378378

379+
// Lowering tts.unstructured_atomic_rmw to linalg.generic { memref.atomic_rmw }.
380+
struct UnstructuredAtomicRMWConverter
381+
: public OpConversionPattern<tts::UnstructuredAtomicRMWOp> {
382+
using OpConversionPattern<tts::UnstructuredAtomicRMWOp>::OpConversionPattern;
383+
384+
UnstructuredAtomicRMWConverter(const TypeConverter &typeConverter,
385+
MLIRContext *context)
386+
: OpConversionPattern<tts::UnstructuredAtomicRMWOp>(typeConverter,
387+
context) {}
388+
389+
UnstructuredAtomicRMWConverter(MLIRContext *context)
390+
: OpConversionPattern<tts::UnstructuredAtomicRMWOp>(context) {}
391+
392+
static arith::AtomicRMWKind convertRMWKind(triton::RMWOp rmwOp) {
393+
switch (rmwOp) {
394+
case triton::RMWOp::FADD: return arith::AtomicRMWKind::addf;
395+
case triton::RMWOp::ADD: return arith::AtomicRMWKind::addi;
396+
case triton::RMWOp::MAX: return arith::AtomicRMWKind::maxs;
397+
case triton::RMWOp::MIN: return arith::AtomicRMWKind::mins;
398+
case triton::RMWOp::UMAX: return arith::AtomicRMWKind::maxu;
399+
case triton::RMWOp::UMIN: return arith::AtomicRMWKind::minu;
400+
case triton::RMWOp::AND: return arith::AtomicRMWKind::andi;
401+
case triton::RMWOp::OR: return arith::AtomicRMWKind::ori;
402+
case triton::RMWOp::XOR: return arith::AtomicRMWKind::xori;
403+
case triton::RMWOp::XCHG: return arith::AtomicRMWKind::assign;
404+
}
405+
llvm_unreachable("Unknown triton::RMWOp");
406+
}
407+
408+
LogicalResult
409+
matchAndRewrite(tts::UnstructuredAtomicRMWOp atomicOp, OpAdaptor adaptor,
410+
ConversionPatternRewriter &rewriter) const override {
411+
auto loc = atomicOp->getLoc();
412+
413+
auto ptr = adaptor.getPtr();
414+
auto offsetTensor = adaptor.getOffset();
415+
auto valueTensor = adaptor.getValue();
416+
417+
// Must be a tensor (not scalar) offset.
418+
if (!isa<ShapedType>(offsetTensor.getType()))
419+
return failure();
420+
421+
auto valueType = dyn_cast<RankedTensorType>(atomicOp.getValue().getType());
422+
if (!valueType)
423+
return failure();
424+
425+
// Treat the base pointer (memref) as 1D because the offsets are all
426+
// relative to a single base pointer (already collapsed).
427+
auto baseMemref =
428+
memref::CastOp::create(
429+
rewriter, loc,
430+
MemRefType::get({ShapedType::kDynamic}, valueType.getElementType()),
431+
ptr)
432+
.getResult();
433+
434+
SmallVector<Value> inputs{offsetTensor, valueTensor};
435+
if (atomicOp.getMask())
436+
inputs.push_back(atomicOp.getMask());
437+
438+
SmallVector<AffineMap> affineMaps(
439+
inputs.size() + 1,
440+
rewriter.getMultiDimIdentityMap(valueType.getRank()));
441+
442+
SmallVector<utils::IteratorType> iteratorTypes(
443+
valueType.getRank(), utils::IteratorType::parallel);
444+
445+
Value emptyTensor =
446+
tensor::EmptyOp::create(rewriter, loc, valueType.getShape(),
447+
valueType.getElementType())
448+
.getResult();
449+
450+
arith::AtomicRMWKind kind = convertRMWKind(atomicOp.getAtomicRmwOp());
451+
452+
auto genericOp = linalg::GenericOp::create(
453+
rewriter, loc, TypeRange{valueType}, inputs, ValueRange{emptyTensor},
454+
affineMaps, iteratorTypes,
455+
[&](OpBuilder &b, Location loc, ValueRange args) {
456+
auto offset = args[0];
457+
auto value = args[1];
458+
459+
auto doAtomic = [&](Value idx, Value val) -> Value {
460+
Value i =
461+
arith::IndexCastOp::create(b, loc, b.getIndexType(), idx);
462+
return memref::AtomicRMWOp::create(b, loc, kind, val, baseMemref,
463+
ValueRange{i})
464+
.getResult();
465+
};
466+
467+
if (!atomicOp.getMask()) {
468+
linalg::YieldOp::create(b, loc, doAtomic(offset, value));
469+
} else {
470+
auto mask = args[2];
471+
auto passThru = args.back(); // output iter arg
472+
auto ifOp = scf::IfOp::create(
473+
b, loc, mask,
474+
[&](OpBuilder &b, Location loc) {
475+
scf::YieldOp::create(b, loc, doAtomic(offset, value));
476+
},
477+
[&](OpBuilder &b, Location loc) {
478+
scf::YieldOp::create(b, loc, passThru);
479+
});
480+
linalg::YieldOp::create(b, loc, ifOp.getResult(0));
481+
}
482+
});
483+
484+
rewriter.replaceOp(atomicOp, genericOp.getResult(0));
485+
return success();
486+
}
487+
};
488+
379489
class UnstructuredToMemrefPass
380490
: public ::impl::UnstructuredToMemrefBase<UnstructuredToMemrefPass> {
381491

@@ -401,11 +511,13 @@ class UnstructuredToMemrefPass
401511
bufferization::BufferizationDialect, memref::MemRefDialect,
402512
ttx::TritonTilingExtDialect>();
403513

404-
target.addIllegalOp<tts::GatherOp, tts::ScatterOp>();
514+
target.addIllegalOp<tts::GatherOp, tts::ScatterOp,
515+
tts::UnstructuredAtomicRMWOp>();
405516

406517
PtrToUnrankedMemrefConverter typeConverter;
407518

408-
patterns.add<GatherConverter, ScatterConverter, ScalarLoadConverter,
519+
patterns.add<GatherConverter, ScatterConverter,
520+
UnstructuredAtomicRMWConverter, ScalarLoadConverter,
409521
ScalarStoreConverter>(typeConverter, patterns.getContext());
410522

411523
if (failed(applyPartialConversion(moduleOp, target, std::move(patterns))))
Lines changed: 74 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,74 @@
1+
// RUN: triton-shared-opt --triton-to-unstructured --canonicalize --unstructured-to-memref --canonicalize %s | FileCheck %s
2+
//
3+
// CHECK-DAG: #[[MAP:.+]] = affine_map<(d0) -> (d0)>
4+
5+
module {
6+
7+
// CHECK-LABEL: tt.func public @atomic_fadd_no_mask
8+
// CHECK-NOT: tt.atomic_rmw
9+
// CHECK-NOT: tts.unstructured_atomic_rmw
10+
// CHECK: [[CAST:%.+]] = builtin.unrealized_conversion_cast %arg0 : !tt.ptr<f32> to memref<*xf32>
11+
// CHECK: [[BASE:%.+]] = memref.cast [[CAST]] : memref<*xf32> to memref<?xf32>
12+
// CHECK: linalg.generic {indexing_maps = [#[[MAP]], #[[MAP]], #[[MAP]]], iterator_types = ["parallel"]} ins(%arg2, %arg1 : tensor<64xi32>, tensor<64xf32>)
13+
// CHECK: ^bb0([[IDX:%.+]]: i32, [[VAL:%.+]]: f32, {{.*}}: f32):
14+
// CHECK: [[I:%.+]] = arith.index_cast [[IDX]] : i32 to index
15+
// CHECK: [[OLD:%.+]] = memref.atomic_rmw addf [[VAL]], [[BASE]]{{\[}}[[I]]{{\]}} : (f32, memref<?xf32>) -> f32
16+
// CHECK: linalg.yield [[OLD]] : f32
17+
tt.func public @atomic_fadd_no_mask(
18+
%out_ptr: !tt.ptr<f32>,
19+
%values: tensor<64xf32>,
20+
%offsets: tensor<64xi32>
21+
) -> tensor<64xf32> {
22+
%splat = tt.splat %out_ptr : !tt.ptr<f32> -> tensor<64x!tt.ptr<f32>>
23+
%ptr = tt.addptr %splat, %offsets : tensor<64x!tt.ptr<f32>>, tensor<64xi32>
24+
%old = tt.atomic_rmw fadd, acq_rel, gpu, %ptr, %values
25+
: (tensor<64x!tt.ptr<f32>>, tensor<64xf32>) -> tensor<64xf32>
26+
tt.return %old : tensor<64xf32>
27+
}
28+
29+
// CHECK-LABEL: tt.func public @atomic_fadd_with_mask
30+
// CHECK-NOT: tt.atomic_rmw
31+
// CHECK-NOT: tts.unstructured_atomic_rmw
32+
// CHECK: [[CAST2:%.+]] = builtin.unrealized_conversion_cast %arg0 : !tt.ptr<f32> to memref<*xf32>
33+
// CHECK: [[BASE2:%.+]] = memref.cast [[CAST2]] : memref<*xf32> to memref<?xf32>
34+
// CHECK: linalg.generic {indexing_maps = [#[[MAP]], #[[MAP]], #[[MAP]], #[[MAP]]], iterator_types = ["parallel"]} ins(%arg2, %arg1, %arg3 : tensor<64xi32>, tensor<64xf32>, tensor<64xi1>)
35+
// CHECK: ^bb0([[IDX2:%.+]]: i32, [[VAL2:%.+]]: f32, [[MASK2:%.+]]: i1, [[OUT2:%.+]]: f32):
36+
// CHECK: [[RES2:%.+]] = scf.if [[MASK2]] -> (f32) {
37+
// CHECK: [[I2:%.+]] = arith.index_cast [[IDX2]] : i32 to index
38+
// CHECK: [[OLD2:%.+]] = memref.atomic_rmw addf [[VAL2]], [[BASE2]]{{\[}}[[I2]]{{\]}} : (f32, memref<?xf32>) -> f32
39+
// CHECK: scf.yield [[OLD2]] : f32
40+
// CHECK: } else {
41+
// CHECK: scf.yield [[OUT2]] : f32
42+
// CHECK: }
43+
// CHECK: linalg.yield [[RES2]] : f32
44+
tt.func public @atomic_fadd_with_mask(
45+
%out_ptr: !tt.ptr<f32>,
46+
%values: tensor<64xf32>,
47+
%offsets: tensor<64xi32>,
48+
%mask: tensor<64xi1>
49+
) -> tensor<64xf32> {
50+
%splat = tt.splat %out_ptr : !tt.ptr<f32> -> tensor<64x!tt.ptr<f32>>
51+
%ptr = tt.addptr %splat, %offsets : tensor<64x!tt.ptr<f32>>, tensor<64xi32>
52+
%old = tt.atomic_rmw fadd, acq_rel, gpu, %ptr, %values, %mask
53+
: (tensor<64x!tt.ptr<f32>>, tensor<64xf32>, tensor<64xi1>)
54+
-> tensor<64xf32>
55+
tt.return %old : tensor<64xf32>
56+
}
57+
58+
// CHECK-LABEL: tt.func public @atomic_addi_no_mask
59+
// CHECK-NOT: tt.atomic_rmw
60+
// CHECK-NOT: tts.unstructured_atomic_rmw
61+
// CHECK: memref.atomic_rmw addi {{.*}}, {{.*}}[{{.*}}] : (i32, memref<?xi32>) -> i32
62+
tt.func public @atomic_addi_no_mask(
63+
%out_ptr: !tt.ptr<i32>,
64+
%values: tensor<64xi32>,
65+
%offsets: tensor<64xi32>
66+
) -> tensor<64xi32> {
67+
%splat = tt.splat %out_ptr : !tt.ptr<i32> -> tensor<64x!tt.ptr<i32>>
68+
%ptr = tt.addptr %splat, %offsets : tensor<64x!tt.ptr<i32>>, tensor<64xi32>
69+
%old = tt.atomic_rmw add, acq_rel, gpu, %ptr, %values
70+
: (tensor<64x!tt.ptr<i32>>, tensor<64xi32>) -> tensor<64xi32>
71+
tt.return %old : tensor<64xi32>
72+
}
73+
74+
}

0 commit comments

Comments
 (0)