llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT--> @llvm/pr-subscribers-mlir-llvm @llvm/pr-subscribers-backend-nvptx Author: Dharuni R Acharya (DharuniRAcharya) <details> <summary>Changes</summary> This patch adds support for `pzo` variants and `rz` rounding mode to existing `f32/f16x2/bf16x2` to `FP8` (`e4m3x2`, `e5m2x2`) and `FP6` (`e2m3x2`, `e3m2x2`) 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 --- Patch is 79.15 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/222511.diff 11 Files Affected: - (modified) clang/include/clang/Basic/BuiltinsNVPTX.td (+78) - (modified) clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp (+50) - (modified) clang/test/CodeGen/builtins-nvptx.c (+39-20) - (modified) llvm/include/llvm/IR/IntrinsicsNVVM.td (+32-16) - (modified) llvm/lib/Target/NVPTX/NVPTXInstrInfo.td (+12-6) - (modified) llvm/lib/Target/NVPTX/NVPTXIntrinsics.td (+65-50) - (added) llvm/test/CodeGen/NVPTX/convert-fp6-pzo.ll (+407) - (added) llvm/test/CodeGen/NVPTX/convert-fp8-pzo.ll (+407) - (modified) mlir/include/mlir/Dialect/LLVMIR/NVVMOps.td (+14-5) - (modified) mlir/test/Target/LLVMIR/nvvm/convert_fp6x2.mlir (+18-18) - (modified) mlir/test/Target/LLVMIR/nvvm/convert_fp8x2.mlir (+18-18) ``````````diff diff --git a/clang/include/clang/Basic/BuiltinsNVPTX.td b/clang/include/clang/Basic/BuiltinsNVPTX.td index 3475b721e95e0..8f1b43f9af65d 100644 --- a/clang/include/clang/Basic/BuiltinsNVPTX.td +++ b/clang/include/clang/Basic/BuiltinsNVPTX.td @@ -783,6 +783,84 @@ def __nvvm_f32x4_to_e3m2x4_rs_relu_satfinite : NVPTXBuiltinSMAndPTX<"_Vector<4, char>(_Vector<4, float>, uint32_t)", SMa<[100, 103]>, PTX87>; +def __nvvm_ff_to_e4m3x2_rz : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e4m3x2_rz_relu : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e5m2x2_rz : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e5m2x2_rz_relu : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e4m3x2_rn_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e4m3x2_rn_relu_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e4m3x2_rz_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e4m3x2_rz_relu_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e5m2x2_rn_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e5m2x2_rn_relu_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e5m2x2_rz_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e5m2x2_rz_relu_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; + +def __nvvm_f16x2_to_e4m3x2_rz : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e4m3x2_rz_relu : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e5m2x2_rz : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e5m2x2_rz_relu : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e4m3x2_rn_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e4m3x2_rn_relu_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e4m3x2_rz_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e4m3x2_rz_relu_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e5m2x2_rn_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e5m2x2_rn_relu_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e5m2x2_rz_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e5m2x2_rz_relu_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; + +def __nvvm_bf16x2_to_e4m3x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e4m3x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e5m2x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e5m2x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e4m3x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e4m3x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e4m3x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e4m3x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e5m2x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e5m2x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e5m2x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e5m2x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; + +def __nvvm_ff_to_e2m3x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m3x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e3m2x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e3m2x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m3x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m3x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m3x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e2m3x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e3m2x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e3m2x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e3m2x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; +def __nvvm_ff_to_e3m2x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(float, float)", SM_107f, PTX94>; + +def __nvvm_f16x2_to_e2m3x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m3x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e3m2x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e3m2x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m3x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m3x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m3x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e2m3x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e3m2x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e3m2x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e3m2x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; +def __nvvm_f16x2_to_e3m2x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __fp16>)", SM_107f, PTX94>; + +def __nvvm_bf16x2_to_e2m3x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m3x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e3m2x2_rz_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e3m2x2_rz_relu_satfinite : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m3x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m3x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m3x2_rz_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e2m3x2_rz_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e3m2x2_rn_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +def __nvvm_bf16x2_to_e3m2x2_rn_relu_satfinite_pzo : NVPTXBuiltinSMAndPTX<"short(_Vector<2, __bf16>)", SM_107f, PTX94>; +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_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 06c5069d6f984..1c4ea2d5e9143 100644 --- a/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp +++ b/clang/lib/CodeGen/TargetBuiltins/NVPTX.cpp @@ -1028,6 +1028,56 @@ Value *CodeGenFunction::EmitNVPTXBuiltinExpr(unsigned BuiltinID, PZO_CVT(f2bf16_rz_satfinite); PZO_CVT(f2bf16_rz_relu_satfinite); + PZO_CVT(ff_to_e4m3x2_rn); + PZO_CVT(ff_to_e4m3x2_rn_relu); + PZO_CVT(ff_to_e4m3x2_rz); + PZO_CVT(ff_to_e4m3x2_rz_relu); + PZO_CVT(ff_to_e5m2x2_rn); + PZO_CVT(ff_to_e5m2x2_rn_relu); + PZO_CVT(ff_to_e5m2x2_rz); + PZO_CVT(ff_to_e5m2x2_rz_relu); + PZO_CVT(f16x2_to_e4m3x2_rn); + PZO_CVT(f16x2_to_e4m3x2_rn_relu); + PZO_CVT(f16x2_to_e4m3x2_rz); + PZO_CVT(f16x2_to_e4m3x2_rz_relu); + PZO_CVT(f16x2_to_e5m2x2_rn); + PZO_CVT(f16x2_to_e5m2x2_rn_relu); + PZO_CVT(f16x2_to_e5m2x2_rz); + PZO_CVT(f16x2_to_e5m2x2_rz_relu); + PZO_CVT(bf16x2_to_e4m3x2_rn_satfinite); + PZO_CVT(bf16x2_to_e4m3x2_rn_relu_satfinite); + PZO_CVT(bf16x2_to_e4m3x2_rz_satfinite); + PZO_CVT(bf16x2_to_e4m3x2_rz_relu_satfinite); + PZO_CVT(bf16x2_to_e5m2x2_rn_satfinite); + PZO_CVT(bf16x2_to_e5m2x2_rn_relu_satfinite); + PZO_CVT(bf16x2_to_e5m2x2_rz_satfinite); + PZO_CVT(bf16x2_to_e5m2x2_rz_relu_satfinite); + + PZO_CVT(ff_to_e2m3x2_rn_satfinite); + PZO_CVT(ff_to_e2m3x2_rn_relu_satfinite); + PZO_CVT(ff_to_e2m3x2_rz_satfinite); + PZO_CVT(ff_to_e2m3x2_rz_relu_satfinite); + PZO_CVT(ff_to_e3m2x2_rn_satfinite); + PZO_CVT(ff_to_e3m2x2_rn_relu_satfinite); + PZO_CVT(ff_to_e3m2x2_rz_satfinite); + PZO_CVT(ff_to_e3m2x2_rz_relu_satfinite); + PZO_CVT(f16x2_to_e2m3x2_rn_satfinite); + PZO_CVT(f16x2_to_e2m3x2_rn_relu_satfinite); + PZO_CVT(f16x2_to_e2m3x2_rz_satfinite); + PZO_CVT(f16x2_to_e2m3x2_rz_relu_satfinite); + PZO_CVT(f16x2_to_e3m2x2_rn_satfinite); + PZO_CVT(f16x2_to_e3m2x2_rn_relu_satfinite); + PZO_CVT(f16x2_to_e3m2x2_rz_satfinite); + PZO_CVT(f16x2_to_e3m2x2_rz_relu_satfinite); + PZO_CVT(bf16x2_to_e2m3x2_rn_satfinite); + PZO_CVT(bf16x2_to_e2m3x2_rn_relu_satfinite); + PZO_CVT(bf16x2_to_e2m3x2_rz_satfinite); + PZO_CVT(bf16x2_to_e2m3x2_rz_relu_satfinite); + PZO_CVT(bf16x2_to_e3m2x2_rn_satfinite); + PZO_CVT(bf16x2_to_e3m2x2_rn_relu_satfinite); + PZO_CVT(bf16x2_to_e3m2x2_rz_satfinite); + PZO_CVT(bf16x2_to_e3m2x2_rz_relu_satfinite); + #undef PZO_CVT case NVPTX::BI__nvvm_fma_rn_f16: diff --git a/clang/test/CodeGen/builtins-nvptx.c b/clang/test/CodeGen/builtins-nvptx.c index bed1498236b06..b74e9c74d84a9 100644 --- a/clang/test/CodeGen/builtins-nvptx.c +++ b/clang/test/CodeGen/builtins-nvptx.c @@ -1187,6 +1187,25 @@ __device__ void nvvm_cvt_pzo_sm107f() { __nvvm_f2f16_rz_satfinite_pzo(1); // CHECK_PTX94_SM107f: call half @llvm.nvvm.f2f16.rz.relu.satfinite(float 1.000000e+00, i1 true) __nvvm_f2f16_rz_relu_satfinite_pzo(1); + + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.ff.to.e4m3x2.rz.relu(float 1.000000e+00, float 1.000000e+00, i1 false) + __nvvm_ff_to_e4m3x2_rz_relu(1, 1); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.ff.to.e4m3x2.rz.relu(float 1.000000e+00, float 1.000000e+00, i1 true) + __nvvm_ff_to_e4m3x2_rz_relu_pzo(1, 1); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.f16x2.to.e5m2x2.rn(<2 x half> zeroinitializer, i1 true) + __nvvm_f16x2_to_e5m2x2_rn_pzo({0, 0}); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.bf16x2.to.e4m3x2.rz.satfinite(<2 x bfloat> zeroinitializer, i1 true) + __nvvm_bf16x2_to_e4m3x2_rz_satfinite_pzo({0, 0}); + + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.ff.to.e2m3x2.rz.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + __nvvm_ff_to_e2m3x2_rz_relu_satfinite(1, 1); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.ff.to.e2m3x2.rz.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 true) + __nvvm_ff_to_e2m3x2_rz_relu_satfinite_pzo(1, 1); + // CHECK_PTX94_SM107f: call i16 @llvm.nvvm.f16x2.to.e3m2x2.rn.satfinite(<2 x half> zeroinitializer, i1 true) + __nvvm_f16x2_to_e3m2x2_rn_satfinite_pzo({0, 0}); + // 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}); + #endif // CHECK: ret void } @@ -1194,22 +1213,22 @@ __device__ void nvvm_cvt_pzo_sm107f() { // CHECK-LABEL: nvvm_cvt_sm89 __device__ void nvvm_cvt_sm89() { #if (PTX >= 81) && (__CUDA_ARCH__ >= 890) - // CHECK_PTX81_SM89: call i16 @llvm.nvvm.ff.to.e4m3x2.rn(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX81_SM89: call i16 @llvm.nvvm.ff.to.e4m3x2.rn(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e4m3x2_rn(1.0f, 1.0f); - // CHECK_PTX81_SM89: call i16 @llvm.nvvm.ff.to.e4m3x2.rn.relu(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX81_SM89: call i16 @llvm.nvvm.ff.to.e4m3x2.rn.relu(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e4m3x2_rn_relu(1.0f, 1.0f); - // CHECK_PTX81_SM89: call i16 @llvm.nvvm.ff.to.e5m2x2.rn(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX81_SM89: call i16 @llvm.nvvm.ff.to.e5m2x2.rn(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e5m2x2_rn(1.0f, 1.0f); - // CHECK_PTX81_SM89: call i16 @llvm.nvvm.ff.to.e5m2x2.rn.relu(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX81_SM89: call i16 @llvm.nvvm.ff.to.e5m2x2.rn.relu(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e5m2x2_rn_relu(1.0f, 1.0f); - // CHECK_PTX81_SM89: call i16 @llvm.nvvm.f16x2.to.e4m3x2.rn(<2 x half> splat (half 1.000000e+00)) + // CHECK_PTX81_SM89: call i16 @llvm.nvvm.f16x2.to.e4m3x2.rn(<2 x half> splat (half 1.000000e+00), i1 false) __nvvm_f16x2_to_e4m3x2_rn({1.0f16, 1.0f16}); - // CHECK_PTX81_SM89: call i16 @llvm.nvvm.f16x2.to.e4m3x2.rn.relu(<2 x half> splat (half 1.000000e+00)) + // CHECK_PTX81_SM89: call i16 @llvm.nvvm.f16x2.to.e4m3x2.rn.relu(<2 x half> splat (half 1.000000e+00), i1 false) __nvvm_f16x2_to_e4m3x2_rn_relu({1.0f16, 1.0f16}); - // CHECK_PTX81_SM89: call i16 @llvm.nvvm.f16x2.to.e5m2x2.rn(<2 x half> splat (half 1.000000e+00)) + // CHECK_PTX81_SM89: call i16 @llvm.nvvm.f16x2.to.e5m2x2.rn(<2 x half> splat (half 1.000000e+00), i1 false) __nvvm_f16x2_to_e5m2x2_rn({1.0f16, 1.0f16}); - // CHECK_PTX81_SM89: call i16 @llvm.nvvm.f16x2.to.e5m2x2.rn.relu(<2 x half> splat (half 1.000000e+00)) + // CHECK_PTX81_SM89: call i16 @llvm.nvvm.f16x2.to.e5m2x2.rn.relu(<2 x half> splat (half 1.000000e+00), i1 false) __nvvm_f16x2_to_e5m2x2_rn_relu({1.0f16, 1.0f16}); // CHECK_PTX81_SM89: call <2 x half> @llvm.nvvm.e4m3x2.to.f16x2.rn(i16 18504) @@ -1259,24 +1278,24 @@ __device__ void nvvm_cvt_sm100a_sm101a_sm120a() { #if (PTX >= 86) && \ (__CUDA_ARCH_FEAT_SM100_ALL || __CUDA_ARCH_FEAT_SM101_ALL || \ __CUDA_ARCH_FEAT_SM120_ALL) - // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e2m3x2_rn_satfinite(1.0f, 1.0f); - // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e2m3x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e2m3x2_rn_relu_satfinite(1.0f, 1.0f); - // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e3m2x2_rn_satfinite(1.0f, 1.0f); - // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) - // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00) + // CHECK_PTX86_SM100a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM101a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) + // CHECK_PTX86_SM120a: call i16 @llvm.nvvm.ff.to.e3m2x2.rn.relu.satfinite(float 1.000000e+00, float 1.000000e+00, i1 false) __nvvm_ff_to_e3m2x2_rn_relu_satfinite(1.0f, 1.0f); // CHECK_PTX86_SM100a: call <2 x half> @llvm.nvvm.e2m3x2.to.f16x2.rn(i16 19532) diff --git a/llvm/include/llvm/IR/IntrinsicsNVVM.td b/llvm/include/llvm/IR/IntrinsicsNVVM.td index 87ae664ce4c17..bc7b6d58653d6 100644 --- a/llvm/include/llvm/IR/IntrinsicsNVVM.td +++ b/llvm/include/llvm/IR/IntrinsicsNVVM.td @@ -2057,18 +2057,25 @@ let TargetPrefix = "nvvm" in { foreach type = ["e4m3x2", "e5m2x2"] in { foreach relu = ["", "_relu"] in { - def int_nvvm_ff_to_ # type # _rn # relu : NVVMBuiltin, - PureIntrinsic<[llvm_i16_ty], [llvm_float_ty, llvm_float_ty]>; + foreach rnd = ["rn", "rz"] in { + def int_nvvm_ff_to_ # type # _ # rnd # relu : NVVMBuiltin, + PureIntrinsic<[llvm_i16_ty], + [llvm_float_ty, llvm_float_ty, llvm_i1_ty], + [ImmArg<ArgIndex<2>, DefaultValue<0>>]>; - def int_nvvm_f16x2_to_ # type # _rn # relu : NVVMBuiltin, - PureIntrinsic<[llvm_i16_ty], [llvm_v2f16_ty]>; + def int_nvvm_f16x2_to_ # type # _ # rnd # relu : NVVMBuiltin, + PureIntrinsic<[llvm_i16_ty], [llvm_v2f16_ty, llvm_i1_ty], + [ImmArg<ArgIndex<1>, DefaultValue<0>>]>; + + def int_nvvm_bf16x2_to_ # type # _ # rnd # relu # _satfinite : + NVVMBuiltin, + PureIntrinsic<[llvm_i16_ty], [llvm_v2bf16_ty, llvm_i1_ty], + [ImmArg<ArgIndex<1>, DefaultValue<0>>]>; + } def int_nvvm_ # type # _to_f16x2_rn # relu : NVVMBuiltin, PureIntrinsic<[llvm_v2f16_ty], [llvm_i16_ty]>; - - def int_nvvm_bf16x2_to_ # type # _rn # relu # _satfinite - : PureIntrinsic<[llvm_i16_ty], [llvm_v2bf16_ty]>; - + foreach satfinite = ["", "_satfinite"] in { def int_nvvm_ # type # _to_bf16x2_rn # relu # satfinite # _scale_n2_ue8m0 : PureIntrinsic<[llvm_v2bf16_ty], [llvm_i16_ty, llvm_i16_ty]>; @@ -2115,18 +2122,27 @@ let TargetPrefix = "nvvm" in { // F... [truncated] `````````` </details> https://github.com/llvm/llvm-project/pull/222511 _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
