[Mlir-commits] [mlir] [MLIR][NVGPU] Add convert.fpext and convert.fptrunc Ops (PR #199700)
Srinivasa Ravi
llvmlistbot at llvm.org
Tue May 26 07:50:02 PDT 2026
https://github.com/Wolfram70 updated https://github.com/llvm/llvm-project/pull/199700
>From 4cfc000d78239519ff6f35a77f648972a50df867 Mon Sep 17 00:00:00 2001
From: Srinivasa Ravi <srinivasar at nvidia.com>
Date: Tue, 26 May 2026 10:09:58 +0000
Subject: [PATCH 1/2] [MLIR][NVGPU] Add convert.fpext and convert.fptrunc Ops
This change adds the `convert.fpext` and `convert.fptrunc` Ops to the
NVGPU dialect to support floating-point conversion operations.
These Ops support scalar, vector and tensor inputs and lower to the
corresponding NVVM Ops after padding and chunking the input into i32
registers.
For tensor inputs, the tensors must be converted to vectors by other means
before they can be successfully lowered.
---
mlir/include/mlir/Conversion/Passes.td | 3 +-
mlir/include/mlir/Dialect/NVGPU/IR/NVGPU.td | 14 +
.../include/mlir/Dialect/NVGPU/IR/NVGPUOps.td | 76 ++
.../lib/Conversion/NVGPUToNVVM/CMakeLists.txt | 1 +
.../Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp | 670 ++++++++++++++++++
mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp | 156 ++++
.../NVGPUToNVVM/nvgpu-convert-fpext.mlir | 466 ++++++++++++
.../NVGPUToNVVM/nvgpu-convert-fptrunc.mlir | 327 +++++++++
.../NVGPU/nvgpu-convert-fpext-invalid.mlir | 106 +++
.../NVGPU/nvgpu-convert-fptrunc-invalid.mlir | 126 ++++
10 files changed, 1944 insertions(+), 1 deletion(-)
create mode 100644 mlir/test/Conversion/NVGPUToNVVM/nvgpu-convert-fpext.mlir
create mode 100644 mlir/test/Conversion/NVGPUToNVVM/nvgpu-convert-fptrunc.mlir
create mode 100644 mlir/test/Dialect/NVGPU/nvgpu-convert-fpext-invalid.mlir
create mode 100644 mlir/test/Dialect/NVGPU/nvgpu-convert-fptrunc-invalid.mlir
diff --git a/mlir/include/mlir/Conversion/Passes.td b/mlir/include/mlir/Conversion/Passes.td
index dda756ddab152..5490f779675ba 100644
--- a/mlir/include/mlir/Conversion/Passes.td
+++ b/mlir/include/mlir/Conversion/Passes.td
@@ -1092,7 +1092,8 @@ def ConvertNVGPUToNVVMPass : Pass<"convert-nvgpu-to-nvvm"> {
"arith::ArithDialect",
"LLVM::LLVMDialect",
"memref::MemRefDialect",
- "NVVM::NVVMDialect"
+ "NVVM::NVVMDialect",
+ "vector::VectorDialect"
];
}
diff --git a/mlir/include/mlir/Dialect/NVGPU/IR/NVGPU.td b/mlir/include/mlir/Dialect/NVGPU/IR/NVGPU.td
index 1c0d7bd1113ea..e56aca779cd1e 100644
--- a/mlir/include/mlir/Dialect/NVGPU/IR/NVGPU.td
+++ b/mlir/include/mlir/Dialect/NVGPU/IR/NVGPU.td
@@ -101,6 +101,20 @@ def TensorMapInterleaveKind : I32EnumAttr<"TensorMapInterleaveKind",
let cppNamespace = "::mlir::nvgpu";
}
+def SubBytesPackedCompact : I32EnumAttrCase<"COMPACT", 0, "compact">;
+def SubBytesPackedU6UnpackU8E3M2 : I32EnumAttrCase<"U6_UNPACK_U8_E3M2", 1, "u6_unpack_u8_e3m2">;
+def SubBytesPackedU6UnpackU8E2M3 : I32EnumAttrCase<"U6_UNPACK_U8_E2M3", 2, "u6_unpack_u8_e2m3">;
+def SubBytesPackedKind : I32EnumAttr<"SubBytesPackedKind",
+ "Sub-bytes packed kind type",
+ [SubBytesPackedCompact, SubBytesPackedU6UnpackU8E3M2,
+ SubBytesPackedU6UnpackU8E2M3]> {
+ let genSpecializedAttr = 0;
+ let cppNamespace = "::mlir::nvgpu";
+}
+def SubBytesPackedKindAttr : EnumAttr<NVGPU_Dialect, SubBytesPackedKind, "subbytes_packedkind"> {
+ let assemblyFormat = "`<` $value `>`";
+}
+
def RcpApprox : I32EnumAttrCase<"APPROX", 0, "approx">;
def RcpRN : I32EnumAttrCase<"RN", 1, "rn">;
def RcpRZ : I32EnumAttrCase<"RZ", 2, "rz">;
diff --git a/mlir/include/mlir/Dialect/NVGPU/IR/NVGPUOps.td b/mlir/include/mlir/Dialect/NVGPU/IR/NVGPUOps.td
index 5a1771eecd0f6..07408ac18f95b 100644
--- a/mlir/include/mlir/Dialect/NVGPU/IR/NVGPUOps.td
+++ b/mlir/include/mlir/Dialect/NVGPU/IR/NVGPUOps.td
@@ -671,4 +671,80 @@ def NVGPU_RcpOp : NVGPU_Op<"rcp", [Pure,
let hasVerifier = 1;
}
+//===----------------------------------------------------------------------===//
+// NVGPU Conversion Ops
+//===----------------------------------------------------------------------===//
+
+def Int8OrFloatLike : TypeConstraint<
+ Or<[FloatLike.predicate,
+ I8.predicate,
+ ValueSemanticsContainerOf<[I8]>.predicate]>,
+ "scalar, vector, or tensor of i8 or floats">;
+def AnyI32Like : TypeOrValueSemanticsContainer<I32, "scalar i32 or vector of i32">;
+
+def NVGPU_FPTruncOp : NVGPU_Op<"convert.fptrunc",
+ [Pure]> {
+ let summary = "Truncate floating-point to narrower floating-point";
+ let description = [{
+ Truncate a floating-point value to a smaller floating-point type.
+ Destination must be strictly narrower than source.
+
+ Supported paths: f32->f16, f32->bf16, f32->f8, f32->f6, f32->f4,
+ f16->f8, f16->f4, bf16->f8, bf16->f4.
+
+ The `random_bits` operand enables stochastic rounding (RS mode)
+ for f32->f16/bf16 conversions. When provided, `rnd` must be RS.
+
+ For f6 types (f6E3M2FN/f6E2M3FN), result is i8 with 2-bit MSB
+ padding; use `packed_kind` to specify the f6 variant. `compact` is not
+ a supported `packed_kind` for f6 types.
+
+ Example:
+ ```mlir
+ %r = nvgpu.cvt_fptrunc %in : vector<8xf32> to vector<8xf8E4M3FN>
+ %r = nvgpu.cvt_fptrunc %in : f32 to f16
+ %r = nvgpu.cvt_fptrunc %in : vector<2x4xf16> to vector<2x4xf8E5M2>
+ ```
+ }];
+ let arguments = (ins FloatLike:$in,
+ DefaultValuedAttr<FPRoundingModeAttr, "NVVM::FPRoundingMode::RN">:$rnd,
+ DefaultValuedAttr<SubBytesPackedKindAttr, "SubBytesPackedKind::COMPACT">:$packed_kind,
+ DefaultValuedAttr<SaturationModeAttr, "NVVM::SaturationMode::SATFINITE">:$sat,
+ DefaultValuedAttr<BoolAttr, "false">:$relu,
+ Optional<I32>:$random_bits
+ );
+ let results = (outs Int8OrFloatLike:$out);
+ let assemblyFormat = "$in (`,` $random_bits^)? attr-dict `:` type($in) `to` type($out)";
+ let hasVerifier = 1;
+}
+
+def NVGPU_FPExtOp : NVGPU_Op<"convert.fpext", [Pure]> {
+ let summary = "Extend floating-point to wider floating-point";
+ let description = [{
+ Extend a floating-point value to a wider floating-point type.
+ Destination must be strictly wider than source.
+
+ Supported paths: f8->f16, f8->bf16, f6->f16, f4->f16, f16->f32,
+ bf16->f32. Narrow->f32 uses two-step lowering (NVVM op to f16/bf16
+ intermediate, then LLVM FPExt). e8m0->bf16 and e8m0->f32 supported
+ (f32 goes through bf16 intermediate).
+
+ For f6 types, input is i8 with `packed_kind` specifying the f6 variant.
+
+ Example:
+ ```mlir
+ %r = nvgpu.cvt_fpext %in : vector<8xf8E4M3FN> to vector<8xf16>
+ %r = nvgpu.cvt_fpext %in : vector<4xf8E5M2> to vector<4xf32>
+ %r = nvgpu.cvt_fpext %in : f8E4M3FN to f32
+ ```
+ }];
+ let arguments = (ins Int8OrFloatLike:$in,
+ DefaultValuedAttr<FPRoundingModeAttr, "NVVM::FPRoundingMode::RN">:$rnd,
+ DefaultValuedAttr<SubBytesPackedKindAttr, "SubBytesPackedKind::COMPACT">:$packed_kind,
+ DefaultValuedAttr<BoolAttr, "false">:$relu);
+ let results = (outs FloatLike:$out);
+ let assemblyFormat = "$in attr-dict `:` type($in) `to` type($out)";
+ let hasVerifier = 1;
+}
+
#endif // MLIR_DIALECT_NVGPU_IR_NVGPUOPS_TD
diff --git a/mlir/lib/Conversion/NVGPUToNVVM/CMakeLists.txt b/mlir/lib/Conversion/NVGPUToNVVM/CMakeLists.txt
index a050749eb7da8..2e0a417bd823d 100644
--- a/mlir/lib/Conversion/NVGPUToNVVM/CMakeLists.txt
+++ b/mlir/lib/Conversion/NVGPUToNVVM/CMakeLists.txt
@@ -18,6 +18,7 @@ add_mlir_conversion_library(MLIRNVGPUToNVVM
MLIRNVGPUDialect
MLIRNVVMDialect
MLIRArithDialect
+ MLIRVectorDialect
MLIRPass
MLIRSCFTransforms
MLIRTransforms
diff --git a/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp b/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
index e566449ffadff..8b992b67aed29 100644
--- a/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
+++ b/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
@@ -20,6 +20,7 @@
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/NVGPU/IR/NVGPUDialect.h"
#include "mlir/Dialect/SCF/Transforms/Patterns.h"
+#include "mlir/Dialect/Vector/IR/VectorOps.h"
#include "mlir/IR/BuiltinTypes.h"
#include "mlir/IR/PatternMatch.h"
#include "mlir/IR/TypeUtilities.h"
@@ -461,6 +462,7 @@ struct ConvertNVGPUToNVVMPass
target.addLegalDialect<::mlir::arith::ArithDialect>();
target.addLegalDialect<::mlir::memref::MemRefDialect>();
target.addLegalDialect<::mlir::NVVM::NVVMDialect>();
+ target.addLegalDialect<::mlir::vector::VectorDialect>();
mlir::scf::populateSCFStructuralTypeConversionsAndLegality(
converter, patterns, target);
if (failed(applyPartialConversion(getOperation(), target,
@@ -1709,6 +1711,669 @@ struct NVGPURcpOpLowering : public ConvertOpToLLVMPattern<nvgpu::RcpOp> {
rewriter);
}
};
+
+//===----------------------------------------------------------------------===//
+// FPTruncOp Lowering
+//===----------------------------------------------------------------------===//
+
+/// Conversion op identifier for nvgpu.convert.fptrunc lowering dispatch table.
+enum class TruncConvOp {
+ F32x2_TO_F16x2,
+ F32x2_TO_BF16x2,
+ F32x2_TO_F8x2,
+ F32x2_TO_F6x2,
+ F32x2_TO_F4x2,
+ F16x2_TO_F8x2,
+ F16x2_TO_F4x2,
+ BF16x2_TO_F8x2,
+ BF16x2_TO_F4x2,
+};
+
+enum class TruncSrcKind { F32, F16, BF16 };
+
+enum class TruncDstKind { F16, BF16, F8, F6, F4 };
+
+struct TruncTableEntry {
+ TruncSrcKind src;
+ TruncDstKind dst;
+ TruncConvOp convOp;
+ int srcStepDecrement; // 2 for f32 pairs, 1 for f16x2/bf16x2
+};
+
+static constexpr TruncTableEntry kTruncTable[] = {
+ // f32 source
+ {TruncSrcKind::F32, TruncDstKind::F16, TruncConvOp::F32x2_TO_F16x2, 2},
+ {TruncSrcKind::F32, TruncDstKind::BF16, TruncConvOp::F32x2_TO_BF16x2, 2},
+ {TruncSrcKind::F32, TruncDstKind::F8, TruncConvOp::F32x2_TO_F8x2, 2},
+ {TruncSrcKind::F32, TruncDstKind::F6, TruncConvOp::F32x2_TO_F6x2, 2},
+ {TruncSrcKind::F32, TruncDstKind::F4, TruncConvOp::F32x2_TO_F4x2, 2},
+ // f16 source
+ {TruncSrcKind::F16, TruncDstKind::F8, TruncConvOp::F16x2_TO_F8x2, 1},
+ {TruncSrcKind::F16, TruncDstKind::F4, TruncConvOp::F16x2_TO_F4x2, 1},
+ // bf16 source
+ {TruncSrcKind::BF16, TruncDstKind::F8, TruncConvOp::BF16x2_TO_F8x2, 1},
+ {TruncSrcKind::BF16, TruncDstKind::F4, TruncConvOp::BF16x2_TO_F4x2, 1},
+};
+
+static bool isConvertibleF8Type(Type t) {
+ return isa<Float8E4M3FNType, Float8E5M2Type, Float8E8M0FNUType>(t);
+}
+
+static std::optional<TruncSrcKind> classifySrcType(Type t) {
+ if (t.isF32())
+ return TruncSrcKind::F32;
+ if (t.isF16())
+ return TruncSrcKind::F16;
+ if (t.isBF16())
+ return TruncSrcKind::BF16;
+ return std::nullopt;
+}
+
+static std::optional<TruncDstKind> classifyDstType(Type t) {
+ if (t.isF16())
+ return TruncDstKind::F16;
+ if (t.isBF16())
+ return TruncDstKind::BF16;
+ if (isConvertibleF8Type(t))
+ return TruncDstKind::F8;
+ int bitWidth = t.getIntOrFloatBitWidth();
+ if (isa<IntegerType>(t) && bitWidth == 8)
+ return TruncDstKind::F6;
+ if (bitWidth == 4)
+ return TruncDstKind::F4;
+ return std::nullopt;
+}
+
+static std::optional<std::pair<TruncConvOp, int>>
+lookupTruncConvOp(Type srcElemType, Type dstElemType) {
+ auto srcKind = classifySrcType(srcElemType);
+ auto dstKind = classifyDstType(dstElemType);
+ if (!srcKind || !dstKind)
+ return std::nullopt;
+ for (const auto &entry : kTruncTable) {
+ if (entry.src == *srcKind && entry.dst == *dstKind)
+ return {{entry.convOp, entry.srcStepDecrement}};
+ }
+ return std::nullopt;
+}
+
+static Value extractElement(ImplicitLocOpBuilder &b, Value srcVec, int idx) {
+ IntegerType i64Ty = b.getI64Type();
+ return b.create<LLVM::ExtractElementOp>(
+ srcVec, b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(idx)));
+}
+
+/// Extract a pair of f32 values from an i32 vector at the given base index.
+/// Returns {f32_lo (lower index), f32_hi (higher index)}.
+static std::pair<Value, Value> extractF32Pair(ImplicitLocOpBuilder &b,
+ Value srcI32Vec, int baseIdx) {
+ FloatType f32Ty = b.getF32Type();
+ Value elem0 = extractElement(b, srcI32Vec, baseIdx);
+ Value elem1 = extractElement(b, srcI32Vec, baseIdx + 1);
+ return {b.create<LLVM::BitcastOp>(f32Ty, elem0),
+ b.create<LLVM::BitcastOp>(f32Ty, elem1)};
+}
+
+static Value extractAndBitcast(ImplicitLocOpBuilder &b, Value srcI32Vec,
+ int idx, VectorType vecTy) {
+ Value elem = extractElement(b, srcI32Vec, idx);
+ return b.create<LLVM::BitcastOp>(vecTy, elem);
+}
+
+/// Dispatch to the specific NVVM conversion op based on TruncConvOp,
+/// extract source operands, and return the native result.
+/// - f16/bf16 destinations: returns i32 (bitcast from vector<2xf16/bf16>)
+/// - f8/f6 destinations: returns i16 (packed 2 x f8/f6)
+/// - f4 destinations: returns i8 (packed 2 x f4)
+/// Create a sub-byte conversion from an f32 pair source.
+template <typename ConvertOp, typename... Args>
+static Value convertFromF32Pair(ImplicitLocOpBuilder &b, Value srcI32Vec,
+ int srcBaseIdx, Type resultTy, Args &&...args) {
+ auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
+ return b.create<ConvertOp>(resultTy, hi, lo, std::forward<Args>(args)...);
+}
+
+/// Create a sub-byte conversion from a packed f16x2/bf16x2 source.
+template <typename ConvertOp, typename... Args>
+static Value convertFromPacked(ImplicitLocOpBuilder &b, Value srcI32Vec,
+ int srcBaseIdx, Type srcElemTy, Type resultTy,
+ Args &&...args) {
+ Value src = extractAndBitcast(b, srcI32Vec, srcBaseIdx,
+ VectorType::get(2, srcElemTy));
+ return b.create<ConvertOp>(resultTy, src, std::forward<Args>(args)...);
+}
+
+static Value createTruncConversion(
+ ImplicitLocOpBuilder &b, MLIRContext *ctx, TruncConvOp convOp,
+ Value srcI32Vec, int srcBaseIdx, NVVM::FPRoundingModeAttr rndAttr,
+ NVVM::SaturationModeAttr satAttr, BoolAttr reluAttr, Type dstElemType,
+ Type actualDstFloatType, Value randomBits = Value()) {
+ IntegerType i8Ty = b.getI8Type();
+ IntegerType i16Ty = b.getI16Type();
+ IntegerType i32Ty = b.getI32Type();
+ auto dstTyAttr = TypeAttr::get(dstElemType);
+ auto actualDstTyAttr = TypeAttr::get(actualDstFloatType);
+
+ switch (convOp) {
+ case TruncConvOp::F32x2_TO_F16x2: {
+ auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
+ Value r = b.create<NVVM::ConvertF32x2ToF16x2Op>(
+ VectorType::get(2, b.getF16Type()), hi, lo, randomBits, rndAttr,
+ satAttr, reluAttr);
+ return b.create<LLVM::BitcastOp>(i32Ty, r);
+ }
+ case TruncConvOp::F32x2_TO_BF16x2: {
+ auto [lo, hi] = extractF32Pair(b, srcI32Vec, srcBaseIdx);
+ Value r = b.create<NVVM::ConvertF32x2ToBF16x2Op>(
+ VectorType::get(2, b.getBF16Type()), hi, lo, randomBits, rndAttr,
+ satAttr, reluAttr);
+ return b.create<LLVM::BitcastOp>(i32Ty, r);
+ }
+ case TruncConvOp::F32x2_TO_F8x2:
+ return convertFromF32Pair<NVVM::ConvertF32x2ToF8x2Op>(
+ b, srcI32Vec, srcBaseIdx, i16Ty, rndAttr, satAttr, reluAttr, dstTyAttr);
+ case TruncConvOp::F32x2_TO_F6x2:
+ return convertFromF32Pair<NVVM::ConvertF32x2ToF6x2Op>(
+ b, srcI32Vec, srcBaseIdx, i16Ty, reluAttr, actualDstTyAttr);
+ case TruncConvOp::F32x2_TO_F4x2:
+ return convertFromF32Pair<NVVM::ConvertF32x2ToF4x2Op>(
+ b, srcI32Vec, srcBaseIdx, i8Ty, reluAttr, dstTyAttr);
+ case TruncConvOp::F16x2_TO_F8x2:
+ return convertFromPacked<NVVM::ConvertF16x2ToF8x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getF16Type(), i16Ty, reluAttr, dstTyAttr);
+ case TruncConvOp::F16x2_TO_F4x2:
+ return convertFromPacked<NVVM::ConvertF16x2ToF4x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getF16Type(), i8Ty, reluAttr,
+ actualDstTyAttr);
+ case TruncConvOp::BF16x2_TO_F8x2:
+ return convertFromPacked<NVVM::ConvertBF16x2ToF8x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i16Ty, rndAttr, satAttr,
+ reluAttr, dstTyAttr);
+ case TruncConvOp::BF16x2_TO_F4x2:
+ return convertFromPacked<NVVM::ConvertBF16x2ToF4x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i8Ty, reluAttr,
+ actualDstTyAttr);
+ }
+ llvm_unreachable("unhandled TruncConvOp");
+}
+
+struct NVGPUFPTruncOpLowering
+ : public ConvertOpToLLVMPattern<nvgpu::FPTruncOp> {
+ using ConvertOpToLLVMPattern<nvgpu::FPTruncOp>::ConvertOpToLLVMPattern;
+
+ static Type getActualDstFloatType(MLIRContext *ctx, Type elemType,
+ nvgpu::SubBytesPackedKind packedKind) {
+ if (isa<IntegerType>(elemType)) {
+ if (packedKind == nvgpu::SubBytesPackedKind::U6_UNPACK_U8_E3M2)
+ return Float6E3M2FNType::get(ctx);
+ return Float6E2M3FNType::get(ctx);
+ }
+ return elemType;
+ }
+
+ LogicalResult
+ matchAndRewrite(nvgpu::FPTruncOp op, OpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ // Tensor inputs are not handled here; they must be converted to vectors
+ // by other means before this pattern runs.
+ if (isa<RankedTensorType>(op.getIn().getType()))
+ return rewriter.notifyMatchFailure(
+ op, "tensor inputs not handled; type converter should lower first");
+
+ MLIRContext *ctx = getContext();
+ ImplicitLocOpBuilder b(op->getLoc(), rewriter);
+ IntegerType i32Ty = b.getI32Type();
+ IntegerType i64Ty = b.getI64Type();
+
+ static constexpr int regBits = 32;
+ auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
+ auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
+ if (!srcType || srcType.getRank() != 1 || !dstType ||
+ dstType.getRank() != 1)
+ return rewriter.notifyMatchFailure(
+ op, "expected 1-D vector; canonicalize pattern handles other shapes");
+
+ auto srcElemType = srcType.getElementType();
+ auto dstElemType = dstType.getElementType();
+ int srcBW = srcType.getElementTypeBitWidth();
+ int dstBW = dstType.getElementTypeBitWidth();
+ int numElems = srcType.getNumElements();
+
+ auto packedKind = op.getPackedKind();
+ NVVM::FPRoundingModeAttr rndModeAttr = op.getRndAttr();
+ NVVM::SaturationModeAttr satModeAttr = op.getSatAttr();
+ auto reluBoolAttr = op.getReluAttr();
+ Value randomBits = adaptor.getRandomBits();
+ Type actualDstFloatType =
+ getActualDstFloatType(ctx, dstElemType, packedKind);
+
+ // STEP 1: bitcast input vector to i32 register vector.
+ // f64 source: decompose to f64->f32 (LLVM fptrunc), then lower f32->dst.
+ Value input = adaptor.getIn();
+ if (srcBW == 64) {
+ auto f32VecTy = VectorType::get(srcType.getShape(), b.getF32Type());
+ input = b.create<LLVM::FPTruncOp>(f32VecTy, input);
+ srcType = f32VecTy;
+ srcElemType = b.getF32Type();
+ srcBW = 32;
+ if (dstBW == 32) {
+ rewriter.replaceOp(op, input);
+ return success();
+ }
+ }
+
+ int srcI32Elems = numElems * srcBW / regBits;
+ int dstI32Elems = numElems * dstBW / regBits;
+ Value srcI32Vec =
+ b.create<LLVM::BitcastOp>(VectorType::get(srcI32Elems, i32Ty), input);
+ Value dstI32Vec =
+ b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
+
+ // STEP 2: look up the conversion op from the (srcType, dstType) table.
+ auto convEntry = lookupTruncConvOp(srcElemType, dstElemType);
+ if (!convEntry)
+ return rewriter.notifyMatchFailure(
+ op, "unsupported type combination for truncation");
+ auto [convOp, srcStepDecrement] = *convEntry;
+
+ // STEP 3: pack conversion results into destination i32 vector.
+ const int srcStep = srcBW / dstBW;
+ const int resultBW =
+ dstBW * 2; // each conversion produces 2 (packed) elements
+ const int numConvsPerI32 = regBits / resultBW;
+
+ for (int srcIdx = 0, dstIdx = 0; dstIdx < dstI32Elems;
+ srcIdx += srcStep, dstIdx++) {
+ Value dstIdxConst =
+ b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
+ Value dstValue;
+
+ if (numConvsPerI32 == 1) {
+ // f16/bf16 destinations
+ dstValue = createTruncConversion(
+ b, ctx, convOp, srcI32Vec, srcIdx, rndModeAttr, satModeAttr,
+ reluBoolAttr, dstElemType, actualDstFloatType, randomBits);
+ } else {
+ // f8/f6 destinations: pack sub-results via vector insert + bitcast.
+ auto subResultType = IntegerType::get(ctx, resultBW);
+ auto subVecTy = VectorType::get(numConvsPerI32, subResultType);
+ Value subVec = b.create<LLVM::UndefOp>(subVecTy);
+
+ int insertIdx = numConvsPerI32 - 1;
+ int curStep = srcStep;
+ while (curStep > 0) {
+ curStep -= srcStepDecrement;
+ Value subResult = createTruncConversion(
+ b, ctx, convOp, srcI32Vec, srcIdx + curStep, rndModeAttr,
+ satModeAttr, reluBoolAttr, dstElemType, actualDstFloatType,
+ /*randomBits=*/Value());
+ subVec = b.create<LLVM::InsertElementOp>(
+ subVec, subResult,
+ b.create<LLVM::ConstantOp>(i64Ty,
+ b.getI64IntegerAttr(insertIdx)));
+ insertIdx--;
+ }
+
+ dstValue = b.create<LLVM::BitcastOp>(i32Ty, subVec);
+ }
+
+ dstI32Vec =
+ b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
+ }
+
+ Type convertedType = getTypeConverter()->convertType(dstType);
+ assert(convertedType && "failed to convert type");
+ auto dstVec = b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
+ rewriter.replaceOp(op, dstVec);
+ return success();
+ }
+};
+
+//===----------------------------------------------------------------------===//
+// NVGPUFPExtOpLowering
+//===----------------------------------------------------------------------===//
+
+/// Extension conversion op identifier for nvgpu.convert.fpext lowering.
+enum class ExtConvOp {
+ F8x2_TO_F16x2,
+ F8x2_TO_BF16x2,
+ F6x2_TO_F16x2,
+ F4x2_TO_F16x2,
+};
+
+enum class ExtSrcKind { F8, F6, F4 };
+enum class ExtDstKind { F16, BF16 };
+
+struct ExtTableEntry {
+ ExtSrcKind src;
+ ExtDstKind dst;
+ ExtConvOp convOp;
+};
+
+static constexpr ExtTableEntry kExtTable[] = {
+ {ExtSrcKind::F8, ExtDstKind::F16, ExtConvOp::F8x2_TO_F16x2},
+ {ExtSrcKind::F8, ExtDstKind::BF16, ExtConvOp::F8x2_TO_BF16x2},
+ {ExtSrcKind::F6, ExtDstKind::F16, ExtConvOp::F6x2_TO_F16x2},
+ {ExtSrcKind::F4, ExtDstKind::F16, ExtConvOp::F4x2_TO_F16x2},
+};
+
+static bool isExtConvertibleF8Type(Type t) {
+ return isa<Float8E4M3FNType, Float8E5M2Type, Float8E8M0FNUType>(t);
+}
+
+static std::optional<ExtSrcKind> classifyExtSrcType(Type t) {
+ if (isExtConvertibleF8Type(t))
+ return ExtSrcKind::F8;
+ int bitWidth = t.getIntOrFloatBitWidth();
+ if (isa<IntegerType>(t) && bitWidth == 8)
+ return ExtSrcKind::F6;
+ if (bitWidth == 4)
+ return ExtSrcKind::F4;
+ return std::nullopt;
+}
+
+static std::optional<ExtDstKind> classifyExtDstType(Type t) {
+ if (t.isF16())
+ return ExtDstKind::F16;
+ if (t.isBF16())
+ return ExtDstKind::BF16;
+ return std::nullopt;
+}
+
+static std::optional<ExtConvOp> lookupExtConvOp(Type srcElemType,
+ Type dstElemType) {
+ auto srcKind = classifyExtSrcType(srcElemType);
+ auto dstKind = classifyExtDstType(dstElemType);
+ if (!srcKind || !dstKind)
+ return std::nullopt;
+ for (const auto &entry : kExtTable) {
+ if (entry.src == *srcKind && entry.dst == *dstKind)
+ return entry.convOp;
+ }
+ return std::nullopt;
+}
+
+static Type getActualSrcFloatType(MLIRContext *ctx, Type elemType,
+ nvgpu::SubBytesPackedKind packedKind) {
+ if (isa<IntegerType>(elemType)) {
+ if (packedKind == nvgpu::SubBytesPackedKind::U6_UNPACK_U8_E3M2)
+ return Float6E3M2FNType::get(ctx);
+ return Float6E2M3FNType::get(ctx);
+ }
+ return elemType;
+}
+
+/// Create a typed NVVM extension conversion.
+/// For f8/f6: src is vector<2xi8>. For f4: src is i8.
+/// Returns i32 (bitcast from vector<2xf16> or vector<2xbf16>).
+static Value createExtConversion(ImplicitLocOpBuilder &b, MLIRContext *ctx,
+ ExtConvOp convOp, Value src, BoolAttr reluAttr,
+ Type actualSrcFloatType,
+ Value extScaleFactor = Value()) {
+ IntegerType i32Ty = b.getI32Type();
+ auto srcTyAttr = TypeAttr::get(actualSrcFloatType);
+
+ switch (convOp) {
+ case ExtConvOp::F8x2_TO_F16x2: {
+ Value r = NVVM::ConvertF8x2ToF16x2Op::create(
+ b, VectorType::get(2, b.getF16Type()), src, reluAttr, srcTyAttr);
+ return b.create<LLVM::BitcastOp>(i32Ty, r);
+ }
+ case ExtConvOp::F8x2_TO_BF16x2: {
+ Value r = NVVM::ConvertF8x2ToBF16x2Op::create(
+ b, VectorType::get(2, b.getBF16Type()), src, srcTyAttr);
+ return b.create<LLVM::BitcastOp>(i32Ty, r);
+ }
+ case ExtConvOp::F6x2_TO_F16x2: {
+ Value r = NVVM::ConvertF6x2ToF16x2Op::create(
+ b, VectorType::get(2, b.getF16Type()), src, reluAttr, srcTyAttr);
+ return b.create<LLVM::BitcastOp>(i32Ty, r);
+ }
+ case ExtConvOp::F4x2_TO_F16x2: {
+ Value r = NVVM::ConvertF4x2ToF16x2Op::create(
+ b, VectorType::get(2, b.getF16Type()), src, reluAttr, srcTyAttr);
+ return b.create<LLVM::BitcastOp>(i32Ty, r);
+ }
+ }
+ llvm_unreachable("unhandled ExtConvOp");
+}
+
+struct NVGPUFPExtOpLowering : public ConvertOpToLLVMPattern<nvgpu::FPExtOp> {
+ using ConvertOpToLLVMPattern<nvgpu::FPExtOp>::ConvertOpToLLVMPattern;
+
+ LogicalResult
+ matchAndRewrite(nvgpu::FPExtOp op, OpAdaptor adaptor,
+ ConversionPatternRewriter &rewriter) const override {
+ // Tensor inputs are not handled here; they must be converted to vectors
+ // by other means before this pattern runs.
+ if (isa<RankedTensorType>(op.getIn().getType()))
+ return rewriter.notifyMatchFailure(
+ op, "tensor inputs not handled; type converter should lower first");
+
+ MLIRContext *ctx = getContext();
+ ImplicitLocOpBuilder b(op->getLoc(), rewriter);
+ IntegerType i8Ty = b.getI8Type();
+ IntegerType i16Ty = b.getI16Type();
+ IntegerType i32Ty = b.getI32Type();
+ IntegerType i64Ty = b.getI64Type();
+
+ static constexpr int regBits = 32;
+ auto srcType = llvm::dyn_cast<VectorType>(op.getIn().getType());
+ auto dstType = llvm::dyn_cast<VectorType>(op.getOut().getType());
+ if (!srcType || srcType.getRank() != 1 || !dstType ||
+ dstType.getRank() != 1)
+ return rewriter.notifyMatchFailure(
+ op, "expected 1-D vector; canonicalize pattern handles other shapes");
+
+ auto srcElemType = srcType.getElementType();
+ auto dstElemType = dstType.getElementType();
+ int srcBW = srcType.getElementTypeBitWidth();
+ int dstBW = dstType.getElementTypeBitWidth();
+ int numElems = srcType.getNumElements();
+
+ auto packedKind = op.getPackedKind();
+ auto reluBoolAttr = op.getReluAttr();
+ Type actualSrcFloatType =
+ getActualSrcFloatType(ctx, srcElemType, packedKind);
+
+ assert(dstBW == 16 || dstBW == 32 || dstBW == 64);
+
+ // Wide source (f16/bf16/f32) to wide destination (f32/f64): single FPExt.
+ if (srcBW >= 16 && dstBW >= 32) {
+ Value result = adaptor.getIn();
+ if (srcElemType != dstElemType) {
+ Type convertedType = getTypeConverter()->convertType(dstType);
+ assert(convertedType && "failed to convert type");
+ result = b.create<LLVM::FPExtOp>(convertedType, result);
+ }
+ rewriter.replaceOp(op, result);
+ return success();
+ }
+
+ // Narrow source (f8/f6/f4): NVVM typed op produces f16/bf16; optionally
+ // followed by FPExt to the final f32/f64 destination.
+ bool needsFinalFPExt = (dstBW >= 32);
+ // e8m0 must go through bf16 (only available NVVM op); others use f16.
+ Type intermediateDstElem = dstElemType;
+ if (needsFinalFPExt && llvm::isa<Float8E8M0FNUType>(srcElemType))
+ intermediateDstElem = b.getBF16Type();
+ else if (needsFinalFPExt)
+ intermediateDstElem = b.getF16Type();
+ int intermediateDstBW = needsFinalFPExt ? 16 : dstBW;
+
+ // STEP 1: bitcast input vector to i32 register vector.
+ int srcI32Elems = numElems * srcBW / regBits;
+ int dstI32Elems = numElems * intermediateDstBW / regBits;
+ Value srcI32Vec = b.create<LLVM::BitcastOp>(
+ VectorType::get(srcI32Elems, i32Ty), adaptor.getIn());
+ Value dstI32Vec =
+ b.create<LLVM::UndefOp>(VectorType::get(dstI32Elems, i32Ty));
+
+ // STEP 2: look up the conversion op from the (srcType, dstType) table.
+ auto convOpOpt = lookupExtConvOp(srcElemType, intermediateDstElem);
+ if (!convOpOpt)
+ return rewriter.notifyMatchFailure(
+ op, "unsupported type combination for extension");
+ ExtConvOp convOp = *convOpOpt;
+ Value extScaleFactor;
+
+ // STEP 3: iterate over source i32 elements, producing destination i32s.
+ for (int srcIdx = 0, dstIdx = 0; srcIdx < srcI32Elems; srcIdx++) {
+ Value srcI32 = b.create<LLVM::ExtractElementOp>(
+ srcI32Vec,
+ b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(srcIdx)));
+
+ if (srcBW == 8) {
+ // f8/f6: one i32 holds 4 bytes -> split into 2 pairs of i16 -> 2 convs.
+ Value i16Vec =
+ b.create<LLVM::BitcastOp>(VectorType::get(2, i16Ty), srcI32);
+ for (int half = 0; half < 2; half++) {
+ Value halfI16 = b.create<LLVM::ExtractElementOp>(
+ i16Vec,
+ b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(half)));
+ Value src =
+ b.create<LLVM::BitcastOp>(VectorType::get(2, i8Ty), halfI16);
+ Value dstValue =
+ createExtConversion(b, ctx, convOp, src, reluBoolAttr,
+ actualSrcFloatType, extScaleFactor);
+ Value dstIdxConst =
+ b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
+ dstI32Vec =
+ b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
+ dstIdx++;
+ }
+ } else {
+ // f4: one i32 holds 4 bytes -> each byte is one conversion input.
+ Value i8Vec =
+ b.create<LLVM::BitcastOp>(VectorType::get(4, i8Ty), srcI32);
+ for (int byteIdx = 0; byteIdx < 4; byteIdx++) {
+ Value src = b.create<LLVM::ExtractElementOp>(
+ i8Vec,
+ b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(byteIdx)));
+ Value dstValue =
+ createExtConversion(b, ctx, convOp, src, reluBoolAttr,
+ actualSrcFloatType, extScaleFactor);
+ Value dstIdxConst =
+ b.create<LLVM::ConstantOp>(i64Ty, b.getI64IntegerAttr(dstIdx));
+ dstI32Vec =
+ b.create<LLVM::InsertElementOp>(dstI32Vec, dstValue, dstIdxConst);
+ dstIdx++;
+ }
+ }
+ }
+
+ // STEP 4: produce final result.
+ Type convertedType = getTypeConverter()->convertType(dstType);
+ assert(convertedType && "failed to convert type");
+ Value result;
+ if (needsFinalFPExt) {
+ auto intermediateVecTy = VectorType::get(numElems, intermediateDstElem);
+ Value intermediateVec =
+ b.create<LLVM::BitcastOp>(intermediateVecTy, dstI32Vec);
+ result = b.create<LLVM::FPExtOp>(convertedType, intermediateVec);
+ } else {
+ result = b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
+ }
+ rewriter.replaceOp(op, result);
+ return success();
+ }
+};
+
+static int64_t computePaddedElems(int64_t numElems, int srcBW, int dstBW,
+ int step) {
+ static constexpr int regBits = 32;
+ auto ceilDiv = [](int64_t x, int64_t y) { return (x + y - 1) / y; };
+ int64_t padded =
+ std::max(ceilDiv(numElems * srcBW, regBits) * regBits / srcBW,
+ ceilDiv(numElems * dstBW, regBits) * regBits / dstBW);
+ return ceilDiv(padded, step) * step;
+}
+
+/// Canonicalization pattern for nvgpu.convert.fptrunc / nvgpu.convert.fpext: handles
+/// scalar inputs, non-32-bit-aligned vectors, and multi-rank vectors. Runs as
+/// an OpRewritePattern on MLIR types before LLVM type conversion.
+template <typename CvtOp, bool IsTrunc>
+struct NVGPUFPCanonicalizePattern : public OpRewritePattern<CvtOp> {
+ using OpRewritePattern<CvtOp>::OpRewritePattern;
+
+ LogicalResult matchAndRewrite(CvtOp op,
+ PatternRewriter &rewriter) const override {
+ Type inType = op.getIn().getType();
+ Type outType = op.getOut().getType();
+
+ // Tensor inputs are not handled here; they must be converted to vectors
+ // by other means before this pattern runs.
+ if (isa<RankedTensorType>(inType))
+ return failure();
+
+ Type srcElemTy = getElementTypeOrSelf(inType);
+ Type dstElemTy = getElementTypeOrSelf(outType);
+ int srcBW = srcElemTy.getIntOrFloatBitWidth();
+ int dstBW = dstElemTy.getIntOrFloatBitWidth();
+
+ bool isScalar = !isa<VectorType>(inType);
+ auto srcVecTy = dyn_cast<VectorType>(inType);
+ bool isMultiRank = srcVecTy && srcVecTy.getRank() > 1;
+ int64_t numElems = isScalar ? 1 : srcVecTy.getNumElements();
+ int step = IsTrunc ? srcBW / dstBW : dstBW / srcBW;
+ int64_t paddedElems = computePaddedElems(numElems, srcBW, dstBW, step);
+ bool needsPad = (paddedElems != numElems);
+
+ if (!isScalar && !isMultiRank && !needsPad)
+ return failure();
+
+ ImplicitLocOpBuilder b(op->getLoc(), rewriter);
+ Value input = op.getIn();
+
+ if (isScalar)
+ input = vector::BroadcastOp::create(b, VectorType::get({1}, srcElemTy),
+ input);
+
+ if (isMultiRank)
+ input = vector::ShapeCastOp::create(
+ b, VectorType::get({numElems}, srcElemTy), input);
+
+ if (needsPad) {
+ auto paddedTy = VectorType::get({paddedElems}, srcElemTy);
+ Value zero = arith::ConstantOp::create(
+ b, DenseElementsAttr::get(paddedTy, b.getZeroAttr(srcElemTy)));
+ input = vector::InsertStridedSliceOp::create(
+ b, input, zero, SmallVector<int64_t>{0}, SmallVector<int64_t>{1});
+ }
+
+ auto cvtDstTy =
+ VectorType::get({needsPad ? paddedElems : numElems}, dstElemTy);
+ Value cvt;
+ if constexpr (IsTrunc)
+ cvt = CvtOp::create(b, cvtDstTy, input, op.getRndAttr(),
+ op.getPackedKindAttr(), op.getSatAttr(),
+ op.getReluAttr(), op.getRandomBits());
+ else
+ cvt = CvtOp::create(b, cvtDstTy, input, op.getRndAttr(),
+ op.getPackedKindAttr(), op.getReluAttr());
+ Value result = cvt;
+
+ if (needsPad)
+ result = vector::ExtractStridedSliceOp::create(
+ b, result, SmallVector<int64_t>{0}, SmallVector<int64_t>{numElems},
+ SmallVector<int64_t>{1});
+
+ if (isMultiRank)
+ result =
+ vector::ShapeCastOp::create(b, cast<VectorType>(outType), result);
+
+ if (isScalar)
+ result = vector::ExtractOp::create(b, result, SmallVector<int64_t>{0});
+
+ rewriter.replaceOp(op, result);
+ return success();
+ }
+};
+
+using NVGPUFPTruncCanonicalizePattern =
+ NVGPUFPCanonicalizePattern<nvgpu::FPTruncOp, true>;
+using NVGPUFPExtCanonicalizePattern =
+ NVGPUFPCanonicalizePattern<nvgpu::FPExtOp, false>;
} // namespace
void mlir::nvgpu::populateCommonGPUTypeAndAttributeConversions(
@@ -1753,7 +2418,12 @@ void mlir::populateNVGPUToNVVMConversionPatterns(
NVGPUWarpgroupMmaOpLowering, // nvgpu.warpgroup.mma
NVGPUWarpgroupMmaStoreOpLowering, // nvgpu.warpgroup.mma.store
NVGPUWarpgroupMmaInitAccumulatorOpLowering, // nvgpu.warpgroup.mma.init.accumulator
+ NVGPUFPTruncOpLowering, // nvgpu.convert.fptrunc
+ NVGPUFPExtOpLowering, // nvgpu.convert.fpext
MmaSyncOptoNVVM, MmaLdMatrixOpToNVVM, NVGPUAsyncCopyLowering,
NVGPUAsyncCreateGroupLowering, NVGPUAsyncWaitLowering,
NVGPUMmaSparseSyncLowering, NVGPURcpOpLowering>(converter);
+
+ patterns.add<NVGPUFPTruncCanonicalizePattern, NVGPUFPExtCanonicalizePattern>(
+ patterns.getContext());
}
diff --git a/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp b/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp
index 6e1ed05c9d0d0..9aa28ca825f73 100644
--- a/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp
+++ b/mlir/lib/Dialect/NVGPU/IR/NVGPUDialect.cpp
@@ -699,6 +699,162 @@ LogicalResult RcpOp::verify() {
return success();
}
+//===----------------------------------------------------------------------===//
+// NVGPU_CvtFPTruncOp
+//===----------------------------------------------------------------------===//
+
+static bool isShapedContainerType(Type t) {
+ return llvm::isa<VectorType, RankedTensorType, UnrankedTensorType>(t);
+}
+
+static LogicalResult verifyConversionShapes(Operation *op, Type inType,
+ Type outType) {
+ if (llvm::isa<UnrankedTensorType>(inType) ||
+ llvm::isa<UnrankedTensorType>(outType))
+ return op->emitOpError("unranked tensor types are not supported, got ")
+ << inType << " and " << outType;
+ bool srcIsShaped = isShapedContainerType(inType);
+ bool dstIsShaped = isShapedContainerType(outType);
+ if (srcIsShaped != dstIsShaped)
+ return op->emitOpError("input and output must both be scalars or both be "
+ "vectors/tensors, got ")
+ << inType << " and " << outType;
+ if (srcIsShaped) {
+ auto srcShaped = llvm::cast<ShapedType>(inType);
+ auto dstShaped = llvm::cast<ShapedType>(outType);
+ if (srcShaped.getRank() == 0 || dstShaped.getRank() == 0)
+ return op->emitOpError("rank-0 shaped types are not supported, use "
+ "scalar type instead");
+ if (srcShaped.getShape() != dstShaped.getShape())
+ return op->emitOpError("input and output shapes must match, got ")
+ << inType << " and " << outType;
+ if (llvm::isa<VectorType>(inType) != llvm::isa<VectorType>(outType))
+ return op->emitOpError("input and output must be the same container "
+ "type (both vector or both tensor), got ")
+ << inType << " and " << outType;
+ }
+ return success();
+}
+
+LogicalResult FPTruncOp::verify() {
+ Type inType = getIn().getType();
+ Type outType = getType();
+ Type srcType = getElementTypeOrSelf(inType);
+ Type dstType = getElementTypeOrSelf(outType);
+ SubBytesPackedKind packedKind = getPackedKind();
+ int srcBitWidth = srcType.getIntOrFloatBitWidth();
+ int dstBitWidth = dstType.getIntOrFloatBitWidth();
+ auto rnd = getRnd();
+
+ if (auto result = verifyConversionShapes(getOperation(), inType, outType);
+ failed(result))
+ return result;
+
+ if (srcBitWidth <= dstBitWidth)
+ return emitOpError("result type ")
+ << dstType << " must be narrower than operand type " << srcType;
+
+ if (!(srcBitWidth == 64 || srcBitWidth == 32 || srcBitWidth == 16))
+ return emitOpError("input type must be 64/32/16 bitwidth, but got ")
+ << srcBitWidth;
+
+ if (dstBitWidth == 6)
+ return emitOpError("currently doesn't support fp6 compact result type");
+
+ if ((packedKind == SubBytesPackedKind::U6_UNPACK_U8_E3M2 ||
+ packedKind == SubBytesPackedKind::U6_UNPACK_U8_E2M3) &&
+ !dstType.isInteger(8))
+ return emitOpError("result type expects `i8` with `u6_unpack_u8` packed "
+ "kind, but got ")
+ << dstType;
+
+ if (packedKind == SubBytesPackedKind::COMPACT &&
+ !llvm::isa<FloatType>(dstType))
+ return emitOpError("result type expects float type with `compact` packed "
+ "kind, but got ")
+ << dstType;
+
+ if (llvm::isa<Float8E8M0FNUType>(dstType)) {
+ if (rnd != mlir::NVVM::FPRoundingMode::RZ &&
+ rnd != mlir::NVVM::FPRoundingMode::RP)
+ return emitOpError("expects RZ or RP rounding mode when result type is "
+ "e8m0, but got ")
+ << getRndAttr();
+ } else if (rnd == mlir::NVVM::FPRoundingMode::RS) {
+ // TODO: Currently, we only support conversions which fit into a single i32
+ // register. Support f32->f8/f6/f4 conversions with RS rounding.
+ if (!(srcBitWidth == 32 && (dstBitWidth == 16)))
+ return emitOpError("RS (stochastic) rounding is only supported for "
+ "f32->f16/bf16, got ")
+ << srcType << " -> " << dstType;
+ if (!getRandomBits())
+ return emitOpError("random_bits operand is required with RS rounding");
+ } else if (rnd != mlir::NVVM::FPRoundingMode::RN) {
+ return emitOpError("expects RN rounding mode, but got ") << getRndAttr();
+ }
+
+ if (getRandomBits() && rnd != mlir::NVVM::FPRoundingMode::RS)
+ return emitOpError("random_bits can only be used with RS rounding mode");
+
+ return success();
+}
+
+//===----------------------------------------------------------------------===//
+// NVGPU_CvtFPExtOp
+//===----------------------------------------------------------------------===//
+
+LogicalResult FPExtOp::verify() {
+ Type inType = getIn().getType();
+ Type outType = getType();
+ Type srcType = getElementTypeOrSelf(inType);
+ Type dstType = getElementTypeOrSelf(outType);
+ SubBytesPackedKind packedKind = getPackedKind();
+ int srcBitWidth = srcType.getIntOrFloatBitWidth();
+ int dstBitWidth = dstType.getIntOrFloatBitWidth();
+ auto rnd = getRnd();
+
+ if (auto result = verifyConversionShapes(getOperation(), inType, outType);
+ failed(result))
+ return result;
+
+ if (srcBitWidth >= dstBitWidth)
+ return emitOpError("result type ")
+ << dstType << " must be wider than operand type " << srcType;
+
+ if (dstBitWidth != 16 && dstBitWidth != 32 && dstBitWidth != 64)
+ return emitOpError("result type must be 16, 32, or 64 bitwidth, but got ")
+ << dstBitWidth;
+
+ if (srcBitWidth == 6)
+ return emitOpError("currently doesn't support fp6 compact input type");
+
+ if ((packedKind == SubBytesPackedKind::U6_UNPACK_U8_E3M2 ||
+ packedKind == SubBytesPackedKind::U6_UNPACK_U8_E2M3) &&
+ !srcType.isInteger(8))
+ return emitOpError("input type expects `i8` with `u6_unpack_u8` packed "
+ "kind, but got ")
+ << srcType;
+
+ if (packedKind == SubBytesPackedKind::COMPACT &&
+ !llvm::isa<FloatType>(srcType))
+ return emitOpError("input type expects float type with `compact` packed "
+ "kind, but got ")
+ << srcType;
+
+ if (llvm::isa<Float8E8M0FNUType>(srcType) &&
+ !llvm::isa<BFloat16Type>(dstType) && !dstType.isF32())
+ return emitOpError("expects bf16 or f32 output type when input type is "
+ "e8m0.");
+
+ if (rnd != mlir::NVVM::FPRoundingMode::RN)
+ return emitOpError("expects RN rounding mode, but got ") << getRndAttr();
+
+ if (getRelu() && llvm::isa<BFloat16Type>(dstType))
+ return emitOpError("relu is not supported for bf16 destination");
+
+ return success();
+}
+
//===----------------------------------------------------------------------===//
// TableGen'd dialect, type, and op definitions
//===----------------------------------------------------------------------===//
diff --git a/mlir/test/Conversion/NVGPUToNVVM/nvgpu-convert-fpext.mlir b/mlir/test/Conversion/NVGPUToNVVM/nvgpu-convert-fpext.mlir
new file mode 100644
index 0000000000000..08407fe4bebf1
--- /dev/null
+++ b/mlir/test/Conversion/NVGPUToNVVM/nvgpu-convert-fpext.mlir
@@ -0,0 +1,466 @@
+// RUN: mlir-opt %s -convert-nvgpu-to-nvvm | FileCheck %s
+// RUN: mlir-opt %s -convert-nvgpu-to-nvvm -convert-vector-to-llvm | FileCheck %s --check-prefix=CHECK-E2E
+
+// Basic aligned vector extension to f16/bf16.
+
+// CHECK-LABEL: @cvt_float_e4m3fn_to_f16(
+// CHECK-SAME: %[[IN:.+]]: vector<8xf8E4M3FN>
+func.func @cvt_float_e4m3fn_to_f16(%in : vector<8xf8E4M3FN>) {
+ // CHECK: %[[CAST:.*]] = builtin.unrealized_conversion_cast %[[IN]] : vector<8xf8E4M3FN> to vector<8xi8>
+ // CHECK: llvm.bitcast %[[CAST]] : vector<8xi8> to vector<2xi32>
+ // CHECK: llvm.mlir.undef : vector<4xi32>
+ // CHECK: llvm.bitcast {{.*}} : i32 to vector<2xi16>
+ // CHECK: llvm.bitcast {{.*}} : i16 to vector<2xi8>
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK-SAME: : vector<2xi8>(f8E4M3FN)
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK-SAME: : vector<2xi8>(f8E4M3FN)
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<4xi32> to vector<8xf16>
+ %out = nvgpu.convert.fpext %in : vector<8xf8E4M3FN> to vector<8xf16>
+ return
+}
+
+// -----
+
+// CHECK-LABEL: @cvt_float_e5m2_to_f16(
+// CHECK-SAME: %[[IN:.+]]: vector<8xf8E5M2>
+func.func @cvt_float_e5m2_to_f16(%in : vector<8xf8E5M2>) {
+ // CHECK: %[[CAST:.*]] = builtin.unrealized_conversion_cast %[[IN]] : vector<8xf8E5M2> to vector<8xi8>
+ // CHECK: %[[IN_I32:.+]] = llvm.bitcast %[[CAST]] : vector<8xi8> to vector<2xi32>
+ // CHECK: %[[OUT_I32:.+]] = llvm.mlir.undef : vector<4xi32>
+ // CHECK: llvm.extractelement %[[IN_I32]]
+ // CHECK: llvm.bitcast {{.*}} : i32 to vector<2xi16>
+ // CHECK: llvm.extractelement {{.*}} : vector<2xi16>
+ // CHECK: llvm.bitcast {{.*}} : i16 to vector<2xi8>
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK-SAME: : vector<2xi8>(f8E5M2)
+ // CHECK: llvm.bitcast {{.*}} : vector<2xf16> to i32
+ // CHECK: llvm.insertelement {{.*}} : vector<4xi32>
+ // CHECK: llvm.extractelement {{.*}} : vector<2xi16>
+ // CHECK: llvm.bitcast {{.*}} : i16 to vector<2xi8>
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK-SAME: : vector<2xi8>(f8E5M2)
+ // CHECK: llvm.bitcast {{.*}} : vector<2xf16> to i32
+ // CHECK: llvm.insertelement {{.*}} : vector<4xi32>
+ // CHECK: llvm.bitcast {{.*}} : i32 to vector<2xi16>
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<4xi32> to vector<8xf16>
+ %out = nvgpu.convert.fpext %in : vector<8xf8E5M2> to vector<8xf16>
+ return
+}
+
+// -----
+
+// CHECK-LABEL: @cvt_float_e8m0_to_bf16(
+// CHECK-SAME: %[[IN:.+]]: vector<8xf8E8M0FNU>
+func.func @cvt_float_e8m0_to_bf16(%in : vector<8xf8E8M0FNU>) {
+ // CHECK: %[[CAST:.*]] = builtin.unrealized_conversion_cast %[[IN]] : vector<8xf8E8M0FNU> to vector<8xi8>
+ // CHECK: llvm.bitcast %[[CAST]] : vector<8xi8> to vector<2xi32>
+ // CHECK: llvm.mlir.undef : vector<4xi32>
+ // CHECK: llvm.bitcast {{.*}} : i32 to vector<2xi16>
+ // CHECK: llvm.bitcast {{.*}} : i16 to vector<2xi8>
+ // CHECK: nvvm.convert.f8x2.to.bf16x2
+ // CHECK-SAME: : vector<2xi8>(f8E8M0FNU)
+ // CHECK: llvm.bitcast {{.*}} : vector<2xbf16> to i32
+ // CHECK: nvvm.convert.f8x2.to.bf16x2
+ // CHECK-SAME: : vector<2xi8>(f8E8M0FNU)
+ // CHECK: nvvm.convert.f8x2.to.bf16x2
+ // CHECK: nvvm.convert.f8x2.to.bf16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<4xi32> to vector<8xbf16>
+ %out = nvgpu.convert.fpext %in : vector<8xf8E8M0FNU> to vector<8xbf16>
+ return
+}
+
+// -----
+
+// CHECK-LABEL: @cvt_float_e2m3_to_f16(
+// CHECK: %[[IN:.+]]: vector<8xi8>
+func.func @cvt_float_e2m3_to_f16(%in : vector<8xi8>) {
+ // CHECK: %[[IN_I32:.+]] = llvm.bitcast %[[IN]] : vector<8xi8> to vector<2xi32>
+ // CHECK: llvm.mlir.undef : vector<4xi32>
+ // CHECK: llvm.bitcast {{.*}} : i32 to vector<2xi16>
+ // CHECK: llvm.bitcast {{.*}} : i16 to vector<2xi8>
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK-SAME: : vector<2xi8>(f6E2M3FN)
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK-SAME: : vector<2xi8>(f6E2M3FN)
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<4xi32> to vector<8xf16>
+ %out = nvgpu.convert.fpext %in {packed_kind = #nvgpu.subbytes_packedkind<u6_unpack_u8_e2m3>} : vector<8xi8> to vector<8xf16>
+ return
+}
+
+// -----
+
+// CHECK-LABEL: @cvt_float_e3m2_to_f16(
+// CHECK: %[[IN:.+]]: vector<8xi8>
+func.func @cvt_float_e3m2_to_f16(%in : vector<8xi8>) {
+ // CHECK: %[[IN_I32:.+]] = llvm.bitcast %[[IN]] : vector<8xi8> to vector<2xi32>
+ // CHECK: %[[OUT_I32:.+]] = llvm.mlir.undef : vector<4xi32>
+ // CHECK: llvm.bitcast {{.*}} : i32 to vector<2xi16>
+ // CHECK: llvm.bitcast {{.*}} : i16 to vector<2xi8>
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK-SAME: : vector<2xi8>(f6E3M2FN)
+ // CHECK: llvm.bitcast {{.*}} : vector<2xf16> to i32
+ // CHECK: llvm.insertelement {{.*}} : vector<4xi32>
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK-SAME: : vector<2xi8>(f6E3M2FN)
+ // CHECK: llvm.insertelement {{.*}} : vector<4xi32>
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<4xi32> to vector<8xf16>
+ %out = nvgpu.convert.fpext %in {packed_kind = #nvgpu.subbytes_packedkind<u6_unpack_u8_e3m2>} : vector<8xi8> to vector<8xf16>
+ return
+}
+
+// -----
+
+// CHECK-LABEL: @cvt_float_e2m1_to_f16(
+// CHECK-SAME: %[[IN:.+]]: vector<16xf4E2M1FN>
+func.func @cvt_float_e2m1_to_f16(%in : vector<16xf4E2M1FN>) {
+ // CHECK: %[[CAST:.*]] = builtin.unrealized_conversion_cast %[[IN]] : vector<16xf4E2M1FN> to vector<16xi4>
+ // CHECK: %[[IN_I32:.+]] = llvm.bitcast %[[CAST]] : vector<16xi4> to vector<2xi32>
+ // CHECK: %[[OUT_I32:.+]] = llvm.mlir.undef : vector<8xi32>
+ // CHECK: llvm.extractelement %[[IN_I32]]
+ // CHECK: llvm.bitcast {{.*}} : i32 to vector<4xi8>
+ // CHECK: llvm.extractelement {{.*}} : vector<4xi8>
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK-SAME: : i8(f4E2M1FN)
+ // CHECK: llvm.bitcast {{.*}} : vector<2xf16> to i32
+ // CHECK: llvm.insertelement {{.*}} : vector<8xi32>
+ // CHECK: llvm.extractelement {{.*}} : vector<4xi8>
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK-SAME: : i8(f4E2M1FN)
+ // CHECK: llvm.insertelement {{.*}} : vector<8xi32>
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK: llvm.insertelement {{.*}} : vector<8xi32>
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK: llvm.insertelement {{.*}} : vector<8xi32>
+ // CHECK: llvm.bitcast {{.*}} : i32 to vector<4xi8>
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<8xi32> to vector<16xf16>
+ %out = nvgpu.convert.fpext %in : vector<16xf4E2M1FN> to vector<16xf16>
+ return
+}
+
+// Extension to f32 (two-step: narrow -> f16/bf16 -> f32).
+
+// -----
+
+// CHECK-LABEL: @fpext_e4m3fn_to_f32(
+// CHECK-SAME: %[[IN:.+]]: vector<4xf8E4M3FN>
+func.func @fpext_e4m3fn_to_f32(%in : vector<4xf8E4M3FN>) -> vector<4xf32> {
+ // CHECK: builtin.unrealized_conversion_cast %[[IN]] : vector<4xf8E4M3FN> to vector<4xi8>
+ // CHECK: llvm.bitcast {{.*}} : vector<4xi8> to vector<1xi32>
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi32> to vector<4xf16>
+ // CHECK: llvm.fpext {{.*}} : vector<4xf16> to vector<4xf32>
+ %out = nvgpu.convert.fpext %in : vector<4xf8E4M3FN> to vector<4xf32>
+ return %out : vector<4xf32>
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_e5m2_to_f32(
+// CHECK-SAME: %[[IN:.+]]: vector<4xf8E5M2>
+func.func @fpext_e5m2_to_f32(%in : vector<4xf8E5M2>) -> vector<4xf32> {
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi32> to vector<4xf16>
+ // CHECK: llvm.fpext {{.*}} : vector<4xf16> to vector<4xf32>
+ %out = nvgpu.convert.fpext %in : vector<4xf8E5M2> to vector<4xf32>
+ return %out : vector<4xf32>
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_e8m0_to_f32(
+// CHECK-SAME: %[[IN:.+]]: vector<4xf8E8M0FNU>
+func.func @fpext_e8m0_to_f32(%in : vector<4xf8E8M0FNU>) -> vector<4xf32> {
+ // CHECK: nvvm.convert.f8x2.to.bf16x2
+ // CHECK: nvvm.convert.f8x2.to.bf16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi32> to vector<4xbf16>
+ // CHECK: llvm.fpext {{.*}} : vector<4xbf16> to vector<4xf32>
+ %out = nvgpu.convert.fpext %in : vector<4xf8E8M0FNU> to vector<4xf32>
+ return %out : vector<4xf32>
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_e2m3_to_f32(
+// CHECK-SAME: %[[IN:.+]]: vector<4xi8>
+func.func @fpext_e2m3_to_f32(%in : vector<4xi8>) -> vector<4xf32> {
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi32> to vector<4xf16>
+ // CHECK: llvm.fpext {{.*}} : vector<4xf16> to vector<4xf32>
+ %out = nvgpu.convert.fpext %in {packed_kind = #nvgpu.subbytes_packedkind<u6_unpack_u8_e2m3>} : vector<4xi8> to vector<4xf32>
+ return %out : vector<4xf32>
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_e2m1_to_f32(
+// CHECK-SAME: %[[IN:.+]]: vector<8xf4E2M1FN>
+func.func @fpext_e2m1_to_f32(%in : vector<8xf4E2M1FN>) -> vector<8xf32> {
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<4xi32> to vector<8xf16>
+ // CHECK: llvm.fpext {{.*}} : vector<8xf16> to vector<8xf32>
+ %out = nvgpu.convert.fpext %in : vector<8xf4E2M1FN> to vector<8xf32>
+ return %out : vector<8xf32>
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_f16_to_f32(
+// CHECK-SAME: %[[IN:.+]]: vector<4xf16>
+func.func @fpext_f16_to_f32(%in : vector<4xf16>) -> vector<4xf32> {
+ // CHECK: llvm.fpext %[[IN]] : vector<4xf16> to vector<4xf32>
+ %out = nvgpu.convert.fpext %in : vector<4xf16> to vector<4xf32>
+ return %out : vector<4xf32>
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_bf16_to_f32(
+// CHECK-SAME: %[[IN:.+]]: vector<4xbf16>
+func.func @fpext_bf16_to_f32(%in : vector<4xbf16>) -> vector<4xf32> {
+ // CHECK: llvm.fpext %[[IN]] : vector<4xbf16> to vector<4xf32>
+ %out = nvgpu.convert.fpext %in : vector<4xbf16> to vector<4xf32>
+ return %out : vector<4xf32>
+}
+
+// Extension to f64.
+
+// CHECK-LABEL: @fpext_f16_to_f64
+func.func @fpext_f16_to_f64(%arg0: vector<4xf16>) -> vector<4xf64> {
+ // CHECK: llvm.fpext %{{.*}} : vector<4xf16> to vector<4xf64>
+ // CHECK-NOT: llvm.fpext
+ %out = nvgpu.convert.fpext %arg0 : vector<4xf16> to vector<4xf64>
+ return %out : vector<4xf64>
+}
+
+// CHECK-LABEL: @fpext_bf16_to_f64
+func.func @fpext_bf16_to_f64(%arg0: vector<4xbf16>) -> vector<4xf64> {
+ // CHECK: llvm.fpext %{{.*}} : vector<4xbf16> to vector<4xf64>
+ // CHECK-NOT: llvm.fpext
+ %out = nvgpu.convert.fpext %arg0 : vector<4xbf16> to vector<4xf64>
+ return %out : vector<4xf64>
+}
+
+// CHECK-LABEL: @fpext_f32_to_f64
+func.func @fpext_f32_to_f64(%arg0: vector<4xf32>) -> vector<4xf64> {
+ // CHECK-NOT: nvvm
+ // CHECK: llvm.fpext %{{.*}} : vector<4xf32> to vector<4xf64>
+ %out = nvgpu.convert.fpext %arg0 : vector<4xf32> to vector<4xf64>
+ return %out : vector<4xf64>
+}
+
+// CHECK-LABEL: @fpext_f8_to_f64
+func.func @fpext_f8_to_f64(%arg0: vector<4xf8E4M3FN>) -> vector<4xf64> {
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: llvm.fpext {{.*}} to {{.*}}f64
+ // CHECK-NOT: llvm.fpext {{.*}} to {{.*}}f32
+ // CHECK: vector.extract_strided_slice
+ %out = nvgpu.convert.fpext %arg0 : vector<4xf8E4M3FN> to vector<4xf64>
+ return %out : vector<4xf64>
+}
+
+// CHECK-LABEL: @fpext_f4_to_f64
+func.func @fpext_f4_to_f64(%arg0: vector<8xf4E2M1FN>) -> vector<8xf64> {
+ // CHECK: nvvm.convert.f4x2.to.f16x2
+ // CHECK: llvm.fpext {{.*}} to {{.*}}f64
+ // CHECK-NOT: llvm.fpext {{.*}} to {{.*}}f32
+ %out = nvgpu.convert.fpext %arg0 : vector<8xf4E2M1FN> to vector<8xf64>
+ return %out : vector<8xf64>
+}
+
+// CHECK-LABEL: @fpext_e2m3_to_f64
+func.func @fpext_e2m3_to_f64(%arg0: vector<4xi8>) -> vector<4xf64> {
+ // CHECK: nvvm.convert.f6x2.to.f16x2
+ // CHECK: llvm.fpext {{.*}} to {{.*}}f64
+ // CHECK-NOT: llvm.fpext {{.*}} to {{.*}}f32
+ %out = nvgpu.convert.fpext %arg0 {packed_kind = #nvgpu.subbytes_packedkind<u6_unpack_u8_e2m3>} : vector<4xi8> to vector<4xf64>
+ return %out : vector<4xf64>
+}
+
+// Scalar inputs (canonicalize: broadcast + pad + extract).
+
+// CHECK-LABEL: @fpext_scalar_f8_to_f16
+// CHECK-SAME: %[[IN:.+]]: f8E4M3FN
+func.func @fpext_scalar_f8_to_f16(%in : f8E4M3FN) -> f16 {
+ // CHECK: vector.broadcast %[[IN]] : f8E4M3FN to vector<1xf8E4M3FN>
+ // CHECK: vector.insert_strided_slice
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: vector.extract_strided_slice
+ // CHECK: vector.extract {{.*}}[0] : f16 from vector<1xf16>
+ %out = nvgpu.convert.fpext %in : f8E4M3FN to f16
+ return %out : f16
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_scalar_f8_to_f32
+// CHECK-SAME: %[[IN:.+]]: f8E4M3FN
+func.func @fpext_scalar_f8_to_f32(%in : f8E4M3FN) -> f32 {
+ // CHECK: vector.broadcast
+ // CHECK: vector.insert_strided_slice
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: llvm.fpext
+ // CHECK: vector.extract_strided_slice
+ // CHECK: vector.extract {{.*}}[0] : f32 from vector<1xf32>
+ %out = nvgpu.convert.fpext %in : f8E4M3FN to f32
+ return %out : f32
+}
+
+// Multi-rank vectors (canonicalize: shape_cast flatten).
+
+// CHECK-LABEL: @fpext_v2x4_f8_to_f16
+// CHECK-SAME: %[[IN:.+]]: vector<2x4xf8E4M3FN>
+func.func @fpext_v2x4_f8_to_f16(%in : vector<2x4xf8E4M3FN>) -> vector<2x4xf16> {
+ // CHECK: vector.shape_cast %[[IN]] : vector<2x4xf8E4M3FN> to vector<8xf8E4M3FN>
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: vector.shape_cast {{.*}} to vector<2x4xf16>
+ %out = nvgpu.convert.fpext %in : vector<2x4xf8E4M3FN> to vector<2x4xf16>
+ return %out : vector<2x4xf16>
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_v2x4_f8_to_f32
+// CHECK-SAME: %[[IN:.+]]: vector<2x4xf8E4M3FN>
+func.func @fpext_v2x4_f8_to_f32(%in : vector<2x4xf8E4M3FN>) -> vector<2x4xf32> {
+ // CHECK: vector.shape_cast {{.*}} to vector<8xf8E4M3FN>
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: llvm.fpext
+ // CHECK: vector.shape_cast {{.*}} to vector<2x4xf32>
+ %out = nvgpu.convert.fpext %in : vector<2x4xf8E4M3FN> to vector<2x4xf32>
+ return %out : vector<2x4xf32>
+}
+
+// Non-aligned 1-D vectors (canonicalize: pad via insert/extract_strided_slice).
+
+// CHECK-LABEL: @fpext_v3f8_to_v3f16
+// CHECK-SAME: %[[IN:.+]]: vector<3xf8E5M2>
+func.func @fpext_v3f8_to_v3f16(%in : vector<3xf8E5M2>) -> vector<3xf16> {
+ // CHECK: vector.insert_strided_slice %[[IN]]
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: vector.extract_strided_slice
+ %out = nvgpu.convert.fpext %in : vector<3xf8E5M2> to vector<3xf16>
+ return %out : vector<3xf16>
+}
+
+// -----
+
+// CHECK-LABEL: @fpext_v3_f8_to_f32
+// CHECK-SAME: %[[IN:.+]]: vector<3xf8E5M2>
+func.func @fpext_v3_f8_to_f32(%in : vector<3xf8E5M2>) -> vector<3xf32> {
+ // CHECK: vector.insert_strided_slice
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: llvm.fpext
+ // CHECK: vector.extract_strided_slice
+ %out = nvgpu.convert.fpext %in : vector<3xf8E5M2> to vector<3xf32>
+ return %out : vector<3xf32>
+}
+
+// Multi-rank + padding combined.
+
+// CHECK-LABEL: @fpext_v3x1_f8_to_f16
+// CHECK-SAME: %[[IN:.+]]: vector<3x1xf8E4M3FN>
+func.func @fpext_v3x1_f8_to_f16(%in : vector<3x1xf8E4M3FN>) -> vector<3x1xf16> {
+ // CHECK: vector.shape_cast %[[IN]] : vector<3x1xf8E4M3FN> to vector<3xf8E4M3FN>
+ // CHECK: vector.insert_strided_slice
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK: vector.extract_strided_slice
+ // CHECK: vector.shape_cast {{.*}} to vector<3x1xf16>
+ %out = nvgpu.convert.fpext %in : vector<3x1xf8E4M3FN> to vector<3x1xf16>
+ return %out : vector<3x1xf16>
+}
+
+// Relu attribute.
+
+// CHECK-LABEL: @fpext_f8_to_f16_relu
+func.func @fpext_f8_to_f16_relu(%in : vector<8xf8E4M3FN>) {
+ // CHECK: nvvm.convert.f8x2.to.f16x2
+ // CHECK-SAME: relu = true
+ %out = nvgpu.convert.fpext %in {relu = true}
+ : vector<8xf8E4M3FN> to vector<8xf16>
+ return
+}
+
+// End-to-end: no residual vector ops after full lowering.
+
+// CHECK-E2E-LABEL: @e2e_scalar_f8_to_f16
+// CHECK-E2E-NOT: vector.broadcast
+// CHECK-E2E-NOT: vector.insert_strided_slice
+// CHECK-E2E-NOT: vector.extract_strided_slice
+// CHECK-E2E-NOT: vector.extract
+// CHECK-E2E-NOT: vector.shape_cast
+// CHECK-E2E: nvvm.convert.f8x2.to.f16x2
+// CHECK-E2E: return
+func.func @e2e_scalar_f8_to_f16(%in : f8E4M3FN) -> f16 {
+ %out = nvgpu.convert.fpext %in : f8E4M3FN to f16
+ return %out : f16
+}
+
+// CHECK-E2E-LABEL: @e2e_v2x4_f8_to_f16
+// CHECK-E2E-NOT: vector.shape_cast
+// CHECK-E2E: nvvm.convert.f8x2.to.f16x2
+// CHECK-E2E: return
+func.func @e2e_v2x4_f8_to_f16(%in : vector<2x4xf8E4M3FN>) -> vector<2x4xf16> {
+ %out = nvgpu.convert.fpext %in : vector<2x4xf8E4M3FN> to vector<2x4xf16>
+ return %out : vector<2x4xf16>
+}
+
+// CHECK-E2E-LABEL: @e2e_v3f8_to_v3f16
+// CHECK-E2E-NOT: vector.insert_strided_slice
+// CHECK-E2E-NOT: vector.extract_strided_slice
+// CHECK-E2E: nvvm.convert.f8x2.to.f16x2
+// CHECK-E2E: return
+func.func @e2e_v3f8_to_v3f16(%in : vector<3xf8E5M2>) -> vector<3xf16> {
+ %out = nvgpu.convert.fpext %in : vector<3xf8E5M2> to vector<3xf16>
+ return %out : vector<3xf16>
+}
+
+// CHECK-E2E-LABEL: @e2e_scalar_f8_to_f32
+// CHECK-E2E-NOT: vector.broadcast
+// CHECK-E2E-NOT: vector.insert_strided_slice
+// CHECK-E2E-NOT: vector.extract_strided_slice
+// CHECK-E2E-NOT: vector.extract
+// CHECK-E2E-NOT: vector.shape_cast
+// CHECK-E2E: nvvm.convert.f8x2.to.f16x2
+// CHECK-E2E: llvm.fpext
+// CHECK-E2E: return
+func.func @e2e_scalar_f8_to_f32(%in : f8E4M3FN) -> f32 {
+ %out = nvgpu.convert.fpext %in : f8E4M3FN to f32
+ return %out : f32
+}
+
+// CHECK-E2E-LABEL: @e2e_v2x4_f8_to_f32
+// CHECK-E2E-NOT: vector.shape_cast
+// CHECK-E2E: nvvm.convert.f8x2.to.f16x2
+// CHECK-E2E: llvm.fpext
+// CHECK-E2E: return
+func.func @e2e_v2x4_f8_to_f32(%in : vector<2x4xf8E4M3FN>) -> vector<2x4xf32> {
+ %out = nvgpu.convert.fpext %in : vector<2x4xf8E4M3FN> to vector<2x4xf32>
+ return %out : vector<2x4xf32>
+}
+
+// CHECK-E2E-LABEL: @e2e_v3f8_to_v3f32
+// CHECK-E2E-NOT: vector.insert_strided_slice
+// CHECK-E2E-NOT: vector.extract_strided_slice
+// CHECK-E2E: nvvm.convert.f8x2.to.f16x2
+// CHECK-E2E: llvm.fpext
+// CHECK-E2E: return
+func.func @e2e_v3f8_to_v3f32(%in : vector<3xf8E5M2>) -> vector<3xf32> {
+ %out = nvgpu.convert.fpext %in : vector<3xf8E5M2> to vector<3xf32>
+ return %out : vector<3xf32>
+}
diff --git a/mlir/test/Conversion/NVGPUToNVVM/nvgpu-convert-fptrunc.mlir b/mlir/test/Conversion/NVGPUToNVVM/nvgpu-convert-fptrunc.mlir
new file mode 100644
index 0000000000000..e8a5c27f22917
--- /dev/null
+++ b/mlir/test/Conversion/NVGPUToNVVM/nvgpu-convert-fptrunc.mlir
@@ -0,0 +1,327 @@
+// RUN: mlir-opt %s -convert-nvgpu-to-nvvm | FileCheck %s
+// RUN: mlir-opt %s -convert-nvgpu-to-nvvm -convert-vector-to-llvm | FileCheck %s --check-prefix=CHECK-E2E
+
+// Basic aligned vector inputs.
+
+// CHECK-LABEL: @cvt_float_f32_to_f16(
+// CHECK-SAME: %[[IN:.+]]: vector<4xf32>
+func.func @cvt_float_f32_to_f16(%in : vector<4xf32>) {
+ // CHECK: llvm.bitcast %[[IN]] : vector<4xf32> to vector<4xi32>
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK-SAME: rnd = #nvvm.fp_rnd_mode<rn>
+ // CHECK-SAME: : vector<2xf16>
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK-SAME: rnd = #nvvm.fp_rnd_mode<rn>
+ // CHECK-SAME: : vector<2xf16>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi32> to vector<4xf16>
+ %out = nvgpu.convert.fptrunc %in : vector<4xf32> to vector<4xf16>
+ return
+}
+
+// CHECK-LABEL: @cvt_float_f32_to_f16_v8(
+// CHECK-SAME: %[[IN:.+]]: vector<8xf32>
+func.func @cvt_float_f32_to_f16_v8(%in : vector<8xf32>) {
+ // CHECK: llvm.bitcast %[[IN]] : vector<8xf32> to vector<8xi32>
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<4xi32> to vector<8xf16>
+ %out = nvgpu.convert.fptrunc %in : vector<8xf32> to vector<8xf16>
+ return
+}
+
+// CHECK-LABEL: @cvt_float_f32_to_bf16(
+// CHECK-SAME: %[[IN:.+]]: vector<4xf32>
+func.func @cvt_float_f32_to_bf16(%in : vector<4xf32>) {
+ // CHECK: llvm.bitcast %[[IN]] : vector<4xf32> to vector<4xi32>
+ // CHECK: nvvm.convert.f32x2.to.bf16x2
+ // CHECK-SAME: rnd = #nvvm.fp_rnd_mode<rn>
+ // CHECK-SAME: : vector<2xbf16>
+ // CHECK: nvvm.convert.f32x2.to.bf16x2
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi32> to vector<4xbf16>
+ %out = nvgpu.convert.fptrunc %in : vector<4xf32> to vector<4xbf16>
+ return
+}
+
+// CHECK-LABEL: @cvt_float_f32_to_e4m3(
+// CHECK-SAME: %[[IN:.+]]: vector<8xf32>
+func.func @cvt_float_f32_to_e4m3(%in : vector<8xf32>) {
+ // CHECK: %[[IN_I32:.+]] = llvm.bitcast %[[IN]] : vector<8xf32> to vector<8xi32>
+ // CHECK: %[[OUT_I32:.+]] = llvm.mlir.undef : vector<2xi32>
+ // CHECK: %[[IDX_0:.+]] = llvm.mlir.constant(0 : i64) : i64
+ // CHECK: %[[SUB_VEC_0:.+]] = llvm.mlir.undef : vector<2xi16>
+ // CHECK: nvvm.convert.f32x2.to.f8x2
+ // CHECK-SAME: sat = #nvvm.sat_mode<satfinite>
+ // CHECK-SAME: : i16(f8E4M3FN)
+ // CHECK: llvm.insertelement {{.*}} : vector<2xi16>
+ // CHECK: nvvm.convert.f32x2.to.f8x2
+ // CHECK-SAME: : i16(f8E4M3FN)
+ // CHECK: llvm.insertelement {{.*}} : vector<2xi16>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi16> to i32
+ // CHECK: llvm.insertelement {{.*}} : vector<2xi32>
+ // CHECK: llvm.mlir.undef : vector<2xi16>
+ // CHECK: nvvm.convert.f32x2.to.f8x2
+ // CHECK: nvvm.convert.f32x2.to.f8x2
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi16> to i32
+ // CHECK: llvm.insertelement {{.*}} : vector<2xi32>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi32> to vector<8xi8>
+ %out = nvgpu.convert.fptrunc %in : vector<8xf32> to vector<8xf8E4M3FN>
+ return
+}
+
+// CHECK-LABEL: @cvt_float_f32_to_e2m3(
+// CHECK-SAME: %[[IN:.+]]: vector<8xf32>
+func.func @cvt_float_f32_to_e2m3(%in : vector<8xf32>) {
+ // CHECK: llvm.bitcast %[[IN]] : vector<8xf32> to vector<8xi32>
+ // CHECK: llvm.mlir.undef : vector<2xi32>
+ // CHECK: llvm.mlir.undef : vector<2xi16>
+ // CHECK: nvvm.convert.f32x2.to.f6x2
+ // CHECK-SAME: : i16(f6E2M3FN)
+ // CHECK: llvm.insertelement {{.*}} : vector<2xi16>
+ // CHECK: nvvm.convert.f32x2.to.f6x2
+ // CHECK-SAME: : i16(f6E2M3FN)
+ // CHECK: llvm.insertelement {{.*}} : vector<2xi16>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi16> to i32
+ // CHECK: llvm.insertelement {{.*}} : vector<2xi32>
+ // CHECK: llvm.mlir.undef : vector<2xi16>
+ // CHECK: nvvm.convert.f32x2.to.f6x2
+ // CHECK: nvvm.convert.f32x2.to.f6x2
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi16> to i32
+ // CHECK: llvm.insertelement {{.*}} : vector<2xi32>
+ // CHECK: llvm.bitcast {{.*}} : vector<2xi32> to vector<8xi8>
+ %out = nvgpu.convert.fptrunc %in {packed_kind = #nvgpu.subbytes_packedkind<u6_unpack_u8_e2m3>}: vector<8xf32> to vector<8xi8>
+ return
+}
+
+// Scalar inputs (canonicalize: broadcast + pad + extract).
+
+// CHECK-LABEL: @fptrunc_scalar_f32_to_f16
+// CHECK-SAME: %[[IN:.+]]: f32
+func.func @fptrunc_scalar_f32_to_f16(%in : f32) -> f16 {
+ // CHECK: vector.broadcast %[[IN]] : f32 to vector<1xf32>
+ // CHECK: vector.insert_strided_slice
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: vector.extract_strided_slice
+ // CHECK: vector.extract {{.*}}[0] : f16 from vector<1xf16>
+ %out = nvgpu.convert.fptrunc %in : f32 to f16
+ return %out : f16
+}
+
+// CHECK-LABEL: @fptrunc_scalar_f32_to_bf16
+// CHECK-SAME: %[[IN:.+]]: f32
+func.func @fptrunc_scalar_f32_to_bf16(%in : f32) -> bf16 {
+ // CHECK: vector.broadcast %[[IN]]
+ // CHECK: nvvm.convert.f32x2.to.bf16x2
+ // CHECK: vector.extract
+ %out = nvgpu.convert.fptrunc %in : f32 to bf16
+ return %out : bf16
+}
+
+// CHECK-LABEL: @fptrunc_scalar_f64_to_f32
+func.func @fptrunc_scalar_f64_to_f32(%arg0: f64) -> f32 {
+ // CHECK: vector.broadcast
+ // CHECK: llvm.fptrunc
+ // CHECK: vector.extract
+ %out = nvgpu.convert.fptrunc %arg0 : f64 to f32
+ return %out : f32
+}
+
+// Multi-rank vectors (canonicalize: shape_cast flatten).
+
+// CHECK-LABEL: @fptrunc_v2x4_f32_to_f16
+// CHECK-SAME: %[[IN:.+]]: vector<2x4xf32>
+func.func @fptrunc_v2x4_f32_to_f16(%in : vector<2x4xf32>) -> vector<2x4xf16> {
+ // CHECK: vector.shape_cast %[[IN]] : vector<2x4xf32> to vector<8xf32>
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: vector.shape_cast {{.*}} to vector<2x4xf16>
+ %out = nvgpu.convert.fptrunc %in : vector<2x4xf32> to vector<2x4xf16>
+ return %out : vector<2x4xf16>
+}
+
+// CHECK-LABEL: @fptrunc_v4x2_f32_to_f8
+// CHECK-SAME: %[[IN:.+]]: vector<4x2xf32>
+func.func @fptrunc_v4x2_f32_to_f8(%in : vector<4x2xf32>) -> vector<4x2xf8E4M3FN> {
+ // CHECK: vector.shape_cast %[[IN]] : vector<4x2xf32> to vector<8xf32>
+ // CHECK: nvvm.convert.f32x2.to.f8x2
+ // CHECK: vector.shape_cast {{.*}} to vector<4x2xf8E4M3FN>
+ %out = nvgpu.convert.fptrunc %in : vector<4x2xf32> to vector<4x2xf8E4M3FN>
+ return %out : vector<4x2xf8E4M3FN>
+}
+
+// Non-aligned 1-D vectors (canonicalize: pad via insert/extract_strided_slice).
+
+// CHECK-LABEL: @fptrunc_v1f32_to_v1f16
+// CHECK-SAME: %[[IN:.+]]: vector<1xf32>
+func.func @fptrunc_v1f32_to_v1f16(%in : vector<1xf32>) -> vector<1xf16> {
+ // CHECK: vector.insert_strided_slice %[[IN]]
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: vector.extract_strided_slice
+ %out = nvgpu.convert.fptrunc %in : vector<1xf32> to vector<1xf16>
+ return %out : vector<1xf16>
+}
+
+// CHECK-LABEL: @fptrunc_v3f16_to_v3f8
+// CHECK-SAME: %[[IN:.+]]: vector<3xf16>
+func.func @fptrunc_v3f16_to_v3f8(%in : vector<3xf16>) -> vector<3xf8E4M3FN> {
+ // CHECK: vector.insert_strided_slice %[[IN]]
+ // CHECK: nvvm.convert.f16x2.to.f8x2
+ // CHECK: vector.extract_strided_slice
+ %out = nvgpu.convert.fptrunc %in : vector<3xf16> to vector<3xf8E4M3FN>
+ return %out : vector<3xf8E4M3FN>
+}
+
+// Multi-rank + padding combined.
+
+// CHECK-LABEL: @fptrunc_v3x1_f32_to_f16
+// CHECK-SAME: %[[IN:.+]]: vector<3x1xf32>
+func.func @fptrunc_v3x1_f32_to_f16(%in : vector<3x1xf32>) -> vector<3x1xf16> {
+ // CHECK: vector.shape_cast %[[IN]] : vector<3x1xf32> to vector<3xf32>
+ // CHECK: vector.insert_strided_slice
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: vector.extract_strided_slice
+ // CHECK: vector.shape_cast {{.*}} to vector<3x1xf16>
+ %out = nvgpu.convert.fptrunc %in : vector<3x1xf32> to vector<3x1xf16>
+ return %out : vector<3x1xf16>
+}
+
+// f64 source truncation.
+
+// CHECK-LABEL: @fptrunc_f64_to_f32
+func.func @fptrunc_f64_to_f32(%arg0: vector<4xf64>) -> vector<4xf32> {
+ // CHECK: llvm.fptrunc %{{.*}} : vector<4xf64> to vector<4xf32>
+ // CHECK-NOT: nvvm
+ %out = nvgpu.convert.fptrunc %arg0 : vector<4xf64> to vector<4xf32>
+ return %out : vector<4xf32>
+}
+
+// CHECK-LABEL: @fptrunc_f64_to_f16
+func.func @fptrunc_f64_to_f16(%arg0: vector<2xf64>) -> vector<2xf16> {
+ // CHECK: vector.insert_strided_slice
+ // CHECK: llvm.fptrunc %{{.*}} : vector<4xf64> to vector<4xf32>
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK: vector.extract_strided_slice
+ %out = nvgpu.convert.fptrunc %arg0 : vector<2xf64> to vector<2xf16>
+ return %out : vector<2xf16>
+}
+
+// CHECK-LABEL: @fptrunc_f64_to_bf16
+func.func @fptrunc_f64_to_bf16(%arg0: vector<2xf64>) -> vector<2xbf16> {
+ // CHECK: vector.insert_strided_slice
+ // CHECK: llvm.fptrunc %{{.*}} : vector<4xf64> to vector<4xf32>
+ // CHECK: nvvm.convert.f32x2.to.bf16x2
+ // CHECK: vector.extract_strided_slice
+ %out = nvgpu.convert.fptrunc %arg0 : vector<2xf64> to vector<2xbf16>
+ return %out : vector<2xbf16>
+}
+
+// CHECK-LABEL: @fptrunc_f64_to_f8
+func.func @fptrunc_f64_to_f8(%arg0: vector<4xf64>) -> vector<4xf8E4M3FN> {
+ // CHECK: vector.insert_strided_slice
+ // CHECK: llvm.fptrunc %{{.*}} : vector<8xf64> to vector<8xf32>
+ // CHECK: nvvm.convert.f32x2.to.f8x2
+ // CHECK: vector.extract_strided_slice
+ %out = nvgpu.convert.fptrunc %arg0 : vector<4xf64> to vector<4xf8E4M3FN>
+ return %out : vector<4xf8E4M3FN>
+}
+
+// Saturation and relu attributes.
+
+// CHECK-LABEL: @fptrunc_f32_to_f8_satfinite
+func.func @fptrunc_f32_to_f8_satfinite(%in : vector<8xf32>) {
+ // CHECK: nvvm.convert.f32x2.to.f8x2
+ // CHECK-SAME: sat = #nvvm.sat_mode<satfinite>
+ %out = nvgpu.convert.fptrunc %in {sat = #nvvm.sat_mode<satfinite>}
+ : vector<8xf32> to vector<8xf8E4M3FN>
+ return
+}
+
+// CHECK-LABEL: @fptrunc_f32_to_f16_relu
+func.func @fptrunc_f32_to_f16_relu(%in : vector<4xf32>) {
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK-SAME: relu = true
+ %out = nvgpu.convert.fptrunc %in {relu = true}
+ : vector<4xf32> to vector<4xf16>
+ return
+}
+
+// Default SATFINITE behavior.
+
+// CHECK-LABEL: @fptrunc_f32_to_f8_default_sat
+func.func @fptrunc_f32_to_f8_default_sat(%in : vector<8xf32>) {
+ // CHECK: nvvm.convert.f32x2.to.f8x2
+ // CHECK-SAME: sat = #nvvm.sat_mode<satfinite>
+ %out = nvgpu.convert.fptrunc %in : vector<8xf32> to vector<8xf8E4M3FN>
+ return
+}
+
+// CHECK-LABEL: @fptrunc_f32_to_f16_default_sat
+func.func @fptrunc_f32_to_f16_default_sat(%in : vector<4xf32>) {
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK-SAME: sat = #nvvm.sat_mode<satfinite>
+ %out = nvgpu.convert.fptrunc %in : vector<4xf32> to vector<4xf16>
+ return
+}
+
+// CHECK-LABEL: @fptrunc_f32_to_f16_explicit_none
+func.func @fptrunc_f32_to_f16_explicit_none(%in : vector<4xf32>) {
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK-NOT: satfinite
+ %out = nvgpu.convert.fptrunc %in {sat = #nvvm.sat_mode<none>}
+ : vector<4xf32> to vector<4xf16>
+ return
+}
+
+// Stochastic rounding (RS + random_bits).
+
+// CHECK-LABEL: @fptrunc_f32_to_f16_rs
+func.func @fptrunc_f32_to_f16_rs(%in : vector<4xf32>, %rbits : i32) {
+ // CHECK: nvvm.convert.f32x2.to.f16x2
+ // CHECK-SAME: rnd = #nvvm.fp_rnd_mode<rs>
+ %out = nvgpu.convert.fptrunc %in, %rbits {rnd = #nvvm.fp_rnd_mode<rs>}
+ : vector<4xf32> to vector<4xf16>
+ return
+}
+
+// CHECK-LABEL: @fptrunc_f32_to_bf16_rs
+func.func @fptrunc_f32_to_bf16_rs(%in : vector<4xf32>, %rbits : i32) {
+ // CHECK: nvvm.convert.f32x2.to.bf16x2
+ // CHECK-SAME: rnd = #nvvm.fp_rnd_mode<rs>
+ %out = nvgpu.convert.fptrunc %in, %rbits {rnd = #nvvm.fp_rnd_mode<rs>}
+ : vector<4xf32> to vector<4xbf16>
+ return
+}
+
+// End-to-end: no residual vector ops after full lowering.
+
+// CHECK-E2E-LABEL: @e2e_scalar_f32_to_f16
+// CHECK-E2E-NOT: vector.broadcast
+// CHECK-E2E-NOT: vector.insert_strided_slice
+// CHECK-E2E-NOT: vector.extract_strided_slice
+// CHECK-E2E-NOT: vector.extract
+// CHECK-E2E-NOT: vector.shape_cast
+// CHECK-E2E: nvvm.convert.f32x2.to.f16x2
+// CHECK-E2E: return
+func.func @e2e_scalar_f32_to_f16(%in : f32) -> f16 {
+ %out = nvgpu.convert.fptrunc %in : f32 to f16
+ return %out : f16
+}
+
+// CHECK-E2E-LABEL: @e2e_v2x4_f32_to_f16
+// CHECK-E2E-NOT: vector.shape_cast
+// CHECK-E2E: nvvm.convert.f32x2.to.f16x2
+// CHECK-E2E: return
+func.func @e2e_v2x4_f32_to_f16(%in : vector<2x4xf32>) -> vector<2x4xf16> {
+ %out = nvgpu.convert.fptrunc %in : vector<2x4xf32> to vector<2x4xf16>
+ return %out : vector<2x4xf16>
+}
+
+// CHECK-E2E-LABEL: @e2e_v3f16_to_v3f8
+// CHECK-E2E-NOT: vector.insert_strided_slice
+// CHECK-E2E-NOT: vector.extract_strided_slice
+// CHECK-E2E: nvvm.convert.f16x2.to.f8x2
+// CHECK-E2E: return
+func.func @e2e_v3f16_to_v3f8(%in : vector<3xf16>) -> vector<3xf8E4M3FN> {
+ %out = nvgpu.convert.fptrunc %in : vector<3xf16> to vector<3xf8E4M3FN>
+ return %out : vector<3xf8E4M3FN>
+}
diff --git a/mlir/test/Dialect/NVGPU/nvgpu-convert-fpext-invalid.mlir b/mlir/test/Dialect/NVGPU/nvgpu-convert-fpext-invalid.mlir
new file mode 100644
index 0000000000000..a2f8fd6ec48f8
--- /dev/null
+++ b/mlir/test/Dialect/NVGPU/nvgpu-convert-fpext-invalid.mlir
@@ -0,0 +1,106 @@
+// RUN: mlir-opt -split-input-file -verify-diagnostics %s
+
+// -----
+
+func.func @fpext_wider(%in : vector<16xf8E5M2>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op result type 'f4E2M1FN' must be wider than operand type 'f8E5M2'}}
+ %out = nvgpu.convert.fpext %in : vector<16xf8E5M2> to vector<16xf4E2M1FN>
+ return
+}
+
+// -----
+
+func.func @fpext_dst_bitwidth(%in : vector<16xf4E2M1FN>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op result type must be 16, 32, or 64 bitwidth, but got 8}}
+ %out = nvgpu.convert.fpext %in : vector<16xf4E2M1FN> to vector<16xf8E4M3FN>
+ return
+}
+
+// -----
+
+func.func @fpext_fp6_compact(%in : vector<16xf6E2M3FN>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op currently doesn't support fp6 compact input type}}
+ %out = nvgpu.convert.fpext %in : vector<16xf6E2M3FN> to vector<16xf16>
+ return
+}
+
+// -----
+
+func.func @fpext_u6_packed_not_i8(%in : vector<16xf8E4M3FN>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op input type expects `i8` with `u6_unpack_u8` packed kind, but got 'f8E4M3FN'}}
+ %out = nvgpu.convert.fpext %in {packed_kind = #nvgpu.subbytes_packedkind<u6_unpack_u8_e3m2>} : vector<16xf8E4M3FN> to vector<16xf16>
+ return
+}
+
+// -----
+
+func.func @fpext_compact_not_float(%in : vector<16xi8>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op input type expects float type with `compact` packed kind, but got 'i8'}}
+ %out = nvgpu.convert.fpext %in : vector<16xi8> to vector<16xf16>
+ return
+}
+
+// -----
+
+func.func @fpext_e8m0_to_f16(%in : vector<16xf8E8M0FNU>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op expects bf16 or f32 output type when input type is e8m0.}}
+ %out = nvgpu.convert.fpext %in : vector<16xf8E8M0FNU> to vector<16xf16>
+ return
+}
+
+// -----
+
+func.func @fpext_bad_rounding(%in : vector<16xf8E5M2>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op expects RN rounding mode, but got #nvvm.fp_rnd_mode<rz>}}
+ %out = nvgpu.convert.fpext %in {rnd = #nvvm.fp_rnd_mode<rz>}
+ : vector<16xf8E5M2> to vector<16xf16>
+ return
+}
+
+// -----
+
+func.func @fpext_relu_bf16(%in : vector<8xf8E5M2>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op relu is not supported for bf16 destination}}
+ %out = nvgpu.convert.fpext %in {relu = true} : vector<8xf8E5M2> to vector<8xbf16>
+ return
+}
+
+// -----
+
+func.func @fpext_shape_mismatch(%in : vector<16xf8E5M2>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op input and output shapes must match}}
+ %out = nvgpu.convert.fpext %in : vector<16xf8E5M2> to vector<8xf16>
+ return
+}
+
+// -----
+
+func.func @fpext_scalar_vector_mismatch(%in : f8E4M3FN) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op input and output must both be scalars or both be vectors/tensors}}
+ %out = nvgpu.convert.fpext %in : f8E4M3FN to vector<1xf16>
+ return
+}
+
+// -----
+
+func.func @fpext_rank0_tensor(%in : tensor<f8E4M3FN>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op rank-0 shaped types are not supported, use scalar type instead}}
+ %out = nvgpu.convert.fpext %in : tensor<f8E4M3FN> to tensor<f16>
+ return
+}
+
+// -----
+
+func.func @fpext_container_mismatch(%in : vector<4xf8E4M3FN>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op input and output must be the same container type (both vector or both tensor)}}
+ %out = nvgpu.convert.fpext %in : vector<4xf8E4M3FN> to tensor<4xf16>
+ return
+}
+
+// -----
+
+func.func @fpext_unranked_tensor(%in : tensor<*xf8E4M3FN>) {
+ // expected-error @+1 {{'nvgpu.convert.fpext' op unranked tensor types are not supported}}
+ %out = nvgpu.convert.fpext %in : tensor<*xf8E4M3FN> to tensor<*xf16>
+ return
+}
diff --git a/mlir/test/Dialect/NVGPU/nvgpu-convert-fptrunc-invalid.mlir b/mlir/test/Dialect/NVGPU/nvgpu-convert-fptrunc-invalid.mlir
new file mode 100644
index 0000000000000..9e295b4ad25a0
--- /dev/null
+++ b/mlir/test/Dialect/NVGPU/nvgpu-convert-fptrunc-invalid.mlir
@@ -0,0 +1,126 @@
+// RUN: mlir-opt -split-input-file -verify-diagnostics %s
+
+// -----
+
+func.func @fptrunc_narrower(%in : vector<16xf16>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op result type 'f32' must be narrower than operand type 'f16'}}
+ %out = nvgpu.convert.fptrunc %in : vector<16xf16> to vector<16xf32>
+ return
+}
+
+// -----
+
+func.func @fptrunc_src_bitwidth(%in : vector<16xf8E5M2>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op input type must be 64/32/16 bitwidth, but got 8}}
+ %out = nvgpu.convert.fptrunc %in : vector<16xf8E5M2> to vector<16xf4E2M1FN>
+ return
+}
+
+// -----
+
+func.func @fptrunc_fp6_compact(%in : vector<16xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op currently doesn't support fp6 compact result type}}
+ %out = nvgpu.convert.fptrunc %in : vector<16xf32> to vector<16xf6E2M3FN>
+ return
+}
+
+// -----
+
+func.func @fptrunc_u6_packed_not_i8(%in : vector<16xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op result type expects `i8` with `u6_unpack_u8` packed kind, but got 'f8E4M3FN'}}
+ %out = nvgpu.convert.fptrunc %in {packed_kind = #nvgpu.subbytes_packedkind<u6_unpack_u8_e3m2>} : vector<16xf32> to vector<16xf8E4M3FN>
+ return
+}
+
+// -----
+
+func.func @fptrunc_compact_not_float(%in : vector<16xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op result type expects float type with `compact` packed kind, but got 'i8'}}
+ %out = nvgpu.convert.fptrunc %in : vector<16xf32> to vector<16xi8>
+ return
+}
+
+// -----
+
+func.func @fptrunc_e8m0_bad_rounding(%in : vector<16xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op expects RZ or RP rounding mode when result type is e8m0, but got #nvvm.fp_rnd_mode<rn>}}
+ %out = nvgpu.convert.fptrunc %in {rnd = #nvvm.fp_rnd_mode<rn>}
+ : vector<16xf32> to vector<16xf8E8M0FNU>
+ return
+}
+
+// -----
+
+func.func @fptrunc_rs_unsupported_types(%in : vector<16xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op RS (stochastic) rounding is only supported for f32->f16/bf16, got 'f32' -> 'f8E4M3FN'}}
+ %out = nvgpu.convert.fptrunc %in {rnd = #nvvm.fp_rnd_mode<rs>}
+ : vector<16xf32> to vector<16xf8E4M3FN>
+ return
+}
+
+// -----
+
+func.func @fptrunc_rs_no_random_bits(%in : vector<4xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op random_bits operand is required with RS rounding}}
+ %out = nvgpu.convert.fptrunc %in {rnd = #nvvm.fp_rnd_mode<rs>}
+ : vector<4xf32> to vector<4xf16>
+ return
+}
+
+// -----
+
+func.func @fptrunc_bad_rounding(%in : vector<16xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op expects RN rounding mode, but got #nvvm.fp_rnd_mode<rp>}}
+ %out = nvgpu.convert.fptrunc %in {rnd = #nvvm.fp_rnd_mode<rp>}
+ : vector<16xf32> to vector<16xf8E4M3FN>
+ return
+}
+
+// -----
+
+func.func @fptrunc_random_bits_no_rs(%in : vector<4xf32>, %rbits : i32) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op random_bits can only be used with RS rounding mode}}
+ %out = nvgpu.convert.fptrunc %in, %rbits
+ : vector<4xf32> to vector<4xf16>
+ return
+}
+
+// -----
+
+func.func @fptrunc_shape_mismatch(%in : vector<16xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op input and output shapes must match}}
+ %out = nvgpu.convert.fptrunc %in : vector<16xf32> to vector<8xf8E4M3FN>
+ return
+}
+
+// -----
+
+func.func @fptrunc_scalar_vector_mismatch(%in : f32) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op input and output must both be scalars or both be vectors/tensors}}
+ %out = nvgpu.convert.fptrunc %in : f32 to vector<1xf16>
+ return
+}
+
+// -----
+
+func.func @fptrunc_rank0_tensor(%in : tensor<f32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op rank-0 shaped types are not supported, use scalar type instead}}
+ %out = nvgpu.convert.fptrunc %in : tensor<f32> to tensor<f16>
+ return
+}
+
+// -----
+
+func.func @fptrunc_container_mismatch(%in : vector<4xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op input and output must be the same container type (both vector or both tensor)}}
+ %out = nvgpu.convert.fptrunc %in : vector<4xf32> to tensor<4xf16>
+ return
+}
+
+// -----
+
+func.func @fptrunc_unranked_tensor(%in : tensor<*xf32>) {
+ // expected-error @+1 {{'nvgpu.convert.fptrunc' op unranked tensor types are not supported}}
+ %out = nvgpu.convert.fptrunc %in : tensor<*xf32> to tensor<*xf16>
+ return
+}
>From 01a01f6bb714850ece708d2d329f89a17d251610 Mon Sep 17 00:00:00 2001
From: Srinivasa Ravi <srinivasar at nvidia.com>
Date: Tue, 26 May 2026 14:49:27 +0000
Subject: [PATCH 2/2] fix formatting
---
mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp | 6 +++---
1 file changed, 3 insertions(+), 3 deletions(-)
diff --git a/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp b/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
index 8b992b67aed29..08cbefee015c2 100644
--- a/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
+++ b/mlir/lib/Conversion/NVGPUToNVVM/NVGPUToNVVM.cpp
@@ -2289,9 +2289,9 @@ static int64_t computePaddedElems(int64_t numElems, int srcBW, int dstBW,
return ceilDiv(padded, step) * step;
}
-/// Canonicalization pattern for nvgpu.convert.fptrunc / nvgpu.convert.fpext: handles
-/// scalar inputs, non-32-bit-aligned vectors, and multi-rank vectors. Runs as
-/// an OpRewritePattern on MLIR types before LLVM type conversion.
+/// Canonicalization pattern for nvgpu.convert.fptrunc / nvgpu.convert.fpext:
+/// handles scalar inputs, non-32-bit-aligned vectors, and multi-rank vectors.
+/// Runs as an OpRewritePattern on MLIR types before LLVM type conversion.
template <typename CvtOp, bool IsTrunc>
struct NVGPUFPCanonicalizePattern : public OpRewritePattern<CvtOp> {
using OpRewritePattern<CvtOp>::OpRewritePattern;
More information about the Mlir-commits
mailing list