@@ -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+
379489class 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))))
0 commit comments