1//===-- X86LowerAMXIntrinsics.cpp -X86 Scalarize AMX Intrinsics------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9/// \file Pass to transform amx intrinsics to scalar operations.
10/// This pass is always enabled and it skips when it is not -O0 and has no
11/// optnone attributes. With -O0 or optnone attribute, the def of shape to amx
12/// intrinsics is near the amx intrinsics code. We are not able to find a
13/// point which post-dominate all the shape and dominate all amx intrinsics.
14/// To decouple the dependency of the shape, we transform amx intrinsics
15/// to scalar operation, so that compiling doesn't fail. In long term, we
16/// should improve fast register allocation to allocate amx register.
17//===----------------------------------------------------------------------===//
18//
19#include "X86.h"
20#include "X86TargetMachine.h"
21#include "llvm/Analysis/DomTreeUpdater.h"
22#include "llvm/Analysis/LoopInfo.h"
23#include "llvm/Analysis/TargetTransformInfo.h"
24#include "llvm/CodeGen/Passes.h"
25#include "llvm/CodeGen/TargetPassConfig.h"
26#include "llvm/CodeGen/ValueTypes.h"
27#include "llvm/IR/Analysis.h"
28#include "llvm/IR/DataLayout.h"
29#include "llvm/IR/Dominators.h"
30#include "llvm/IR/Function.h"
31#include "llvm/IR/IRBuilder.h"
32#include "llvm/IR/Instructions.h"
33#include "llvm/IR/IntrinsicInst.h"
34#include "llvm/IR/IntrinsicsX86.h"
35#include "llvm/IR/MDBuilder.h"
36#include "llvm/IR/PassManager.h"
37#include "llvm/IR/PatternMatch.h"
38#include "llvm/IR/ProfDataUtils.h"
39#include "llvm/InitializePasses.h"
40#include "llvm/Pass.h"
41#include "llvm/Support/CommandLine.h"
42#include "llvm/Target/TargetMachine.h"
43#include "llvm/Transforms/Utils/BasicBlockUtils.h"
44#include "llvm/Transforms/Utils/LoopUtils.h"
45
46using namespace llvm;
47using namespace PatternMatch;
48
49namespace llvm {
50extern cl::opt<bool> ProfcheckDisableMetadataFixes;
51} // end namespace llvm
52
53#define DEBUG_TYPE "x86-lower-amx-intrinsics"
54
55#ifndef NDEBUG
56static bool isV256I32Ty(Type *Ty) {
57 if (auto *FVT = dyn_cast<FixedVectorType>(Ty))
58 return FVT->getNumElements() == 256 &&
59 FVT->getElementType()->isIntegerTy(32);
60 return false;
61}
62#endif
63
64namespace {
65class X86LowerAMXIntrinsics {
66 Function &Func;
67
68public:
69 X86LowerAMXIntrinsics(Function &F, DomTreeUpdater &DomTU, LoopInfo *LoopI)
70 : Func(F), DTU(DomTU), LI(LoopI) {}
71 bool visit();
72
73private:
74 DomTreeUpdater &DTU;
75 LoopInfo *LI;
76 BasicBlock *createLoop(BasicBlock *Preheader, BasicBlock *Exit, Value *Bound,
77 ConstantInt *Step, StringRef Name, IRBuilderBase &B,
78 Loop *L);
79 template <bool IsTileLoad>
80 Value *createTileLoadStoreLoops(BasicBlock *Start, BasicBlock *End,
81 IRBuilderBase &B, Value *Row, Value *Col,
82 Value *Ptr, Value *Stride, Value *Tile);
83 template <Intrinsic::ID IntrID>
84 std::enable_if_t<IntrID == Intrinsic::x86_tdpbssd_internal ||
85 IntrID == Intrinsic::x86_tdpbsud_internal ||
86 IntrID == Intrinsic::x86_tdpbusd_internal ||
87 IntrID == Intrinsic::x86_tdpbuud_internal ||
88 IntrID == Intrinsic::x86_tdpbf16ps_internal,
89 Value *>
90 createTileDPLoops(BasicBlock *Start, BasicBlock *End, IRBuilderBase &B,
91 Value *Row, Value *Col, Value *K, Value *Acc, Value *LHS,
92 Value *RHS);
93 template <bool IsTileLoad>
94 bool lowerTileLoadStore(Instruction *TileLoadStore);
95 template <Intrinsic::ID IntrID>
96 std::enable_if_t<IntrID == Intrinsic::x86_tdpbssd_internal ||
97 IntrID == Intrinsic::x86_tdpbsud_internal ||
98 IntrID == Intrinsic::x86_tdpbusd_internal ||
99 IntrID == Intrinsic::x86_tdpbuud_internal ||
100 IntrID == Intrinsic::x86_tdpbf16ps_internal,
101 bool>
102 lowerTileDP(Instruction *TileDP);
103 bool lowerTileZero(Instruction *TileZero);
104};
105} // anonymous namespace
106
107BasicBlock *X86LowerAMXIntrinsics::createLoop(BasicBlock *Preheader,
108 BasicBlock *Exit, Value *Bound,
109 ConstantInt *Step, StringRef Name,
110 IRBuilderBase &B, Loop *L) {
111 LLVMContext &Ctx = Preheader->getContext();
112 BasicBlock *Header =
113 BasicBlock::Create(Context&: Ctx, Name: Name + ".header", Parent: Preheader->getParent(), InsertBefore: Exit);
114 BasicBlock *Body =
115 BasicBlock::Create(Context&: Ctx, Name: Name + ".body", Parent: Header->getParent(), InsertBefore: Exit);
116 BasicBlock *Latch =
117 BasicBlock::Create(Context&: Ctx, Name: Name + ".latch", Parent: Header->getParent(), InsertBefore: Exit);
118
119 Type *I16Ty = Type::getInt16Ty(C&: Ctx);
120 UncondBrInst::Create(Target: Body, InsertBefore: Header);
121 UncondBrInst::Create(Target: Latch, InsertBefore: Body);
122 PHINode *IV =
123 PHINode::Create(Ty: I16Ty, NumReservedValues: 2, NameStr: Name + ".iv", InsertBefore: Header->getTerminator()->getIterator());
124 IV->addIncoming(V: ConstantInt::get(Ty: I16Ty, V: 0), BB: Preheader);
125
126 B.SetInsertPoint(Latch);
127 Value *Inc = B.CreateAdd(LHS: IV, RHS: Step, Name: Name + ".step");
128 Value *Cond = B.CreateICmpNE(LHS: Inc, RHS: Bound, Name: Name + ".cond");
129 auto *BR = CondBrInst::Create(Cond, IfTrue: Header, IfFalse: Exit, InsertBefore: Latch);
130 if (!ProfcheckDisableMetadataFixes) {
131 if (auto *BoundInt = dyn_cast<ConstantInt>(Val: Bound)) {
132 assert(Step->getZExtValue() != 0 &&
133 "Expected a non-zero step size. This is chosen by the pass and "
134 "should always be non-zero to imply a finite loop.");
135 MDBuilder MDB(Preheader->getContext());
136 setFittedBranchWeights(
137 I&: *BR, Weights: {BoundInt->getZExtValue() / Step->getZExtValue(), 1}, IsExpected: false);
138 } else {
139 setExplicitlyUnknownBranchWeightsIfProfiled(I&: *BR, DEBUG_TYPE);
140 }
141 }
142 IV->addIncoming(V: Inc, BB: Latch);
143
144 UncondBrInst *PreheaderBr = cast<UncondBrInst>(Val: Preheader->getTerminator());
145 BasicBlock *Tmp = PreheaderBr->getSuccessor();
146 PreheaderBr->setSuccessor(Header);
147 DTU.applyUpdatesPermissive(Updates: {
148 {DominatorTree::Delete, Preheader, Tmp},
149 {DominatorTree::Insert, Header, Body},
150 {DominatorTree::Insert, Body, Latch},
151 {DominatorTree::Insert, Latch, Header},
152 {DominatorTree::Insert, Latch, Exit},
153 {DominatorTree::Insert, Preheader, Header},
154 });
155 if (LI) {
156 L->addBasicBlockToLoop(NewBB: Header, LI&: *LI);
157 L->addBasicBlockToLoop(NewBB: Body, LI&: *LI);
158 L->addBasicBlockToLoop(NewBB: Latch, LI&: *LI);
159 }
160 return Body;
161}
162
163template <bool IsTileLoad>
164Value *X86LowerAMXIntrinsics::createTileLoadStoreLoops(
165 BasicBlock *Start, BasicBlock *End, IRBuilderBase &B, Value *Row,
166 Value *Col, Value *Ptr, Value *Stride, Value *Tile) {
167 std::string IntrinName = IsTileLoad ? "tileload" : "tilestore";
168 Loop *RowLoop = nullptr;
169 Loop *ColLoop = nullptr;
170 if (LI) {
171 RowLoop = LI->AllocateLoop();
172 ColLoop = LI->AllocateLoop();
173 RowLoop->addChildLoop(NewChild: ColLoop);
174 if (Loop *ParentL = LI->getLoopFor(BB: Start))
175 ParentL->addChildLoop(NewChild: RowLoop);
176 else
177 LI->addTopLevelLoop(New: RowLoop);
178 }
179
180 BasicBlock *RowBody = createLoop(Preheader: Start, Exit: End, Bound: Row, Step: B.getInt16(C: 1),
181 Name: IntrinName + ".scalarize.rows", B, L: RowLoop);
182 BasicBlock *RowLatch = RowBody->getSingleSuccessor();
183
184 BasicBlock *ColBody = createLoop(Preheader: RowBody, Exit: RowLatch, Bound: Col, Step: B.getInt16(C: 1),
185 Name: IntrinName + ".scalarize.cols", B, L: ColLoop);
186
187 BasicBlock *ColLoopLatch = ColBody->getSingleSuccessor();
188 BasicBlock *ColLoopHeader = ColBody->getSinglePredecessor();
189 BasicBlock *RowLoopHeader = RowBody->getSinglePredecessor();
190 Value *CurrentRow = &*RowLoopHeader->begin();
191 Value *CurrentCol = &*ColLoopHeader->begin();
192 Type *EltTy = B.getInt32Ty();
193 FixedVectorType *V256I32Ty = FixedVectorType::get(ElementType: EltTy, NumElts: 256);
194
195 // Common part for tileload and tilestore
196 // *.scalarize.cols.body:
197 // Calculate %idxmem and %idxvec
198 B.SetInsertPoint(ColBody->getTerminator());
199 Value *CurrentRowZExt = B.CreateZExt(V: CurrentRow, DestTy: Stride->getType());
200 Value *CurrentColZExt = B.CreateZExt(V: CurrentCol, DestTy: Stride->getType());
201 Value *Offset =
202 B.CreateAdd(LHS: B.CreateMul(LHS: CurrentRowZExt, RHS: Stride), RHS: CurrentColZExt);
203 Value *EltPtr = B.CreateGEP(Ty: EltTy, Ptr, IdxList: Offset);
204 Value *Idx = B.CreateAdd(LHS: B.CreateMul(LHS: CurrentRow, RHS: B.getInt16(C: 16)), RHS: CurrentCol);
205 if (IsTileLoad) {
206 // tileload.scalarize.rows.header:
207 // %vec.phi.row = phi <256 x i32> [ zeroinitializer, %entry ], [ %ResVec,
208 // %tileload.scalarize.rows.latch ]
209 B.SetInsertPoint(RowLoopHeader->getTerminator());
210 Value *VecZero = Constant::getNullValue(Ty: V256I32Ty);
211 PHINode *VecCPhiRowLoop = B.CreatePHI(Ty: V256I32Ty, NumReservedValues: 2, Name: "vec.phi.row");
212 VecCPhiRowLoop->addIncoming(V: VecZero, BB: Start);
213
214 // tileload.scalarize.cols.header:
215 // %vec.phi = phi <256 x i32> [ %vec.phi.row, %tileload.scalarize.rows.body
216 // ], [ %ResVec, %tileload.scalarize.cols.latch ]
217 B.SetInsertPoint(ColLoopHeader->getTerminator());
218 PHINode *VecPhi = B.CreatePHI(Ty: V256I32Ty, NumReservedValues: 2, Name: "vec.phi");
219 VecPhi->addIncoming(V: VecCPhiRowLoop, BB: RowBody);
220
221 // tileload.scalarize.cols.body:
222 // Calculate %idxmem and %idxvec
223 // %eltptr = getelementptr i32, i32* %base, i64 %idxmem
224 // %elt = load i32, i32* %ptr
225 // %ResVec = insertelement <256 x i32> %vec.phi, i32 %elt, i16 %idxvec
226 B.SetInsertPoint(ColBody->getTerminator());
227 Value *Elt = B.CreateLoad(Ty: EltTy, Ptr: EltPtr);
228 Value *ResVec = B.CreateInsertElement(Vec: VecPhi, NewElt: Elt, Idx);
229 VecPhi->addIncoming(V: ResVec, BB: ColLoopLatch);
230 VecCPhiRowLoop->addIncoming(V: ResVec, BB: RowLatch);
231
232 return ResVec;
233 } else {
234 auto *BitCast = cast<BitCastInst>(Val: Tile);
235 Value *Vec = BitCast->getOperand(i_nocapture: 0);
236 assert(isV256I32Ty(Vec->getType()) && "bitcast from non-v256i32 to x86amx");
237 // tilestore.scalarize.cols.body:
238 // %mul = mul i16 %row.iv, i16 16
239 // %idx = add i16 %mul, i16 %col.iv
240 // %vec = extractelement <16 x i32> %vec, i16 %idx
241 // store i32 %vec, i32* %ptr
242 B.SetInsertPoint(ColBody->getTerminator());
243 Value *Elt = B.CreateExtractElement(Vec, Idx);
244
245 B.CreateStore(Val: Elt, Ptr: EltPtr);
246 return nullptr;
247 }
248}
249
250template <Intrinsic::ID IntrID>
251std::enable_if_t<IntrID == Intrinsic::x86_tdpbssd_internal ||
252 IntrID == Intrinsic::x86_tdpbsud_internal ||
253 IntrID == Intrinsic::x86_tdpbusd_internal ||
254 IntrID == Intrinsic::x86_tdpbuud_internal ||
255 IntrID == Intrinsic::x86_tdpbf16ps_internal,
256 Value *>
257X86LowerAMXIntrinsics::createTileDPLoops(BasicBlock *Start, BasicBlock *End,
258 IRBuilderBase &B, Value *Row,
259 Value *Col, Value *K, Value *Acc,
260 Value *LHS, Value *RHS) {
261 std::string IntrinName;
262 switch (IntrID) {
263 case Intrinsic::x86_tdpbssd_internal:
264 IntrinName = "tiledpbssd";
265 break;
266 case Intrinsic::x86_tdpbsud_internal:
267 IntrinName = "tiledpbsud";
268 break;
269 case Intrinsic::x86_tdpbusd_internal:
270 IntrinName = "tiledpbusd";
271 break;
272 case Intrinsic::x86_tdpbuud_internal:
273 IntrinName = "tiledpbuud";
274 break;
275 case Intrinsic::x86_tdpbf16ps_internal:
276 IntrinName = "tiledpbf16ps";
277 break;
278 }
279 Loop *RowLoop = nullptr;
280 Loop *ColLoop = nullptr;
281 Loop *InnerLoop = nullptr;
282 if (LI) {
283 RowLoop = LI->AllocateLoop();
284 ColLoop = LI->AllocateLoop();
285 InnerLoop = LI->AllocateLoop();
286 ColLoop->addChildLoop(NewChild: InnerLoop);
287 RowLoop->addChildLoop(NewChild: ColLoop);
288 if (Loop *ParentL = LI->getLoopFor(BB: Start))
289 ParentL->addChildLoop(NewChild: RowLoop);
290 else
291 LI->addTopLevelLoop(New: RowLoop);
292 }
293
294 BasicBlock *RowBody = createLoop(Preheader: Start, Exit: End, Bound: Row, Step: B.getInt16(C: 1),
295 Name: IntrinName + ".scalarize.rows", B, L: RowLoop);
296 BasicBlock *RowLatch = RowBody->getSingleSuccessor();
297
298 BasicBlock *ColBody = createLoop(Preheader: RowBody, Exit: RowLatch, Bound: Col, Step: B.getInt16(C: 1),
299 Name: IntrinName + ".scalarize.cols", B, L: ColLoop);
300
301 BasicBlock *ColLoopLatch = ColBody->getSingleSuccessor();
302
303 B.SetInsertPoint(ColBody->getTerminator());
304 BasicBlock *InnerBody =
305 createLoop(Preheader: ColBody, Exit: ColLoopLatch, Bound: K, Step: B.getInt16(C: 1),
306 Name: IntrinName + ".scalarize.inner", B, L: InnerLoop);
307
308 BasicBlock *ColLoopHeader = ColBody->getSinglePredecessor();
309 BasicBlock *RowLoopHeader = RowBody->getSinglePredecessor();
310 BasicBlock *InnerLoopHeader = InnerBody->getSinglePredecessor();
311 BasicBlock *InnerLoopLatch = InnerBody->getSingleSuccessor();
312 Value *CurrentRow = &*RowLoopHeader->begin();
313 Value *CurrentCol = &*ColLoopHeader->begin();
314 Value *CurrentInner = &*InnerLoopHeader->begin();
315
316 FixedVectorType *V256I32Ty = FixedVectorType::get(ElementType: B.getInt32Ty(), NumElts: 256);
317 auto *BitCastAcc = cast<BitCastInst>(Val: Acc);
318 Value *VecC = BitCastAcc->getOperand(i_nocapture: 0);
319 assert(isV256I32Ty(VecC->getType()) && "bitcast from non-v256i32 to x86amx");
320 // TODO else create BitCast from x86amx to v256i32.
321 // Store x86amx to memory, and reload from memory
322 // to vector. However with -O0, it doesn't happen.
323 auto *BitCastLHS = cast<BitCastInst>(Val: LHS);
324 Value *VecA = BitCastLHS->getOperand(i_nocapture: 0);
325 assert(isV256I32Ty(VecA->getType()) && "bitcast from non-v256i32 to x86amx");
326 auto *BitCastRHS = cast<BitCastInst>(Val: RHS);
327 Value *VecB = BitCastRHS->getOperand(i_nocapture: 0);
328 assert(isV256I32Ty(VecB->getType()) && "bitcast from non-v256i32 to x86amx");
329
330 // tiledpbssd.scalarize.rows.header:
331 // %vec.c.phi.row = phi <256 x i32> [ %VecC, %continue ], [ %NewVecC,
332 // %tiledpbssd.scalarize.rows.latch ]
333
334 // %vec.d.phi.row = phi <256 x i32> [ zeroinitializer, %continue ], [
335 // %NewVecD, %tiledpbssd.scalarize.rows.latch ]
336 B.SetInsertPoint(RowLoopHeader->getTerminator());
337 PHINode *VecCPhiRowLoop = B.CreatePHI(Ty: V256I32Ty, NumReservedValues: 2, Name: "vec.c.phi.row");
338 VecCPhiRowLoop->addIncoming(V: VecC, BB: Start);
339 Value *VecZero = Constant::getNullValue(Ty: V256I32Ty);
340 PHINode *VecDPhiRowLoop = B.CreatePHI(Ty: V256I32Ty, NumReservedValues: 2, Name: "vec.d.phi.row");
341 VecDPhiRowLoop->addIncoming(V: VecZero, BB: Start);
342
343 // tiledpbssd.scalarize.cols.header:
344 // %vec.c.phi.col = phi <256 x i32> [ %vec.c.phi.row,
345 // %tiledpbssd.scalarize.rows.body ], [ %NewVecC,
346 // %tiledpbssd.scalarize.cols.latch ]
347
348 // %vec.d.phi.col = phi <256 x i32> [
349 // %vec.d.phi.row, %tiledpbssd.scalarize.rows.body ], [ %NewVecD,
350 // %tiledpbssd.scalarize.cols.latch ]
351
352 // calculate idxc.
353 B.SetInsertPoint(ColLoopHeader->getTerminator());
354 PHINode *VecCPhiColLoop = B.CreatePHI(Ty: V256I32Ty, NumReservedValues: 2, Name: "vec.c.phi.col");
355 VecCPhiColLoop->addIncoming(V: VecCPhiRowLoop, BB: RowBody);
356 PHINode *VecDPhiColLoop = B.CreatePHI(Ty: V256I32Ty, NumReservedValues: 2, Name: "vec.d.phi.col");
357 VecDPhiColLoop->addIncoming(V: VecDPhiRowLoop, BB: RowBody);
358 Value *IdxC =
359 B.CreateAdd(LHS: B.CreateMul(LHS: CurrentRow, RHS: B.getInt16(C: 16)), RHS: CurrentCol);
360
361 // tiledpbssd.scalarize.inner.header:
362 // %vec.c.inner.phi = phi <256 x i32> [ %vec.c.phi.col,
363 // %tiledpbssd.scalarize.cols.body ], [ %NewVecC,
364 // %tiledpbssd.scalarize.inner.latch ]
365
366 B.SetInsertPoint(InnerLoopHeader->getTerminator());
367 PHINode *VecCPhi = B.CreatePHI(Ty: V256I32Ty, NumReservedValues: 2, Name: "vec.c.inner.phi");
368 VecCPhi->addIncoming(V: VecCPhiColLoop, BB: ColBody);
369
370 B.SetInsertPoint(InnerBody->getTerminator());
371 Value *IdxA =
372 B.CreateAdd(LHS: B.CreateMul(LHS: CurrentRow, RHS: B.getInt16(C: 16)), RHS: CurrentInner);
373 Value *IdxB =
374 B.CreateAdd(LHS: B.CreateMul(LHS: CurrentInner, RHS: B.getInt16(C: 16)), RHS: CurrentCol);
375 Value *NewVecC = nullptr;
376
377 if (IntrID != Intrinsic::x86_tdpbf16ps_internal) {
378 // tiledpbssd.scalarize.inner.body:
379 // calculate idxa, idxb
380 // %eltc = extractelement <256 x i32> %vec.c.inner.phi, i16 %idxc
381 // %elta = extractelement <256 x i32> %veca, i16 %idxa
382 // %eltav4i8 = bitcast i32 %elta to <4 x i8>
383 // %eltb = extractelement <256 x i32> %vecb, i16 %idxb
384 // %eltbv4i8 = bitcast i32 %eltb to <4 x i8>
385 // %eltav4i32 = sext <4 x i8> %eltav4i8 to <4 x i32>
386 // %eltbv4i32 = sext <4 x i8> %eltbv4i8 to <4 x i32>
387 // %mulab = mul <4 x i32> %eltbv4i32, %eltav4i32
388 // %acc = call i32 @llvm.vector.reduce.add.v4i32(<4 x i32> %131)
389 // %neweltc = add i32 %elt, %acc
390 // %NewVecC = insertelement <256 x i32> %vec.c.inner.phi, i32 %neweltc,
391 // i16 %idxc
392 FixedVectorType *V4I8Ty = FixedVectorType::get(ElementType: B.getInt8Ty(), NumElts: 4);
393 FixedVectorType *V4I32Ty = FixedVectorType::get(ElementType: B.getInt32Ty(), NumElts: 4);
394 Value *EltC = B.CreateExtractElement(Vec: VecCPhi, Idx: IdxC);
395 Value *EltA = B.CreateExtractElement(Vec: VecA, Idx: IdxA);
396 Value *SubVecA = B.CreateBitCast(V: EltA, DestTy: V4I8Ty);
397 Value *EltB = B.CreateExtractElement(Vec: VecB, Idx: IdxB);
398 Value *SubVecB = B.CreateBitCast(V: EltB, DestTy: V4I8Ty);
399 Value *SEXTSubVecB = nullptr;
400 Value *SEXTSubVecA = nullptr;
401 switch (IntrID) {
402 case Intrinsic::x86_tdpbssd_internal:
403 SEXTSubVecB = B.CreateSExt(V: SubVecB, DestTy: V4I32Ty);
404 SEXTSubVecA = B.CreateSExt(V: SubVecA, DestTy: V4I32Ty);
405 break;
406 case Intrinsic::x86_tdpbsud_internal:
407 SEXTSubVecB = B.CreateZExt(V: SubVecB, DestTy: V4I32Ty);
408 SEXTSubVecA = B.CreateSExt(V: SubVecA, DestTy: V4I32Ty);
409 break;
410 case Intrinsic::x86_tdpbusd_internal:
411 SEXTSubVecB = B.CreateSExt(V: SubVecB, DestTy: V4I32Ty);
412 SEXTSubVecA = B.CreateZExt(V: SubVecA, DestTy: V4I32Ty);
413 break;
414 case Intrinsic::x86_tdpbuud_internal:
415 SEXTSubVecB = B.CreateZExt(V: SubVecB, DestTy: V4I32Ty);
416 SEXTSubVecA = B.CreateZExt(V: SubVecA, DestTy: V4I32Ty);
417 break;
418 default:
419 llvm_unreachable("Invalid intrinsic ID!");
420 }
421 Value *SubVecR = B.CreateAddReduce(Src: B.CreateMul(LHS: SEXTSubVecA, RHS: SEXTSubVecB));
422 Value *ResElt = B.CreateAdd(LHS: EltC, RHS: SubVecR);
423 NewVecC = B.CreateInsertElement(Vec: VecCPhi, NewElt: ResElt, Idx: IdxC);
424 } else {
425 // tiledpbf16ps.scalarize.inner.body:
426 // calculate idxa, idxb, idxc
427 // %eltc = extractelement <256 x i32> %vec.c.inner.phi, i16 %idxc
428 // %eltcf32 = bitcast i32 %eltc to float
429 // %elta = extractelement <256 x i32> %veca, i16 %idxa
430 // %eltav2i16 = bitcast i32 %elta to <2 x i16>
431 // %eltb = extractelement <256 x i32> %vecb, i16 %idxb
432 // %eltbv2i16 = bitcast i32 %eltb to <2 x i16>
433 // %shufflea = shufflevector <2 x i16> %elta, <2 x i16> zeroinitializer, <4
434 // x i32> <i32 2, i32 0, i32 3, i32 1>
435 // %eltav2f32 = bitcast <4 x i16> %shufflea to <2 x float>
436 // %shuffleb = shufflevector <2 x i16> %eltb, <2 xi16> zeroinitializer, <4 x
437 // i32> <i32 2, i32 0, i32 3, i32 1>
438 // %eltbv2f32 = bitcast <4 x i16> %shuffleb to <2 x float>
439 // %mulab = fmul <2 x float> %eltav2f32, %eltbv2f32
440 // %acc = call float
441 // @llvm.vector.reduce.fadd.v2f32(float %eltcf32, <2 x float> %mulab)
442 // %neweltc = bitcast float %acc to i32
443 // %NewVecC = insertelement <256 x i32> %vec.c.inner.phi, i32 %neweltc,
444 // i16 %idxc
445 // %NewVecD = insertelement <256 x i32> %vec.d.inner.phi, i32 %neweltc,
446 // i16 %idxc
447 FixedVectorType *V2I16Ty = FixedVectorType::get(ElementType: B.getInt16Ty(), NumElts: 2);
448 FixedVectorType *V2F32Ty = FixedVectorType::get(ElementType: B.getFloatTy(), NumElts: 2);
449 Value *EltC = B.CreateExtractElement(Vec: VecCPhi, Idx: IdxC);
450 Value *EltCF32 = B.CreateBitCast(V: EltC, DestTy: B.getFloatTy());
451 Value *EltA = B.CreateExtractElement(Vec: VecA, Idx: IdxA);
452 Value *SubVecA = B.CreateBitCast(V: EltA, DestTy: V2I16Ty);
453 Value *EltB = B.CreateExtractElement(Vec: VecB, Idx: IdxB);
454 Value *SubVecB = B.CreateBitCast(V: EltB, DestTy: V2I16Ty);
455 Value *ZeroV2I16 = Constant::getNullValue(Ty: V2I16Ty);
456 int ShuffleMask[4] = {2, 0, 3, 1};
457 auto ShuffleArray = ArrayRef(ShuffleMask);
458 Value *AV2F32 = B.CreateBitCast(
459 V: B.CreateShuffleVector(V1: SubVecA, V2: ZeroV2I16, Mask: ShuffleArray), DestTy: V2F32Ty);
460 Value *BV2F32 = B.CreateBitCast(
461 V: B.CreateShuffleVector(V1: SubVecB, V2: ZeroV2I16, Mask: ShuffleArray), DestTy: V2F32Ty);
462 Value *SubVecR = B.CreateFAddReduce(Acc: EltCF32, Src: B.CreateFMul(L: AV2F32, R: BV2F32));
463 Value *ResElt = B.CreateBitCast(V: SubVecR, DestTy: B.getInt32Ty());
464 NewVecC = B.CreateInsertElement(Vec: VecCPhi, NewElt: ResElt, Idx: IdxC);
465 }
466
467 // tiledpbssd.scalarize.cols.latch:
468 // %NewEltC = extractelement <256 x i32> %vec.c.phi.col, i16 %idxc
469 // %NewVecD = insertelement <256 x i32> %vec.d.phi.col, i32 %NewEltC,
470 // i16 %idxc
471 B.SetInsertPoint(ColLoopLatch->getTerminator());
472 Value *NewEltC = B.CreateExtractElement(Vec: NewVecC, Idx: IdxC);
473 Value *NewVecD = B.CreateInsertElement(Vec: VecDPhiColLoop, NewElt: NewEltC, Idx: IdxC);
474
475 VecCPhi->addIncoming(V: NewVecC, BB: InnerLoopLatch);
476 VecCPhiRowLoop->addIncoming(V: NewVecC, BB: RowLatch);
477 VecCPhiColLoop->addIncoming(V: NewVecC, BB: ColLoopLatch);
478 VecDPhiRowLoop->addIncoming(V: NewVecD, BB: RowLatch);
479 VecDPhiColLoop->addIncoming(V: NewVecD, BB: ColLoopLatch);
480
481 return NewVecD;
482}
483
484template <Intrinsic::ID IntrID>
485std::enable_if_t<IntrID == Intrinsic::x86_tdpbssd_internal ||
486 IntrID == Intrinsic::x86_tdpbsud_internal ||
487 IntrID == Intrinsic::x86_tdpbusd_internal ||
488 IntrID == Intrinsic::x86_tdpbuud_internal ||
489 IntrID == Intrinsic::x86_tdpbf16ps_internal,
490 bool>
491X86LowerAMXIntrinsics::lowerTileDP(Instruction *TileDP) {
492 Value *M, *N, *K, *C, *A, *B;
493 match(TileDP, m_Intrinsic<IntrID>(m_Value(V&: M), m_Value(V&: N), m_Value(V&: K),
494 m_Value(V&: C), m_Value(V&: A), m_Value(V&: B)));
495 Instruction *InsertI = TileDP;
496 IRBuilder<> PreBuilder(TileDP);
497 PreBuilder.SetInsertPoint(TileDP);
498 // We visit the loop with (m, n/4, k/4):
499 // %n_dword = lshr i16 %n, 2
500 // %k_dword = lshr i16 %k, 2
501 Value *NDWord = PreBuilder.CreateLShr(LHS: N, RHS: PreBuilder.getInt16(C: 2));
502 Value *KDWord = PreBuilder.CreateLShr(LHS: K, RHS: PreBuilder.getInt16(C: 2));
503 BasicBlock *Start = InsertI->getParent();
504 BasicBlock *End =
505 SplitBlock(Old: InsertI->getParent(), SplitPt: InsertI, DTU: &DTU, LI, MSSAU: nullptr, BBName: "continue");
506 IRBuilder<> Builder(TileDP);
507 Value *ResVec = createTileDPLoops<IntrID>(Start, End, Builder, M, NDWord,
508 KDWord, C, A, B);
509 // we cannot assume there always be bitcast after tiledpbssd. So we need to
510 // insert one bitcast as required
511 Builder.SetInsertPoint(End->getFirstNonPHIIt());
512 Value *ResAMX =
513 Builder.CreateBitCast(V: ResVec, DestTy: Type::getX86_AMXTy(C&: Builder.getContext()));
514 // Delete TileDP intrinsic and do some clean-up.
515 for (Use &U : llvm::make_early_inc_range(Range: TileDP->uses())) {
516 Instruction *I = cast<Instruction>(Val: U.getUser());
517 Value *Vec;
518 if (match(V: I, P: m_BitCast(Op: m_Value(V&: Vec)))) {
519 I->replaceAllUsesWith(V: ResVec);
520 I->eraseFromParent();
521 }
522 }
523 TileDP->replaceAllUsesWith(V: ResAMX);
524 TileDP->eraseFromParent();
525 return true;
526}
527
528template <bool IsTileLoad>
529bool X86LowerAMXIntrinsics::lowerTileLoadStore(Instruction *TileLoadStore) {
530 Value *M, *N, *Ptr, *Stride, *Tile;
531 if (IsTileLoad)
532 match(V: TileLoadStore,
533 P: m_Intrinsic<Intrinsic::x86_tileloadd64_internal>(
534 Ops: m_Value(V&: M), Ops: m_Value(V&: N), Ops: m_Value(V&: Ptr), Ops: m_Value(V&: Stride)));
535 else
536 match(V: TileLoadStore, P: m_Intrinsic<Intrinsic::x86_tilestored64_internal>(
537 Ops: m_Value(V&: M), Ops: m_Value(V&: N), Ops: m_Value(V&: Ptr),
538 Ops: m_Value(V&: Stride), Ops: m_Value(V&: Tile)));
539
540 Instruction *InsertI = TileLoadStore;
541 IRBuilder<> PreBuilder(TileLoadStore);
542 PreBuilder.SetInsertPoint(TileLoadStore);
543 Value *NDWord = PreBuilder.CreateLShr(LHS: N, RHS: PreBuilder.getInt16(C: 2));
544 Value *StrideDWord = PreBuilder.CreateLShr(LHS: Stride, RHS: PreBuilder.getInt64(C: 2));
545 BasicBlock *Start = InsertI->getParent();
546 BasicBlock *End =
547 SplitBlock(Old: InsertI->getParent(), SplitPt: InsertI, DTU: &DTU, LI, MSSAU: nullptr, BBName: "continue");
548 IRBuilder<> Builder(TileLoadStore);
549 Value *ResVec = createTileLoadStoreLoops<IsTileLoad>(
550 Start, End, Builder, M, NDWord, Ptr, StrideDWord,
551 IsTileLoad ? nullptr : Tile);
552 if (IsTileLoad) {
553 // we cannot assume there always be bitcast after tileload. So we need to
554 // insert one bitcast as required
555 Builder.SetInsertPoint(End->getFirstNonPHIIt());
556 Value *ResAMX =
557 Builder.CreateBitCast(V: ResVec, DestTy: Type::getX86_AMXTy(C&: Builder.getContext()));
558 // Delete tileloadd6 intrinsic and do some clean-up
559 for (Use &U : llvm::make_early_inc_range(Range: TileLoadStore->uses())) {
560 Instruction *I = cast<Instruction>(Val: U.getUser());
561 Value *Vec;
562 if (match(V: I, P: m_BitCast(Op: m_Value(V&: Vec)))) {
563 I->replaceAllUsesWith(V: ResVec);
564 I->eraseFromParent();
565 }
566 }
567 TileLoadStore->replaceAllUsesWith(V: ResAMX);
568 }
569 TileLoadStore->eraseFromParent();
570 return true;
571}
572
573bool X86LowerAMXIntrinsics::lowerTileZero(Instruction *TileZero) {
574 IRBuilder<> Builder(TileZero);
575 FixedVectorType *V256I32Ty = FixedVectorType::get(ElementType: Builder.getInt32Ty(), NumElts: 256);
576 Value *VecZero = Constant::getNullValue(Ty: V256I32Ty);
577 for (Use &U : llvm::make_early_inc_range(Range: TileZero->uses())) {
578 Instruction *I = cast<Instruction>(Val: U.getUser());
579 Value *Vec;
580 if (match(V: I, P: m_BitCast(Op: m_Value(V&: Vec)))) {
581 I->replaceAllUsesWith(V: VecZero);
582 I->eraseFromParent();
583 }
584 }
585 TileZero->eraseFromParent();
586 return true;
587}
588
589bool X86LowerAMXIntrinsics::visit() {
590 bool C = false;
591 SmallVector<IntrinsicInst *, 8> WorkList;
592 for (BasicBlock *BB : depth_first(G: &Func)) {
593 for (BasicBlock::iterator II = BB->begin(), IE = BB->end(); II != IE;) {
594 if (auto *Inst = dyn_cast<IntrinsicInst>(Val: &*II++)) {
595 switch (Inst->getIntrinsicID()) {
596 case Intrinsic::x86_tdpbssd_internal:
597 case Intrinsic::x86_tdpbsud_internal:
598 case Intrinsic::x86_tdpbusd_internal:
599 case Intrinsic::x86_tdpbuud_internal:
600 case Intrinsic::x86_tileloadd64_internal:
601 case Intrinsic::x86_tilestored64_internal:
602 case Intrinsic::x86_tilezero_internal:
603 case Intrinsic::x86_tdpbf16ps_internal:
604 WorkList.push_back(Elt: Inst);
605 break;
606 default:
607 break;
608 }
609 }
610 }
611 }
612
613 for (auto *Inst : WorkList) {
614 switch (Inst->getIntrinsicID()) {
615 case Intrinsic::x86_tdpbssd_internal:
616 C = lowerTileDP<Intrinsic::x86_tdpbssd_internal>(TileDP: Inst) || C;
617 break;
618 case Intrinsic::x86_tdpbsud_internal:
619 C = lowerTileDP<Intrinsic::x86_tdpbsud_internal>(TileDP: Inst) || C;
620 break;
621 case Intrinsic::x86_tdpbusd_internal:
622 C = lowerTileDP<Intrinsic::x86_tdpbusd_internal>(TileDP: Inst) || C;
623 break;
624 case Intrinsic::x86_tdpbuud_internal:
625 C = lowerTileDP<Intrinsic::x86_tdpbuud_internal>(TileDP: Inst) || C;
626 break;
627 case Intrinsic::x86_tdpbf16ps_internal:
628 C = lowerTileDP<Intrinsic::x86_tdpbf16ps_internal>(TileDP: Inst) || C;
629 break;
630 case Intrinsic::x86_tileloadd64_internal:
631 C = lowerTileLoadStore<true>(TileLoadStore: Inst) || C;
632 break;
633 case Intrinsic::x86_tilestored64_internal:
634 C = lowerTileLoadStore<false>(TileLoadStore: Inst) || C;
635 break;
636 case Intrinsic::x86_tilezero_internal:
637 C = lowerTileZero(TileZero: Inst) || C;
638 break;
639 default:
640 llvm_unreachable("invalid amx intrinsics!");
641 }
642 }
643
644 return C;
645}
646
647namespace {
648bool shouldRunLowerAMXIntrinsics(const Function &F, const TargetMachine *TM) {
649 const X86Options &CLOpts =
650 static_cast<const X86TargetMachine *>(TM)->getCLOpts();
651 return CLOpts.enable_x86_scalar_amx &&
652 (F.hasFnAttribute(Kind: Attribute::OptimizeNone) ||
653 TM->getOptLevel() == CodeGenOptLevel::None);
654}
655
656bool runLowerAMXIntrinsics(Function &F, DominatorTree *DT, LoopInfo *LI) {
657 DomTreeUpdater DTU(DT, DomTreeUpdater::UpdateStrategy::Lazy);
658
659 X86LowerAMXIntrinsics LAT(F, DTU, LI);
660 return LAT.visit();
661}
662} // namespace
663
664PreservedAnalyses X86LowerAMXIntrinsicsPass::run(Function &F,
665 FunctionAnalysisManager &FAM) {
666 if (!shouldRunLowerAMXIntrinsics(F, TM))
667 return PreservedAnalyses::all();
668
669 DominatorTree &DT = FAM.getResult<DominatorTreeAnalysis>(IR&: F);
670 LoopInfo &LI = FAM.getResult<LoopAnalysis>(IR&: F);
671 bool Changed = runLowerAMXIntrinsics(F, DT: &DT, LI: &LI);
672 if (!Changed)
673 return PreservedAnalyses::all();
674
675 PreservedAnalyses PA = PreservedAnalyses::none();
676 PA.preserve<DominatorTreeAnalysis>();
677 PA.preserve<LoopAnalysis>();
678 return PA;
679}
680
681namespace {
682class X86LowerAMXIntrinsicsLegacyPass : public FunctionPass {
683public:
684 static char ID;
685
686 X86LowerAMXIntrinsicsLegacyPass() : FunctionPass(ID) {}
687
688 bool runOnFunction(Function &F) override {
689 TargetMachine *TM = &getAnalysis<TargetPassConfig>().getTM<TargetMachine>();
690 if (!shouldRunLowerAMXIntrinsics(F, TM))
691 return false;
692
693 auto *DTWP = getAnalysisIfAvailable<DominatorTreeWrapperPass>();
694 auto *DT = DTWP ? &DTWP->getDomTree() : nullptr;
695 auto *LIWP = getAnalysisIfAvailable<LoopInfoWrapperPass>();
696 auto *LI = LIWP ? &LIWP->getLoopInfo() : nullptr;
697 return runLowerAMXIntrinsics(F, DT, LI);
698 }
699 StringRef getPassName() const override { return "Lower AMX intrinsics"; }
700
701 void getAnalysisUsage(AnalysisUsage &AU) const override {
702 AU.addPreserved<DominatorTreeWrapperPass>();
703 AU.addPreserved<LoopInfoWrapperPass>();
704 AU.addRequired<TargetPassConfig>();
705 }
706};
707} // namespace
708
709static const char PassName[] = "Lower AMX intrinsics";
710char X86LowerAMXIntrinsicsLegacyPass::ID = 0;
711INITIALIZE_PASS_BEGIN(X86LowerAMXIntrinsicsLegacyPass, DEBUG_TYPE, PassName,
712 false, false)
713INITIALIZE_PASS_DEPENDENCY(TargetPassConfig)
714INITIALIZE_PASS_END(X86LowerAMXIntrinsicsLegacyPass, DEBUG_TYPE, PassName,
715 false, false)
716
717FunctionPass *llvm::createX86LowerAMXIntrinsicsLegacyPass() {
718 return new X86LowerAMXIntrinsicsLegacyPass();
719}
720