Skip to content

Commit e6fdc7a

Browse files
author
Dmitry Sidorov
committed
[Backport to 14] Support for SPV_INTEL_shader_atomic_bfloat16 extension (KhronosGroup#3343)
Spec is available here: intel/llvm#20009 Author: "Ratajewski, Andrzej" <andrzej.ratajewski@intel.com> Signed-off-by: Sidorov, Dmitry <dmitry.sidorov@intel.com>
1 parent b7b12ab commit e6fdc7a

10 files changed

Lines changed: 209 additions & 3 deletions

File tree

include/LLVMSPIRVExtensions.inc

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,3 +73,4 @@ EXT(SPV_INTEL_2d_block_io)
7373
EXT(SPV_INTEL_subgroup_matrix_multiply_accumulate)
7474
EXT(SPV_KHR_bfloat16)
7575
EXT(SPV_INTEL_bfloat16_arithmetic)
76+
EXT(SPV_INTEL_shader_atomic_bfloat16)

lib/SPIRV/libSPIRV/SPIRVInstruction.h

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2866,40 +2866,48 @@ class SPIRVAtomicFAddEXTInst : public SPIRVAtomicInstBase {
28662866
public:
28672867
llvm::Optional<ExtensionID> getRequiredExtension() const override {
28682868
assert(hasType());
2869+
if (getType()->isTypeFloat(16, FPEncodingBFloat16KHR))
2870+
return ExtensionID::SPV_INTEL_shader_atomic_bfloat16;
28692871
if (getType()->isTypeFloat(16))
28702872
return ExtensionID::SPV_EXT_shader_atomic_float16_add;
28712873
return ExtensionID::SPV_EXT_shader_atomic_float_add;
28722874
}
28732875

28742876
SPIRVCapVec getRequiredCapability() const override {
28752877
assert(hasType());
2878+
if (getType()->isTypeFloat(16, FPEncodingBFloat16KHR))
2879+
return {internal::CapabilityAtomicBFloat16AddINTEL};
28762880
if (getType()->isTypeFloat(16))
28772881
return {CapabilityAtomicFloat16AddEXT};
28782882
if (getType()->isTypeFloat(32))
28792883
return {CapabilityAtomicFloat32AddEXT};
28802884
if (getType()->isTypeFloat(64))
28812885
return {CapabilityAtomicFloat64AddEXT};
28822886
llvm_unreachable(
2883-
"AtomicFAddEXT can only be generated for f16, f32, f64 types");
2887+
"AtomicFAddEXT can only be generated for bf16, f16, f32, f64 types");
28842888
}
28852889
};
28862890

28872891
class SPIRVAtomicFMinMaxEXTBase : public SPIRVAtomicInstBase {
28882892
public:
28892893
llvm::Optional<ExtensionID> getRequiredExtension() const override {
2894+
if (getType()->isTypeFloat(16, FPEncodingBFloat16KHR))
2895+
return ExtensionID::SPV_INTEL_shader_atomic_bfloat16;
28902896
return ExtensionID::SPV_EXT_shader_atomic_float_min_max;
28912897
}
28922898

28932899
SPIRVCapVec getRequiredCapability() const override {
28942900
assert(hasType());
2901+
if (getType()->isTypeFloat(16, FPEncodingBFloat16KHR))
2902+
return {internal::CapabilityAtomicBFloat16MinMaxINTEL};
28952903
if (getType()->isTypeFloat(16))
28962904
return {CapabilityAtomicFloat16MinMaxEXT};
28972905
if (getType()->isTypeFloat(32))
28982906
return {CapabilityAtomicFloat32MinMaxEXT};
28992907
if (getType()->isTypeFloat(64))
29002908
return {CapabilityAtomicFloat64MinMaxEXT};
2901-
llvm_unreachable(
2902-
"AtomicF(Min|Max)EXT can only be generated for f16, f32, f64 types");
2909+
llvm_unreachable("AtomicF(Min|Max)EXT can only be generated for bf16, f16, "
2910+
"f32, f64 types");
29032911
}
29042912
};
29052913

