[llvm] [Windows][Arm64EC] Enable thunk generation for bfloat16 (PR #206710)
via llvm-commits
llvm-commits at lists.llvm.org
Tue Jun 30 04:21:31 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-backend-aarch64
Author: David Truby (DavidTruby)
<details>
<summary>Changes</summary>
This patch enables thunk generation for functions that take and return bfloat16 types. These types live in the same registers as fp16 types in both Arm64 and x86 so we just need the same behaviour as for fp16.
Assisted-by: codex (gpt-5.5)
Co-authored-by: nick.dingle@<!-- -->arm.com
---
Full diff: https://github.com/llvm/llvm-project/pull/206710.diff
4 Files Affected:
- (modified) llvm/lib/Target/AArch64/AArch64Arm64ECCallLowering.cpp (+15-7)
- (modified) llvm/lib/Target/AArch64/AArch64CallingConvention.td (+5-1)
- (modified) llvm/test/CodeGen/AArch64/arm64ec-entry-thunks.ll (+46-1)
- (modified) llvm/test/CodeGen/AArch64/arm64ec-exit-thunks.ll (+54-1)
``````````diff
diff --git a/llvm/lib/Target/AArch64/AArch64Arm64ECCallLowering.cpp b/llvm/lib/Target/AArch64/AArch64Arm64ECCallLowering.cpp
index 44f76f9e8772b..6a9f859896d9b 100644
--- a/llvm/lib/Target/AArch64/AArch64Arm64ECCallLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64Arm64ECCallLowering.cpp
@@ -323,6 +323,11 @@ ThunkArgInfo AArch64Arm64ECCallLowering::canonicalizeThunkType(
return direct(T);
}
+ if (T->isBFloatTy()) {
+ Out << "bf16";
+ return direct(T);
+ }
+
if (T->isFloatTy()) {
Out << "f";
return direct(T);
@@ -334,8 +339,9 @@ ThunkArgInfo AArch64Arm64ECCallLowering::canonicalizeThunkType(
}
if (T->isFloatingPointTy()) {
- report_fatal_error("Only 16, 32, and 64 bit floating points are supported "
- "for ARM64EC thunks");
+ report_fatal_error(
+ "Only half, bfloat16, float, and double are supported for ARM64EC "
+ "thunks");
}
auto &DL = M->getDataLayout();
@@ -349,11 +355,13 @@ ThunkArgInfo AArch64Arm64ECCallLowering::canonicalizeThunkType(
uint64_t ElementCnt = T->getArrayNumElements();
uint64_t ElementSizePerBytes = DL.getTypeSizeInBits(ElementTy) / 8;
uint64_t TotalSizeBytes = ElementCnt * ElementSizePerBytes;
- if (ElementTy->isHalfTy() || ElementTy->isFloatTy() ||
- ElementTy->isDoubleTy()) {
+ if (ElementTy->isHalfTy() || ElementTy->isBFloatTy() ||
+ ElementTy->isFloatTy() || ElementTy->isDoubleTy()) {
if (ElementTy->isHalfTy())
// Prefix with `llvm` since MSVC doesn't specify `_Float16`
Out << "__llvm_H__";
+ else if (ElementTy->isBFloatTy())
+ Out << "BF16";
else if (ElementTy->isFloatTy())
Out << "F";
else if (ElementTy->isDoubleTy())
@@ -369,10 +377,10 @@ ThunkArgInfo AArch64Arm64ECCallLowering::canonicalizeThunkType(
// Struct is passed directly on Arm64, but indirectly on X64.
return pointerIndirection(T);
}
- } else if (T->isFloatingPointTy()) {
+ } else if (ElementTy->isFloatingPointTy()) {
report_fatal_error(
- "Only 16, 32, and 64 bit floating points are supported "
- "for ARM64EC thunks");
+ "Only half, bfloat16, float, and double are supported for ARM64EC "
+ "thunks");
}
}
diff --git a/llvm/lib/Target/AArch64/AArch64CallingConvention.td b/llvm/lib/Target/AArch64/AArch64CallingConvention.td
index 8223984aed8a0..4683f166386e7 100644
--- a/llvm/lib/Target/AArch64/AArch64CallingConvention.td
+++ b/llvm/lib/Target/AArch64/AArch64CallingConvention.td
@@ -274,6 +274,9 @@ def CC_AArch64_Arm64EC_Thunk : CallingConv<[
CCIfType<[f16],
CCAssignToRegWithShadow<[H0, H1, H2, H3],
[X0, X1, X2, X3]>>,
+ CCIfType<[bf16],
+ CCAssignToRegWithShadow<[H0, H1, H2, H3],
+ [X0, X1, X2, X3]>>,
CCIfType<[f32],
CCAssignToRegWithShadow<[S0, S1, S2, S3],
[X0, X1, X2, X3]>>,
@@ -294,7 +297,7 @@ def CC_AArch64_Arm64EC_Thunk : CallingConv<[
// Integer/FP values get stored in stack slots that are 8 bytes in size and
// 8-byte aligned if there are no more registers to hold them.
- CCIfType<[i8, i16, i32, i64, f16, f32, f64], CCAssignToStack<8, 8>>
+ CCIfType<[i8, i16, i32, i64, f16, bf16, f32, f64], CCAssignToStack<8, 8>>
]>;
// The native side of ARM64EC thunks
@@ -313,6 +316,7 @@ def RetCC_AArch64_Arm64EC_Thunk : CallingConv<[
// The X86-64 calling convention always returns FP values in XMM0.
CCIfType<[f16], CCAssignToReg<[H0, H1]>>,
+ CCIfType<[bf16], CCAssignToReg<[H0, H1]>>,
CCIfType<[f32], CCAssignToReg<[S0, S1]>>,
CCIfType<[f64], CCAssignToReg<[D0, D1]>>,
CCIfType<[f128], CCAssignToReg<[Q0, Q1]>>,
diff --git a/llvm/test/CodeGen/AArch64/arm64ec-entry-thunks.ll b/llvm/test/CodeGen/AArch64/arm64ec-entry-thunks.ll
index 09fe884940c6a..984582686a47f 100644
--- a/llvm/test/CodeGen/AArch64/arm64ec-entry-thunks.ll
+++ b/llvm/test/CodeGen/AArch64/arm64ec-entry-thunks.ll
@@ -85,7 +85,7 @@ define i64 @simple_integers(i8, i16, i32, i64) nounwind {
ret i64 0
}
-; NOTE: Only half, float, and double are supported.
+; NOTE: Half, bfloat16, float, and double are supported.
define double @simple_floats(half, float, double) nounwind {
; CHECK-LABEL: .def $ientry_thunk$cdecl$d$__llvm_h__fd;
; CHECK: .section .wowthk$aa,"xr",discard,$ientry_thunk$cdecl$d$__llvm_h__fd
@@ -128,6 +128,48 @@ define double @simple_floats(half, float, double) nounwind {
ret double 0.0
}
+define bfloat @simple_bfloat(bfloat %0, bfloat %1) nounwind {
+; CHECK-LABEL: .def $ientry_thunk$cdecl$bf16$bf16bf16;
+; CHECK: .section .wowthk$aa,"xr",discard,$ientry_thunk$cdecl$bf16$bf16bf16
+; CHECK: // %bb.0:
+; CHECK-NEXT: stp q6, q7, [sp, #-176]! // 32-byte Folded Spill
+; CHECK-NEXT: .seh_save_any_reg_px q6, 176
+; CHECK-NEXT: stp q8, q9, [sp, #32] // 32-byte Folded Spill
+; CHECK-NEXT: .seh_save_any_reg_p q8, 32
+; CHECK-NEXT: stp q10, q11, [sp, #64] // 32-byte Folded Spill
+; CHECK-NEXT: .seh_save_any_reg_p q10, 64
+; CHECK-NEXT: stp q12, q13, [sp, #96] // 32-byte Folded Spill
+; CHECK-NEXT: .seh_save_any_reg_p q12, 96
+; CHECK-NEXT: stp q14, q15, [sp, #128] // 32-byte Folded Spill
+; CHECK-NEXT: .seh_save_any_reg_p q14, 128
+; CHECK-NEXT: stp x29, x30, [sp, #160] // 16-byte Folded Spill
+; CHECK-NEXT: .seh_save_fplr 160
+; CHECK-NEXT: add x29, sp, #160
+; CHECK-NEXT: .seh_add_fp 160
+; CHECK-NEXT: .seh_endprologue
+; CHECK-NEXT: blr x9
+; CHECK-NEXT: adrp x8, __os_arm64x_dispatch_ret
+; CHECK-NEXT: ldr x0, [x8, :lo12:__os_arm64x_dispatch_ret]
+; CHECK-NEXT: .seh_startepilogue
+; CHECK-NEXT: ldp x29, x30, [sp, #160] // 16-byte Folded Reload
+; CHECK-NEXT: .seh_save_fplr 160
+; CHECK-NEXT: ldp q14, q15, [sp, #128] // 32-byte Folded Reload
+; CHECK-NEXT: .seh_save_any_reg_p q14, 128
+; CHECK-NEXT: ldp q12, q13, [sp, #96] // 32-byte Folded Reload
+; CHECK-NEXT: .seh_save_any_reg_p q12, 96
+; CHECK-NEXT: ldp q10, q11, [sp, #64] // 32-byte Folded Reload
+; CHECK-NEXT: .seh_save_any_reg_p q10, 64
+; CHECK-NEXT: ldp q8, q9, [sp, #32] // 32-byte Folded Reload
+; CHECK-NEXT: .seh_save_any_reg_p q8, 32
+; CHECK-NEXT: ldp q6, q7, [sp], #176 // 32-byte Folded Reload
+; CHECK-NEXT: .seh_save_any_reg_px q6, 176
+; CHECK-NEXT: .seh_endepilogue
+; CHECK-NEXT: br x0
+; CHECK-NEXT: .seh_endfunclet
+; CHECK-NEXT: .seh_endproc
+ ret bfloat %0
+}
+
define void @has_varargs(...) nounwind {
; CHECK-LABEL: .def $ientry_thunk$cdecl$v$varargs;
; CHECK: .section .wowthk$aa,"xr",discard,$ientry_thunk$cdecl$v$varargs
@@ -604,6 +646,9 @@ start:
; CHECK-NEXT: .symidx "#simple_floats"
; CHECK-NEXT: .symidx $ientry_thunk$cdecl$d$__llvm_h__fd
; CHECK-NEXT: .word 1
+; CHECK-NEXT: .symidx "#simple_bfloat"
+; CHECK-NEXT: .symidx $ientry_thunk$cdecl$bf16$bf16bf16
+; CHECK-NEXT: .word 1
; CHECK-NEXT: .symidx "#has_varargs"
; CHECK-NEXT: .symidx $ientry_thunk$cdecl$v$varargs
; CHECK-NEXT: .word 1
diff --git a/llvm/test/CodeGen/AArch64/arm64ec-exit-thunks.ll b/llvm/test/CodeGen/AArch64/arm64ec-exit-thunks.ll
index 6fba6a3974574..4b7d1790f0a20 100644
--- a/llvm/test/CodeGen/AArch64/arm64ec-exit-thunks.ll
+++ b/llvm/test/CodeGen/AArch64/arm64ec-exit-thunks.ll
@@ -93,7 +93,7 @@ declare i64 @simple_integers(i8, i16, i32, i64) nounwind;
; CHECK-NEXT: .seh_endfunclet
; CHECK-NEXT: .seh_endproc
-; NOTE: Only half, float, and double are supported.
+; NOTE: Half, bfloat16, float, and double are supported.
declare double @simple_floats(half, float, double) nounwind;
; CHECK-LABEL: .def $iexit_thunk$cdecl$d$__llvm_h__fd;
; CHECK: .section .wowthk$aa,"xr",discard,$iexit_thunk$cdecl$d$__llvm_h__fd
@@ -140,6 +140,52 @@ declare double @simple_floats(half, float, double) nounwind;
; CHECK-NEXT: .seh_endfunclet
; CHECK-NEXT: .seh_endproc
+declare bfloat @simple_bfloat(bfloat, bfloat) nounwind;
+; CHECK-LABEL: .def $iexit_thunk$cdecl$bf16$bf16bf16;
+; CHECK: .section .wowthk$aa,"xr",discard,$iexit_thunk$cdecl$bf16$bf16bf16
+; CHECK: // %bb.0:
+; CHECK-NEXT: sub sp, sp, #48
+; CHECK-NEXT: .seh_stackalloc 48
+; CHECK-NEXT: stp x29, x30, [sp, #32] // 16-byte Folded Spill
+; CHECK-NEXT: .seh_save_fplr 32
+; CHECK-NEXT: add x29, sp, #32
+; CHECK-NEXT: .seh_add_fp 32
+; CHECK-NEXT: .seh_endprologue
+; CHECK-NEXT: adrp x8, __os_arm64x_dispatch_call_no_redirect
+; CHECK-NEXT: ldr x16, [x8, :lo12:__os_arm64x_dispatch_call_no_redirect]
+; CHECK-NEXT: blr x16
+; CHECK-NEXT: .seh_startepilogue
+; CHECK-NEXT: ldp x29, x30, [sp, #32] // 16-byte Folded Reload
+; CHECK-NEXT: .seh_save_fplr 32
+; CHECK-NEXT: add sp, sp, #48
+; CHECK-NEXT: .seh_stackalloc 48
+; CHECK-NEXT: .seh_endepilogue
+; CHECK-NEXT: ret
+; CHECK-NEXT: .seh_endfunclet
+; CHECK-NEXT: .seh_endproc
+; CHECK-LABEL: .def "#simple_bfloat$exit_thunk";
+; CHECK: .section .wowthk$aa,"xr",discard,"#simple_bfloat$exit_thunk"
+; CHECK: .weak_anti_dep simple_bfloat
+; CHECK: .weak_anti_dep "#simple_bfloat"
+; CHECK: // %bb.0:
+; CHECK-NEXT: str x30, [sp, #-16]! // 8-byte Folded Spill
+; CHECK-NEXT: .seh_save_reg_x x30, 16
+; CHECK-NEXT: .seh_endprologue
+; CHECK-NEXT: adrp x8, __os_arm64x_check_icall
+; CHECK-NEXT: adrp x11, simple_bfloat
+; CHECK-NEXT: add x11, x11, :lo12:simple_bfloat
+; CHECK-NEXT: ldr x8, [x8, :lo12:__os_arm64x_check_icall]
+; CHECK-NEXT: adrp x10, $iexit_thunk$cdecl$bf16$bf16bf16
+; CHECK-NEXT: add x10, x10, :lo12:$iexit_thunk$cdecl$bf16$bf16bf16
+; CHECK-NEXT: blr x8
+; CHECK-NEXT: .seh_startepilogue
+; CHECK-NEXT: ldr x30, [sp], #16 // 8-byte Folded Reload
+; CHECK-NEXT: .seh_save_reg_x x30, 16
+; CHECK-NEXT: .seh_endepilogue
+; CHECK-NEXT: br x11
+; CHECK-NEXT: .seh_endfunclet
+; CHECK-NEXT: .seh_endproc
+
declare void @has_varargs(...) nounwind;
; CHECK-LABEL: .def $iexit_thunk$cdecl$v$varargs;
; CHECK: .section .wowthk$aa,"xr",discard,$iexit_thunk$cdecl$v$varargs
@@ -620,6 +666,12 @@ declare void @"??@md5mangleaaaaaaaaaaaaaaaaaaaaaaa@"()
; CHECK-NEXT: .symidx "#simple_floats$exit_thunk"
; CHECK-NEXT: .symidx simple_floats
; CHECK-NEXT: .word 0
+; CHECK-NEXT: .symidx simple_bfloat
+; CHECK-NEXT: .symidx $iexit_thunk$cdecl$bf16$bf16bf16
+; CHECK-NEXT: .word 4
+; CHECK-NEXT: .symidx "#simple_bfloat$exit_thunk"
+; CHECK-NEXT: .symidx simple_bfloat
+; CHECK-NEXT: .word 0
; CHECK-NEXT: .symidx has_varargs
; CHECK-NEXT: .symidx $iexit_thunk$cdecl$v$varargs
; CHECK-NEXT: .word 4
@@ -679,6 +731,7 @@ define void @func_caller() nounwind {
call void @no_op()
call i64 @simple_integers(i8 0, i16 0, i32 0, i64 0)
call double @simple_floats(half 0.0, float 0.0, double 0.0)
+ call bfloat @simple_bfloat(bfloat 0xR0000, bfloat 0xR0000)
call void (...) @has_varargs()
%c = alloca i8
call void @has_sret(ptr sret([100 x i8]) %c)
``````````
</details>
https://github.com/llvm/llvm-project/pull/206710
More information about the llvm-commits
mailing list