[Mlir-commits] [mlir] [MLIR][XeGPU] Lower arith.truncf wider than one xevm.truncf group (PR #217130)
llvmlistbot at llvm.org
llvmlistbot at llvm.org
Tue Aug 18 13:52:40 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir
Author: Sang Ik Lee (silee2)
<details>
<summary>Changes</summary>
xevm.truncf lowers to device builtins that convert a fixed number of elements at a time: the fp4 path feeds __builtin_IB_dnscl_bf16 with a vector<8xi32>, and the fp8 paths call the _16 flavours of the bf16/half conversion builtins. All of them take exactly 16 f16/bf16 elements, so TruncfToXeVMPattern only matched an arith.truncf on a 16-element source and left anything wider alone.
Wider sources do occur. In a workgroup-level mxfp GEMM that quantizes A in place, A's inst_data and lane layout give each lane 32 bf16 per instruction, and arith.scaling_truncf expands to an arith.truncf from vector<32xbf16> to vector<32xf4E2M1FN>. Nothing else lowers it, so it survived to LLVM translation and the module failed to build with "cannot be converted to LLVM IR: missing `LLVMTranslationDialectInterface` registration ... for op: arith.truncf".
Convert a whole number of groups at a time instead: slice the source into 16-element groups, emit one xevm.truncf per group, and concatenate the packed i8 results before the final bitcast to the converted result type. Sources that are not a whole number of groups are still not matched and are left to the regular arith-to-LLVM path.
A source of exactly one group takes the same path as before and its generated code is unchanged, so this only adds the slicing when a wider vector is actually converted.
arith.extf has the same restriction on its result, but no wider case has come up yet, so it is left as is.
---
Full diff: https://github.com/llvm/llvm-project/pull/217130.diff
2 Files Affected:
- (modified) mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp (+41-14)
- (modified) mlir/test/Conversion/XeGPUToXeVM/extf_truncf.mlir (+38)
``````````diff
diff --git a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
index 78d99cf88b768..98bf354cce1e3 100644
--- a/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
+++ b/mlir/lib/Conversion/XeGPUToXeVM/XeGPUToXeVM.cpp
@@ -1257,14 +1257,16 @@ static bool isXeVMExtf(arith::ExtFOp op) {
}
// Returns true if `op` is an arith.truncf that can be lowered to xevm.truncf,
-// i.e. a rank-1 truncation from a 16-element f16/bf16 vector to an MX narrow
-// float.
+// i.e. a rank-1 truncation from an f16/bf16 vector to an MX narrow float. The
+// source has to hold a whole number of the fixed-size groups xevm.truncf
+// converts at a time; wider vectors are converted in several steps.
static bool isXeVMTruncf(arith::TruncFOp op) {
auto srcTy = dyn_cast<VectorType>(op.getIn().getType());
auto dstTy = dyn_cast<VectorType>(op.getType());
if (!srcTy || !dstTy || srcTy.getRank() != 1 || dstTy.getRank() != 1)
return false;
- if (srcTy.getNumElements() != kXeVMExtfTruncfNumElems)
+ int64_t numElems = srcTy.getNumElements();
+ if (numElems == 0 || numElems % kXeVMExtfTruncfNumElems != 0)
return false;
Type srcETy = srcTy.getElementType();
if (!srcETy.isF16() && !srcETy.isBF16())
@@ -1324,19 +1326,44 @@ class TruncfToXeVMPattern : public OpConversionPattern<arith::TruncFOp> {
: xevm::TruncfSrcElemTypes::BF16;
xevm::TruncfDstElemTypes dstEnum =
*getTruncfNarrowType(dstVecTy.getElementType());
- // xevm.truncf produces the narrow floats packed into an i8 vector.
- int64_t numNarrowBits =
- dstVecTy.getNumElements() * dstVecTy.getElementTypeBitWidth();
- Type packedTy = VectorType::get(numNarrowBits / 8, rewriter.getI8Type());
- Value res =
- xevm::TruncfOp::create(rewriter, loc, packedTy, adaptor.getIn(),
- xevm::TruncfSrcElemTypeAttr::get(ctx, srcEnum),
- xevm::TruncfDstElemTypeAttr::get(ctx, dstEnum));
+ auto srcEnumAttr = xevm::TruncfSrcElemTypeAttr::get(ctx, srcEnum);
+ auto dstEnumAttr = xevm::TruncfDstElemTypeAttr::get(ctx, dstEnum);
+
+ // xevm.truncf lowers to instructions that convert a fixed number of
+ // elements at a time, so a wider source is converted one group at a time
+ // and the packed results are concatenated. Each group produces the narrow
+ // floats packed into an i8 vector.
+ int64_t numGroups = srcVecTy.getNumElements() / kXeVMExtfTruncfNumElems;
+ int64_t groupBytes =
+ kXeVMExtfTruncfNumElems * dstVecTy.getElementTypeBitWidth() / 8;
+ Type groupTy = VectorType::get(groupBytes, rewriter.getI8Type());
+
+ Value src = adaptor.getIn();
+ Value packed;
+ if (numGroups == 1) {
+ packed = xevm::TruncfOp::create(rewriter, loc, groupTy, src, srcEnumAttr,
+ dstEnumAttr);
+ } else {
+ auto packedTy =
+ VectorType::get(groupBytes * numGroups, rewriter.getI8Type());
+ packed = arith::ConstantOp::create(rewriter, loc, packedTy,
+ rewriter.getZeroAttr(packedTy));
+ for (int64_t group = 0; group < numGroups; group++) {
+ Value slice = vector::ExtractStridedSliceOp::create(
+ rewriter, loc, src, group * kXeVMExtfTruncfNumElems,
+ kXeVMExtfTruncfNumElems, /*strides=*/1);
+ Value converted = xevm::TruncfOp::create(rewriter, loc, groupTy, slice,
+ srcEnumAttr, dstEnumAttr);
+ packed = vector::InsertStridedSliceOp::create(
+ rewriter, loc, converted, packed, group * groupBytes,
+ /*strides=*/1);
+ }
+ }
// Re-shape to the type-converted result type (i4 vector for fp4).
Type resTy = getTypeConverter()->convertType(dstVecTy);
- if (res.getType() != resTy)
- res = vector::BitCastOp::create(rewriter, loc, resTy, res);
- rewriter.replaceOp(op, res);
+ if (packed.getType() != resTy)
+ packed = vector::BitCastOp::create(rewriter, loc, resTy, packed);
+ rewriter.replaceOp(op, packed);
return success();
}
};
diff --git a/mlir/test/Conversion/XeGPUToXeVM/extf_truncf.mlir b/mlir/test/Conversion/XeGPUToXeVM/extf_truncf.mlir
index b74c92f26b2f6..51db50c782845 100644
--- a/mlir/test/Conversion/XeGPUToXeVM/extf_truncf.mlir
+++ b/mlir/test/Conversion/XeGPUToXeVM/extf_truncf.mlir
@@ -82,6 +82,44 @@ gpu.module @truncf_bf16_f8 [#xevm.target<chip = "cri">] {
// -----
+// A source wider than one conversion group is converted a group at a time and
+// the packed results are concatenated. 32 bf16 -> 32 fp4 is two groups of 16,
+// each packed into vector<8xi8>.
+
+// CHECK-LABEL: gpu.func @truncf_bf16_e2m1_wide
+// CHECK-SAME: (%[[ARG0:.*]]: vector<32xbf16>)
+gpu.module @truncf_bf16_e2m1_wide [#xevm.target<chip = "cri">] {
+ gpu.func @truncf_bf16_e2m1_wide(%a: vector<32xbf16>) kernel {
+ // CHECK: %[[ZERO:.*]] = arith.constant dense<0> : vector<16xi8>
+ // CHECK: %[[S0:.*]] = vector.extract_strided_slice %[[ARG0]] {offsets = [0], sizes = [16], strides = [1]} : vector<32xbf16> to vector<16xbf16>
+ // CHECK: %[[T0:.*]] = xevm.truncf %[[S0]] {src_etype = bf16, dst_etype = e2m1} : (vector<16xbf16>) -> vector<8xi8>
+ // CHECK: %[[P0:.*]] = vector.insert_strided_slice %[[T0]], %[[ZERO]] {offsets = [0], strides = [1]} : vector<8xi8> into vector<16xi8>
+ // CHECK: %[[S1:.*]] = vector.extract_strided_slice %[[ARG0]] {offsets = [16], sizes = [16], strides = [1]} : vector<32xbf16> to vector<16xbf16>
+ // CHECK: %[[T1:.*]] = xevm.truncf %[[S1]] {src_etype = bf16, dst_etype = e2m1} : (vector<16xbf16>) -> vector<8xi8>
+ // CHECK: %[[P1:.*]] = vector.insert_strided_slice %[[T1]], %[[P0]] {offsets = [8], strides = [1]} : vector<8xi8> into vector<16xi8>
+ // CHECK: %{{.*}} = vector.bitcast %[[P1]] : vector<16xi8> to vector<32xi4>
+ %r = arith.truncf %a : vector<32xbf16> to vector<32xf4E2M1FN>
+ gpu.return
+ }
+}
+
+// -----
+
+// A source that is not a whole number of conversion groups has no xevm.truncf
+// lowering and is left for the regular arith-to-LLVM path.
+
+// CHECK-LABEL: gpu.func @truncf_partial_group
+gpu.module @truncf_partial_group [#xevm.target<chip = "cri">] {
+ gpu.func @truncf_partial_group(%a: vector<24xbf16>) kernel {
+ // CHECK: %{{.*}} = arith.truncf %{{.*}} : vector<24xbf16> to vector<24xf4E2M1FN>
+ // CHECK-NOT: xevm.truncf
+ %r = arith.truncf %a : vector<24xbf16> to vector<24xf4E2M1FN>
+ gpu.return
+ }
+}
+
+// -----
+
// Plain float extensions/truncations are not micro-scaling and must be left
// untouched for the regular arith-to-LLVM lowering.
``````````
</details>
https://github.com/llvm/llvm-project/pull/217130
More information about the Mlir-commits
mailing list