[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