Skip to content

Commit 52b2a01

Browse files
committed
[Metax][TLE] Metax TLE support local pointer.
1 parent 552b8fd commit 52b2a01

9 files changed

Lines changed: 160 additions & 4 deletions

File tree

.github/workflows/metax3.6-build-and-test.yml

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -144,4 +144,7 @@ jobs:
144144
test_extract_tile_dynamic_index.py \
145145
test_extract_tile_static_index.py \
146146
test_insert_tile_dynamic_index.py \
147-
test_insert_tile_static_index.py
147+
test_insert_tile_static_index.py \
148+
test_tle_gpu_local_ptr.py \
149+
-k "not test_local_pointer_tiled_matmul_matches_torch and \
150+
not test_local_pointer_full_view_dot_avoids_pointer_convert_layout"

python/setup_tools/utils/metax.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@ def register_cache(cache, flagtree_backend, check_env, set_llvm_env):
3535
condition=is_metax and not os.environ.get("FLAGTREE_PLUGIN"),
3636
url="https://baai-cp-web.ks3-cn-beijing.ksyuncs.com/trans/metaxTritonPlugin-cpython3.12-x86_64_v0.6.1.tar.gz",
3737
copy_dst_path=f"third_party/{flagtree_backend}",
38-
md5_digest="afb7ab8f",
38+
md5_digest="de37eb01",
3939
)
4040

4141

third_party/metax/backend/compiler.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -272,7 +272,7 @@ def make_ttgir(mod, metadata, opt, capability):
272272
passes.ttgpuir.add_remove_layout_conversions(pm)
273273
passes.ttgpuir.add_optimize_thread_locality(pm)
274274
if enable_mctle:
275-
#mctle.passes.add_reject_dot_op(pm)
275+
mctle.passes.add_reject_dot_op(pm)
276276
mctle.passes.add_early_assign_memory_space(pm)
277277
mctle.passes.add_select_encodings(pm)
278278
mctle.passes.add_insert_local_pointer_barriers(pm)

third_party/metax/include/triton/Dialect/Triton/IR/Dialect.h

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,12 @@ struct GlobalMemory : public SideEffects::Resource::Base<GlobalMemory> {
2828
StringRef getName() final { return "<GlobalMemory>"; }
2929
};
3030

