LLVM 24.0.0git
X86LowerAMXType.cpp
Go to the documentation of this file.
1//===- Target/X86/X86LowerAMXType.cpp - -------------------------*- C++ -*-===//
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 <256 x i32> load/store
10/// <256 x i32> is bitcasted to x86_amx on X86, and AMX instruction set only
11/// provides simple operation on x86_amx. The basic elementwise operation
12/// is not supported by AMX. Since x86_amx is bitcasted from vector <256 x i32>
13/// and only AMX intrinsics can operate on the type, we need transform
14/// load/store <256 x i32> instruction to AMX load/store. If the bitcast can
15/// not be combined with load/store, we transform the bitcast to amx load/store
16/// and <256 x i32> store/load.
17///
18/// If Front End not use O0 but the Mid/Back end use O0, (e.g. "Clang -O2 -S
19/// -emit-llvm t.c" + "llc t.ll") we should make sure the amx data is volatile,
20/// because that is necessary for AMX fast register allocation. (In Fast
21/// registera allocation, register will be allocated before spill/reload, so
22/// there is no additional register for amx to identify the step in spill.)
23/// The volatileTileData() will handle this case.
24/// e.g.
25/// ----------------------------------------------------------
26/// | def %td = ... |
27/// | ... |
28/// | "use %td" |
29/// ----------------------------------------------------------
30/// will transfer to -->
31/// ----------------------------------------------------------
32/// | def %td = ... |
33/// | call void @llvm.x86.tilestored64.internal(mem, %td) |
34/// | ... |
35/// | %td2 = call x86_amx @llvm.x86.tileloadd64.internal(mem)|
36/// | "use %td2" |
37/// ----------------------------------------------------------
38//
39//===----------------------------------------------------------------------===//
40//
41#include "X86.h"
43#include "llvm/ADT/SetVector.h"
46#include "llvm/CodeGen/Passes.h"
49#include "llvm/IR/Analysis.h"
50#include "llvm/IR/DataLayout.h"
51#include "llvm/IR/Function.h"
52#include "llvm/IR/IRBuilder.h"
55#include "llvm/IR/IntrinsicsX86.h"
56#include "llvm/IR/PassManager.h"
59#include "llvm/Pass.h"
63
64#include <map>
65
66using namespace llvm;
67using namespace PatternMatch;
68
69#define DEBUG_TYPE "x86-lower-amx-type"
70
76
77static bool isAMXIntrinsic(Value *I) {
79 if (!II)
80 return false;
81 if (isAMXCast(II))
82 return false;
83 // Check if return type or parameter is x86_amx. If it is x86_amx
84 // the intrinsic must be x86 amx intrinsics.
85 if (II->getType()->isX86_AMXTy())
86 return true;
87 for (Value *V : II->args()) {
88 if (V->getType()->isX86_AMXTy())
89 return true;
90 }
91
92 return false;
93}
94
95static bool containsAMXCode(Function &F) {
96 for (BasicBlock &BB : F)
97 for (Instruction &I : BB)
98 if (I.getType()->isX86_AMXTy())
99 return true;
100 return false;
101}
102
104 Type *Ty) {
105 Function &F = *BB->getParent();
106 const DataLayout &DL = F.getDataLayout();
107
108 LLVMContext &Ctx = Builder.getContext();
109 auto AllocaAlignment = DL.getPrefTypeAlign(Type::getX86_AMXTy(Ctx));
110 unsigned AllocaAS = DL.getAllocaAddrSpace();
111 AllocaInst *AllocaRes =
112 new AllocaInst(Ty, AllocaAS, "", F.getEntryBlock().begin());
113 AllocaRes->setAlignment(AllocaAlignment);
114 return AllocaRes;
115}
116
118 for (Instruction &I : F.getEntryBlock())
119 if (!isa<AllocaInst>(&I))
120 return &I;
121 llvm_unreachable("No terminator in the entry block!");
122}
123
124static Value *getRowFromCol(Instruction *II, Value *V, unsigned Granularity) {
125 IRBuilder<> Builder(II);
126 Value *RealRow = nullptr;
127 if (isa<ConstantInt>(V))
128 RealRow =
129 Builder.getInt16((cast<ConstantInt>(V)->getSExtValue()) / Granularity);
130 else if (isa<Instruction>(V)) {
131 // When it is not a const value and it is not a function argument, we
132 // create Row after the definition of V instead of
133 // before II. For example, II is %118, we try to getshape for %117:
134 // %117 = call x86_amx @llvm.x86.cast.vector.to.tile.v256i32(<256 x
135 // i32> %115).
136 // %118 = call x86_amx @llvm.x86.tdpbf16ps.internal(i16
137 // %104, i16 %105, i16 %106, x86_amx %110, x86_amx %114, x86_amx
138 // %117).
139 // If we create %row = udiv i16 %106, 4 before %118(aka. II), then its
140 // definition is after its user(new tileload for %117).
141 // So, the best choice is to create %row right after the definition of
142 // %106.
143 Builder.SetInsertPoint(cast<Instruction>(V));
144 RealRow = Builder.CreateUDiv(V, Builder.getInt16(4));
145 cast<Instruction>(RealRow)->moveAfter(cast<Instruction>(V));
146 } else {
147 // When it is not a const value and it is a function argument, we create
148 // Row at the entry bb.
149 IRBuilder<> NewBuilder(
150 getFirstNonAllocaInTheEntryBlock(*II->getFunction()));
151 RealRow = NewBuilder.CreateUDiv(V, NewBuilder.getInt16(Granularity));
152 }
153 return RealRow;
154}
155
156// TODO: Refine the row and col-in-bytes of tile to row and col of matrix.
157std::pair<Value *, Value *> getShape(IntrinsicInst *II, unsigned OpNo) {
158 IRBuilder<> Builder(II);
159 Value *Row = nullptr, *Col = nullptr;
160 switch (II->getIntrinsicID()) {
161 default:
162 llvm_unreachable("Expect amx intrinsics");
163 case Intrinsic::x86_tileloadd64_internal:
164 case Intrinsic::x86_tileloaddt164_internal:
165 case Intrinsic::x86_tilestored64_internal:
166 case Intrinsic::x86_tileloaddrs64_internal:
167 case Intrinsic::x86_tileloaddrst164_internal: {
168 Row = II->getArgOperand(0);
169 Col = II->getArgOperand(1);
170 break;
171 }
172 // a * b + c
173 // The shape depends on which operand.
174 case Intrinsic::x86_tcmmimfp16ps_internal:
175 case Intrinsic::x86_tcmmrlfp16ps_internal:
176 case Intrinsic::x86_tdpbssd_internal:
177 case Intrinsic::x86_tdpbsud_internal:
178 case Intrinsic::x86_tdpbusd_internal:
179 case Intrinsic::x86_tdpbuud_internal:
180 case Intrinsic::x86_tdpbf16ps_internal:
181 case Intrinsic::x86_tdpfp16ps_internal:
182 case Intrinsic::x86_tdpbf8ps_internal:
183 case Intrinsic::x86_tdpbhf8ps_internal:
184 case Intrinsic::x86_tdphbf8ps_internal:
185 case Intrinsic::x86_tdphf8ps_internal: {
186 switch (OpNo) {
187 case 3:
188 Row = II->getArgOperand(0);
189 Col = II->getArgOperand(1);
190 break;
191 case 4:
192 Row = II->getArgOperand(0);
193 Col = II->getArgOperand(2);
194 break;
195 case 5:
196 Row = getRowFromCol(II, II->getArgOperand(2), 4);
197 Col = II->getArgOperand(1);
198 break;
199 }
200 break;
201 }
202 case Intrinsic::x86_tcvtrowd2ps_internal:
203 case Intrinsic::x86_tcvtrowps2bf16h_internal:
204 case Intrinsic::x86_tcvtrowps2bf16l_internal:
205 case Intrinsic::x86_tcvtrowps2phh_internal:
206 case Intrinsic::x86_tcvtrowps2phl_internal:
207 case Intrinsic::x86_tilemovrow_internal: {
208 assert(OpNo == 2 && "Illegal Operand Number.");
209 Row = II->getArgOperand(0);
210 Col = II->getArgOperand(1);
211 break;
212 }
213 }
214
215 return std::make_pair(Row, Col);
216}
217
218static std::pair<Value *, Value *> getShape(PHINode *Phi) {
219 Use &U = *(Phi->use_begin());
220 unsigned OpNo = U.getOperandNo();
221 User *V = U.getUser();
222 // TODO We don't traverse all users. To make the algorithm simple, here we
223 // just traverse the first user. If we can find shape, then return the shape,
224 // otherwise just return nullptr and the optimization for undef/zero will be
225 // abandoned.
226 while (V) {
228 if (V->use_empty())
229 break;
230 Use &U = *(V->use_begin());
231 OpNo = U.getOperandNo();
232 V = U.getUser();
233 } else if (isAMXIntrinsic(V)) {
234 return getShape(cast<IntrinsicInst>(V), OpNo);
235 } else if (isa<PHINode>(V)) {
236 if (V->use_empty())
237 break;
238 Use &U = *(V->use_begin());
239 V = U.getUser();
240 } else {
241 break;
242 }
243 }
244
245 return std::make_pair(nullptr, nullptr);
246}
247
248namespace {
249class X86LowerAMXType {
250 Function &Func;
251
252 // In AMX intrinsics we let Shape = {Row, Col}, but the
253 // RealCol = Col / ElementSize. We may use the RealCol
254 // as a new Row for other new created AMX intrinsics.
255 std::map<Value *, Value *> Col2Row;
256
257public:
258 X86LowerAMXType(Function &F) : Func(F) {}
259 bool visit();
260 void combineLoadBitcast(LoadInst *LD, BitCastInst *Bitcast);
261 void combineBitcastStore(BitCastInst *Bitcast, StoreInst *ST);
262 bool transformBitcast(BitCastInst *Bitcast);
263};
264
265// %src = load <256 x i32>, <256 x i32>* %addr, align 64
266// %2 = bitcast <256 x i32> %src to x86_amx
267// -->
268// %2 = call x86_amx @llvm.x86.tileloadd64.internal(i16 %row, i16 %col,
269// i8* %addr, i64 %stride64)
270void X86LowerAMXType::combineLoadBitcast(LoadInst *LD, BitCastInst *Bitcast) {
271 Value *Row = nullptr, *Col = nullptr;
272 Use &U = *(Bitcast->use_begin());
273 unsigned OpNo = U.getOperandNo();
274 auto *II = cast<IntrinsicInst>(U.getUser());
275 std::tie(Row, Col) = getShape(II, OpNo);
276 IRBuilder<> Builder(Bitcast);
277 // Use the maximun column as stride.
278 Value *Stride = Builder.getInt64(64);
279 Value *I8Ptr = LD->getOperand(0);
280 std::array<Value *, 4> Args = {Row, Col, I8Ptr, Stride};
281
282 Value *NewInst =
283 Builder.CreateIntrinsic(Intrinsic::x86_tileloadd64_internal, Args);
284 Bitcast->replaceAllUsesWith(NewInst);
285}
286
287// %src = call x86_amx @llvm.x86.tileloadd64.internal(%row, %col, %addr,
288// %stride);
289// %13 = bitcast x86_amx %src to <256 x i32>
290// store <256 x i32> %13, <256 x i32>* %addr, align 64
291// -->
292// call void @llvm.x86.tilestored64.internal(%row, %col, %addr,
293// %stride64, %13)
294void X86LowerAMXType::combineBitcastStore(BitCastInst *Bitcast, StoreInst *ST) {
295
296 Value *Tile = Bitcast->getOperand(0);
297 auto *II = cast<IntrinsicInst>(Tile);
298 // Tile is output from AMX intrinsic. The first operand of the
299 // intrinsic is row, the second operand of the intrinsic is column.
300 Value *Row = II->getOperand(0);
301 Value *Col = II->getOperand(1);
302 IRBuilder<> Builder(ST);
303 // Use the maximum column as stride. It must be the same with load
304 // stride.
305 Value *Stride = Builder.getInt64(64);
306 Value *I8Ptr = ST->getOperand(1);
307 std::array<Value *, 5> Args = {Row, Col, I8Ptr, Stride, Tile};
308 Builder.CreateIntrinsic(Intrinsic::x86_tilestored64_internal, Args);
309 if (Bitcast->hasOneUse())
310 return;
311 // %13 = bitcast x86_amx %src to <256 x i32>
312 // store <256 x i32> %13, <256 x i32>* %addr, align 64
313 // %add = <256 x i32> %13, <256 x i32> %src2
314 // -->
315 // %13 = bitcast x86_amx %src to <256 x i32>
316 // call void @llvm.x86.tilestored64.internal(%row, %col, %addr,
317 // %stride64, %13)
318 // %14 = load <256 x i32>, %addr
319 // %add = <256 x i32> %14, <256 x i32> %src2
320 Value *Vec = Builder.CreateLoad(Bitcast->getType(), ST->getOperand(1));
321 Bitcast->replaceAllUsesWith(Vec);
322}
323
324// transform bitcast to <store, load> instructions.
325bool X86LowerAMXType::transformBitcast(BitCastInst *Bitcast) {
326 IRBuilder<> Builder(Bitcast);
327 AllocaInst *AllocaAddr;
328 Value *I8Ptr, *Stride;
329 auto *Src = Bitcast->getOperand(0);
330
331 auto Prepare = [&](Type *MemTy) {
332 AllocaAddr = createAllocaInstAtEntry(Builder, Bitcast->getParent(), MemTy);
333 I8Ptr = AllocaAddr;
334 Stride = Builder.getInt64(64);
335 };
336
337 if (Bitcast->getType()->isX86_AMXTy()) {
338 // %2 = bitcast <256 x i32> %src to x86_amx
339 // -->
340 // %addr = alloca <256 x i32>, align 64
341 // store <256 x i32> %src, <256 x i32>* %addr, align 64
342 // %addr2 = bitcast <256 x i32>* to i8*
343 // %2 = call x86_amx @llvm.x86.tileloadd64.internal(i16 %row, i16 %col,
344 // i8* %addr2,
345 // i64 64)
346 Use &U = *(Bitcast->use_begin());
347 unsigned OpNo = U.getOperandNo();
348 auto *II = dyn_cast<IntrinsicInst>(U.getUser());
349 if (!II)
350 return false; // May be bitcast from x86amx to <256 x i32>.
351 Prepare(Bitcast->getOperand(0)->getType());
352 Builder.CreateStore(Src, AllocaAddr);
353 // TODO we can pick an constant operand for the shape.
354 Value *Row = nullptr, *Col = nullptr;
355 std::tie(Row, Col) = getShape(II, OpNo);
356 std::array<Value *, 4> Args = {Row, Col, I8Ptr, Stride};
357 Value *NewInst =
358 Builder.CreateIntrinsic(Intrinsic::x86_tileloadd64_internal, Args);
359 Bitcast->replaceAllUsesWith(NewInst);
360 } else {
361 // %2 = bitcast x86_amx %src to <256 x i32>
362 // -->
363 // %addr = alloca <256 x i32>, align 64
364 // %addr2 = bitcast <256 x i32>* to i8*
365 // call void @llvm.x86.tilestored64.internal(i16 %row, i16 %col,
366 // i8* %addr2, i64 %stride)
367 // %2 = load <256 x i32>, <256 x i32>* %addr, align 64
368 auto *II = dyn_cast<IntrinsicInst>(Src);
369 if (!II)
370 return false; // May be bitcast from <256 x i32> to x86amx.
371 Prepare(Bitcast->getType());
372 Value *Row = II->getOperand(0);
373 Value *Col = II->getOperand(1);
374 std::array<Value *, 5> Args = {Row, Col, I8Ptr, Stride, Src};
375 Builder.CreateIntrinsic(Intrinsic::x86_tilestored64_internal, Args);
376 Value *NewInst = Builder.CreateLoad(Bitcast->getType(), AllocaAddr);
377 Bitcast->replaceAllUsesWith(NewInst);
378 }
379
380 return true;
381}
382
383bool X86LowerAMXType::visit() {
384 SmallVector<Instruction *, 8> DeadInsts;
385 Col2Row.clear();
386
387 for (BasicBlock *BB : post_order(&Func)) {
388 for (Instruction &Inst : llvm::make_early_inc_range(llvm::reverse(*BB))) {
389 auto *Bitcast = dyn_cast<BitCastInst>(&Inst);
390 if (!Bitcast)
391 continue;
392
393 Value *Src = Bitcast->getOperand(0);
394 if (Bitcast->getType()->isX86_AMXTy()) {
395 if (Bitcast->user_empty()) {
396 DeadInsts.push_back(Bitcast);
397 continue;
398 }
399 LoadInst *LD = dyn_cast<LoadInst>(Src);
400 if (!LD) {
401 if (transformBitcast(Bitcast))
402 DeadInsts.push_back(Bitcast);
403 continue;
404 }
405 // If load has multi-user, duplicate a vector load.
406 // %src = load <256 x i32>, <256 x i32>* %addr, align 64
407 // %2 = bitcast <256 x i32> %src to x86_amx
408 // %add = add <256 x i32> %src, <256 x i32> %src2
409 // -->
410 // %src = load <256 x i32>, <256 x i32>* %addr, align 64
411 // %2 = call x86_amx @llvm.x86.tileloadd64.internal(i16 %row, i16 %col,
412 // i8* %addr, i64 %stride64)
413 // %add = add <256 x i32> %src, <256 x i32> %src2
414
415 // If load has one user, the load will be eliminated in DAG ISel.
416 // %src = load <256 x i32>, <256 x i32>* %addr, align 64
417 // %2 = bitcast <256 x i32> %src to x86_amx
418 // -->
419 // %2 = call x86_amx @llvm.x86.tileloadd64.internal(i16 %row, i16 %col,
420 // i8* %addr, i64 %stride64)
421 combineLoadBitcast(LD, Bitcast);
422 DeadInsts.push_back(Bitcast);
423 if (LD->hasOneUse())
424 DeadInsts.push_back(LD);
425 } else if (Src->getType()->isX86_AMXTy()) {
426 if (Bitcast->user_empty()) {
427 DeadInsts.push_back(Bitcast);
428 continue;
429 }
430 StoreInst *ST = nullptr;
431 for (Use &U : Bitcast->uses()) {
432 ST = dyn_cast<StoreInst>(U.getUser());
433 if (ST)
434 break;
435 }
436 if (!ST) {
437 if (transformBitcast(Bitcast))
438 DeadInsts.push_back(Bitcast);
439 continue;
440 }
441 // If bitcast (%13) has one use, combine bitcast and store to amx store.
442 // %src = call x86_amx @llvm.x86.tileloadd64.internal(%row, %col, %addr,
443 // %stride);
444 // %13 = bitcast x86_amx %src to <256 x i32>
445 // store <256 x i32> %13, <256 x i32>* %addr, align 64
446 // -->
447 // call void @llvm.x86.tilestored64.internal(%row, %col, %addr,
448 // %stride64, %13)
449 //
450 // If bitcast (%13) has multi-use, transform as below.
451 // %13 = bitcast x86_amx %src to <256 x i32>
452 // store <256 x i32> %13, <256 x i32>* %addr, align 64
453 // %add = <256 x i32> %13, <256 x i32> %src2
454 // -->
455 // %13 = bitcast x86_amx %src to <256 x i32>
456 // call void @llvm.x86.tilestored64.internal(%row, %col, %addr,
457 // %stride64, %13)
458 // %14 = load <256 x i32>, %addr
459 // %add = <256 x i32> %14, <256 x i32> %src2
460 //
461 combineBitcastStore(Bitcast, ST);
462 // Delete user first.
463 DeadInsts.push_back(ST);
464 DeadInsts.push_back(Bitcast);
465 }
466 }
467 }
468
469 bool C = !DeadInsts.empty();
470
471 for (auto *Inst : DeadInsts)
472 Inst->eraseFromParent();
473
474 return C;
475}
476} // anonymous namespace
477
479 Function *F = BB->getParent();
480 IRBuilder<> Builder(&F->getEntryBlock().front());
481 const DataLayout &DL = F->getDataLayout();
482 unsigned AllocaAS = DL.getAllocaAddrSpace();
483 Type *V256I32Ty = VectorType::get(Builder.getInt32Ty(), 256, false);
484 AllocaInst *AllocaRes =
485 new AllocaInst(V256I32Ty, AllocaAS, "", F->getEntryBlock().begin());
486 BasicBlock::iterator Iter = AllocaRes->getIterator();
487 ++Iter;
488 Builder.SetInsertPoint(&*Iter);
489 Value *I8Ptr = Builder.CreateBitCast(AllocaRes, Builder.getPtrTy());
490 return I8Ptr;
491}
492
494 assert(TileDef->getType()->isX86_AMXTy() && "Not define tile!");
495 auto *II = cast<IntrinsicInst>(TileDef);
496
497 assert(II && "Not tile intrinsic!");
498 Value *Row = II->getOperand(0);
499 Value *Col = II->getOperand(1);
500
501 BasicBlock::iterator Iter = TileDef->getIterator();
502 IRBuilder<> Builder(++Iter);
503 Value *Stride = Builder.getInt64(64);
504 std::array<Value *, 5> Args = {Row, Col, Ptr, Stride, TileDef};
505
506 Instruction *TileStore = Builder.CreateIntrinsicWithoutFolding(
507 Intrinsic::x86_tilestored64_internal, Args);
508 return TileStore;
509}
510
511static void replaceWithTileLoad(Use &U, Value *Ptr, bool IsPHI = false) {
512 Value *V = U.get();
513 assert(V->getType()->isX86_AMXTy() && "Not define tile!");
514
515 // Get tile shape.
516 IntrinsicInst *II = nullptr;
517 if (IsPHI) {
518 Value *PhiOp = cast<PHINode>(V)->getIncomingValue(0);
519 II = cast<IntrinsicInst>(PhiOp);
520 } else {
522 }
523 Value *Row = II->getOperand(0);
524 Value *Col = II->getOperand(1);
525
526 Instruction *UserI = cast<Instruction>(U.getUser());
527 IRBuilder<> Builder(UserI);
528 Value *Stride = Builder.getInt64(64);
529 std::array<Value *, 4> Args = {Row, Col, Ptr, Stride};
530
531 Value *TileLoad =
532 Builder.CreateIntrinsic(Intrinsic::x86_tileloadd64_internal, Args);
533 UserI->replaceUsesOfWith(V, TileLoad);
534}
535
537 for (Use &U : I->uses()) {
538 User *V = U.getUser();
539 if (isa<PHINode>(V))
540 return true;
541 }
542 return false;
543}
544
545// Let all AMX tile data become volatile data, shorten the life range
546// of each tile register before fast register allocation.
547namespace {
548class X86VolatileTileData {
549 Function &F;
550
551public:
552 X86VolatileTileData(Function &Func) : F(Func) {}
553 Value *updatePhiIncomings(BasicBlock *BB,
554 SmallVector<Instruction *, 2> &Incomings);
555 void replacePhiDefWithLoad(Instruction *PHI, Value *StorePtr);
556 bool volatileTileData();
557 void volatileTilePHI(PHINode *PHI);
558 void volatileTileNonPHI(Instruction *I);
559};
560
561Value *X86VolatileTileData::updatePhiIncomings(
562 BasicBlock *BB, SmallVector<Instruction *, 2> &Incomings) {
563 Value *I8Ptr = getAllocaPos(BB);
564
565 for (auto *I : Incomings) {
566 User *Store = createTileStore(I, I8Ptr);
567
568 // All its uses (except phi) should load from stored mem.
569 for (Use &U : I->uses()) {
570 User *V = U.getUser();
571 if (isa<PHINode>(V) || V == Store)
572 continue;
573 replaceWithTileLoad(U, I8Ptr);
574 }
575 }
576 return I8Ptr;
577}
578
579void X86VolatileTileData::replacePhiDefWithLoad(Instruction *PHI,
580 Value *StorePtr) {
581 for (Use &U : PHI->uses())
582 replaceWithTileLoad(U, StorePtr, true);
583 PHI->eraseFromParent();
584}
585
586// Smilar with volatileTileNonPHI, this function only handle PHI Nodes
587// and their related AMX intrinsics.
588// 1) PHI Def should change to tileload.
589// 2) PHI Incoming Values should tilestored in just after their def.
590// 3) The mem of these tileload and tilestores should be same.
591// e.g.
592// ------------------------------------------------------
593// bb_dom:
594// ...
595// br i1 %bool.cond, label %if.else, label %if.then
596//
597// if.then:
598// def %t0 = ...
599// ...
600// use %t0
601// ...
602// br label %if.end
603//
604// if.else:
605// def %t1 = ...
606// br label %if.end
607//
608// if.end:
609// %td = phi x86_amx [ %t1, %if.else ], [ %t0, %if.then ]
610// ...
611// use %td
612// ------------------------------------------------------
613// -->
614// ------------------------------------------------------
615// bb_entry:
616// %mem = alloca <256 x i32>, align 1024 *
617// ...
618// bb_dom:
619// ...
620// br i1 %bool.cond, label %if.else, label %if.then
621//
622// if.then:
623// def %t0 = ...
624// call void @llvm.x86.tilestored64.internal(mem, %t0) *
625// ...
626// %t0` = call x86_amx @llvm.x86.tileloadd64.internal(mem)*
627// use %t0` *
628// ...
629// br label %if.end
630//
631// if.else:
632// def %t1 = ...
633// call void @llvm.x86.tilestored64.internal(mem, %t1) *
634// br label %if.end
635//
636// if.end:
637// ...
638// %td = call x86_amx @llvm.x86.tileloadd64.internal(mem) *
639// use %td
640// ------------------------------------------------------
641void X86VolatileTileData::volatileTilePHI(PHINode *PHI) {
642 BasicBlock *BB = PHI->getParent();
643 SmallVector<Instruction *, 2> Incomings;
644
645 for (unsigned I = 0, E = PHI->getNumIncomingValues(); I != E; ++I) {
646 Value *Op = PHI->getIncomingValue(I);
648 assert(Inst && "We shouldn't fold AMX instrution!");
649 Incomings.push_back(Inst);
650 }
651
652 Value *StorePtr = updatePhiIncomings(BB, Incomings);
653 replacePhiDefWithLoad(PHI, StorePtr);
654}
655
656// Store the defined tile and load it before use.
657// All its users are not PHI.
658// e.g.
659// ------------------------------------------------------
660// def %td = ...
661// ...
662// "use %td"
663// ------------------------------------------------------
664// -->
665// ------------------------------------------------------
666// def %td = ...
667// call void @llvm.x86.tilestored64.internal(mem, %td)
668// ...
669// %td2 = call x86_amx @llvm.x86.tileloadd64.internal(mem)
670// "use %td2"
671// ------------------------------------------------------
672void X86VolatileTileData::volatileTileNonPHI(Instruction *I) {
673 BasicBlock *BB = I->getParent();
674 Value *I8Ptr = getAllocaPos(BB);
675 User *Store = createTileStore(I, I8Ptr);
676
677 // All its uses should load from stored mem.
678 for (Use &U : I->uses()) {
679 User *V = U.getUser();
680 assert(!isa<PHINode>(V) && "PHI Nodes should be excluded!");
681 if (V != Store)
682 replaceWithTileLoad(U, I8Ptr);
683 }
684}
685
686// Volatile Tile Model:
687// 1) All the uses of tile data comes from tileload in time.
688// 2) All the defs of tile data tilestore into mem immediately.
689// For example:
690// --------------------------------------------------------------------------
691// %t1 = call x86_amx @llvm.x86.tileloadd64.internal(m, k, ...) key
692// %t2 = call x86_amx @llvm.x86.tileloadd64.internal(k, n, ...)
693// %t3 = call x86_amx @llvm.x86.tileloadd64.internal(m, n, ...) amx
694// %td = tail call x86_amx @llvm.x86.tdpbssd.internal(m, n, k, t1, t2, t3)
695// call void @llvm.x86.tilestored64.internal(... td) area
696// --------------------------------------------------------------------------
697// 3) No terminator, call or other amx instructions in the key amx area.
698bool X86VolatileTileData::volatileTileData() {
699 bool Changed = false;
700 for (BasicBlock &BB : F) {
701 SmallVector<Instruction *, 2> PHIInsts;
702 SmallVector<Instruction *, 8> AMXDefInsts;
703
704 for (Instruction &I : BB) {
705 if (!I.getType()->isX86_AMXTy())
706 continue;
707 if (isa<PHINode>(&I))
708 PHIInsts.push_back(&I);
709 else
710 AMXDefInsts.push_back(&I);
711 }
712
713 // First we "volatile" the non-phi related amx intrinsics.
714 for (Instruction *I : AMXDefInsts) {
715 if (isIncomingOfPHI(I))
716 continue;
717 volatileTileNonPHI(I);
718 Changed = true;
719 }
720
721 for (Instruction *I : PHIInsts) {
722 volatileTilePHI(dyn_cast<PHINode>(I));
723 Changed = true;
724 }
725 }
726 return Changed;
727}
728
729} // anonymous namespace
730
731namespace {
732
733class X86LowerAMXCast {
734 Function &Func;
735 std::unique_ptr<DominatorTree> DT;
736
737public:
738 X86LowerAMXCast(Function &F) : Func(F), DT(nullptr) {}
739 bool combineCastStore(IntrinsicInst *Cast, StoreInst *ST);
740 bool combineLoadCast(IntrinsicInst *Cast, LoadInst *LD);
741 bool combineTilezero(IntrinsicInst *Cast);
742 bool combineLdSt(SmallVectorImpl<Instruction *> &Casts);
743 bool combineAMXcast(TargetLibraryInfo *TLI);
744 bool transformAMXCast(IntrinsicInst *AMXCast);
745 bool transformAllAMXCast();
746 bool optimizeAMXCastFromPhi(IntrinsicInst *CI, PHINode *PN,
747 SmallSetVector<Instruction *, 16> &DeadInst);
748};
749
750static bool DCEInstruction(Instruction *I,
751 SmallSetVector<Instruction *, 16> &WorkList,
752 const TargetLibraryInfo *TLI) {
753 if (isInstructionTriviallyDead(I, TLI)) {
756
757 // Null out all of the instruction's operands to see if any operand becomes
758 // dead as we go.
759 for (unsigned i = 0, e = I->getNumOperands(); i != e; ++i) {
760 Value *OpV = I->getOperand(i);
761 I->setOperand(i, nullptr);
762
763 if (!OpV->use_empty() || I == OpV)
764 continue;
765
766 // If the operand is an instruction that became dead as we nulled out the
767 // operand, and if it is 'trivially' dead, delete it in a future loop
768 // iteration.
769 if (Instruction *OpI = dyn_cast<Instruction>(OpV)) {
770 if (isInstructionTriviallyDead(OpI, TLI)) {
771 WorkList.insert(OpI);
772 }
773 }
774 }
775 I->eraseFromParent();
776 return true;
777 }
778 return false;
779}
780
781/// This function handles following case
782///
783/// A -> B amxcast
784/// PHI
785/// B -> A amxcast
786///
787/// All the related PHI nodes can be replaced by new PHI nodes with type A.
788/// The uses of \p CI can be changed to the new PHI node corresponding to \p PN.
789bool X86LowerAMXCast::optimizeAMXCastFromPhi(
790 IntrinsicInst *CI, PHINode *PN,
791 SmallSetVector<Instruction *, 16> &DeadInst) {
792 IRBuilder<> Builder(CI);
793 Value *Src = CI->getOperand(0);
794 Type *SrcTy = Src->getType(); // Type B
795 Type *DestTy = CI->getType(); // Type A
796
797 SmallVector<PHINode *, 4> PhiWorklist;
798 SmallSetVector<PHINode *, 4> OldPhiNodes;
799
800 // Find all of the A->B casts and PHI nodes.
801 // We need to inspect all related PHI nodes, but PHIs can be cyclic, so
802 // OldPhiNodes is used to track all known PHI nodes, before adding a new
803 // PHI to PhiWorklist, it is checked against and added to OldPhiNodes first.
804 PhiWorklist.push_back(PN);
805 OldPhiNodes.insert(PN);
806 while (!PhiWorklist.empty()) {
807 auto *OldPN = PhiWorklist.pop_back_val();
808 for (unsigned I = 0; I < OldPN->getNumOperands(); ++I) {
809 Value *IncValue = OldPN->getIncomingValue(I);
810 // TODO: currently, We ignore cases where it is a const. In the future, we
811 // might support const.
812 if (isa<Constant>(IncValue)) {
813 auto *IncConst = dyn_cast<Constant>(IncValue);
814 if (!isa<UndefValue>(IncValue) && !IncConst->isNullValue())
815 return false;
816 Value *Row = nullptr, *Col = nullptr;
817 std::tie(Row, Col) = getShape(OldPN);
818 // TODO: If it is not constant the Row and Col must domoniate tilezero
819 // that we are going to create.
820 if (!Row || !Col || !isa<Constant>(Row) || !isa<Constant>(Col))
821 return false;
822 // Create tilezero at the end of incoming block.
823 auto *Block = OldPN->getIncomingBlock(I);
824 BasicBlock::iterator Iter = Block->getTerminator()->getIterator();
825 Instruction *NewInst = Builder.CreateIntrinsicWithoutFolding(
826 Intrinsic::x86_tilezero_internal, {}, {Row, Col});
827 NewInst->moveBefore(Iter);
828 NewInst = Builder.CreateIntrinsicWithoutFolding(
829 Intrinsic::x86_cast_tile_to_vector, {IncValue->getType()},
830 {NewInst});
831 NewInst->moveBefore(Iter);
832 // Replace InValue with new Value.
833 OldPN->setIncomingValue(I, NewInst);
834 IncValue = NewInst;
835 }
836
837 if (auto *PNode = dyn_cast<PHINode>(IncValue)) {
838 if (OldPhiNodes.insert(PNode))
839 PhiWorklist.push_back(PNode);
840 continue;
841 }
842 Instruction *ACI = dyn_cast<Instruction>(IncValue);
843 if (ACI && isAMXCast(ACI)) {
844 // Verify it's a A->B cast.
845 Type *TyA = ACI->getOperand(0)->getType();
846 Type *TyB = ACI->getType();
847 if (TyA != DestTy || TyB != SrcTy)
848 return false;
849 continue;
850 }
851 return false;
852 }
853 }
854
855 // Check that each user of each old PHI node is something that we can
856 // rewrite, so that all of the old PHI nodes can be cleaned up afterwards.
857 for (auto *OldPN : OldPhiNodes) {
858 for (User *V : OldPN->users()) {
860 if (ACI && isAMXCast(ACI)) {
861 // Verify it's a B->A cast.
862 Type *TyB = ACI->getOperand(0)->getType();
863 Type *TyA = ACI->getType();
864 if (TyA != DestTy || TyB != SrcTy)
865 return false;
866 } else if (auto *PHI = dyn_cast<PHINode>(V)) {
867 // As long as the user is another old PHI node, then even if we don't
868 // rewrite it, the PHI web we're considering won't have any users
869 // outside itself, so it'll be dead.
870 // example:
871 // bb.0:
872 // %0 = amxcast ...
873 // bb.1:
874 // %1 = amxcast ...
875 // bb.2:
876 // %goodphi = phi %0, %1
877 // %3 = amxcast %goodphi
878 // bb.3:
879 // %goodphi2 = phi %0, %goodphi
880 // %4 = amxcast %goodphi2
881 // When optimizeAMXCastFromPhi process %3 and %goodphi, %goodphi2 is
882 // outside the phi-web, so the combination stop When
883 // optimizeAMXCastFromPhi process %4 and %goodphi2, the optimization
884 // will be done.
885 if (OldPhiNodes.count(PHI) == 0)
886 return false;
887 } else
888 return false;
889 }
890 }
891
892 // For each old PHI node, create a corresponding new PHI node with a type A.
893 SmallDenseMap<PHINode *, PHINode *> NewPNodes;
894 for (auto *OldPN : OldPhiNodes) {
895 Builder.SetInsertPoint(OldPN);
896 PHINode *NewPN = Builder.CreatePHI(DestTy, OldPN->getNumOperands());
897 NewPNodes[OldPN] = NewPN;
898 }
899
900 // Fill in the operands of new PHI nodes.
901 for (auto *OldPN : OldPhiNodes) {
902 PHINode *NewPN = NewPNodes[OldPN];
903 for (unsigned j = 0, e = OldPN->getNumOperands(); j != e; ++j) {
904 Value *V = OldPN->getOperand(j);
905 Value *NewV = nullptr;
907 // There should not be a AMXcast from a const.
908 if (ACI && isAMXCast(ACI))
909 NewV = ACI->getOperand(0);
910 else if (auto *PrevPN = dyn_cast<PHINode>(V))
911 NewV = NewPNodes[PrevPN];
912 assert(NewV);
913 NewPN->addIncoming(NewV, OldPN->getIncomingBlock(j));
914 }
915 }
916
917 // Traverse all accumulated PHI nodes and process its users,
918 // which are Stores and BitcCasts. Without this processing
919 // NewPHI nodes could be replicated and could lead to extra
920 // moves generated after DeSSA.
921 // If there is a store with type B, change it to type A.
922
923 // Replace users of BitCast B->A with NewPHI. These will help
924 // later to get rid of a closure formed by OldPHI nodes.
925 for (auto *OldPN : OldPhiNodes) {
926 PHINode *NewPN = NewPNodes[OldPN];
927 for (User *V : make_early_inc_range(OldPN->users())) {
929 if (ACI && isAMXCast(ACI)) {
930 Type *TyB = ACI->getOperand(0)->getType();
931 Type *TyA = ACI->getType();
932 assert(TyA == DestTy && TyB == SrcTy);
933 (void)TyA;
934 (void)TyB;
935 ACI->replaceAllUsesWith(NewPN);
936 DeadInst.insert(ACI);
937 } else if (auto *PHI = dyn_cast<PHINode>(V)) {
938 // We don't need to push PHINode into DeadInst since they are operands
939 // of rootPN DCE can safely delete rootPN's operands if rootPN is dead.
940 assert(OldPhiNodes.contains(PHI));
941 (void)PHI;
942 } else
943 llvm_unreachable("all uses should be handled");
944 }
945 }
946 return true;
947}
948
949// %43 = call <256 x i32> @llvm.x86.cast.tile.to.vector.v256i32(x86_amx %42)
950// store <256 x i32> %43, <256 x i32>* %p, align 64
951// -->
952// call void @llvm.x86.tilestored64.internal(i16 %row, i16 %col, i8* %p,
953// i64 64, x86_amx %42)
954bool X86LowerAMXCast::combineCastStore(IntrinsicInst *Cast, StoreInst *ST) {
955 Value *Tile = Cast->getOperand(0);
956
957 assert(Tile->getType()->isX86_AMXTy() && "Not Tile Operand!");
958
959 // TODO: Specially handle the multi-use case.
960 if (!Tile->hasOneUse())
961 return false;
962
963 auto *II = cast<IntrinsicInst>(Tile);
964 // Tile is output from AMX intrinsic. The first operand of the
965 // intrinsic is row, the second operand of the intrinsic is column.
966 Value *Row = II->getOperand(0);
967 Value *Col = II->getOperand(1);
968
969 IRBuilder<> Builder(ST);
970
971 // Stride should be equal to col(measured by bytes)
972 Value *Stride = Builder.CreateSExt(Col, Builder.getInt64Ty());
973 Value *I8Ptr = Builder.CreateBitCast(ST->getOperand(1), Builder.getPtrTy());
974 std::array<Value *, 5> Args = {Row, Col, I8Ptr, Stride, Tile};
975 Builder.CreateIntrinsic(Intrinsic::x86_tilestored64_internal, Args);
976 return true;
977}
978
979// %65 = load <256 x i32>, <256 x i32>* %p, align 64
980// %66 = call x86_amx @llvm.x86.cast.vector.to.tile(<256 x i32> %65)
981// -->
982// %66 = call x86_amx @llvm.x86.tileloadd64.internal(i16 %row, i16 %col,
983// i8* %p, i64 64)
984bool X86LowerAMXCast::combineLoadCast(IntrinsicInst *Cast, LoadInst *LD) {
985 bool EraseLoad = true;
986 Value *Row = nullptr, *Col = nullptr;
987 Use &U = *(Cast->use_begin());
988 unsigned OpNo = U.getOperandNo();
989 auto *II = cast<IntrinsicInst>(U.getUser());
990 // TODO: If it is cast intrinsic or phi node, we can propagate the
991 // shape information through def-use chain.
992 if (!isAMXIntrinsic(II))
993 return false;
994 std::tie(Row, Col) = getShape(II, OpNo);
995 IRBuilder<> Builder(LD);
996 Value *I8Ptr;
997
998 // To save compiling time, we create dominator tree when it is really needed.
999 if (!DT)
1000 DT.reset(new DominatorTree(Func));
1001 if (!DT->dominates(Row, LD) || !DT->dominates(Col, LD)) {
1002 // store the value to stack and reload it from stack before cast.
1003 auto *AllocaAddr =
1004 createAllocaInstAtEntry(Builder, Cast->getParent(), LD->getType());
1005 Builder.SetInsertPoint(&*std::next(LD->getIterator()));
1006 Builder.CreateStore(LD, AllocaAddr);
1007
1008 Builder.SetInsertPoint(Cast);
1009 I8Ptr = Builder.CreateBitCast(AllocaAddr, Builder.getPtrTy());
1010 EraseLoad = false;
1011 } else {
1012 I8Ptr = Builder.CreateBitCast(LD->getOperand(0), Builder.getPtrTy());
1013 }
1014 // Stride should be equal to col(measured by bytes)
1015 Value *Stride = Builder.CreateSExt(Col, Builder.getInt64Ty());
1016 std::array<Value *, 4> Args = {Row, Col, I8Ptr, Stride};
1017
1018 Value *NewInst =
1019 Builder.CreateIntrinsic(Intrinsic::x86_tileloadd64_internal, Args);
1020 Cast->replaceAllUsesWith(NewInst);
1021
1022 return EraseLoad;
1023}
1024
1025// %19 = tail call x86_amx @llvm.x86.cast.vector.to.tile.v256i32(<256 x i32> zeroinitializer)
1026// -->
1027// %19 = tail call x86_amx @llvm.x86.tilezero.internal(i16 %row, i16 %col)
1028bool X86LowerAMXCast::combineTilezero(IntrinsicInst *Cast) {
1029 Value *Row = nullptr, *Col = nullptr;
1030 Use &U = *(Cast->use_begin());
1031 unsigned OpNo = U.getOperandNo();
1032 auto *II = cast<IntrinsicInst>(U.getUser());
1033 if (!isAMXIntrinsic(II))
1034 return false;
1035
1036 std::tie(Row, Col) = getShape(II, OpNo);
1037
1038 IRBuilder<> Builder(Cast);
1039 Value *NewInst =
1040 Builder.CreateIntrinsic(Intrinsic::x86_tilezero_internal, {}, {Row, Col});
1041 Cast->replaceAllUsesWith(NewInst);
1042 return true;
1043}
1044
1045bool X86LowerAMXCast::combineLdSt(SmallVectorImpl<Instruction *> &Casts) {
1046 bool Change = false;
1047 for (auto *Cast : Casts) {
1048 auto *II = cast<IntrinsicInst>(Cast);
1049 // %43 = call <256 x i32> @llvm.x86.cast.tile.to.vector(x86_amx %42)
1050 // store <256 x i32> %43, <256 x i32>* %p, align 64
1051 // -->
1052 // call void @llvm.x86.tilestored64.internal(i16 %row, i16 %col, i8* %p,
1053 // i64 64, x86_amx %42)
1054 if (II->getIntrinsicID() == Intrinsic::x86_cast_tile_to_vector) {
1055 SmallVector<Instruction *, 2> DeadStores;
1056 for (User *U : Cast->users()) {
1057 StoreInst *Store = dyn_cast<StoreInst>(U);
1058 if (!Store)
1059 continue;
1060 if (combineCastStore(cast<IntrinsicInst>(Cast), Store)) {
1061 DeadStores.push_back(Store);
1062 Change = true;
1063 }
1064 }
1065 for (auto *Store : DeadStores)
1066 Store->eraseFromParent();
1067 } else { // x86_cast_vector_to_tile
1068 // %19 = tail call x86_amx @llvm.x86.cast.vector.to.tile.v256i32(<256 x i32> zeroinitializer)
1069 // -->
1070 // %19 = tail call x86_amx @llvm.x86.tilezero.internal(i16 %row, i16 %col)
1071 if (isa<ConstantAggregateZero>(Cast->getOperand(0))) {
1072 Change |= combineTilezero(cast<IntrinsicInst>(Cast));
1073 continue;
1074 }
1075
1076 auto *Load = dyn_cast<LoadInst>(Cast->getOperand(0));
1077 if (!Load || !Load->hasOneUse())
1078 continue;
1079 // %65 = load <256 x i32>, <256 x i32>* %p, align 64
1080 // %66 = call x86_amx @llvm.x86.cast.vector.to.tile(<256 x i32> %65)
1081 // -->
1082 // %66 = call x86_amx @llvm.x86.tileloadd64.internal(i16 %row, i16 %col,
1083 // i8* %p, i64 64)
1084 if (combineLoadCast(cast<IntrinsicInst>(Cast), Load)) {
1085 // Set the operand is null so that load instruction can be erased.
1086 Cast->setOperand(0, nullptr);
1087 Load->eraseFromParent();
1088 Change = true;
1089 }
1090 }
1091 }
1092 return Change;
1093}
1094
1095bool X86LowerAMXCast::combineAMXcast(TargetLibraryInfo *TLI) {
1096 bool Change = false;
1097 // Collect tile cast instruction.
1098 SmallVector<Instruction *, 8> Vec2TileInsts;
1099 SmallVector<Instruction *, 8> Tile2VecInsts;
1100 SmallVector<Instruction *, 8> PhiCastWorkList;
1101 SmallSetVector<Instruction *, 16> DeadInst;
1102 for (BasicBlock &BB : Func) {
1103 for (Instruction &I : BB) {
1104 Value *Vec;
1105 if (match(&I,
1107 Vec2TileInsts.push_back(&I);
1109 m_Value(Vec))))
1110 Tile2VecInsts.push_back(&I);
1111 }
1112 }
1113
1114 auto Convert = [&](SmallVectorImpl<Instruction *> &Insts, Intrinsic::ID IID) {
1115 for (auto *Inst : Insts) {
1116 for (User *U : Inst->users()) {
1117 IntrinsicInst *II = dyn_cast<IntrinsicInst>(U);
1118 if (!II || II->getIntrinsicID() != IID)
1119 continue;
1120 // T1 = vec2tile V0
1121 // V2 = tile2vec T1
1122 // V3 = OP V2
1123 // -->
1124 // T1 = vec2tile V0
1125 // V2 = tile2vec T1
1126 // V3 = OP V0
1127 II->replaceAllUsesWith(Inst->getOperand(0));
1128 Change = true;
1129 }
1130 }
1131 };
1132
1133 Convert(Vec2TileInsts, Intrinsic::x86_cast_tile_to_vector);
1134 Convert(Tile2VecInsts, Intrinsic::x86_cast_vector_to_tile);
1135
1136 SmallVector<Instruction *, 8> LiveCasts;
1137 auto EraseInst = [&](SmallVectorImpl<Instruction *> &Insts) {
1138 for (auto *Inst : Insts) {
1139 if (Inst->use_empty()) {
1140 Inst->eraseFromParent();
1141 Change = true;
1142 } else {
1143 LiveCasts.push_back(Inst);
1144 }
1145 }
1146 };
1147
1148 EraseInst(Vec2TileInsts);
1149 EraseInst(Tile2VecInsts);
1150 LLVM_DEBUG(dbgs() << "[LowerAMXTYpe][combineAMXcast] IR dump after combine "
1151 "Vec2Tile and Tile2Vec:\n";
1152 Func.dump());
1153 Change |= combineLdSt(LiveCasts);
1154 EraseInst(LiveCasts);
1155 LLVM_DEBUG(dbgs() << "[LowerAMXTYpe][combineAMXcast] IR dump after combine "
1156 "AMXCast and load/store:\n";
1157 Func.dump());
1158
1159 // Handle the A->B->A cast, and there is an intervening PHI node.
1160 for (BasicBlock &BB : Func) {
1161 for (Instruction &I : BB) {
1162 if (isAMXCast(&I)) {
1163 if (isa<PHINode>(I.getOperand(0)))
1164 PhiCastWorkList.push_back(&I);
1165 }
1166 }
1167 }
1168 for (auto *I : PhiCastWorkList) {
1169 // We skip the dead Amxcast.
1170 if (DeadInst.contains(I))
1171 continue;
1172 PHINode *PN = cast<PHINode>(I->getOperand(0));
1173 if (optimizeAMXCastFromPhi(cast<IntrinsicInst>(I), PN, DeadInst)) {
1174 DeadInst.insert(PN);
1175 Change = true;
1176 }
1177 }
1178
1179 // Since we create new phi and merge AMXCast, some old phis and AMXCast might
1180 // have no uses. We do some DeadCodeElimination for them.
1181 while (!DeadInst.empty()) {
1182 Instruction *I = DeadInst.pop_back_val();
1183 Change |= DCEInstruction(I, DeadInst, TLI);
1184 }
1185 LLVM_DEBUG(dbgs() << "[LowerAMXTYpe][combineAMXcast] IR dump after "
1186 "optimizeAMXCastFromPhi:\n";
1187 Func.dump());
1188 return Change;
1189}
1190
1191// There might be remaining AMXcast after combineAMXcast and they should be
1192// handled elegantly.
1193bool X86LowerAMXCast::transformAMXCast(IntrinsicInst *AMXCast) {
1194 IRBuilder<> Builder(AMXCast);
1195 AllocaInst *AllocaAddr;
1196 Value *I8Ptr, *Stride;
1197 auto *Src = AMXCast->getOperand(0);
1198
1199 auto Prepare = [&](Type *MemTy) {
1200 AllocaAddr = createAllocaInstAtEntry(Builder, AMXCast->getParent(), MemTy);
1201 I8Ptr = Builder.CreateBitCast(AllocaAddr, Builder.getPtrTy());
1202 Stride = Builder.getInt64(64);
1203 };
1204
1205 if (AMXCast->getType()->isX86_AMXTy()) {
1206 // %2 = amxcast <225 x i32> %src to x86_amx
1207 // call void @llvm.x86.tilestored64.internal(i16 15, i16 60,
1208 // i8* %addr3, i64 60, x86_amx %2)
1209 // -->
1210 // %addr = alloca <225 x i32>, align 64
1211 // store <225 x i32> %src, <225 x i32>* %addr, align 64
1212 // %addr2 = bitcast <225 x i32>* %addr to i8*
1213 // %2 = call x86_amx @llvm.x86.tileloadd64.internal(i16 15, i16 60,
1214 // i8* %addr2,
1215 // i64 60)
1216 // call void @llvm.x86.tilestored64.internal(i16 15, i16 60,
1217 // i8* %addr3, i64 60, x86_amx %2)
1218 if (AMXCast->use_empty()) {
1219 AMXCast->eraseFromParent();
1220 return true;
1221 }
1222 Use &U = *(AMXCast->use_begin());
1223 unsigned OpNo = U.getOperandNo();
1224 auto *II = dyn_cast<IntrinsicInst>(U.getUser());
1225 if (!II)
1226 return false; // May be bitcast from x86amx to <256 x i32>.
1227 Prepare(AMXCast->getOperand(0)->getType());
1228 Builder.CreateStore(Src, AllocaAddr);
1229 // TODO we can pick an constant operand for the shape.
1230 Value *Row = nullptr, *Col = nullptr;
1231 std::tie(Row, Col) = getShape(II, OpNo);
1232 std::array<Value *, 4> Args = {
1233 Row, Col, I8Ptr, Builder.CreateSExt(Col, Builder.getInt64Ty())};
1234 Value *NewInst =
1235 Builder.CreateIntrinsic(Intrinsic::x86_tileloadd64_internal, Args);
1236 AMXCast->replaceAllUsesWith(NewInst);
1237 AMXCast->eraseFromParent();
1238 } else {
1239 // %2 = amxcast x86_amx %src to <225 x i32>
1240 // -->
1241 // %addr = alloca <225 x i32>, align 64
1242 // %addr2 = bitcast <225 x i32>* to i8*
1243 // call void @llvm.x86.tilestored64.internal(i16 %row, i16 %col,
1244 // i8* %addr2, i64 %stride)
1245 // %2 = load <225 x i32>, <225 x i32>* %addr, align 64
1246 auto *II = dyn_cast<IntrinsicInst>(Src);
1247 if (!II)
1248 return false; // May be bitcast from <256 x i32> to x86amx.
1249 Prepare(AMXCast->getType());
1250 Value *Row = II->getOperand(0);
1251 Value *Col = II->getOperand(1);
1252 std::array<Value *, 5> Args = {
1253 Row, Col, I8Ptr, Builder.CreateSExt(Col, Builder.getInt64Ty()), Src};
1254 Builder.CreateIntrinsic(Intrinsic::x86_tilestored64_internal, Args);
1255 Value *NewInst = Builder.CreateLoad(AMXCast->getType(), AllocaAddr);
1256 AMXCast->replaceAllUsesWith(NewInst);
1257 AMXCast->eraseFromParent();
1258 }
1259
1260 return true;
1261}
1262
1263bool X86LowerAMXCast::transformAllAMXCast() {
1264 bool Change = false;
1265 // Collect tile cast instruction.
1266 SmallVector<Instruction *, 8> WorkLists;
1267 for (BasicBlock &BB : Func) {
1268 for (Instruction &I : BB) {
1269 if (isAMXCast(&I))
1270 WorkLists.push_back(&I);
1271 }
1272 }
1273
1274 for (auto *Inst : WorkLists) {
1275 Change |= transformAMXCast(cast<IntrinsicInst>(Inst));
1276 }
1277
1278 return Change;
1279}
1280
1281bool lowerAmxType(Function &F, const TargetMachine *TM,
1282 TargetLibraryInfo *TLI) {
1283 // Performance optimization: most code doesn't use AMX, so return early if
1284 // there are no instructions that produce AMX values. This is sufficient, as
1285 // AMX arguments and constants are not allowed -- so any producer of an AMX
1286 // value must be an instruction.
1287 // TODO: find a cheaper way for this, without looking at all instructions.
1288 if (!containsAMXCode(F))
1289 return false;
1290
1291 bool C = false;
1292 X86LowerAMXCast LAC(F);
1293 C |= LAC.combineAMXcast(TLI);
1294 // There might be remaining AMXcast after combineAMXcast and they should be
1295 // handled elegantly.
1296 C |= LAC.transformAllAMXCast();
1297
1298 X86LowerAMXType LAT(F);
1299 C |= LAT.visit();
1300
1301 // Prepare for fast register allocation at O0.
1302 // Todo: May better check the volatile model of AMX code, not just
1303 // by checking Attribute::OptimizeNone and CodeGenOptLevel::None.
1304 if (TM->getOptLevel() == CodeGenOptLevel::None) {
1305 // If Front End not use O0 but the Mid/Back end use O0, (e.g.
1306 // "Clang -O2 -S -emit-llvm t.c" + "llc t.ll") we should make
1307 // sure the amx data is volatile, that is necessary for AMX fast
1308 // register allocation.
1309 if (!F.hasFnAttribute(Attribute::OptimizeNone)) {
1310 X86VolatileTileData VTD(F);
1311 C = VTD.volatileTileData() || C;
1312 }
1313 }
1314
1315 return C;
1316}
1317
1318} // anonymous namespace
1319
1322 TargetLibraryInfo &TLI = FAM.getResult<TargetLibraryAnalysis>(F);
1323 bool Changed = lowerAmxType(F, TM, &TLI);
1324 if (!Changed)
1325 return PreservedAnalyses::all();
1326
1329 return PA;
1330}
1331
1332namespace {
1333
1334class X86LowerAMXTypeLegacyPass : public FunctionPass {
1335public:
1336 static char ID;
1337
1338 X86LowerAMXTypeLegacyPass() : FunctionPass(ID) {}
1339
1340 bool runOnFunction(Function &F) override {
1341 TargetMachine *TM = &getAnalysis<TargetPassConfig>().getTM<TargetMachine>();
1342 TargetLibraryInfo *TLI =
1343 &getAnalysis<TargetLibraryInfoWrapperPass>().getTLI(F);
1344 return lowerAmxType(F, TM, TLI);
1345 }
1346
1347 void getAnalysisUsage(AnalysisUsage &AU) const override {
1348 AU.setPreservesCFG();
1349 AU.addRequired<TargetPassConfig>();
1350 AU.addRequired<TargetLibraryInfoWrapperPass>();
1351 }
1352};
1353
1354} // anonymous namespace
1355
1356static const char PassName[] = "Lower AMX type for load/store";
1357char X86LowerAMXTypeLegacyPass::ID = 0;
1358INITIALIZE_PASS_BEGIN(X86LowerAMXTypeLegacyPass, DEBUG_TYPE, PassName, false,
1359 false)
1362INITIALIZE_PASS_END(X86LowerAMXTypeLegacyPass, DEBUG_TYPE, PassName, false,
1363 false)
1364
1366 return new X86LowerAMXTypeLegacyPass();
1367}
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
Rewrite undef for PHI
MachineBasicBlock MachineBasicBlock::iterator DebugLoc DL
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
static GCRegistry::Add< CoreCLRGC > E("coreclr", "CoreCLR-compatible GC")
static bool DCEInstruction(Instruction *I, SmallSetVector< Instruction *, 16 > &WorkList, const TargetLibraryInfo *TLI)
Definition DCE.cpp:55
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 builds on the ADT/GraphTraits.h file to build a generic graph post order iterator.
static void visit(BasicBlock &Start, std::function< bool(BasicBlock *)> op)
This file implements a set that has insertion order iteration characteristics.
#define LLVM_DEBUG(...)
Definition Debug.h:119
Target-Independent Code Generator Pass Configuration Options pass.
This pass exposes codegen information to IR-level passes.
static ShapeT getShape(MachineRegisterInfo *MRI, Register TileReg)
static const char PassName[]
static bool isAMXCast(Instruction *II)
static Value * getRowFromCol(Instruction *II, Value *V, unsigned Granularity)
static void replaceWithTileLoad(Use &U, Value *Ptr, bool IsPHI=false)
static Instruction * createTileStore(Instruction *TileDef, Value *Ptr)
static Value * getAllocaPos(BasicBlock *BB)
static bool containsAMXCode(Function &F)
std::pair< Value *, Value * > getShape(IntrinsicInst *II, unsigned OpNo)
static bool isIncomingOfPHI(Instruction *I)
static bool isAMXIntrinsic(Value *I)
static Instruction * getFirstNonAllocaInTheEntryBlock(Function &F)
static AllocaInst * createAllocaInstAtEntry(IRBuilder<> &Builder, BasicBlock *BB, Type *Ty)
an instruction to allocate memory on the stack
void setAlignment(Align Align)
AnalysisUsage & addRequired()
LLVM_ABI void setPreservesCFG()
This function should be called by the pass, iff they do not:
Definition Pass.cpp:278
LLVM Basic Block Representation.
Definition BasicBlock.h:62
const Function * getParent() const
Return the enclosing method, or null if none.
Definition BasicBlock.h:213
InstListType::iterator iterator
Instruction iterators...
Definition BasicBlock.h:170
This class represents a no-op cast from one type to another.
Represents analyses that only rely on functions' control flow.
Definition Analysis.h:73
A parsed version of the target data layout string in and methods for querying it.
Definition DataLayout.h:64
FunctionPass class - This class is used to implement most global optimizations.
Definition Pass.h:314
Value * CreateUDiv(Value *LHS, Value *RHS, const Twine &Name="", bool isExact=false)
Definition IRBuilder.h:1468
ConstantInt * getInt16(uint16_t C)
Get a constant 16-bit value.
Definition IRBuilder.h:459
This provides a uniform API for creating instructions and inserting them into a basic block: either a...
Definition IRBuilder.h:2908
LLVM_ABI void moveBefore(InstListType::iterator InsertPos)
Unlink this instruction from its current basic block and insert it into the basic block that MovePos ...
LLVM_ABI InstListType::iterator eraseFromParent()
This method unlinks 'this' from the containing basic block and deletes it.
iterator_range< user_iterator > users()
A wrapper class for inspecting calls to intrinsic functions.
This is an important class for using LLVM in a threaded context.
Definition LLVMContext.h:68
An instruction for reading from memory.
void addIncoming(Value *V, BasicBlock *BB)
Add an incoming value to the end of the PHI list.
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 & preserveSet()
Mark an analysis set as preserved.
Definition Analysis.h:151
bool contains(const_arg_type key) const
Check if the SetVector contains the given key.
Definition SetVector.h:258
bool empty() const
Determine if the SetVector is empty or not.
Definition SetVector.h:100
bool insert(const value_type &X)
Insert a new element into the SetVector.
Definition SetVector.h:157
value_type pop_back_val()
Definition SetVector.h:285
void push_back(const T &Elt)
Analysis pass providing the TargetLibraryInfo.
Provides information about what library functions are available for the current target.
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
static LLVM_ABI Type * getX86_AMXTy(LLVMContext &C)
Definition Type.cpp:283
bool isX86_AMXTy() const
Return true if this is X86 AMX.
Definition Type.h:202
A Use represents the edge between a Value definition and its users.
Definition Use.h:35
LLVM_ABI unsigned getOperandNo() const
Return the operand # of this use in its User.
Definition Use.cpp:35
User * getUser() const
Returns the User that contains this Use.
Definition Use.h:61
void setOperand(unsigned i, Value *Val)
Definition User.h:212
LLVM_ABI bool replaceUsesOfWith(Value *From, Value *To)
Replace uses of one Value with another.
Definition User.cpp:25
Value * getOperand(unsigned i) const
Definition User.h:207
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
use_iterator use_begin()
Definition Value.h:366
bool use_empty() const
Definition Value.h:348
static LLVM_ABI VectorType * get(Type *ElementType, ElementCount EC)
This static method is the primary way to construct an VectorType.
PreservedAnalyses run(Function &F, FunctionAnalysisManager &FAM)
const ParentTy * getParent() const
Definition ilist_node.h:34
self_iterator getIterator()
Definition ilist_node.h:123
Changed
Pass manager infrastructure for declaring and invalidating analyses.
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
constexpr char Args[]
Key for Kernel::Metadata::mArgs.
@ BasicBlock
Various leaf nodes.
Definition ISDOpcodes.h:83
@ Bitcast
Perform the operation on a different, but equivalently sized type.
bool match(Val *V, const Pattern &P)
auto m_Value()
Match an arbitrary value and ignore it.
auto m_Intrinsic(const Ts &...Ops)
Match intrinsic calls like this: m_Intrinsic<Intrinsic::fabs>(m_Value(X))
@ User
could "use" a pointer
NodeAddr< UseNode * > Use
Definition RDFGraph.h:385
NodeAddr< FuncNode * > Func
Definition RDFGraph.h:393
friend class Instruction
Iterator for Instructions in a `BasicBlock.
Definition BasicBlock.h:73
This is an optimization pass for GlobalISel generic memory operations.
decltype(auto) dyn_cast(const From &Val)
dyn_cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:643
LLVM_ABI void salvageDebugInfo(const MachineRegisterInfo &MRI, MachineInstr &MI)
Assuming the instruction MI is going to be deleted, attempt to salvage debug users of MI by writing t...
Definition Utils.cpp:1676
@ Load
The value being inserted comes from a load (InsertElement only).
@ Store
The extracted value is stored (ExtractElement only).
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
LLVM_ABI bool isInstructionTriviallyDead(Instruction *I, const TargetLibraryInfo *TLI=nullptr)
Return true if the result produced by the instruction is not used, and the instruction will return.
Definition Local.cpp:402
auto reverse(ContainerTy &&C)
Definition STLExtras.h:408
LLVM_ABI raw_ostream & dbgs()
dbgs() - This returns a reference to a raw_ostream for debugging messages.
Definition Debug.cpp:209
IRBuilder(LLVMContext &, FolderTy, InserterTy) -> IRBuilder< FolderTy, InserterTy >
class LLVM_GSL_OWNER SmallVector
Forward declaration of SmallVector so that calculateSmallVectorDefaultInlinedElements can reference s...
auto post_order(const T &G)
Post-order traversal of a graph.
bool isa(const From &Val)
isa<X> - Return true if the parameter to the template is an instance of one of the template type argu...
Definition Casting.h:547
LLVM_ABI bool salvageKnowledge(Instruction *I, AssumptionCache *AC=nullptr, DominatorTree *DT=nullptr)
Calls BuildAssumeFromInst and if the resulting llvm.assume is valid insert if before I.
DWARFExpression::Operation Op
FunctionPass * createX86LowerAMXTypeLegacyPass()
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
AnalysisManager< Function > FunctionAnalysisManager
Convenience typedef for the Function analysis manager.