[llvm] [NVPTX] Add support for f32x2 mixed-precision add/sub (PR #221957)
via llvm-commits
llvm-commits at lists.llvm.org
Tue Sep 8 03:59:39 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-nvptx
Author: Srinivasa Ravi (Wolfram70)
<details>
<summary>Changes</summary>
This change adds support for mixed precision addition and
subtraction of `f16x2` and `bf16x2` with `f32x2`, where the
following upconverting patterns:
```
%e = fpext <2 x half> %h to <2 x float>
%res = fp-operation(%e, ...)
...
%e = fpext <2 x bfloat> %b to <2 x float>
%res = fp-operation(%e, ...)
where the fp-operation can be any of:
- fadd
- fsub
- llvm.nvvm.fadd.v2f32
```
are lowered to `add/sub.{rnd}.f32x2.{f16x2/bf16x2}.f32x2`, and the
following downconverting pattern:
```
%sum = llvm.nvvm.fadd{.ftz}.v2f32(%a, %b, rz)
%lo = extractelement <2 x float> %sum, i32 0
%hi = extractelement <2 x float> %sum, i32 1
%res = llvm.nvvm.ff2{f16x2/bf16x2}.rz(%hi, %lo, false)
```
is lowered to `add/sub.rz{.ftz}.{f16x2/bf16x2}.f32x2.f32x2`. These
instructions combine the conversion and the addition into one
instruction from `sm_107f` onwards. `fneg` followed by one of the
above addition patterns is lowered to the corresponding `sub`
instruction.
The tests have been verified through ptxas-13.4.
PTX spec reference: https://docs.nvidia.com/cuda/developer-preview/13.4/parallel-thread-execution/index.html#mixed-precision-floating-point-instructions-add
---
Patch is 89.38 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/221957.diff
5 Files Affected:
- (modified) llvm/lib/Target/NVPTX/NVPTXIntrinsics.td (+119)
- (added) llvm/test/CodeGen/NVPTX/mixed-precision-add-f32x2-downconvert.ll (+330)
- (added) llvm/test/CodeGen/NVPTX/mixed-precision-add-f32x2-upconvert.ll (+675)
- (added) llvm/test/CodeGen/NVPTX/mixed-precision-sub-f32x2-downconvert.ll (+346)
- (added) llvm/test/CodeGen/NVPTX/mixed-precision-sub-f32x2-upconvert.ll (+524)
``````````diff
diff --git a/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td b/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
index 220ef64732830..a653c70050e8b 100644
--- a/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
+++ b/llvm/lib/Target/NVPTX/NVPTXIntrinsics.td
@@ -2495,6 +2495,19 @@ def INT_NVVM_ADD_D :
F_MATH_2_RNDOP_TY<"add.${rnd}.f64", F64RT, int_nvvm_fadd>;
// mixed precision
+
+// FP_EXTEND to v2f32 is scalarized before isel, leaving one of two shapes
+// behind depending on whether the source lanes share a register.
+def fpextend_v2f32_from_packed : PatFrag<(ops node:$a),
+ (v2f32 (build_vector
+ (f32 (fpextend (extractelt node:$a, 0))),
+ (f32 (fpextend (extractelt node:$a, 1)))))>;
+
+def fpextend_v2f32_from_lanes : PatFrag<(ops node:$a0, node:$a1),
+ (v2f32 (build_vector
+ (f32 (fpextend node:$a0)),
+ (f32 (fpextend node:$a1))))>;
+
foreach rnd = FPRoundingModes in {
defvar rnd_imm = !cast<TImmLeaf>(StrJoin<"_", ["fp_rnd", rnd, "imm"]>.ret);
@@ -2508,6 +2521,24 @@ foreach rnd = FPRoundingModes in {
(f32 (fpextend type:$a)),
f32:$b, rnd_imm))]>,
Requires<[SM100]>;
+
+ foreach t = [F16X2RT, BF16X2RT] in {
+ def INT_NVVM_MIXED_ADD_ # rnd # _f32x2_ # t.PtxType :
+ BasicNVPTXInst<(outs B64:$dst), (ins t.RC:$a, B64:$b),
+ StrJoin<".", ["add", rnd, "f32x2", t.PtxType, "f32x2"]>.ret,
+ [(set v2f32:$dst,
+ (int_nvvm_fadd (fpextend_v2f32_from_packed t.Ty:$a), v2f32:$b,
+ rnd_imm))]>,
+ Requires<[SM107f]>;
+
+ def : Pat<(v2f32 (int_nvvm_fadd
+ (fpextend_v2f32_from_lanes t.Ty.ElementType:$a0,
+ t.Ty.ElementType:$a1),
+ v2f32:$b, rnd_imm)),
+ (!cast<Instruction>("INT_NVVM_MIXED_ADD_" # rnd # "_f32x2_" #
+ t.PtxType) (t.Ty (V2I16toI32 $a0, $a1)), $b)>,
+ Requires<[SM107f]>;
+ }
}
// Pattern for fadd when there is no FTZ flag
@@ -2518,6 +2549,40 @@ let Predicates = [SM100, doNoF32FTZ] in {
(INT_NVVM_MIXED_ADD_rn_f32_bf16 B16:$a, B32:$b)>;
}
+let Predicates = [SM107f, doNoF32FTZ] in
+ foreach t = [F16X2RT, BF16X2RT] in {
+ defvar Inst = !cast<Instruction>("INT_NVVM_MIXED_ADD_rn_f32x2_" # t.PtxType);
+
+ def : Pat<(v2f32 (fadd (fpextend_v2f32_from_packed t.Ty:$a), v2f32:$b)),
+ (Inst $a, $b)>;
+
+ def : Pat<(v2f32 (fadd (fpextend_v2f32_from_lanes t.Ty.ElementType:$a0,
+ t.Ty.ElementType:$a1),
+ v2f32:$b)),
+ (Inst (t.Ty (V2I16toI32 $a0, $a1)), $b)>;
+ }
+
+
+// mixed precision - downconverting
+
+// f16x2 - rz rounding mode, only with ftz
+// bf16x2 - rz rounding mode, without ftz
+foreach t = [F16X2RT, BF16X2RT] in {
+ defvar ftz = !if(!eq(t, F16X2RT), "ftz", "");
+ defvar AddOp = !cast<Intrinsic>(StrJoin<"_", ["int_nvvm_fadd", ftz]>.ret);
+ defvar CvtOp = !cast<Intrinsic>("int_nvvm_ff2" # t.PtxType # "_rz");
+
+ def INT_NVVM_MIXED_ADD_rz_ # t.PtxType # _f32x2 :
+ BasicNVPTXInst<(outs t.RC:$dst), (ins B64:$a, B64:$b),
+ StrJoin<".", ["add.rz", ftz, t.PtxType, "f32x2", "f32x2"]>.ret,
+ // The cvt packs its first argument into the high half of the result.
+ [(set t.Ty:$dst,
+ (CvtOp (extractelt (AddOp v2f32:$a, v2f32:$b, fp_rnd_rz_imm), 1),
+ (extractelt (AddOp v2f32:$a, v2f32:$b, fp_rnd_rz_imm), 0),
+ /*pzo=*/0))]>,
+ Requires<[SM107f]>;
+}
+
//
// Sub
//
@@ -2579,6 +2644,7 @@ def INT_NVVM_SUB_D :
[(set f64:$dst, (int_nvvm_fadd f64:$a, (f64 (fneg f64:$b)), timm:$rnd))]>;
// mixed precision
+
foreach rnd = FPRoundingModes in {
defvar rnd_imm = !cast<TImmLeaf>(StrJoin<"_", ["fp_rnd", rnd, "imm"]>.ret);
@@ -2592,6 +2658,26 @@ foreach rnd = FPRoundingModes in {
(f32 (fpextend type:$a)),
(f32 (fneg f32:$b)), rnd_imm))]>,
Requires<[SM100]>;
+
+ foreach t = [F16X2RT, BF16X2RT] in {
+ // combineFAddWithNeg has already folded the fneg into a sub node.
+ defvar SubOp = !cast<SDNode>("sub_" # rnd);
+
+ def INT_NVVM_MIXED_SUB_ # rnd # _f32x2_ # t.PtxType :
+ BasicNVPTXInst<(outs B64:$dst), (ins t.RC:$a, B64:$b),
+ StrJoin<".", ["sub", rnd, "f32x2", t.PtxType, "f32x2"]>.ret,
+ [(set v2f32:$dst,
+ (SubOp (fpextend_v2f32_from_packed t.Ty:$a), v2f32:$b))]>,
+ Requires<[SM107f]>;
+
+ def : Pat<(v2f32 (SubOp
+ (fpextend_v2f32_from_lanes t.Ty.ElementType:$a0,
+ t.Ty.ElementType:$a1),
+ v2f32:$b)),
+ (!cast<Instruction>("INT_NVVM_MIXED_SUB_" # rnd # "_f32x2_" #
+ t.PtxType) (t.Ty (V2I16toI32 $a0, $a1)), $b)>,
+ Requires<[SM107f]>;
+ }
}
// Pattern for fsub when there is no FTZ flag
@@ -2602,6 +2688,39 @@ let Predicates = [SM100, doNoF32FTZ] in {
(INT_NVVM_MIXED_SUB_rn_f32_bf16 B16:$a, B32:$b)>;
}
+let Predicates = [SM107f, doNoF32FTZ] in
+ foreach t = [F16X2RT, BF16X2RT] in {
+ defvar Inst = !cast<Instruction>("INT_NVVM_MIXED_SUB_rn_f32x2_" # t.PtxType);
+
+ def : Pat<(v2f32 (fsub (fpextend_v2f32_from_packed t.Ty:$a), v2f32:$b)),
+ (Inst $a, $b)>;
+
+ def : Pat<(v2f32 (fsub (fpextend_v2f32_from_lanes t.Ty.ElementType:$a0,
+ t.Ty.ElementType:$a1),
+ v2f32:$b)),
+ (Inst (t.Ty (V2I16toI32 $a0, $a1)), $b)>;
+ }
+
+// mixed precision - downconverting
+
+// f16x2 - rz rounding mode, only with ftz
+// bf16x2 - rz rounding mode, without ftz
+foreach t = [F16X2RT, BF16X2RT] in {
+ defvar ftz = !if(!eq(t, F16X2RT), "ftz", "");
+ defvar SubOp = !cast<SDNode>(StrJoin<"_", ["sub", "rz", ftz]>.ret);
+ defvar CvtOp = !cast<Intrinsic>("int_nvvm_ff2" # t.PtxType # "_rz");
+
+ def INT_NVVM_MIXED_SUB_rz_ # t.PtxType # _f32x2 :
+ BasicNVPTXInst<(outs t.RC:$dst), (ins B64:$a, B64:$b),
+ StrJoin<".", ["sub.rz", ftz, t.PtxType, "f32x2", "f32x2"]>.ret,
+ // The cvt packs its first argument into the high half of the result.
+ [(set t.Ty:$dst,
+ (CvtOp (extractelt (SubOp v2f32:$a, v2f32:$b), 1),
+ (extractelt (SubOp v2f32:$a, v2f32:$b), 0),
+ /*pzo=*/0))]>,
+ Requires<[SM107f]>;
+}
+
//
// BFIND
//
diff --git a/llvm/test/CodeGen/NVPTX/mixed-precision-add-f32x2-downconvert.ll b/llvm/test/CodeGen/NVPTX/mixed-precision-add-f32x2-downconvert.ll
new file mode 100644
index 0000000000000..9dbfe1e3ab014
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/mixed-precision-add-f32x2-downconvert.ll
@@ -0,0 +1,330 @@
+; 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 %}
+
+;
+; F16x2
+;
+
+define <2 x half> @add_rz_f16x2_f32x2(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_f16x2_f32x2(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<2>;
+; CHECK-NEXT: .reg .b64 %rd<3>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_f16x2_f32x2_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_f16x2_f32x2_param_1];
+; CHECK-NEXT: add.rz.ftz.f16x2.f32x2.f32x2 %r1, %rd1, %rd2;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r1;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x half> %r
+}
+
+define <2 x half> @add_rz_f16x2_f32x2_commuted(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_f16x2_f32x2_commuted(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<2>;
+; CHECK-NEXT: .reg .b64 %rd<3>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_f16x2_f32x2_commuted_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_f16x2_f32x2_commuted_param_1];
+; CHECK-NEXT: add.rz.ftz.f16x2.f32x2.f32x2 %r1, %rd2, %rd1;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r1;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %b, <2 x float> %a, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x half> %r
+}
+
+define <2 x half> @add_rz_f16x2_f32x2_extra_use(<2 x float> %a, <2 x float> %b, ptr %p) {
+; CHECK-LABEL: add_rz_f16x2_f32x2_extra_use(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<2>;
+; CHECK-NEXT: .reg .b64 %rd<5>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_f16x2_f32x2_extra_use_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_f16x2_f32x2_extra_use_param_1];
+; CHECK-NEXT: add.rz.ftz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: ld.param::func.b64 %rd4, [add_rz_f16x2_f32x2_extra_use_param_2];
+; CHECK-NEXT: st.b64 [%rd4], %rd3;
+; CHECK-NEXT: add.rz.ftz.f16x2.f32x2.f32x2 %r1, %rd1, %rd2;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r1;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ store <2 x float> %sum, ptr %p
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x half> %r
+}
+
+; invalid patterns
+
+define <2 x half> @add_rz_f16x2_f32x2_swapped_halves(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_f16x2_f32x2_swapped_halves(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<4>;
+; CHECK-NEXT: .reg .b64 %rd<4>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_f16x2_f32x2_swapped_halves_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_f16x2_f32x2_swapped_halves_param_1];
+; CHECK-NEXT: add.rz.ftz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: mov.b64 {%r1, %r2}, %rd3;
+; CHECK-NEXT: cvt.rz.f16x2.f32 %r3, %r1, %r2;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rz(float %lo, float %hi, i1 false)
+ ret <2 x half> %r
+}
+
+define <2 x half> @add_rz_f16x2_f32x2_no_ftz(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_f16x2_f32x2_no_ftz(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<4>;
+; CHECK-NEXT: .reg .b64 %rd<4>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_f16x2_f32x2_no_ftz_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_f16x2_f32x2_no_ftz_param_1];
+; CHECK-NEXT: add.rz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: mov.b64 {%r1, %r2}, %rd3;
+; CHECK-NEXT: cvt.rz.f16x2.f32 %r3, %r2, %r1;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x half> %r
+}
+
+define <2 x half> @add_rn_f16x2_f32x2(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rn_f16x2_f32x2(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<4>;
+; CHECK-NEXT: .reg .b64 %rd<4>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rn_f16x2_f32x2_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rn_f16x2_f32x2_param_1];
+; CHECK-NEXT: add.rn.ftz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: mov.b64 {%r1, %r2}, %rd3;
+; CHECK-NEXT: cvt.rz.f16x2.f32 %r3, %r2, %r1;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %a, <2 x float> %b, i32 1)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x half> %r
+}
+
+define <2 x half> @add_rz_f16x2_f32x2_rn_convert(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_f16x2_f32x2_rn_convert(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<4>;
+; CHECK-NEXT: .reg .b64 %rd<4>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_f16x2_f32x2_rn_convert_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_f16x2_f32x2_rn_convert_param_1];
+; CHECK-NEXT: add.rz.ftz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: mov.b64 {%r1, %r2}, %rd3;
+; CHECK-NEXT: cvt.rn.f16x2.f32 %r3, %r2, %r1;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rn(float %hi, float %lo, i1 false)
+ ret <2 x half> %r
+}
+
+define <2 x half> @add_rz_f16x2_f32x2_relu(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_f16x2_f32x2_relu(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<4>;
+; CHECK-NEXT: .reg .b64 %rd<4>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_f16x2_f32x2_relu_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_f16x2_f32x2_relu_param_1];
+; CHECK-NEXT: add.rz.ftz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: mov.b64 {%r1, %r2}, %rd3;
+; CHECK-NEXT: cvt.rz.relu.f16x2.f32 %r3, %r2, %r1;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rz.relu(float %hi, float %lo, i1 false)
+ ret <2 x half> %r
+}
+
+define <2 x half> @add_rz_f16x2_f32x2_distinct_adds(<2 x float> %a, <2 x float> %b, <2 x float> %c) {
+; CHECK-LABEL: add_rz_f16x2_f32x2_distinct_adds(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<4>;
+; CHECK-NEXT: .reg .b64 %rd<6>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_f16x2_f32x2_distinct_adds_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_f16x2_f32x2_distinct_adds_param_1];
+; CHECK-NEXT: add.rz.ftz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: ld.param::func.b64 %rd4, [add_rz_f16x2_f32x2_distinct_adds_param_2];
+; CHECK-NEXT: add.rz.ftz.f32x2 %rd5, %rd1, %rd4;
+; CHECK-NEXT: mov.b64 {%r1, _}, %rd3;
+; CHECK-NEXT: mov.b64 {_, %r2}, %rd5;
+; CHECK-NEXT: cvt.rz.f16x2.f32 %r3, %r2, %r1;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3;
+; CHECK-NEXT: ret;
+ %sum1 = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ %sum2 = call <2 x float> @llvm.nvvm.fadd.ftz.v2f32(<2 x float> %a, <2 x float> %c, i32 0)
+ %lo = extractelement <2 x float> %sum1, i32 0
+ %hi = extractelement <2 x float> %sum2, i32 1
+ %r = call <2 x half> @llvm.nvvm.ff2f16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x half> %r
+}
+
+;
+; BF16x2
+;
+
+define <2 x bfloat> @add_rz_bf16x2_f32x2(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_bf16x2_f32x2(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<2>;
+; CHECK-NEXT: .reg .b64 %rd<3>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_bf16x2_f32x2_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_bf16x2_f32x2_param_1];
+; CHECK-NEXT: add.rz.bf16x2.f32x2.f32x2 %r1, %rd1, %rd2;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r1;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x bfloat> @llvm.nvvm.ff2bf16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x bfloat> %r
+}
+
+define <2 x bfloat> @add_rz_bf16x2_f32x2_commuted(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_bf16x2_f32x2_commuted(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<2>;
+; CHECK-NEXT: .reg .b64 %rd<3>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_bf16x2_f32x2_commuted_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_bf16x2_f32x2_commuted_param_1];
+; CHECK-NEXT: add.rz.bf16x2.f32x2.f32x2 %r1, %rd2, %rd1;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r1;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.v2f32(<2 x float> %b, <2 x float> %a, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x bfloat> @llvm.nvvm.ff2bf16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x bfloat> %r
+}
+
+define <2 x bfloat> @add_rz_bf16x2_f32x2_extra_use(<2 x float> %a, <2 x float> %b, ptr %p) {
+; CHECK-LABEL: add_rz_bf16x2_f32x2_extra_use(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<2>;
+; CHECK-NEXT: .reg .b64 %rd<5>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_bf16x2_f32x2_extra_use_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_bf16x2_f32x2_extra_use_param_1];
+; CHECK-NEXT: add.rz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: ld.param::func.b64 %rd4, [add_rz_bf16x2_f32x2_extra_use_param_2];
+; CHECK-NEXT: st.b64 [%rd4], %rd3;
+; CHECK-NEXT: add.rz.bf16x2.f32x2.f32x2 %r1, %rd1, %rd2;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r1;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ store <2 x float> %sum, ptr %p
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x bfloat> @llvm.nvvm.ff2bf16x2.rz(float %hi, float %lo, i1 false)
+ ret <2 x bfloat> %r
+}
+
+; invalid patterns
+
+define <2 x bfloat> @add_rz_bf16x2_f32x2_swapped_halves(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_bf16x2_f32x2_swapped_halves(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<4>;
+; CHECK-NEXT: .reg .b64 %rd<4>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_bf16x2_f32x2_swapped_halves_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_bf16x2_f32x2_swapped_halves_param_1];
+; CHECK-NEXT: add.rz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: mov.b64 {%r1, %r2}, %rd3;
+; CHECK-NEXT: cvt.rz.bf16x2.f32 %r3, %r1, %r2;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float> @llvm.nvvm.fadd.v2f32(<2 x float> %a, <2 x float> %b, i32 0)
+ %lo = extractelement <2 x float> %sum, i32 0
+ %hi = extractelement <2 x float> %sum, i32 1
+ %r = call <2 x bfloat> @llvm.nvvm.ff2bf16x2.rz(float %lo, float %hi, i1 false)
+ ret <2 x bfloat> %r
+}
+
+define <2 x bfloat> @add_rz_bf16x2_f32x2_ftz(<2 x float> %a, <2 x float> %b) {
+; CHECK-LABEL: add_rz_bf16x2_f32x2_ftz(
+; CHECK: {
+; CHECK-NEXT: .reg .b32 %r<4>;
+; CHECK-NEXT: .reg .b64 %rd<4>;
+; CHECK-EMPTY:
+; CHECK-NEXT: // %bb.0:
+; CHECK-NEXT: ld.param::func.b64 %rd1, [add_rz_bf16x2_f32x2_ftz_param_0];
+; CHECK-NEXT: ld.param::func.b64 %rd2, [add_rz_bf16x2_f32x2_ftz_param_1];
+; CHECK-NEXT: add.rz.ftz.f32x2 %rd3, %rd1, %rd2;
+; CHECK-NEXT: mov.b64 {%r1, %r2}, %rd3;
+; CHECK-NEXT: cvt.rz.bf16x2.f32 %r3, %r2, %r1;
+; CHECK-NEXT: st.param::func.b32 [func_retval0], %r3;
+; CHECK-NEXT: ret;
+ %sum = call <2 x float>...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/221957
More information about the llvm-commits
mailing list