[Mlir-commits] [mlir] [MLIR][NVGPU] Add convert.fpext and convert.fptrunc Ops (PR #199700)
Durgadoss R
llvmlistbot at llvm.org
Tue Jun 30 01:07:13 PDT 2026
================
@@ -1709,6 +1711,685 @@ struct NVGPURcpOpLowering : public ConvertOpToLLVMPattern<nvgpu::RcpOp> {
rewriter);
}
};
+
+//===----------------------------------------------------------------------===//
+// NVGPUConvertFPTruncOp Lowering
+//===----------------------------------------------------------------------===//
+
+enum class FPKind { F32, BF16, F16, F8, F6, F4 };
+
+static int getEffectiveBitWidth(int bitWidth) {
+ // f6 types are 6-bit but NVVM Ops expect 8-bit (i8) containers.
+ return bitWidth == 6 ? 8 : bitWidth;
+}
+
+static std::optional<FPKind> classifyFPType(Type t) {
+ static constexpr auto isConvertibleF8Type = [](Type t) {
+ return isa<Float8E4M3FNType, Float8E5M2Type, Float8E8M0FNUType>(t);
+ };
+ static constexpr auto isConvertibleF6Type = [](Type t) {
+ return isa<Float6E2M3FNType, Float6E3M2FNType>(t);
+ };
+ static constexpr auto isConvertibleF4Type = [](Type t) {
+ return isa<Float4E2M1FNType>(t);
+ };
+
+ if (t.isF32())
+ return FPKind::F32;
+ if (t.isBF16())
+ return FPKind::BF16;
+ if (t.isF16())
+ return FPKind::F16;
+ if (isConvertibleF8Type(t))
+ return FPKind::F8;
+ if (isConvertibleF6Type(t))
+ return FPKind::F6;
+ if (isConvertibleF4Type(t))
+ return FPKind::F4;
+ return std::nullopt;
+}
+
+/// Number of source-side i32 register slots consumed by each NVVM convert Op.
+static int getNumSrcI32PerConv(FPKind src) {
+ return src == FPKind::F32 ? 2 : 1;
+}
+
+/// Conversion op identifier for nvgpu.convert.fptrunc lowering dispatch table.
+enum class FPTruncConvOp {
+ F32x2_TO_F16x2,
+ F32x2_TO_BF16x2,
+ F32x2_TO_F8x2,
+ F32x2_TO_F6x2,
+ F32x2_TO_F4x2,
+ F16x2_TO_F8x2,
+ F16x2_TO_F6x2,
+ F16x2_TO_F4x2,
+ BF16x2_TO_F8x2,
+ BF16x2_TO_F6x2,
+ BF16x2_TO_F4x2,
+};
+
+struct FPTruncTableEntry {
+ FPKind src;
+ FPKind dst;
+ FPTruncConvOp convOp;
+};
+
+static constexpr FPTruncTableEntry kFPTruncTable[] = {
+ // f32 source
+ {FPKind::F32, FPKind::F16, FPTruncConvOp::F32x2_TO_F16x2},
+ {FPKind::F32, FPKind::BF16, FPTruncConvOp::F32x2_TO_BF16x2},
+ {FPKind::F32, FPKind::F8, FPTruncConvOp::F32x2_TO_F8x2},
+ {FPKind::F32, FPKind::F6, FPTruncConvOp::F32x2_TO_F6x2},
+ {FPKind::F32, FPKind::F4, FPTruncConvOp::F32x2_TO_F4x2},
+ // f16 source
+ {FPKind::F16, FPKind::F8, FPTruncConvOp::F16x2_TO_F8x2},
+ {FPKind::F16, FPKind::F6, FPTruncConvOp::F16x2_TO_F6x2},
+ {FPKind::F16, FPKind::F4, FPTruncConvOp::F16x2_TO_F4x2},
+ // bf16 source
+ {FPKind::BF16, FPKind::F8, FPTruncConvOp::BF16x2_TO_F8x2},
+ {FPKind::BF16, FPKind::F6, FPTruncConvOp::BF16x2_TO_F6x2},
+ {FPKind::BF16, FPKind::F4, FPTruncConvOp::BF16x2_TO_F4x2},
+};
+
+static std::optional<FPTruncTableEntry> lookupTruncConvOp(Type srcElemType,
+ Type dstElemType) {
+ auto srcKind = classifyFPType(srcElemType);
+ auto dstKind = classifyFPType(dstElemType);
+ if (!srcKind || !dstKind)
+ return std::nullopt;
+ for (const auto &entry : kFPTruncTable) {
+ if (entry.src == *srcKind && entry.dst == *dstKind)
+ return entry;
+ }
+ return std::nullopt;
+}
+
+/// Extract a single element from a vector.
+static Value extractElement(ImplicitLocOpBuilder &b, Value srcVec, int idx) {
+ auto vecTy = cast<VectorType>(srcVec.getType());
+ assert(idx >= 0 && idx < vecTy.getNumElements() &&
+ "extractElement: index out of bounds");
+ 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.
+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)};
+}
+
+/// Extract a vector of elements of size i32 from an i32 vector and bitcast to
+/// the specified vector type.
+static Value extractAndBitcast(ImplicitLocOpBuilder &b, Value srcI32Vec,
+ int idx, VectorType vecTy) {
+ Value elem = extractElement(b, srcI32Vec, idx);
+ return b.create<LLVM::BitcastOp>(vecTy, elem);
+}
+
+/// Create a sub-byte conversion from an f32 pair source and return the native
+/// result.
+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 and return
+/// the native result.
+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)...);
+}
+
+/// Create a typed NVVM truncation conversion.
+static Value createTruncConversion(
+ ImplicitLocOpBuilder &b, MLIRContext *ctx, FPTruncConvOp 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 FPTruncConvOp::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 FPTruncConvOp::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 FPTruncConvOp::F32x2_TO_F8x2:
+ return convertFromF32Pair<NVVM::ConvertF32x2ToF8x2Op>(
+ b, srcI32Vec, srcBaseIdx, i16Ty, rndAttr, satAttr, reluAttr, dstTyAttr);
+ case FPTruncConvOp::F32x2_TO_F6x2:
+ return convertFromF32Pair<NVVM::ConvertF32x2ToF6x2Op>(
+ b, srcI32Vec, srcBaseIdx, i16Ty, reluAttr, actualDstTyAttr);
+ case FPTruncConvOp::F32x2_TO_F4x2:
+ return convertFromF32Pair<NVVM::ConvertF32x2ToF4x2Op>(
+ b, srcI32Vec, srcBaseIdx, i8Ty, reluAttr, dstTyAttr);
+ case FPTruncConvOp::F16x2_TO_F8x2:
+ return convertFromPacked<NVVM::ConvertF16x2ToF8x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getF16Type(), i16Ty, reluAttr, dstTyAttr);
+ case FPTruncConvOp::F16x2_TO_F6x2:
+ return convertFromPacked<NVVM::ConvertF16x2ToF6x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getF16Type(), i16Ty, reluAttr,
+ actualDstTyAttr);
+ case FPTruncConvOp::F16x2_TO_F4x2:
+ return convertFromPacked<NVVM::ConvertF16x2ToF4x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getF16Type(), i8Ty, reluAttr,
+ actualDstTyAttr);
+ case FPTruncConvOp::BF16x2_TO_F8x2:
+ return convertFromPacked<NVVM::ConvertBF16x2ToF8x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i16Ty, rndAttr, satAttr,
+ reluAttr, dstTyAttr);
+ case FPTruncConvOp::BF16x2_TO_F6x2:
+ return convertFromPacked<NVVM::ConvertBF16x2ToF6x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i16Ty, reluAttr,
+ actualDstTyAttr);
+ case FPTruncConvOp::BF16x2_TO_F4x2:
+ return convertFromPacked<NVVM::ConvertBF16x2ToF4x2Op>(
+ b, srcI32Vec, srcBaseIdx, b.getBF16Type(), i8Ty, reluAttr,
+ actualDstTyAttr);
+ }
+ llvm_unreachable("unhandled FPTruncConvOp");
+}
+
+static LogicalResult lowerFPTrunc(nvgpu::ConvertFPTruncOp op,
+ nvgpu::ConvertFPTruncOp::Adaptor adaptor,
+ ConversionPatternRewriter &rewriter,
+ const LLVMTypeConverter *typeConverter) {
+ MLIRContext *ctx = op.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();
+
+ NVVM::FPRoundingModeAttr rndModeAttr = op.getRndAttr();
+ NVVM::SaturationModeAttr satModeAttr = op.getSatAttr();
+ auto reluBoolAttr = op.getReluAttr();
+ Value randomBits = adaptor.getRandomBits();
+ Type actualDstFloatType = dstElemType;
+
+ // STEP 1: bitcast input vector to i32 register vector.
----------------
durga4github wrote:
bitcast input vector to i32 vector type
https://github.com/llvm/llvm-project/pull/199700
More information about the Mlir-commits
mailing list