[llvm] [AArch64] Mark masked load/store data types that get legalized to scalars illegal (PR #228122)

Mattéo Rizza Murgier via llvm-commits llvm-commits at lists.llvm.org
Thu Oct 1 09:07:22 PDT 2026


https://github.com/matteo-rm created https://github.com/llvm/llvm-project/pull/228122

Masked load/stores cannot handle data types that get scalarized during type legalization. This PR marks them illegal in AArch64 TTI so they get expanded before reaching SDAG.

Fixes e.g. the crash:

```llvm
; llc -mtriple=aarch64 -mattr=+sve crash.ll

define <1 x half> @masked_load_v1f16(ptr %p, <1 x i1> %mask) vscale_range(2,2) {
  %load = call <1 x half> @llvm.masked.load.v1f16.p0(ptr %p, <1 x i1> %mask, <1 x half> poison)
  ret <1 x half> %load
}
```

>From 866695c19f5353458a0ab74a9f32cca929a64f3d Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?Matt=C3=A9o=20Rizza=20Murgier?=
 <matteo.rizza-murgier at sipearl.com>
Date: Thu, 1 Oct 2026 13:58:15 +0200
Subject: [PATCH] [AArch64] Mark masked load/store data types that get
 legalized to scalars illegal

---
 .../AArch64/AArch64TargetTransformInfo.h      |  6 +++
 .../AArch64/expand-masked-load.ll             | 30 ++++++++++++++
 .../AArch64/expand-masked-store.ll            | 39 +++++++++++++++++++
 3 files changed, 75 insertions(+)

diff --git a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.h b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.h
index d090f69c1476a4b..2ea69433b28d2bb 100644
--- a/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.h
+++ b/llvm/lib/Target/AArch64/AArch64TargetTransformInfo.h
@@ -321,6 +321,12 @@ class AArch64TTIImpl final : public BasicTTIImplBase<AArch64TTIImpl> {
         return false; // Fall back to scalarization of masked operations.
     }
 
+    // Type legalization cannot scalarize the data of a masked load/store.
+    // Reject types that would be legalized to a scalar (e.g. <1 x half>).
+    if (isa<FixedVectorType>(DataType) &&
+        !getTypeLegalizationCost(DataType).second.isVector())
+      return false;
+
     return isElementTypeLegalForScalableVector(DataType->getScalarType());
   }
 
diff --git a/llvm/test/Transforms/ScalarizeMaskedMemIntrin/AArch64/expand-masked-load.ll b/llvm/test/Transforms/ScalarizeMaskedMemIntrin/AArch64/expand-masked-load.ll
index c0055aff4df705b..e08fbc987858b0e 100644
--- a/llvm/test/Transforms/ScalarizeMaskedMemIntrin/AArch64/expand-masked-load.ll
+++ b/llvm/test/Transforms/ScalarizeMaskedMemIntrin/AArch64/expand-masked-load.ll
@@ -221,6 +221,36 @@ define <2 x i48> @scalarize_v2i48(ptr %p, <2 x i1> %mask, <2 x i48> %passthru) {
   ret <2 x i48> %ret
 }
 
