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"
42#include "llvm/ADT/PostOrderIterator.h"
43#include "llvm/ADT/SetVector.h"
44#include "llvm/Analysis/TargetLibraryInfo.h"
45#include "llvm/Analysis/TargetTransformInfo.h"
46#include "llvm/CodeGen/Passes.h"
47#include "llvm/CodeGen/TargetPassConfig.h"
48#include "llvm/CodeGen/ValueTypes.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"
53#include "llvm/IR/Instructions.h"
54#include "llvm/IR/IntrinsicInst.h"
55#include "llvm/IR/IntrinsicsX86.h"
56#include "llvm/IR/PassManager.h"
57#include "llvm/IR/PatternMatch.h"
58#include "llvm/InitializePasses.h"
59#include "llvm/Pass.h"
60#include "llvm/Target/TargetMachine.h"
61#include "llvm/Transforms/Utils/AssumeBundleBuilder.h"
62#include "llvm/Transforms/Utils/Local.h"
63
64#include <map>
65
66using namespace llvm;
67using namespace PatternMatch;
68
69#define DEBUG_TYPE "x86-lower-amx-type"
70
71static bool isAMXCast(Instruction *II) {
72 return match(V: II,
73 P: m_Intrinsic<Intrinsic::x86_cast_vector_to_tile>(Ops: m_Value())) ||
74 match(V: II, P: m_Intrinsic<Intrinsic::x86_cast_tile_to_vector>(Ops: m_Value()));
75}
76
77static bool isAMXIntrinsic(Value *I) {
78 auto *II = dyn_cast<IntrinsicInst>(Val: 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
103static AllocaInst *createAllocaInstAtEntry(IRBuilder<> &Builder, BasicBlock *BB,
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(Ty: Type::getX86_AMXTy(C&: 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
117static Instruction *getFirstNonAllocaInTheEntryBlock(Function &F) {
118 for (Instruction &I : F.getEntryBlock())
119 if (!isa<AllocaInst>(Val: &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>(Val: V))
128 RealRow =
129 Builder.getInt16(C: (cast<ConstantInt>(Val: V)->getSExtValue()) / Granularity);
130 else if (isa<Instruction>(Val: 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>(Val: V));
144 RealRow = Builder.CreateUDiv(LHS: V, RHS: Builder.getInt16(C: 4));
145 cast<Instruction>(Val: RealRow)->moveAfter(MovePos: cast<Instruction>(Val: 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(F&: *II->getFunction()));
151 RealRow = NewBuilder.CreateUDiv(LHS: V, RHS: NewBuilder.getInt16(C: 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(i: 0);
169 Col = II->getArgOperand(i: 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(i: 0);
189 Col = II->getArgOperand(i: 1);
190 break;
191 case 4:
192 Row = II->getArgOperand(i: 0);
193 Col = II->getArgOperand(i: 2);
194 break;
195 case 5:
196 Row = getRowFromCol(II, V: II->getArgOperand(i: 2), Granularity: 4);
197 Col = II->getArgOperand(i: 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(i: 0);
210 Col = II->getArgOperand(i: 1);
211 break;
212 }
213 }
214
215 return std::make_pair(x&: Row, y&: 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) {
227 if (isAMXCast(II: dyn_cast<Instruction>(Val: 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(I: V)) {
234 return getShape(II: cast<IntrinsicInst>(Val: V), OpNo);
235 } else if (isa<PHINode>(Val: 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(x: nullptr, y: 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>(Val: U.getUser());
275 std::tie(args&: Row, args&: Col) = getShape(II, OpNo);
276 IRBuilder<> Builder(Bitcast);
277 // Use the maximun column as stride.
278 Value *Stride = Builder.getInt64(C: 64);
279 Value *I8Ptr = LD->getOperand(i_nocapture: 0);
280 std::array<Value *, 4> Args = {Row, Col, I8Ptr, Stride};
281
282 Value *NewInst =
283 Builder.CreateIntrinsic(ID: Intrinsic::x86_tileloadd64_internal, Args);
284 Bitcast->replaceAllUsesWith(V: 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(i_nocapture: 0);
297 auto *II = cast<IntrinsicInst>(Val: 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(i_nocapture: 0);
301 Value *Col = II->getOperand(i_nocapture: 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(C: 64);
306 Value *I8Ptr = ST->getOperand(i_nocapture: 1);
307 std::array<Value *, 5> Args = {Row, Col, I8Ptr, Stride, Tile};
308 Builder.CreateIntrinsic(ID: 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(Ty: Bitcast->getType(), Ptr: ST->getOperand(i_nocapture: 1));
321 Bitcast->replaceAllUsesWith(V: 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(i_nocapture: 0);
330
331 auto Prepare = [&](Type *MemTy) {
332 AllocaAddr = createAllocaInstAtEntry(Builder, BB: Bitcast->getParent(), Ty: MemTy);
333 I8Ptr = AllocaAddr;
334 Stride = Builder.getInt64(C: 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>(Val: U.getUser());
349 if (!II)
350 return false; // May be bitcast from x86amx to <256 x i32>.
351 Prepare(Bitcast->getOperand(i_nocapture: 0)->getType());
352 Builder.CreateStore(Val: Src, Ptr: AllocaAddr);
353 // TODO we can pick an constant operand for the shape.
354 Value *Row = nullptr, *Col = nullptr;
355 std::tie(args&: Row, args&: Col) = getShape(II, OpNo);
356 std::array<Value *, 4> Args = {Row, Col, I8Ptr, Stride};
357 Value *NewInst =
358 Builder.CreateIntrinsic(ID: Intrinsic::x86_tileloadd64_internal, Args);
359 Bitcast->replaceAllUsesWith(V: 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>(Val: 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(i_nocapture: 0);
373 Value *Col = II->getOperand(i_nocapture: 1);
374 std::array<Value *, 5> Args = {Row, Col, I8Ptr, Stride, Src};
375 Builder.CreateIntrinsic(ID: Intrinsic::x86_tilestored64_internal, Args);
376 Value *NewInst = Builder.CreateLoad(Ty: Bitcast->getType(), Ptr: AllocaAddr);
377 Bitcast->replaceAllUsesWith(V: 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(G: &Func)) {
388 for (Instruction &Inst : llvm::make_early_inc_range(Range: llvm::reverse(C&: *BB))) {
389 auto *Bitcast = dyn_cast<BitCastInst>(Val: &Inst);
390 if (!Bitcast)
391 continue;
392
393 Value *Src = Bitcast->getOperand(i_nocapture: 0);
394 if (Bitcast->getType()->isX86_AMXTy()) {
395 if (Bitcast->user_empty()) {
396 DeadInsts.push_back(Elt: Bitcast);
397 continue;
398 }
399 LoadInst *LD = dyn_cast<LoadInst>(Val: Src);
400 if (!LD) {
401 if (transformBitcast(Bitcast))
402 DeadInsts.push_back(Elt: 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(Elt: Bitcast);
423 if (LD->hasOneUse())
424 DeadInsts.push_back(Elt: LD);
425 } else if (Src->getType()->isX86_AMXTy()) {
426 if (Bitcast->user_empty()) {
427 DeadInsts.push_back(Elt: Bitcast);
428 continue;
429 }
430 StoreInst *ST = nullptr;
431 for (Use &U : Bitcast->uses()) {
432 ST = dyn_cast<StoreInst>(Val: U.getUser());
433 if (ST)
434 break;
435 }
436 if (!ST) {
437 if (transformBitcast(Bitcast))
438 DeadInsts.push_back(Elt: 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(Elt: ST);
464 DeadInsts.push_back(Elt: 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
478static Value *getAllocaPos(BasicBlock *BB) {
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(ElementType: Builder.getInt32Ty(), NumElements: 256, Scalable: 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(V: AllocaRes, DestTy: Builder.getPtrTy());
490 return I8Ptr;
491}
492
493static Instruction *createTileStore(Instruction *TileDef, Value *Ptr) {
494 assert(TileDef->getType()->isX86_AMXTy() && "Not define tile!");
495 auto *II = cast<IntrinsicInst>(Val: TileDef);
496
497 assert(II && "Not tile intrinsic!");
498 Value *Row = II->getOperand(i_nocapture: 0);
499 Value *Col = II->getOperand(i_nocapture: 1);
500
501 BasicBlock::iterator Iter = TileDef->getIterator();
502 IRBuilder<> Builder(++Iter);
503 Value *Stride = Builder.getInt64(C: 64);
504 std::array<Value *, 5> Args = {Row, Col, Ptr, Stride, TileDef};
505
506 Instruction *TileStore = Builder.CreateIntrinsicWithoutFolding(
507 ID: 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>(Val: V)->getIncomingValue(i: 0);
519 II = cast<IntrinsicInst>(Val: PhiOp);
520 } else {
521 II = cast<IntrinsicInst>(Val: V);
522 }
523 Value *Row = II->getOperand(i_nocapture: 0);
524 Value *Col = II->getOperand(i_nocapture: 1);
525
526 Instruction *UserI = cast<Instruction>(Val: U.getUser());
527 IRBuilder<> Builder(UserI);
528 Value *Stride = Builder.getInt64(C: 64);
529 std::array<Value *, 4> Args = {Row, Col, Ptr, Stride};
530
531 Value *TileLoad =
532 Builder.CreateIntrinsic(ID: Intrinsic::x86_tileloadd64_internal, Args);
533 UserI->replaceUsesOfWith(From: V, To: TileLoad);
534}
535
536static bool isIncomingOfPHI(Instruction *I) {
537 for (Use &U : I->uses()) {
538 User *V = U.getUser();
539 if (isa<PHINode>(Val: 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(TileDef: I, Ptr: 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>(Val: V) || V == Store)
572 continue;
573 replaceWithTileLoad(U, Ptr: 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, Ptr: StorePtr, IsPHI: 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: I);
647 Instruction *Inst = dyn_cast<Instruction>(Val: Op);
648 assert(Inst && "We shouldn't fold AMX instrution!");
649 Incomings.push_back(Elt: 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(TileDef: I, Ptr: 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, Ptr: 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>(Val: &I))
708 PHIInsts.push_back(Elt: &I);
709 else
710 AMXDefInsts.push_back(Elt: &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(PHI: dyn_cast<PHINode>(Val: 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)) {
754 salvageDebugInfo(I&: *I);
755 salvageKnowledge(I);
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, Val: 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>(Val: OpV)) {
770 if (isInstructionTriviallyDead(I: OpI, TLI)) {
771 WorkList.insert(X: 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(i_nocapture: 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(Elt: PN);
805 OldPhiNodes.insert(X: 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: I);
810 // TODO: currently, We ignore cases where it is a const. In the future, we
811 // might support const.
812 if (isa<Constant>(Val: IncValue)) {
813 auto *IncConst = dyn_cast<Constant>(Val: IncValue);
814 if (!isa<UndefValue>(Val: IncValue) && !IncConst->isNullValue())
815 return false;
816 Value *Row = nullptr, *Col = nullptr;
817 std::tie(args&: Row, args&: Col) = getShape(Phi: 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>(Val: Row) || !isa<Constant>(Val: Col))
821 return false;
822 // Create tilezero at the end of incoming block.
823 auto *Block = OldPN->getIncomingBlock(i: I);
824 BasicBlock::iterator Iter = Block->getTerminator()->getIterator();
825 Instruction *NewInst = Builder.CreateIntrinsicWithoutFolding(
826 ID: Intrinsic::x86_tilezero_internal, OverloadTypes: {}, Args: {Row, Col});
827 NewInst->moveBefore(InsertPos: Iter);
828 NewInst = Builder.CreateIntrinsicWithoutFolding(
829 ID: Intrinsic::x86_cast_tile_to_vector, OverloadTypes: {IncValue->getType()},
830 Args: {NewInst});
831 NewInst->moveBefore(InsertPos: Iter);
832 // Replace InValue with new Value.
833 OldPN->setIncomingValue(i: I, V: NewInst);
834 IncValue = NewInst;
835 }
836
837 if (auto *PNode = dyn_cast<PHINode>(Val: IncValue)) {
838 if (OldPhiNodes.insert(X: PNode))
839 PhiWorklist.push_back(Elt: PNode);
840 continue;
841 }
842 Instruction *ACI = dyn_cast<Instruction>(Val: IncValue);
843 if (ACI && isAMXCast(II: ACI)) {
844 // Verify it's a A->B cast.
845 Type *TyA = ACI->getOperand(i: 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()) {
859 Instruction *ACI = dyn_cast<Instruction>(Val: V);
860 if (ACI && isAMXCast(II: ACI)) {
861 // Verify it's a B->A cast.
862 Type *TyB = ACI->getOperand(i: 0)->getType();
863 Type *TyA = ACI->getType();
864 if (TyA != DestTy || TyB != SrcTy)
865 return false;
866 } else if (auto *PHI = dyn_cast<PHINode>(Val: 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(key: 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(Ty: DestTy, NumReservedValues: 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(i_nocapture: j);
905 Value *NewV = nullptr;
906 Instruction *ACI = dyn_cast<Instruction>(Val: V);
907 // There should not be a AMXcast from a const.
908 if (ACI && isAMXCast(II: ACI))
909 NewV = ACI->getOperand(i: 0);
910 else if (auto *PrevPN = dyn_cast<PHINode>(Val: V))
911 NewV = NewPNodes[PrevPN];
912 assert(NewV);
913 NewPN->addIncoming(V: NewV, BB: OldPN->getIncomingBlock(i: 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(Range: OldPN->users())) {
928 Instruction *ACI = dyn_cast<Instruction>(Val: V);
929 if (ACI && isAMXCast(II: ACI)) {
930 Type *TyB = ACI->getOperand(i: 0)->getType();
931 Type *TyA = ACI->getType();
932 assert(TyA == DestTy && TyB == SrcTy);
933 (void)TyA;
934 (void)TyB;
935 ACI->replaceAllUsesWith(V: NewPN);
936 DeadInst.insert(X: ACI);
937 } else if (auto *PHI = dyn_cast<PHINode>(Val: 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(i_nocapture: 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>(Val: 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(i_nocapture: 0);
967 Value *Col = II->getOperand(i_nocapture: 1);
968
969 IRBuilder<> Builder(ST);
970
971 // Stride should be equal to col(measured by bytes)
972 Value *Stride = Builder.CreateSExt(V: Col, DestTy: Builder.getInt64Ty());
973 Value *I8Ptr = Builder.CreateBitCast(V: ST->getOperand(i_nocapture: 1), DestTy: Builder.getPtrTy());
974 std::array<Value *, 5> Args = {Row, Col, I8Ptr, Stride, Tile};
975 Builder.CreateIntrinsic(ID: 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>(Val: 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(I: II))
993 return false;
994 std::tie(args&: Row, args&: 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(p: new DominatorTree(Func));
1001 if (!DT->dominates(Def: Row, User: LD) || !DT->dominates(Def: Col, User: LD)) {
1002 // store the value to stack and reload it from stack before cast.
1003 auto *AllocaAddr =
1004 createAllocaInstAtEntry(Builder, BB: Cast->getParent(), Ty: LD->getType());
1005 Builder.SetInsertPoint(&*std::next(x: LD->getIterator()));
1006 Builder.CreateStore(Val: LD, Ptr: AllocaAddr);
1007
1008 Builder.SetInsertPoint(Cast);
1009 I8Ptr = Builder.CreateBitCast(V: AllocaAddr, DestTy: Builder.getPtrTy());
1010 EraseLoad = false;
1011 } else {
1012 I8Ptr = Builder.CreateBitCast(V: LD->getOperand(i_nocapture: 0), DestTy: Builder.getPtrTy());
1013 }
1014 // Stride should be equal to col(measured by bytes)
1015 Value *Stride = Builder.CreateSExt(V: Col, DestTy: Builder.getInt64Ty());
1016 std::array<Value *, 4> Args = {Row, Col, I8Ptr, Stride};
1017
1018 Value *NewInst =
1019 Builder.CreateIntrinsic(ID: Intrinsic::x86_tileloadd64_internal, Args);
1020 Cast->replaceAllUsesWith(V: 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>(Val: U.getUser());
1033 if (!isAMXIntrinsic(I: II))
1034 return false;
1035
1036 std::tie(args&: Row, args&: Col) = getShape(II, OpNo);
1037
1038 IRBuilder<> Builder(Cast);
1039 Value *NewInst =
1040 Builder.CreateIntrinsic(ID: Intrinsic::x86_tilezero_internal, OverloadTypes: {}, Args: {Row, Col});
1041 Cast->replaceAllUsesWith(V: 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>(Val: 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>(Val: U);
1058 if (!Store)
1059 continue;
1060 if (combineCastStore(Cast: cast<IntrinsicInst>(Val: Cast), ST: Store)) {
1061 DeadStores.push_back(Elt: 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>(Val: Cast->getOperand(i: 0))) {
1072 Change |= combineTilezero(Cast: cast<IntrinsicInst>(Val: Cast));
1073 continue;
1074 }
1075
1076 auto *Load = dyn_cast<LoadInst>(Val: Cast->getOperand(i: 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: cast<IntrinsicInst>(Val: Cast), LD: Load)) {
1085 // Set the operand is null so that load instruction can be erased.
1086 Cast->setOperand(i: 0, Val: 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(V: &I,
1106 P: m_Intrinsic<Intrinsic::x86_cast_vector_to_tile>(Ops: m_Value(V&: Vec))))
1107 Vec2TileInsts.push_back(Elt: &I);
1108 else if (match(V: &I, P: m_Intrinsic<Intrinsic::x86_cast_tile_to_vector>(
1109 Ops: m_Value(V&: Vec))))
1110 Tile2VecInsts.push_back(Elt: &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>(Val: 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(V: Inst->getOperand(i: 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(Elt: 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(Casts&: 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(II: &I)) {
1163 if (isa<PHINode>(Val: I.getOperand(i: 0)))
1164 PhiCastWorkList.push_back(Elt: &I);
1165 }
1166 }
1167 }
1168 for (auto *I : PhiCastWorkList) {
1169 // We skip the dead Amxcast.
1170 if (DeadInst.contains(key: I))
1171 continue;
1172 PHINode *PN = cast<PHINode>(Val: I->getOperand(i: 0));
1173 if (optimizeAMXCastFromPhi(CI: cast<IntrinsicInst>(Val: I), PN, DeadInst)) {
1174 DeadInst.insert(X: 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, WorkList&: 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(i_nocapture: 0);
1198
1199 auto Prepare = [&](Type *MemTy) {
1200 AllocaAddr = createAllocaInstAtEntry(Builder, BB: AMXCast->getParent(), Ty: MemTy);
1201 I8Ptr = Builder.CreateBitCast(V: AllocaAddr, DestTy: Builder.getPtrTy());
1202 Stride = Builder.getInt64(C: 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>(Val: U.getUser());
1225 if (!II)
1226 return false; // May be bitcast from x86amx to <256 x i32>.
1227 Prepare(AMXCast->getOperand(i_nocapture: 0)->getType());
1228 Builder.CreateStore(Val: Src, Ptr: AllocaAddr);
1229 // TODO we can pick an constant operand for the shape.
1230 Value *Row = nullptr, *Col = nullptr;
1231 std::tie(args&: Row, args&: Col) = getShape(II, OpNo);
1232 std::array<Value *, 4> Args = {
1233 Row, Col, I8Ptr, Builder.CreateSExt(V: Col, DestTy: Builder.getInt64Ty())};
1234 Value *NewInst =
1235 Builder.CreateIntrinsic(ID: Intrinsic::x86_tileloadd64_internal, Args);
1236 AMXCast->replaceAllUsesWith(V: 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>(Val: 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(i_nocapture: 0);
1251 Value *Col = II->getOperand(i_nocapture: 1);
1252 std::array<Value *, 5> Args = {
1253 Row, Col, I8Ptr, Builder.CreateSExt(V: Col, DestTy: Builder.getInt64Ty()), Src};
1254 Builder.CreateIntrinsic(ID: Intrinsic::x86_tilestored64_internal, Args);
1255 Value *NewInst = Builder.CreateLoad(Ty: AMXCast->getType(), Ptr: AllocaAddr);
1256 AMXCast->replaceAllUsesWith(V: 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(II: &I))
1270 WorkLists.push_back(Elt: &I);
1271 }
1272 }
1273
1274 for (auto *Inst : WorkLists) {
1275 Change |= transformAMXCast(AMXCast: cast<IntrinsicInst>(Val: 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(Kind: Attribute::OptimizeNone)) {
1310 X86VolatileTileData VTD(F);
1311 C = VTD.volatileTileData() || C;
1312 }
1313 }
1314
1315 return C;
1316}
1317
1318} // anonymous namespace
1319
1320PreservedAnalyses X86LowerAMXTypePass::run(Function &F,
1321 FunctionAnalysisManager &FAM) {
1322 TargetLibraryInfo &TLI = FAM.getResult<TargetLibraryAnalysis>(IR&: F);
1323 bool Changed = lowerAmxType(F, TM, TLI: &TLI);
1324 if (!Changed)
1325 return PreservedAnalyses::all();
1326
1327 PreservedAnalyses PA = PreservedAnalyses::none();
1328 PA.preserveSet<CFGAnalyses>();
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)
1360INITIALIZE_PASS_DEPENDENCY(TargetPassConfig)
1361INITIALIZE_PASS_DEPENDENCY(TargetLibraryInfoWrapperPass)
1362INITIALIZE_PASS_END(X86LowerAMXTypeLegacyPass, DEBUG_TYPE, PassName, false,
1363 false)
1364
1365FunctionPass *llvm::createX86LowerAMXTypeLegacyPass() {
1366 return new X86LowerAMXTypeLegacyPass();
1367}
1368