31+
#ifdef __MCTLE__
32+
struct SharedMemory : public SideEffects::Resource::Base<SharedMemory> {
33+
StringRef getName() final { return "<SharedMemory>"; }
34+
};
35+
#endif
36+
3137
class DialectInferLayoutInterface
3238
: public DialectInterface::Base<DialectInferLayoutInterface> {
3339
public:

third_party/metax/include/triton/Dialect/Triton/IR/TritonOps.td

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,9 @@ include "triton/Dialect/Triton/IR/TritonOpInterfaces.td"
2020
// Interfaces
2121
//
2222
def GlobalMemory : Resource<"::mlir::triton::GlobalMemory">;
23+
#ifdef __MCTLE__
24+
def SharedMemory : Resource<"::mlir::triton::SharedMemory">;
25+
#endif
2326

2427
//
2528
// Op Base
@@ -351,8 +354,13 @@ def TT_StoreOp : TT_Op<"store", [
351354
def TT_AtomicRMWOp : TT_Op<"atomic_rmw", [
352355
SameOperandsAndResultShape,
353356
SameOperandsAndResultEncoding,
357+
#ifdef __MCTLE__
358+
TypesMatchWith<"value type matches ptr type", "ptr", "val",
359+
"getPointeeType($_self)">,
360+
#else
354361
TypesMatchWith<"ptr type matches value type", "val", "ptr",
355362
"getPointerTypeSameShape($_self)">,
363+
#endif
356364
TypesMatchWith<"mask type matches value type",
357365
"val", "mask", "getI1SameShape($_self)",
358366
"($_op.getOperands().size() <= 2) || std::equal_to<>()">
@@ -367,7 +375,12 @@ def TT_AtomicRMWOp : TT_Op<"atomic_rmw", [
367375

368376
let arguments = (ins
369377
TT_AtomicRMWAttr:$atomic_rmw_op,
378+
#ifdef __MCTLE__
379+
Arg<TT_PtrLike, "", [MemRead<GlobalMemory>, MemWrite<GlobalMemory>,
380+
MemRead<SharedMemory>, MemWrite<SharedMemory>]>:$ptr,
381+
#else
370382
Arg<TT_PtrLike, "", [MemRead<GlobalMemory>, MemWrite<GlobalMemory>]>:$ptr,
383+
#endif
371384
TT_Type:$val,
372385
Optional<TT_BoolLike>:$mask,
373386
TT_MemSemanticAttr:$sem,
@@ -387,10 +400,17 @@ def TT_AtomicRMWOp : TT_Op<"atomic_rmw", [
387400
def TT_AtomicCASOp : TT_Op<"atomic_cas", [
388401
SameOperandsAndResultShape,
389402
SameOperandsAndResultEncoding,
403+
#ifdef __MCTLE__
404+
TypesMatchWith<"cmp type matches ptr type", "ptr", "cmp",
405+
"getPointeeType($_self)">,
406+
TypesMatchWith<"value type matches ptr type", "ptr", "val",
407+
"getPointeeType($_self)">
408+
#else
390409
TypesMatchWith<"ptr type matches cmp type", "cmp", "ptr",
391410
"getPointerTypeSameShape($_self)">,
392411
TypesMatchWith<"ptr type matches value type", "val", "ptr",
393412
"getPointerTypeSameShape($_self)">
413+
#endif
394414
]> {
395415
let summary = "atomic cas";
396416

@@ -405,7 +425,12 @@ def TT_AtomicCASOp : TT_Op<"atomic_cas", [
405425
}];
406426

407427
let arguments = (ins
428+
#ifdef __MCTLE__
429+
Arg<TT_PtrLike, "", [MemRead<GlobalMemory>, MemWrite<GlobalMemory>,
430+
MemRead<SharedMemory>, MemWrite<SharedMemory>]>:$ptr,
431+
#else
408432
Arg<TT_PtrLike, "", [MemRead<GlobalMemory>, MemWrite<GlobalMemory>]>:$ptr,
433+
#endif // __MCTLE__
409434
TT_Type:$cmp,
410435
TT_Type:$val,
411436
TT_MemSemanticAttr:$sem,

third_party/metax/plugin/lib/TritonMETAXGPUToLLVM/LoadStoreOpToLLVM.cpp

Lines changed: 114 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,10 +7,13 @@
77
#include "Utility.h"
88
#include "triton/Dialect/TritonGPU/Transforms/Utility.h"
99
#include "triton/Tools/LayoutUtils.h"
10+
#include <optional>
1011

1112
using namespace mlir;
1213
using namespace mlir::triton;
1314

15+
namespace ttg = mlir::triton::gpu;
16+
1417
using ::mlir::LLVM::delinearize;
1518
using ::mlir::LLVM::getSharedMemoryBase;
1619
using ::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) {

third_party/metax/plugin/mctle/dialect/include/Transforms/Passes.td

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -122,4 +122,10 @@ def TritonTleLowerInsertTile
122122
"mlir::triton::gpu::TritonGPUDialect"];
123123
}
124124

125+
def TritonTleRejectDotOp
126+
: Pass<"triton-tle-reject-dot-op", "mlir::ModuleOp"> {
127+
let summary = "reject tt.dot in TLE kernels";
128+
let dependentDialects = ["mlir::triton::TritonDialect"];
129+
}
130+
125131
#endif // TRITON_TLE_PASSES

third_party/metax/plugin/mctle/triton_mctle.cc

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -131,6 +131,7 @@ void init_triton_mctle_ir(py::module &&m) {
131131
}
132132

133133
void init_triton_mctle_passes(py::module &&m) {
134+
ADD_PASS_WRAPPER_0("add_reject_dot_op", tle::createTritonTleRejectDotOp);
134135
ADD_PASS_WRAPPER_0("add_early_assign_memory_space",
135136
tle::createTritonTleEarlyAssignMemorySpace);
136137
ADD_PASS_WRAPPER_0("add_select_encodings",

third_party/metax/spec/triton/language/semantic.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1078,7 +1078,8 @@ def _load_legacy(self, ptr, mask, other, boundary_check, padding, cache, evictio
10781078
return ret
10791079

10801080
def load(self, ptr: TensorTy, mask: Optional[TensorTy], other: Optional[TensorTy], boundary_check: Tuple,
1081-
padding_option: str, cache_modifier: str, eviction_policy: str, is_volatile: bool) -> TensorTy:
1081+
padding_option: str, cache_modifier: str, eviction_policy: str, is_volatile: bool,
1082+
smem_hints: str = None) -> TensorTy:
10821083
# Cache, eviction and padding options
10831084
cache = self._str_to_load_cache_modifier(cache_modifier)
10841085
eviction = self._str_to_eviction_policy(eviction_policy)

0 commit comments

Comments
 (0)