LLVM 24.0.0git
X86LowerAMXIntrinsics.cpp
Go to the documentation of this file.
1//===-- X86LowerAMXIntrinsics.cpp -X86 Scalarize AMX Intrinsics------------===//
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 Pass to transform amx intrinsics to scalar operations.
10/// This pass is always enabled and it skips when it is not -O0 and has no
11/// optnone attributes. With -O0 or optnone attribute, the def of shape to amx
12/// intrinsics is near the amx intrinsics code. We are not able to find a
13/// point which post-dominate all the shape and dominate all amx intrinsics.
14/// To decouple the dependency of the shape, we transform amx intrinsics
15/// to scalar operation, so that compiling doesn't fail. In long term, we
16/// should improve fast register allocation to allocate amx register.
17//===----------------------------------------------------------------------===//
18//
19#include "X86.h"
20#include "X86TargetMachine.h"
24#include "llvm/CodeGen/Passes.h"
27#include "llvm/IR/Analysis.h"
28#include "llvm/IR/DataLayout.h"
29#include "llvm/IR/Dominators.h"
30#include "llvm/IR/Function.h"
31#include "llvm/IR/IRBuilder.h"
34#include "llvm/IR/IntrinsicsX86.h"
35#include "llvm/IR/MDBuilder.h"
36#include "llvm/IR/PassManager.h"
40#include "llvm/Pass.h"
45
46using namespace llvm;
47using namespace PatternMatch;
48
49namespace llvm {
51} // end namespace llvm
52
53#define DEBUG_TYPE "x86-lower-amx-intrinsics"
54
55#ifndef NDEBUG
56static bool isV256I32Ty(Type *Ty) {
57 if (auto *FVT = dyn_cast<FixedVectorType>(Ty))
58 return FVT->getNumElements() == 256 &&
59 FVT->getElementType()->isIntegerTy(32);
60 return false;
61}
62#endif
63
64namespace {
65class X86LowerAMXIntrinsics {
66 Function &Func;
67
68public:
69 X86LowerAMXIntrinsics(Function &F, DomTreeUpdater &DomTU, LoopInfo *LoopI)
70 : Func(F), DTU(DomTU), LI(LoopI) {}
71 bool visit();
72
73private:
74 DomTreeUpdater &DTU;
75 LoopInfo *LI;
76 BasicBlock *createLoop(BasicBlock *Preheader, BasicBlock *Exit, Value *Bound,
77 ConstantInt *Step, StringRef Name, IRBuilderBase &B,
78 Loop *L);
79 template <bool IsTileLoad>
80 Value *createTileLoadStoreLoops(BasicBlock *Start, BasicBlock *End,
81 IRBuilderBase &B, Value *Row, Value *Col,
82 Value *Ptr, Value *Stride, Value *Tile);
83 template <Intrinsic::ID IntrID>
84 std::enable_if_t<IntrID == Intrinsic::x86_tdpbssd_internal ||
85 IntrID == Intrinsic::x86_tdpbsud_internal ||
86 IntrID == Intrinsic::x86_tdpbusd_internal ||
87 IntrID == Intrinsic::x86_tdpbuud_internal ||
88 IntrID == Intrinsic::x86_tdpbf16ps_internal,
89 Value *>
90 createTileDPLoops(BasicBlock *Start, BasicBlock *End, IRBuilderBase &B,
91 Value *Row, Value *Col, Value *K, Value *Acc, Value *LHS,
92 Value *RHS);
93 template <bool IsTileLoad>
94 bool lowerTileLoadStore(Instruction *TileLoadStore);
95 template <Intrinsic::ID IntrID>
96 std::enable_if_t<IntrID == Intrinsic::x86_tdpbssd_internal ||
97 IntrID == Intrinsic::x86_tdpbsud_internal ||
98 IntrID == Intrinsic::x86_tdpbusd_internal ||
99 IntrID == Intrinsic::x86_tdpbuud_internal ||
100 IntrID == Intrinsic::x86_tdpbf16ps_internal,
101 bool>
102 lowerTileDP(Instruction *TileDP);
103 bool lowerTileZero(Instruction *TileZero);
104};
105} // anonymous namespace
106
107BasicBlock *X86LowerAMXIntrinsics::createLoop(BasicBlock *Preheader,
108 BasicBlock *Exit, Value *Bound,
109 ConstantInt *Step, StringRef Name,
110 IRBuilderBase &B, Loop *L) {
111 LLVMContext &Ctx = Preheader->getContext();
112 BasicBlock *Header =
113 BasicBlock::Create(Ctx, Name + ".header", Preheader->getParent(), Exit);
114 BasicBlock *Body =
115 BasicBlock::Create(Ctx, Name + ".body", Header->getParent(), Exit);
116 BasicBlock *Latch =
117 BasicBlock::Create(Ctx, Name + ".latch", Header->getParent(), Exit);
118
119 Type *I16Ty = Type::getInt16Ty(Ctx);
120 UncondBrInst::Create(Body, Header);
121 UncondBrInst::Create(Latch, Body);
122 PHINode *IV =
123 PHINode::Create(I16Ty, 2, Name + ".iv", Header->getTerminator()->getIterator());
124 IV->addIncoming(ConstantInt::get(I16Ty, 0), Preheader);
125
126 B.SetInsertPoint(Latch);
127 Value *Inc = B.CreateAdd(IV, Step, Name + ".step");
128 Value *Cond = B.CreateICmpNE(Inc, Bound, Name + ".cond");
129 auto *BR = CondBrInst::Create(Cond, Header, Exit, Latch);
131 if (auto *BoundInt = dyn_cast<ConstantInt>(Bound)) {
132 assert(Step->getZExtValue() != 0 &&
133 "Expected a non-zero step size. This is chosen by the pass and "
134 "should always be non-zero to imply a finite loop.");
135 MDBuilder MDB(Preheader->getContext());
137 *BR, {BoundInt->getZExtValue() / Step->getZExtValue(), 1}, false);
138 } else {
140 }
141 }
142 IV->addIncoming(Inc, Latch);
143
144 UncondBrInst *PreheaderBr = cast<UncondBrInst>(Preheader->getTerminator());
145 BasicBlock *Tmp = PreheaderBr->getSuccessor();
146 PreheaderBr->setSuccessor(Header);
148 {DominatorTree::Delete, Preheader, Tmp},
149 {DominatorTree::Insert, Header, Body},
150 {DominatorTree::Insert, Body, Latch},
151 {DominatorTree::Insert, Latch, Header},
152 {DominatorTree::Insert, Latch, Exit},
153 {DominatorTree::Insert, Preheader, Header},
154 });
155 if (LI) {
156 L->addBasicBlockToLoop(Header, *LI);
157 L->addBasicBlockToLoop(Body, *LI);
158 L->addBasicBlockToLoop(Latch, *LI);
159 }
160 return Body;
161}
162
163template <bool IsTileLoad>
164Value *X86LowerAMXIntrinsics::createTileLoadStoreLoops(
165 BasicBlock *Start, BasicBlock *End, IRBuilderBase &B, Value *Row,
166 Value *Col, Value *Ptr, Value *Stride, Value *Tile) {
167 std::string IntrinName = IsTileLoad ? "tileload" : "tilestore";
168 Loop *RowLoop = nullptr;
169 Loop *ColLoop = nullptr;
170 if (LI) {
171 RowLoop = LI->AllocateLoop();
172 ColLoop = LI->AllocateLoop();
173 RowLoop->addChildLoop(ColLoop);
174 if (Loop *ParentL = LI->getLoopFor(Start))
175 ParentL->addChildLoop(RowLoop);
176 else
177 LI->addTopLevelLoop(RowLoop);
178 }
179
180 BasicBlock *RowBody = createLoop(Start, End, Row, B.getInt16(1),
181 IntrinName + ".scalarize.rows", B, RowLoop);
182 BasicBlock *RowLatch = RowBody->getSingleSuccessor();
183
184 BasicBlock *ColBody = createLoop(RowBody, RowLatch, Col, B.getInt16(1),
185 IntrinName + ".scalarize.cols", B, ColLoop);
186
187 BasicBlock *ColLoopLatch = ColBody->getSingleSuccessor();
188 BasicBlock *ColLoopHeader = ColBody->getSinglePredecessor();
189 BasicBlock *RowLoopHeader = RowBody->getSinglePredecessor();
190 Value *CurrentRow = &*RowLoopHeader->begin();
191 Value *CurrentCol = &*ColLoopHeader->begin();
192 Type *EltTy = B.getInt32Ty();
193 FixedVectorType *V256I32Ty = FixedVectorType::get(EltTy, 256);
194
195 // Common part for tileload and tilestore
196 // *.scalarize.cols.body:
197 // Calculate %idxmem and %idxvec
198 B.SetInsertPoint(ColBody->getTerminator());
199 Value *CurrentRowZExt = B.CreateZExt(CurrentRow, Stride->getType());
200 Value *CurrentColZExt = B.CreateZExt(CurrentCol, Stride->getType());
201 Value *Offset =
202 B.CreateAdd(B.CreateMul(CurrentRowZExt, Stride), CurrentColZExt);
203 Value *EltPtr = B.CreateGEP(EltTy, Ptr, Offset);
204 Value *Idx = B.CreateAdd(B.CreateMul(CurrentRow, B.getInt16(16)), CurrentCol);
205 if (IsTileLoad) {
206 // tileload.scalarize.rows.header:
207 // %vec.phi.row = phi <256 x i32> [ zeroinitializer, %entry ], [ %ResVec,
208 // %tileload.scalarize.rows.latch ]
209 B.SetInsertPoint(RowLoopHeader->getTerminator());
210 Value *VecZero = Constant::getNullValue(V256I32Ty);
211 PHINode *VecCPhiRowLoop = B.CreatePHI(V256I32Ty, 2, "vec.phi.row");
212 VecCPhiRowLoop->addIncoming(VecZero, Start);
213
214 // tileload.scalarize.cols.header:
215 // %vec.phi = phi <256 x i32> [ %vec.phi.row, %tileload.scalarize.rows.body
216 // ], [ %ResVec, %tileload.scalarize.cols.latch ]
217 B.SetInsertPoint(ColLoopHeader->getTerminator());
218 PHINode *VecPhi = B.CreatePHI(V256I32Ty, 2, "vec.phi");
219 VecPhi->addIncoming(VecCPhiRowLoop, RowBody);
220
221 // tileload.scalarize.cols.body:
222 // Calculate %idxmem and %idxvec
223 // %eltptr = getelementptr i32, i32* %base, i64 %idxmem
224 // %elt = load i32, i32* %ptr
225 // %ResVec = insertelement <256 x i32> %vec.phi, i32 %elt, i16 %idxvec
226 B.SetInsertPoint(ColBody->getTerminator());
227 Value *Elt = B.CreateLoad(EltTy, EltPtr);
228 Value *ResVec = B.CreateInsertElement(VecPhi, Elt, Idx);
229 VecPhi->addIncoming(ResVec, ColLoopLatch);
230 VecCPhiRowLoop->addIncoming(ResVec, RowLatch);
231
232 return ResVec;
233 } else {
234 auto *BitCast = cast<BitCastInst>(Tile);
235 Value *Vec = BitCast->getOperand(0);
236 assert(isV256I32Ty(Vec->getType()) && "bitcast from non-v256i32 to x86amx");
237 // tilestore.scalarize.cols.body:
238 // %mul = mul i16 %row.iv, i16 16
239 // %idx = add i16 %mul, i16 %col.iv
240 // %vec = extractelement <16 x i32> %vec, i16 %idx
241 // store i32 %vec, i32* %ptr
242 B.SetInsertPoint(ColBody->getTerminator());
243 Value *Elt = B.CreateExtractElement(Vec, Idx);
244
245 B.CreateStore(Elt, EltPtr);
246 return nullptr;
247 }
248}
249
250template <Intrinsic::ID IntrID>
251std::enable_if_t<IntrID == Intrinsic::x86_tdpbssd_internal ||
252 IntrID == Intrinsic::x86_tdpbsud_internal ||
253 IntrID == Intrinsic::x86_tdpbusd_internal ||
254 IntrID == Intrinsic::x86_tdpbuud_internal ||
255 IntrID == Intrinsic::x86_tdpbf16ps_internal,
256 Value *>
257X86LowerAMXIntrinsics::createTileDPLoops(BasicBlock *Start, BasicBlock *End,
258 IRBuilderBase &B, Value *Row,
259 Value *Col, Value *K, Value *Acc,
260 Value *LHS, Value *RHS) {
261 std::string IntrinName;
262 switch (IntrID) {
263 case Intrinsic::x86_tdpbssd_internal:
264 IntrinName = "tiledpbssd";
265 break;
266 case Intrinsic::x86_tdpbsud_internal:
267 IntrinName = "tiledpbsud";
268 break;
269 case Intrinsic::x86_tdpbusd_internal:
270 IntrinName = "tiledpbusd";
271 break;
272 case Intrinsic::x86_tdpbuud_internal:
273 IntrinName = "tiledpbuud";
274 break;
275 case Intrinsic::x86_tdpbf16ps_internal:
276 IntrinName = "tiledpbf16ps";
277 break;
278 }
279 Loop *RowLoop = nullptr;
280 Loop *ColLoop = nullptr;
281 Loop *InnerLoop = nullptr;
282 if (LI) {
283 RowLoop = LI->AllocateLoop();
284 ColLoop = LI->AllocateLoop();
285 InnerLoop = LI->AllocateLoop();
286 ColLoop->addChildLoop(InnerLoop);
287 RowLoop->addChildLoop(ColLoop);
288 if (Loop *ParentL = LI->getLoopFor(Start))
289 ParentL->addChildLoop(RowLoop);
290 else
291 LI->addTopLevelLoop(RowLoop);
292 }
293
294 BasicBlock *RowBody = createLoop(Start, End, Row, B.getInt16(1),
295 IntrinName + ".scalarize.rows", B, RowLoop);
296 BasicBlock *RowLatch = RowBody->getSingleSuccessor();
297
298 BasicBlock *ColBody = createLoop(RowBody, RowLatch, Col, B.getInt16(1),
299 IntrinName + ".scalarize.cols", B, ColLoop);
300
301 BasicBlock *ColLoopLatch = ColBody->getSingleSuccessor();
302
303 B.SetInsertPoint(ColBody->getTerminator());
304 BasicBlock *InnerBody =
305 createLoop(ColBody, ColLoopLatch, K, B.getInt16(1),
306 IntrinName + ".scalarize.inner", B, InnerLoop);
307
308 BasicBlock *ColLoopHeader = ColBody->getSinglePredecessor();
309 BasicBlock *RowLoopHeader = RowBody->getSinglePredecessor();
310 BasicBlock *InnerLoopHeader = InnerBody->getSinglePredecessor();
311 BasicBlock *InnerLoopLatch = InnerBody->getSingleSuccessor();
312 Value *CurrentRow = &*RowLoopHeader->begin();
313 Value *CurrentCol = &*ColLoopHeader->begin();
314 Value *CurrentInner = &*InnerLoopHeader->begin();
315
316 FixedVectorType *V256I32Ty = FixedVectorType::get(B.getInt32Ty(), 256);
317 auto *BitCastAcc = cast<BitCastInst>(Acc);
318 Value *VecC = BitCastAcc->getOperand(0);
319 assert(isV256I32Ty(VecC->getType()) && "bitcast from non-v256i32 to x86amx");
320 // TODO else create BitCast from x86amx to v256i32.
321 // Store x86amx to memory, and reload from memory
322 // to vector. However with -O0, it doesn't happen.
323 auto *BitCastLHS = cast<BitCastInst>(LHS);
324 Value *VecA = BitCastLHS->getOperand(0);
325 assert(isV256I32Ty(VecA->getType()) && "bitcast from non-v256i32 to x86amx");
326 auto *BitCastRHS = cast<BitCastInst>(RHS);
327 Value *VecB = BitCastRHS->getOperand(0);
328 assert(isV256I32Ty(VecB->getType()) && "bitcast from non-v256i32 to x86amx");
329
330 // tiledpbssd.scalarize.rows.header:
331 // %vec.c.phi.row = phi <256 x i32> [ %VecC, %continue ], [ %NewVecC,
332 // %tiledpbssd.scalarize.rows.latch ]
333
334 // %vec.d.phi.row = phi <256 x i32> [ zeroinitializer, %continue ], [
335 // %NewVecD, %tiledpbssd.scalarize.rows.latch ]
336 B.SetInsertPoint(RowLoopHeader->getTerminator());
337 PHINode *VecCPhiRowLoop = B.CreatePHI(V256I32Ty, 2, "vec.c.phi.row");
338 VecCPhiRowLoop->addIncoming(VecC, Start);
339 Value *VecZero = Constant::getNullValue(V256I32Ty);
340 PHINode *VecDPhiRowLoop = B.CreatePHI(V256I32Ty, 2, "vec.d.phi.row");
341 VecDPhiRowLoop->addIncoming(VecZero, Start);
342
343 // tiledpbssd.scalarize.cols.header:
344 // %vec.c.phi.col = phi <256 x i32> [ %vec.c.phi.row,
345 // %tiledpbssd.scalarize.rows.body ], [ %NewVecC,
346 // %tiledpbssd.scalarize.cols.latch ]
347
348 // %vec.d.phi.col = phi <256 x i32> [
349 // %vec.d.phi.row, %tiledpbssd.scalarize.rows.body ], [ %NewVecD,
350 // %tiledpbssd.scalarize.cols.latch ]
351
352 // calculate idxc.
353 B.SetInsertPoint(ColLoopHeader->getTerminator());
354 PHINode *VecCPhiColLoop = B.CreatePHI(V256I32Ty, 2, "vec.c.phi.col");
355 VecCPhiColLoop->addIncoming(VecCPhiRowLoop, RowBody);
356 PHINode *VecDPhiColLoop = B.CreatePHI(V256I32Ty, 2, "vec.d.phi.col");
357 VecDPhiColLoop->addIncoming(VecDPhiRowLoop, RowBody);
358 Value *IdxC =
359 B.CreateAdd(B.CreateMul(CurrentRow, B.getInt16(16)), CurrentCol);
360
361 // tiledpbssd.scalarize.inner.header:
362 // %vec.c.inner.phi = phi <256 x i32> [ %vec.c.phi.col,
363 // %tiledpbssd.scalarize.cols.body ], [ %NewVecC,
364 // %tiledpbssd.scalarize.inner.latch ]
365
366 B.SetInsertPoint(InnerLoopHeader->getTerminator());
367 PHINode *VecCPhi = B.CreatePHI(V256I32Ty, 2, "vec.c.inner.phi");
368 VecCPhi->addIncoming(VecCPhiColLoop, ColBody);
369
370 B.SetInsertPoint(InnerBody->getTerminator());
371 Value *IdxA =
372 B.CreateAdd(B.CreateMul(CurrentRow, B.getInt16(16)), CurrentInner);
373 Value *IdxB =
374 B.CreateAdd(B.CreateMul(CurrentInner, B.getInt16(16)), CurrentCol);
375 Value *NewVecC = nullptr;
376
377 if (IntrID != Intrinsic::x86_tdpbf16ps_internal) {
378 // tiledpbssd.scalarize.inner.body:
379 // calculate idxa, idxb
380 // %eltc = extractelement <256 x i32> %vec.c.inner.phi, i16 %idxc
381 // %elta = extractelement <256 x i32> %veca, i16 %idxa
382 // %eltav4i8 = bitcast i32 %elta to <4 x i8>
383 // %eltb = extractelement <256 x i32> %vecb, i16 %idxb
384 // %eltbv4i8 = bitcast i32 %eltb to <4 x i8>
385 // %eltav4i32 = sext <4 x i8> %eltav4i8 to <4 x i32>
386 // %eltbv4i32 = sext <4 x i8> %eltbv4i8 to <4 x i32>
387 // %mulab = mul <4 x i32> %eltbv4i32, %eltav4i32
388 // %acc = call i32 @llvm.vector.reduce.add.v4i32(<4 x i32> %131)
389 // %neweltc = add i32 %elt, %acc
390 // %NewVecC = insertelement <256 x i32> %vec.c.inner.phi, i32 %neweltc,
391 // i16 %idxc
392 FixedVectorType *V4I8Ty = FixedVectorType::get(B.getInt8Ty(), 4);
393 FixedVectorType *V4I32Ty = FixedVectorType::get(B.getInt32Ty(), 4);
394 Value *EltC = B.CreateExtractElement(VecCPhi, IdxC);
395 Value *EltA = B.CreateExtractElement(VecA, IdxA);
396 Value *SubVecA = B.CreateBitCast(EltA, V4I8Ty);
397 Value *EltB = B.CreateExtractElement(VecB, IdxB);
398 Value *SubVecB = B.CreateBitCast(EltB, V4I8Ty);
399 Value *SEXTSubVecB = nullptr;
400 Value *SEXTSubVecA = nullptr;
401 switch (IntrID) {
402 case Intrinsic::x86_tdpbssd_internal:
403 SEXTSubVecB = B.CreateSExt(SubVecB, V4I32Ty);
404 SEXTSubVecA = B.CreateSExt(SubVecA, V4I32Ty);
405 break;
406 case Intrinsic::x86_tdpbsud_internal:
407 SEXTSubVecB = B.CreateZExt(SubVecB, V4I32Ty);
408 SEXTSubVecA = B.CreateSExt(SubVecA, V4I32Ty);
409 break;
410 case Intrinsic::x86_tdpbusd_internal:
411 SEXTSubVecB = B.CreateSExt(SubVecB, V4I32Ty);
412 SEXTSubVecA = B.CreateZExt(SubVecA, V4I32Ty);
413 break;
414 case Intrinsic::x86_tdpbuud_internal:
415 SEXTSubVecB = B.CreateZExt(SubVecB, V4I32Ty);
416 SEXTSubVecA = B.CreateZExt(SubVecA, V4I32Ty);
417 break;
418 default:
419 llvm_unreachable("Invalid intrinsic ID!");
420 }
421 Value *SubVecR = B.CreateAddReduce(B.CreateMul(SEXTSubVecA, SEXTSubVecB));
422 Value *ResElt = B.CreateAdd(EltC, SubVecR);
423 NewVecC = B.CreateInsertElement(VecCPhi, ResElt, IdxC);
424 } else {
425 // tiledpbf16ps.scalarize.inner.body:
426 // calculate idxa, idxb, idxc
427 // %eltc = extractelement <256 x i32> %vec.c.inner.phi, i16 %idxc
428 // %eltcf32 = bitcast i32 %eltc to float
429 // %elta = extractelement <256 x i32> %veca, i16 %idxa
430 // %eltav2i16 = bitcast i32 %elta to <2 x i16>
431 // %eltb = extractelement <256 x i32> %vecb, i16 %idxb
432 // %eltbv2i16 = bitcast i32 %eltb to <2 x i16>
433 // %shufflea = shufflevector <2 x i16> %elta, <2 x i16> zeroinitializer, <4
434 // x i32> <i32 2, i32 0, i32 3, i32 1>
435 // %eltav2f32 = bitcast <4 x i16> %shufflea to <2 x float>
436 // %shuffleb = shufflevector <2 x i16> %eltb, <2 xi16> zeroinitializer, <4 x
437 // i32> <i32 2, i32 0, i32 3, i32 1>
438 // %eltbv2f32 = bitcast <4 x i16> %shuffleb to <2 x float>
439 // %mulab = fmul <2 x float> %eltav2f32, %eltbv2f32
440 // %acc = call float
441 // @llvm.vector.reduce.fadd.v2f32(float %eltcf32, <2 x float> %mulab)
442 // %neweltc = bitcast float %acc to i32
443 // %NewVecC = insertelement <256 x i32> %vec.c.inner.phi, i32 %neweltc,
444 // i16 %idxc
445 // %NewVecD = insertelement <256 x i32> %vec.d.inner.phi, i32 %neweltc,
446 // i16 %idxc
447 FixedVectorType *V2I16Ty = FixedVectorType::get(B.getInt16Ty(), 2);
448 FixedVectorType *V2F32Ty = FixedVectorType::get(B.getFloatTy(), 2);
449 Value *EltC = B.CreateExtractElement(VecCPhi, IdxC);
450 Value *EltCF32 = B.CreateBitCast(EltC, B.getFloatTy());
451 Value *EltA = B.CreateExtractElement(VecA, IdxA);
452 Value *SubVecA = B.CreateBitCast(EltA, V2I16Ty);
453 Value *EltB = B.CreateExtractElement(VecB, IdxB);
454 Value *SubVecB = B.CreateBitCast(EltB, V2I16Ty);
455 Value *ZeroV2I16 = Constant::getNullValue(V2I16Ty);
456 int ShuffleMask[4] = {2, 0, 3, 1};
457 auto ShuffleArray = ArrayRef(ShuffleMask);
458 Value *AV2F32 = B.CreateBitCast(
459 B.CreateShuffleVector(SubVecA, ZeroV2I16, ShuffleArray), V2F32Ty);
460 Value *BV2F32 = B.CreateBitCast(
461 B.CreateShuffleVector(SubVecB, ZeroV2I16, ShuffleArray), V2F32Ty);
462 Value *SubVecR = B.CreateFAddReduce(EltCF32, B.CreateFMul(AV2F32, BV2F32));
463 Value *ResElt = B.CreateBitCast(SubVecR, B.getInt32Ty());
464 NewVecC = B.CreateInsertElement(VecCPhi, ResElt, IdxC);
465 }
466
467 // tiledpbssd.scalarize.cols.latch:
468 // %NewEltC = extractelement <256 x i32> %vec.c.phi.col, i16 %idxc
469 // %NewVecD = insertelement <256 x i32> %vec.d.phi.col, i32 %NewEltC,
470 // i16 %idxc
471 B.SetInsertPoint(ColLoopLatch->getTerminator());
472 Value *NewEltC = B.CreateExtractElement(NewVecC, IdxC);
473 Value *NewVecD = B.CreateInsertElement(VecDPhiColLoop, NewEltC, IdxC);
474
475 VecCPhi->addIncoming(NewVecC, InnerLoopLatch);
476 VecCPhiRowLoop->addIncoming(NewVecC, RowLatch);
477 VecCPhiColLoop->addIncoming(NewVecC, ColLoopLatch);
478 VecDPhiRowLoop->addIncoming(NewVecD, RowLatch);
479 VecDPhiColLoop->addIncoming(NewVecD, ColLoopLatch);
480
481 return NewVecD;
482}
483
484template <Intrinsic::ID IntrID>
485std::enable_if_t<IntrID == Intrinsic::x86_tdpbssd_internal ||
486 IntrID == Intrinsic::x86_tdpbsud_internal ||
487 IntrID == Intrinsic::x86_tdpbusd_internal ||
488 IntrID == Intrinsic::x86_tdpbuud_internal ||
489 IntrID == Intrinsic::x86_tdpbf16ps_internal,
490 bool>
491X86LowerAMXIntrinsics::lowerTileDP(Instruction *TileDP) {
492 Value *M, *N, *K, *C, *A, *B;
494 m_Value(C), m_Value(A), m_Value(B)));
495 Instruction *InsertI = TileDP;
496 IRBuilder<> PreBuilder(TileDP);
497 PreBuilder.SetInsertPoint(TileDP);
498 // We visit the loop with (m, n/4, k/4):
499 // %n_dword = lshr i16 %n, 2
500 // %k_dword = lshr i16 %k, 2
501 Value *NDWord = PreBuilder.CreateLShr(N, PreBuilder.getInt16(2));
502 Value *KDWord = PreBuilder.CreateLShr(K, PreBuilder.getInt16(2));
503 BasicBlock *Start = InsertI->getParent();
504 BasicBlock *End =
505 SplitBlock(InsertI->getParent(), InsertI, &DTU, LI, nullptr, "continue");
506 IRBuilder<> Builder(TileDP);
507 Value *ResVec = createTileDPLoops<IntrID>(Start, End, Builder, M, NDWord,
508 KDWord, C, A, B);
509 // we cannot assume there always be bitcast after tiledpbssd. So we need to
510 // insert one bitcast as required
511 Builder.SetInsertPoint(End->getFirstNonPHIIt());
512 Value *ResAMX =
513 Builder.CreateBitCast(ResVec, Type::getX86_AMXTy(Builder.getContext()));
514 // Delete TileDP intrinsic and do some clean-up.
515 for (Use &U : llvm::make_early_inc_range(TileDP->uses())) {
516 Instruction *I = cast<Instruction>(U.getUser());
517 Value *Vec;
518 if (match(I, m_BitCast(m_Value(Vec)))) {
519 I->replaceAllUsesWith(ResVec);
520 I->eraseFromParent();
521 }
522 }
523 TileDP->replaceAllUsesWith(ResAMX);
524 TileDP->eraseFromParent();
525 return true;
526}
527
528template <bool IsTileLoad>
529bool X86LowerAMXIntrinsics::lowerTileLoadStore(Instruction *TileLoadStore) {
530 Value *M, *N, *Ptr, *Stride, *Tile;
531 if (IsTileLoad)
532 match(TileLoadStore,
534 m_Value(M), m_Value(N), m_Value(Ptr), m_Value(Stride)));
535 else
537 m_Value(M), m_Value(N), m_Value(Ptr),
538 m_Value(Stride), m_Value(Tile)));
539
540 Instruction *InsertI = TileLoadStore;
541 IRBuilder<> PreBuilder(TileLoadStore);
542 PreBuilder.SetInsertPoint(TileLoadStore);
543 Value *NDWord = PreBuilder.CreateLShr(N, PreBuilder.getInt16(2));
544 Value *StrideDWord = PreBuilder.CreateLShr(Stride, PreBuilder.getInt64(2));
545 BasicBlock *Start = InsertI->getParent();
546 BasicBlock *End =
547 SplitBlock(InsertI->getParent(), InsertI, &DTU, LI, nullptr, "continue");
548 IRBuilder<> Builder(TileLoadStore);
549 Value *ResVec = createTileLoadStoreLoops<IsTileLoad>(
550 Start, End, Builder, M, NDWord, Ptr, StrideDWord,
551 IsTileLoad ? nullptr : Tile);
552 if (IsTileLoad) {
553 // we cannot assume there always be bitcast after tileload. So we need to
554 // insert one bitcast as required
555 Builder.SetInsertPoint(End->getFirstNonPHIIt());
556 Value *ResAMX =
557 Builder.CreateBitCast(ResVec, Type::getX86_AMXTy(Builder.getContext()));
558 // Delete tileloadd6 intrinsic and do some clean-up
559 for (Use &U : llvm::make_early_inc_range(TileLoadStore->uses())) {
560 Instruction *I = cast<Instruction>(U.getUser());
561 Value *Vec;
562 if (match(I, m_BitCast(m_Value(Vec)))) {
563 I->replaceAllUsesWith(ResVec);
564 I->eraseFromParent();
565 }
566 }
567 TileLoadStore->replaceAllUsesWith(ResAMX);
568 }
569 TileLoadStore->eraseFromParent();
570 return true;
571}
572
573bool X86LowerAMXIntrinsics::lowerTileZero(Instruction *TileZero) {
574 IRBuilder<> Builder(TileZero);
575 FixedVectorType *V256I32Ty = FixedVectorType::get(Builder.getInt32Ty(), 256);
576 Value *VecZero = Constant::getNullValue(V256I32Ty);
577 for (Use &U : llvm::make_early_inc_range(TileZero->uses())) {
578 Instruction *I = cast<Instruction>(U.getUser());
579 Value *Vec;
580 if (match(I, m_BitCast(m_Value(Vec)))) {
581 I->replaceAllUsesWith(VecZero);
582 I->eraseFromParent();
583 }
584 }
585 TileZero->eraseFromParent();
586 return true;
587}
588
589bool X86LowerAMXIntrinsics::visit() {
590 bool C = false;
592 for (BasicBlock *BB : depth_first(&Func)) {
593 for (BasicBlock::iterator II = BB->begin(), IE = BB->end(); II != IE;) {
594 if (auto *Inst = dyn_cast<IntrinsicInst>(&*II++)) {
595 switch (Inst->getIntrinsicID()) {
596 case Intrinsic::x86_tdpbssd_internal:
597 case Intrinsic::x86_tdpbsud_internal:
598 case Intrinsic::x86_tdpbusd_internal:
599 case Intrinsic::x86_tdpbuud_internal:
600 case Intrinsic::x86_tileloadd64_internal:
601 case Intrinsic::x86_tilestored64_internal:
602 case Intrinsic::x86_tilezero_internal:
603 case Intrinsic::x86_tdpbf16ps_internal:
604 WorkList.push_back(Inst);
605 break;
606 default:
607 break;
608 }
609 }
610 }
611 }
612
613 for (auto *Inst : WorkList) {
614 switch (Inst->getIntrinsicID()) {
615 case Intrinsic::x86_tdpbssd_internal:
616 C = lowerTileDP<Intrinsic::x86_tdpbssd_internal>(Inst) || C;
617 break;
618 case Intrinsic::x86_tdpbsud_internal:
619 C = lowerTileDP<Intrinsic::x86_tdpbsud_internal>(Inst) || C;
620 break;
621 case Intrinsic::x86_tdpbusd_internal:
622 C = lowerTileDP<Intrinsic::x86_tdpbusd_internal>(Inst) || C;
623 break;
624 case Intrinsic::x86_tdpbuud_internal:
625 C = lowerTileDP<Intrinsic::x86_tdpbuud_internal>(Inst) || C;
626 break;
627 case Intrinsic::x86_tdpbf16ps_internal:
628 C = lowerTileDP<Intrinsic::x86_tdpbf16ps_internal>(Inst) || C;
629 break;
630 case Intrinsic::x86_tileloadd64_internal:
631 C = lowerTileLoadStore<true>(Inst) || C;
632 break;
633 case Intrinsic::x86_tilestored64_internal:
634 C = lowerTileLoadStore<false>(Inst) || C;
635 break;
636 case Intrinsic::x86_tilezero_internal:
637 C = lowerTileZero(Inst) || C;
638 break;
639 default:
640 llvm_unreachable("invalid amx intrinsics!");
641 }
642 }
643
644 return C;
645}
646
647namespace {
648bool shouldRunLowerAMXIntrinsics(const Function &F, const TargetMachine *TM) {
649 const X86Options &CLOpts =
650 static_cast<const X86TargetMachine *>(TM)->getCLOpts();
651 return CLOpts.enable_x86_scalar_amx &&
652 (F.hasFnAttribute(Attribute::OptimizeNone) ||
653 TM->getOptLevel() == CodeGenOptLevel::None);
654}
655
656bool runLowerAMXIntrinsics(Function &F, DominatorTree *DT, LoopInfo *LI) {
657 DomTreeUpdater DTU(DT, DomTreeUpdater::UpdateStrategy::Lazy);
658
659 X86LowerAMXIntrinsics LAT(F, DTU, LI);
660 return LAT.visit();
661}
662} // namespace
663
666 if (!shouldRunLowerAMXIntrinsics(F, TM))
667 return PreservedAnalyses::all();
668
669 DominatorTree &DT = FAM.getResult<DominatorTreeAnalysis>(F);
670 LoopInfo &LI = FAM.getResult<LoopAnalysis>(F);
671 bool Changed = runLowerAMXIntrinsics(F, &DT, &LI);
672 if (!Changed)
673 return PreservedAnalyses::all();
674
678 return PA;
679}
680
681namespace {
682class X86LowerAMXIntrinsicsLegacyPass : public FunctionPass {
683public:
684 static char ID;
685
686 X86LowerAMXIntrinsicsLegacyPass() : FunctionPass(ID) {}
687
688 bool runOnFunction(Function &F) override {
689 TargetMachine *TM = &getAnalysis<TargetPassConfig>().getTM<TargetMachine>();
690 if (!shouldRunLowerAMXIntrinsics(F, TM))
691 return false;
692
693 auto *DTWP = getAnalysisIfAvailable<DominatorTreeWrapperPass>();
694 auto *DT = DTWP ? &DTWP->getDomTree() : nullptr;
695 auto *LIWP = getAnalysisIfAvailable<LoopInfoWrapperPass>();
696 auto *LI = LIWP ? &LIWP->getLoopInfo() : nullptr;
697 return runLowerAMXIntrinsics(F, DT, LI);
698 }
699 StringRef getPassName() const override { return "Lower AMX intrinsics"; }
700
701 void getAnalysisUsage(AnalysisUsage &AU) const override {
702 AU.addPreserved<DominatorTreeWrapperPass>();
703 AU.addPreserved<LoopInfoWrapperPass>();
704 AU.addRequired<TargetPassConfig>();
705 }
706};
707} // namespace
708
709static const char PassName[] = "Lower AMX intrinsics";
710char X86LowerAMXIntrinsicsLegacyPass::ID = 0;
711INITIALIZE_PASS_BEGIN(X86LowerAMXIntrinsicsLegacyPass, DEBUG_TYPE, PassName,
712 false, false)
714INITIALIZE_PASS_END(X86LowerAMXIntrinsicsLegacyPass, DEBUG_TYPE, PassName,
716
718 return new X86LowerAMXIntrinsicsLegacyPass();
719}
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
static GCRegistry::Add< ErlangGC > A("erlang", "erlang-compatible garbage collector")
static GCRegistry::Add< OcamlGC > B("ocaml", "ocaml 3.10-compatible GC")
static bool runOnFunction(Function &F, bool PostInlining)
#define DEBUG_TYPE
This header defines various interfaces for pass management in LLVM.
#define F(x, y, z)
Definition MD5.cpp:54
#define I(x, y, z)
Definition MD5.cpp:57
uint64_t IntrinsicInst * II
FunctionAnalysisManager FAM
#define INITIALIZE_PASS_DEPENDENCY(depName)
Definition PassSupport.h:42
#define INITIALIZE_PASS_END(passName, arg, name, cfg, analysis)
Definition PassSupport.h:44
#define INITIALIZE_PASS_BEGIN(passName, arg, name, cfg, analysis)
Definition PassSupport.h:39
This file contains the declarations for profiling metadata utility functions.
const SmallVectorImpl< MachineOperand > & Cond
static void visit(BasicBlock &Start, std::function< bool(BasicBlock *)> op)
Target-Independent Code Generator Pass Configuration Options pass.
This pass exposes codegen information to IR-level passes.
static bool isV256I32Ty(Type *Ty)
static const char PassName[]
Value * RHS
Value * LHS
static const uint32_t IV[8]
Definition blake3_impl.h:83
AnalysisUsage & addRequired()
AnalysisUsage & addPreserved()
Add the specified Pass class to the set of analyses preserved by this pass.
LLVM Basic Block Representation.
Definition BasicBlock.h:62
iterator begin()
Instruction iterator methods.
Definition BasicBlock.h:446
const Function * getParent() const
Return the enclosing method, or null if none.
Definition BasicBlock.h:213
LLVM_ABI InstListType::const_iterator getFirstNonPHIIt() const
Returns an iterator to the first instruction in this block that is not a PHINode instruction.
static BasicBlock * Create(LLVMContext &Context, const Twine &Name="", Function *Parent=nullptr, BasicBlock *InsertBefore=nullptr)
Creates a new BasicBlock.
Definition BasicBlock.h:206
LLVM_ABI const BasicBlock * getSinglePredecessor() const
Return the predecessor of this block if it has a single predecessor block.
LLVM_ABI const BasicBlock * getSingleSuccessor() const
Return the successor of this block if it has a single successor.
InstListType::iterator iterator
Instruction iterators...
Definition BasicBlock.h:170
LLVM_ABI LLVMContext & getContext() const
Get the context in which this basic block lives.
const Instruction * getTerminator() const LLVM_READONLY
Returns the terminator instruction; assumes that the block is well-formed.
Definition BasicBlock.h:237
static CondBrInst * Create(Value *Cond, BasicBlock *IfTrue, BasicBlock *IfFalse, InsertPosition InsertBefore=nullptr)
This is the shared class of boolean and integer constants.
Definition Constants.h:87
uint64_t getZExtValue() const
Return the constant as a 64-bit unsigned integer value after it has been zero extended as appropriate...
Definition Constants.h:168
static LLVM_ABI Constant * getNullValue(Type *Ty)
Constructor to create a '0' constant of arbitrary type.
Analysis pass which computes a DominatorTree.
Definition Dominators.h:241
Concrete subclass of DominatorTreeBase that is used to compute a normal dominator tree.
Definition Dominators.h:122
static LLVM_ABI FixedVectorType * get(Type *ElementType, unsigned NumElts)
Definition Type.cpp:843
FunctionPass class - This class is used to implement most global optimizations.
Definition Pass.h:314
void applyUpdatesPermissive(ArrayRef< UpdateT > Updates)
Submit updates to all available trees.
Common base class shared among various IRBuilders.
Definition IRBuilder.h:111
LLVM_ABI InstListType::iterator eraseFromParent()
This method unlinks 'this' from the containing basic block and deletes it.
Analysis pass that exposes the LoopInfo for a function.
Definition LoopInfo.h:594
void addChildLoop(LoopT *NewChild)
Add the specified loop to be a child of this loop.
void addTopLevelLoop(LoopT *New)
This adds the specified loop to the collection of top-level loops.
LoopT * getLoopFor(const BlockT *BB) const
Return the inner most loop that BB lives in.
Represents a single loop in the control flow graph.
Definition LoopInfo.h:40
void addIncoming(Value *V, BasicBlock *BB)
Add an incoming value to the end of the PHI list.
static PHINode * Create(Type *Ty, unsigned NumReservedValues, const Twine &NameStr="", InsertPosition InsertBefore=nullptr)
Constructors - NumReservedValues is a hint for the number of incoming edges that this phi node will h...
A set of analyses that are preserved following a run of a transformation pass.
Definition Analysis.h:112
static PreservedAnalyses none()
Convenience factory function for the empty preserved set.
Definition Analysis.h:115
static PreservedAnalyses all()
Construct a special preserved set that preserves all passes.
Definition Analysis.h:118
PreservedAnalyses & preserve()
Mark an analysis as preserved.
Definition Analysis.h:132
void push_back(const T &Elt)
Represent a constant reference to a string, i.e.
Definition StringRef.h:56
Primary interface to the complete machine description for the target machine.
CodeGenOptLevel getOptLevel() const
Returns the optimization level: None, Less, Default, or Aggressive.
Target-Independent Code Generator Pass Configuration Options.
The instances of the Type class are immutable: once they are created, they are never changed.
Definition Type.h:46
void setSuccessor(BasicBlock *NewSucc)
static UncondBrInst * Create(BasicBlock *Target, InsertPosition InsertBefore=nullptr)
BasicBlock * getSuccessor(unsigned i=0) const
LLVM Value Representation.
Definition Value.h:75
Type * getType() const
All values are typed, get the type of this value.
Definition Value.h:257
LLVM_ABI void replaceAllUsesWith(Value *V)
Change all uses of this to point to a new Value.
Definition Value.cpp:553
iterator_range< use_iterator > uses()
Definition Value.h:382
PreservedAnalyses run(Function &F, FunctionAnalysisManager &FAM)
const ParentTy * getParent() const
Definition ilist_node.h:34
Changed
Pass manager infrastructure for declaring and invalidating analyses.
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
@ BR
Control flow instructions. These all have token chains.
@ BasicBlock
Various leaf nodes.
Definition ISDOpcodes.h:83
bool match(Val *V, const Pattern &P)
auto m_Value()
Match an arbitrary value and ignore it.
CastOperator_match< OpTy, Instruction::BitCast > m_BitCast(const OpTy &Op)
Matches BitCast.
auto m_Intrinsic(const Ts &...Ops)
Match intrinsic calls like this: m_Intrinsic<Intrinsic::fabs>(m_Value(X))
friend class Instruction
Iterator for Instructions in a `BasicBlock.
Definition BasicBlock.h:73
This is an optimization pass for GlobalISel generic memory operations.
@ Offset
Definition DWP.cpp:577
FunctionPass * createX86LowerAMXIntrinsicsLegacyPass()
LLVM_ABI cl::opt< bool > ProfcheckDisableMetadataFixes
Definition LoopInfo.cpp:60
LLVM_ABI void setExplicitlyUnknownBranchWeightsIfProfiled(Instruction &I, StringRef PassName, const Function *F=nullptr)
Like setExplicitlyUnknownBranchWeights(...), but only sets unknown branch weights in the new instruct...
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:643
iterator_range< early_inc_iterator_impl< detail::IterOfRange< RangeT > > > make_early_inc_range(RangeT &&Range)
Make a range that does early increment to allow mutation of the underlying range without disrupting i...
Definition STLExtras.h:649
RelativeUniformCounterPtr ValuesPtrExpr VTableAddr Value
Definition InstrProf.h:143
IRBuilder(LLVMContext &, FolderTy, InserterTy) -> IRBuilder< FolderTy, InserterTy >
class LLVM_GSL_OWNER SmallVector
Forward declaration of SmallVector so that calculateSmallVectorDefaultInlinedElements can reference s...
LLVM_ABI BasicBlock * SplitBlock(BasicBlock *Old, BasicBlock::iterator SplitPt, DominatorTree *DT, LoopInfo *LI=nullptr, MemorySSAUpdater *MSSAU=nullptr, const Twine &BBName="")
Split the specified block at the specified instruction.
ArrayRef(const T &OneElt) -> ArrayRef< T >
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
iterator_range< df_iterator< T > > depth_first(const T &G)
AnalysisManager< Function > FunctionAnalysisManager
Convenience typedef for the Function analysis manager.
LLVM_ABI void setFittedBranchWeights(Instruction &I, ArrayRef< uint64_t > Weights, bool IsExpected, bool ElideAllZero=false)
Variant of setBranchWeights where the Weights will be fit first to uint32_t by shifting right.
#define N