From fd95f98ccc5bc4ca314ccd769e58749efc4650a1 Mon Sep 17 00:00:00 2001 From: cc-c122 <2513326504@qq.com> Date: Tue, 18 Aug 2026 18:35:38 +0800 Subject: [PATCH] [Metax][MCTLE] Add NVMMAShared encoding builder binding --- .../metax/plugin/mctle/triton_mctle.cc | 31 +++++++++++++++++++ 1 file changed, 31 insertions(+) diff --git a/third_party/metax/plugin/mctle/triton_mctle.cc b/third_party/metax/plugin/mctle/triton_mctle.cc index 18406b88f4..0f1306af5c 100644 --- a/third_party/metax/plugin/mctle/triton_mctle.cc +++ b/third_party/metax/plugin/mctle/triton_mctle.cc @@ -95,6 +95,37 @@ void init_triton_mctle_ir(py::module &&m) { return mlir::cast(ttg::SwizzledSharedEncodingAttr::get( context, vectorSize, perPhase, maxPhase, order, CTALayout)); }) + .def("make_nv_mma_shared_encoding_attr", + [](TritonOpBuilder &self, std::vector shape, + std::vector order, Type &elemType, + std::vector CTAsPerCGA, + std::vector CTASplitNum, + std::vector CTAOrder, bool fp4Padded, + bool swizzled) { + assert(shape.size() == order.size()); + assert(order.size() == CTAsPerCGA.size()); + assert(CTAsPerCGA.size() == CTASplitNum.size()); + assert(CTASplitNum.size() == CTAOrder.size()); + + auto context = self.getBuilder().getContext(); + auto CTALayout = ttg::CTAEncodingAttr::fromSplitParams( + context, CTAsPerCGA, CTASplitNum, CTAOrder); + + if (swizzled) { + return mlir::cast( + ttg::NVMMASharedEncodingAttr::get( + context, shape, order, CTALayout, elemType, + fp4Padded)); + } + + return mlir::cast( + ttg::NVMMASharedEncodingAttr::get( + context, + /*swizzlingByteWidth=*/0, + /*transposed=*/order[0] == 0, + elemType.getIntOrFloatBitWidth(), + fp4Padded, CTALayout)); + }) .def("create_local_alloc", [](TritonOpBuilder &self, std::vector shape, Type &elementType, Attribute &encoding) -> mlir::Value {