[llvm] [LV] Add support for absolute difference partial reductions (PR #188043)

Benjamin Maxwell via llvm-commits llvm-commits at lists.llvm.org
Wed Apr 15 03:45:50 PDT 2026


================
@@ -0,0 +1,214 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --check-globals none --version 6
+; RUN: opt -passes=loop-vectorize -mattr=+neon,+dotprod -force-vector-interleave=1 -enable-epilogue-vectorization=false -S < %s | FileCheck %s
+
+target triple = "aarch64-none-unknown-elf"
+
+define i32 @unsigned_absolute_difference(ptr noalias %x, ptr noalias %y) {
+; CHECK-LABEL: define i32 @unsigned_absolute_difference(
+; CHECK-SAME: ptr noalias [[X:%.*]], ptr noalias [[Y:%.*]]) #[[ATTR0:[0-9]+]] {
+; CHECK-NEXT:  [[FOR_BODY:.*:]]
+; CHECK-NEXT:    br label %[[VECTOR_BODY:.*]]
+; CHECK:       [[VECTOR_BODY]]:
+; CHECK-NEXT:    br label %[[EXIT1:.*]]
+; CHECK:       [[EXIT1]]:
+; CHECK-NEXT:    [[IV1:%.*]] = phi i64 [ 0, %[[VECTOR_BODY]] ], [ [[IV_NEXT:%.*]], %[[EXIT1]] ]
+; CHECK-NEXT:    [[VEC_PHI:%.*]] = phi <4 x i32> [ zeroinitializer, %[[VECTOR_BODY]] ], [ [[PARTIAL_REDUCE:%.*]], %[[EXIT1]] ]
+; CHECK-NEXT:    [[X_PTR1:%.*]] = getelementptr inbounds nuw i8, ptr [[X]], i64 [[IV1]]
+; CHECK-NEXT:    [[WIDE_LOAD:%.*]] = load <16 x i8>, ptr [[X_PTR1]], align 1
+; CHECK-NEXT:    [[Y_PTR1:%.*]] = getelementptr inbounds nuw i8, ptr [[Y]], i64 [[IV1]]
+; CHECK-NEXT:    [[WIDE_LOAD1:%.*]] = load <16 x i8>, ptr [[Y_PTR1]], align 1
+; CHECK-NEXT:    [[TMP2:%.*]] = freeze <16 x i8> [[WIDE_LOAD]]
+; CHECK-NEXT:    [[TMP3:%.*]] = freeze <16 x i8> [[WIDE_LOAD1]]
+; CHECK-NEXT:    [[TMP4:%.*]] = call <16 x i8> @llvm.umax.v16i8(<16 x i8> [[TMP2]], <16 x i8> [[TMP3]])
+; CHECK-NEXT:    [[TMP5:%.*]] = call <16 x i8> @llvm.umin.v16i8(<16 x i8> [[TMP2]], <16 x i8> [[TMP3]])
+; CHECK-NEXT:    [[TMP6:%.*]] = sub <16 x i8> [[TMP4]], [[TMP5]]
+; CHECK-NEXT:    [[TMP7:%.*]] = zext <16 x i8> [[TMP6]] to <16 x i32>
+; CHECK-NEXT:    [[PARTIAL_REDUCE]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> [[VEC_PHI]], <16 x i32> [[TMP7]])
+; CHECK-NEXT:    [[IV_NEXT]] = add nuw i64 [[IV1]], 16
+; CHECK-NEXT:    [[EXITCOND_NOT:%.*]] = icmp eq i64 [[IV_NEXT]], 8000
+; CHECK-NEXT:    br i1 [[EXITCOND_NOT]], label %[[EXIT2:.*]], label %[[EXIT1]], !llvm.loop [[LOOP0:![0-9]+]]
+; CHECK:       [[EXIT2]]:
+; CHECK-NEXT:    [[SUM_1_LCSSA:%.*]] = call i32 @llvm.vector.reduce.add.v4i32(<4 x i32> [[PARTIAL_REDUCE]])
+; CHECK-NEXT:    br label %[[EXIT:.*]]
+; CHECK:       [[EXIT]]:
+; CHECK-NEXT:    ret i32 [[SUM_1_LCSSA]]
+;
+entry:
+  br label %for.body
+
+for.body:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %for.body ]
+  %sum.0 = phi i32 [ 0, %entry ], [ %sum.1, %for.body ]
+  %x.ptr = getelementptr inbounds nuw i8, ptr %x, i64 %iv
+  %x.val = load i8, ptr %x.ptr, align 1
+  %ext.x = zext i8 %x.val to i32
+  %y.ptr = getelementptr inbounds nuw i8, ptr %y, i64 %iv
+  %y.val = load i8, ptr %y.ptr, align 1
+  %ext.y = zext i8 %y.val to i32
+  %sub = sub nsw i32 %ext.x, %ext.y
+  %abs.diff = tail call i32 @llvm.abs.i32(i32 %sub, i1 true)
+  %sum.1 = add nuw nsw i32 %abs.diff, %sum.0
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond.not = icmp eq i64 %iv.next, 8000
+  br i1 %exitcond.not, label %exit, label %for.body
+
+exit:
+  ret i32 %sum.1
+}
+
+define i32 @signed_absolute_difference(ptr noalias %x, ptr noalias %y) {
+; CHECK-LABEL: define i32 @signed_absolute_difference(
+; CHECK-SAME: ptr noalias [[X:%.*]], ptr noalias [[Y:%.*]]) #[[ATTR0]] {
+; CHECK-NEXT:  [[FOR_BODY:.*:]]
+; CHECK-NEXT:    br label %[[VECTOR_BODY:.*]]
+; CHECK:       [[VECTOR_BODY]]:
+; CHECK-NEXT:    br label %[[EXIT1:.*]]
+; CHECK:       [[EXIT1]]:
+; CHECK-NEXT:    [[IV1:%.*]] = phi i64 [ 0, %[[VECTOR_BODY]] ], [ [[IV_NEXT:%.*]], %[[EXIT1]] ]
+; CHECK-NEXT:    [[VEC_PHI:%.*]] = phi <4 x i32> [ zeroinitializer, %[[VECTOR_BODY]] ], [ [[PARTIAL_REDUCE:%.*]], %[[EXIT1]] ]
+; CHECK-NEXT:    [[X_PTR1:%.*]] = getelementptr inbounds nuw i8, ptr [[X]], i64 [[IV1]]
+; CHECK-NEXT:    [[WIDE_LOAD:%.*]] = load <16 x i8>, ptr [[X_PTR1]], align 1
+; CHECK-NEXT:    [[Y_PTR1:%.*]] = getelementptr inbounds nuw i8, ptr [[Y]], i64 [[IV1]]
+; CHECK-NEXT:    [[WIDE_LOAD1:%.*]] = load <16 x i8>, ptr [[Y_PTR1]], align 1
+; CHECK-NEXT:    [[TMP2:%.*]] = freeze <16 x i8> [[WIDE_LOAD]]
+; CHECK-NEXT:    [[TMP3:%.*]] = freeze <16 x i8> [[WIDE_LOAD1]]
+; CHECK-NEXT:    [[TMP4:%.*]] = call <16 x i8> @llvm.smax.v16i8(<16 x i8> [[TMP2]], <16 x i8> [[TMP3]])
+; CHECK-NEXT:    [[TMP5:%.*]] = call <16 x i8> @llvm.smin.v16i8(<16 x i8> [[TMP2]], <16 x i8> [[TMP3]])
+; CHECK-NEXT:    [[TMP6:%.*]] = sub <16 x i8> [[TMP4]], [[TMP5]]
+; CHECK-NEXT:    [[TMP7:%.*]] = zext <16 x i8> [[TMP6]] to <16 x i32>
+; CHECK-NEXT:    [[PARTIAL_REDUCE]] = call <4 x i32> @llvm.vector.partial.reduce.add.v4i32.v16i32(<4 x i32> [[VEC_PHI]], <16 x i32> [[TMP7]])
+; CHECK-NEXT:    [[IV_NEXT]] = add nuw i64 [[IV1]], 16
+; CHECK-NEXT:    [[EXITCOND_NOT:%.*]] = icmp eq i64 [[IV_NEXT]], 8000
+; CHECK-NEXT:    br i1 [[EXITCOND_NOT]], label %[[EXIT2:.*]], label %[[EXIT1]], !llvm.loop [[LOOP3:![0-9]+]]
+; CHECK:       [[EXIT2]]:
+; CHECK-NEXT:    [[SUM_1_LCSSA:%.*]] = call i32 @llvm.vector.reduce.add.v4i32(<4 x i32> [[PARTIAL_REDUCE]])
+; CHECK-NEXT:    br label %[[EXIT:.*]]
+; CHECK:       [[EXIT]]:
+; CHECK-NEXT:    ret i32 [[SUM_1_LCSSA]]
+;
+entry:
+  br label %for.body
+
+for.body:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %for.body ]
+  %sum.0 = phi i32 [ 0, %entry ], [ %sum.1, %for.body ]
+  %x.ptr = getelementptr inbounds nuw i8, ptr %x, i64 %iv
+  %x.val = load i8, ptr %x.ptr, align 1
+  %ext.x = sext i8 %x.val to i32
+  %y.ptr = getelementptr inbounds nuw i8, ptr %y, i64 %iv
+  %y.val = load i8, ptr %y.ptr, align 1
+  %ext.y = sext i8 %y.val to i32
+  %sub = sub nsw i32 %ext.x, %ext.y
+  %abs.diff = tail call i32 @llvm.abs.i32(i32 %sub, i1 true)
+  %sum.1 = add nuw nsw i32 %abs.diff, %sum.0
+  %iv.next = add nuw nsw i64 %iv, 1
+  %exitcond.not = icmp eq i64 %iv.next, 8000
+  br i1 %exitcond.not, label %exit, label %for.body
+
+exit:
+  ret i32 %sum.1
+}
+
+; Negative test: Mismatched sign and zero extend.
+define i32 @mismatched_extends(ptr noalias %x, ptr noalias %y) {
+; CHECK-LABEL: define i32 @mismatched_extends(
+; CHECK-SAME: ptr noalias [[X:%.*]], ptr noalias [[Y:%.*]]) #[[ATTR0]] {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    br label %[[VECTOR_PH:.*]]
+; CHECK:       [[VECTOR_PH]]:
+; CHECK-NEXT:    br label %[[VECTOR_BODY:.*]]
+; CHECK:       [[VECTOR_BODY]]:
+; CHECK-NEXT:    [[INDEX:%.*]] = phi i64 [ 0, %[[VECTOR_PH]] ], [ [[INDEX_NEXT:%.*]], %[[VECTOR_BODY]] ]
+; CHECK-NEXT:    [[VEC_PHI:%.*]] = phi <16 x i32> [ zeroinitializer, %[[VECTOR_PH]] ], [ [[TMP6:%.*]], %[[VECTOR_BODY]] ]
+; CHECK-NEXT:    [[TMP0:%.*]] = getelementptr inbounds nuw i8, ptr [[X]], i64 [[INDEX]]
+; CHECK-NEXT:    [[WIDE_LOAD:%.*]] = load <16 x i8>, ptr [[TMP0]], align 1
+; CHECK-NEXT:    [[TMP1:%.*]] = sext <16 x i8> [[WIDE_LOAD]] to <16 x i32>
+; CHECK-NEXT:    [[TMP2:%.*]] = getelementptr inbounds nuw i8, ptr [[Y]], i64 [[INDEX]]
+; CHECK-NEXT:    [[WIDE_LOAD1:%.*]] = load <16 x i8>, ptr [[TMP2]], align 1
+; CHECK-NEXT:    [[TMP3:%.*]] = zext <16 x i8> [[WIDE_LOAD1]] to <16 x i32>
+; CHECK-NEXT:    [[TMP4:%.*]] = sub nsw <16 x i32> [[TMP1]], [[TMP3]]
+; CHECK-NEXT:    [[TMP5:%.*]] = call <16 x i32> @llvm.abs.v16i32(<16 x i32> [[TMP4]], i1 true)
+; CHECK-NEXT:    [[TMP6]] = add <16 x i32> [[TMP5]], [[VEC_PHI]]
+; CHECK-NEXT:    [[INDEX_NEXT]] = add nuw i64 [[INDEX]], 16
+; CHECK-NEXT:    [[TMP7:%.*]] = icmp eq i64 [[INDEX_NEXT]], 8000
+; CHECK-NEXT:    br i1 [[TMP7]], label %[[MIDDLE_BLOCK:.*]], label %[[VECTOR_BODY]], !llvm.loop [[LOOP4:![0-9]+]]
+; CHECK:       [[MIDDLE_BLOCK]]:
+; CHECK-NEXT:    [[TMP8:%.*]] = call i32 @llvm.vector.reduce.add.v16i32(<16 x i32> [[TMP6]])
+; CHECK-NEXT:    br label %[[EXIT:.*]]
+; CHECK:       [[EXIT]]:
+; CHECK-NEXT:    ret i32 [[TMP8]]
+;
+entry:
+  br label %for.body
+
+for.body:
+  %iv = phi i64 [ 0, %entry ], [ %iv.next, %for.body ]
+  %sum.0 = phi i32 [ 0, %entry ], [ %sum.1, %for.body ]
+  %x.ptr = getelementptr inbounds nuw i8, ptr %x, i64 %iv
+  %x.val = load i8, ptr %x.ptr, align 1
+  %ext.x = sext i8 %x.val to i32
+  %y.ptr = getelementptr inbounds nuw i8, ptr %y, i64 %iv
+  %y.val = load i8, ptr %y.ptr, align 1
+  %ext.y = zext i8 %y.val to i32
+  %sub = sub nsw i32 %ext.x, %ext.y
+  %abs.diff = tail call i32 @llvm.abs.i32(i32 %sub, i1 true)
----------------
MacDue wrote:

Added some additional test cases :+1: 

https://github.com/llvm/llvm-project/pull/188043


More information about the llvm-commits mailing list