77#include " Utility.h"
88#include " triton/Dialect/TritonGPU/Transforms/Utility.h"
99#include " triton/Tools/LayoutUtils.h"
10+ #include < optional>
1011
1112using namespace mlir ;
1213using namespace mlir ::triton;
1314
15+ namespace ttg = mlir::triton::gpu;
16+
1417using ::mlir::LLVM ::delinearize;
1518using ::mlir::LLVM ::getSharedMemoryBase;
1619using ::mlir::LLVM ::getSharedMemoryObjectFromStruct;
@@ -31,6 +34,25 @@ Value maybeAnd(RewriterBase &rewriter, Location loc, Value a, Value b) {
3134 return a ? a : b;
3235}
3336
37+ #ifdef __MCTLE__
38+ std::optional<unsigned > inferPtrAddrSpace (llvm::ArrayRef<Value> ptrElems) {
39+ for (Value elem : ptrElems) {
40+ if (auto ptrTy = dyn_cast<LLVM ::LLVMPointerType>(elem.getType ()))
41+ return ptrTy.getAddressSpace ();
42+ }
43+ return std::nullopt ;
44+ }
45+
46+ bool isSharedFamilyAddressSpace (unsigned addressSpace) {
47+ return addressSpace == 3 ;
48+ }
49+
50+ bool isSharedPointerValue (llvm::ArrayRef<Value> ptrElems,
51+ unsigned defaultAddrSpace = 1 ) {
52+ return inferPtrAddrSpace (ptrElems).value_or (defaultAddrSpace) == 3 ;
53+ }
54+ #endif
55+
3456// Return a predicate that is true only if the current thread holds unique data,
3557// according to freeVarsMask. The predicate may be null to indicate no
3658// predication is required.
@@ -205,6 +227,31 @@ struct LoadStoreConversionBase {
205227 return std::min<unsigned >(128 / pointeeBitWidth, contiguity);
206228 }
207229
230+ #ifdef __MCTLE__
231+ unsigned getMaxVectorSizeByAlignment (Value ptr) const {
232+ auto tensorTy = dyn_cast<RankedTensorType>(ptr.getType ());
233+ if (!tensorTy)
234+ return 1 ;
235+ auto *axisInfo = axisAnalysisPass.getAxisInfo (ptr);
236+ if (!axisInfo || axisInfo->getRank () == 0 )
237+ return 1 ;
238+
239+ auto linAttr = ttg::toLinearEncoding (tensorTy);
240+ auto order = linAttr.getOrder ();
241+ if (order.empty () || order[0 ] >= axisInfo->getRank ())
242+ return 1 ;
243+
244+ unsigned pointeeBitWidth = triton::getPointeeBitWidth (tensorTy);
245+ if (pointeeBitWidth == 0 )
246+ return 1 ;
247+ unsigned elemBytes = std::max<unsigned >(pointeeBitWidth / 8 , 1 );
248+ unsigned maxMultipleBytes = axisInfo->getDivisibility (order[0 ]);
249+ unsigned maxMultiple = std::max<unsigned >(maxMultipleBytes / elemBytes, 1 );
250+
251+ return std::min<unsigned >(128 / pointeeBitWidth, maxMultiple);
252+ }
253+ #endif
254+
208255 unsigned getMaskAlignment (Value mask) const {
209256 return axisAnalysisPass.getMaskAlignment (mask);
210257 }
@@ -258,6 +305,23 @@ struct LoadOpConversion : public ConvertOpToLLVMPattern<triton::LoadOp>,
258305 } else {
259306 vec = getVectorSize (ptr);
260307 }
308+ #ifdef __MCTLE__
309+ auto ptrTensorTy = dyn_cast<RankedTensorType>(ptr.getType ());
310+ auto ptrElemTy = ptrTensorTy
311+ ? dyn_cast<PointerType>(ptrTensorTy.getElementType ())
312+ : PointerType ();
313+ bool isSharedTensorPtr =
314+ ptrElemTy && isSharedFamilyAddressSpace (ptrElemTy.getAddressSpace ());
315+ if (!llMask && isSharedTensorPtr) {
316+ // For TLE local/shared pointer chains, AxisInfo contiguity can be
317+ // conservative on packed contiguous lanes. The layout hint may recover
318+ // the lane grouping, but it is only legal up to the vector width whose
319+ // first element alignment is proven by AxisInfo divisibility.
320+ // unsigned hint = ttg::inferTilePointerLayoutVectorHint(ptr);
321+ unsigned alignmentBound = getMaxVectorSizeByAlignment (ptr);
322+ vec = std::max (vec, alignmentBound);
323+ }
324+ #endif
261325 unsigned numElems = getTotalElemsPerThread (ptr.getType ());
262326 unsigned constRepeatPerThread = getThreadConstRepeatTimes (ptr);
263327 if (llMask) {
@@ -272,6 +336,9 @@ struct LoadOpConversion : public ConvertOpToLLVMPattern<triton::LoadOp>,
272336 // Get the LLVM values for pointers
273337 auto ptrElems = unpackLLElements (loc, llPtr, rewriter);
274338 assert (ptrElems.size () == numElems);
339+ #ifdef __MCTLE__
340+ const bool isSharedPtr = isSharedPointerValue (ptrElems);
341+ #endif
275342
276343 // Get the LLVM values for mask
277344 SmallVector<Value> maskElems;
@@ -358,6 +425,30 @@ struct LoadOpConversion : public ConvertOpToLLVMPattern<triton::LoadOp>,
358425 assert (wordNElems * nWords * numVecs == numElems);
359426 Value pred = mask ? maskElems[vecStart] : b.int_val (1 , 1 );
360427 Value zeroVal = b.bitcast (b.int_val (valueElemNBits, 0 ), valueElemTy);
428+ #ifdef __MCTLE__
429+ if (isSharedPtr) {
430+ Value addrVal = ptrElems[vecStart];
431+ Type retTy = vec > 1 ? vec_ty (valueElemTy, vec) : valueElemTy;
432+ Value lds_values =
433+ targetInfo.loadDShared (rewriter, loc, addrVal, std::nullopt , retTy,
434+ /* pred=*/ pred);
435+
436+ if (vec == 1 ) {
437+ for (int i = 0 ; i < constRepeatPerThread; i++) {
438+ loadedVals.push_back (lds_values);
439+ }
440+ } else {
441+ for (size_t elemIndex = 0 ; elemIndex < vec; elemIndex++) {
442+ Value curr = b.extract_element (valueElemTy, lds_values,
443+ b.i32_val (elemIndex));
444+ for (int i = 0 ; i < constRepeatPerThread; i++) {
445+ loadedVals.push_back (curr);
446+ }
447+ }
448+ }
449+ continue ;
450+ }
451+ #endif
361452 if (!disableOptFlag && !op.getIsVolatile ()) {
362453 if (!disableLdgPred && isOtherValid &&
363454 (totalWidth == 128 || totalWidth == 64 || totalWidth == 32 ||
@@ -534,6 +625,9 @@ struct StoreOpConversion : public ConvertOpToLLVMPattern<triton::StoreOp>,
534625
535626 auto ptrElems = unpackLLElements (loc, llPtr, rewriter);
536627 auto valueElems = unpackLLElements (loc, llValue, rewriter);
628+ #ifdef __MCTLE__
629+ const bool isSharedPtr = isSharedPointerValue (ptrElems);
630+ #endif
537631 assert (ptrElems.size () == valueElems.size ());
538632
539633 if (valueElemTy.isFloat (8 )) {
@@ -604,6 +698,26 @@ struct StoreOpConversion : public ConvertOpToLLVMPattern<triton::StoreOp>,
604698 // TODO(Superjomn) Add cache policy fields to StoreOp.
605699 // TODO(Superjomn) Deal with cache policy here.
606700
701+ #ifdef __MCTLE__
702+ if (isSharedPtr) {
703+ Value pred = threadPred;
704+ if (llMask) {
705+ auto mask = maskElems[vecStart];
706+ pred = maybeAnd (rewriter, loc, pred, mask);
707+ }
708+ Type retTy = vec_ty (valueElemTy, vec);
709+ Value vec_values = b.undef (retTy);
710+ for (size_t elemIndex = 0 ; elemIndex < vec; elemIndex++) {
711+ Value elem = valueElems[vecStart + elemIndex];
712+ vec_values =
713+ b.insert_element (retTy, vec_values, elem, b.i32_val (elemIndex));
714+ }
715+ targetInfo.storeDShared (rewriter, loc, ptrElems[vecStart], std::nullopt ,
716+ vec_values, pred);
717+ continue ;
718+ }
719+ #endif
720+
607721 Type valArgTy = IntegerType::get (ctx, width);
608722 auto wordTy = vec_ty (valueElemTy, wordNElems);
609723 if (totalWidth == 128 || totalWidth == 64 || totalWidth == 32 ) {
0 commit comments