https://github.com/DharuniRAcharya created https://github.com/llvm/llvm-project/pull/225318
This patch adds support for `pzo` variants and `rz` rounding mode to existing `f32/f16x2/bf16x2` to `FP4` (`e2m1x2`) conversion intrinsics. Also adds `clang builtins` for the new variants. Tests have been verified through `ptxas-13.4`. PTX ISA Reference: https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#data-movement-and-conversion-instructions-cvt >From 6b7b3210508bcf5855a5d3dd6d8f2887d3620610 Mon Sep 17 00:00:00 2001 From: DharuniRAcharya <[email protected]> Date: Tue, 22 Sep 2026 07:47:14 +0000 Subject: [PATCH] [clang][NVPTX] Add support for pzo in f32/f16x2/bf16x2 to FP4 conversions This patch adds support for pzo variants and rz rounding mode to existing f32/f16x2/bf16x2 to FP4 (e2m1x2) conversion intrinsics. Also adds clang builtins for the new variants. Tests have been verified through ptxas-13.4. PTX ISA Reference: https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#data-movement-and-conversion-instructions-cvt Signed-off-by: DharuniRAcharya <[email protected]> --- clang/include/clang/Basic/BuiltinsNVPTX.td | 21 ++ clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp | 12 + clang/test/CodeGen/builtins-nvptx.c | 21 +- llvm/include/llvm/IR/IntrinsicsNVVM.td | 22 +- llvm/lib/Target/NVPTX/NVPTXInstrInfo.td | 6 +- llvm/lib/Target/NVPTX/NVPTXIntrinsics.td | 44 +-- llvm/test/CodeGen/NVPTX/convert-fp4-pzo.ll | 253 ++++++++++++++++++ mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp | 3 + .../Target/LLVMIR/nvvm/convert_fp4x2.mlir | 12 +- 9 files changed, 355 insertions(+), 39 deletions(-) create mode 100644 llvm/test/CodeGen/NVPTX/convert-fp4-pzo.ll diff --git a/clang/include/clang/Basic/BuiltinsNVPTX.td b/clang/include/clang/Basic/BuiltinsNVPTX.td index 8f1b43f9af65d4..05d47cb893f1b8 100644 --- a/clang/include/clang/Basic/BuiltinsNVPTX.td +++ b/clang/include/clang/Basic/BuiltinsNVPTX.td @@ -861,6 +861,27 @@ def __nvvm_bf16x2_to_e3m2x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short( def __nvvm_bf16x2_to_e3m2x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; def __nvvm_bf16x2_to_e3m2x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m1x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m1x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m1x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m1x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m1x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m1x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; + +def __nvvm_f16x2_to_e2m1x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m1x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m1x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m1x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m1x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m1x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; + +def __nvvm_bf16x2_to_e2m1x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m1x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m1x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m1x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m1x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m1x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; + def __nvvm_ff_to_e2m1x2_rn_satfinite : NVPTXBuiltinSMAndPTX<"short(float, float)", SMa<[100, 101, 120]>, PTX86>; def __nvvm_ff_to_e2m1x2_rn_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(float, float)", SMa<[100, 101, 120]>, PTX86>; diff --git a/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp b/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp index 76f6757326ecab..0b927d958ab2c7 100644 --- a/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp +++ b/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp @@ -1077,6 +1077,18 @@ Value *CodeGenFunction::EmitNVPTXBuiltinExpr(unsigned BuiltinID, PZO_CVT(bf16x2_to_e3m2x2_rn_relu_satfinite); PZO_CVT(bf16x2_to_e3m2x2_rz_satfinite); PZO_CVT(bf16x2_to_e3m2x2_rz_relu_satfinite); + PZO_CVT(ff_to_e2m1x2_rn_satfinite); + PZO_CVT(ff_to_e2m1x2_rn_relu_satfinite); + PZO_CVT(ff_to_e2m1x2_rz_satfinite); + PZO_CVT(ff_to_e2m1x2_rz_relu_satfinite); + PZO_CVT(f16x2_to_e2m1x2_rn_satfinite); + PZO_CVT(f16x2_to_e2m1x2_rn_relu_satfinite); + PZO_CVT(f16x2_to_e2m1x2_rz_satfinite); + PZO_CVT(f16x2_to_e2m1x2_rz_relu_satfinite); + PZO_CVT(bf16x2_to_e2m1x2_rn_satfinite); + PZO_CVT(bf16x2_to_e2m1x2_rn_relu_satfinite); + PZO_CVT(bf16x2_to_e2m1x2_rz_satfinite); + PZO_CVT(bf16x2_to_e2m1x2_rz_relu_satfinite); #undef PZO_CVT diff --git a/clang/test/CodeGen/builtins-nvptx.c b/clang/test/CodeGen/builtins-nvptx.c index 607095bce481c7..111770a0010f86 100644 --- a/clang/test/CodeGen/builtins-nvptx.c +++ b/clang/test/CodeGen/builtins-nvptx.c @@ -1206,6 +1206,15 @@ __device__ void nvvm_cvt_pzo_sm107f() { // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.bf16x2.to.e2m3x2.rz.satfinite(<2 x bfloat> zeroinitializer, i1 true) __nvvm_bf16x2_to_e2m3x2_rz_satfinite_pzo({0, 0}); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.ff.to.e2m1x2.rz.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + __nvvm_ff_to_e2m1x2_rz_relu_satfinite(1, 1); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.ff.to.e2m1x2.rz.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 true) + __nvvm_ff_to_e2m1x2_rz_relu_satfinite_pzo(1, 1); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.f16x2.to.e2m1x2.rn.satfinite(<2 x half> zeroinitializer, i1 true) + __nvvm_f16x2_to_e2m1x2_rn_satfinite_pzo({0, 0}); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rz.satfinite(<2 x bfloat> zeroinitializer, i1 true) + __nvvm_bf16x2_to_e2m1x2_rz_satfinite_pzo({0, 0}); + #endif // CHECK: ret void } @@ -1318,14 +1327,14 @@ __device__ void nvvm_cvt_sm100a_sm101a_sm120a() { // CHECK_PTX86_SM120a: call <2 x half> @llvm.nvvm.e3m2x2.to.f16x2.rn.relu(i16 19532) __nvvm_e3m2x2_to_f16x2_rn_relu(0x4C4C); - // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e2m1x2_rn_satfinite(1.0f, 1.0f); - // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e2m1x2_rn_relu_satfinite(1.0f, 1.0f); // CHECK_PTX86_SM100a: call <2 x half> @llvm.nvvm.e2m1x2.to.f16x2.rn(i16 76) diff --git a/llvm/include/llvm/IR/IntrinsicsNVVM.td b/llvm/include/llvm/IR/IntrinsicsNVVM.td index d9b73418394ed6..d1777b85eabef3 100644 --- a/llvm/include/llvm/IR/IntrinsicsNVVM.td +++ b/llvm/include/llvm/IR/IntrinsicsNVVM.td @@ -2136,18 +2136,24 @@ let TargetPrefix = "nvvm" in { // FP4 conversions. foreach relu = ["", "_relu"] in { - def int_nvvm_ff_to_e2m1x2_rn # relu # _satfinite : NVVMBuiltin, - PureIntrinsic<[llvm_i16_ty], [llvm_float_ty, llvm_float_ty]>; + foreach rnd = ["rn", "rz"] in { + def int_nvvm_ff_to_e2m1x2_ # rnd # relu # _satfinite : NVVMBuiltin, + PureIntrinsic<[llvm_i16_ty], + [llvm_float_ty, llvm_float_ty, llvm_i1_ty], + [ImmArg<ArgIndex<2>, DefaultValue<0>>]>; + + def int_nvvm_f16x2_to_e2m1x2_ # rnd # relu # _satfinite : NVVMBuiltin, + PureIntrinsic<[llvm_i16_ty], [llvm_v2f16_ty, llvm_i1_ty], + [ImmArg<ArgIndex<1>, DefaultValue<0>>]>; + + def int_nvvm_bf16x2_to_e2m1x2_ # rnd # relu # _satfinite : NVVMBuiltin, + PureIntrinsic<[llvm_i16_ty], [llvm_v2bf16_ty, llvm_i1_ty], + [ImmArg<ArgIndex<1>, DefaultValue<0>>]>; + } def int_nvvm_e2m1x2_to_f16x2_rn # relu : NVVMBuiltin, PureIntrinsic<[llvm_v2f16_ty], [llvm_i16_ty]>; - - def int_nvvm_f16x2_to_e2m1x2_rn # relu # _satfinite - : PureIntrinsic<[llvm_i16_ty], [llvm_v2f16_ty]>; - def int_nvvm_bf16x2_to_e2m1x2_rn # relu # _satfinite - : PureIntrinsic<[llvm_i16_ty], [llvm_v2bf16_ty]>; - foreach satfinite = ["", "_satfinite"] in { def int_nvvm_e2m1x2_to_bf16x2_rn # relu # satfinite # _scale_n2_ue8m0 : PureIntrinsic<[llvm_v2bf16_ty], [llvm_i16_ty, llvm_i16_ty]>; diff --git a/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td b/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td index f7940029719053..755c7884d6c1aa 100644 --- a/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td +++ b/llvm/lib/Target/NVPTX/NVPTXInstrInfo.td @@ -891,7 +891,7 @@ let Predicates = [hasS2F6X2ConversionSupport] in { (ins B32:$src1, B32:$src2, CvtMode:$mode), !strconcat("{{ \n\t", ".reg .b8 \t%e2m1x2_out; \n\t", - "cvt${mode:base}.satfinite${mode:relu}.e2m1x2.f32 \t%e2m1x2_out, $src1, $src2; \n\t", + "cvt${mode:base}.satfinite${mode:relu}${mode:pzo}.e2m1x2.f32 \t%e2m1x2_out, $src1, $src2; \n\t", "cvt.u16.u8 \t$dst, %e2m1x2_out; \n\t", "}}"), []>; @@ -923,7 +923,7 @@ let Predicates = [hasS2F6X2ConversionSupport] in { (ins B32:$src, CvtMode:$mode), "{{ \n\t" # ".reg .b8 \t%e2m1x2_out; \n\t" # - "cvt${mode:base}.satfinite${mode:relu}.e2m1x2.f16x2 \t%e2m1x2_out, $src; \n\t" # + "cvt${mode:base}.satfinite${mode:relu}${mode:pzo}.e2m1x2.f16x2 \t%e2m1x2_out, $src; \n\t" # "cvt.u16.u8 \t$dst, %e2m1x2_out; \n\t" # "}}", []>, Requires<[hasFP16X2ToNarrowFPConversionSupport]>; @@ -932,7 +932,7 @@ let Predicates = [hasS2F6X2ConversionSupport] in { (ins B32:$src, CvtMode:$mode), "{{ \n\t" # ".reg .b8 \t%e2m1x2_out; \n\t" # - "cvt${mode:base}.satfinite${mode:relu}.e2m1x2.bf16x2 \t%e2m1x2_out, $src; \n\t" # + "cvt${mode:base}.satfinite${mode:relu}${mode:pzo}.e2m1x2.bf16x2 \t%e2m1x2_out, $src; \n\t" # "cvt.u16.u8 \t$dst, %e2m1x2_out; \n\t" # "}}", []>, Requires<[hasFP16X2ToNarrowFPConversionSupport]>; diff --git a/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td b/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td index 89a9adf4426d80..492783cf2d9816 100644 --- a/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td +++ b/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td @@ -3244,29 +3244,41 @@ let Predicates = [hasNarrowFPToBF16x2ConversionSupport] in { } } // let Predicates = [hasNarrowFPToBF16x2ConversionSupport] -let Predicates = [hasNarrowFPConversionSupport] in { - def : Pat<(int_nvvm_ff_to_e2m1x2_rn_satfinite f32:$a, f32:$b), - (CVT_e2m1x2_f32_sf $a, $b, CvtRN)>; - def : Pat<(int_nvvm_ff_to_e2m1x2_rn_relu_satfinite f32:$a, f32:$b), - (CVT_e2m1x2_f32_sf $a, $b, CvtRN_RELU)>; +foreach relu = ["", "relu"] in { + foreach rnd = ["rn", "rz"] in { + defvar Relu = !if(!empty(relu), "", "_" # relu); + defvar BaseMode = !cast<PatLeaf>("Cvt" # !toupper(rnd # Relu)); + defvar Suffix = "e2m1x2_" # rnd # Relu; + defvar F32X2 = ToCvtIntrinsics<"ff_to_" # Suffix>; + defvar F16X2 = ToCvtIntrinsics<"f16x2_to_" # Suffix>; + defvar BF16X2 = ToCvtIntrinsics<"bf16x2_to_" # Suffix>; + + foreach pzo = [0, 1] in { + defvar PZO = !if(pzo, -1, 0); + defvar Mode = CvtModeWithPZO<BaseMode, pzo>.Value; + defvar NeedsPZOSupport = !or(!eq(rnd, "rz"), !eq(pzo, 1)); + defvar F32X2Preds = !if(NeedsPZOSupport, [hasConvertWithPZOSupport]<Predicate>, + [hasNarrowFPConversionSupport]<Predicate>); + defvar FPX2Preds = !if(NeedsPZOSupport, [hasConvertWithPZOSupport]<Predicate>, + [hasFP16X2ToNarrowFPConversionSupport]<Predicate>); + + def : Pat<(F32X2.sf f32:$a, f32:$b, PZO), (CVT_e2m1x2_f32_sf $a, $b, Mode)>, + Requires<F32X2Preds>; + def : Pat<(F16X2.sf v2f16:$a, PZO), (CVT_e2m1x2_f16x2_sf $a, Mode)>, + Requires<FPX2Preds>; + def : Pat<(BF16X2.sf v2bf16:$a, PZO), (CVT_e2m1x2_bf16x2_sf $a, Mode)>, + Requires<FPX2Preds>; + } + } +} +let Predicates = [hasNarrowFPConversionSupport] in { def : Pat<(int_nvvm_e2m1x2_to_f16x2_rn i16:$a), (CVT_f16x2_e2m1x2 $a, CvtRN)>; def : Pat<(int_nvvm_e2m1x2_to_f16x2_rn_relu i16:$a), (CVT_f16x2_e2m1x2 $a, CvtRN_RELU)>; } -let Predicates = [hasFP16X2ToNarrowFPConversionSupport] in { - foreach src_type = ["f16x2", "bf16x2"] in { - foreach relu = ["", "_relu"] in { - defvar intrin = !cast<Intrinsic>("int_nvvm_" # src_type # "_to_e2m1x2_rn" # relu # "_satfinite"); - defvar cvt_inst = !cast<NVPTXInst>("CVT_e2m1x2_" # src_type # "_sf"); - defvar cvt_mode = !cast<PatLeaf>("CvtRN" # !toupper(relu)); - def : Pat<(intrin B32:$a), (cvt_inst $a, cvt_mode)>; - } - } -} - let Predicates = [hasNarrowFPToBF16x2ConversionSupport] in { foreach relu = ["", "_relu"] in { foreach satfinite = ["", "_satfinite"] in { diff --git a/llvm/test/CodeGen/NVPTX/convert-fp4-pzo.ll b/llvm/test/CodeGen/NVPTX/convert-fp4-pzo.ll new file mode 100644 index 00000000000000..f276f664b4c16c --- /dev/null +++ b/llvm/test/CodeGen/NVPTX/convert-fp4-pzo.ll @@ -0,0 +1,253 @@ +; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6 +; RUN: llc < %s -mtriple=nvptx64 -mcpu=sm_107f -mattr=+ptx94 | FileCheck %s +; RUN: %if ptxas-sm_107f && ptxas-isa-9.4 %{ llc < %s -mtriple=nvptx64 -mcpu=sm_107f -mattr=+ptx94 | %ptxas-verify -arch=sm_107f %} + +; E2M1X2 conversions from f32 + +define i16 @cvt_rn_pzo_e2m1x2_f32(float %f1, float %f2) { +; CHECK-LABEL: cvt_rn_pzo_e2m1x2_f32( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<4>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rn_pzo_e2m1x2_f32_param_0]; +; CHECK-NEXT: ld.param::func.b32 %r2, [cvt_rn_pzo_e2m1x2_f32_param_1]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rn.satfinite.pzo.e2m1x2.f32 %e2m1x2_out, %r1, %r2; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r3, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float %f1, float %f2, i1 true) + ret i16 %val +} + +define i16 @cvt_rz_pzo_e2m1x2_f32(float %f1, float %f2) { +; CHECK-LABEL: cvt_rz_pzo_e2m1x2_f32( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<4>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rz_pzo_e2m1x2_f32_param_0]; +; CHECK-NEXT: ld.param::func.b32 %r2, [cvt_rz_pzo_e2m1x2_f32_param_1]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rz.satfinite.pzo.e2m1x2.f32 %e2m1x2_out, %r1, %r2; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r3, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.ff.to.e2m1x2.rz.satfinite(float %f1, float %f2, i1 true) + ret i16 %val +} + +define i16 @cvt_rn_relu_pzo_e2m1x2_f32(float %f1, float %f2) { +; CHECK-LABEL: cvt_rn_relu_pzo_e2m1x2_f32( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<4>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rn_relu_pzo_e2m1x2_f32_param_0]; +; CHECK-NEXT: ld.param::func.b32 %r2, [cvt_rn_relu_pzo_e2m1x2_f32_param_1]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rn.satfinite.relu.pzo.e2m1x2.f32 %e2m1x2_out, %r1, %r2; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r3, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float %f1, float %f2, i1 true) + ret i16 %val +} + +define i16 @cvt_rz_relu_pzo_e2m1x2_f32(float %f1, float %f2) { +; CHECK-LABEL: cvt_rz_relu_pzo_e2m1x2_f32( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<4>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rz_relu_pzo_e2m1x2_f32_param_0]; +; CHECK-NEXT: ld.param::func.b32 %r2, [cvt_rz_relu_pzo_e2m1x2_f32_param_1]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rz.satfinite.relu.pzo.e2m1x2.f32 %e2m1x2_out, %r1, %r2; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r3, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.ff.to.e2m1x2.rz.relu.satfinite(float %f1, float %f2, i1 true) + ret i16 %val +} + +; E2M1X2 conversions from f16x2 + +define i16 @cvt_rn_pzo_e2m1x2_f16x2(<2 x half> %a) { +; CHECK-LABEL: cvt_rn_pzo_e2m1x2_f16x2( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<3>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rn_pzo_e2m1x2_f16x2_param_0]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rn.satfinite.pzo.e2m1x2.f16x2 %e2m1x2_out, %r1; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r2, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r2; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.f16x2.to.e2m1x2.rn.satfinite(<2 x half> %a, i1 true) + ret i16 %val +} + +define i16 @cvt_rz_pzo_e2m1x2_f16x2(<2 x half> %a) { +; CHECK-LABEL: cvt_rz_pzo_e2m1x2_f16x2( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<3>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rz_pzo_e2m1x2_f16x2_param_0]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rz.satfinite.pzo.e2m1x2.f16x2 %e2m1x2_out, %r1; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r2, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r2; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.f16x2.to.e2m1x2.rz.satfinite(<2 x half> %a, i1 true) + ret i16 %val +} + +define i16 @cvt_rn_relu_pzo_e2m1x2_f16x2(<2 x half> %a) { +; CHECK-LABEL: cvt_rn_relu_pzo_e2m1x2_f16x2( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<3>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rn_relu_pzo_e2m1x2_f16x2_param_0]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rn.satfinite.relu.pzo.e2m1x2.f16x2 %e2m1x2_out, %r1; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r2, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r2; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.f16x2.to.e2m1x2.rn.relu.satfinite(<2 x half> %a, i1 true) + ret i16 %val +} + +define i16 @cvt_rz_relu_pzo_e2m1x2_f16x2(<2 x half> %a) { +; CHECK-LABEL: cvt_rz_relu_pzo_e2m1x2_f16x2( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<3>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rz_relu_pzo_e2m1x2_f16x2_param_0]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rz.satfinite.relu.pzo.e2m1x2.f16x2 %e2m1x2_out, %r1; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r2, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r2; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.f16x2.to.e2m1x2.rz.relu.satfinite(<2 x half> %a, i1 true) + ret i16 %val +} + +; E2M1X2 conversions from bf16x2 + +define i16 @cvt_rn_pzo_e2m1x2_bf16x2(<2 x bfloat> %a) { +; CHECK-LABEL: cvt_rn_pzo_e2m1x2_bf16x2( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<3>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rn_pzo_e2m1x2_bf16x2_param_0]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rn.satfinite.pzo.e2m1x2.bf16x2 %e2m1x2_out, %r1; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r2, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r2; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rn.satfinite(<2 x bfloat> %a, i1 true) + ret i16 %val +} + +define i16 @cvt_rz_pzo_e2m1x2_bf16x2(<2 x bfloat> %a) { +; CHECK-LABEL: cvt_rz_pzo_e2m1x2_bf16x2( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<3>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rz_pzo_e2m1x2_bf16x2_param_0]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rz.satfinite.pzo.e2m1x2.bf16x2 %e2m1x2_out, %r1; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r2, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r2; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rz.satfinite(<2 x bfloat> %a, i1 true) + ret i16 %val +} + +define i16 @cvt_rn_relu_pzo_e2m1x2_bf16x2(<2 x bfloat> %a) { +; CHECK-LABEL: cvt_rn_relu_pzo_e2m1x2_bf16x2( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<3>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rn_relu_pzo_e2m1x2_bf16x2_param_0]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rn.satfinite.relu.pzo.e2m1x2.bf16x2 %e2m1x2_out, %r1; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r2, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r2; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rn.relu.satfinite(<2 x bfloat> %a, i1 true) + ret i16 %val +} + +define i16 @cvt_rz_relu_pzo_e2m1x2_bf16x2(<2 x bfloat> %a) { +; CHECK-LABEL: cvt_rz_relu_pzo_e2m1x2_bf16x2( +; CHECK: { +; CHECK-NEXT: .reg .b16 %rs<2>; +; CHECK-NEXT: .reg .b32 %r<3>; +; CHECK-EMPTY: +; CHECK-NEXT: // %bb.0: +; CHECK-NEXT: ld.param::func.b32 %r1, [cvt_rz_relu_pzo_e2m1x2_bf16x2_param_0]; +; CHECK-NEXT: { +; CHECK-NEXT: .reg .b8 %e2m1x2_out; +; CHECK-NEXT: cvt.rz.satfinite.relu.pzo.e2m1x2.bf16x2 %e2m1x2_out, %r1; +; CHECK-NEXT: cvt.u16.u8 %rs1, %e2m1x2_out; +; CHECK-NEXT: } +; CHECK-NEXT: cvt.u32.u16 %r2, %rs1; +; CHECK-NEXT: st.param::func.b32 [func_retval0], %r2; +; CHECK-NEXT: ret; + %val = call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rz.relu.satfinite(<2 x bfloat> %a, i1 true) + ret i16 %val +} diff --git a/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp b/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp index 46e9ab29bc4a04..7d7bdfa9465627 100644 --- a/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp +++ b/mlir/lib/Dialect/LLVMIR/IR/NVVMDialect.cpp @@ -5086,6 +5086,7 @@ ConvertF32x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF32x2ToF4x2Op op, llvm::SmallVector<llvm::Value *> args; args.push_back(mt.lookupValue(op.getA())); args.push_back(mt.lookupValue(op.getB())); + args.push_back(builder.getInt1(false)); bool hasRelu = op.getRelu(); @@ -5130,6 +5131,7 @@ ConvertF16x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertF16x2ToF4x2Op &op, llvm::SmallVector<llvm::Value *> args; args.push_back(mt.lookupValue(op.getSrc())); + args.push_back(builder.getInt1(false)); return {intId, std::move(args)}; } @@ -5149,6 +5151,7 @@ ConvertBF16x2ToF4x2Op::getIntrinsicIDAndArgs(NVVM::ConvertBF16x2ToF4x2Op &op, llvm::SmallVector<llvm::Value *> args; args.push_back(mt.lookupValue(op.getSrc())); + args.push_back(builder.getInt1(false)); return {intId, std::move(args)}; } diff --git a/mlir/test/Target/LLVMIR/nvvm/convert_fp4x2.mlir b/mlir/test/Target/LLVMIR/nvvm/convert_fp4x2.mlir index 75feb7176b7b38..9538b7d798f7ae 100644 --- a/mlir/test/Target/LLVMIR/nvvm/convert_fp4x2.mlir +++ b/mlir/test/Target/LLVMIR/nvvm/convert_fp4x2.mlir @@ -2,10 +2,10 @@ // CHECK-LABEL: @convert_f32x2_to_f4x2_e2m1 llvm.func @convert_f32x2_to_f4x2_e2m1(%srcA : f32, %srcB : f32) { - // CHECK: %[[res1:.*]] = call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float %{{.*}}, float %{{.*}}) + // CHECK: %[[res1:.*]] = call i16 @llvm.nvvm.ff.to.e2m1x2.rn.satfinite(float %{{.*}}, float %{{.*}}, i1 false) // CHECK-NEXT: %{{.*}} = trunc i16 %[[res1]] to i8 %res1 = nvvm.convert.f32x2.to.f4x2 %srcA, %srcB : i8 (f4E2M1FN) - // CHECK: %[[res2:.*]] = call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float %{{.*}}, float %{{.*}}) + // CHECK: %[[res2:.*]] = call i16 @llvm.nvvm.ff.to.e2m1x2.rn.relu.satfinite(float %{{.*}}, float %{{.*}}, i1 false) // CHECK-NEXT: %{{.*}} = trunc i16 %[[res2]] to i8 %res2 = nvvm.convert.f32x2.to.f4x2 %srcA, %srcB relu = true : i8 (f4E2M1FN) llvm.return @@ -15,10 +15,10 @@ llvm.func @convert_f32x2_to_f4x2_e2m1(%srcA : f32, %srcB : f32) { // CHECK-LABEL: @convert_f16x2_to_f4x2 llvm.func @convert_f16x2_to_f4x2(%srcA : vector<2xf16>) { - // CHECK: %[[res1:.*]] = call i16 @llvm.nvvm.f16x2.to.e2m1x2.rn.satfinite(<2 x half> %{{.*}}) + // CHECK: %[[res1:.*]] = call i16 @llvm.nvvm.f16x2.to.e2m1x2.rn.satfinite(<2 x half> %{{.*}}, i1 false) // CHECK-NEXT: %{{.*}} = trunc i16 %[[res1]] to i8 %res1 = nvvm.convert.f16x2.to.f4x2 %srcA : vector<2xf16> -> i8 (f4E2M1FN) - // CHECK: %[[res2:.*]] = call i16 @llvm.nvvm.f16x2.to.e2m1x2.rn.relu.satfinite(<2 x half> %{{.*}}) + // CHECK: %[[res2:.*]] = call i16 @llvm.nvvm.f16x2.to.e2m1x2.rn.relu.satfinite(<2 x half> %{{.*}}, i1 false) // CHECK-NEXT: %{{.*}} = trunc i16 %[[res2]] to i8 %res2 = nvvm.convert.f16x2.to.f4x2 %srcA relu = true : vector<2xf16> -> i8 (f4E2M1FN) llvm.return @@ -28,10 +28,10 @@ llvm.func @convert_f16x2_to_f4x2(%srcA : vector<2xf16>) { // CHECK-LABEL: @convert_bf16x2_to_f4x2 llvm.func @convert_bf16x2_to_f4x2(%srcA : vector<2xbf16>) { - // CHECK: %[[res1:.*]] = call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rn.satfinite(<2 x bfloat> %{{.*}}) + // CHECK: %[[res1:.*]] = call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rn.satfinite(<2 x bfloat> %{{.*}}, i1 false) // CHECK-NEXT: %{{.*}} = trunc i16 %[[res1]] to i8 %res1 = nvvm.convert.bf16x2.to.f4x2 %srcA : vector<2xbf16> -> i8 (f4E2M1FN) - // CHECK: %[[res2:.*]] = call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rn.relu.satfinite(<2 x bfloat> %{{.*}}) + // CHECK: %[[res2:.*]] = call i16 @llvm.nvvm.bf16x2.to.e2m1x2.rn.relu.satfinite(<2 x bfloat> %{{.*}}, i1 false) // CHECK-NEXT: %{{.*}} = trunc i16 %[[res2]] to i8 %res2 = nvvm.convert.bf16x2.to.f4x2 %srcA relu = true : vector<2xbf16> -> i8 (f4E2M1FN) llvm.return _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
