| 1 | //===- DXILIntrinsicExpansion.cpp - Prepare LLVM Module for DXIL encoding--===// |
| 2 | // |
| 3 | // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. |
| 4 | // See https://llvm.org/LICENSE.txt for license information. |
| 5 | // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception |
| 6 | // |
| 7 | //===----------------------------------------------------------------------===// |
| 8 | /// |
| 9 | /// \file This file contains DXIL intrinsic expansions for those that don't have |
| 10 | // opcodes in DirectX Intermediate Language (DXIL). |
| 11 | //===----------------------------------------------------------------------===// |
| 12 | |
| 13 | #include "DXILIntrinsicExpansion.h" |
| 14 | #include "DirectX.h" |
| 15 | #include "llvm/ADT/APInt.h" |
| 16 | #include "llvm/ADT/STLExtras.h" |
| 17 | #include "llvm/ADT/SmallVector.h" |
| 18 | #include "llvm/CodeGen/Passes.h" |
| 19 | #include "llvm/IR/Constants.h" |
| 20 | #include "llvm/IR/IRBuilder.h" |
| 21 | #include "llvm/IR/InstrTypes.h" |
| 22 | #include "llvm/IR/Instruction.h" |
| 23 | #include "llvm/IR/Instructions.h" |
| 24 | #include "llvm/IR/Intrinsics.h" |
| 25 | #include "llvm/IR/IntrinsicsDirectX.h" |
| 26 | #include "llvm/IR/MatrixBuilder.h" |
| 27 | #include "llvm/IR/Module.h" |
| 28 | #include "llvm/IR/PassManager.h" |
| 29 | #include "llvm/IR/Type.h" |
| 30 | #include "llvm/Pass.h" |
| 31 | #include "llvm/Support/Casting.h" |
| 32 | #include "llvm/Support/ErrorHandling.h" |
| 33 | #include "llvm/Support/MathExtras.h" |
| 34 | |
| 35 | #define DEBUG_TYPE "dxil-intrinsic-expansion" |
| 36 | |
| 37 | using namespace llvm; |
| 38 | |
| 39 | class DXILIntrinsicExpansionLegacy : public ModulePass { |
| 40 | |
| 41 | public: |
| 42 | bool runOnModule(Module &M) override; |
| 43 | DXILIntrinsicExpansionLegacy() : ModulePass(ID) {} |
| 44 | |
| 45 | static char ID; // Pass identification. |
| 46 | }; |
| 47 | |
| 48 | static bool resourceAccessNeeds64BitExpansion(Module *M, Type *OverloadTy, |
| 49 | bool IsRaw) { |
| 50 | if (IsRaw && M->getTargetTriple().getDXILVersion() > VersionTuple(1, 2)) |
| 51 | return false; |
| 52 | |
| 53 | Type *ScalarTy = OverloadTy->getScalarType(); |
| 54 | return ScalarTy->isDoubleTy() || ScalarTy->isIntegerTy(BitWidth: 64); |
| 55 | } |
| 56 | |
| 57 | static Value *expand16BitIsInf(CallInst *Orig) { |
| 58 | Module *M = Orig->getModule(); |
| 59 | if (M->getTargetTriple().getDXILVersion() >= VersionTuple(1, 9)) |
| 60 | return nullptr; |
| 61 | |
| 62 | Value *Val = Orig->getOperand(i_nocapture: 0); |
| 63 | Type *ValTy = Val->getType(); |
| 64 | if (!ValTy->getScalarType()->isHalfTy()) |
| 65 | return nullptr; |
| 66 | |
| 67 | IRBuilder<> Builder(Orig); |
| 68 | Type *IType = Type::getInt16Ty(C&: M->getContext()); |
| 69 | Constant *PosInf = |
| 70 | ValTy->isVectorTy() |
| 71 | ? ConstantVector::getSplat( |
| 72 | EC: ElementCount::getFixed( |
| 73 | MinVal: cast<FixedVectorType>(Val: ValTy)->getNumElements()), |
| 74 | Elt: ConstantInt::get(Ty: IType, V: 0x7c00)) |
| 75 | : ConstantInt::get(Ty: IType, V: 0x7c00); |
| 76 | |
| 77 | Constant *NegInf = |
| 78 | ValTy->isVectorTy() |
| 79 | ? ConstantVector::getSplat( |
| 80 | EC: ElementCount::getFixed( |
| 81 | MinVal: cast<FixedVectorType>(Val: ValTy)->getNumElements()), |
| 82 | Elt: ConstantInt::get(Ty: IType, V: 0xfc00)) |
| 83 | : ConstantInt::get(Ty: IType, V: 0xfc00); |
| 84 | |
| 85 | Value *IVal = Builder.CreateBitCast(V: Val, DestTy: PosInf->getType()); |
| 86 | Value *B1 = Builder.CreateICmpEQ(LHS: IVal, RHS: PosInf); |
| 87 | Value *B2 = Builder.CreateICmpEQ(LHS: IVal, RHS: NegInf); |
| 88 | Value *B3 = Builder.CreateOr(LHS: B1, RHS: B2); |
| 89 | return B3; |
| 90 | } |
| 91 | |
| 92 | static Value *expand16BitIsNaN(CallInst *Orig) { |
| 93 | Module *M = Orig->getModule(); |
| 94 | if (M->getTargetTriple().getDXILVersion() >= VersionTuple(1, 9)) |
| 95 | return nullptr; |
| 96 | |
| 97 | Value *Val = Orig->getOperand(i_nocapture: 0); |
| 98 | Type *ValTy = Val->getType(); |
| 99 | if (!ValTy->getScalarType()->isHalfTy()) |
| 100 | return nullptr; |
| 101 | |
| 102 | IRBuilder<> Builder(Orig); |
| 103 | Type *IType = Type::getInt16Ty(C&: M->getContext()); |
| 104 | |
| 105 | Constant *ExpBitMask = |
| 106 | ValTy->isVectorTy() |
| 107 | ? ConstantVector::getSplat( |
| 108 | EC: ElementCount::getFixed( |
| 109 | MinVal: cast<FixedVectorType>(Val: ValTy)->getNumElements()), |
| 110 | Elt: ConstantInt::get(Ty: IType, V: 0x7c00)) |
| 111 | : ConstantInt::get(Ty: IType, V: 0x7c00); |
| 112 | Constant *SigBitMask = |
| 113 | ValTy->isVectorTy() |
| 114 | ? ConstantVector::getSplat( |
| 115 | EC: ElementCount::getFixed( |
| 116 | MinVal: cast<FixedVectorType>(Val: ValTy)->getNumElements()), |
| 117 | Elt: ConstantInt::get(Ty: IType, V: 0x3ff)) |
| 118 | : ConstantInt::get(Ty: IType, V: 0x3ff); |
| 119 | |
| 120 | Constant *Zero = |
| 121 | ValTy->isVectorTy() |
| 122 | ? ConstantVector::getSplat( |
| 123 | EC: ElementCount::getFixed( |
| 124 | MinVal: cast<FixedVectorType>(Val: ValTy)->getNumElements()), |
| 125 | Elt: ConstantInt::get(Ty: IType, V: 0)) |
| 126 | : ConstantInt::get(Ty: IType, V: 0); |
| 127 | |
| 128 | Value *IVal = Builder.CreateBitCast(V: Val, DestTy: ExpBitMask->getType()); |
| 129 | Value *Exp = Builder.CreateAnd(LHS: IVal, RHS: ExpBitMask); |
| 130 | Value *B1 = Builder.CreateICmpEQ(LHS: Exp, RHS: ExpBitMask); |
| 131 | |
| 132 | Value *Sig = Builder.CreateAnd(LHS: IVal, RHS: SigBitMask); |
| 133 | Value *B2 = Builder.CreateICmpNE(LHS: Sig, RHS: Zero); |
| 134 | Value *B3 = Builder.CreateAnd(LHS: B1, RHS: B2); |
| 135 | return B3; |
| 136 | } |
| 137 | |
| 138 | static Value *expand16BitIsFinite(CallInst *Orig) { |
| 139 | Module *M = Orig->getModule(); |
| 140 | if (M->getTargetTriple().getDXILVersion() >= VersionTuple(1, 9)) |
| 141 | return nullptr; |
| 142 | |
| 143 | Value *Val = Orig->getOperand(i_nocapture: 0); |
| 144 | Type *ValTy = Val->getType(); |
| 145 | if (!ValTy->getScalarType()->isHalfTy()) |
| 146 | return nullptr; |
| 147 | |
| 148 | IRBuilder<> Builder(Orig); |
| 149 | Type *IType = Type::getInt16Ty(C&: M->getContext()); |
| 150 | |
| 151 | Constant *ExpBitMask = |
| 152 | ValTy->isVectorTy() |
| 153 | ? ConstantVector::getSplat( |
| 154 | EC: ElementCount::getFixed( |
| 155 | MinVal: cast<FixedVectorType>(Val: ValTy)->getNumElements()), |
| 156 | Elt: ConstantInt::get(Ty: IType, V: 0x7c00)) |
| 157 | : ConstantInt::get(Ty: IType, V: 0x7c00); |
| 158 | |
| 159 | Value *IVal = Builder.CreateBitCast(V: Val, DestTy: ExpBitMask->getType()); |
| 160 | Value *Exp = Builder.CreateAnd(LHS: IVal, RHS: ExpBitMask); |
| 161 | Value *B1 = Builder.CreateICmpNE(LHS: Exp, RHS: ExpBitMask); |
| 162 | return B1; |
| 163 | } |
| 164 | |
| 165 | static Value *expand16BitIsNormal(CallInst *Orig) { |
| 166 | Module *M = Orig->getModule(); |
| 167 | if (M->getTargetTriple().getDXILVersion() >= VersionTuple(1, 9)) |
| 168 | return nullptr; |
| 169 | |
| 170 | Value *Val = Orig->getOperand(i_nocapture: 0); |
| 171 | Type *ValTy = Val->getType(); |
| 172 | if (!ValTy->getScalarType()->isHalfTy()) |
| 173 | return nullptr; |
| 174 | |
| 175 | IRBuilder<> Builder(Orig); |
| 176 | Type *IType = Type::getInt16Ty(C&: M->getContext()); |
| 177 | |
| 178 | Constant *ExpBitMask = |
| 179 | ValTy->isVectorTy() |
| 180 | ? ConstantVector::getSplat( |
| 181 | EC: ElementCount::getFixed( |
| 182 | MinVal: cast<FixedVectorType>(Val: ValTy)->getNumElements()), |
| 183 | Elt: ConstantInt::get(Ty: IType, V: 0x7c00)) |
| 184 | : ConstantInt::get(Ty: IType, V: 0x7c00); |
| 185 | Constant *Zero = |
| 186 | ValTy->isVectorTy() |
| 187 | ? ConstantVector::getSplat( |
| 188 | EC: ElementCount::getFixed( |
| 189 | MinVal: cast<FixedVectorType>(Val: ValTy)->getNumElements()), |
| 190 | Elt: ConstantInt::get(Ty: IType, V: 0)) |
| 191 | : ConstantInt::get(Ty: IType, V: 0); |
| 192 | |
| 193 | Value *IVal = Builder.CreateBitCast(V: Val, DestTy: ExpBitMask->getType()); |
| 194 | Value *Exp = Builder.CreateAnd(LHS: IVal, RHS: ExpBitMask); |
| 195 | Value *NotAllZeroes = Builder.CreateICmpNE(LHS: Exp, RHS: Zero); |
| 196 | Value *NotAllOnes = Builder.CreateICmpNE(LHS: Exp, RHS: ExpBitMask); |
| 197 | Value *B1 = Builder.CreateAnd(LHS: NotAllZeroes, RHS: NotAllOnes); |
| 198 | return B1; |
| 199 | } |
| 200 | |
| 201 | static bool shouldExpandFloatDotIntrinsic(Function &F) { |
| 202 | assert(F.getIntrinsicID() == Intrinsic::dx_fdot && |
| 203 | "Function is not a dx.fdot intrinsic" ); |
| 204 | auto *ParamTy = cast<FixedVectorType>(Val: F.getFunctionType()->getParamType(i: 0)); |
| 205 | return ParamTy->getNumElements() <= 4 || |
| 206 | F.getParent()->getTargetTriple().getOSVersion() < VersionTuple(6, 9); |
| 207 | } |
| 208 | |
| 209 | static bool isIntrinsicExpansion(Function &F) { |
| 210 | switch (F.getIntrinsicID()) { |
| 211 | case Intrinsic::assume: |
| 212 | case Intrinsic::abs: |
| 213 | case Intrinsic::atan2: |
| 214 | case Intrinsic::copysign: |
| 215 | case Intrinsic::fshl: |
| 216 | case Intrinsic::fshr: |
| 217 | case Intrinsic::exp: |
| 218 | case Intrinsic::is_fpclass: |
| 219 | case Intrinsic::log: |
| 220 | case Intrinsic::log10: |
| 221 | case Intrinsic::pow: |
| 222 | case Intrinsic::powi: |
| 223 | case Intrinsic::dx_all: |
| 224 | case Intrinsic::dx_any: |
| 225 | case Intrinsic::dx_uclamp: |
| 226 | case Intrinsic::dx_sclamp: |
| 227 | case Intrinsic::dx_nclamp: |
| 228 | case Intrinsic::dx_isfinite: |
| 229 | case Intrinsic::dx_isinf: |
| 230 | case Intrinsic::dx_isnan: |
| 231 | case Intrinsic::dx_sdot: |
| 232 | case Intrinsic::dx_udot: |
| 233 | case Intrinsic::dx_sign: |
| 234 | case Intrinsic::usub_sat: |
| 235 | case Intrinsic::vector_reduce_add: |
| 236 | case Intrinsic::vector_reduce_fadd: |
| 237 | case Intrinsic::matrix_multiply: |
| 238 | case Intrinsic::matrix_transpose: |
| 239 | case Intrinsic::umul_with_overflow: |
| 240 | case Intrinsic::smul_with_overflow: |
| 241 | case Intrinsic::dx_load_input: |
| 242 | case Intrinsic::dx_store_output: |
| 243 | return true; |
| 244 | case Intrinsic::dx_fdot: |
| 245 | return shouldExpandFloatDotIntrinsic(F); |
| 246 | case Intrinsic::dx_resource_load_rawbuffer: |
| 247 | return resourceAccessNeeds64BitExpansion( |
| 248 | M: F.getParent(), OverloadTy: F.getReturnType()->getStructElementType(N: 0), |
| 249 | /*IsRaw*/ true); |
| 250 | case Intrinsic::dx_resource_load_typedbuffer: |
| 251 | return resourceAccessNeeds64BitExpansion( |
| 252 | M: F.getParent(), OverloadTy: F.getReturnType()->getStructElementType(N: 0), |
| 253 | /*IsRaw*/ false); |
| 254 | case Intrinsic::dx_resource_store_rawbuffer: |
| 255 | return resourceAccessNeeds64BitExpansion( |
| 256 | M: F.getParent(), OverloadTy: F.getFunctionType()->getParamType(i: 3), /*IsRaw*/ true); |
| 257 | case Intrinsic::dx_resource_store_typedbuffer: |
| 258 | return resourceAccessNeeds64BitExpansion( |
| 259 | M: F.getParent(), OverloadTy: F.getFunctionType()->getParamType(i: 2), /*IsRaw*/ false); |
| 260 | } |
| 261 | return false; |
| 262 | } |
| 263 | |
| 264 | static Value *expandUsubSat(CallInst *Orig) { |
| 265 | Value *A = Orig->getArgOperand(i: 0); |
| 266 | Value *B = Orig->getArgOperand(i: 1); |
| 267 | Type *Ty = A->getType(); |
| 268 | |
| 269 | IRBuilder<> Builder(Orig); |
| 270 | |
| 271 | Value *Cmp = Builder.CreateICmpULT(LHS: A, RHS: B, Name: "usub.cmp" ); |
| 272 | Value *Sub = Builder.CreateSub(LHS: A, RHS: B, Name: "usub.sub" ); |
| 273 | Value *Zero = ConstantInt::get(Ty, V: 0); |
| 274 | return Builder.CreateSelect(C: Cmp, True: Zero, False: Sub, Name: "usub.sat" ); |
| 275 | } |
| 276 | |
| 277 | // Compute the high N bits of the 2N-bit unsigned product of two N-bit values |
| 278 | // using only N-bit arithmetic, so we don't introduce a wider integer type that |
| 279 | // may be unsupported in DXIL. |
| 280 | static Value *createMulHighUnsigned(IRBuilder<> &Builder, Value *A, Value *B, |
| 281 | Type *Ty, unsigned BW) { |
| 282 | assert(BW % 2 == 0 && "high-half split needs symmetric halves" ); |
| 283 | unsigned Half = BW / 2; |
| 284 | Value *HalfShift = ConstantInt::get(Ty, V: Half); |
| 285 | Value *LoMask = ConstantInt::get(Ty, V: APInt::getLowBitsSet(numBits: BW, loBitsSet: Half)); |
| 286 | |
| 287 | Value *U0 = Builder.CreateAnd(LHS: A, RHS: LoMask); |
| 288 | Value *U1 = Builder.CreateLShr(LHS: A, RHS: HalfShift); |
| 289 | Value *V0 = Builder.CreateAnd(LHS: B, RHS: LoMask); |
| 290 | Value *V1 = Builder.CreateLShr(LHS: B, RHS: HalfShift); |
| 291 | |
| 292 | Value *W0 = Builder.CreateMul(LHS: U0, RHS: V0); |
| 293 | Value *T = Builder.CreateAdd(LHS: Builder.CreateMul(LHS: U1, RHS: V0), |
| 294 | RHS: Builder.CreateLShr(LHS: W0, RHS: HalfShift)); |
| 295 | Value *W1 = Builder.CreateAnd(LHS: T, RHS: LoMask); |
| 296 | Value *W2 = Builder.CreateLShr(LHS: T, RHS: HalfShift); |
| 297 | W1 = Builder.CreateAdd(LHS: Builder.CreateMul(LHS: U0, RHS: V1), RHS: W1); |
| 298 | return Builder.CreateAdd(LHS: Builder.CreateAdd(LHS: Builder.CreateMul(LHS: U1, RHS: V1), RHS: W2), |
| 299 | RHS: Builder.CreateLShr(LHS: W1, RHS: HalfShift)); |
| 300 | } |
| 301 | |
| 302 | // Expand a {u,s}mul.with.overflow intrinsic. The low half of the result is a |
| 303 | // plain multiply; overflow is derived from the high half of the double-width |
| 304 | // product. |
| 305 | static Value *expandMulWithOverflow(CallInst *Orig, bool Signed) { |
| 306 | IRBuilder<> Builder(Orig); |
| 307 | Value *A = Orig->getArgOperand(i: 0); |
| 308 | Value *B = Orig->getArgOperand(i: 1); |
| 309 | Type *Ty = A->getType(); |
| 310 | unsigned BW = Ty->getScalarSizeInBits(); |
| 311 | |
| 312 | Value *Lo; |
| 313 | Value *Ov; |
| 314 | |
| 315 | // A plain double-width multiply is simplest, but we avoid it once it would |
| 316 | // introduce a 64-bit (or wider) integer, which DXIL does not always support. |
| 317 | // For i32 we use the native DXIL IMul/UMul ops, which return the full product |
| 318 | // as two i32s; wider types fall back to a same-width high-half computation. |
| 319 | if (2 * BW <= 32) { |
| 320 | Lo = Builder.CreateMul(LHS: A, RHS: B); |
| 321 | Type *WideTy = Ty->getWithNewBitWidth(NewBitWidth: 2 * BW); |
| 322 | Value *WideA = |
| 323 | Signed ? Builder.CreateSExt(V: A, DestTy: WideTy) : Builder.CreateZExt(V: A, DestTy: WideTy); |
| 324 | Value *WideB = |
| 325 | Signed ? Builder.CreateSExt(V: B, DestTy: WideTy) : Builder.CreateZExt(V: B, DestTy: WideTy); |
| 326 | Value *Wide = Builder.CreateMul(LHS: WideA, RHS: WideB); |
| 327 | if (Signed) { |
| 328 | // Overflow when the full product doesn't fit back into BW signed bits. |
| 329 | Ov = Builder.CreateICmpNE(LHS: Wide, RHS: Builder.CreateSExt(V: Lo, DestTy: WideTy)); |
| 330 | } else { |
| 331 | Value *Hi = Builder.CreateLShr(LHS: Wide, RHS: ConstantInt::get(Ty: WideTy, V: BW)); |
| 332 | Ov = Builder.CreateICmpNE(LHS: Hi, RHS: ConstantInt::get(Ty: WideTy, V: 0)); |
| 333 | } |
| 334 | } else if (BW == 32) { |
| 335 | // IMul/UMul return {high, low}; index 0 is the high 32 bits. |
| 336 | Type *ResTy = StructType::get(elt1: Ty, elts: Ty); |
| 337 | Intrinsic::ID IntrinsicID = |
| 338 | Signed ? Intrinsic::dx_imul : Intrinsic::dx_umul; |
| 339 | Value *Mul = Builder.CreateIntrinsic(RetTy: ResTy, ID: IntrinsicID, Args: {A, B}); |
| 340 | Value *Hi = Builder.CreateExtractValue(Agg: Mul, Idxs: 0); |
| 341 | Lo = Builder.CreateExtractValue(Agg: Mul, Idxs: 1); |
| 342 | if (Signed) |
| 343 | Ov = Builder.CreateICmpNE( |
| 344 | LHS: Hi, RHS: Builder.CreateAShr(LHS: Lo, RHS: ConstantInt::get(Ty, V: BW - 1))); |
| 345 | else |
| 346 | Ov = Builder.CreateICmpNE(LHS: Hi, RHS: ConstantInt::get(Ty, V: 0)); |
| 347 | } else { |
| 348 | Lo = Builder.CreateMul(LHS: A, RHS: B); |
| 349 | Value *Hi = createMulHighUnsigned(Builder, A, B, Ty, BW); |
| 350 | if (Signed) { |
| 351 | // Turn the unsigned high half into the signed one, then overflow means it |
| 352 | // isn't the sign extension of the low half. |
| 353 | Value *SignShift = ConstantInt::get(Ty, V: BW - 1); |
| 354 | Value *ASign = Builder.CreateAShr(LHS: A, RHS: SignShift); |
| 355 | Value *BSign = Builder.CreateAShr(LHS: B, RHS: SignShift); |
| 356 | Hi = Builder.CreateSub(LHS: Hi, RHS: Builder.CreateAnd(LHS: ASign, RHS: B)); |
| 357 | Hi = Builder.CreateSub(LHS: Hi, RHS: Builder.CreateAnd(LHS: BSign, RHS: A)); |
| 358 | Ov = Builder.CreateICmpNE(LHS: Hi, RHS: Builder.CreateAShr(LHS: Lo, RHS: SignShift)); |
| 359 | } else { |
| 360 | Ov = Builder.CreateICmpNE(LHS: Hi, RHS: ConstantInt::get(Ty, V: 0)); |
| 361 | } |
| 362 | } |
| 363 | |
| 364 | Value *Agg = PoisonValue::get(T: Orig->getType()); |
| 365 | Agg = Builder.CreateInsertValue(Agg, Val: Lo, Idxs: 0); |
| 366 | return Builder.CreateInsertValue(Agg, Val: Ov, Idxs: 1); |
| 367 | } |
| 368 | |
| 369 | static Value *expandVecReduceAdd(CallInst *Orig, Intrinsic::ID IntrinsicId) { |
| 370 | assert(IntrinsicId == Intrinsic::vector_reduce_add || |
| 371 | IntrinsicId == Intrinsic::vector_reduce_fadd); |
| 372 | |
| 373 | IRBuilder<> Builder(Orig); |
| 374 | bool IsFAdd = (IntrinsicId == Intrinsic::vector_reduce_fadd); |
| 375 | |
| 376 | Value *X = Orig->getOperand(i_nocapture: IsFAdd ? 1 : 0); |
| 377 | Type *Ty = X->getType(); |
| 378 | auto *XVec = dyn_cast<FixedVectorType>(Val: Ty); |
| 379 | unsigned XVecSize = XVec->getNumElements(); |
| 380 | Value *Sum = Builder.CreateExtractElement(Vec: X, Idx: static_cast<uint64_t>(0)); |
| 381 | |
| 382 | // Handle the initial start value for floating-point addition. |
| 383 | if (IsFAdd) { |
| 384 | Constant *StartValue = dyn_cast<Constant>(Val: Orig->getOperand(i_nocapture: 0)); |
| 385 | if (StartValue && !StartValue->isNullValue()) |
| 386 | Sum = Builder.CreateFAdd(L: Sum, R: StartValue); |
| 387 | } |
| 388 | |
| 389 | // Accumulate the remaining vector elements. |
| 390 | for (unsigned I = 1; I < XVecSize; I++) { |
| 391 | Value *Elt = Builder.CreateExtractElement(Vec: X, Idx: I); |
| 392 | if (IsFAdd) |
| 393 | Sum = Builder.CreateFAdd(L: Sum, R: Elt); |
| 394 | else |
| 395 | Sum = Builder.CreateAdd(LHS: Sum, RHS: Elt); |
| 396 | } |
| 397 | |
| 398 | return Sum; |
| 399 | } |
| 400 | |
| 401 | static Value *expandAbs(CallInst *Orig) { |
| 402 | Value *X = Orig->getOperand(i_nocapture: 0); |
| 403 | IRBuilder<> Builder(Orig); |
| 404 | Type *Ty = X->getType(); |
| 405 | Type *EltTy = Ty->getScalarType(); |
| 406 | Constant *Zero = Ty->isVectorTy() |
| 407 | ? ConstantVector::getSplat( |
| 408 | EC: ElementCount::getFixed( |
| 409 | MinVal: cast<FixedVectorType>(Val: Ty)->getNumElements()), |
| 410 | Elt: ConstantInt::get(Ty: EltTy, V: 0)) |
| 411 | : ConstantInt::get(Ty: EltTy, V: 0); |
| 412 | auto *V = Builder.CreateSub(LHS: Zero, RHS: X); |
| 413 | return Builder.CreateIntrinsic(RetTy: Ty, ID: Intrinsic::smax, Args: {X, V}, FMFSource: nullptr, |
| 414 | Name: "dx.max" ); |
| 415 | } |
| 416 | |
| 417 | // Create a DXIL dot2, dot3, or dot4 for the given operands. |
| 418 | static Value *expandFloatDotChunk(CallInst *Orig, Value *A, Value *B) { |
| 419 | Type *ATy = A->getType(); |
| 420 | [[maybe_unused]] Type *BTy = B->getType(); |
| 421 | assert(ATy->isVectorTy() && BTy->isVectorTy()); |
| 422 | |
| 423 | IRBuilder<> Builder(Orig); |
| 424 | |
| 425 | auto *AVec = dyn_cast<FixedVectorType>(Val: ATy); |
| 426 | |
| 427 | assert(ATy->getScalarType()->isFloatingPointTy()); |
| 428 | |
| 429 | unsigned NumElts = AVec->getNumElements(); |
| 430 | Intrinsic::ID DotIntrinsic; |
| 431 | switch (NumElts) { |
| 432 | case 2: |
| 433 | DotIntrinsic = Intrinsic::dx_dot2; |
| 434 | break; |
| 435 | case 3: |
| 436 | DotIntrinsic = Intrinsic::dx_dot3; |
| 437 | break; |
| 438 | case 4: |
| 439 | DotIntrinsic = Intrinsic::dx_dot4; |
| 440 | break; |
| 441 | default: |
| 442 | reportFatalUsageError( |
| 443 | reason: "Invalid dot product input vector: length is outside 2-4" ); |
| 444 | } |
| 445 | |
| 446 | SmallVector<Value *> Args; |
| 447 | for (unsigned I = 0; I < NumElts; ++I) |
| 448 | Args.push_back(Elt: Builder.CreateExtractElement(Vec: A, Idx: Builder.getInt32(C: I))); |
| 449 | for (unsigned I = 0; I < NumElts; ++I) |
| 450 | Args.push_back(Elt: Builder.CreateExtractElement(Vec: B, Idx: Builder.getInt32(C: I))); |
| 451 | return Builder.CreateIntrinsic(RetTy: ATy->getScalarType(), ID: DotIntrinsic, Args, |
| 452 | FMFSource: nullptr, Name: "dot" ); |
| 453 | } |
| 454 | |
| 455 | // Expand an arbitrary-width float dot into the minimum number of legal DXIL |
| 456 | // dot2, dot3, and dot4 operations. |
| 457 | static Value *expandFloatDotIntrinsic(CallInst *Orig) { |
| 458 | Value *A = Orig->getOperand(i_nocapture: 0); |
| 459 | Value *B = Orig->getOperand(i_nocapture: 1); |
| 460 | unsigned NumElts = cast<FixedVectorType>(Val: A->getType())->getNumElements(); |
| 461 | |
| 462 | // We return early here to avoid constructing unnecessary identity shuffles. |
| 463 | if (NumElts <= 4) |
| 464 | return expandFloatDotChunk(Orig, A, B); |
| 465 | |
| 466 | assert(Orig->getModule()->getTargetTriple().getOSVersion() < |
| 467 | VersionTuple(6, 9) && |
| 468 | "long fdot must not be expanded for shader model 6.9 or later" ); |
| 469 | |
| 470 | IRBuilder<> Builder(Orig); |
| 471 | Value *Result = nullptr; |
| 472 | for (unsigned Offset = 0; Offset < NumElts;) { |
| 473 | unsigned Remaining = NumElts - Offset; |
| 474 | // Taking four is optimal unless it would leave an illegal one-element |
| 475 | // tail. In that case, take three and finish with dot2. |
| 476 | unsigned ChunkSize = Remaining == 5 ? 3 : std::min(a: Remaining, b: 4u); |
| 477 | SmallVector<int, 4> Mask; |
| 478 | for (unsigned I = 0; I < ChunkSize; ++I) |
| 479 | Mask.push_back(Elt: Offset + I); |
| 480 | Value *AChunk = Builder.CreateShuffleVector(V: A, Mask); |
| 481 | Value *BChunk = Builder.CreateShuffleVector(V: B, Mask); |
| 482 | Value *Chunk = expandFloatDotChunk(Orig, A: AChunk, B: BChunk); |
| 483 | Result = Result ? Builder.CreateFAdd(L: Result, R: Chunk, Name: "dot.add" ) : Chunk; |
| 484 | Offset += ChunkSize; |
| 485 | } |
| 486 | return Result; |
| 487 | } |
| 488 | |
| 489 | // Expand integer dot product to multiply and add ops |
| 490 | static Value *expandIntegerDotIntrinsic(CallInst *Orig, |
| 491 | Intrinsic::ID DotIntrinsic) { |
| 492 | assert(DotIntrinsic == Intrinsic::dx_sdot || |
| 493 | DotIntrinsic == Intrinsic::dx_udot); |
| 494 | Value *A = Orig->getOperand(i_nocapture: 0); |
| 495 | Value *B = Orig->getOperand(i_nocapture: 1); |
| 496 | Type *ATy = A->getType(); |
| 497 | [[maybe_unused]] Type *BTy = B->getType(); |
| 498 | assert(ATy->isVectorTy() && BTy->isVectorTy()); |
| 499 | |
| 500 | IRBuilder<> Builder(Orig); |
| 501 | |
| 502 | auto *AVec = dyn_cast<FixedVectorType>(Val: ATy); |
| 503 | |
| 504 | assert(ATy->getScalarType()->isIntegerTy()); |
| 505 | |
| 506 | Value *Result; |
| 507 | Intrinsic::ID MadIntrinsic = DotIntrinsic == Intrinsic::dx_sdot |
| 508 | ? Intrinsic::dx_imad |
| 509 | : Intrinsic::dx_umad; |
| 510 | Value *Elt0 = Builder.CreateExtractElement(Vec: A, Idx: (uint64_t)0); |
| 511 | Value *Elt1 = Builder.CreateExtractElement(Vec: B, Idx: (uint64_t)0); |
| 512 | Result = Builder.CreateMul(LHS: Elt0, RHS: Elt1); |
| 513 | for (unsigned I = 1; I < AVec->getNumElements(); I++) { |
| 514 | Elt0 = Builder.CreateExtractElement(Vec: A, Idx: I); |
| 515 | Elt1 = Builder.CreateExtractElement(Vec: B, Idx: I); |
| 516 | Result = Builder.CreateIntrinsic(RetTy: Result->getType(), ID: MadIntrinsic, |
| 517 | Args: ArrayRef<Value *>{Elt0, Elt1, Result}, |
| 518 | FMFSource: nullptr, Name: "dx.mad" ); |
| 519 | } |
| 520 | return Result; |
| 521 | } |
| 522 | |
| 523 | static Value *expandExpIntrinsic(CallInst *Orig) { |
| 524 | Value *X = Orig->getOperand(i_nocapture: 0); |
| 525 | IRBuilder<> Builder(Orig); |
| 526 | Type *Ty = X->getType(); |
| 527 | Type *EltTy = Ty->getScalarType(); |
| 528 | Constant *Log2eConst = |
| 529 | Ty->isVectorTy() ? ConstantVector::getSplat( |
| 530 | EC: ElementCount::getFixed( |
| 531 | MinVal: cast<FixedVectorType>(Val: Ty)->getNumElements()), |
| 532 | Elt: ConstantFP::get(Ty: EltTy, V: numbers::log2ef)) |
| 533 | : ConstantFP::get(Ty: EltTy, V: numbers::log2ef); |
| 534 | Value *NewX = Builder.CreateFMul(L: Log2eConst, R: X); |
| 535 | CallInst *Exp2Call = Builder.CreateIntrinsicWithoutFolding( |
| 536 | RetTy: Ty, ID: Intrinsic::exp2, Args: {NewX}, FMFSource: nullptr, Name: "dx.exp2" ); |
| 537 | Exp2Call->setTailCall(Orig->isTailCall()); |
| 538 | Exp2Call->setAttributes(Orig->getAttributes()); |
| 539 | return Exp2Call; |
| 540 | } |
| 541 | |
| 542 | static Value *expandIsFPClass(CallInst *Orig) { |
| 543 | Value *T = Orig->getArgOperand(i: 1); |
| 544 | auto *TCI = dyn_cast<ConstantInt>(Val: T); |
| 545 | |
| 546 | // These FPClassTest cases have DXIL opcodes, so they will be handled in |
| 547 | // DXIL Op Lowering instead for all non f16 cases. |
| 548 | switch (TCI->getZExtValue()) { |
| 549 | case FPClassTest::fcInf: |
| 550 | return expand16BitIsInf(Orig); |
| 551 | case FPClassTest::fcNan: |
| 552 | return expand16BitIsNaN(Orig); |
| 553 | case FPClassTest::fcNormal: |
| 554 | return expand16BitIsNormal(Orig); |
| 555 | case FPClassTest::fcFinite: |
| 556 | return expand16BitIsFinite(Orig); |
| 557 | } |
| 558 | |
| 559 | IRBuilder<> Builder(Orig); |
| 560 | |
| 561 | Value *F = Orig->getArgOperand(i: 0); |
| 562 | Type *FTy = F->getType(); |
| 563 | unsigned FNumElem = 0; // 0 => F is not a vector |
| 564 | |
| 565 | unsigned BitWidth; // Bit width of F or the ElemTy of F |
| 566 | Type *BitCastTy; // An IntNTy of the same bitwidth as F or ElemTy of F |
| 567 | |
| 568 | if (auto *FVecTy = dyn_cast<FixedVectorType>(Val: FTy)) { |
| 569 | Type *ElemTy = FVecTy->getElementType(); |
| 570 | FNumElem = FVecTy->getNumElements(); |
| 571 | BitWidth = ElemTy->getPrimitiveSizeInBits(); |
| 572 | BitCastTy = FixedVectorType::get(ElementType: Builder.getIntNTy(N: BitWidth), NumElts: FNumElem); |
| 573 | } else { |
| 574 | BitWidth = FTy->getPrimitiveSizeInBits(); |
| 575 | BitCastTy = Builder.getIntNTy(N: BitWidth); |
| 576 | } |
| 577 | |
| 578 | Value *FBitCast = Builder.CreateBitCast(V: F, DestTy: BitCastTy); |
| 579 | switch (TCI->getZExtValue()) { |
| 580 | case FPClassTest::fcNegZero: { |
| 581 | Value *NegZero = |
| 582 | ConstantInt::get(Ty: Builder.getIntNTy(N: BitWidth), V: 1 << (BitWidth - 1), |
| 583 | /*IsSigned=*/true); |
| 584 | Value *RetVal; |
| 585 | if (FNumElem) { |
| 586 | Value *NegZeroSplat = Builder.CreateVectorSplat(NumElts: FNumElem, V: NegZero); |
| 587 | RetVal = |
| 588 | Builder.CreateICmpEQ(LHS: FBitCast, RHS: NegZeroSplat, Name: "is.fpclass.negzero" ); |
| 589 | } else |
| 590 | RetVal = Builder.CreateICmpEQ(LHS: FBitCast, RHS: NegZero, Name: "is.fpclass.negzero" ); |
| 591 | return RetVal; |
| 592 | } |
| 593 | default: |
| 594 | reportFatalUsageError(reason: "Unsupported FPClassTest" ); |
| 595 | } |
| 596 | } |
| 597 | |
| 598 | static Value *expandAnyOrAllIntrinsic(CallInst *Orig, |
| 599 | Intrinsic::ID IntrinsicId) { |
| 600 | Value *X = Orig->getOperand(i_nocapture: 0); |
| 601 | IRBuilder<> Builder(Orig); |
| 602 | Type *Ty = X->getType(); |
| 603 | Type *EltTy = Ty->getScalarType(); |
| 604 | |
| 605 | auto ApplyOp = [&Builder](Intrinsic::ID IntrinsicId, Value *Result, |
| 606 | Value *Elt) { |
| 607 | if (IntrinsicId == Intrinsic::dx_any) |
| 608 | return Builder.CreateOr(LHS: Result, RHS: Elt); |
| 609 | assert(IntrinsicId == Intrinsic::dx_all); |
| 610 | return Builder.CreateAnd(LHS: Result, RHS: Elt); |
| 611 | }; |
| 612 | |
| 613 | Value *Result = nullptr; |
| 614 | if (!Ty->isVectorTy()) { |
| 615 | Result = EltTy->isFloatingPointTy() |
| 616 | ? Builder.CreateFCmpUNE(LHS: X, RHS: ConstantFP::get(Ty: EltTy, V: 0)) |
| 617 | : Builder.CreateICmpNE(LHS: X, RHS: ConstantInt::get(Ty: EltTy, V: 0)); |
| 618 | } else { |
| 619 | auto *XVec = dyn_cast<FixedVectorType>(Val: Ty); |
| 620 | Value *Cond = |
| 621 | EltTy->isFloatingPointTy() |
| 622 | ? Builder.CreateFCmpUNE( |
| 623 | LHS: X, RHS: ConstantVector::getSplat( |
| 624 | EC: ElementCount::getFixed(MinVal: XVec->getNumElements()), |
| 625 | Elt: ConstantFP::get(Ty: EltTy, V: 0))) |
| 626 | : Builder.CreateICmpNE( |
| 627 | LHS: X, RHS: ConstantVector::getSplat( |
| 628 | EC: ElementCount::getFixed(MinVal: XVec->getNumElements()), |
| 629 | Elt: ConstantInt::get(Ty: EltTy, V: 0))); |
| 630 | Result = Builder.CreateExtractElement(Vec: Cond, Idx: (uint64_t)0); |
| 631 | for (unsigned I = 1; I < XVec->getNumElements(); I++) { |
| 632 | Value *Elt = Builder.CreateExtractElement(Vec: Cond, Idx: I); |
| 633 | Result = ApplyOp(IntrinsicId, Result, Elt); |
| 634 | } |
| 635 | } |
| 636 | return Result; |
| 637 | } |
| 638 | |
| 639 | static Value *expandLogIntrinsic(CallInst *Orig, |
| 640 | float LogConstVal = numbers::ln2f) { |
| 641 | Value *X = Orig->getOperand(i_nocapture: 0); |
| 642 | IRBuilder<> Builder(Orig); |
| 643 | Type *Ty = X->getType(); |
| 644 | Type *EltTy = Ty->getScalarType(); |
| 645 | Constant *Ln2Const = |
| 646 | Ty->isVectorTy() ? ConstantVector::getSplat( |
| 647 | EC: ElementCount::getFixed( |
| 648 | MinVal: cast<FixedVectorType>(Val: Ty)->getNumElements()), |
| 649 | Elt: ConstantFP::get(Ty: EltTy, V: LogConstVal)) |
| 650 | : ConstantFP::get(Ty: EltTy, V: LogConstVal); |
| 651 | CallInst *Log2Call = Builder.CreateIntrinsicWithoutFolding( |
| 652 | RetTy: Ty, ID: Intrinsic::log2, Args: {X}, FMFSource: nullptr, Name: "elt.log2" ); |
| 653 | Log2Call->setTailCall(Orig->isTailCall()); |
| 654 | Log2Call->setAttributes(Orig->getAttributes()); |
| 655 | return Builder.CreateFMul(L: Ln2Const, R: Log2Call); |
| 656 | } |
| 657 | static Value *expandLog10Intrinsic(CallInst *Orig) { |
| 658 | return expandLogIntrinsic(Orig, LogConstVal: numbers::ln2f / numbers::ln10f); |
| 659 | } |
| 660 | |
| 661 | static Value *expandAtan2Intrinsic(CallInst *Orig) { |
| 662 | Value *Y = Orig->getOperand(i_nocapture: 0); |
| 663 | Value *X = Orig->getOperand(i_nocapture: 1); |
| 664 | Type *Ty = X->getType(); |
| 665 | IRBuilder<> Builder(Orig); |
| 666 | Builder.setFastMathFlags(Orig->getFastMathFlags()); |
| 667 | |
| 668 | Value *Tan = Builder.CreateFDiv(L: Y, R: X); |
| 669 | |
| 670 | CallInst *Atan = Builder.CreateIntrinsicWithoutFolding( |
| 671 | RetTy: Ty, ID: Intrinsic::atan, Args: {Tan}, FMFSource: nullptr, Name: "Elt.Atan" ); |
| 672 | Atan->setTailCall(Orig->isTailCall()); |
| 673 | Atan->setAttributes(Orig->getAttributes()); |
| 674 | |
| 675 | // Modify atan result based on https://en.wikipedia.org/wiki/Atan2. |
| 676 | Constant *Pi = ConstantFP::get(Ty, V: llvm::numbers::pi); |
| 677 | Constant *HalfPi = ConstantFP::get(Ty, V: llvm::numbers::pi / 2); |
| 678 | Constant *NegHalfPi = ConstantFP::get(Ty, V: -llvm::numbers::pi / 2); |
| 679 | Constant *Zero = ConstantFP::get(Ty, V: 0); |
| 680 | Value *AtanAddPi = Builder.CreateFAdd(L: Atan, R: Pi); |
| 681 | Value *AtanSubPi = Builder.CreateFSub(L: Atan, R: Pi); |
| 682 | |
| 683 | // x > 0 -> atan. |
| 684 | Value *Result = Atan; |
| 685 | Value *XLt0 = Builder.CreateFCmpOLT(LHS: X, RHS: Zero); |
| 686 | Value *XEq0 = Builder.CreateFCmpOEQ(LHS: X, RHS: Zero); |
| 687 | Value *YGe0 = Builder.CreateFCmpOGE(LHS: Y, RHS: Zero); |
| 688 | Value *YLt0 = Builder.CreateFCmpOLT(LHS: Y, RHS: Zero); |
| 689 | |
| 690 | // x < 0, y >= 0 -> atan + pi. |
| 691 | Value *XLt0AndYGe0 = Builder.CreateAnd(LHS: XLt0, RHS: YGe0); |
| 692 | Result = Builder.CreateSelect(C: XLt0AndYGe0, True: AtanAddPi, False: Result); |
| 693 | |
| 694 | // x < 0, y < 0 -> atan - pi. |
| 695 | Value *XLt0AndYLt0 = Builder.CreateAnd(LHS: XLt0, RHS: YLt0); |
| 696 | Result = Builder.CreateSelect(C: XLt0AndYLt0, True: AtanSubPi, False: Result); |
| 697 | |
| 698 | // x == 0, y < 0 -> -pi/2 |
| 699 | Value *XEq0AndYLt0 = Builder.CreateAnd(LHS: XEq0, RHS: YLt0); |
| 700 | Result = Builder.CreateSelect(C: XEq0AndYLt0, True: NegHalfPi, False: Result); |
| 701 | |
| 702 | // x == 0, y > 0 -> pi/2 |
| 703 | Value *XEq0AndYGe0 = Builder.CreateAnd(LHS: XEq0, RHS: YGe0); |
| 704 | Result = Builder.CreateSelect(C: XEq0AndYGe0, True: HalfPi, False: Result); |
| 705 | |
| 706 | return Result; |
| 707 | } |
| 708 | |
| 709 | template <bool LeftFunnel> |
| 710 | static Value *expandFunnelShiftIntrinsic(CallInst *Orig) { |
| 711 | Type *Ty = Orig->getType(); |
| 712 | Value *A = Orig->getOperand(i_nocapture: 0); |
| 713 | Value *B = Orig->getOperand(i_nocapture: 1); |
| 714 | Value *Shift = Orig->getOperand(i_nocapture: 2); |
| 715 | |
| 716 | IRBuilder<> Builder(Orig); |
| 717 | |
| 718 | assert(llvm::isPowerOf2_32(Ty->getScalarSizeInBits()) && |
| 719 | "Can't use Mask to compute modulo and inverse" ); |
| 720 | |
| 721 | // Note: if (Shift % BitWidth) == 0 then (BitWidth - Shift) == BitWidth, |
| 722 | // shifting by the bitwidth for shl/lshr returns a poisoned result. As such, |
| 723 | // we implement the same formula as LegalizerHelper::lowerFunnelShiftAsShifts. |
| 724 | // |
| 725 | // The funnel shift is expanded like so: |
| 726 | // fshl |
| 727 | // -> msb_extract((concat(A, B) << (Shift % BitWidth)), BitWidth) |
| 728 | // -> A << (Shift % BitWidth) | B >> 1 >> (BitWidth - 1 - (Shift % BitWidth)) |
| 729 | // fshr |
| 730 | // -> lsb_extract((concat(A, B) >> (Shift % BitWidth), BitWidth)) |
| 731 | // -> A << 1 << (BitWidth - 1 - (Shift % BitWidth)) | B >> (Shift % BitWidth) |
| 732 | |
| 733 | // (BitWidth - 1) -> Mask |
| 734 | Constant *Mask = ConstantInt::get(Ty, V: Ty->getScalarSizeInBits() - 1); |
| 735 | |
| 736 | // Shift % BitWidth |
| 737 | // -> Shift & (BitWidth - 1) |
| 738 | // -> Shift & Mask |
| 739 | Value *MaskedShift = Builder.CreateAnd(LHS: Shift, RHS: Mask); |
| 740 | |
| 741 | // (BitWidth - 1) - (Shift % BitWidth) |
| 742 | // -> ~Shift & (BitWidth - 1) |
| 743 | // -> ~Shift & Mask |
| 744 | Value *NotShift = Builder.CreateNot(V: Shift); |
| 745 | Value *InverseShift = Builder.CreateAnd(LHS: NotShift, RHS: Mask); |
| 746 | |
| 747 | Constant *One = ConstantInt::get(Ty, V: 1); |
| 748 | Value *ShiftedA; |
| 749 | Value *ShiftedB; |
| 750 | |
| 751 | if (LeftFunnel) { |
| 752 | ShiftedA = Builder.CreateShl(LHS: A, RHS: MaskedShift); |
| 753 | Value *ShiftB1 = Builder.CreateLShr(LHS: B, RHS: One); |
| 754 | ShiftedB = Builder.CreateLShr(LHS: ShiftB1, RHS: InverseShift); |
| 755 | } else { |
| 756 | Value *ShiftA1 = Builder.CreateShl(LHS: A, RHS: One); |
| 757 | ShiftedA = Builder.CreateShl(LHS: ShiftA1, RHS: InverseShift); |
| 758 | ShiftedB = Builder.CreateLShr(LHS: B, RHS: MaskedShift); |
| 759 | } |
| 760 | |
| 761 | Value *Result = Builder.CreateOr(LHS: ShiftedA, RHS: ShiftedB); |
| 762 | return Result; |
| 763 | } |
| 764 | |
| 765 | static Value *expandPowIntrinsic(CallInst *Orig, Intrinsic::ID IntrinsicId) { |
| 766 | |
| 767 | Value *X = Orig->getOperand(i_nocapture: 0); |
| 768 | Value *Y = Orig->getOperand(i_nocapture: 1); |
| 769 | Type *Ty = X->getType(); |
| 770 | IRBuilder<> Builder(Orig); |
| 771 | |
| 772 | if (IntrinsicId == Intrinsic::powi) |
| 773 | Y = Builder.CreateSIToFP(V: Y, DestTy: Ty); |
| 774 | |
| 775 | Value *Log2Call = |
| 776 | Builder.CreateIntrinsic(RetTy: Ty, ID: Intrinsic::log2, Args: {X}, FMFSource: nullptr, Name: "elt.log2" ); |
| 777 | auto *Mul = Builder.CreateFMul(L: Log2Call, R: Y); |
| 778 | CallInst *Exp2Call = Builder.CreateIntrinsicWithoutFolding( |
| 779 | RetTy: Ty, ID: Intrinsic::exp2, Args: {Mul}, FMFSource: nullptr, Name: "elt.exp2" ); |
| 780 | Exp2Call->setTailCall(Orig->isTailCall()); |
| 781 | Exp2Call->setAttributes(Orig->getAttributes()); |
| 782 | return Exp2Call; |
| 783 | } |
| 784 | |
| 785 | static bool expandBufferLoadIntrinsic(CallInst *Orig, bool IsRaw) { |
| 786 | IRBuilder<> Builder(Orig); |
| 787 | |
| 788 | Type *BufferTy = Orig->getType()->getStructElementType(N: 0); |
| 789 | Type *ScalarTy = BufferTy->getScalarType(); |
| 790 | bool IsDouble = ScalarTy->isDoubleTy(); |
| 791 | assert(IsDouble || ScalarTy->isIntegerTy(64) && |
| 792 | "Only expand double or int64 scalars or vectors" ); |
| 793 | bool IsVector = false; |
| 794 | unsigned = 2; |
| 795 | if (auto *VT = dyn_cast<FixedVectorType>(Val: BufferTy)) { |
| 796 | ExtractNum = 2 * VT->getNumElements(); |
| 797 | IsVector = true; |
| 798 | assert(IsRaw || ExtractNum == 4 && "TypedBufferLoad vector must be size 2" ); |
| 799 | } |
| 800 | |
| 801 | SmallVector<Value *, 2> Loads; |
| 802 | Value *Result = PoisonValue::get(T: BufferTy); |
| 803 | unsigned Base = 0; |
| 804 | // If we need to extract more than 4 i32; we need to break it up into |
| 805 | // more than one load. LoadNum tells us how many i32s we are loading in |
| 806 | // each load |
| 807 | while (ExtractNum > 0) { |
| 808 | unsigned LoadNum = std::min(a: ExtractNum, b: 4u); |
| 809 | Type *Ty = VectorType::get(ElementType: Builder.getInt32Ty(), NumElements: LoadNum, Scalable: false); |
| 810 | |
| 811 | Type *LoadType = StructType::get(elt1: Ty, elts: Builder.getInt1Ty()); |
| 812 | Intrinsic::ID LoadIntrinsic = Intrinsic::dx_resource_load_typedbuffer; |
| 813 | SmallVector<Value *, 3> Args = {Orig->getOperand(i_nocapture: 0), Orig->getOperand(i_nocapture: 1)}; |
| 814 | if (IsRaw) { |
| 815 | LoadIntrinsic = Intrinsic::dx_resource_load_rawbuffer; |
| 816 | Value *Tmp = Builder.getInt32(C: 4 * Base * 2); |
| 817 | Value *Offset = Orig->getOperand(i_nocapture: 2); |
| 818 | Args.push_back(Elt: Offset); |
| 819 | unsigned AddressArg = isa<PoisonValue>(Val: Offset) ? 1 : 2; |
| 820 | if (Base != 0) |
| 821 | Args[AddressArg] = Builder.CreateAdd(LHS: Args[AddressArg], RHS: Tmp); |
| 822 | } |
| 823 | |
| 824 | Value *Load = Builder.CreateIntrinsic(RetTy: LoadType, ID: LoadIntrinsic, Args); |
| 825 | Loads.push_back(Elt: Load); |
| 826 | |
| 827 | // extract the buffer load's result |
| 828 | Value * = Builder.CreateExtractValue(Agg: Load, Idxs: {0}); |
| 829 | |
| 830 | SmallVector<Value *> ; |
| 831 | for (unsigned I = 0; I < LoadNum; ++I) |
| 832 | ExtractElements.push_back( |
| 833 | Elt: Builder.CreateExtractElement(Vec: Extract, Idx: Builder.getInt32(C: I))); |
| 834 | |
| 835 | // combine into double(s) or int64(s) |
| 836 | for (unsigned I = 0; I < LoadNum; I += 2) { |
| 837 | Value *Combined = nullptr; |
| 838 | if (IsDouble) |
| 839 | // For doubles, use dx_asdouble intrinsic |
| 840 | Combined = Builder.CreateIntrinsic( |
| 841 | RetTy: Builder.getDoubleTy(), ID: Intrinsic::dx_asdouble, |
| 842 | Args: {ExtractElements[I], ExtractElements[I + 1]}); |
| 843 | else { |
| 844 | // For int64, manually combine two int32s |
| 845 | // First, zero-extend both values to i64 |
| 846 | Value *Lo = |
| 847 | Builder.CreateZExt(V: ExtractElements[I], DestTy: Builder.getInt64Ty()); |
| 848 | Value *Hi = |
| 849 | Builder.CreateZExt(V: ExtractElements[I + 1], DestTy: Builder.getInt64Ty()); |
| 850 | // Shift the high bits left by 32 bits |
| 851 | Value *ShiftedHi = Builder.CreateShl(LHS: Hi, RHS: Builder.getInt64(C: 32)); |
| 852 | // OR the high and low bits together |
| 853 | Combined = Builder.CreateOr(LHS: Lo, RHS: ShiftedHi); |
| 854 | } |
| 855 | |
| 856 | if (IsVector) |
| 857 | Result = Builder.CreateInsertElement(Vec: Result, NewElt: Combined, |
| 858 | Idx: Builder.getInt32(C: (I / 2) + Base)); |
| 859 | else |
| 860 | Result = Combined; |
| 861 | } |
| 862 | |
| 863 | ExtractNum -= LoadNum; |
| 864 | Base += LoadNum / 2; |
| 865 | } |
| 866 | |
| 867 | Value *CheckBit = nullptr; |
| 868 | for (User *U : make_early_inc_range(Range: Orig->users())) { |
| 869 | // If it's not a ExtractValueInst, we don't know how to |
| 870 | // handle it |
| 871 | auto *EVI = dyn_cast<ExtractValueInst>(Val: U); |
| 872 | if (!EVI) |
| 873 | llvm_unreachable("Unexpected user of typedbufferload" ); |
| 874 | |
| 875 | ArrayRef<unsigned> Indices = EVI->getIndices(); |
| 876 | assert(Indices.size() == 1); |
| 877 | |
| 878 | if (Indices[0] == 0) { |
| 879 | // Use of the value(s) |
| 880 | EVI->replaceAllUsesWith(V: Result); |
| 881 | } else { |
| 882 | // Use of the check bit |
| 883 | assert(Indices[0] == 1 && "Unexpected type for typedbufferload" ); |
| 884 | // Note: This does not always match the historical behaviour of DXC. |
| 885 | // See https://github.com/microsoft/DirectXShaderCompiler/issues/7622 |
| 886 | if (!CheckBit) { |
| 887 | SmallVector<Value *, 2> CheckBits; |
| 888 | for (Value *L : Loads) |
| 889 | CheckBits.push_back(Elt: Builder.CreateExtractValue(Agg: L, Idxs: {1})); |
| 890 | CheckBit = Builder.CreateAnd(Ops: CheckBits); |
| 891 | } |
| 892 | EVI->replaceAllUsesWith(V: CheckBit); |
| 893 | } |
| 894 | EVI->eraseFromParent(); |
| 895 | } |
| 896 | Orig->eraseFromParent(); |
| 897 | return true; |
| 898 | } |
| 899 | |
| 900 | static bool expandBufferStoreIntrinsic(CallInst *Orig, bool IsRaw) { |
| 901 | IRBuilder<> Builder(Orig); |
| 902 | |
| 903 | unsigned ValIndex = IsRaw ? 3 : 2; |
| 904 | Type *BufferTy = Orig->getFunctionType()->getParamType(i: ValIndex); |
| 905 | Type *ScalarTy = BufferTy->getScalarType(); |
| 906 | bool IsDouble = ScalarTy->isDoubleTy(); |
| 907 | assert((IsDouble || ScalarTy->isIntegerTy(64)) && |
| 908 | "Only expand double or int64 scalars or vectors" ); |
| 909 | |
| 910 | // Determine if we're dealing with a vector or scalar |
| 911 | bool IsVector = false; |
| 912 | unsigned = 2; |
| 913 | unsigned VecLen = 0; |
| 914 | if (auto *VT = dyn_cast<FixedVectorType>(Val: BufferTy)) { |
| 915 | VecLen = VT->getNumElements(); |
| 916 | assert(IsRaw || VecLen == 2 && "TypedBufferStore vector must be size 2" ); |
| 917 | ExtractNum = VecLen * 2; |
| 918 | IsVector = true; |
| 919 | } |
| 920 | |
| 921 | // Create the appropriate vector type for the result |
| 922 | Type *Int32Ty = Builder.getInt32Ty(); |
| 923 | Type *ResultTy = VectorType::get(ElementType: Int32Ty, NumElements: ExtractNum, Scalable: false); |
| 924 | Value *Val = PoisonValue::get(T: ResultTy); |
| 925 | |
| 926 | Type *SplitElementTy = Int32Ty; |
| 927 | if (IsVector) |
| 928 | SplitElementTy = VectorType::get(ElementType: SplitElementTy, NumElements: VecLen, Scalable: false); |
| 929 | |
| 930 | Value *LowBits = nullptr; |
| 931 | Value *HighBits = nullptr; |
| 932 | // Split the 64-bit values into 32-bit components |
| 933 | if (IsDouble) { |
| 934 | auto *SplitTy = llvm::StructType::get(elt1: SplitElementTy, elts: SplitElementTy); |
| 935 | Value *Split = Builder.CreateIntrinsic(RetTy: SplitTy, ID: Intrinsic::dx_splitdouble, |
| 936 | Args: {Orig->getOperand(i_nocapture: ValIndex)}); |
| 937 | LowBits = Builder.CreateExtractValue(Agg: Split, Idxs: 0); |
| 938 | HighBits = Builder.CreateExtractValue(Agg: Split, Idxs: 1); |
| 939 | } else { |
| 940 | // Handle int64 type(s) |
| 941 | Value *InputVal = Orig->getOperand(i_nocapture: ValIndex); |
| 942 | Constant *ShiftAmt = Builder.getInt64(C: 32); |
| 943 | if (IsVector) |
| 944 | ShiftAmt = |
| 945 | ConstantVector::getSplat(EC: ElementCount::getFixed(MinVal: VecLen), Elt: ShiftAmt); |
| 946 | |
| 947 | // Split into low and high 32-bit parts |
| 948 | LowBits = Builder.CreateTrunc(V: InputVal, DestTy: SplitElementTy); |
| 949 | Value *ShiftedVal = Builder.CreateLShr(LHS: InputVal, RHS: ShiftAmt); |
| 950 | HighBits = Builder.CreateTrunc(V: ShiftedVal, DestTy: SplitElementTy); |
| 951 | } |
| 952 | |
| 953 | if (IsVector) { |
| 954 | SmallVector<int, 8> Mask; |
| 955 | for (unsigned I = 0; I < VecLen; ++I) { |
| 956 | Mask.push_back(Elt: I); |
| 957 | Mask.push_back(Elt: I + VecLen); |
| 958 | } |
| 959 | Val = Builder.CreateShuffleVector(V1: LowBits, V2: HighBits, Mask); |
| 960 | } else { |
| 961 | Val = Builder.CreateInsertElement(Vec: Val, NewElt: LowBits, Idx: Builder.getInt32(C: 0)); |
| 962 | Val = Builder.CreateInsertElement(Vec: Val, NewElt: HighBits, Idx: Builder.getInt32(C: 1)); |
| 963 | } |
| 964 | |
| 965 | // If we need to extract more than 4 i32; we need to break it up into |
| 966 | // more than one store. StoreNum tells us how many i32s we are storing in |
| 967 | // each store |
| 968 | unsigned Base = 0; |
| 969 | while (ExtractNum > 0) { |
| 970 | unsigned StoreNum = std::min(a: ExtractNum, b: 4u); |
| 971 | |
| 972 | Intrinsic::ID StoreIntrinsic = Intrinsic::dx_resource_store_typedbuffer; |
| 973 | SmallVector<Value *, 4> Args = {Orig->getOperand(i_nocapture: 0), Orig->getOperand(i_nocapture: 1)}; |
| 974 | if (IsRaw) { |
| 975 | StoreIntrinsic = Intrinsic::dx_resource_store_rawbuffer; |
| 976 | Value *Tmp = Builder.getInt32(C: 4 * Base); |
| 977 | Value *Offset = Orig->getOperand(i_nocapture: 2); |
| 978 | Args.push_back(Elt: Offset); |
| 979 | unsigned AddressArg = isa<PoisonValue>(Val: Offset) ? 1 : 2; |
| 980 | if (Base != 0) |
| 981 | Args[AddressArg] = Builder.CreateAdd(LHS: Args[AddressArg], RHS: Tmp); |
| 982 | } |
| 983 | |
| 984 | SmallVector<int, 4> Mask; |
| 985 | for (unsigned I = 0; I < StoreNum; ++I) { |
| 986 | Mask.push_back(Elt: Base + I); |
| 987 | } |
| 988 | |
| 989 | Value *SubVal = Val; |
| 990 | if (VecLen > 2) |
| 991 | SubVal = Builder.CreateShuffleVector(V: Val, Mask); |
| 992 | |
| 993 | Args.push_back(Elt: SubVal); |
| 994 | // Create the final intrinsic call |
| 995 | Builder.CreateIntrinsic(RetTy: Builder.getVoidTy(), ID: StoreIntrinsic, Args); |
| 996 | |
| 997 | ExtractNum -= StoreNum; |
| 998 | Base += StoreNum; |
| 999 | } |
| 1000 | Orig->eraseFromParent(); |
| 1001 | return true; |
| 1002 | } |
| 1003 | |
| 1004 | static Intrinsic::ID getMaxForClamp(Intrinsic::ID ClampIntrinsic) { |
| 1005 | if (ClampIntrinsic == Intrinsic::dx_uclamp) |
| 1006 | return Intrinsic::umax; |
| 1007 | if (ClampIntrinsic == Intrinsic::dx_sclamp) |
| 1008 | return Intrinsic::smax; |
| 1009 | assert(ClampIntrinsic == Intrinsic::dx_nclamp); |
| 1010 | return Intrinsic::maxnum; |
| 1011 | } |
| 1012 | |
| 1013 | static Intrinsic::ID getMinForClamp(Intrinsic::ID ClampIntrinsic) { |
| 1014 | if (ClampIntrinsic == Intrinsic::dx_uclamp) |
| 1015 | return Intrinsic::umin; |
| 1016 | if (ClampIntrinsic == Intrinsic::dx_sclamp) |
| 1017 | return Intrinsic::smin; |
| 1018 | assert(ClampIntrinsic == Intrinsic::dx_nclamp); |
| 1019 | return Intrinsic::minnum; |
| 1020 | } |
| 1021 | |
| 1022 | static Value *expandClampIntrinsic(CallInst *Orig, |
| 1023 | Intrinsic::ID ClampIntrinsic) { |
| 1024 | Value *X = Orig->getOperand(i_nocapture: 0); |
| 1025 | Value *Min = Orig->getOperand(i_nocapture: 1); |
| 1026 | Value *Max = Orig->getOperand(i_nocapture: 2); |
| 1027 | Type *Ty = X->getType(); |
| 1028 | IRBuilder<> Builder(Orig); |
| 1029 | auto *MaxCall = Builder.CreateIntrinsic(RetTy: Ty, ID: getMaxForClamp(ClampIntrinsic), |
| 1030 | Args: {X, Min}, FMFSource: nullptr, Name: "dx.max" ); |
| 1031 | return Builder.CreateIntrinsic(RetTy: Ty, ID: getMinForClamp(ClampIntrinsic), |
| 1032 | Args: {MaxCall, Max}, FMFSource: nullptr, Name: "dx.min" ); |
| 1033 | } |
| 1034 | |
| 1035 | static Value *expandSignIntrinsic(CallInst *Orig) { |
| 1036 | Value *X = Orig->getOperand(i_nocapture: 0); |
| 1037 | Type *Ty = X->getType(); |
| 1038 | Type *ScalarTy = Ty->getScalarType(); |
| 1039 | Type *RetTy = Orig->getType(); |
| 1040 | Constant *Zero = Constant::getNullValue(Ty); |
| 1041 | |
| 1042 | IRBuilder<> Builder(Orig); |
| 1043 | |
| 1044 | Value *GT; |
| 1045 | Value *LT; |
| 1046 | if (ScalarTy->isFloatingPointTy()) { |
| 1047 | GT = Builder.CreateFCmpOLT(LHS: Zero, RHS: X); |
| 1048 | LT = Builder.CreateFCmpOLT(LHS: X, RHS: Zero); |
| 1049 | } else { |
| 1050 | assert(ScalarTy->isIntegerTy()); |
| 1051 | GT = Builder.CreateICmpSLT(LHS: Zero, RHS: X); |
| 1052 | LT = Builder.CreateICmpSLT(LHS: X, RHS: Zero); |
| 1053 | } |
| 1054 | |
| 1055 | Value *ZextGT = Builder.CreateZExt(V: GT, DestTy: RetTy); |
| 1056 | Value *ZextLT = Builder.CreateZExt(V: LT, DestTy: RetTy); |
| 1057 | |
| 1058 | return Builder.CreateSub(LHS: ZextGT, RHS: ZextLT); |
| 1059 | } |
| 1060 | |
| 1061 | // Expand llvm.copysign by combining the sign bit with the magnitude bits using |
| 1062 | // bitwise operations. |
| 1063 | static Value *expandCopySignIntrinsic(CallInst *Orig) { |
| 1064 | Value *Magnitude = Orig->getOperand(i_nocapture: 0); |
| 1065 | Value *Sign = Orig->getOperand(i_nocapture: 1); |
| 1066 | Type *Ty = Orig->getType(); |
| 1067 | |
| 1068 | IRBuilder<> Builder(Orig); |
| 1069 | |
| 1070 | bool IsDouble = Ty->getScalarType()->isDoubleTy(); |
| 1071 | unsigned BitWidth = IsDouble ? 32 : Ty->getScalarSizeInBits(); |
| 1072 | Type *IntTy = Ty->getWithNewType(EltTy: Builder.getIntNTy(N: BitWidth)); |
| 1073 | |
| 1074 | auto CopySignBit = [&](Value *MagnitudeInt, Value *SignInt) { |
| 1075 | APInt SignMaskVal = APInt::getSignMask(BitWidth); |
| 1076 | // `ConstantInt::get` broadcasts to a splat when `IntTy` is a vector. |
| 1077 | Constant *SignMask = ConstantInt::get(Ty: IntTy, V: SignMaskVal); |
| 1078 | Constant *NotSignMask = ConstantInt::get(Ty: IntTy, V: ~SignMaskVal); |
| 1079 | |
| 1080 | Value *MagnitudeBits = Builder.CreateAnd(LHS: MagnitudeInt, RHS: NotSignMask); |
| 1081 | Value *SignBits = Builder.CreateAnd(LHS: SignInt, RHS: SignMask); |
| 1082 | return Builder.CreateOr(LHS: MagnitudeBits, RHS: SignBits); |
| 1083 | }; |
| 1084 | |
| 1085 | // Avoid i64 bitwise ops, which require the Int64Ops shader feature. |
| 1086 | if (IsDouble) { |
| 1087 | auto *SplitTy = StructType::get(elt1: IntTy, elts: IntTy); |
| 1088 | Value *MagnitudeHalves = Builder.CreateIntrinsic( |
| 1089 | RetTy: SplitTy, ID: Intrinsic::dx_splitdouble, Args: {Magnitude}); |
| 1090 | Value *SignHalves = |
| 1091 | Builder.CreateIntrinsic(RetTy: SplitTy, ID: Intrinsic::dx_splitdouble, Args: {Sign}); |
| 1092 | Value *MagnitudeLow = Builder.CreateExtractValue(Agg: MagnitudeHalves, Idxs: 0); |
| 1093 | Value *MagnitudeHigh = Builder.CreateExtractValue(Agg: MagnitudeHalves, Idxs: 1); |
| 1094 | Value *SignHigh = Builder.CreateExtractValue(Agg: SignHalves, Idxs: 1); |
| 1095 | |
| 1096 | Value *CombinedHigh = CopySignBit(MagnitudeHigh, SignHigh); |
| 1097 | return Builder.CreateIntrinsic(RetTy: Ty, ID: Intrinsic::dx_asdouble, |
| 1098 | Args: {MagnitudeLow, CombinedHigh}); |
| 1099 | } |
| 1100 | |
| 1101 | Value *MagnitudeInt = Builder.CreateBitCast(V: Magnitude, DestTy: IntTy); |
| 1102 | Value *SignInt = Builder.CreateBitCast(V: Sign, DestTy: IntTy); |
| 1103 | Value *CombinedInt = CopySignBit(MagnitudeInt, SignInt); |
| 1104 | return Builder.CreateBitCast(V: CombinedInt, DestTy: Ty); |
| 1105 | } |
| 1106 | |
| 1107 | // Expand llvm.matrix.multiply by extracting row/column vectors and computing |
| 1108 | // dot products. |
| 1109 | // Result[r,c] = dot(row_r(LHS), col_c(RHS)) |
| 1110 | // Element (r,c) is at index c*NumRows + r (column-major). |
| 1111 | static Value *expandMatrixMultiply(CallInst *Orig) { |
| 1112 | Value *LHS = Orig->getArgOperand(i: 0); |
| 1113 | Value *RHS = Orig->getArgOperand(i: 1); |
| 1114 | unsigned LHSRows = cast<ConstantInt>(Val: Orig->getArgOperand(i: 2))->getZExtValue(); |
| 1115 | unsigned LHSCols = cast<ConstantInt>(Val: Orig->getArgOperand(i: 3))->getZExtValue(); |
| 1116 | unsigned RHSCols = cast<ConstantInt>(Val: Orig->getArgOperand(i: 4))->getZExtValue(); |
| 1117 | |
| 1118 | auto *RetTy = cast<FixedVectorType>(Val: Orig->getType()); |
| 1119 | Type *EltTy = RetTy->getElementType(); |
| 1120 | bool IsFP = EltTy->isFloatingPointTy(); |
| 1121 | |
| 1122 | IRBuilder<> Builder(Orig); |
| 1123 | |
| 1124 | // Column-major indexing: |
| 1125 | // LHS row R, element K: index = K * LHSRows + R |
| 1126 | // RHS col C, element K: index = C * LHSCols + K |
| 1127 | Value *Result = PoisonValue::get(T: RetTy); |
| 1128 | |
| 1129 | // Extract all scalar elements from LHS and RHS once, then reuse them. |
| 1130 | unsigned LHSSize = LHSRows * LHSCols; |
| 1131 | unsigned RHSSize = LHSCols * RHSCols; |
| 1132 | SmallVector<Value *, 16> LHSElts(LHSSize); |
| 1133 | SmallVector<Value *, 16> RHSElts(RHSSize); |
| 1134 | for (unsigned I = 0; I < LHSSize; ++I) |
| 1135 | LHSElts[I] = Builder.CreateExtractElement(Vec: LHS, Idx: I); |
| 1136 | for (unsigned I = 0; I < RHSSize; ++I) |
| 1137 | RHSElts[I] = Builder.CreateExtractElement(Vec: RHS, Idx: I); |
| 1138 | |
| 1139 | // Choose the appropriate scalar-arg dot intrinsic for floats. |
| 1140 | // K=1 and double types use scalar expansion instead. |
| 1141 | Intrinsic::ID FloatDotID = Intrinsic::not_intrinsic; |
| 1142 | bool UseScalarFP = IsFP && (EltTy->isDoubleTy() || LHSCols == 1); |
| 1143 | if (IsFP && !UseScalarFP) { |
| 1144 | switch (LHSCols) { |
| 1145 | case 2: |
| 1146 | FloatDotID = Intrinsic::dx_dot2; |
| 1147 | break; |
| 1148 | case 3: |
| 1149 | FloatDotID = Intrinsic::dx_dot3; |
| 1150 | break; |
| 1151 | case 4: |
| 1152 | FloatDotID = Intrinsic::dx_dot4; |
| 1153 | break; |
| 1154 | default: |
| 1155 | reportFatalUsageError( |
| 1156 | reason: "Invalid matrix inner dimension for dot product: must be 2-4" ); |
| 1157 | return nullptr; |
| 1158 | } |
| 1159 | } |
| 1160 | |
| 1161 | for (unsigned C = 0; C < RHSCols; ++C) { |
| 1162 | for (unsigned R = 0; R < LHSRows; ++R) { |
| 1163 | // Gather row R from LHS and column C from RHS. |
| 1164 | SmallVector<Value *, 4> RowElts, ColElts; |
| 1165 | for (unsigned K = 0; K < LHSCols; ++K) { |
| 1166 | RowElts.push_back(Elt: LHSElts[K * LHSRows + R]); |
| 1167 | ColElts.push_back(Elt: RHSElts[C * LHSCols + K]); |
| 1168 | } |
| 1169 | |
| 1170 | Value *Dot; |
| 1171 | if (UseScalarFP) { |
| 1172 | // Scalar fmul+fmuladd expansion for double types and K=1. |
| 1173 | Dot = Builder.CreateFMul(L: RowElts[0], R: ColElts[0]); |
| 1174 | for (unsigned K = 1; K < LHSCols; ++K) |
| 1175 | Dot = Builder.CreateIntrinsic(RetTy: EltTy, ID: Intrinsic::fmuladd, |
| 1176 | Args: {RowElts[K], ColElts[K], Dot}); |
| 1177 | } else if (IsFP) { |
| 1178 | // Emit scalar-arg DXIL dot directly (dx.dot2/dx.dot3/dx.dot4). |
| 1179 | SmallVector<Value *, 8> Args; |
| 1180 | Args.append(in_start: RowElts.begin(), in_end: RowElts.end()); |
| 1181 | Args.append(in_start: ColElts.begin(), in_end: ColElts.end()); |
| 1182 | Dot = Builder.CreateIntrinsic(RetTy: EltTy, ID: FloatDotID, Args); |
| 1183 | } else { |
| 1184 | // Integer: emit multiply + imad chain. |
| 1185 | Dot = Builder.CreateMul(LHS: RowElts[0], RHS: ColElts[0]); |
| 1186 | for (unsigned K = 1; K < LHSCols; ++K) |
| 1187 | Dot = Builder.CreateIntrinsic(RetTy: EltTy, ID: Intrinsic::dx_imad, |
| 1188 | Args: {RowElts[K], ColElts[K], Dot}); |
| 1189 | } |
| 1190 | unsigned ResIdx = C * LHSRows + R; |
| 1191 | Result = Builder.CreateInsertElement(Vec: Result, NewElt: Dot, Idx: ResIdx); |
| 1192 | } |
| 1193 | } |
| 1194 | return Result; |
| 1195 | } |
| 1196 | |
| 1197 | // Expand llvm.matrix.transpose as a shufflevector that permutes elements |
| 1198 | // from column-major source to column-major transposed layout. |
| 1199 | // Element (r,c) at index c*Rows + r moves to index r*Cols + c. |
| 1200 | static Value *expandMatrixTranspose(CallInst *Orig) { |
| 1201 | Value *Mat = Orig->getArgOperand(i: 0); |
| 1202 | unsigned Rows = cast<ConstantInt>(Val: Orig->getArgOperand(i: 1))->getZExtValue(); |
| 1203 | unsigned Cols = cast<ConstantInt>(Val: Orig->getArgOperand(i: 2))->getZExtValue(); |
| 1204 | |
| 1205 | unsigned NumElts = Rows * Cols; |
| 1206 | SmallVector<int, 16> Mask(NumElts); |
| 1207 | for (unsigned I = 0; I < NumElts; ++I) |
| 1208 | Mask[I] = (I % Cols) * Rows + (I / Cols); |
| 1209 | |
| 1210 | IRBuilder<> Builder(Orig); |
| 1211 | return Builder.CreateShuffleVector(V: Mat, Mask); |
| 1212 | } |
| 1213 | |
| 1214 | // Scalarize a vector int_dx_store_output call into per-component scalar calls. |
| 1215 | // The DXIL StoreOutput op is per-component; vector intrinsics are split here |
| 1216 | // so that DXILOpLowering sees only scalar variants. |
| 1217 | static bool expandStoreOutput(CallInst *Orig) { |
| 1218 | auto *VT = dyn_cast<FixedVectorType>(Val: Orig->getArgOperand(i: 3)->getType()); |
| 1219 | if (!VT) |
| 1220 | return false; // already scalar, nothing to expand |
| 1221 | |
| 1222 | IRBuilder<> Builder(Orig); |
| 1223 | Module *M = Orig->getModule(); |
| 1224 | Type *Int8Ty = Builder.getInt8Ty(); |
| 1225 | Type *Int32Ty = Builder.getInt32Ty(); |
| 1226 | Type *ScalarTy = VT->getElementType(); |
| 1227 | unsigned NumElems = VT->getNumElements(); |
| 1228 | |
| 1229 | Value *SigElementId = Orig->getArgOperand(i: 0); |
| 1230 | Value *RowIndex = Orig->getArgOperand(i: 1); |
| 1231 | Value *StartCol = Orig->getArgOperand(i: 2); // i8 |
| 1232 | Value *Data = Orig->getArgOperand(i: 3); |
| 1233 | Value *StartColI32 = Builder.CreateZExt(V: StartCol, DestTy: Int32Ty); |
| 1234 | |
| 1235 | Function *ScalarFn = Intrinsic::getOrInsertDeclaration( |
| 1236 | M, id: Intrinsic::dx_store_output, OverloadTys: {ScalarTy}); |
| 1237 | |
| 1238 | for (unsigned I = 0; I < NumElems; ++I) { |
| 1239 | Value *Scalar = |
| 1240 | Builder.CreateExtractElement(Vec: Data, Idx: ConstantInt::get(Ty: Int32Ty, V: I)); |
| 1241 | Value *ColIdx = |
| 1242 | Builder.CreateAdd(LHS: StartColI32, RHS: ConstantInt::get(Ty: Int32Ty, V: I)); |
| 1243 | Value *ColI8 = Builder.CreateTrunc(V: ColIdx, DestTy: Int8Ty); |
| 1244 | Builder.CreateCall(Callee: ScalarFn, Args: {SigElementId, RowIndex, ColI8, Scalar}); |
| 1245 | } |
| 1246 | |
| 1247 | Orig->eraseFromParent(); |
| 1248 | return true; |
| 1249 | } |
| 1250 | |
| 1251 | // Scalarize a vector int_dx_load_input call into per-component scalar calls |
| 1252 | // and reassemble the vector. The DXIL LoadInput op is per-component. |
| 1253 | static Value *expandLoadInput(CallInst *Orig) { |
| 1254 | auto *VT = dyn_cast<FixedVectorType>(Val: Orig->getType()); |
| 1255 | if (!VT) |
| 1256 | return nullptr; // already scalar, nothing to expand |
| 1257 | |
| 1258 | IRBuilder<> Builder(Orig); |
| 1259 | Module *M = Orig->getModule(); |
| 1260 | Type *Int8Ty = Builder.getInt8Ty(); |
| 1261 | Type *Int32Ty = Builder.getInt32Ty(); |
| 1262 | Type *ScalarTy = VT->getElementType(); |
| 1263 | unsigned NumElems = VT->getNumElements(); |
| 1264 | |
| 1265 | Value *SigElementId = Orig->getArgOperand(i: 0); |
| 1266 | Value *RowIndex = Orig->getArgOperand(i: 1); |
| 1267 | Value *StartCol = Orig->getArgOperand(i: 2); // i8 |
| 1268 | Value *GsVertexOrPrimIndex = Orig->getArgOperand(i: 3); |
| 1269 | Value *StartColI32 = Builder.CreateZExt(V: StartCol, DestTy: Int32Ty); |
| 1270 | |
| 1271 | Function *ScalarFn = Intrinsic::getOrInsertDeclaration( |
| 1272 | M, id: Intrinsic::dx_load_input, OverloadTys: {ScalarTy}); |
| 1273 | |
| 1274 | Value *Vec = PoisonValue::get(T: VT); |
| 1275 | for (unsigned I = 0; I < NumElems; ++I) { |
| 1276 | Value *ColIdx = |
| 1277 | Builder.CreateAdd(LHS: StartColI32, RHS: ConstantInt::get(Ty: Int32Ty, V: I)); |
| 1278 | Value *ColI8 = Builder.CreateTrunc(V: ColIdx, DestTy: Int8Ty); |
| 1279 | Value *Scalar = Builder.CreateCall( |
| 1280 | Callee: ScalarFn, Args: {SigElementId, RowIndex, ColI8, GsVertexOrPrimIndex}); |
| 1281 | Vec = |
| 1282 | Builder.CreateInsertElement(Vec, NewElt: Scalar, Idx: ConstantInt::get(Ty: Int32Ty, V: I)); |
| 1283 | } |
| 1284 | |
| 1285 | return Vec; |
| 1286 | } |
| 1287 | |
| 1288 | static bool expandIntrinsic(Function &F, CallInst *Orig) { |
| 1289 | Value *Result = nullptr; |
| 1290 | Intrinsic::ID IntrinsicId = F.getIntrinsicID(); |
| 1291 | switch (IntrinsicId) { |
| 1292 | case Intrinsic::abs: |
| 1293 | Result = expandAbs(Orig); |
| 1294 | break; |
| 1295 | case Intrinsic::assume: |
| 1296 | Orig->eraseFromParent(); |
| 1297 | return true; |
| 1298 | case Intrinsic::atan2: |
| 1299 | Result = expandAtan2Intrinsic(Orig); |
| 1300 | break; |
| 1301 | case Intrinsic::copysign: |
| 1302 | Result = expandCopySignIntrinsic(Orig); |
| 1303 | break; |
| 1304 | case Intrinsic::fshl: |
| 1305 | Result = expandFunnelShiftIntrinsic<true>(Orig); |
| 1306 | break; |
| 1307 | case Intrinsic::fshr: |
| 1308 | Result = expandFunnelShiftIntrinsic<false>(Orig); |
| 1309 | break; |
| 1310 | case Intrinsic::exp: |
| 1311 | Result = expandExpIntrinsic(Orig); |
| 1312 | break; |
| 1313 | case Intrinsic::is_fpclass: |
| 1314 | Result = expandIsFPClass(Orig); |
| 1315 | break; |
| 1316 | case Intrinsic::log: |
| 1317 | Result = expandLogIntrinsic(Orig); |
| 1318 | break; |
| 1319 | case Intrinsic::log10: |
| 1320 | Result = expandLog10Intrinsic(Orig); |
| 1321 | break; |
| 1322 | case Intrinsic::pow: |
| 1323 | case Intrinsic::powi: |
| 1324 | Result = expandPowIntrinsic(Orig, IntrinsicId); |
| 1325 | break; |
| 1326 | case Intrinsic::dx_all: |
| 1327 | case Intrinsic::dx_any: |
| 1328 | Result = expandAnyOrAllIntrinsic(Orig, IntrinsicId); |
| 1329 | break; |
| 1330 | case Intrinsic::dx_uclamp: |
| 1331 | case Intrinsic::dx_sclamp: |
| 1332 | case Intrinsic::dx_nclamp: |
| 1333 | Result = expandClampIntrinsic(Orig, ClampIntrinsic: IntrinsicId); |
| 1334 | break; |
| 1335 | case Intrinsic::dx_isfinite: |
| 1336 | Result = expand16BitIsFinite(Orig); |
| 1337 | break; |
| 1338 | case Intrinsic::dx_isinf: |
| 1339 | Result = expand16BitIsInf(Orig); |
| 1340 | break; |
| 1341 | case Intrinsic::dx_isnan: |
| 1342 | Result = expand16BitIsNaN(Orig); |
| 1343 | break; |
| 1344 | case Intrinsic::dx_fdot: |
| 1345 | Result = expandFloatDotIntrinsic(Orig); |
| 1346 | break; |
| 1347 | case Intrinsic::dx_sdot: |
| 1348 | case Intrinsic::dx_udot: |
| 1349 | Result = expandIntegerDotIntrinsic(Orig, DotIntrinsic: IntrinsicId); |
| 1350 | break; |
| 1351 | case Intrinsic::dx_sign: |
| 1352 | Result = expandSignIntrinsic(Orig); |
| 1353 | break; |
| 1354 | case Intrinsic::dx_load_input: |
| 1355 | Result = expandLoadInput(Orig); |
| 1356 | break; |
| 1357 | case Intrinsic::dx_store_output: |
| 1358 | if (expandStoreOutput(Orig)) |
| 1359 | return true; |
| 1360 | break; |
| 1361 | case Intrinsic::dx_resource_load_rawbuffer: |
| 1362 | if (expandBufferLoadIntrinsic(Orig, /*IsRaw*/ true)) |
| 1363 | return true; |
| 1364 | break; |
| 1365 | case Intrinsic::dx_resource_store_rawbuffer: |
| 1366 | if (expandBufferStoreIntrinsic(Orig, /*IsRaw*/ true)) |
| 1367 | return true; |
| 1368 | break; |
| 1369 | case Intrinsic::dx_resource_load_typedbuffer: |
| 1370 | if (expandBufferLoadIntrinsic(Orig, /*IsRaw*/ false)) |
| 1371 | return true; |
| 1372 | break; |
| 1373 | case Intrinsic::dx_resource_store_typedbuffer: |
| 1374 | if (expandBufferStoreIntrinsic(Orig, /*IsRaw*/ false)) |
| 1375 | return true; |
| 1376 | break; |
| 1377 | case Intrinsic::usub_sat: |
| 1378 | Result = expandUsubSat(Orig); |
| 1379 | break; |
| 1380 | case Intrinsic::umul_with_overflow: |
| 1381 | case Intrinsic::smul_with_overflow: |
| 1382 | Result = expandMulWithOverflow(Orig, /*Signed=*/IntrinsicId == |
| 1383 | Intrinsic::smul_with_overflow); |
| 1384 | break; |
| 1385 | case Intrinsic::vector_reduce_add: |
| 1386 | case Intrinsic::vector_reduce_fadd: |
| 1387 | Result = expandVecReduceAdd(Orig, IntrinsicId); |
| 1388 | break; |
| 1389 | case Intrinsic::matrix_multiply: |
| 1390 | Result = expandMatrixMultiply(Orig); |
| 1391 | break; |
| 1392 | case Intrinsic::matrix_transpose: |
| 1393 | Result = expandMatrixTranspose(Orig); |
| 1394 | break; |
| 1395 | } |
| 1396 | if (Result) { |
| 1397 | Orig->replaceAllUsesWith(V: Result); |
| 1398 | Orig->eraseFromParent(); |
| 1399 | return true; |
| 1400 | } |
| 1401 | return false; |
| 1402 | } |
| 1403 | |
| 1404 | static bool expansionIntrinsics(Module &M) { |
| 1405 | for (auto &F : make_early_inc_range(Range: M.functions())) { |
| 1406 | if (!isIntrinsicExpansion(F)) |
| 1407 | continue; |
| 1408 | bool IntrinsicExpanded = false; |
| 1409 | for (User *U : make_early_inc_range(Range: F.users())) { |
| 1410 | auto *IntrinsicCall = dyn_cast<CallInst>(Val: U); |
| 1411 | if (!IntrinsicCall) |
| 1412 | continue; |
| 1413 | IntrinsicExpanded = expandIntrinsic(F, Orig: IntrinsicCall); |
| 1414 | } |
| 1415 | if (F.user_empty() && IntrinsicExpanded) |
| 1416 | F.eraseFromParent(); |
| 1417 | } |
| 1418 | return true; |
| 1419 | } |
| 1420 | |
| 1421 | PreservedAnalyses DXILIntrinsicExpansion::run(Module &M, |
| 1422 | ModuleAnalysisManager &) { |
| 1423 | if (expansionIntrinsics(M)) |
| 1424 | return PreservedAnalyses::none(); |
| 1425 | return PreservedAnalyses::all(); |
| 1426 | } |
| 1427 | |
| 1428 | bool DXILIntrinsicExpansionLegacy::runOnModule(Module &M) { |
| 1429 | return expansionIntrinsics(M); |
| 1430 | } |
| 1431 | |
| 1432 | char DXILIntrinsicExpansionLegacy::ID = 0; |
| 1433 | |
| 1434 | INITIALIZE_PASS_BEGIN(DXILIntrinsicExpansionLegacy, DEBUG_TYPE, |
| 1435 | "DXIL Intrinsic Expansion" , false, false) |
| 1436 | INITIALIZE_PASS_END(DXILIntrinsicExpansionLegacy, DEBUG_TYPE, |
| 1437 | "DXIL Intrinsic Expansion" , false, false) |
| 1438 | |
| 1439 | ModulePass *llvm::createDXILIntrinsicExpansionLegacyPass() { |
| 1440 | return new DXILIntrinsicExpansionLegacy(); |
| 1441 | } |
| 1442 | |