lib/SPIRV/libSPIRV/SPIRVNameMapEnum.h

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -602,6 +602,9 @@ template <> inline void SPIRVMap<Capability, std::string>::init() {
602602
add(CapabilityLongCompositesINTEL, "LongCompositesINTEL");
603603
add(CapabilityOptNoneINTEL, "OptNoneINTEL");
604604
add(CapabilityAtomicFloat16AddEXT, "AtomicFloat16AddEXT");
605+
add(internal::CapabilityAtomicBFloat16AddINTEL, "AtomicBFloat16AddINTEL");
606+
add(internal::CapabilityAtomicBFloat16MinMaxINTEL,
607+
"AtomicBFloat16MinMaxINTEL");
605608
add(CapabilityDebugInfoModuleINTEL, "DebugInfoModuleINTEL");
606609
add(CapabilitySplitBarrierINTEL, "SplitBarrierINTEL");
607610
add(CapabilityGlobalVariableFPGADecorationsINTEL,

lib/SPIRV/libSPIRV/spirv_internal.hpp

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -115,6 +115,8 @@ enum InternalCapability {
115115
ICapGlobalVariableDecorationsINTEL = 6146,
116116
ICapabilityCooperativeMatrixCheckedInstructionsINTEL = 6192,
117117
ICapabilityBFloat16ArithmeticINTEL = 6226,
118+
ICapabilityAtomicBFloat16AddINTEL = 6255,
119+
ICapabilityAtomicBFloat16MinMaxINTEL = 6256,
118120
ICapabilityCooperativeMatrixPrefetchINTEL = 6411,
119121
ICapabilityComplexFloatMulDivINTEL = 6414,
120122
ICapabilityTensorFloat32RoundingINTEL = 6425,
@@ -217,6 +219,9 @@ _SPIRV_OP(Capability, BindlessImagesINTEL)
217219
_SPIRV_OP(Op, ConvertHandleToImageINTEL)
218220
_SPIRV_OP(Op, ConvertHandleToSamplerINTEL)
219221
_SPIRV_OP(Op, ConvertHandleToSampledImageINTEL)
222+
223+
_SPIRV_OP(Capability, AtomicBFloat16AddINTEL)
224+
_SPIRV_OP(Capability, AtomicBFloat16MinMaxINTEL)
220225
#undef _SPIRV_OP
221226

222227
constexpr SourceLanguage SourceLanguagePython =
Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
; RUN: llvm-as %s -o %t.bc
2+
; RUN: llvm-spirv %t.bc --spirv-ext=+SPV_INTEL_shader_atomic_bfloat16,+SPV_KHR_bfloat16 -o %t.spv
3+
; RUN: llvm-spirv -to-text %t.spv -o %t.spt
4+
; RUN: FileCheck < %t.spt %s --check-prefix=CHECK-SPIRV
5+
6+
; RUN: llvm-spirv --spirv-target-env=SPV-IR -r %t.spv -o %t.rev.bc
7+
; RUN: llvm-dis %t.rev.bc -o - | FileCheck %s --check-prefixes=CHECK-LLVM-SPV
8+
9+
target datalayout = "e-i64:64-v16:16-v24:32-v32:32-v48:64-v96:128-v192:256-v256:256-v512:512-v1024:1024-n8:16:32:64"
10+
target triple = "spir64-unknown-unknown"
11+
12+
; CHECK-SPIRV-DAG: Capability AtomicBFloat16AddINTEL
13+
; CHECK-SPIRV-DAG: Capability BFloat16TypeKHR
14+
; CHECK-SPIRV-DAG: Extension "SPV_INTEL_shader_atomic_bfloat16"
15+
; CHECK-SPIRV-DAG: Extension "SPV_KHR_bfloat16"
16+
17+
; CHECK-SPIRV: TypeFloat [[BFLOAT:[0-9]+]] 16 0
18+
19+
; Function Attrs: convergent norecurse nounwind
20+
define dso_local spir_func bfloat @test_AtomicFAddEXT_bfloat(ptr addrspace(4) align 2 dereferenceable(4) %Arg) {
21+
entry:
22+
%0 = addrspacecast ptr addrspace(4) %Arg to ptr addrspace(1)
23+
; CHECK-SPIRV: AtomicFAddEXT [[BFLOAT]]
24+
; CHECK-LLVM-SPV: call spir_func bfloat @_Z21__spirv_AtomicFAddEXTPU3AS1u6__bf16iiu6__bf16({{.*}}bfloat
25+
%ret = tail call spir_func bfloat @_Z21__spirv_AtomicFAddEXTPU3AS1u6__bf16iiu6__bf16(ptr addrspace(1) %0, i32 1, i32 896, bfloat 1.000000e+00)
26+
ret bfloat %ret
27+
}
28+
29+
; Function Attrs: convergent
30+
declare dso_local spir_func bfloat @_Z21__spirv_AtomicFAddEXTPU3AS1u6__bf16iiu6__bf16(ptr addrspace(1), i32, i32, bfloat)
Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
; RUN: llvm-as %s -o %t.bc
2+
; RUN: llvm-spirv %t.bc --spirv-ext=+SPV_INTEL_shader_atomic_bfloat16,+SPV_KHR_bfloat16 -o %t.spv
3+
; RUN: llvm-spirv -to-text %t.spv -o %t.spt
4+
; RUN: FileCheck < %t.spt %s --check-prefix=CHECK-SPIRV
5+
6+
; RUN: llvm-spirv --spirv-target-env=SPV-IR -r %t.spv -o %t.rev.bc
7+
; RUN: llvm-dis %t.rev.bc -o - | FileCheck %s --check-prefixes=CHECK-LLVM-SPV
8+
9+
target datalayout = "e-i64:64-v16:16-v24:32-v32:32-v48:64-v96:128-v192:256-v256:256-v512:512-v1024:1024-n8:16:32:64"
10+
target triple = "spir64-unknown-unknown"
11+
12+
; CHECK-SPIRV-DAG: Capability AtomicBFloat16MinMaxINTEL
13+
; CHECK-SPIRV-DAG: Capability BFloat16TypeKHR
14+
; CHECK-SPIRV-DAG: Extension "SPV_INTEL_shader_atomic_bfloat16"
15+
; CHECK-SPIRV-DAG: Extension "SPV_KHR_bfloat16"
16+
17+
; CHECK-SPIRV: TypeFloat [[BFLOAT:[0-9]+]] 16 0
18+
19+
; Function Attrs: convergent norecurse nounwind
20+
define dso_local spir_func bfloat @test_AtomicFMaxEXT_bfloat(ptr addrspace(4) align 2 dereferenceable(4) %Arg) {
21+
entry:
22+
%0 = addrspacecast ptr addrspace(4) %Arg to ptr addrspace(1)
23+
; CHECK-SPIRV: AtomicFMaxEXT [[BFLOAT]]
24+
; CHECK-LLVM-SPV: call spir_func bfloat @_Z21__spirv_AtomicFMaxEXTPU3AS1u6__bf16iiu6__bf16({{.*}}bfloat
25+
%ret = tail call spir_func bfloat @_Z21__spirv_AtomicFMaxEXTPU3AS1u6__bf16iiu6__bf16(ptr addrspace(1) %0, i32 1, i32 896, bfloat 1.000000e+00)
26+
ret bfloat %ret
27+
}
28+
29+
; Function Attrs: convergent
30+
declare dso_local spir_func bfloat @_Z21__spirv_AtomicFMaxEXTPU3AS1u6__bf16iiu6__bf16(ptr addrspace(1), i32, i32, bfloat)
Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
; RUN: llvm-as %s -o %t.bc
2+
; RUN: llvm-spirv %t.bc --spirv-ext=+SPV_INTEL_shader_atomic_bfloat16,+SPV_KHR_bfloat16 -o %t.spv
3+
; RUN: llvm-spirv -to-text %t.spv -o %t.spt
4+
; RUN: FileCheck < %t.spt %s --check-prefix=CHECK-SPIRV
5+
6+
; RUN: llvm-spirv --spirv-target-env=SPV-IR -r %t.spv -o %t.rev.bc
7+
; RUN: llvm-dis %t.rev.bc -o - | FileCheck %s --check-prefixes=CHECK-LLVM-SPV
8+
9+
target datalayout = "e-i64:64-v16:16-v24:32-v32:32-v48:64-v96:128-v192:256-v256:256-v512:512-v1024:1024-n8:16:32:64"
10+
target triple = "spir64-unknown-unknown"
11+
12+
; CHECK-SPIRV-DAG: Capability AtomicBFloat16MinMaxINTEL
13+
; CHECK-SPIRV-DAG: Capability BFloat16TypeKHR
14+
; CHECK-SPIRV-DAG: Extension "SPV_INTEL_shader_atomic_bfloat16"
15+
; CHECK-SPIRV-DAG: Extension "SPV_KHR_bfloat16"
16+
17+
; CHECK-SPIRV: TypeFloat [[BFLOAT:[0-9]+]] 16 0
18+
19+
; Function Attrs: convergent norecurse nounwind
20+
define dso_local spir_func bfloat @test_AtomicFMinEXT_bfloat(ptr addrspace(4) align 2 dereferenceable(4) %Arg) {
21+
entry:
22+
%0 = addrspacecast ptr addrspace(4) %Arg to ptr addrspace(1)
23+
; CHECK-SPIRV: AtomicFMinEXT [[BFLOAT]]
24+
; CHECK-LLVM-SPV: call spir_func bfloat @_Z21__spirv_AtomicFMinEXTPU3AS1u6__bf16iiu6__bf16({{.*}}bfloat
25+
%ret = tail call spir_func bfloat @_Z21__spirv_AtomicFMinEXTPU3AS1u6__bf16iiu6__bf16(ptr addrspace(1) %0, i32 1, i32 896, bfloat 1.000000e+00)
26+
ret bfloat %ret
27+
}
28+
29+
; Function Attrs: convergent
30+
declare dso_local spir_func bfloat @_Z21__spirv_AtomicFMinEXTPU3AS1u6__bf16iiu6__bf16(ptr addrspace(1), i32, i32, bfloat)
Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,34 @@
1+
; RUN: llvm-as < %s -o %t.bc
2+
; RUN: llvm-spirv --spirv-ext=+SPV_INTEL_shader_atomic_bfloat16,+SPV_KHR_bfloat16 %t.bc -o %t.spv
3+
; RUN: llvm-spirv -to-text %t.spv -o - | FileCheck %s
4+
5+
; CHECK-DAG: Extension "SPV_INTEL_shader_atomic_bfloat16"
6+
; CHECK-DAG: Extension "SPV_KHR_bfloat16"
7+
; CHECK-DAG: Capability AtomicBFloat16AddINTEL
8+
; CHECK-DAG: Capability BFloat16TypeKHR
9+
; CHECK: TypeInt [[Int:[0-9]+]] 32 0
10+
; CHECK-DAG: Constant [[Int]] [[Scope_CrossDevice:[0-9]+]] 0 {{$}}
11+
; CHECK-DAG: Constant [[Int]] [[MemSem_SequentiallyConsistent:[0-9]+]] 16
12+
; CHECK: TypeFloat [[BFloat:[0-9]+]] 16 0
13+
; CHECK: Variable {{[0-9]+}} [[BFloatPointer:[0-9]+]]
14+
; CHECK: Constant [[BFloat]] [[BFloatValue:[0-9]+]] 16936
15+
16+
target datalayout = "e-i64:64-v16:16-v24:32-v32:32-v48:64-v96:128-v192:256-v256:256-v512:512-v1024:1024"
17+
target triple = "spir64"
18+
19+
@f = common dso_local local_unnamed_addr addrspace(1) global bfloat 0.000000e+00, align 8
20+
21+
; Function Attrs: nounwind
22+
define dso_local spir_func void @test_atomicrmw_fadd() local_unnamed_addr #0 {
23+
entry:
24+
%0 = atomicrmw fadd ptr addrspace(1) @f, bfloat 42.000000e+00 seq_cst
25+
; CHECK: AtomicFAddEXT [[BFloat]] {{[0-9]+}} [[BFloatPointer]] [[Scope_CrossDevice]] [[MemSem_SequentiallyConsistent]] [[BFloatValue]]
26+
27+
ret void
28+
}
29+
30+
attributes #0 = { nounwind "correctly-rounded-divide-sqrt-fp-math"="false" "disable-tail-calls"="false" "frame-pointer"="all" "less-precise-fpmad"="false" "min-legal-vector-width"="0" "no-infs-fp-math"="false" "no-jump-tables"="false" "no-nans-fp-math"="false" "no-signed-zeros-fp-math"="false" "no-trapping-math"="false" "stack-protector-buffer-size"="8" "unsafe-fp-math"="false" "use-soft-float"="false" }
31+
32+
!llvm.module.flags = !{!0}
33+
34+
!0 = !{i32 1, !"wchar_size", i32 4}
Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
; RUN: llvm-as < %s -o %t.bc
2+
; RUN: llvm-spirv --spirv-ext=+SPV_INTEL_shader_atomic_bfloat16,+SPV_KHR_bfloat16 %t.bc -o %t.spv
3+
; RUN: llvm-spirv -to-text %t.spv -o - | FileCheck %s
4+
5+
; CHECK-DAG: Extension "SPV_INTEL_shader_atomic_bfloat16"
6+
; CHECK-DAG: Extension "SPV_KHR_bfloat16"
7+
; CHECK-DAG: AtomicBFloat16MinMaxINTEL
8+
; CHECK-DAG: Capability BFloat16TypeKHR
9+
; CHECK: TypeInt [[Int:[0-9]+]] 32 0
10+
; CHECK-DAG: Constant [[Int]] [[Scope_CrossDevice:[0-9]+]] 0 {{$}}
11+
; CHECK-DAG: Constant [[Int]] [[MemSem_SequentiallyConsistent:[0-9]+]] 16
12+
; CHECK: TypeFloat [[BFloat:[0-9]+]] 16 0
13+
; CHECK: Variable {{[0-9]+}} [[BFloatPointer:[0-9]+]]
14+
; CHECK: Constant [[BFloat]] [[BFloatValue:[0-9]+]] 16936
15+
16+
target datalayout = "e-i64:64-v16:16-v24:32-v32:32-v48:64-v96:128-v192:256-v256:256-v512:512-v1024:1024"
17+
target triple = "spir64"
18+
19+
@f = common dso_local local_unnamed_addr addrspace(1) global bfloat 0.000000e+00, align 4
20+
21+
; Function Attrs: nounwind
22+
define dso_local spir_func void @test_atomicrmw_fadd() local_unnamed_addr #0 {
23+
entry:
24+
%0 = atomicrmw fmin ptr addrspace(1) @f, bfloat 42.000000e+00 seq_cst
25+
; CHECK: AtomicFMinEXT [[BFloat]] {{[0-9]+}} [[BFloatPointer]] [[Scope_CrossDevice]] [[MemSem_SequentiallyConsistent]] [[BFloatValue]]
26+
%1 = atomicrmw fmax ptr addrspace(1) @f, bfloat 42.000000e+00 seq_cst
27+
; CHECK: AtomicFMaxEXT [[BFloat]] {{[0-9]+}} [[BFloatPointer]] [[Scope_CrossDevice]] [[MemSem_SequentiallyConsistent]] [[BFloatValue]]
28+
29+
ret void
30+
}
31+
32+
attributes #0 = { nounwind "correctly-rounded-divide-sqrt-fp-math"="false" "disable-tail-calls"="false" "frame-pointer"="all" "less-precise-fpmad"="false" "min-legal-vector-width"="0" "no-infs-fp-math"="false" "no-jump-tables"="false" "no-nans-fp-math"="false" "no-signed-zeros-fp-math"="false" "no-trapping-math"="false" "stack-protector-buffer-size"="8" "unsafe-fp-math"="false" "use-soft-float"="false" }
33+
34+
!llvm.module.flags = !{!0}
35+
36+
!0 = !{i32 1, !"wchar_size", i32 4}
Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,29 @@
1+
; RUN: llvm-as < %s -o %t.bc
2+
; RUN: not llvm-spirv --spirv-ext=+SPV_INTEL_shader_atomic_bfloat16 %t.bc 2>&1 | FileCheck %s --check-prefix=CHECK-NO-BF
3+
; RUN: not llvm-spirv --spirv-ext=+SPV_KHR_bfloat16 %t.bc 2>&1 | FileCheck %s --check-prefix=CHECK-NO-ATOM
4+
5+
; CHECK-NO-BF: RequiresExtension: Feature requires the following SPIR-V extension:
6+
; CHECK-NO-BF-NEXT: SPV_KHR_bfloat16
7+
; CHECK-NO-BF-NEXT: NOTE: LLVM module contains bfloat type, translation of which requires this extension
8+
9+
; CHECK-NO-ATOM: RequiresExtension: Feature requires the following SPIR-V extension:
10+
; CHECK-NO-ATOM-NEXT: SPV_INTEL_shader_atomic_bfloat16
11+
12+
target datalayout = "e-i64:64-v16:16-v24:32-v32:32-v48:64-v96:128-v192:256-v256:256-v512:512-v1024:1024"
13+
target triple = "spir64"
14+
15+
@f = common dso_local local_unnamed_addr addrspace(1) global bfloat 0.000000e+00, align 8
16+
17+
; Function Attrs: nounwind
18+
define dso_local spir_func void @test_atomicrmw_fadd() local_unnamed_addr #0 {
19+
entry:
20+
%0 = atomicrmw fadd ptr addrspace(1) @f, bfloat 42.000000e+00 seq_cst
21+
22+
ret void
23+
}
24+
25+
attributes #0 = { nounwind "correctly-rounded-divide-sqrt-fp-math"="false" "disable-tail-calls"="false" "frame-pointer"="all" "less-precise-fpmad"="false" "min-legal-vector-width"="0" "no-infs-fp-math"="false" "no-jump-tables"="false" "no-nans-fp-math"="false" "no-signed-zeros-fp-math"="false" "no-trapping-math"="false" "stack-protector-buffer-size"="8" "unsafe-fp-math"="false" "use-soft-float"="false" }
26+
27+
!llvm.module.flags = !{!0}
28+
29+
!0 = !{i32 1, !"wchar_size", i32 4}

0 commit comments

Comments
 (0)