--- a/lib/Dialect/TritonGPU/Transforms/AccelerateMatmul.cpp +++ b/lib/Dialect/TritonGPU/Transforms/AccelerateMatmul.cpp @@ -877,6 +877,11 @@ bool isFp4MMA = isAFP4 && isBFP4; bool requiresFp4Padding = nvidia_gpu::TargetFeatures(computeCapability).requiresFp4Padding(); + // mxf4/mxf4nvf4 do not support MN-major operands, so FP4 x FP4 uses + // mxf8f6f4 if either operand is not K-packed. The mxf8f6f4 shared-memory + // packing format requires padding for every FP4 operand. + bool isFp4MMAUsingMxf8f6f4 = + isFp4MMA && (!dotOp.getLhsKPack() || !dotOp.getRhsKPack()); // On Blackwell, if we use mixed-precision MMA we need to pad the fp4 // operand @@ -888,11 +893,11 @@ auto blockK = dotOp.getA().getType().getShape().back() * (isAFP4 ? 2 : 1); auto blockM = dotOp.getA().getType().getShape()[0]; bool isMMAv5Fp4PaddedLhs = - (isFp4MMA && !dotOp.getLhsKPack()) || // fp4 x fp4 with M-major A + isFp4MMAUsingMxf8f6f4 || // fp4 x fp4 using mxf8f6f4 (IsAMixedPrecFp4 && (requiresFp4Padding || blockM == 64 || blockK == 32 || !dotOp.getLhsKPack())); bool isMMAv5Fp4PaddedRhs = - (isFp4MMA && !dotOp.getRhsKPack()) || // fp4 x fp4 with N-major B + isFp4MMAUsingMxf8f6f4 || // fp4 x fp4 using mxf8f6f4 (IsBMixedPrecFp4 && (requiresFp4Padding || blockK == 32 || !dotOp.getRhsKPack())); --- a/test/Conversion/tritongpu_to_llvm_blackwell_256bit.mlir +++ b/test/Conversion/tritongpu_to_llvm_blackwell_256bit.mlir @@ -5,7 +5,7 @@ #blocked_8xf32 = #ttg.blocked<{sizePerThread = [8], threadsPerWarp = [32], warpsPerCTA = [1], order = [0]}> module attributes {"ttg.num-ctas" = 1 : i32, "ttg.num-warps" = 1 : i32} { // BW256-LABEL: global_load_v8_b32 - // BW256: ld.global.v8.b32 + // BW256: ld.global.v4.b32 // PRE_BW-LABEL: global_load_v8_b32 // PRE_BW-NOT: ld.global.v8.b32 // PRE_BW: ld.global.v4.b32 @@ -29,7 +29,7 @@ #blocked_8xf32 = #ttg.blocked<{sizePerThread = [8], threadsPerWarp = [32], warpsPerCTA = [1], order = [0]}> module attributes {"ttg.num-ctas" = 1 : i32, "ttg.num-warps" = 1 : i32} { // BW256-LABEL: global_store_v8_b32 - // BW256: st.global.v8.b32 + // BW256: st.global.v4.b32 // PRE_BW-LABEL: global_store_v8_b32 // PRE_BW-NOT: st.global.v8.b32 // PRE_BW: st.global.v4.b32 @@ -54,7 +54,7 @@ #blocked_4xf64 = #ttg.blocked<{sizePerThread = [4], threadsPerWarp = [32], warpsPerCTA = [1], order = [0]}> module attributes {"ttg.num-ctas" = 1 : i32, "ttg.num-warps" = 1 : i32} { // BW256-LABEL: global_load_v4_b64 - // BW256: ld.global.v4.b64 + // BW256: ld.global.v2.b64 // PRE_BW-LABEL: global_load_v4_b64 // PRE_BW-NOT: ld.global.v4.b64 // PRE_BW: ld.global.v2.b64 --- a/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/LoadStoreOpToLLVM.cpp +++ b/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/LoadStoreOpToLLVM.cpp @@ -95,12 +95,9 @@ auto pointeeBitWidth = triton::getPointeeBitWidth(tensorTy); LDBG("getVectorSize contiguity = " << contiguity << " pointeeBitWidth = " << pointeeBitWidth); - // Blackwell (sm_100+) with PTX 8.8+ supports 256-bit global load/store. + // Blackwell (sm_100+) with PTX 8.8+ supports 256-bit global load/store, + // but we cap it at 128 bits to avoid compilation failures with ptxas. unsigned maxVecBits = 128; - if (targetInfo.getComputeCapability() >= 100 && - targetInfo.getPtxVersion() >= 88) { - maxVecBits = 256; - } return std::min(maxVecBits / pointeeBitWidth, contiguity); }