[Mlir-commits] [mlir] [MLIR][XeVM] Add xevm.extf op as the inverse of xevm.truncf (PR #203124)
Sang Ik Lee
llvmlistbot at llvm.org
Wed Jun 10 16:47:41 PDT 2026
https://github.com/silee2 updated https://github.com/llvm/llvm-project/pull/203124
>From 656f1ef9d475fad4b24979c2d6ee4676b0540c9e Mon Sep 17 00:00:00 2001
From: "Lee, Sang Ik" <sang.ik.lee at intel.com>
Date: Wed, 10 Jun 2026 19:01:31 +0000
Subject: [PATCH 1/2] [MLIR][XeVM] Add xevm.extf op as the inverse of
xevm.truncf
Add a new xevm.extf operation that extends f8/bf8/f4 values to f16/bf16, mirroring the existing xevm.truncf op, together with its lowering in XeVMToLLVM.
Lowering details (XeVMToLLVM):
- bf8/f8 -> f16 via __builtin_IB_bf8tohf_16 / __builtin_IB_hf8tohf_16.
- bf8/f8 -> bf16 via f16 -> f32 (convert_float16) -> bf16 (__builtin_IB_ftobf_16).
- e2m1 (fp4) -> f16/bf16 via __builtin_IB_shfl_idx4_lut and __builtin_IB_shfl_idx4_to_fp16_8_packed (LUT 7 for f16, 5 for bf16).
Adds the op definition and verifier, conversion/roundtrip/invalid unit tests, and f8 and fp4 GPU round-trip integration tests.
---
mlir/include/mlir/Dialect/LLVMIR/XeVMOps.td | 45 +++++++
mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp | 127 +++++++++++++++++-
mlir/lib/Dialect/LLVMIR/IR/XeVMDialect.cpp | 20 +++
.../XeVMToLLVM/xevm_mx-to-llvm.mlir | 114 ++++++++++++++++
mlir/test/Dialect/LLVMIR/invalid.mlir | 8 ++
mlir/test/Dialect/LLVMIR/xevm.mlir | 20 +++
.../XeVM/GPU/xevm_truncf_extf_roundtrip.mlir | 106 +++++++++++++++
.../GPU/xevm_truncf_extf_roundtrip_bf16.mlir | 107 +++++++++++++++
.../GPU/xevm_truncf_extf_roundtrip_fp4.mlir | 118 ++++++++++++++++
.../xevm_truncf_extf_roundtrip_fp4_bf16.mlir | 119 ++++++++++++++++
10 files changed, 782 insertions(+), 2 deletions(-)
create mode 100644 mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip.mlir
create mode 100644 mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_bf16.mlir
create mode 100644 mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_fp4.mlir
create mode 100644 mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_fp4_bf16.mlir
diff --git a/mlir/include/mlir/Dialect/LLVMIR/XeVMOps.td b/mlir/include/mlir/Dialect/LLVMIR/XeVMOps.td
index d2dc51198be32..8c4e409e9c395 100644
--- a/mlir/include/mlir/Dialect/LLVMIR/XeVMOps.td
+++ b/mlir/include/mlir/Dialect/LLVMIR/XeVMOps.td
@@ -648,6 +648,51 @@ def XeVM_TruncfOp
let hasVerifier = 1;
}
+def XeVM_ExtfSrcElemTypes
+ : I32EnumAttr<"ExtfSrcElemTypes",
+ "Source element type for xevm.extf",
+ [XeVM_ET_BF8, XeVM_ET_F8, XeVM_ET_E2M1]> {
+ let cppNamespace = "::mlir::xevm";
+}
+
+def XeVM_ExtfSrcElemTypeAttr : XeVM_Attr<"ExtfSrcElemType", "ext_src_etype"> {
+ let parameters = (ins "xevm::ExtfSrcElemTypes":$etype);
+ let assemblyFormat = "`src_etype` `=` $etype";
+}
+
+def XeVM_ExtfDstElemTypes
+ : I32EnumAttr<
+ "ExtfDstElemTypes",
+ "Destination element type for xevm.extf", [XeVM_ET_F16, XeVM_ET_BF16]> {
+ let cppNamespace = "::mlir::xevm";
+}
+
+def XeVM_ExtfDstElemTypeAttr : XeVM_Attr<"ExtfDstElemType", "ext_dst_etype"> {
+ let parameters = (ins "xevm::ExtfDstElemTypes":$etype);
+ let assemblyFormat = "`dst_etype` `=` $etype";
+}
+
+def XeVM_ExtfOp
+ : XeVM_Op<"extf">,
+ Results<(outs AnyTypeOf<[FixedVectorOfRankAndType<[1], [F16, BF16]>, F16,
+ BF16]>:$dst)>,
+ Arguments<(ins AnyTypeOf<[FixedVectorOfRankAndType<[1], [I8, I<4>]>, I8,
+ I<4>]>:$src,
+ XeVM_ExtfSrcElemTypeAttr:$src_etype,
+ XeVM_ExtfDstElemTypeAttr:$dst_etype)> {
+ let summary = "Floating point extension from f8/bf8/f4 to f16/bf16";
+ let description = [{
+ The `xevm.extf` operation extends a floating point value from
+ f8/bf8/f4 format to f16/bf16 format. It is the inverse of `xevm.truncf`.
+ }];
+
+ let assemblyFormat = [{
+ $src ` ` `{` $src_etype `,` $dst_etype `}` attr-dict `:` functional-type(operands, results)
+ }];
+
+ let hasVerifier = 1;
+}
+
def XeVM_MMAMxOp
: XeVM_Op<"mma_mx">,
Results<(outs FixedVectorOfRankAndType<[1], [XeVM_MatrixElemType]>:$d)>,
diff --git a/mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp b/mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp
index cd638cb610b0f..c833771c69286 100644
--- a/mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp
+++ b/mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp
@@ -1307,6 +1307,129 @@ class TruncfToOCLPattern : public OpConversionPattern<TruncfOp> {
}
};
+class ExtfToOCLPattern : public OpConversionPattern<ExtfOp> {
+ using OpConversionPattern::OpConversionPattern;
+ LogicalResult
+ matchAndRewrite(ExtfOp op, ExtfOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ // `xevm.extf` is the inverse of `xevm.truncf`. Supported source and result
+ // types are restricted for now, mirroring the truncf lowering.
+ auto srcEtype = op.getSrcEtype().getEtype();
+ auto dstEtype = op.getDstEtype().getEtype();
+ // Scalar case is not supported until usage case become clear.
+ auto vecSrcTy = dyn_cast<VectorType>(op.getSrc().getType());
+ if (!vecSrcTy)
+ return rewriter.notifyMatchFailure(op, "Scalar src is not supported.");
+ auto vecDstTy = dyn_cast<VectorType>(op.getDst().getType());
+ if (!vecDstTy)
+ return rewriter.notifyMatchFailure(op, "Scalar dst is not supported.");
+ Value src = op.getSrc();
+ auto memAttr = rewriter.getAttr<LLVM::MemoryEffectsAttr>(
+ /*other=*/LLVM::ModRefInfo::NoModRef,
+ /*argMem=*/LLVM::ModRefInfo::NoModRef,
+ /*inaccessibleMem=*/LLVM::ModRefInfo::NoModRef,
+ /*errnoMem=*/LLVM::ModRefInfo::NoModRef,
+ /*targetMem0=*/LLVM::ModRefInfo::NoModRef,
+ /*targetMem1=*/LLVM::ModRefInfo::NoModRef);
+ auto funcAttrs = convergentNoUnwindWillReturnAttrs;
+ funcAttrs.memEffectsAttr = memAttr;
+
+ // Handle the case where src type is fp4 (e2m1) first.
+ if (srcEtype == ExtfSrcElemTypes::E2M1) {
+ // 16 fp4 values are packed into vector<8xi8>, the result is a
+ // vector<16xf16> or vector<16xbf16>.
+ // Use:
+ // uint16 __builtin_IB_shfl_idx4_lut(int lut_index)
+ // uint8 __builtin_IB_shfl_idx4_to_fp16_8_packed(uint16 lut,
+ // char8 source)
+ // The lookup table selects the target format:
+ // 7 = e2m1 -> f16, 5 = e2m1 -> bf16.
+ if (vecSrcTy.getNumElements() != 8 || vecDstTy.getNumElements() != 16)
+ return rewriter.notifyMatchFailure(
+ op, "fp4 src expects a vector<8xi8> src and a 16 element dst");
+ constexpr int kLutE2M1ToF16 = 7;
+ constexpr int kLutE2M1ToBF16 = 5;
+ int lutIndex =
+ (dstEtype == ExtfDstElemTypes::F16) ? kLutE2M1ToF16 : kLutE2M1ToBF16;
+ Value lutIdx = LLVM::ConstantOp::create(
+ rewriter, op.getLoc(), rewriter.getI32Type(), lutIndex);
+ Type lutTy = VectorType::get(16, rewriter.getI32Type());
+ Value lut = createDeviceFunctionCall(
+ rewriter, "__builtin_IB_shfl_idx4_lut", lutTy,
+ {lutIdx.getType()}, {lutIdx}, {}, funcAttrs,
+ op.getOperation())
+ ->getResult(0);
+ Type packedResTy = VectorType::get(8, rewriter.getI32Type());
+ SmallVector<Type> convArgTypes{lut.getType(), src.getType()};
+ SmallVector<Value> convArgs{lut, src};
+ Value result =
+ createDeviceFunctionCall(
+ rewriter, "__builtin_IB_shfl_idx4_to_fp16_8_packed", packedResTy,
+ convArgTypes, convArgs, {}, funcAttrs, op.getOperation())
+ ->getResult(0);
+ // The builtin returns the f16/bf16 bits packed as i32, bitcast to the
+ // f16/bf16 dst type.
+ result = LLVM::BitcastOp::create(rewriter, op.getLoc(), vecDstTy, result);
+ rewriter.replaceOp(op, result);
+ return success();
+ }
+
+ // Handle the case where src type is fp8 (bf8/hf8).
+ // Only 16 input elements are supported, see TruncfToOCLPattern for details.
+ if (vecSrcTy.getNumElements() != 16)
+ return rewriter.notifyMatchFailure(
+ op, "Only vector src of 16 elements is supported");
+
+ // Step 1: Extend fp8 (bf8/hf8) to F16.
+ // bf8 -> half: half16 __builtin_IB_bf8tohf_16(char16)
+ // hf8 -> half: half16 __builtin_IB_hf8tohf_16(char16)
+ std::string fnName = (srcEtype == ExtfSrcElemTypes::BF8)
+ ? "__builtin_IB_bf8tohf_16"
+ : "__builtin_IB_hf8tohf_16";
+ Type f16Ty = VectorType::get(vecSrcTy.getShape(), rewriter.getF16Type());
+ SmallVector<Type> argTypes{src.getType()};
+ SmallVector<Value> args{src};
+ Value result =
+ createDeviceFunctionCall(rewriter, fnName, f16Ty, argTypes, args, {},
+ funcAttrs, op.getOperation())
+ ->getResult(0);
+
+ // When the destination is F16, we are done.
+ if (dstEtype == ExtfDstElemTypes::F16) {
+ rewriter.replaceOp(op, result);
+ return success();
+ }
+
+ // BF16 destination needs some postprocessing.
+ // First extend F16 to F32 and then truncate to BF16.
+ // Step 2: Extend to F32.
+ // Use float16 convert_float16(half16)
+ std::string convFnName = "convert_float16";
+ SmallVector<Type> convArgTypes{result.getType()};
+ SmallVector<Value> convArgs{result};
+ convFnName = mangle(convFnName, convArgTypes);
+ Type f32Ty = VectorType::get(vecSrcTy.getShape(), rewriter.getF32Type());
+ result =
+ createDeviceFunctionCall(rewriter, convFnName, f32Ty, convArgTypes,
+ convArgs, {}, funcAttrs, op.getOperation())
+ ->getResult(0);
+ // Step 3: Truncate F32 to BF16.
+ // Use short16 __builtin_IB_ftobf_16(float16)
+ std::string ftobfFnName = "__builtin_IB_ftobf_16";
+ SmallVector<Type> ftobfArgTypes{result.getType()};
+ SmallVector<Value> ftobfArgs{result};
+ Type i16Ty = VectorType::get(vecSrcTy.getShape(), rewriter.getI16Type());
+ result =
+ createDeviceFunctionCall(rewriter, ftobfFnName, i16Ty, ftobfArgTypes,
+ ftobfArgs, {}, funcAttrs, op.getOperation())
+ ->getResult(0);
+ // The builtin returns the bf16 bits as i16, bitcast to the bf16 dst type.
+ result = LLVM::BitcastOp::create(rewriter, op.getLoc(), vecDstTy, result);
+ rewriter.replaceOp(op, result);
+ return success();
+ }
+};
+
class MMAMxToOCLPattern : public OpConversionPattern<MMAMxOp> {
using OpConversionPattern::OpConversionPattern;
LogicalResult
@@ -1687,6 +1810,6 @@ void ::mlir::populateXeVMToLLVMConversionPatterns(ConversionTarget &target,
SubgroupOpWorkitemOpToOCLPattern<LaneIdOp>,
SubgroupOpWorkitemOpToOCLPattern<SubgroupIdOp>,
SubgroupOpWorkitemOpToOCLPattern<SubgroupSizeOp>,
- TruncfToOCLPattern, MMAMxToOCLPattern, AllocaToGlobalPattern>(
- patterns.getContext());
+ TruncfToOCLPattern, ExtfToOCLPattern, MMAMxToOCLPattern,
+ AllocaToGlobalPattern>(patterns.getContext());
}
diff --git a/mlir/lib/Dialect/LLVMIR/IR/XeVMDialect.cpp b/mlir/lib/Dialect/LLVMIR/IR/XeVMDialect.cpp
index e14c8253baa3f..d127655d40c04 100644
--- a/mlir/lib/Dialect/LLVMIR/IR/XeVMDialect.cpp
+++ b/mlir/lib/Dialect/LLVMIR/IR/XeVMDialect.cpp
@@ -377,6 +377,26 @@ LogicalResult TruncfOp::verify() {
return success();
}
+LogicalResult ExtfOp::verify() {
+ Type srcTy = getSrc().getType();
+ Type dstTy = getDst().getType();
+ if (isa<VectorType>(srcTy) && !isa<VectorType>(dstTy))
+ return emitOpError("both src and dst should be vector types or both should "
+ "be scalar types");
+ if (isa<VectorType>(srcTy)) {
+ VectorType srcVecTy = dyn_cast<VectorType>(srcTy);
+ VectorType dstVecTy = dyn_cast<VectorType>(dstTy);
+ if (srcVecTy.getElementTypeBitWidth() >= dstVecTy.getElementTypeBitWidth())
+ return emitError(
+ "dst element bitwidth should be greater than src element bitwidth");
+ } else {
+ if (srcTy.getIntOrFloatBitWidth() >= dstTy.getIntOrFloatBitWidth())
+ return emitError(
+ "dst element bitwidth should be greater than src element bitwidth");
+ }
+ return success();
+}
+
LogicalResult
XeVMTargetAttr::verify(function_ref<InFlightDiagnostic()> emitError, int O,
StringRef triple, StringRef chip, DictionaryAttr flags,
diff --git a/mlir/test/Conversion/XeVMToLLVM/xevm_mx-to-llvm.mlir b/mlir/test/Conversion/XeVMToLLVM/xevm_mx-to-llvm.mlir
index 9706887182003..2ac76ac1e73a0 100644
--- a/mlir/test/Conversion/XeVMToLLVM/xevm_mx-to-llvm.mlir
+++ b/mlir/test/Conversion/XeVMToLLVM/xevm_mx-to-llvm.mlir
@@ -162,6 +162,120 @@ llvm.func @truncf_bf16_to_e2m1(%src: vector<16xbf16>) -> vector<8xi8> {
// -----
+// CHECK: llvm.func spir_funccc @__builtin_IB_bf8tohf_16(vector<16xi8>) -> vector<16xf16>
+// CHECK-SAME: attributes {convergent, memory_effects = #llvm.memory_effects<other = none,
+// CHECK-SAME: argMem = none, inaccessibleMem = none, errnoMem = none,
+// CHECK-SAME: targetMem0 = none, targetMem1 = none>, no_unwind, will_return}
+// CHECK-LABEL: llvm.func @extf_bf8_to_f16
+// CHECK-SAME: %[[ARG0:.*]]: vector<16xi8>
+llvm.func @extf_bf8_to_f16(%src: vector<16xi8>) -> vector<16xf16> {
+ // CHECK: %[[VAR0:.*]] = llvm.call spir_funccc @__builtin_IB_bf8tohf_16(%[[ARG0]])
+ // CHECK-SAME: {convergent, function_type = !llvm.func<vector<16xf16> (vector<16xi8>)>,
+ // CHECK-SAME: linkage = #llvm.linkage<external>, memory_effects = #llvm.memory_effects<other = none,
+ // CHECK-SAME: argMem = none, inaccessibleMem = none, errnoMem = none,
+ // CHECK-SAME: targetMem0 = none, targetMem1 = none>,
+ // CHECK-SAME: no_unwind, sym_name = "__builtin_IB_bf8tohf_16",
+ // CHECK-SAME: visibility_ = 0 : i64, will_return} :
+ // CHECK-SAME: (vector<16xi8>) -> vector<16xf16>
+ %dst = xevm.extf %src { src_etype = bf8, dst_etype = f16 } : (vector<16xi8>) -> vector<16xf16>
+ llvm.return %dst : vector<16xf16>
+}
+
+// -----
+
+// CHECK: llvm.func spir_funccc @__builtin_IB_hf8tohf_16(vector<16xi8>) -> vector<16xf16>
+// CHECK-SAME: attributes {convergent, memory_effects = #llvm.memory_effects<other = none,
+// CHECK-SAME: argMem = none, inaccessibleMem = none, errnoMem = none,
+// CHECK-SAME: targetMem0 = none, targetMem1 = none>, no_unwind, will_return}
+// CHECK-LABEL: llvm.func @extf_f8_to_f16
+// CHECK-SAME: %[[ARG0:.*]]: vector<16xi8>
+llvm.func @extf_f8_to_f16(%src: vector<16xi8>) -> vector<16xf16> {
+ // CHECK: %[[VAR0:.*]] = llvm.call spir_funccc @__builtin_IB_hf8tohf_16(%[[ARG0]])
+ // CHECK-SAME: {convergent, function_type = !llvm.func<vector<16xf16> (vector<16xi8>)>,
+ // CHECK-SAME: linkage = #llvm.linkage<external>, memory_effects = #llvm.memory_effects<other = none,
+ // CHECK-SAME: argMem = none, inaccessibleMem = none, errnoMem = none,
+ // CHECK-SAME: targetMem0 = none, targetMem1 = none>,
+ // CHECK-SAME: no_unwind, sym_name = "__builtin_IB_hf8tohf_16",
+ // CHECK-SAME: visibility_ = 0 : i64, will_return} :
+ // CHECK-SAME: (vector<16xi8>) -> vector<16xf16>
+ %dst = xevm.extf %src { src_etype = f8, dst_etype = f16 } : (vector<16xi8>) -> vector<16xf16>
+ llvm.return %dst : vector<16xf16>
+}
+
+// -----
+
+// CHECK: llvm.func spir_funccc @__builtin_IB_ftobf_16(vector<16xf32>) -> vector<16xi16>
+// CHECK: llvm.func spir_funccc @_Z15convert_float16Dv16_Dh(vector<16xf16>) -> vector<16xf32>
+// CHECK: llvm.func spir_funccc @__builtin_IB_bf8tohf_16(vector<16xi8>) -> vector<16xf16>
+// CHECK-LABEL: llvm.func @extf_bf8_to_bf16
+// CHECK-SAME: %[[ARG0:.*]]: vector<16xi8>
+llvm.func @extf_bf8_to_bf16(%src: vector<16xi8>) -> vector<16xbf16> {
+ // CHECK: %[[VAR0:.*]] = llvm.call spir_funccc @__builtin_IB_bf8tohf_16(%[[ARG0]])
+ // CHECK-SAME: : (vector<16xi8>) -> vector<16xf16>
+ // CHECK: %[[VAR1:.*]] = llvm.call spir_funccc @_Z15convert_float16Dv16_Dh(%[[VAR0]])
+ // CHECK-SAME: : (vector<16xf16>) -> vector<16xf32>
+ // CHECK: %[[VAR2:.*]] = llvm.call spir_funccc @__builtin_IB_ftobf_16(%[[VAR1]])
+ // CHECK-SAME: : (vector<16xf32>) -> vector<16xi16>
+ // CHECK: %[[VAR3:.*]] = llvm.bitcast %[[VAR2]] : vector<16xi16> to vector<16xbf16>
+ %dst = xevm.extf %src { src_etype = bf8, dst_etype = bf16 } : (vector<16xi8>) -> vector<16xbf16>
+ llvm.return %dst : vector<16xbf16>
+}
+
+// -----
+
+// CHECK: llvm.func spir_funccc @__builtin_IB_ftobf_16(vector<16xf32>) -> vector<16xi16>
+// CHECK: llvm.func spir_funccc @_Z15convert_float16Dv16_Dh(vector<16xf16>) -> vector<16xf32>
+// CHECK: llvm.func spir_funccc @__builtin_IB_hf8tohf_16(vector<16xi8>) -> vector<16xf16>
+// CHECK-LABEL: llvm.func @extf_f8_to_bf16
+// CHECK-SAME: %[[ARG0:.*]]: vector<16xi8>
+llvm.func @extf_f8_to_bf16(%src: vector<16xi8>) -> vector<16xbf16> {
+ // CHECK: %[[VAR0:.*]] = llvm.call spir_funccc @__builtin_IB_hf8tohf_16(%[[ARG0]])
+ // CHECK-SAME: : (vector<16xi8>) -> vector<16xf16>
+ // CHECK: %[[VAR1:.*]] = llvm.call spir_funccc @_Z15convert_float16Dv16_Dh(%[[VAR0]])
+ // CHECK-SAME: : (vector<16xf16>) -> vector<16xf32>
+ // CHECK: %[[VAR2:.*]] = llvm.call spir_funccc @__builtin_IB_ftobf_16(%[[VAR1]])
+ // CHECK-SAME: : (vector<16xf32>) -> vector<16xi16>
+ // CHECK: %[[VAR3:.*]] = llvm.bitcast %[[VAR2]] : vector<16xi16> to vector<16xbf16>
+ %dst = xevm.extf %src { src_etype = f8, dst_etype = bf16 } : (vector<16xi8>) -> vector<16xbf16>
+ llvm.return %dst : vector<16xbf16>
+}
+
+// -----
+
+// CHECK: llvm.func spir_funccc @__builtin_IB_shfl_idx4_to_fp16_8_packed(vector<16xi32>, vector<8xi8>) -> vector<8xi32>
+// CHECK: llvm.func spir_funccc @__builtin_IB_shfl_idx4_lut(i32) -> vector<16xi32>
+// CHECK-LABEL: llvm.func @extf_e2m1_to_f16
+// CHECK-SAME: %[[ARG0:.*]]: vector<8xi8>
+llvm.func @extf_e2m1_to_f16(%src: vector<8xi8>) -> vector<16xf16> {
+ // CHECK: %[[LUTIDX:.*]] = llvm.mlir.constant(7 : i32) : i32
+ // CHECK: %[[LUT:.*]] = llvm.call spir_funccc @__builtin_IB_shfl_idx4_lut(%[[LUTIDX]])
+ // CHECK-SAME: : (i32) -> vector<16xi32>
+ // CHECK: %[[CONV:.*]] = llvm.call spir_funccc @__builtin_IB_shfl_idx4_to_fp16_8_packed(%[[LUT]], %[[ARG0]])
+ // CHECK-SAME: : (vector<16xi32>, vector<8xi8>) -> vector<8xi32>
+ // CHECK: %[[RES:.*]] = llvm.bitcast %[[CONV]] : vector<8xi32> to vector<16xf16>
+ %dst = xevm.extf %src { src_etype = e2m1, dst_etype = f16 } : (vector<8xi8>) -> vector<16xf16>
+ llvm.return %dst : vector<16xf16>
+}
+
+// -----
+
+// CHECK: llvm.func spir_funccc @__builtin_IB_shfl_idx4_to_fp16_8_packed(vector<16xi32>, vector<8xi8>) -> vector<8xi32>
+// CHECK: llvm.func spir_funccc @__builtin_IB_shfl_idx4_lut(i32) -> vector<16xi32>
+// CHECK-LABEL: llvm.func @extf_e2m1_to_bf16
+// CHECK-SAME: %[[ARG0:.*]]: vector<8xi8>
+llvm.func @extf_e2m1_to_bf16(%src: vector<8xi8>) -> vector<16xbf16> {
+ // CHECK: %[[LUTIDX:.*]] = llvm.mlir.constant(5 : i32) : i32
+ // CHECK: %[[LUT:.*]] = llvm.call spir_funccc @__builtin_IB_shfl_idx4_lut(%[[LUTIDX]])
+ // CHECK-SAME: : (i32) -> vector<16xi32>
+ // CHECK: %[[CONV:.*]] = llvm.call spir_funccc @__builtin_IB_shfl_idx4_to_fp16_8_packed(%[[LUT]], %[[ARG0]])
+ // CHECK-SAME: : (vector<16xi32>, vector<8xi8>) -> vector<8xi32>
+ // CHECK: %[[RES:.*]] = llvm.bitcast %[[CONV]] : vector<8xi32> to vector<16xbf16>
+ %dst = xevm.extf %src { src_etype = e2m1, dst_etype = bf16 } : (vector<8xi8>) -> vector<16xbf16>
+ llvm.return %dst : vector<16xbf16>
+}
+
+// -----
+
// CHECK: llvm.func spir_funccc @__builtin_IB_sub_group16_bdpas_f_f_bf8_bf8_8_8
// CHECK-SAME: (vector<8xf32>, vector<8xi16>, vector<8xi32>, i8, i8) -> vector<8xf32>
// CHECK-SAME: attributes {convergent, memory_effects = #llvm.memory_effects<other = none,
diff --git a/mlir/test/Dialect/LLVMIR/invalid.mlir b/mlir/test/Dialect/LLVMIR/invalid.mlir
index e80094df1eed2..3aa37c57b756c 100644
--- a/mlir/test/Dialect/LLVMIR/invalid.mlir
+++ b/mlir/test/Dialect/LLVMIR/invalid.mlir
@@ -2037,6 +2037,14 @@ llvm.func @invalid_xevm_truncf_1(%arg0: vector<8xf16>) {
// -----
+llvm.func @invalid_xevm_extf_1(%arg0: vector<8xi8>) {
+ // expected-error at +1 {{op both src and dst should be vector types or both}}
+ %0 = xevm.extf %arg0 { src_etype = bf8, dst_etype = f16 } : (vector<8xi8>) -> f16
+ llvm.return
+}
+
+// -----
+
llvm.func @invalid_xevm_mma_mx(%loaded_c_casted: vector<4xf32>, %loaded_a: vector<8xi16>, %loaded_b_casted: vector<8xi32>, %scale_a: vector<2xi8>, %scale_b: vector<2xi8>) -> vector<8xf32> {
// expected-error at +1 {{op type of C operand must match result type}}
%c_result = xevm.mma_mx %loaded_a, %loaded_b_casted, %scale_a, %scale_b, %loaded_c_casted { shape=<m=8, n=16, k=64>,
diff --git a/mlir/test/Dialect/LLVMIR/xevm.mlir b/mlir/test/Dialect/LLVMIR/xevm.mlir
index abd315a2dd8c3..fbce5355611d4 100644
--- a/mlir/test/Dialect/LLVMIR/xevm.mlir
+++ b/mlir/test/Dialect/LLVMIR/xevm.mlir
@@ -124,6 +124,26 @@ func.func @truncf_vector() -> vector<8xi4> {
return %2 : vector<8xi4>
}
+// -----
+// CHECK-LABEL: func.func @extf_scalar
+func.func @extf_scalar() -> f16 {
+ // CHECK: %[[VAR0:.*]] = arith.constant
+ %0 = arith.constant 1 : i8
+ // CHECK: xevm.extf %[[VAR0]] {src_etype = bf8, dst_etype = f16} : (i8) -> f16
+ %2 = xevm.extf %0 { src_etype=bf8, dst_etype=f16 } : (i8) -> f16
+ return %2 : f16
+}
+
+// -----
+// CHECK-LABEL: func.func @extf_vector
+func.func @extf_vector() -> vector<8xbf16> {
+ // CHECK: %[[VAR0:.*]] = arith.constant
+ %0 = arith.constant dense<1> : vector<8xi4>
+ // CHECK: xevm.extf %[[VAR0]] {src_etype = e2m1, dst_etype = bf16} : (vector<8xi4>) -> vector<8xbf16>
+ %2 = xevm.extf %0 { src_etype=e2m1, dst_etype=bf16 } : (vector<8xi4>) -> vector<8xbf16>
+ return %2 : vector<8xbf16>
+}
+
// -----
// CHECK-LABEL: func.func @memfence()
func.func @memfence() {
diff --git a/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip.mlir b/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip.mlir
new file mode 100644
index 0000000000000..a21878af51f15
--- /dev/null
+++ b/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip.mlir
@@ -0,0 +1,106 @@
+// RUN: mlir-opt %s --gpu-lower-to-xevm-pipeline="xegpu-op-level=lane zebin-chip=cri" \
+// RUN: | mlir-runner \
+// RUN: --shared-libs=%mlir_levelzero_runtime \
+// RUN: --shared-libs=%mlir_runner_utils \
+// RUN: --shared-libs=%mlir_c_runner_utils \
+// RUN: --entry-point-result=void \
+// RUN: | FileCheck %s
+
+// Round trip test for xevm.truncf followed by xevm.extf.
+// Each of the 16 lanes owns a vector<16xf16>, truncates it to f8 (E4M3FN) with
+// xevm.truncf and extends it back to f16 with xevm.extf. The integers 1..16 are
+// exactly representable in f8E4M3FN, so the round trip must reproduce the input.
+module @roundtrip attributes {gpu.container_module} {
+
+ gpu.module @kernel {
+ gpu.func @truncf_extf_roundtrip(%ptr: !llvm.ptr<1>) kernel {
+ // Each lane processes 16 contiguous f16 values: lane L owns [L*16, L*16+16).
+ %lane = gpu.lane_id
+ %lane_i64 = arith.index_cast %lane : index to i64
+ %c16 = arith.constant 16 : i64
+ %offset = arith.muli %lane_i64, %c16 : i64
+ %lane_ptr = llvm.getelementptr %ptr[%offset]
+ : (!llvm.ptr<1>, i64) -> !llvm.ptr<1>, f16
+ %vec = llvm.load %lane_ptr : !llvm.ptr<1> -> vector<16xf16>
+ // f16 -> f8 (E4M3FN) -> f16 round trip.
+ %trunc = xevm.truncf %vec { src_etype = f16, dst_etype = f8 }
+ : (vector<16xf16>) -> vector<16xi8>
+ %ext = xevm.extf %trunc { src_etype = f8, dst_etype = f16 }
+ : (vector<16xi8>) -> vector<16xf16>
+ llvm.store %ext, %lane_ptr : vector<16xf16>, !llvm.ptr<1>
+ gpu.return
+ }
+ }
+
+ func.func @test(%src : memref<16x16xf16>) -> memref<16x16xf16> attributes {llvm.emit_c_interface} {
+ %c1 = arith.constant 1 : index
+ %c16 = arith.constant 16 : index
+ %memref_0 = gpu.alloc() : memref<16x16xf16>
+ gpu.memcpy %memref_0, %src : memref<16x16xf16>, memref<16x16xf16>
+ %0 = memref.extract_aligned_pointer_as_index %memref_0 : memref<16x16xf16> -> index
+ %1 = arith.index_cast %0 : index to i64
+ %2 = llvm.inttoptr %1 : i64 to !llvm.ptr
+ %src_casted = llvm.addrspacecast %2 : !llvm.ptr to !llvm.ptr<1>
+ gpu.launch_func @kernel::@truncf_extf_roundtrip blocks in (%c1, %c1, %c1) threads in (%c16, %c1, %c1)
+ args(%src_casted : !llvm.ptr<1>)
+ %dst = memref.alloc() : memref<16x16xf16>
+ gpu.memcpy %dst, %memref_0 : memref<16x16xf16>, memref<16x16xf16>
+ gpu.dealloc %memref_0 : memref<16x16xf16>
+ return %dst : memref<16x16xf16>
+ }
+
+ func.func @main() attributes {llvm.emit_c_interface} {
+ %c0 = arith.constant 0 : index
+ %c1 = arith.constant 1 : index
+ %c16 = arith.constant 16 : index
+ %c1_i32 = arith.constant 1 : i32
+
+ // Fill row L with the value (L + 1), exactly representable in f8E4M3FN.
+ %A = memref.alloc() : memref<16x16xf16>
+ scf.for %i = %c0 to %c16 step %c1 {
+ %i_i32 = arith.index_cast %i : index to i32
+ %v_i32 = arith.addi %i_i32, %c1_i32 : i32
+ %v = arith.sitofp %v_i32 : i32 to f16
+ scf.for %j = %c0 to %c16 step %c1 {
+ memref.store %v, %A[%i, %j] : memref<16x16xf16>
+ }
+ }
+
+ %B = call @test(%A) : (memref<16x16xf16>) -> memref<16x16xf16>
+
+ // Convert the f16 result to f32 so it can be printed with printMemrefF32.
+ %Bf32 = memref.alloc() : memref<16x16xf32>
+ scf.for %i = %c0 to %c16 step %c1 {
+ scf.for %j = %c0 to %c16 step %c1 {
+ %v = memref.load %B[%i, %j] : memref<16x16xf16>
+ %vf = arith.extf %v : f16 to f32
+ memref.store %vf, %Bf32[%i, %j] : memref<16x16xf32>
+ }
+ }
+ %B_cast = memref.cast %Bf32 : memref<16x16xf32> to memref<*xf32>
+ call @printMemrefF32(%B_cast) : (memref<*xf32>) -> ()
+
+ // CHECK: Unranked Memref base@ = 0x{{[0-9a-f]+}}
+ // CHECK: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]
+ // CHECK: [2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2]
+ // CHECK: [3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3]
+ // CHECK: [4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4]
+ // CHECK: [5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5]
+ // CHECK: [6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6]
+ // CHECK: [7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7]
+ // CHECK: [8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8]
+ // CHECK: [9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9]
+ // CHECK: [10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10]
+ // CHECK: [11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11]
+ // CHECK: [12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12]
+ // CHECK: [13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13]
+ // CHECK: [14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14]
+ // CHECK: [15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15]
+ // CHECK: [16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16]
+ memref.dealloc %A : memref<16x16xf16>
+ memref.dealloc %B : memref<16x16xf16>
+ memref.dealloc %Bf32 : memref<16x16xf32>
+ return
+ }
+ func.func private @printMemrefF32(%ptr : memref<*xf32>) attributes { llvm.emit_c_interface }
+}
diff --git a/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_bf16.mlir b/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_bf16.mlir
new file mode 100644
index 0000000000000..4174d8afa65e7
--- /dev/null
+++ b/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_bf16.mlir
@@ -0,0 +1,107 @@
+// RUN: mlir-opt %s --gpu-lower-to-xevm-pipeline="xegpu-op-level=lane zebin-chip=cri" \
+// RUN: | mlir-runner \
+// RUN: --shared-libs=%mlir_levelzero_runtime \
+// RUN: --shared-libs=%mlir_runner_utils \
+// RUN: --shared-libs=%mlir_c_runner_utils \
+// RUN: --entry-point-result=void \
+// RUN: | FileCheck %s
+
+// Round trip test for xevm.truncf followed by xevm.extf.
+// Each of the 16 lanes owns a vector<16xbf16>, truncates it to f8 (E4M3FN) with
+// xevm.truncf and extends it back to bf16 with xevm.extf. The integers 1..16 are
+// exactly representable in both f8E4M3FN and bf16, so the round trip must
+// reproduce the input.
+module @roundtrip attributes {gpu.container_module} {
+
+ gpu.module @kernel {
+ gpu.func @truncf_extf_roundtrip(%ptr: !llvm.ptr<1>) kernel {
+ // Each lane processes 16 contiguous bf16 values: lane L owns [L*16, L*16+16).
+ %lane = gpu.lane_id
+ %lane_i64 = arith.index_cast %lane : index to i64
+ %c16 = arith.constant 16 : i64
+ %offset = arith.muli %lane_i64, %c16 : i64
+ %lane_ptr = llvm.getelementptr %ptr[%offset]
+ : (!llvm.ptr<1>, i64) -> !llvm.ptr<1>, bf16
+ %vec = llvm.load %lane_ptr : !llvm.ptr<1> -> vector<16xbf16>
+ // bf16 -> f8 (E4M3FN) -> bf16 round trip.
+ %trunc = xevm.truncf %vec { src_etype = bf16, dst_etype = f8 }
+ : (vector<16xbf16>) -> vector<16xi8>
+ %ext = xevm.extf %trunc { src_etype = f8, dst_etype = bf16 }
+ : (vector<16xi8>) -> vector<16xbf16>
+ llvm.store %ext, %lane_ptr : vector<16xbf16>, !llvm.ptr<1>
+ gpu.return
+ }
+ }
+
+ func.func @test(%src : memref<16x16xbf16>) -> memref<16x16xbf16> attributes {llvm.emit_c_interface} {
+ %c1 = arith.constant 1 : index
+ %c16 = arith.constant 16 : index
+ %memref_0 = gpu.alloc() : memref<16x16xbf16>
+ gpu.memcpy %memref_0, %src : memref<16x16xbf16>, memref<16x16xbf16>
+ %0 = memref.extract_aligned_pointer_as_index %memref_0 : memref<16x16xbf16> -> index
+ %1 = arith.index_cast %0 : index to i64
+ %2 = llvm.inttoptr %1 : i64 to !llvm.ptr
+ %src_casted = llvm.addrspacecast %2 : !llvm.ptr to !llvm.ptr<1>
+ gpu.launch_func @kernel::@truncf_extf_roundtrip blocks in (%c1, %c1, %c1) threads in (%c16, %c1, %c1)
+ args(%src_casted : !llvm.ptr<1>)
+ %dst = memref.alloc() : memref<16x16xbf16>
+ gpu.memcpy %dst, %memref_0 : memref<16x16xbf16>, memref<16x16xbf16>
+ gpu.dealloc %memref_0 : memref<16x16xbf16>
+ return %dst : memref<16x16xbf16>
+ }
+
+ func.func @main() attributes {llvm.emit_c_interface} {
+ %c0 = arith.constant 0 : index
+ %c1 = arith.constant 1 : index
+ %c16 = arith.constant 16 : index
+ %c1_i32 = arith.constant 1 : i32
+
+ // Fill row L with the value (L + 1), exactly representable in f8E4M3FN.
+ %A = memref.alloc() : memref<16x16xbf16>
+ scf.for %i = %c0 to %c16 step %c1 {
+ %i_i32 = arith.index_cast %i : index to i32
+ %v_i32 = arith.addi %i_i32, %c1_i32 : i32
+ %v = arith.sitofp %v_i32 : i32 to bf16
+ scf.for %j = %c0 to %c16 step %c1 {
+ memref.store %v, %A[%i, %j] : memref<16x16xbf16>
+ }
+ }
+
+ %B = call @test(%A) : (memref<16x16xbf16>) -> memref<16x16xbf16>
+
+ // Convert the bf16 result to f32 so it can be printed with printMemrefF32.
+ %Bf32 = memref.alloc() : memref<16x16xf32>
+ scf.for %i = %c0 to %c16 step %c1 {
+ scf.for %j = %c0 to %c16 step %c1 {
+ %v = memref.load %B[%i, %j] : memref<16x16xbf16>
+ %vf = arith.extf %v : bf16 to f32
+ memref.store %vf, %Bf32[%i, %j] : memref<16x16xf32>
+ }
+ }
+ %B_cast = memref.cast %Bf32 : memref<16x16xf32> to memref<*xf32>
+ call @printMemrefF32(%B_cast) : (memref<*xf32>) -> ()
+
+ // CHECK: Unranked Memref base@ = 0x{{[0-9a-f]+}}
+ // CHECK: [1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1]
+ // CHECK: [2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2]
+ // CHECK: [3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3]
+ // CHECK: [4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4]
+ // CHECK: [5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5]
+ // CHECK: [6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6]
+ // CHECK: [7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7]
+ // CHECK: [8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8, 8]
+ // CHECK: [9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9, 9]
+ // CHECK: [10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10, 10]
+ // CHECK: [11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11, 11]
+ // CHECK: [12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12, 12]
+ // CHECK: [13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13, 13]
+ // CHECK: [14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14, 14]
+ // CHECK: [15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15, 15]
+ // CHECK: [16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16, 16]
+ memref.dealloc %A : memref<16x16xbf16>
+ memref.dealloc %B : memref<16x16xbf16>
+ memref.dealloc %Bf32 : memref<16x16xf32>
+ return
+ }
+ func.func private @printMemrefF32(%ptr : memref<*xf32>) attributes { llvm.emit_c_interface }
+}
diff --git a/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_fp4.mlir b/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_fp4.mlir
new file mode 100644
index 0000000000000..53ad29d31adca
--- /dev/null
+++ b/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_fp4.mlir
@@ -0,0 +1,118 @@
+// RUN: mlir-opt %s --gpu-lower-to-xevm-pipeline="xegpu-op-level=lane zebin-chip=cri" \
+// RUN: | mlir-runner \
+// RUN: --shared-libs=%mlir_levelzero_runtime \
+// RUN: --shared-libs=%mlir_runner_utils \
+// RUN: --shared-libs=%mlir_c_runner_utils \
+// RUN: --entry-point-result=void \
+// RUN: | FileCheck %s
+
+// Round trip test for xevm.truncf followed by xevm.extf with the fp4 (e2m1)
+// format. Each of the 16 lanes owns a vector<16xf16>, truncates it to e2m1 (16
+// fp4 values packed into vector<8xi8>) with xevm.truncf and extends it back to
+// f16 with xevm.extf. Every value used is exactly representable in e2m1
+// (0, 0.5, 1, 1.5, 2, 3, 4, 6), so the round trip must reproduce the input.
+module @roundtrip attributes {gpu.container_module} {
+
+ gpu.module @kernel {
+ gpu.func @truncf_extf_roundtrip_fp4(%ptr: !llvm.ptr<1>) kernel {
+ // Each lane processes 16 contiguous f16 values: lane L owns [L*16, L*16+16).
+ %lane = gpu.lane_id
+ %lane_i64 = arith.index_cast %lane : index to i64
+ %c16 = arith.constant 16 : i64
+ %offset = arith.muli %lane_i64, %c16 : i64
+ %lane_ptr = llvm.getelementptr %ptr[%offset]
+ : (!llvm.ptr<1>, i64) -> !llvm.ptr<1>, f16
+ %vec = llvm.load %lane_ptr : !llvm.ptr<1> -> vector<16xf16>
+ // f16 -> e2m1 (fp4, 16 values packed in vector<8xi8>) -> f16 round trip.
+ %trunc = xevm.truncf %vec { src_etype = f16, dst_etype = e2m1 }
+ : (vector<16xf16>) -> vector<8xi8>
+ %ext = xevm.extf %trunc { src_etype = e2m1, dst_etype = f16 }
+ : (vector<8xi8>) -> vector<16xf16>
+ llvm.store %ext, %lane_ptr : vector<16xf16>, !llvm.ptr<1>
+ gpu.return
+ }
+ }
+
+ func.func @test(%src : memref<16x16xf16>) -> memref<16x16xf16> attributes {llvm.emit_c_interface} {
+ %c1 = arith.constant 1 : index
+ %c16 = arith.constant 16 : index
+ %memref_0 = gpu.alloc() : memref<16x16xf16>
+ gpu.memcpy %memref_0, %src : memref<16x16xf16>, memref<16x16xf16>
+ %0 = memref.extract_aligned_pointer_as_index %memref_0 : memref<16x16xf16> -> index
+ %1 = arith.index_cast %0 : index to i64
+ %2 = llvm.inttoptr %1 : i64 to !llvm.ptr
+ %src_casted = llvm.addrspacecast %2 : !llvm.ptr to !llvm.ptr<1>
+ gpu.launch_func @kernel::@truncf_extf_roundtrip_fp4 blocks in (%c1, %c1, %c1) threads in (%c16, %c1, %c1)
+ args(%src_casted : !llvm.ptr<1>)
+ %dst = memref.alloc() : memref<16x16xf16>
+ gpu.memcpy %dst, %memref_0 : memref<16x16xf16>, memref<16x16xf16>
+ gpu.dealloc %memref_0 : memref<16x16xf16>
+ return %dst : memref<16x16xf16>
+ }
+
+ func.func @main() attributes {llvm.emit_c_interface} {
+ %c0 = arith.constant 0 : index
+ %c1 = arith.constant 1 : index
+ %c2 = arith.constant 2 : index
+ %c3 = arith.constant 3 : index
+ %c4 = arith.constant 4 : index
+ %c5 = arith.constant 5 : index
+ %c6 = arith.constant 6 : index
+ %c7 = arith.constant 7 : index
+ %c8 = arith.constant 8 : index
+ %c16 = arith.constant 16 : index
+
+ // Lookup table of the 8 magnitudes exactly representable in e2m1.
+ %lut = memref.alloc() : memref<8xf16>
+ %v0 = arith.constant 0.0 : f16
+ %v1 = arith.constant 0.5 : f16
+ %v2 = arith.constant 1.0 : f16
+ %v3 = arith.constant 1.5 : f16
+ %v4 = arith.constant 2.0 : f16
+ %v5 = arith.constant 3.0 : f16
+ %v6 = arith.constant 4.0 : f16
+ %v7 = arith.constant 6.0 : f16
+ memref.store %v0, %lut[%c0] : memref<8xf16>
+ memref.store %v1, %lut[%c1] : memref<8xf16>
+ memref.store %v2, %lut[%c2] : memref<8xf16>
+ memref.store %v3, %lut[%c3] : memref<8xf16>
+ memref.store %v4, %lut[%c4] : memref<8xf16>
+ memref.store %v5, %lut[%c5] : memref<8xf16>
+ memref.store %v6, %lut[%c6] : memref<8xf16>
+ memref.store %v7, %lut[%c7] : memref<8xf16>
+
+ // Fill every row with the repeating pattern of representable values so each
+ // lane exercises the full set of e2m1 values.
+ %A = memref.alloc() : memref<16x16xf16>
+ scf.for %i = %c0 to %c16 step %c1 {
+ scf.for %j = %c0 to %c16 step %c1 {
+ %jm8 = arith.remui %j, %c8 : index
+ %val = memref.load %lut[%jm8] : memref<8xf16>
+ memref.store %val, %A[%i, %j] : memref<16x16xf16>
+ }
+ }
+
+ %B = call @test(%A) : (memref<16x16xf16>) -> memref<16x16xf16>
+
+ // Convert the f16 result to f32 so it can be printed with printMemrefF32.
+ %Bf32 = memref.alloc() : memref<16x16xf32>
+ scf.for %i = %c0 to %c16 step %c1 {
+ scf.for %j = %c0 to %c16 step %c1 {
+ %v = memref.load %B[%i, %j] : memref<16x16xf16>
+ %vf = arith.extf %v : f16 to f32
+ memref.store %vf, %Bf32[%i, %j] : memref<16x16xf32>
+ }
+ }
+ %B_cast = memref.cast %Bf32 : memref<16x16xf32> to memref<*xf32>
+ call @printMemrefF32(%B_cast) : (memref<*xf32>) -> ()
+
+ // CHECK: Unranked Memref base@ = 0x{{[0-9a-f]+}}
+ // CHECK-COUNT-16: [0, 0.5, 1, 1.5, 2, 3, 4, 6, 0, 0.5, 1, 1.5, 2, 3, 4, 6]
+ memref.dealloc %A : memref<16x16xf16>
+ memref.dealloc %B : memref<16x16xf16>
+ memref.dealloc %Bf32 : memref<16x16xf32>
+ memref.dealloc %lut : memref<8xf16>
+ return
+ }
+ func.func private @printMemrefF32(%ptr : memref<*xf32>) attributes { llvm.emit_c_interface }
+}
diff --git a/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_fp4_bf16.mlir b/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_fp4_bf16.mlir
new file mode 100644
index 0000000000000..d913205a42d4c
--- /dev/null
+++ b/mlir/test/Integration/Dialect/XeVM/GPU/xevm_truncf_extf_roundtrip_fp4_bf16.mlir
@@ -0,0 +1,119 @@
+// RUN: mlir-opt %s --gpu-lower-to-xevm-pipeline="xegpu-op-level=lane zebin-chip=cri" \
+// RUN: | mlir-runner \
+// RUN: --shared-libs=%mlir_levelzero_runtime \
+// RUN: --shared-libs=%mlir_runner_utils \
+// RUN: --shared-libs=%mlir_c_runner_utils \
+// RUN: --entry-point-result=void \
+// RUN: | FileCheck %s
+
+// Round trip test for xevm.truncf followed by xevm.extf with the fp4 (e2m1)
+// format. Each of the 16 lanes owns a vector<16xbf16>, truncates it to e2m1 (16
+// fp4 values packed into vector<8xi8>) with xevm.truncf and extends it back to
+// bf16 with xevm.extf. Every value used is exactly representable in both e2m1
+// and bf16 (0, 0.5, 1, 1.5, 2, 3, 4, 6), so the round trip must reproduce the
+// input.
+module @roundtrip attributes {gpu.container_module} {
+
+ gpu.module @kernel {
+ gpu.func @truncf_extf_roundtrip_fp4(%ptr: !llvm.ptr<1>) kernel {
+ // Each lane processes 16 contiguous bf16 values: lane L owns [L*16, L*16+16).
+ %lane = gpu.lane_id
+ %lane_i64 = arith.index_cast %lane : index to i64
+ %c16 = arith.constant 16 : i64
+ %offset = arith.muli %lane_i64, %c16 : i64
+ %lane_ptr = llvm.getelementptr %ptr[%offset]
+ : (!llvm.ptr<1>, i64) -> !llvm.ptr<1>, bf16
+ %vec = llvm.load %lane_ptr : !llvm.ptr<1> -> vector<16xbf16>
+ // bf16 -> e2m1 (fp4, 16 values packed in vector<8xi8>) -> bf16 round trip.
+ %trunc = xevm.truncf %vec { src_etype = bf16, dst_etype = e2m1 }
+ : (vector<16xbf16>) -> vector<8xi8>
+ %ext = xevm.extf %trunc { src_etype = e2m1, dst_etype = bf16 }
+ : (vector<8xi8>) -> vector<16xbf16>
+ llvm.store %ext, %lane_ptr : vector<16xbf16>, !llvm.ptr<1>
+ gpu.return
+ }
+ }
+
+ func.func @test(%src : memref<16x16xbf16>) -> memref<16x16xbf16> attributes {llvm.emit_c_interface} {
+ %c1 = arith.constant 1 : index
+ %c16 = arith.constant 16 : index
+ %memref_0 = gpu.alloc() : memref<16x16xbf16>
+ gpu.memcpy %memref_0, %src : memref<16x16xbf16>, memref<16x16xbf16>
+ %0 = memref.extract_aligned_pointer_as_index %memref_0 : memref<16x16xbf16> -> index
+ %1 = arith.index_cast %0 : index to i64
+ %2 = llvm.inttoptr %1 : i64 to !llvm.ptr
+ %src_casted = llvm.addrspacecast %2 : !llvm.ptr to !llvm.ptr<1>
+ gpu.launch_func @kernel::@truncf_extf_roundtrip_fp4 blocks in (%c1, %c1, %c1) threads in (%c16, %c1, %c1)
+ args(%src_casted : !llvm.ptr<1>)
+ %dst = memref.alloc() : memref<16x16xbf16>
+ gpu.memcpy %dst, %memref_0 : memref<16x16xbf16>, memref<16x16xbf16>
+ gpu.dealloc %memref_0 : memref<16x16xbf16>
+ return %dst : memref<16x16xbf16>
+ }
+
+ func.func @main() attributes {llvm.emit_c_interface} {
+ %c0 = arith.constant 0 : index
+ %c1 = arith.constant 1 : index
+ %c2 = arith.constant 2 : index
+ %c3 = arith.constant 3 : index
+ %c4 = arith.constant 4 : index
+ %c5 = arith.constant 5 : index
+ %c6 = arith.constant 6 : index
+ %c7 = arith.constant 7 : index
+ %c8 = arith.constant 8 : index
+ %c16 = arith.constant 16 : index
+
+ // Lookup table of the 8 magnitudes exactly representable in e2m1.
+ %lut = memref.alloc() : memref<8xbf16>
+ %v0 = arith.constant 0.0 : bf16
+ %v1 = arith.constant 0.5 : bf16
+ %v2 = arith.constant 1.0 : bf16
+ %v3 = arith.constant 1.5 : bf16
+ %v4 = arith.constant 2.0 : bf16
+ %v5 = arith.constant 3.0 : bf16
+ %v6 = arith.constant 4.0 : bf16
+ %v7 = arith.constant 6.0 : bf16
+ memref.store %v0, %lut[%c0] : memref<8xbf16>
+ memref.store %v1, %lut[%c1] : memref<8xbf16>
+ memref.store %v2, %lut[%c2] : memref<8xbf16>
+ memref.store %v3, %lut[%c3] : memref<8xbf16>
+ memref.store %v4, %lut[%c4] : memref<8xbf16>
+ memref.store %v5, %lut[%c5] : memref<8xbf16>
+ memref.store %v6, %lut[%c6] : memref<8xbf16>
+ memref.store %v7, %lut[%c7] : memref<8xbf16>
+
+ // Fill every row with the repeating pattern of representable values so each
+ // lane exercises the full set of e2m1 values.
+ %A = memref.alloc() : memref<16x16xbf16>
+ scf.for %i = %c0 to %c16 step %c1 {
+ scf.for %j = %c0 to %c16 step %c1 {
+ %jm8 = arith.remui %j, %c8 : index
+ %val = memref.load %lut[%jm8] : memref<8xbf16>
+ memref.store %val, %A[%i, %j] : memref<16x16xbf16>
+ }
+ }
+
+ %B = call @test(%A) : (memref<16x16xbf16>) -> memref<16x16xbf16>
+
+ // Convert the bf16 result to f32 so it can be printed with printMemrefF32.
+ %Bf32 = memref.alloc() : memref<16x16xf32>
+ scf.for %i = %c0 to %c16 step %c1 {
+ scf.for %j = %c0 to %c16 step %c1 {
+ %v = memref.load %B[%i, %j] : memref<16x16xbf16>
+ %vf = arith.extf %v : bf16 to f32
+ memref.store %vf, %Bf32[%i, %j] : memref<16x16xf32>
+ }
+ }
+ %B_cast = memref.cast %Bf32 : memref<16x16xf32> to memref<*xf32>
+ call @printMemrefF32(%B_cast) : (memref<*xf32>) -> ()
+
+ // CHECK: Unranked Memref base@ = 0x{{[0-9a-f]+}}
+ // CHECK-COUNT-16: [0, 0.5, 1, 1.5, 2, 3, 4, 6, 0, 0.5, 1, 1.5, 2, 3, 4, 6]
+ memref.dealloc %A : memref<16x16xbf16>
+ memref.dealloc %B : memref<16x16xbf16>
+ memref.dealloc %Bf32 : memref<16x16xf32>
+ memref.dealloc %lut : memref<8xbf16>
+ return
+ }
+ func.func private @printMemrefF32(%ptr : memref<*xf32>) attributes { llvm.emit_c_interface }
+}
>From 4d2966873e4681857c45da0b94a5165fc19fe220 Mon Sep 17 00:00:00 2001
From: "Lee, Sang Ik" <sang.ik.lee at intel.com>
Date: Wed, 10 Jun 2026 23:47:29 +0000
Subject: [PATCH 2/2] Run clang-format.
---
mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp | 65 ++++++++++---------
1 file changed, 33 insertions(+), 32 deletions(-)
diff --git a/mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp b/mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp
index c833771c69286..cff5effd17bfb 100644
--- a/mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp
+++ b/mlir/lib/Conversion/XeVMToLLVM/XeVMToLLVM.cpp
@@ -1351,14 +1351,14 @@ class ExtfToOCLPattern : public OpConversionPattern<ExtfOp> {
constexpr int kLutE2M1ToBF16 = 5;
int lutIndex =
(dstEtype == ExtfDstElemTypes::F16) ? kLutE2M1ToF16 : kLutE2M1ToBF16;
- Value lutIdx = LLVM::ConstantOp::create(
- rewriter, op.getLoc(), rewriter.getI32Type(), lutIndex);
+ Value lutIdx = LLVM::ConstantOp::create(rewriter, op.getLoc(),
+ rewriter.getI32Type(), lutIndex);
Type lutTy = VectorType::get(16, rewriter.getI32Type());
- Value lut = createDeviceFunctionCall(
- rewriter, "__builtin_IB_shfl_idx4_lut", lutTy,
- {lutIdx.getType()}, {lutIdx}, {}, funcAttrs,
- op.getOperation())
- ->getResult(0);
+ Value lut =
+ createDeviceFunctionCall(rewriter, "__builtin_IB_shfl_idx4_lut",
+ lutTy, {lutIdx.getType()}, {lutIdx}, {},
+ funcAttrs, op.getOperation())
+ ->getResult(0);
Type packedResTy = VectorType::get(8, rewriter.getI32Type());
SmallVector<Type> convArgTypes{lut.getType(), src.getType()};
SmallVector<Value> convArgs{lut, src};
@@ -1787,29 +1787,30 @@ void ::mlir::populateXeVMToLLVMConversionPatterns(ConversionTarget &target,
return !op->hasAttr("cache_control");
});
target.addIllegalDialect<XeVMDialect>();
- patterns.add<LoadStorePrefetchToOCLPattern<BlockLoad2dOp>,
- LoadStorePrefetchToOCLPattern<BlockStore2dOp>,
- LoadStorePrefetchToOCLPattern<BlockPrefetch2dOp>,
- MMAToOCLPattern, MemfenceToOCLPattern, PrefetchToOCLPattern,
- LLVMLoadStoreToOCLPattern<LLVM::LoadOp>,
- LLVMLoadStoreToOCLPattern<LLVM::StoreOp>,
- BlockLoadStore1DToOCLPattern<BlockLoadOp>,
- BlockLoadStore1DToOCLPattern<BlockStoreOp>,
- LaunchConfigOpToOCLPattern<WorkitemIdXOp>,
- LaunchConfigOpToOCLPattern<WorkitemIdYOp>,
- LaunchConfigOpToOCLPattern<WorkitemIdZOp>,
- LaunchConfigOpToOCLPattern<WorkgroupDimXOp>,
- LaunchConfigOpToOCLPattern<WorkgroupDimYOp>,
- LaunchConfigOpToOCLPattern<WorkgroupDimZOp>,
- LaunchConfigOpToOCLPattern<WorkgroupIdXOp>,
- LaunchConfigOpToOCLPattern<WorkgroupIdYOp>,
- LaunchConfigOpToOCLPattern<WorkgroupIdZOp>,
- LaunchConfigOpToOCLPattern<GridDimXOp>,
- LaunchConfigOpToOCLPattern<GridDimYOp>,
- LaunchConfigOpToOCLPattern<GridDimZOp>,
- SubgroupOpWorkitemOpToOCLPattern<LaneIdOp>,
- SubgroupOpWorkitemOpToOCLPattern<SubgroupIdOp>,
- SubgroupOpWorkitemOpToOCLPattern<SubgroupSizeOp>,
- TruncfToOCLPattern, ExtfToOCLPattern, MMAMxToOCLPattern,
- AllocaToGlobalPattern>(patterns.getContext());
+ patterns
+ .add<LoadStorePrefetchToOCLPattern<BlockLoad2dOp>,
+ LoadStorePrefetchToOCLPattern<BlockStore2dOp>,
+ LoadStorePrefetchToOCLPattern<BlockPrefetch2dOp>, MMAToOCLPattern,
+ MemfenceToOCLPattern, PrefetchToOCLPattern,
+ LLVMLoadStoreToOCLPattern<LLVM::LoadOp>,
+ LLVMLoadStoreToOCLPattern<LLVM::StoreOp>,
+ BlockLoadStore1DToOCLPattern<BlockLoadOp>,
+ BlockLoadStore1DToOCLPattern<BlockStoreOp>,
+ LaunchConfigOpToOCLPattern<WorkitemIdXOp>,
+ LaunchConfigOpToOCLPattern<WorkitemIdYOp>,
+ LaunchConfigOpToOCLPattern<WorkitemIdZOp>,
+ LaunchConfigOpToOCLPattern<WorkgroupDimXOp>,
+ LaunchConfigOpToOCLPattern<WorkgroupDimYOp>,
+ LaunchConfigOpToOCLPattern<WorkgroupDimZOp>,
+ LaunchConfigOpToOCLPattern<WorkgroupIdXOp>,
+ LaunchConfigOpToOCLPattern<WorkgroupIdYOp>,
+ LaunchConfigOpToOCLPattern<WorkgroupIdZOp>,
+ LaunchConfigOpToOCLPattern<GridDimXOp>,
+ LaunchConfigOpToOCLPattern<GridDimYOp>,
+ LaunchConfigOpToOCLPattern<GridDimZOp>,
+ SubgroupOpWorkitemOpToOCLPattern<LaneIdOp>,
+ SubgroupOpWorkitemOpToOCLPattern<SubgroupIdOp>,
+ SubgroupOpWorkitemOpToOCLPattern<SubgroupSizeOp>, TruncfToOCLPattern,
+ ExtfToOCLPattern, MMAMxToOCLPattern, AllocaToGlobalPattern>(
+ patterns.getContext());
}
More information about the Mlir-commits
mailing list