[llvm] [DAGCombiner] Fold extract_vector(get_active_lane_mask(), 0) to get_active_lane_mask (PR #208978)

via llvm-commits llvm-commits at lists.llvm.org
Sat Jul 11 22:07:43 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-aarch64

@llvm/pr-subscribers-llvm-selectiondag

Author: Usman Nadeem (UsmanNadeem)

<details>
<summary>Changes</summary>

If we are extracting from the start of the vector, then we can just rewrite to `get_active_lane_mask` with a new type and the same operands.

Depends on https://github.com/llvm/llvm-project/pull/208962 and https://github.com/llvm/llvm-project/pull/208977

---
Full diff: https://github.com/llvm/llvm-project/pull/208978.diff


3 Files Affected:

- (modified) llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp (+9) 
- (modified) llvm/test/CodeGen/AArch64/fold-sext-in-reg-predicate-fixed-length.ll (+28) 
- (modified) llvm/test/CodeGen/AArch64/get-active-lane-mask-extract.ll (+138) 


``````````diff
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index ee30495055f47..ce23480fc5c48 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -27869,6 +27869,15 @@ SDValue DAGCombiner::visitEXTRACT_SUBVECTOR(SDNode *N) {
       if (!LegalOperations || TLI.isOperationLegal(ISD::SPLAT_VECTOR, NVT))
         return DAG.getSplatVector(NVT, DL, V.getOperand(0));
 
+  // ty1 extract_vector(ty2 get_active_lane_mask(X, Y), 0) --> ty1
+  // get_active_lane_mask(X, Y)
+  if (ExtIdx == 0 && V.getOpcode() == ISD::GET_ACTIVE_LANE_MASK &&
+      V.hasOneUse() &&
+      (!LegalOperations ||
+       TLI.isOperationLegal(ISD::GET_ACTIVE_LANE_MASK, NVT)))
+    return DAG.getNode(ISD::GET_ACTIVE_LANE_MASK, DL, NVT, V.getOperand(0),
+                       V.getOperand(1));
+
   // extract_subvector(insert_subvector(x,y,c1),c2)
   //  --> extract_subvector(y,c2-c1)
   // iff we're just extracting from the inserted subvector.
diff --git a/llvm/test/CodeGen/AArch64/fold-sext-in-reg-predicate-fixed-length.ll b/llvm/test/CodeGen/AArch64/fold-sext-in-reg-predicate-fixed-length.ll
index 783cd6f693ed8..6d183b2f93534 100644
--- a/llvm/test/CodeGen/AArch64/fold-sext-in-reg-predicate-fixed-length.ll
+++ b/llvm/test/CodeGen/AArch64/fold-sext-in-reg-predicate-fixed-length.ll
@@ -17,3 +17,31 @@ entry:
   %data = call <16 x i8> @llvm.masked.load.v16i8.p0(ptr %p, i32 1, <16 x i1> %mask, <16 x i8> zeroinitializer)
   ret <16 x i8> %data
 }
+
+define void @active_lane_mask_mstore_vscaleX2(ptr %p, i64 %n) #0 {
+; CHECK-LABEL: active_lane_mask_mstore_vscaleX2:
+; CHECK:       // %bb.0: // %entry
+; CHECK-NEXT:    mov z0.h, #123 // =0x7b
+; CHECK-NEXT:    whilelo p0.h, xzr, x1
+; CHECK-NEXT:    st1h { z0.h }, p0, [x0]
+; CHECK-NEXT:    ret
+entry:
+  %mask = call <16 x i1> @llvm.get.active.lane.mask.v16i1(i64 0, i64 %n)
+  call void @llvm.masked.store.v16i16.p0(<16 x i16> splat(i16 123), ptr %p, <16 x i1> %mask)
+  ret void
+}
+
+define void @active_lane_mask_mstore_vscaleX4(ptr %p, i64 %n) #1 {
+; CHECK-LABEL: active_lane_mask_mstore_vscaleX4:
+; CHECK:       // %bb.0:
+; CHECK-NEXT:    mov z0.s, #123 // =0x7b
+; CHECK-NEXT:    whilelo p0.s, xzr, x1
+; CHECK-NEXT:    st1w { z0.s }, p0, [x0]
+; CHECK-NEXT:    ret
+  %mask =  call <16 x i1> @llvm.get.active.lane.mask.v16i1.i64(i64 0, i64 %n)
+  call void @llvm.masked.store.v16i32.p0(<16 x i32> splat(i32 123), ptr %p, <16 x i1> %mask)
+  ret void
+}
+
+attributes #0 = { vscale_range(2,2) }
+attributes #1 = { vscale_range(4,4) }
diff --git a/llvm/test/CodeGen/AArch64/get-active-lane-mask-extract.ll b/llvm/test/CodeGen/AArch64/get-active-lane-mask-extract.ll
index fd8cb06ed02e4..990584c3f2ac8 100644
--- a/llvm/test/CodeGen/AArch64/get-active-lane-mask-extract.ll
+++ b/llvm/test/CodeGen/AArch64/get-active-lane-mask-extract.ll
@@ -598,4 +598,142 @@ if.end:
     ret void
 }
 
