[llvm] cf6335b - [X86] Lower scalar bf16 arithmetic on AVX10.2 via packed ops (#212245)
via llvm-commits
llvm-commits at lists.llvm.org
Thu Jul 30 06:59:30 PDT 2026
Author: tfzee
Date: 2026-07-30T15:59:25+02:00
New Revision: cf6335b275a996adff8334cb35245480ece481e1
URL: https://github.com/llvm/llvm-project/commit/cf6335b275a996adff8334cb35245480ece481e1
DIFF: https://github.com/llvm/llvm-project/commit/cf6335b275a996adff8334cb35245480ece481e1.diff
LOG: [X86] Lower scalar bf16 arithmetic on AVX10.2 via packed ops (#212245)
Currently basic(fadd/fsub/fmul/fdiv/fsqrt/fma) bf16 operations, as they
are not natively supported, are expanded to f32 operations.
However with AVX10.2 there are packed versions for these operations
which should be used instead and are enabled with this PR.
Since there is no native bf16 register class I instead go through f16
since it matches its size and register class.
This conversion will end up getting optimized away leaving only the
intended packed operation.
I used AI to double check and expand on comments.
Added:
llvm/test/CodeGen/X86/avx10_2bf16-arith-scalar.ll
Modified:
llvm/lib/Target/X86/X86ISelLowering.cpp
llvm/test/CodeGen/X86/avx10_2bf16-fma.ll
Removed:
################################################################################
diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index d78c478ab672f..1de4ffa33e960 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -2601,6 +2601,15 @@ X86TargetLowering::X86TargetLowering(const X86TargetMachine &TM,
}
if (!Subtarget.useSoftFloat() && Subtarget.hasAVX10_2()) {
+ // Lower scalar bf16 arithmetic by widening to a vector op and extracting
+ // the low element.
+ setOperationAction(ISD::FADD, MVT::bf16, Custom);
+ setOperationAction(ISD::FSUB, MVT::bf16, Custom);
+ setOperationAction(ISD::FMUL, MVT::bf16, Custom);
+ setOperationAction(ISD::FDIV, MVT::bf16, Custom);
+ setOperationAction(ISD::FSQRT, MVT::bf16, Custom);
+ setOperationAction(ISD::FMA, MVT::bf16, Custom);
+
setOperationAction(ISD::FADD, MVT::v32bf16, Legal);
setOperationAction(ISD::FSUB, MVT::v32bf16, Legal);
setOperationAction(ISD::FMUL, MVT::v32bf16, Legal);
@@ -34531,6 +34540,29 @@ void X86TargetLowering::ReplaceNodeResults(SDNode *N,
N->dump(&DAG);
#endif
llvm_unreachable("Do not know how to custom type legalize this operation!");
+ case ISD::FADD:
+ case ISD::FSUB:
+ case ISD::FMUL:
+ case ISD::FSQRT:
+ case ISD::FDIV:
+ case ISD::FMA: {
+ assert(N->getValueType(0) == MVT::bf16 && "Expected scalar bf16 result");
+ // AVX10.2 has no scalar bf16 arithmetic instructions, and bf16 is a
+ // soft-promoted-half type, so scalar ops would otherwise be promoted to
+ // f32. Instead widen each operand to a v8bf16 vector, perform the legal
+ // packed operation, and extract the low element afterwards.
+ SmallVector<SDValue, 3> VecOps;
+ for (const SDValue &Op : N->ops()) {
+ SDValue AsF16 = DAG.getBitcast(MVT::f16, Op);
+ SDValue VecF16 =
+ DAG.getNode(ISD::SCALAR_TO_VECTOR, dl, MVT::v8f16, AsF16);
+ VecOps.push_back(DAG.getBitcast(MVT::v8bf16, VecF16));
+ }
+ SDValue Vec =
+ DAG.getNode(N->getOpcode(), dl, MVT::v8bf16, VecOps, N->getFlags());
+ Results.push_back(DAG.getExtractVectorElt(dl, MVT::bf16, Vec, 0));
+ return;
+ }
case X86ISD::CVTPH2PS: {
EVT VT = N->getValueType(0);
SDValue Lo, Hi;
diff --git a/llvm/test/CodeGen/X86/avx10_2bf16-arith-scalar.ll b/llvm/test/CodeGen/X86/avx10_2bf16-arith-scalar.ll
new file mode 100644
index 0000000000000..cc7f3ad0eb174
--- /dev/null
+++ b/llvm/test/CodeGen/X86/avx10_2bf16-arith-scalar.ll
@@ -0,0 +1,221 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 6
+; RUN: llc < %s -verify-machineinstrs -mtriple=x86_64-unknown-unknown -mattr=+avx10.2 | FileCheck %s --check-prefixes=AVX10_2
+; RUN: llc < %s -verify-machineinstrs -mtriple=x86_64-unknown-unknown -mattr=+avx512bf16,+avx512vl | FileCheck %s --check-prefixes=AVX512BF16
+; RUN: llc < %s -verify-machineinstrs -mtriple=x86_64-unknown-unknown -mattr=+avxneconvert | FileCheck %s --check-prefixes=AVXNECONVERT
+
+define bfloat @fadd_bf16(bfloat %a, bfloat %b) nounwind {
+; AVX10_2-LABEL: fadd_bf16:
+; AVX10_2: # %bb.0: # %entry
+; AVX10_2-NEXT: vaddbf16 %xmm1, %xmm0, %xmm0
+; AVX10_2-NEXT: retq
+;
+; AVX512BF16-LABEL: fadd_bf16:
+; AVX512BF16: # %bb.0: # %entry
+; AVX512BF16-NEXT: vpextrw $0, %xmm0, %eax
+; AVX512BF16-NEXT: vpextrw $0, %xmm1, %ecx
+; AVX512BF16-NEXT: shll $16, %ecx
+; AVX512BF16-NEXT: vmovd %ecx, %xmm0
+; AVX512BF16-NEXT: shll $16, %eax
+; AVX512BF16-NEXT: vmovd %eax, %xmm1
+; AVX512BF16-NEXT: vaddss %xmm0, %xmm1, %xmm0
+; AVX512BF16-NEXT: vcvtneps2bf16 %xmm0, %xmm0
+; AVX512BF16-NEXT: retq
+;
+; AVXNECONVERT-LABEL: fadd_bf16:
+; AVXNECONVERT: # %bb.0: # %entry
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm0, %eax
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm1, %ecx
+; AVXNECONVERT-NEXT: shll $16, %ecx
+; AVXNECONVERT-NEXT: vmovd %ecx, %xmm0
+; AVXNECONVERT-NEXT: shll $16, %eax
+; AVXNECONVERT-NEXT: vmovd %eax, %xmm1
+; AVXNECONVERT-NEXT: vaddss %xmm0, %xmm1, %xmm0
+; AVXNECONVERT-NEXT: {vex} vcvtneps2bf16 %xmm0, %xmm0
+; AVXNECONVERT-NEXT: retq
+entry:
+ %r = fadd contract bfloat %a, %b
+ ret bfloat %r
+}
+
+
+define bfloat @dont_fuse_bf16(bfloat %a, bfloat %b, bfloat %c) nounwind {
+; AVX10_2-LABEL: dont_fuse_bf16:
+; AVX10_2: # %bb.0: # %entry
+; AVX10_2-NEXT: vmulbf16 %xmm1, %xmm0, %xmm0
+; AVX10_2-NEXT: vaddbf16 %xmm2, %xmm0, %xmm0
+; AVX10_2-NEXT: retq
+;
+; AVX512BF16-LABEL: dont_fuse_bf16:
+; AVX512BF16: # %bb.0: # %entry
+; AVX512BF16-NEXT: vpextrw $0, %xmm2, %eax
+; AVX512BF16-NEXT: vpextrw $0, %xmm0, %ecx
+; AVX512BF16-NEXT: vpextrw $0, %xmm1, %edx
+; AVX512BF16-NEXT: shll $16, %edx
+; AVX512BF16-NEXT: vmovd %edx, %xmm0
+; AVX512BF16-NEXT: shll $16, %ecx
+; AVX512BF16-NEXT: vmovd %ecx, %xmm1
+; AVX512BF16-NEXT: vmulss %xmm0, %xmm1, %xmm0
+; AVX512BF16-NEXT: vcvtneps2bf16 %xmm0, %xmm0
+; AVX512BF16-NEXT: vmovd %xmm0, %ecx
+; AVX512BF16-NEXT: shll $16, %ecx
+; AVX512BF16-NEXT: vmovd %ecx, %xmm0
+; AVX512BF16-NEXT: shll $16, %eax
+; AVX512BF16-NEXT: vmovd %eax, %xmm1
+; AVX512BF16-NEXT: vaddss %xmm1, %xmm0, %xmm0
+; AVX512BF16-NEXT: vcvtneps2bf16 %xmm0, %xmm0
+; AVX512BF16-NEXT: retq
+;
+; AVXNECONVERT-LABEL: dont_fuse_bf16:
+; AVXNECONVERT: # %bb.0: # %entry
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm2, %eax
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm0, %ecx
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm1, %edx
+; AVXNECONVERT-NEXT: shll $16, %edx
+; AVXNECONVERT-NEXT: vmovd %edx, %xmm0
+; AVXNECONVERT-NEXT: shll $16, %ecx
+; AVXNECONVERT-NEXT: vmovd %ecx, %xmm1
+; AVXNECONVERT-NEXT: vmulss %xmm0, %xmm1, %xmm0
+; AVXNECONVERT-NEXT: {vex} vcvtneps2bf16 %xmm0, %xmm0
+; AVXNECONVERT-NEXT: vmovd %xmm0, %ecx
+; AVXNECONVERT-NEXT: shll $16, %ecx
+; AVXNECONVERT-NEXT: vmovd %ecx, %xmm0
+; AVXNECONVERT-NEXT: shll $16, %eax
+; AVXNECONVERT-NEXT: vmovd %eax, %xmm1
+; AVXNECONVERT-NEXT: vaddss %xmm1, %xmm0, %xmm0
+; AVXNECONVERT-NEXT: {vex} vcvtneps2bf16 %xmm0, %xmm0
+; AVXNECONVERT-NEXT: retq
+entry:
+ %m = fmul bfloat %a, %b
+ %r = fadd bfloat %m, %c
+ ret bfloat %r
+}
+
+define bfloat @fsub_bf16(bfloat %a, bfloat %b) nounwind {
+; AVX10_2-LABEL: fsub_bf16:
+; AVX10_2: # %bb.0: # %entry
+; AVX10_2-NEXT: vsubbf16 %xmm1, %xmm0, %xmm0
+; AVX10_2-NEXT: retq
+;
+; AVX512BF16-LABEL: fsub_bf16:
+; AVX512BF16: # %bb.0: # %entry
+; AVX512BF16-NEXT: vpextrw $0, %xmm0, %eax
+; AVX512BF16-NEXT: vpextrw $0, %xmm1, %ecx
+; AVX512BF16-NEXT: shll $16, %ecx
+; AVX512BF16-NEXT: vmovd %ecx, %xmm0
+; AVX512BF16-NEXT: shll $16, %eax
+; AVX512BF16-NEXT: vmovd %eax, %xmm1
+; AVX512BF16-NEXT: vsubss %xmm0, %xmm1, %xmm0
+; AVX512BF16-NEXT: vcvtneps2bf16 %xmm0, %xmm0
+; AVX512BF16-NEXT: retq
+;
+; AVXNECONVERT-LABEL: fsub_bf16:
+; AVXNECONVERT: # %bb.0: # %entry
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm0, %eax
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm1, %ecx
+; AVXNECONVERT-NEXT: shll $16, %ecx
+; AVXNECONVERT-NEXT: vmovd %ecx, %xmm0
+; AVXNECONVERT-NEXT: shll $16, %eax
+; AVXNECONVERT-NEXT: vmovd %eax, %xmm1
+; AVXNECONVERT-NEXT: vsubss %xmm0, %xmm1, %xmm0
+; AVXNECONVERT-NEXT: {vex} vcvtneps2bf16 %xmm0, %xmm0
+; AVXNECONVERT-NEXT: retq
+entry:
+ %r = fsub bfloat %a, %b
+ ret bfloat %r
+}
+
+define bfloat @fmul_bf16(bfloat %a, bfloat %b) nounwind {
+; AVX10_2-LABEL: fmul_bf16:
+; AVX10_2: # %bb.0: # %entry
+; AVX10_2-NEXT: vmulbf16 %xmm1, %xmm0, %xmm0
+; AVX10_2-NEXT: retq
+;
+; AVX512BF16-LABEL: fmul_bf16:
+; AVX512BF16: # %bb.0: # %entry
+; AVX512BF16-NEXT: vpextrw $0, %xmm0, %eax
+; AVX512BF16-NEXT: vpextrw $0, %xmm1, %ecx
+; AVX512BF16-NEXT: shll $16, %ecx
+; AVX512BF16-NEXT: vmovd %ecx, %xmm0
+; AVX512BF16-NEXT: shll $16, %eax
+; AVX512BF16-NEXT: vmovd %eax, %xmm1
+; AVX512BF16-NEXT: vmulss %xmm0, %xmm1, %xmm0
+; AVX512BF16-NEXT: vcvtneps2bf16 %xmm0, %xmm0
+; AVX512BF16-NEXT: retq
+;
+; AVXNECONVERT-LABEL: fmul_bf16:
+; AVXNECONVERT: # %bb.0: # %entry
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm0, %eax
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm1, %ecx
+; AVXNECONVERT-NEXT: shll $16, %ecx
+; AVXNECONVERT-NEXT: vmovd %ecx, %xmm0
+; AVXNECONVERT-NEXT: shll $16, %eax
+; AVXNECONVERT-NEXT: vmovd %eax, %xmm1
+; AVXNECONVERT-NEXT: vmulss %xmm0, %xmm1, %xmm0
+; AVXNECONVERT-NEXT: {vex} vcvtneps2bf16 %xmm0, %xmm0
+; AVXNECONVERT-NEXT: retq
+entry:
+ %m = fmul bfloat %a, %b
+ ret bfloat %m
+}
+
+define bfloat @fdiv_bf16(bfloat %a, bfloat %b) nounwind {
+; AVX10_2-LABEL: fdiv_bf16:
+; AVX10_2: # %bb.0: # %entry
+; AVX10_2-NEXT: vdivbf16 %xmm1, %xmm0, %xmm0
+; AVX10_2-NEXT: retq
+;
+; AVX512BF16-LABEL: fdiv_bf16:
+; AVX512BF16: # %bb.0: # %entry
+; AVX512BF16-NEXT: vpextrw $0, %xmm0, %eax
+; AVX512BF16-NEXT: vpextrw $0, %xmm1, %ecx
+; AVX512BF16-NEXT: shll $16, %ecx
+; AVX512BF16-NEXT: vmovd %ecx, %xmm0
+; AVX512BF16-NEXT: shll $16, %eax
+; AVX512BF16-NEXT: vmovd %eax, %xmm1
+; AVX512BF16-NEXT: vdivss %xmm0, %xmm1, %xmm0
+; AVX512BF16-NEXT: vcvtneps2bf16 %xmm0, %xmm0
+; AVX512BF16-NEXT: retq
+;
+; AVXNECONVERT-LABEL: fdiv_bf16:
+; AVXNECONVERT: # %bb.0: # %entry
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm0, %eax
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm1, %ecx
+; AVXNECONVERT-NEXT: shll $16, %ecx
+; AVXNECONVERT-NEXT: vmovd %ecx, %xmm0
+; AVXNECONVERT-NEXT: shll $16, %eax
+; AVXNECONVERT-NEXT: vmovd %eax, %xmm1
+; AVXNECONVERT-NEXT: vdivss %xmm0, %xmm1, %xmm0
+; AVXNECONVERT-NEXT: {vex} vcvtneps2bf16 %xmm0, %xmm0
+; AVXNECONVERT-NEXT: retq
+entry:
+ %m = fdiv bfloat %a, %b
+ ret bfloat %m
+}
+
+define bfloat @fsqrt_bf16(bfloat %a) nounwind {
+; AVX10_2-LABEL: fsqrt_bf16:
+; AVX10_2: # %bb.0: # %entry
+; AVX10_2-NEXT: vsqrtbf16 %xmm0, %xmm0
+; AVX10_2-NEXT: retq
+;
+; AVX512BF16-LABEL: fsqrt_bf16:
+; AVX512BF16: # %bb.0: # %entry
+; AVX512BF16-NEXT: vpextrw $0, %xmm0, %eax
+; AVX512BF16-NEXT: shll $16, %eax
+; AVX512BF16-NEXT: vmovd %eax, %xmm0
+; AVX512BF16-NEXT: vsqrtss %xmm0, %xmm0, %xmm0
+; AVX512BF16-NEXT: vcvtneps2bf16 %xmm0, %xmm0
+; AVX512BF16-NEXT: retq
+;
+; AVXNECONVERT-LABEL: fsqrt_bf16:
+; AVXNECONVERT: # %bb.0: # %entry
+; AVXNECONVERT-NEXT: vpextrw $0, %xmm0, %eax
+; AVXNECONVERT-NEXT: shll $16, %eax
+; AVXNECONVERT-NEXT: vmovd %eax, %xmm0
+; AVXNECONVERT-NEXT: vsqrtss %xmm0, %xmm0, %xmm0
+; AVXNECONVERT-NEXT: {vex} vcvtneps2bf16 %xmm0, %xmm0
+; AVXNECONVERT-NEXT: retq
+entry:
+ %m = call bfloat @llvm.sqrt.bf16(bfloat %a)
+ ret bfloat %m
+}
diff --git a/llvm/test/CodeGen/X86/avx10_2bf16-fma.ll b/llvm/test/CodeGen/X86/avx10_2bf16-fma.ll
index d79f0cc79b5f4..4c829cc78b3e5 100644
--- a/llvm/test/CodeGen/X86/avx10_2bf16-fma.ll
+++ b/llvm/test/CodeGen/X86/avx10_2bf16-fma.ll
@@ -6,17 +6,7 @@
define bfloat @fuse_bf16(bfloat %a, bfloat %b, bfloat %c) nounwind {
; AVX10_2-LABEL: fuse_bf16:
; AVX10_2: # %bb.0: # %entry
-; AVX10_2-NEXT: vmovw %xmm1, %eax
-; AVX10_2-NEXT: vmovw %xmm0, %ecx
-; AVX10_2-NEXT: vmovw %xmm2, %edx
-; AVX10_2-NEXT: shll $16, %edx
-; AVX10_2-NEXT: vmovd %edx, %xmm0
-; AVX10_2-NEXT: shll $16, %ecx
-; AVX10_2-NEXT: vmovd %ecx, %xmm1
-; AVX10_2-NEXT: shll $16, %eax
-; AVX10_2-NEXT: vmovd %eax, %xmm2
-; AVX10_2-NEXT: vfmadd213ss {{.*#+}} xmm2 = (xmm1 * xmm2) + xmm0
-; AVX10_2-NEXT: vcvtneps2bf16 %xmm2, %xmm0
+; AVX10_2-NEXT: vfmadd213bf16 %xmm2, %xmm1, %xmm0
; AVX10_2-NEXT: retq
;
; AVX512BF16-LABEL: fuse_bf16:
More information about the llvm-commits
mailing list