Backport of https://github.com/triton-lang/triton/pull/11245. --- a/lib/Dialect/TritonGPU/Transforms/AccelerateMatmul.cpp +++ b/lib/Dialect/TritonGPU/Transforms/AccelerateMatmul.cpp @@ -877,11 +877,6 @@ public: 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 @@ -893,11 +888,33 @@ public: // MMA_K = 64 is also not supported when BLOCK_M = 64. auto blockK = dotOp.getA().getType().getShape().back() * (isAFP4 ? 2 : 1); auto blockM = dotOp.getA().getType().getShape()[0]; + bool hasMNMajorFp4Operand = + isFp4MMA && (!dotOp.getLhsKPack() || !dotOp.getRhsKPack()); + auto aScaleType = dotOp.getAScale().getType(); + auto bScaleType = dotOp.getBScale().getType(); + auto aScaleElemType = aScaleType.getElementType(); + auto bScaleElemType = bScaleType.getElementType(); + auto isBlock16Scale = [blockK](RankedTensorType scaleType) { + return scaleType.getShape().back() * 16 == blockK; + }; + bool hasUE4M3Scale = isa(aScaleElemType) || + isa(bScaleElemType); + bool hasUE5M3Scale = + (aScaleElemType.isInteger(8) && isBlock16Scale(aScaleType)) || + (bScaleElemType.isInteger(8) && isBlock16Scale(bScaleType)); + // mxf4nvf4 supports UE4M3/UE5M3 scales but not MN-major operands, while + // the mxf8f6f4 fallback for MN-major FP4 only supports E8M0 scales. + if (hasMNMajorFp4Operand && (hasUE4M3Scale || hasUE5M3Scale)) + return failure(); + // mxf4 does not support MN-major operands, so mxfp4 x mxfp4 falls back to + // mxf8f6f4 if either operand is not K-packed. The mxf8f6f4 shared-memory + // packing format requires padding for every fp4 operand, even if the + // operand is K packed. bool isMMAv5Fp4PaddedLhs = - isFp4MMAUsingMxf8f6f4 || // fp4 x fp4 using mxf8f6f4 + hasMNMajorFp4Operand || (IsAMixedPrecFp4 && (requiresFp4Padding || blockM == 64 || blockK == 32 || !dotOp.getLhsKPack())); bool isMMAv5Fp4PaddedRhs = - isFp4MMAUsingMxf8f6f4 || // fp4 x fp4 using mxf8f6f4 + hasMNMajorFp4Operand || (IsBMixedPrecFp4 && (requiresFp4Padding || blockK == 32 || !dotOp.getRhsKPack()));