+; Test combining of getActiveLaneMask with extract_vector(idx=0) operations.
+
+define void @test_8bit_mask_with_32bit_index_and_trip_count(i32 %i, i32 %n) #0 {
+; CHECK-SVE-LABEL: test_8bit_mask_with_32bit_index_and_trip_count:
+; CHECK-SVE:       // %bb.0:
+; CHECK-SVE-NEXT:    whilelo p0.h, w0, w1
+; CHECK-SVE-NEXT:    // fake_use: $p0
+; CHECK-SVE-NEXT:    ret
+;
+; CHECK-SVE2p1-SME2-LABEL: test_8bit_mask_with_32bit_index_and_trip_count:
+; CHECK-SVE2p1-SME2:       // %bb.0:
+; CHECK-SVE2p1-SME2-NEXT:    whilelo p0.h, w0, w1
+; CHECK-SVE2p1-SME2-NEXT:    // fake_use: $p0
+; CHECK-SVE2p1-SME2-NEXT:    ret
+    %r = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i32(i32 %i, i32 %n)
+    %v0 = call <vscale x 8 x i1> @llvm.vector.extract.nxv8i1.nxv16i1.i64(<vscale x 16 x i1> %r, i64 0)
+    call void (...) @llvm.fake.use(<vscale x 8 x i1> %v0)
+    ret void
+}
+
+define void @test_8bit_mask_extractIdx1_with_32bit_index_and_trip_count(i32 %i, i32 %n) #0 {
+; CHECK-SVE-LABEL: test_8bit_mask_extractIdx1_with_32bit_index_and_trip_count:
+; CHECK-SVE:       // %bb.0:
+; CHECK-SVE-NEXT:    whilelo p0.b, w0, w1
+; CHECK-SVE-NEXT:    punpkhi p0.h, p0.b
+; CHECK-SVE-NEXT:    // fake_use: $p0
+; CHECK-SVE-NEXT:    ret
+;
+; CHECK-SVE2p1-SME2-LABEL: test_8bit_mask_extractIdx1_with_32bit_index_and_trip_count:
+; CHECK-SVE2p1-SME2:       // %bb.0:
+; CHECK-SVE2p1-SME2-NEXT:    whilelo p0.b, w0, w1
+; CHECK-SVE2p1-SME2-NEXT:    punpkhi p0.h, p0.b
+; CHECK-SVE2p1-SME2-NEXT:    // fake_use: $p0
+; CHECK-SVE2p1-SME2-NEXT:    ret
+    %r = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i32(i32 %i, i32 %n)
+    %v0 = call <vscale x 8 x i1> @llvm.vector.extract.nxv8i1.nxv16i1.i64(<vscale x 16 x i1> %r, i64 8)
+    call void (...) @llvm.fake.use(<vscale x 8 x i1> %v0)
+    ret void
+}
+
+define void @test_8bit_mask_with_32bit_index_and_trip_count_multiuse(i32 %i, i32 %n) #0 {
+; CHECK-SVE-LABEL: test_8bit_mask_with_32bit_index_and_trip_count_multiuse:
+; CHECK-SVE:       // %bb.0:
+; CHECK-SVE-NEXT:    whilelo p0.b, w0, w1
+; CHECK-SVE-NEXT:    punpklo p1.h, p0.b
+; CHECK-SVE-NEXT:    // fake_use: $p0
+; CHECK-SVE-NEXT:    // fake_use: $p1
+; CHECK-SVE-NEXT:    ret
+;
+; CHECK-SVE2p1-SME2-LABEL: test_8bit_mask_with_32bit_index_and_trip_count_multiuse:
+; CHECK-SVE2p1-SME2:       // %bb.0:
+; CHECK-SVE2p1-SME2-NEXT:    whilelo p0.b, w0, w1
+; CHECK-SVE2p1-SME2-NEXT:    punpklo p1.h, p0.b
+; CHECK-SVE2p1-SME2-NEXT:    // fake_use: $p0
+; CHECK-SVE2p1-SME2-NEXT:    // fake_use: $p1
+; CHECK-SVE2p1-SME2-NEXT:    ret
+    %r = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i32(i32 %i, i32 %n)
+    %v0 = call <vscale x 8 x i1> @llvm.vector.extract.nxv8i1.nxv16i1.i64(<vscale x 16 x i1> %r, i64 0)
+    call void (...) @llvm.fake.use(<vscale x 16 x i1> %r)
+    call void (...) @llvm.fake.use(<vscale x 8 x i1> %v0)
+    ret void
+}
+
+define void @test_8bit_mask_with_64bit_index_and_trip_count(i64 %i, i64 %n) #0 {
+; CHECK-SVE-LABEL: test_8bit_mask_with_64bit_index_and_trip_count:
+; CHECK-SVE:       // %bb.0:
+; CHECK-SVE-NEXT:    whilelo p0.h, x0, x1
+; CHECK-SVE-NEXT:    // fake_use: $p0
+; CHECK-SVE-NEXT:    ret
+;
+; CHECK-SVE2p1-SME2-LABEL: test_8bit_mask_with_64bit_index_and_trip_count:
+; CHECK-SVE2p1-SME2:       // %bb.0:
+; CHECK-SVE2p1-SME2-NEXT:    whilelo p0.h, x0, x1
+; CHECK-SVE2p1-SME2-NEXT:    // fake_use: $p0
+; CHECK-SVE2p1-SME2-NEXT:    ret
+    %r = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i64(i64 %i, i64 %n)
+    %v0 = call <vscale x 8 x i1> @llvm.vector.extract.nxv8i1.nxv16i1.i64(<vscale x 16 x i1> %r, i64 0)
+    call void (...) @llvm.fake.use(<vscale x 8 x i1> %v0)
+    ret void
+}
+
+define void @test_4bit_mask_with_64bit_index_and_trip_count(i64 %i, i64 %n) #0 {
+; CHECK-SVE-LABEL: test_4bit_mask_with_64bit_index_and_trip_count:
+; CHECK-SVE:       // %bb.0:
+; CHECK-SVE-NEXT:    whilelo p0.s, x0, x1
+; CHECK-SVE-NEXT:    // fake_use: $p0
+; CHECK-SVE-NEXT:    ret
+;
+; CHECK-SVE2p1-SME2-LABEL: test_4bit_mask_with_64bit_index_and_trip_count:
+; CHECK-SVE2p1-SME2:       // %bb.0:
+; CHECK-SVE2p1-SME2-NEXT:    whilelo p0.s, x0, x1
+; CHECK-SVE2p1-SME2-NEXT:    // fake_use: $p0
+; CHECK-SVE2p1-SME2-NEXT:    ret
+    %r = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i64(i64 %i, i64 %n)
+    %v0 = call <vscale x 4 x i1> @llvm.vector.extract.nxv8i1.nxv16i1.i64(<vscale x 16 x i1> %r, i64 0)
+    call void (...) @llvm.fake.use(<vscale x 4 x i1> %v0)
+    ret void
+}
+
+define void @test_4bit_mask_with_2extracts_64bit_index_and_trip_count(i64 %i, i64 %n) #0 {
+; CHECK-SVE-LABEL: test_4bit_mask_with_2extracts_64bit_index_and_trip_count:
+; CHECK-SVE:       // %bb.0:
+; CHECK-SVE-NEXT:    whilelo p0.s, x0, x1
+; CHECK-SVE-NEXT:    // fake_use: $p0
+; CHECK-SVE-NEXT:    ret
+;
+; CHECK-SVE2p1-SME2-LABEL: test_4bit_mask_with_2extracts_64bit_index_and_trip_count:
+; CHECK-SVE2p1-SME2:       // %bb.0:
+; CHECK-SVE2p1-SME2-NEXT:    whilelo p0.s, x0, x1
+; CHECK-SVE2p1-SME2-NEXT:    // fake_use: $p0
+; CHECK-SVE2p1-SME2-NEXT:    ret
+    %r = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i64(i64 %i, i64 %n)
+    %v0 = call <vscale x 8 x i1> @llvm.vector.extract.nxv8i1.nxv16i1.i64(<vscale x 16 x i1> %r, i64 0)
+    %v1 = call <vscale x 4 x i1> @llvm.vector.extract.nxv8i1.nxv16i1.i64(<vscale x 8 x i1> %v0, i64 0)
+    call void (...) @llvm.fake.use(<vscale x 4 x i1> %v1)
+    ret void
+}
+
+define void @test_1bit_mask_with_64bit_index_and_trip_count(i64 %i, i64 %n) #0 {
+; CHECK-SVE-LABEL: test_1bit_mask_with_64bit_index_and_trip_count:
+; CHECK-SVE:       // %bb.0:
+; CHECK-SVE-NEXT:    whilelo p0.d, x0, x1
+; CHECK-SVE-NEXT:    punpklo p0.h, p0.b
+; CHECK-SVE-NEXT:    // fake_use: $p0
+; CHECK-SVE-NEXT:    ret
+;
+; CHECK-SVE2p1-SME2-LABEL: test_1bit_mask_with_64bit_index_and_trip_count:
+; CHECK-SVE2p1-SME2:       // %bb.0:
+; CHECK-SVE2p1-SME2-NEXT:    whilelo p0.d, x0, x1
+; CHECK-SVE2p1-SME2-NEXT:    punpklo p0.h, p0.b
+; CHECK-SVE2p1-SME2-NEXT:    // fake_use: $p0
+; CHECK-SVE2p1-SME2-NEXT:    ret
+    %r = call <vscale x 16 x i1> @llvm.get.active.lane.mask.nxv16i1.i64(i64 %i, i64 %n)
+    %v0 = call <vscale x 1 x i1> @llvm.vector.extract.nxv8i1.nxv16i1.i64(<vscale x 16 x i1> %r, i64 0)
+    call void (...) @llvm.fake.use(<vscale x 1 x i1> %v0)
+    ret void
+}
+
 attributes #0 = { nounwind }

``````````

</details>


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


More information about the llvm-commits mailing list