[llvm] [AArch64][SVE] Support lowering masked loads/stores of <4 x bf16> and <8 x bf16> (PR #208744)
via llvm-commits
llvm-commits at lists.llvm.org
Mon Jul 13 06:26:32 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-aarch64
Author: Shanzhi Chen (chenshanzhi)
<details>
<summary>Changes</summary>
Add support for lowering masked loads/stores of <4 x bf16> and <8 x bf16> when target features contain "+sve".
Fixes: #<!-- -->201149
---
Full diff: https://github.com/llvm/llvm-project/pull/208744.diff
4 Files Affected:
- (modified) llvm/lib/Target/AArch64/AArch64ISelLowering.cpp (+7-4)
- (modified) llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-loads.ll (+37)
- (modified) llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-stores.ll (+35)
- (modified) llvm/test/CodeGen/AArch64/sve2p1-vector-shuffles.ll (+3-25)
``````````diff
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index e883c8bb5e96e..bee755e413425 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -1791,9 +1791,9 @@ AArch64TargetLowering::AArch64TargetLowering(const TargetMachine &TM,
// NEON doesn't support masked loads/stores, but SME and SVE do.
for (auto VT :
- {MVT::v4f16, MVT::v8f16, MVT::v2f32, MVT::v4f32, MVT::v1f64,
- MVT::v2f64, MVT::v8i8, MVT::v16i8, MVT::v4i16, MVT::v8i16,
- MVT::v2i32, MVT::v4i32, MVT::v1i64, MVT::v2i64}) {
+ {MVT::v4f16, MVT::v8f16, MVT::v4bf16, MVT::v8bf16, MVT::v2f32,
+ MVT::v4f32, MVT::v1f64, MVT::v2f64, MVT::v8i8, MVT::v16i8, MVT::v4i16,
+ MVT::v8i16, MVT::v2i32, MVT::v4i32, MVT::v1i64, MVT::v2i64}) {
setOperationAction(ISD::MLOAD, VT, Custom);
setOperationAction(ISD::MSTORE, VT, Custom);
}
@@ -2521,7 +2521,7 @@ void AArch64TargetLowering::addTypeForFixedLengthSVE(MVT VT) {
// Mark floating-point truncating stores/extending loads as having custom
// lowering
- if (VT.isFloatingPoint()) {
+ if (VT.isFloatingPoint() && !VT.isVectorOf(MVT::bf16)) {
MVT InnerVT = VT.changeVectorElementType(MVT::f16);
while (InnerVT != VT) {
setTruncStoreAction(VT, InnerVT, Custom);
@@ -8847,6 +8847,9 @@ bool AArch64TargetLowering::useSVEForFixedLengthVectorVT(
case MVT::i1:
default:
return false;
+ case MVT::bf16:
+ return OverrideNEON && (VT.is128BitVector() || VT.is64BitVector()) &&
+ Subtarget->isSVEorStreamingSVEAvailable();
case MVT::i8:
case MVT::i16:
case MVT::i32:
diff --git a/llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-loads.ll b/llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-loads.ll
index 91af81c617c41..b30a3a9fc933d 100644
--- a/llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-loads.ll
+++ b/llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-loads.ll
@@ -35,6 +35,20 @@ define <8 x half> @masked_load_v8f16(ptr %src, <8 x i1> %mask) {
ret <8 x half> %load
}
+define <8 x bfloat> @masked_load_v8bf16(ptr %src, <8 x i1> %mask) #0 {
+; CHECK-LABEL: masked_load_v8bf16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: ushll v0.8h, v0.8b, #0
+; CHECK-NEXT: ptrue p0.h, vl8
+; CHECK-NEXT: shl v0.8h, v0.8h, #15
+; CHECK-NEXT: cmpne p1.h, p0/z, z0.h, #0
+; CHECK-NEXT: ld1h { z0.h }, p1/z, [x0]
+; CHECK-NEXT: // kill: def $q0 killed $q0 killed $z0
+; CHECK-NEXT: ret
+ %load = call <8 x bfloat> @llvm.masked.load.v8bf16(ptr %src, i32 8, <8 x i1> %mask, <8 x bfloat> zeroinitializer)
+ ret <8 x bfloat> %load
+}
+
define <4 x float> @masked_load_v4f32(ptr %src, <4 x i1> %mask) {
; CHECK-LABEL: masked_load_v4f32:
; CHECK: // %bb.0:
@@ -134,3 +148,26 @@ define <4 x half> @masked_load_v4f16(ptr %ap, ptr %bp) {
%load = call <4 x half> @llvm.masked.load.v4f16(ptr %ap, i32 2, <4 x i1> %mask, <4 x half> zeroinitializer)
ret <4 x half> %load
}
+
+define <4 x bfloat> @masked_load_v4bf16(ptr %ap, ptr %bp) #0 {
+; CHECK-LABEL: masked_load_v4bf16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: ldr d0, [x0]
+; CHECK-NEXT: ldr d1, [x1]
+; CHECK-NEXT: ptrue p0.h, vl4
+; CHECK-NEXT: shll v1.4s, v1.4h, #16
+; CHECK-NEXT: shll v0.4s, v0.4h, #16
+; CHECK-NEXT: fcmeq v0.4s, v0.4s, v1.4s
+; CHECK-NEXT: xtn v0.4h, v0.4s
+; CHECK-NEXT: cmpne p1.h, p0/z, z0.h, #0
+; CHECK-NEXT: ld1h { z0.h }, p1/z, [x0]
+; CHECK-NEXT: // kill: def $d0 killed $d0 killed $z0
+; CHECK-NEXT: ret
+ %a = load <4 x bfloat>, ptr %ap
+ %b = load <4 x bfloat>, ptr %bp
+ %mask = fcmp oeq <4 x bfloat> %a, %b
+ %load = call <4 x bfloat> @llvm.masked.load.v4bf16(ptr %ap, i32 2, <4 x i1> %mask, <4 x bfloat> zeroinitializer)
+ ret <4 x bfloat> %load
+}
+
+attributes #0 = { "target-features"="+sve,+bf16" }
diff --git a/llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-stores.ll b/llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-stores.ll
index 439c61339a631..6f42d4ceeffef 100644
--- a/llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-stores.ll
+++ b/llvm/test/CodeGen/AArch64/sve-fixed-length-masked-64-128bit-stores.ll
@@ -35,6 +35,20 @@ define void @masked_store_v8f16(ptr %dst, <8 x i1> %mask) {
ret void
}
+define void @masked_store_v8bf16(ptr %dst, <8 x i1> %mask) #0 {
+; CHECK-LABEL: masked_store_v8bf16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: ushll v0.8h, v0.8b, #0
+; CHECK-NEXT: ptrue p0.h, vl8
+; CHECK-NEXT: shl v0.8h, v0.8h, #15
+; CHECK-NEXT: cmpne p1.h, p0/z, z0.h, #0
+; CHECK-NEXT: movi v0.2d, #0000000000000000
+; CHECK-NEXT: st1h { z0.h }, p1, [x0]
+; CHECK-NEXT: ret
+ call void @llvm.masked.store.v8bf16(<8 x bfloat> zeroinitializer, ptr %dst, i32 8, <8 x i1> %mask)
+ ret void
+}
+
define void @masked_store_v4f32(ptr %dst, <4 x i1> %mask) {
; CHECK-LABEL: masked_store_v4f32:
; CHECK: // %bb.0:
@@ -116,3 +130,24 @@ define void @masked_store_v4f16(ptr %ap, ptr %bp) {
ret void
}
+define void @masked_store_v4bf16(ptr %ap, ptr %bp) #0 {
+; CHECK-LABEL: masked_store_v4bf16:
+; CHECK: // %bb.0:
+; CHECK-NEXT: ldr d0, [x0]
+; CHECK-NEXT: ldr d1, [x1]
+; CHECK-NEXT: ptrue p0.h, vl4
+; CHECK-NEXT: shll v1.4s, v1.4h, #16
+; CHECK-NEXT: shll v2.4s, v0.4h, #16
+; CHECK-NEXT: fcmeq v1.4s, v2.4s, v1.4s
+; CHECK-NEXT: xtn v1.4h, v1.4s
+; CHECK-NEXT: cmpne p1.h, p0/z, z1.h, #0
+; CHECK-NEXT: st1h { z0.h }, p1, [x1]
+; CHECK-NEXT: ret
+ %a = load <4 x bfloat>, ptr %ap
+ %b = load <4 x bfloat>, ptr %bp
+ %mask = fcmp oeq <4 x bfloat> %a, %b
+ call void @llvm.masked.store.v4bf16(<4 x bfloat> %a, ptr %bp, i32 2, <4 x i1> %mask)
+ ret void
+}
+
+attributes #0 = { "target-features"="+sve,+bf16" }
diff --git a/llvm/test/CodeGen/AArch64/sve2p1-vector-shuffles.ll b/llvm/test/CodeGen/AArch64/sve2p1-vector-shuffles.ll
index a2d4ababd919f..d13837f1c3da6 100644
--- a/llvm/test/CodeGen/AArch64/sve2p1-vector-shuffles.ll
+++ b/llvm/test/CodeGen/AArch64/sve2p1-vector-shuffles.ll
@@ -82,32 +82,10 @@ define void @dupq_bf16_256b(ptr %addr) #0 {
;
; SME-LABEL: dupq_bf16_256b:
; SME: // %bb.0:
-; SME-NEXT: ldp q1, q0, [x0]
-; SME-NEXT: str q0, [sp, #-64]!
-; SME-NEXT: .cfi_def_cfa_offset 64
-; SME-NEXT: ldr h0, [sp, #4]
-; SME-NEXT: str q1, [sp, #32]
-; SME-NEXT: str h0, [sp, #30]
-; SME-NEXT: str h0, [sp, #28]
-; SME-NEXT: str h0, [sp, #26]
-; SME-NEXT: str h0, [sp, #24]
-; SME-NEXT: str h0, [sp, #22]
-; SME-NEXT: str h0, [sp, #20]
-; SME-NEXT: str h0, [sp, #18]
-; SME-NEXT: str h0, [sp, #16]
-; SME-NEXT: ldr h0, [sp, #36]
-; SME-NEXT: ldr q1, [sp, #16]
-; SME-NEXT: str h0, [sp, #62]
-; SME-NEXT: str h0, [sp, #60]
-; SME-NEXT: str h0, [sp, #58]
-; SME-NEXT: str h0, [sp, #56]
-; SME-NEXT: str h0, [sp, #54]
-; SME-NEXT: str h0, [sp, #52]
-; SME-NEXT: str h0, [sp, #50]
-; SME-NEXT: str h0, [sp, #48]
-; SME-NEXT: ldr q0, [sp, #48]
+; SME-NEXT: ldp q0, q1, [x0]
+; SME-NEXT: mov z0.h, z0.h[2]
+; SME-NEXT: mov z1.h, z1.h[2]
; SME-NEXT: stp q0, q1, [x0]
-; SME-NEXT: add sp, sp, #64
; SME-NEXT: ret
%load = load <16 x bfloat>, ptr %addr
%splat.lanes = shufflevector <16 x bfloat> %load, <16 x bfloat> poison, <16 x i32> <i32 2, i32 2, i32 2, i32 2, i32 2, i32 2, i32 2, i32 2,
``````````
</details>
https://github.com/llvm/llvm-project/pull/208744
More information about the llvm-commits
mailing list