+define <1 x half> @scalarize_v1f16(ptr %p, <1 x i1> %mask, <1 x half> %passthru) vscale_range(2,2) {
+; CHECK-LE-COMMON-LABEL: @scalarize_v1f16(
+; CHECK-LE-COMMON-NEXT:    [[TMP1:%.*]] = extractelement <1 x i1> [[MASK:%.*]], i64 0
+; CHECK-LE-COMMON-NEXT:    br i1 [[TMP1]], label [[COND_LOAD:%.*]], label [[ELSE:%.*]]
+; CHECK-LE-COMMON:       cond.load:
+; CHECK-LE-COMMON-NEXT:    [[TMP2:%.*]] = getelementptr inbounds half, ptr [[P:%.*]], i32 0
+; CHECK-LE-COMMON-NEXT:    [[TMP3:%.*]] = load half, ptr [[TMP2]], align 2
+; CHECK-LE-COMMON-NEXT:    [[TMP4:%.*]] = insertelement <1 x half> [[PASSTHRU:%.*]], half [[TMP3]], i64 0
+; CHECK-LE-COMMON-NEXT:    br label [[ELSE]]
+; CHECK-LE-COMMON:       else:
+; CHECK-LE-COMMON-NEXT:    [[RES_PHI_ELSE:%.*]] = phi <1 x half> [ [[TMP4]], [[COND_LOAD]] ], [ [[PASSTHRU]], [[TMP0:%.*]] ]
+; CHECK-LE-COMMON-NEXT:    ret <1 x half> [[RES_PHI_ELSE]]
+;
+; CHECK-BE-LABEL: @scalarize_v1f16(
+; CHECK-BE-NEXT:    [[TMP1:%.*]] = extractelement <1 x i1> [[MASK:%.*]], i64 0
+; CHECK-BE-NEXT:    br i1 [[TMP1]], label [[COND_LOAD:%.*]], label [[ELSE:%.*]]
+; CHECK-BE:       cond.load:
+; CHECK-BE-NEXT:    [[TMP2:%.*]] = getelementptr inbounds half, ptr [[P:%.*]], i32 0
+; CHECK-BE-NEXT:    [[TMP3:%.*]] = load half, ptr [[TMP2]], align 2
+; CHECK-BE-NEXT:    [[TMP4:%.*]] = insertelement <1 x half> [[PASSTHRU:%.*]], half [[TMP3]], i64 0
+; CHECK-BE-NEXT:    br label [[ELSE]]
+; CHECK-BE:       else:
+; CHECK-BE-NEXT:    [[RES_PHI_ELSE:%.*]] = phi <1 x half> [ [[TMP4]], [[COND_LOAD]] ], [ [[PASSTHRU]], [[TMP0:%.*]] ]
+; CHECK-BE-NEXT:    ret <1 x half> [[RES_PHI_ELSE]]
+;
+  %ret = call <1 x half> @llvm.masked.load.v1f16.p0(ptr %p, i32 2, <1 x i1> %mask, <1 x half> %passthru)
+  ret <1 x half> %ret
+}
+
+declare <1 x half> @llvm.masked.load.v1f16.p0(ptr, i32, <1 x i1>, <1 x half>)
 declare <2 x i24> @llvm.masked.load.v2i24.p0(ptr, i32, <2 x i1>, <2 x i24>)
 declare <2 x i48> @llvm.masked.load.v2i48.p0(ptr, i32, <2 x i1>, <2 x i48>)
 declare <2 x i64> @llvm.masked.load.v2i64.p0(ptr, i32, <2 x i1>, <2 x i64>)
diff --git a/llvm/test/Transforms/ScalarizeMaskedMemIntrin/AArch64/expand-masked-store.ll b/llvm/test/Transforms/ScalarizeMaskedMemIntrin/AArch64/expand-masked-store.ll
index a7f2ac5128a270f..5d79d19da8d0685 100644
--- a/llvm/test/Transforms/ScalarizeMaskedMemIntrin/AArch64/expand-masked-store.ll
+++ b/llvm/test/Transforms/ScalarizeMaskedMemIntrin/AArch64/expand-masked-store.ll
@@ -109,4 +109,43 @@ define void @scalarize_v2i64_const_mask(ptr %p, <2 x i64> %data) {
   ret void
 }
 
+define void @scalarize_v1f16(ptr %p, <1 x i1> %mask, <1 x half> %data) vscale_range(2,2) {
+; CHECK-LE-LABEL: @scalarize_v1f16(
+; CHECK-LE-NEXT:    [[TMP1:%.*]] = extractelement <1 x i1> [[MASK:%.*]], i64 0
+; CHECK-LE-NEXT:    br i1 [[TMP1]], label [[COND_STORE:%.*]], label [[ELSE:%.*]]
+; CHECK-LE:       cond.store:
+; CHECK-LE-NEXT:    [[TMP2:%.*]] = extractelement <1 x half> [[DATA:%.*]], i64 0
+; CHECK-LE-NEXT:    [[TMP3:%.*]] = getelementptr inbounds half, ptr [[P:%.*]], i32 0
+; CHECK-LE-NEXT:    store half [[TMP2]], ptr [[TMP3]], align 2
+; CHECK-LE-NEXT:    br label [[ELSE]]
+; CHECK-LE:       else:
+; CHECK-LE-NEXT:    ret void
+;
+; CHECK-SVE-LE-LABEL: @scalarize_v1f16(
+; CHECK-SVE-LE-NEXT:    [[TMP1:%.*]] = extractelement <1 x i1> [[MASK:%.*]], i64 0
+; CHECK-SVE-LE-NEXT:    br i1 [[TMP1]], label [[COND_STORE:%.*]], label [[ELSE:%.*]]
+; CHECK-SVE-LE:       cond.store:
+; CHECK-SVE-LE-NEXT:    [[TMP2:%.*]] = extractelement <1 x half> [[DATA:%.*]], i64 0
+; CHECK-SVE-LE-NEXT:    [[TMP3:%.*]] = getelementptr inbounds half, ptr [[P:%.*]], i32 0
+; CHECK-SVE-LE-NEXT:    store half [[TMP2]], ptr [[TMP3]], align 2
+; CHECK-SVE-LE-NEXT:    br label [[ELSE]]
+; CHECK-SVE-LE:       else:
+; CHECK-SVE-LE-NEXT:    ret void
+;
+; CHECK-BE-LABEL: @scalarize_v1f16(
+; CHECK-BE-NEXT:    [[TMP1:%.*]] = extractelement <1 x i1> [[MASK:%.*]], i64 0
+; CHECK-BE-NEXT:    br i1 [[TMP1]], label [[COND_STORE:%.*]], label [[ELSE:%.*]]
+; CHECK-BE:       cond.store:
+; CHECK-BE-NEXT:    [[TMP2:%.*]] = extractelement <1 x half> [[DATA:%.*]], i64 0
+; CHECK-BE-NEXT:    [[TMP3:%.*]] = getelementptr inbounds half, ptr [[P:%.*]], i32 0
+; CHECK-BE-NEXT:    store half [[TMP2]], ptr [[TMP3]], align 2
+; CHECK-BE-NEXT:    br label [[ELSE]]
+; CHECK-BE:       else:
+; CHECK-BE-NEXT:    ret void
+;
+  call void @llvm.masked.store.v1f16.p0(<1 x half> %data, ptr %p, i32 2, <1 x i1> %mask)
+  ret void
+}
+
+declare void @llvm.masked.store.v1f16.p0(<1 x half>, ptr, i32, <1 x i1>)
 declare void @llvm.masked.store.v2i64.p0(<2 x i64>, ptr, i32, <2 x i1>)



More information about the llvm-commits mailing list