[Mlir-commits] [mlir] [MLIR][NVGPU] Add convert.fpext and convert.fptrunc Ops (PR #199700)
Durgadoss R
llvmlistbot at llvm.org
Tue Jun 30 01:10:42 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.
+ // f64 -> f32/f16/bf16 lowers to a single direct LLVM fptrunc
+ // f64 -> f8/f6/f4 first truncates to f32 and then reuses the narrow
+ // conversion path below.
+ Value input = adaptor.getIn();
+ if (srcBW == 64) {
+ if (dstBW >= 16) {
+ Type convertedType = typeConverter->convertType(dstType);
+ assert(convertedType && "failed to convert type");
+ Value result = b.create<LLVM::FPTruncOp>(convertedType, input);
+ rewriter.replaceOp(op, result);
+ return success();
+ }
+ auto f32VecTy = VectorType::get(srcType.getShape(), b.getF32Type());
+ input = b.create<LLVM::FPTruncOp>(f32VecTy, input);
+ srcType = f32VecTy;
+ srcElemType = b.getF32Type();
+ srcBW = 32;
+ }
+
+ // f6 types are 6-bit in MLIR but NVVM uses 8-bit containers.
+ int effectiveDstBW = getEffectiveBitWidth(dstBW);
+
+ int srcI32Elems = numElems * srcBW / regBits;
+ int dstI32Elems = numElems * effectiveDstBW / 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");
+ FPTruncConvOp convOp = convEntry->convOp;
+ int numSrcI32PerConv = getNumSrcI32PerConv(convEntry->src);
+
+ // STEP 3: pack conversion results into destination i32 vector.
+ const int srcStep = srcBW / effectiveDstBW;
+ const int resultBW =
+ effectiveDstBW * 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/f4 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 -= numSrcI32PerConv;
+ 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);
+ }
+
+ // STEP 4: produce final result.
+ Type convertedType = typeConverter->convertType(dstType);
+ assert(convertedType && "failed to convert type");
+ if (convEntry->dst == FPKind::F6) {
+ IntegerType i8Ty = b.getI8Type();
+ auto i8VecTy = VectorType::get(numElems, i8Ty);
+ Value i8Vec = b.create<LLVM::BitcastOp>(i8VecTy, dstI32Vec);
+ Value truncVec = b.create<LLVM::TruncOp>(convertedType, i8Vec);
+ rewriter.replaceOp(op, truncVec);
+ } else {
+ auto dstVec = b.create<LLVM::BitcastOp>(convertedType, dstI32Vec);
+ rewriter.replaceOp(op, dstVec);
+ }
+ return success();
+}
+
+//===----------------------------------------------------------------------===//
+// NVGPUConvertFPExtOp Lowering
+//===----------------------------------------------------------------------===//
+
+/// Conversion op identifier for nvgpu.convert.fpext lowering dispatch table.
+enum class FPExtConvOp {
+ F8x2_TO_F16x2,
+ F8x2_TO_BF16x2,
+ F6x2_TO_F16x2,
+ F6x2_TO_BF16x2,
+ F4x2_TO_F16x2,
+ F4x2_TO_BF16x2,
+};
+
+struct FPExtTableEntry {
+ FPKind src;
+ FPKind dst;
+ FPExtConvOp convOp;
+};
+
+static constexpr FPExtTableEntry kFPExtTable[] = {
+ {FPKind::F8, FPKind::F16, FPExtConvOp::F8x2_TO_F16x2},
+ {FPKind::F8, FPKind::BF16, FPExtConvOp::F8x2_TO_BF16x2},
+ {FPKind::F6, FPKind::F16, FPExtConvOp::F6x2_TO_F16x2},
+ {FPKind::F6, FPKind::BF16, FPExtConvOp::F6x2_TO_BF16x2},
+ {FPKind::F4, FPKind::F16, FPExtConvOp::F4x2_TO_F16x2},
+ {FPKind::F4, FPKind::BF16, FPExtConvOp::F4x2_TO_BF16x2},
+};
+
+static std::optional<FPExtTableEntry> lookupExtConvOp(Type srcElemType,
+ Type dstElemType) {
+ auto srcKind = classifyFPType(srcElemType);
+ auto dstKind = classifyFPType(dstElemType);
+ if (!srcKind || !dstKind)
+ return std::nullopt;
+ for (const auto &entry : kFPExtTable) {
+ if (entry.src == *srcKind && entry.dst == *dstKind)
+ return entry;
+ }
+ return std::nullopt;
+}
----------------
durga4github wrote:
This lookup seems very similar to the Trunc lookup except for the Ext/Trunc Op kind.
Is it possible to unify these methods using a template or something?
https://github.com/llvm/llvm-project/pull/199700
More information about the Mlir-commits
mailing list