Skip to content

Commit b7b12ab

Browse files
author
Dmitry Sidorov
authored
[Backport to 14] Implement SPV_INTEL_bfloat16_arithmetic (KhronosGroup#3290) (KhronosGroup#3320) (KhronosGroup#3340)
The extension relaxes rules for bf16 type allowing to use it in some arithmetic operations. Spec is available here: intel/llvm#18352 Co-authered by: Michael Aziz <michael.aziz@intel.com> --------- Signed-off-by: Sidorov, Dmitry <dmitry.sidorov@intel.com>
1 parent 3ac1cbf commit b7b12ab

9 files changed

Lines changed: 314 additions & 0 deletions

File tree

include/LLVMSPIRVExtensions.inc

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -72,3 +72,4 @@ EXT(SPV_INTEL_bindless_images)
7272
EXT(SPV_INTEL_2d_block_io)
7373
EXT(SPV_INTEL_subgroup_matrix_multiply_accumulate)
7474
EXT(SPV_KHR_bfloat16)
75+
EXT(SPV_INTEL_bfloat16_arithmetic)

lib/SPIRV/SPIRVUtil.cpp

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -618,6 +618,11 @@ ParamType lastFuncParamType(StringRef MangledName) {
618618
char Mangled = Copy.back();
619619
std::string Mangled2 = Copy.substr(Copy.size() - 2);
620620

621+
std::string Mangled6 = Copy.substr(Copy.size() - 6);
622+
if (Mangled6 == "__bf16") {
623+
return ParamType::FLOAT;
624+
}
625+
621626
if (isMangledTypeFP(Mangled) || isMangledTypeHalf(Mangled2)) {
622627
return ParamType::FLOAT;
623628
} else if (isMangledTypeUnsigned(Mangled)) {
@@ -1847,6 +1852,9 @@ bool checkTypeForSPIRVExtendedInstLowering(IntrinsicInst *II, SPIRVModule *BM) {
18471852
NumElems = VecTy->getNumElements();
18481853
Ty = VecTy->getElementType();
18491854
}
1855+
if (Ty->isBFloatTy() &&
1856+
BM->hasCapability(internal::CapabilityBFloat16ArithmeticINTEL))
1857+
return true;
18501858
if ((!Ty->isFloatTy() && !Ty->isDoubleTy() && !Ty->isHalfTy()) ||
18511859
(!BM->hasCapability(CapabilityVectorAnyINTEL) &&
18521860
((NumElems > 4) && (NumElems != 8) && (NumElems != 16)))) {
@@ -1863,6 +1871,9 @@ bool checkTypeForSPIRVExtendedInstLowering(IntrinsicInst *II, SPIRVModule *BM) {
18631871
NumElems = VecTy->getNumElements();
18641872
Ty = VecTy->getElementType();
18651873
}
1874+
if (Ty->isBFloatTy() &&
1875+
BM->hasCapability(internal::CapabilityBFloat16ArithmeticINTEL))
1876+
return true;
18661877
if ((!Ty->isIntegerTy()) ||
18671878
(!BM->hasCapability(CapabilityVectorAnyINTEL) &&
18681879
((NumElems > 4) && (NumElems != 8) && (NumElems != 16)))) {

lib/SPIRV/SPIRVWriter.cpp

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3630,6 +3630,20 @@ SPIRVValue *LLVMToSPIRVBase::transIntrinsicInst(IntrinsicInst *II,
36303630
// -spirv-allow-unknown-intrinsics work correctly.
36313631
auto IID = II->getIntrinsicID();
36323632
switch (IID) {
3633+
case Intrinsic::fabs:
3634+
case Intrinsic::fma:
3635+
case Intrinsic::maxnum:
3636+
case Intrinsic::minnum:
3637+
case Intrinsic::fmuladd: {
3638+
Type *Ty = II->getType();
3639+
if (Ty->isBFloatTy())
3640+
BM->addCapability(internal::CapabilityBFloat16ArithmeticINTEL);
3641+
break;
3642+
}
3643+
default:
3644+
break;
3645+
}
3646+
switch (IID) {
36333647
case Intrinsic::assume: {
36343648
// llvm.assume translation is currently supported only within
36353649
// SPV_KHR_expect_assume extension, ignore it otherwise, since it's
@@ -4413,6 +4427,11 @@ SPIRVValue *LLVMToSPIRVBase::transDirectCallInst(CallInst *CI,
44134427
SmallVector<std::string, 2> Dec;
44144428
if (isBuiltinTransToExtInst(CI->getCalledFunction(), &ExtSetKind, &ExtOp,
44154429
&Dec)) {
4430+
if (const auto *FirstArg = F->getArg(0)) {
4431+
const auto *Type = FirstArg->getType();
4432+
if (Type->isBFloatTy())
4433+
BM->addCapability(internal::CapabilityBFloat16ArithmeticINTEL);
4434+
}
44164435
if (DemangledName.find("__spirv_ocl_printf") != StringRef::npos) {
44174436
auto *FormatStrPtr = cast<PointerType>(CI->getArgOperand(0)->getType());
44184437
if (FormatStrPtr->getAddressSpace() !=

lib/SPIRV/libSPIRV/SPIRVEntry.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -904,6 +904,8 @@ class SPIRVCapability : public SPIRVEntryNoId<OpCapability> {
904904
case CapabilityVectorComputeINTEL:
905905
case CapabilityVectorAnyINTEL:
906906
return ExtensionID::SPV_INTEL_vector_compute;
907+
case internal::CapabilityBFloat16ArithmeticINTEL:
908+
return ExtensionID::SPV_INTEL_bfloat16_arithmetic;
907909
default:
908910
return {};
909911
}

lib/SPIRV/libSPIRV/SPIRVEnum.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -214,6 +214,8 @@ template <> inline void SPIRVMap<SPIRVCapabilityKind, SPIRVCapVec>::init() {
214214
ADD_VEC_INIT(CapabilityBFloat16DotProductKHR, {CapabilityBFloat16TypeKHR});
215215
ADD_VEC_INIT(CapabilityBFloat16CooperativeMatrixKHR,
216216
{CapabilityBFloat16TypeKHR, CapabilityCooperativeMatrixKHR});
217+
ADD_VEC_INIT(internal::CapabilityBFloat16ArithmeticINTEL,
218+
{CapabilityBFloat16TypeKHR});
217219
}
218220

219221
template <> inline void SPIRVMap<SPIRVExecutionModelKind, SPIRVCapVec>::init() {

lib/SPIRV/libSPIRV/SPIRVModule.cpp

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1543,6 +1543,8 @@ SPIRVInstruction *SPIRVModuleImpl::addBinaryInst(Op TheOpCode, SPIRVType *Type,
15431543
SPIRVValue *Op1,
15441544
SPIRVValue *Op2,
15451545
SPIRVBasicBlock *BB) {
1546+
if (Type->isTypeFloat(16, FPEncodingBFloat16KHR) && TheOpCode != OpDot)
1547+
addCapability(internal::CapabilityBFloat16ArithmeticINTEL);
15461548
return addInstruction(SPIRVInstTemplateBase::create(
15471549
TheOpCode, Type, getId(),
15481550
getVec(Op1->getId(), Op2->getId()), BB, this),
@@ -1566,6 +1568,8 @@ SPIRVInstruction *SPIRVModuleImpl::addUnaryInst(Op TheOpCode,
15661568
SPIRVType *TheType,
15671569
SPIRVValue *Op,
15681570
SPIRVBasicBlock *BB) {
1571+
if (TheType->isTypeFloat(16, FPEncodingBFloat16KHR) && TheOpCode != OpDot)
1572+
addCapability(internal::CapabilityBFloat16ArithmeticINTEL);
15691573
return addInstruction(
15701574
SPIRVInstTemplateBase::create(TheOpCode, TheType, getId(),
15711575
getVec(Op->getId()), BB, this),

lib/SPIRV/libSPIRV/SPIRVNameMapEnum.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -646,6 +646,7 @@ template <> inline void SPIRVMap<Capability, std::string>::init() {
646646
add(internal::CapabilityCooperativeMatrixCheckedInstructionsINTEL,
647647
"CooperativeMatrixCheckedInstructionsINTEL");
648648
add(internal::CapabilityBindlessImagesINTEL, "BindlessImagesINTEL");
649+
add(internal::CapabilityBFloat16ArithmeticINTEL, "BFloat16ArithmeticINTEL");
649650
}
650651
SPIRV_DEF_NAMEMAP(Capability, SPIRVCapabilityNameMap)
651652

lib/SPIRV/libSPIRV/spirv_internal.hpp

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -114,6 +114,7 @@ enum InternalCapability {
114114
ICapFPArithmeticFenceINTEL = 6144,
115115
ICapGlobalVariableDecorationsINTEL = 6146,
116116
ICapabilityCooperativeMatrixCheckedInstructionsINTEL = 6192,
117+
ICapabilityBFloat16ArithmeticINTEL = 6226,
117118
ICapabilityCooperativeMatrixPrefetchINTEL = 6411,
118119
ICapabilityComplexFloatMulDivINTEL = 6414,
119120
ICapabilityTensorFloat32RoundingINTEL = 6425,
@@ -304,6 +305,8 @@ constexpr Capability CapabilityGlobalVariableDecorationsINTEL =
304305
static_cast<Capability>(ICapGlobalVariableDecorationsINTEL);
305306
constexpr Capability CapabilityRegisterLimitsINTEL =
306307
static_cast<Capability>(ICapRegisterLimitsINTEL);
308+
constexpr Capability CapabilityBFloat16ArithmeticINTEL =
309+
static_cast<Capability>(ICapabilityBFloat16ArithmeticINTEL);
307310

308311
constexpr FunctionControlMask FunctionControlOptNoneINTELMask =
309312
static_cast<FunctionControlMask>(IFunctionControlOptNoneINTELMask);

0 commit comments

Comments
 (0)