1//===- LowerMatrixIntrinsics.cpp - Lower matrix intrinsics -----*- 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// Lower matrix intrinsics to vector operations.
10//
11// TODO:
12// * Improve fusion:
13// * Support more cases, e.g. multiply-add, multiply-sub, operands/results
14// transposed.
15// * Improve cost-modeling, e.g. choose different number of rows/columns
16// columns for tiles, consider cost of copies on alias.
17//
18//===----------------------------------------------------------------------===//
19
20#include "llvm/Transforms/Scalar/LowerMatrixIntrinsics.h"
21#include "ScalarOptions.h"
22#include "llvm/ADT/PostOrderIterator.h"
23#include "llvm/ADT/STLExtras.h"
24#include "llvm/ADT/ScopeExit.h"
25#include "llvm/ADT/SmallVector.h"
26#include "llvm/ADT/Statistic.h"
27#include "llvm/Analysis/AliasAnalysis.h"
28#include "llvm/Analysis/DomTreeUpdater.h"
29#include "llvm/Analysis/LoopInfo.h"
30#include "llvm/Analysis/OptimizationRemarkEmitter.h"
31#include "llvm/Analysis/TargetTransformInfo.h"
32#include "llvm/Analysis/ValueTracking.h"
33#include "llvm/Analysis/VectorUtils.h"
34#include "llvm/IR/CFG.h"
35#include "llvm/IR/DataLayout.h"
36#include "llvm/IR/DebugInfoMetadata.h"
37#include "llvm/IR/DerivedTypes.h"
38#include "llvm/IR/Function.h"
39#include "llvm/IR/IRBuilder.h"
40#include "llvm/IR/InstrTypes.h"
41#include "llvm/IR/Instructions.h"
42#include "llvm/IR/IntrinsicInst.h"
43#include "llvm/IR/MatrixBuilder.h"
44#include "llvm/IR/PatternMatch.h"
45#include "llvm/IR/ProfDataUtils.h"
46#include "llvm/Support/Alignment.h"
47#include "llvm/Support/CommandLine.h"
48#include "llvm/Support/Compiler.h"
49#include "llvm/Support/Debug.h"
50#include "llvm/Transforms/Utils/BasicBlockUtils.h"
51#include "llvm/Transforms/Utils/LoopUtils.h"
52#include "llvm/Transforms/Utils/MatrixUtils.h"
53
54#include <cmath>
55
56using namespace llvm;
57using namespace PatternMatch;
58
59#define DEBUG_TYPE "lower-matrix-intrinsics"
60
61STATISTIC(FlattenedMatrices, "Number of matrix flattenings");
62STATISTIC(ReshapedMatrices, "Number of matrix reshapes");
63STATISTIC(SplitMatrices, "Number of matrix splits");
64
65namespace llvm {
66extern cl::opt<bool> ProfcheckDisableMetadataFixes;
67} // end namespace llvm
68
69/// Helper function to either return Scope, if it is a subprogram or the
70/// attached subprogram for a local scope.
71static DISubprogram *getSubprogram(DIScope *Scope) {
72 if (auto *Subprogram = dyn_cast<DISubprogram>(Val: Scope))
73 return Subprogram;
74 return cast<DILocalScope>(Val: Scope)->getSubprogram();
75}
76
77/// Return true if V is a splat of a value (which is used when multiplying a
78/// matrix with a scalar).
79static bool isSplat(Value *V) {
80 if (auto *SV = dyn_cast<ShuffleVectorInst>(Val: V))
81 return SV->isZeroEltSplat();
82 return false;
83}
84
85/// Match any mul operation (fp or integer).
86template <typename LTy, typename RTy>
87static auto m_AnyMul(const LTy &L, const RTy &R) {
88 return m_CombineOr(m_Mul(L, R), m_FMul(L, R));
89}
90
91/// Match any add operation (fp or integer).
92template <typename LTy, typename RTy>
93static auto m_AnyAdd(const LTy &L, const RTy &R) {
94 return m_CombineOr(m_Add(L, R), m_FAdd(L, R));
95}
96
97// Given an element pointer \p BasePtr to the start of a (sub) matrix, compute
98// the start address of vector \p VecIdx with type (\p EltType x \p NumElements)
99// assuming \p Stride elements between start two consecutive vectors.
100// \p Stride must be >= \p NumElements.
101// For column-major matrixes, the function computes the address of a column
102// vectors and \p NumElements must be set to the number of elements in a column
103// (= number of rows of the matrix). For row-major matrixes, the function
104// computes the address of a row vector and \p NumElements must be set to the
105// number of elements in a column (= number of columns of the matrix).
106//
107// Consider a 4x4 matrix in column-mjaor layout like below
108//
109// 0 1 2 3
110// 0 v_0_0 v_0_1 v_0_2 v_0_3
111// 1 v_1_0 v_1_1 v_1_2 v_1_3
112// 2 v_2_0 v_2_1 v_2_2 v_2_3
113// 3 v_3_0 v_3_1 v_3_2 v_3_3
114
115// To compute the column addresses for a 2x3 sub-matrix at row 1 and column 1,
116// we need a pointer to the first element of the submatrix as base pointer.
117// Then we can use computeVectorAddr to compute the addresses for the columns
118// of the sub-matrix.
119//
120// Column 0: computeVectorAddr(Base, 0 (column), 4 (stride), 2 (num rows), ..)
121// -> just returns Base
122// Column 1: computeVectorAddr(Base, 1 (column), 4 (stride), 2 (num rows), ..)
123// -> returns Base + (1 * 4)
124// Column 2: computeVectorAddr(Base, 2 (column), 4 (stride), 2 (num rows), ..)
125// -> returns Base + (2 * 4)
126//
127// The graphic below illustrates the number of elements in a column (marked
128// with |) and the number of skipped elements (marked with }).
129//
130// v_0_0 v_0_1 {v_0_2 {v_0_3
131// Base Col 1 Col 2
132// | | |
133// v_1_0 |v_1_1 |v_1_2 |v_1_3
134// v_2_0 |v_2_1 |v_2_2 |v_2_3
135// v_3_0 {v_3_1 {v_3_2 v_3_3
136//
137static Value *computeVectorAddr(Value *BasePtr, Value *VecIdx, Value *Stride,
138 unsigned NumElements, Type *EltType,
139 IRBuilder<> &Builder) {
140
141 assert((!isa<ConstantInt>(Stride) ||
142 cast<ConstantInt>(Stride)->getZExtValue() >= NumElements) &&
143 "Stride must be >= the number of elements in the result vector.");
144
145 // Compute the start of the vector with index VecIdx as VecIdx * Stride.
146 Value *VecStart = Builder.CreateMul(LHS: VecIdx, RHS: Stride, Name: "vec.start");
147
148 // Get pointer to the start of the selected vector. Skip GEP creation,
149 // if we select vector 0.
150 if (isa<ConstantInt>(Val: VecStart) && cast<ConstantInt>(Val: VecStart)->isZero())
151 VecStart = BasePtr;
152 else
153 VecStart = Builder.CreateInBoundsGEP(Ty: EltType, Ptr: BasePtr, IdxList: VecStart, Name: "vec.gep");
154
155 return VecStart;
156}
157
158namespace {
159struct ShapeInfo {
160 unsigned NumRows;
161 unsigned NumColumns;
162
163 bool IsColumnMajor;
164
165 ShapeInfo(unsigned NumRows = 0, unsigned NumColumns = 0)
166 : NumRows(NumRows), NumColumns(NumColumns),
167 IsColumnMajor(ScalarOptions::Global.matrix_default_layout ==
168 MatrixLayoutTy::ColumnMajor) {}
169
170 ShapeInfo(Value *NumRows, Value *NumColumns)
171 : ShapeInfo(cast<ConstantInt>(Val: NumRows)->getZExtValue(),
172 cast<ConstantInt>(Val: NumColumns)->getZExtValue()) {}
173
174 bool operator==(const ShapeInfo &other) {
175 return NumRows == other.NumRows && NumColumns == other.NumColumns;
176 }
177 bool operator!=(const ShapeInfo &other) { return !(*this == other); }
178
179 /// Returns true if shape-information is defined, meaning both dimensions
180 /// are != 0.
181 operator bool() const {
182 assert(NumRows == 0 || NumColumns != 0);
183 return NumRows != 0;
184 }
185
186 unsigned getStride() const {
187 if (IsColumnMajor)
188 return NumRows;
189 return NumColumns;
190 }
191
192 unsigned getNumVectors() const {
193 if (IsColumnMajor)
194 return NumColumns;
195 return NumRows;
196 }
197
198 /// Returns the transposed shape.
199 ShapeInfo t() const { return ShapeInfo(NumColumns, NumRows); }
200
201 friend raw_ostream &operator<<(raw_ostream &OS, ShapeInfo SI);
202
203 LLVM_DUMP_METHOD void dump() const { dbgs() << *this << '\n'; }
204};
205
206raw_ostream &operator<<(raw_ostream &OS, ShapeInfo SI) {
207 return OS << SI.NumRows << 'x' << SI.NumColumns;
208}
209
210} // namespace
211
212static bool isShapePreserving(Value *V) {
213 Instruction *I = dyn_cast<Instruction>(Val: V);
214 if (!I)
215 return true;
216
217 if (isa<SelectInst>(Val: I))
218 return true;
219
220 if (I->isBinaryOp())
221 return true;
222
223 if (auto *Cast = dyn_cast<CastInst>(Val: V)) {
224 switch (Cast->getOpcode()) {
225 case llvm::Instruction::Trunc:
226 case llvm::Instruction::ZExt:
227 case llvm::Instruction::SExt:
228 case llvm::Instruction::FPToUI:
229 case llvm::Instruction::FPToSI:
230 case llvm::Instruction::UIToFP:
231 case llvm::Instruction::SIToFP:
232 case llvm::Instruction::FPTrunc:
233 case llvm::Instruction::FPExt:
234 return true;
235 case llvm::Instruction::AddrSpaceCast:
236 case CastInst::PtrToAddr:
237 case CastInst::PtrToInt:
238 case CastInst::IntToPtr:
239 return false;
240 case CastInst::BitCast: {
241 if (auto *SrcVTy = dyn_cast<FixedVectorType>(Val: Cast->getSrcTy()))
242 if (auto *DestVTy = dyn_cast<FixedVectorType>(Val: Cast->getDestTy()))
243 return SrcVTy->getNumElements() == DestVTy->getNumElements();
244 return false;
245 }
246 case llvm::Instruction::CastOpsEnd:
247 llvm_unreachable("not an actual cast op");
248 }
249 llvm_unreachable("unhandled cast opcode");
250 }
251
252 if (auto *II = dyn_cast<IntrinsicInst>(Val: V))
253 switch (II->getIntrinsicID()) {
254 case Intrinsic::abs:
255 case Intrinsic::fabs:
256 return true;
257 default:
258 return false;
259 }
260
261 switch (I->getOpcode()) {
262 case Instruction::PHI:
263 case Instruction::FNeg:
264 return true;
265 default:
266 return false;
267 }
268}
269
270/// Return an iterator over the operands of \p I that should share shape
271/// information with \p I.
272static iterator_range<Use *> getShapedOperandsForInst(Instruction *I) {
273 assert(isShapePreserving(I) &&
274 "Can't retrieve shaped operands for an instruction that does not "
275 "preserve shape information");
276 auto Ops = I->operands();
277 return isa<SelectInst>(Val: I) ? drop_begin(RangeOrContainer&: Ops) : Ops;
278}
279
280/// Return the ShapeInfo for the result of \p I, it it can be determined.
281static std::optional<ShapeInfo>
282computeShapeInfoForInst(Instruction *I,
283 const DenseMap<Value *, ShapeInfo> &ShapeMap) {
284 Value *M;
285 Value *N;
286 Value *K;
287 if (match(V: I, P: m_Intrinsic<Intrinsic::matrix_multiply>(
288 Ops: m_Value(), Ops: m_Value(), Ops: m_Value(V&: M), Ops: m_Value(V&: N), Ops: m_Value(V&: K))))
289 return ShapeInfo(M, K);
290 if (match(V: I, P: m_Intrinsic<Intrinsic::matrix_transpose>(Ops: m_Value(), Ops: m_Value(V&: M),
291 Ops: m_Value(V&: N)))) {
292 // Flip dimensions.
293 return ShapeInfo(N, M);
294 }
295 if (match(V: I, P: m_Intrinsic<Intrinsic::matrix_column_major_store>(
296 Ops: m_Value(), Ops: m_Value(), Ops: m_Value(), Ops: m_Value(), Ops: m_Value(V&: M),
297 Ops: m_Value(V&: N))))
298 return ShapeInfo(N, M);
299 if (match(V: I, P: m_Intrinsic<Intrinsic::matrix_column_major_load>(
300 Ops: m_Value(), Ops: m_Value(), Ops: m_Value(), Ops: m_Value(V&: M), Ops: m_Value(V&: N))))
301 return ShapeInfo(M, N);
302 Value *MatrixA;
303 if (match(V: I, P: m_Store(ValueOp: m_Value(V&: MatrixA), PointerOp: m_Value()))) {
304 auto OpShape = ShapeMap.find(Val: MatrixA);
305 if (OpShape != ShapeMap.end())
306 return OpShape->second;
307 }
308
309 if (isShapePreserving(V: I)) {
310 auto ShapedOps = getShapedOperandsForInst(I);
311 // Find the first operand that has a known shape and use that.
312 for (auto &Op : ShapedOps) {
313 auto OpShape = ShapeMap.find(Val: Op.get());
314 if (OpShape != ShapeMap.end())
315 return OpShape->second;
316 }
317 }
318 return std::nullopt;
319}
320
321namespace {
322
323/// LowerMatrixIntrinsics contains the methods used to lower matrix intrinsics.
324///
325/// Currently, the lowering for each matrix intrinsic is done as follows:
326/// 1. Propagate the shape information from intrinsics to connected
327/// instructions.
328/// 2. Lower instructions with shape information (assuming column-major layout).
329/// The lowering works similarly using row-major layout.
330/// 2.1. Get column vectors for each argument. If we already lowered the
331/// definition of an argument, use the produced column vectors directly.
332/// If not, split the operand vector containing an embedded matrix into
333/// a set of column vectors,
334/// 2.2. Lower the instruction in terms of column major operations, which
335/// yields a set of column vectors containing result matrix. Note that we
336/// lower all instructions that have shape information. Besides the
337/// intrinsics, this includes stores for example.
338/// 2.3. Update uses of the lowered instruction. If we have shape information
339/// for a user, there is nothing to do, as we will look up the result
340/// column matrix when lowering the user. For other uses, we embed the
341/// result matrix in a flat vector and update the use.
342/// 2.4. Cache the result column matrix for the instruction we lowered
343/// 3. After we lowered all instructions in a function, remove the now
344/// obsolete instructions.
345///
346class LowerMatrixIntrinsics {
347 const ScalarOptions &Opts;
348 Function &Func;
349 const DataLayout &DL;
350 const TargetTransformInfo &TTI;
351 FunctionAnalysisManager *AM;
352 AliasAnalysis *AA = nullptr;
353 DominatorTree *DT = nullptr;
354 LoopInfo *LI = nullptr;
355 OptimizationRemarkEmitter *ORE = nullptr;
356
357 /// Contains estimates of the number of operations (loads, stores, compute)
358 /// required to lower a matrix operation.
359 struct OpInfoTy {
360 /// Number of stores emitted to generate this matrix.
361 unsigned NumStores = 0;
362 /// Number of loads emitted to generate this matrix.
363 unsigned NumLoads = 0;
364 /// Number of compute operations emitted to generate this matrix.
365 unsigned NumComputeOps = 0;
366 /// Most of the time transposes can be fused with matrix multiplies or can
367 /// be folded away via algebraic simplifications. This is the number of
368 /// transposes that we failed to make "free" via such optimizations.
369 unsigned NumExposedTransposes = 0;
370
371 OpInfoTy &operator+=(const OpInfoTy &RHS) {
372 NumStores += RHS.NumStores;
373 NumLoads += RHS.NumLoads;
374 NumComputeOps += RHS.NumComputeOps;
375 NumExposedTransposes += RHS.NumExposedTransposes;
376 return *this;
377 }
378 };
379
380 /// Wrapper class representing a matrix as a set of vectors, either in row or
381 /// column major layout. All vectors must have the same vector type.
382 class MatrixTy {
383 SmallVector<Value *, 16> Vectors;
384
385 OpInfoTy OpInfo;
386
387 bool IsColumnMajor = ScalarOptions::Global.matrix_default_layout ==
388 MatrixLayoutTy::ColumnMajor;
389
390 public:
391 MatrixTy() = default;
392 MatrixTy(ArrayRef<Value *> Vectors) : Vectors(Vectors) {}
393 MatrixTy(unsigned NumRows, unsigned NumColumns, Type *EltTy) {
394
395 unsigned D = isColumnMajor() ? NumColumns : NumRows;
396 for (unsigned J = 0; J < D; ++J)
397 addVector(V: PoisonValue::get(T: FixedVectorType::get(
398 ElementType: EltTy, NumElts: isColumnMajor() ? NumRows : NumColumns)));
399 }
400
401 Value *getVector(unsigned i) const { return Vectors[i]; }
402 Value *getColumn(unsigned i) const {
403 assert(isColumnMajor() && "only supported for column-major matrixes");
404 return Vectors[i];
405 }
406 Value *getRow(unsigned i) const {
407 assert(!isColumnMajor() && "only supported for row-major matrixes");
408 return Vectors[i];
409 }
410
411 void setVector(unsigned i, Value *V) { Vectors[i] = V; }
412
413 Type *getElementType() const { return getVectorTy()->getElementType(); }
414
415 unsigned getNumVectors() const {
416 if (isColumnMajor())
417 return getNumColumns();
418 return getNumRows();
419 }
420
421 unsigned getNumColumns() const {
422 if (isColumnMajor())
423 return Vectors.size();
424 else {
425 assert(Vectors.size() > 0 && "Cannot call getNumRows without columns");
426 return getVectorTy()->getNumElements();
427 }
428 }
429 unsigned getNumRows() const {
430 if (isColumnMajor()) {
431 assert(Vectors.size() > 0 && "Cannot call getNumRows without columns");
432 return getVectorTy()->getNumElements();
433 } else
434 return Vectors.size();
435 }
436
437 void addVector(Value *V) { Vectors.push_back(Elt: V); }
438 FixedVectorType *getColumnTy() {
439 assert(isColumnMajor() && "only supported for column-major matrixes");
440 return getVectorTy();
441 }
442
443 FixedVectorType *getVectorTy() const {
444 return cast<FixedVectorType>(Val: Vectors[0]->getType());
445 }
446
447 iterator_range<SmallVector<Value *, 8>::iterator> columns() {
448 assert(isColumnMajor() &&
449 "columns() only supported for column-major matrixes");
450 return make_range(x: Vectors.begin(), y: Vectors.end());
451 }
452
453 iterator_range<SmallVector<Value *, 8>::iterator> vectors() {
454 return make_range(x: Vectors.begin(), y: Vectors.end());
455 }
456
457 /// Embed the vectors of the matrix into a flat vector by concatenating
458 /// them.
459 Value *embedInVector(IRBuilder<> &Builder) const {
460 return Vectors.size() == 1 ? Vectors[0]
461 : concatenateVectors(Builder, Vecs: Vectors);
462 }
463
464 MatrixTy &addNumLoads(unsigned N) {
465 OpInfo.NumLoads += N;
466 return *this;
467 }
468
469 void setNumLoads(unsigned N) { OpInfo.NumLoads = N; }
470
471 MatrixTy &addNumStores(unsigned N) {
472 OpInfo.NumStores += N;
473 return *this;
474 }
475
476 MatrixTy &addNumExposedTransposes(unsigned N) {
477 OpInfo.NumExposedTransposes += N;
478 return *this;
479 }
480
481 MatrixTy &addNumComputeOps(unsigned N) {
482 OpInfo.NumComputeOps += N;
483 return *this;
484 }
485
486 unsigned getNumStores() const { return OpInfo.NumStores; }
487 unsigned getNumLoads() const { return OpInfo.NumLoads; }
488 unsigned getNumComputeOps() const { return OpInfo.NumComputeOps; }
489
490 const OpInfoTy &getOpInfo() const { return OpInfo; }
491
492 bool isColumnMajor() const { return IsColumnMajor; }
493
494 unsigned getStride() const {
495 if (isColumnMajor())
496 return getNumRows();
497 return getNumColumns();
498 }
499
500 ShapeInfo shape() const { return {getNumRows(), getNumColumns()}; }
501
502 /// Extract a vector of \p NumElts starting at index (\p I, \p J). If the
503 /// matrix is column-major, the result vector is extracted from a column
504 /// vector, otherwise from a row vector.
505 Value *extractVector(unsigned I, unsigned J, unsigned NumElts,
506 IRBuilder<> &Builder) const {
507 Value *Vec = isColumnMajor() ? getColumn(i: J) : getRow(i: I);
508 assert(cast<FixedVectorType>(Vec->getType())->getNumElements() >=
509 NumElts &&
510 "Extracted vector will contain poison values");
511 return Builder.CreateShuffleVector(
512 V: Vec, Mask: createSequentialMask(Start: isColumnMajor() ? I : J, NumInts: NumElts, NumUndefs: 0),
513 Name: "block");
514 }
515 };
516
517 /// Maps instructions to their shape information. The shape information
518 /// describes the shape to be used while lowering. This matches the shape of
519 /// the result value of the instruction, with the only exceptions being store
520 /// instructions and the matrix_column_major_store intrinsics. For those, the
521 /// shape information indicates that those instructions should be lowered
522 /// using shape information as well. Note that extra care is needed when
523 /// erasing or RAUW'ing a value that is present in ShapeMap. If the
524 /// replacement is also a matrix operation, use
525 /// updateShapeAndReplaceAllUsesWith to make sure the replacement is added to
526 /// ShapeMap. We don't use ValueMap, as there are also cases where we do not
527 /// want to add shape information for a replacement instruction. When directly
528 /// erasing a value with an entry in ShapeMap, use
529 /// eraseFromParentAndRemoveFromShapeMap to make sure ShapeMap is also updated
530 /// accordingly.
531 DenseMap<Value *, ShapeInfo> ShapeMap;
532
533 /// List of instructions to remove. While lowering, we are not replacing all
534 /// users of a lowered instruction, if shape information is available and
535 /// those need to be removed after we finished lowering.
536 SmallVector<Instruction *, 16> ToRemove;
537
538 /// Map from instructions to their produced column matrix.
539 MapVector<Value *, MatrixTy> Inst2ColumnMatrix;
540
541private:
542 FastMathFlags getFastMathFlags(Instruction *Inst) const {
543 FastMathFlags FMF;
544
545 if (isa<FPMathOperator>(Val: *Inst))
546 FMF = Inst->getFastMathFlags();
547
548 FMF.setAllowContract(Opts.matrix_allow_contract || FMF.allowContract());
549
550 return FMF;
551 }
552
553public:
554 LowerMatrixIntrinsics(Function &F, TargetTransformInfo &TTI,
555 FunctionAnalysisManager *AM)
556 : Opts(ScalarOptions::Global), Func(F), DL(F.getDataLayout()), TTI(TTI),
557 AM(AM) {}
558
559 unsigned getNumOps(Type *VT) {
560 assert(isa<FixedVectorType>(VT) && "Expected vector type");
561 return getNumOps(ST: VT->getScalarType(),
562 N: cast<FixedVectorType>(Val: VT)->getNumElements());
563 }
564
565 /// Is this the minimal version executed in the backend pipelines.
566 bool isMinimal() const {
567 return !DT;
568 }
569
570 /// Return the estimated number of vector ops required for an operation on
571 /// \p VT * N.
572 unsigned getNumOps(Type *ST, unsigned N) {
573 return std::ceil(x: (ST->getPrimitiveSizeInBits() * N).getFixedValue() /
574 double(TTI.getRegisterBitWidth(
575 K: TargetTransformInfo::RGK_FixedWidthVector)
576 .getFixedValue()));
577 }
578
579 /// Estimate the number of native vector operations for a multiply of matrices
580 /// with dimensions \p R x \p M and \p M x \p C. Native ops are computed as
581 /// ceil(ElementCount * ElementBits / RegisterBits).
582 ///
583 /// Native vector ops per operation type (VF = native vector elements):
584 /// FMAs: C * ceil(R/VF) * M (one FMA per VF output elements)
585 /// A loads: ceil(R/VF) * M (A has M columns, ceil(R/VF) native loads each)
586 /// B loads: ceil(M/VF) * C (B has C columns, ceil(M/VF) native loads each)
587 /// Stores: C * ceil(R/VF) (one store per VF output elements)
588 unsigned getNumNativeVectorOps(Type *EltType, unsigned R, unsigned M,
589 unsigned C) {
590 unsigned NumFMAs = C * getNumOps(ST: EltType, N: R) * M;
591 unsigned NumALoads = getNumOps(ST: EltType, N: R) * M;
592 unsigned NumBLoads = getNumOps(ST: EltType, N: M) * C;
593 unsigned NumStores = getNumOps(ST: EltType, N: R) * C;
594 return NumFMAs + NumALoads + NumBLoads + NumStores;
595 }
596
597 /// Return the set of vectors that a matrix value is lowered to.
598 ///
599 /// If we lowered \p MatrixVal, just return the cache result matrix. Otherwise
600 /// split the flat vector \p MatrixVal containing a matrix with shape \p SI
601 /// into vectors.
602 MatrixTy getMatrix(Value *MatrixVal, const ShapeInfo &SI,
603 IRBuilder<> &Builder) {
604 FixedVectorType *VType = cast<FixedVectorType>(Val: MatrixVal->getType());
605 assert(VType->getNumElements() == SI.NumRows * SI.NumColumns &&
606 "The vector size must match the number of matrix elements");
607
608 // Check if we lowered MatrixVal using shape information. In that case,
609 // return the existing matrix, if it matches the requested shape
610 // information. If there is a mis-match, embed the result in a flat
611 // vector and split it later.
612 auto Found = Inst2ColumnMatrix.find(Key: MatrixVal);
613 if (Found != Inst2ColumnMatrix.end()) {
614 MatrixTy &M = Found->second;
615 // Return the found matrix, if its shape matches the requested shape
616 // information
617 if (SI.NumRows == M.getNumRows() && SI.NumColumns == M.getNumColumns())
618 return M;
619
620 MatrixVal = M.embedInVector(Builder);
621 }
622
623 // Otherwise split MatrixVal.
624 SmallVector<Value *, 16> SplitVecs;
625 for (unsigned MaskStart = 0; MaskStart < VType->getNumElements();
626 MaskStart += SI.getStride()) {
627 Value *V = Builder.CreateShuffleVector(
628 V: MatrixVal, Mask: createSequentialMask(Start: MaskStart, NumInts: SI.getStride(), NumUndefs: 0),
629 Name: "split");
630 SplitVecs.push_back(Elt: V);
631 }
632
633 if (Instruction *Inst = dyn_cast<Instruction>(Val: MatrixVal)) {
634 if (Found != Inst2ColumnMatrix.end()) {
635 // FIXME: re: "at least": SplitVecs.size() doesn't count the shuffles
636 // that embedInVector created.
637 LLVM_DEBUG(dbgs() << "matrix reshape from " << Found->second.shape()
638 << " to " << SI << " using at least "
639 << SplitVecs.size() << " shuffles on behalf of:\n"
640 << *Inst << '\n');
641 ReshapedMatrices++;
642 } else if (!ShapeMap.contains(Val: MatrixVal)) {
643 LLVM_DEBUG(
644 dbgs()
645 << "splitting a " << SI << " matrix with " << SplitVecs.size()
646 << " shuffles beacuse we do not have a shape-aware lowering for "
647 "its def:\n"
648 << *Inst << '\n');
649 (void)Inst;
650 SplitMatrices++;
651 } else {
652 // The ShapeMap has it, so it's a case where we're being lowered
653 // before the def, and we expect that InstCombine will clean things up
654 // afterward.
655 }
656 }
657
658 return {SplitVecs};
659 }
660
661 /// If \p V already has a known shape return false. Otherwise set the shape
662 /// for instructions that support it.
663 bool setShapeInfo(Value *V, ShapeInfo Shape) {
664 assert(Shape && "Shape not set");
665 if (isa<UndefValue>(Val: V) || !supportsShapeInfo(V))
666 return false;
667
668 auto SIter = ShapeMap.find(Val: V);
669 if (SIter != ShapeMap.end()) {
670 if (Opts.verify_matrix_shapes &&
671 (SIter->second.NumRows != Shape.NumRows ||
672 SIter->second.NumColumns != Shape.NumColumns)) {
673 errs() << "Conflicting shapes (" << SIter->second.NumRows << "x"
674 << SIter->second.NumColumns << " vs " << Shape.NumRows << "x"
675 << Shape.NumColumns << ") for " << *V << "\n";
676 report_fatal_error(
677 reason: "Matrix shape verification failed, compilation aborted!");
678 }
679
680 LLVM_DEBUG(dbgs() << " not overriding existing shape: "
681 << SIter->second.NumRows << " "
682 << SIter->second.NumColumns << " for " << *V << "\n");
683 return false;
684 }
685
686 ShapeMap.insert(KV: {V, Shape});
687 LLVM_DEBUG(dbgs() << " " << Shape.NumRows << " x " << Shape.NumColumns
688 << " for " << *V << "\n");
689 return true;
690 }
691
692 /// Returns true if shape information can be used for \p V. The supported
693 /// instructions must match the instructions that can be lowered by this pass.
694 bool supportsShapeInfo(Value *V) {
695 Instruction *Inst = dyn_cast<Instruction>(Val: V);
696 if (!Inst)
697 return false;
698
699 IntrinsicInst *II = dyn_cast<IntrinsicInst>(Val: Inst);
700 if (II)
701 switch (II->getIntrinsicID()) {
702 case Intrinsic::matrix_multiply:
703 case Intrinsic::matrix_transpose:
704 case Intrinsic::matrix_column_major_load:
705 case Intrinsic::matrix_column_major_store:
706 return true;
707 default:
708 break;
709 }
710 return isShapePreserving(V) || isa<StoreInst>(Val: V) || isa<LoadInst>(Val: V);
711 }
712
713 /// Propagate the shape information of instructions to their users.
714 /// The work list contains instructions for which we can compute the shape,
715 /// either based on the information provided by matrix intrinsics or known
716 /// shapes of operands.
717 SmallVector<Instruction *, 32>
718 propagateShapeForward(SmallVectorImpl<Instruction *> &WorkList) {
719 SmallVector<Instruction *, 32> NewWorkList;
720 // Pop an element for which we guaranteed to have at least one of the
721 // operand shapes. Add the shape for this and then add users to the work
722 // list.
723 LLVM_DEBUG(dbgs() << "Forward-propagate shapes:\n");
724 while (!WorkList.empty()) {
725 Instruction *Inst = WorkList.pop_back_val();
726
727 // New entry, set the value and insert operands
728 bool Propagate = false;
729 if (auto SI = computeShapeInfoForInst(I: Inst, ShapeMap))
730 Propagate = setShapeInfo(V: Inst, Shape: *SI);
731
732 if (Propagate) {
733 NewWorkList.push_back(Elt: Inst);
734 for (auto *User : Inst->users())
735 if (ShapeMap.count(Val: User) == 0)
736 WorkList.push_back(Elt: cast<Instruction>(Val: User));
737 }
738 }
739
740 return NewWorkList;
741 }
742
743 /// Propagate the shape to operands of instructions with shape information.
744 /// \p Worklist contains the instruction for which we already know the shape.
745 SmallVector<Instruction *, 32>
746 propagateShapeBackward(SmallVectorImpl<Instruction *> &WorkList) {
747 SmallVector<Instruction *, 32> NewWorkList;
748
749 auto pushInstruction = [](Value *V,
750 SmallVectorImpl<Instruction *> &WorkList) {
751 Instruction *I = dyn_cast<Instruction>(Val: V);
752 if (I)
753 WorkList.push_back(Elt: I);
754 };
755 // Pop an element with known shape. Traverse the operands, if their shape
756 // derives from the result shape and is unknown, add it and add them to the
757 // worklist.
758 LLVM_DEBUG(dbgs() << "Backward-propagate shapes:\n");
759 while (!WorkList.empty()) {
760 Value *V = WorkList.pop_back_val();
761
762 size_t BeforeProcessingV = WorkList.size();
763 if (!isa<Instruction>(Val: V))
764 continue;
765
766 Value *MatrixA;
767 Value *MatrixB;
768 Value *M;
769 Value *N;
770 Value *K;
771 if (match(V, P: m_Intrinsic<Intrinsic::matrix_multiply>(
772 Ops: m_Value(V&: MatrixA), Ops: m_Value(V&: MatrixB), Ops: m_Value(V&: M),
773 Ops: m_Value(V&: N), Ops: m_Value(V&: K)))) {
774 if (setShapeInfo(V: MatrixA, Shape: {M, N}))
775 pushInstruction(MatrixA, WorkList);
776
777 if (setShapeInfo(V: MatrixB, Shape: {N, K}))
778 pushInstruction(MatrixB, WorkList);
779
780 } else if (match(V, P: m_Intrinsic<Intrinsic::matrix_transpose>(
781 Ops: m_Value(V&: MatrixA), Ops: m_Value(V&: M), Ops: m_Value(V&: N)))) {
782 // Flip dimensions.
783 if (setShapeInfo(V: MatrixA, Shape: {M, N}))
784 pushInstruction(MatrixA, WorkList);
785 } else if (match(V, P: m_Intrinsic<Intrinsic::matrix_column_major_store>(
786 Ops: m_Value(V&: MatrixA), Ops: m_Value(), Ops: m_Value(), Ops: m_Value(),
787 Ops: m_Value(V&: M), Ops: m_Value(V&: N)))) {
788 if (setShapeInfo(V: MatrixA, Shape: {M, N})) {
789 pushInstruction(MatrixA, WorkList);
790 }
791 } else if (isa<LoadInst>(Val: V) ||
792 match(V, P: m_Intrinsic<Intrinsic::matrix_column_major_load>())) {
793 // Nothing to do, no matrix input.
794 } else if (isa<StoreInst>(Val: V)) {
795 // Nothing to do. We forward-propagated to this so we would just
796 // backward propagate to an instruction with an already known shape.
797 } else if (isShapePreserving(V)) {
798 auto ShapedOps = getShapedOperandsForInst(I: cast<Instruction>(Val: V));
799 // Propagate to all operands.
800 ShapeInfo Shape = ShapeMap[V];
801 for (Use &U : ShapedOps) {
802 if (setShapeInfo(V: U.get(), Shape))
803 pushInstruction(U.get(), WorkList);
804 }
805 }
806 // After we discovered new shape info for new instructions in the
807 // worklist, we use their users as seeds for the next round of forward
808 // propagation.
809 for (size_t I = BeforeProcessingV; I != WorkList.size(); I++)
810 for (User *U : WorkList[I]->users())
811 if (isa<Instruction>(Val: U) && V != U)
812 NewWorkList.push_back(Elt: cast<Instruction>(Val: U));
813 }
814 return NewWorkList;
815 }
816
817 /// (Op0 op Op1)^T -> Op0^T op Op1^T
818 /// Transpose \p Op0 and \p Op1 of shape \p Shape0 and \p Shape1, then use
819 /// them on both sides of \p Operation.
820 Instruction *distributeTransposes(
821 Value *Op0, ShapeInfo Shape0, Value *Op1, ShapeInfo Shape1,
822 MatrixBuilder &Builder,
823 function_ref<Instruction *(Value *, ShapeInfo, Value *, ShapeInfo)>
824 Operation) {
825 Value *T0 = Builder.CreateMatrixTranspose(
826 Matrix: Op0, Rows: Shape0.NumRows, Columns: Shape0.NumColumns, Name: Op0->getName() + "_t");
827 // We are being run after shape prop, add shape for newly created
828 // instructions so that we lower them later.
829 setShapeInfo(V: T0, Shape: Shape0.t());
830 Value *T1 = Builder.CreateMatrixTranspose(
831 Matrix: Op1, Rows: Shape1.NumRows, Columns: Shape1.NumColumns, Name: Op1->getName() + "_t");
832 setShapeInfo(V: T1, Shape: Shape1.t());
833 return Operation(T0, Shape0.t(), T1, Shape1.t());
834 }
835
836 /// Erase \p Inst from both ShapeMap (if an entry exists) and erase \p Inst
837 /// itself.
838 void eraseFromParentAndRemoveFromShapeMap(Instruction *Inst) {
839 ShapeMap.erase(Val: Inst);
840 Inst->eraseFromParent();
841 }
842
843 /// Erase \p V from \p BB and move \II forward to avoid invalidating
844 /// iterators.
845 void eraseFromParentAndMove(Value *V, BasicBlock::reverse_iterator &II,
846 BasicBlock &BB) {
847 auto *Inst = cast<Instruction>(Val: V);
848 // Still used, don't erase.
849 if (!Inst->use_empty())
850 return;
851 if (II != BB.rend() && Inst == &*II)
852 ++II;
853 eraseFromParentAndRemoveFromShapeMap(Inst);
854 }
855
856 /// Add a new entry to ShapeMap for \p New with \p Old's shape info, erase the
857 /// entry for \p Old and replace all uses of \p Old with \p New.
858 void updateShapeAndReplaceAllUsesWith(Instruction &Old, Value *New) {
859 // We need to remove Old from the ShapeMap otherwise RAUW will replace it
860 // with New. We should only add New it it supportsShapeInfo so we insert
861 // it conditionally instead.
862 auto S = ShapeMap.find(Val: &Old);
863 if (S != ShapeMap.end()) {
864 ShapeInfo Shape = S->second;
865 ShapeMap.erase(I: S);
866 if (supportsShapeInfo(V: New))
867 ShapeMap.insert(KV: {New, Shape});
868 }
869 Old.replaceAllUsesWith(V: New);
870 }
871
872 /// Sink a top-level transpose inside matmuls and adds.
873 /// This creates and erases instructions as needed, and returns the newly
874 /// created instruction while updating the iterator to avoid invalidation. If
875 /// this returns nullptr, no new instruction was created.
876 Instruction *sinkTranspose(Instruction &I, BasicBlock::reverse_iterator &II,
877 bool &Changed) {
878 BasicBlock &BB = *I.getParent();
879 IRBuilder<> IB(&I);
880 MatrixBuilder Builder(IB);
881
882 Value *TA, *TAMA, *TAMB;
883 ConstantInt *R, *K, *C;
884 if (!match(V: &I, P: m_Intrinsic<Intrinsic::matrix_transpose>(
885 Ops: m_Value(V&: TA), Ops: m_ConstantInt(CI&: R), Ops: m_ConstantInt(CI&: C))))
886 return nullptr;
887
888 // Transpose of a transpose is a nop when the shapes match.
889 Value *TATA;
890 if (match(V: TA, P: m_Intrinsic<Intrinsic::matrix_transpose>(
891 Ops: m_Value(V&: TATA), Ops: m_Specific(V: C), Ops: m_Specific(V: R)))) {
892 updateShapeAndReplaceAllUsesWith(Old&: I, New: TATA);
893 eraseFromParentAndMove(V: &I, II, BB);
894 eraseFromParentAndMove(V: TA, II, BB);
895 Changed = true;
896 return nullptr;
897 }
898
899 // k^T -> k
900 if (isSplat(V: TA)) {
901 updateShapeAndReplaceAllUsesWith(Old&: I, New: TA);
902 eraseFromParentAndMove(V: &I, II, BB);
903 Changed = true;
904 return nullptr;
905 }
906
907 // (A * B)^t -> B^t * A^t
908 // RxK KxC CxK KxR
909 if (match(V: TA, P: m_Intrinsic<Intrinsic::matrix_multiply>(
910 Ops: m_Value(V&: TAMA), Ops: m_Value(V&: TAMB), Ops: m_ConstantInt(CI&: R),
911 Ops: m_ConstantInt(CI&: K), Ops: m_ConstantInt(CI&: C)))) {
912 auto NewInst = distributeTransposes(
913 Op0: TAMB, Shape0: {K, C}, Op1: TAMA, Shape1: {R, K}, Builder,
914 Operation: [&](Value *T0, ShapeInfo Shape0, Value *T1, ShapeInfo Shape1) {
915 return Builder.CreateMatrixMultiply(LHS: T0, RHS: T1, LHSRows: Shape0.NumRows,
916 LHSColumns: Shape0.NumColumns,
917 RHSColumns: Shape1.NumColumns, Name: "mmul");
918 });
919 updateShapeAndReplaceAllUsesWith(Old&: I, New: NewInst);
920 eraseFromParentAndMove(V: &I, II, BB);
921 eraseFromParentAndMove(V: TA, II, BB);
922 Changed = true;
923 return NewInst;
924 }
925
926 // Same as above, but with a mul, which occurs when multiplied
927 // with a scalar.
928 // (A * k)^t -> A^t * k
929 // R x C RxC
930 if (match(V: TA, P: m_AnyMul(L: m_Value(V&: TAMA), R: m_Value(V&: TAMB))) &&
931 (isSplat(V: TAMA) || isSplat(V: TAMB))) {
932 IRBuilder<> LocalBuilder(&I);
933 // We know that the transposed operand is of shape RxC.
934 // An when multiplied with a scalar, the shape is preserved.
935 auto NewInst = distributeTransposes(
936 Op0: TAMA, Shape0: {R, C}, Op1: TAMB, Shape1: {R, C}, Builder,
937 Operation: [&](Value *T0, ShapeInfo Shape0, Value *T1, ShapeInfo Shape1) {
938 bool IsFP = I.getType()->isFPOrFPVectorTy();
939 auto *Mul = IsFP ? LocalBuilder.CreateFMul(L: T0, R: T1, Name: "mmul")
940 : LocalBuilder.CreateMul(LHS: T0, RHS: T1, Name: "mmul");
941 auto *Result = cast<Instruction>(Val: Mul);
942 setShapeInfo(V: Result, Shape: Shape0);
943 return Result;
944 });
945 updateShapeAndReplaceAllUsesWith(Old&: I, New: NewInst);
946 eraseFromParentAndMove(V: &I, II, BB);
947 eraseFromParentAndMove(V: TA, II, BB);
948 Changed = true;
949 return NewInst;
950 }
951
952 // (A + B)^t -> A^t + B^t
953 // RxC RxC CxR CxR
954 if (match(V: TA, P: m_AnyAdd(L: m_Value(V&: TAMA), R: m_Value(V&: TAMB)))) {
955 IRBuilder<> LocalBuilder(&I);
956 auto NewInst = distributeTransposes(
957 Op0: TAMA, Shape0: {R, C}, Op1: TAMB, Shape1: {R, C}, Builder,
958 Operation: [&](Value *T0, ShapeInfo Shape0, Value *T1, ShapeInfo Shape1) {
959 bool IsFP = I.getType()->isFPOrFPVectorTy();
960 auto *Add = IsFP ? LocalBuilder.CreateFAdd(L: T0, R: T1, Name: "madd")
961 : LocalBuilder.CreateAdd(LHS: T0, RHS: T1, Name: "madd");
962
963 auto *Result = cast<Instruction>(Val: Add);
964 setShapeInfo(V: Result, Shape: Shape0);
965 return Result;
966 });
967 updateShapeAndReplaceAllUsesWith(Old&: I, New: NewInst);
968 eraseFromParentAndMove(V: &I, II, BB);
969 eraseFromParentAndMove(V: TA, II, BB);
970 Changed = true;
971 return NewInst;
972 }
973
974 return nullptr;
975 }
976
977 bool liftTranspose(Instruction &I) {
978 // Erase dead Instructions after lifting transposes from binops.
979 auto CleanupBinOp = [this](Instruction &T, Value *A, Value *B) {
980 if (T.use_empty())
981 eraseFromParentAndRemoveFromShapeMap(Inst: &T);
982 if (A->use_empty())
983 eraseFromParentAndRemoveFromShapeMap(Inst: cast<Instruction>(Val: A));
984 if (A != B && B->use_empty())
985 eraseFromParentAndRemoveFromShapeMap(Inst: cast<Instruction>(Val: B));
986 };
987
988 Value *A, *B, *AT, *BT;
989 ConstantInt *R, *K, *C;
990 // A^t * B ^t -> (B * A)^t
991 if (match(V: &I, P: m_Intrinsic<Intrinsic::matrix_multiply>(
992 Ops: m_Value(V&: A), Ops: m_Value(V&: B), Ops: m_ConstantInt(CI&: R),
993 Ops: m_ConstantInt(CI&: K), Ops: m_ConstantInt(CI&: C))) &&
994 match(V: A, P: m_Intrinsic<Intrinsic::matrix_transpose>(Ops: m_Value(V&: AT))) &&
995 match(V: B, P: m_Intrinsic<Intrinsic::matrix_transpose>(Ops: m_Value(V&: (BT))))) {
996 IRBuilder<> IB(&I);
997 MatrixBuilder Builder(IB);
998 Value *M = Builder.CreateMatrixMultiply(
999 LHS: BT, RHS: AT, LHSRows: C->getZExtValue(), LHSColumns: K->getZExtValue(), RHSColumns: R->getZExtValue());
1000 setShapeInfo(V: M, Shape: {C, R});
1001 Instruction *NewInst = Builder.CreateMatrixTranspose(Matrix: M, Rows: C->getZExtValue(),
1002 Columns: R->getZExtValue());
1003 updateShapeAndReplaceAllUsesWith(Old&: I, New: NewInst);
1004 CleanupBinOp(I, A, B);
1005 return true;
1006 }
1007 // A^t + B ^t -> (A + B)^t. Pick rows and columns from first transpose. If
1008 // the shape of the second transpose is different, there's a shape conflict
1009 // which gets resolved by picking the shape of the first operand.
1010 else if (match(V: &I, P: m_FAdd(L: m_Value(V&: A), R: m_Value(V&: B))) &&
1011 match(V: A, P: m_Intrinsic<Intrinsic::matrix_transpose>(
1012 Ops: m_Value(V&: AT), Ops: m_ConstantInt(CI&: R), Ops: m_ConstantInt(CI&: C))) &&
1013 match(V: B, P: m_Intrinsic<Intrinsic::matrix_transpose>(
1014 Ops: m_Value(V&: BT), Ops: m_ConstantInt(), Ops: m_ConstantInt()))) {
1015 IRBuilder<> Builder(&I);
1016 auto *Add = Builder.CreateFAdd(L: AT, R: BT, Name: "mfadd");
1017 MatrixBuilder MBuilder(Builder);
1018 Instruction *NewInst = MBuilder.CreateMatrixTranspose(
1019 Matrix: Add, Rows: R->getZExtValue(), Columns: C->getZExtValue(), Name: "mfadd_t");
1020 updateShapeAndReplaceAllUsesWith(Old&: I, New: NewInst);
1021 assert(computeShapeInfoForInst(NewInst, ShapeMap) ==
1022 computeShapeInfoForInst(&I, ShapeMap) &&
1023 "Shape of new instruction doesn't match original shape.");
1024 CleanupBinOp(I, A, B);
1025 if (auto *AddI = dyn_cast<Instruction>(Val: Add)) {
1026 setShapeInfo(V: AddI, Shape: {R, C});
1027 assert(
1028 computeShapeInfoForInst(AddI, ShapeMap).value_or(ShapeMap[AddI]) ==
1029 ShapeMap[AddI] &&
1030 "Shape of updated addition doesn't match cached shape.");
1031 }
1032 return true;
1033 }
1034 return false;
1035 }
1036
1037 /// Try moving transposes in order to fold them away or into multiplies.
1038 bool optimizeTransposes() {
1039 bool Changed = false;
1040 // First sink all transposes inside matmuls and adds, hoping that we end up
1041 // with NN, NT or TN variants.
1042 for (BasicBlock &BB : reverse(C&: Func)) {
1043 for (auto II = BB.rbegin(); II != BB.rend();) {
1044 Instruction &I = *II;
1045 // We may remove II. By default continue on the next/prev instruction.
1046 ++II;
1047 if (Instruction *NewInst = sinkTranspose(I, II, Changed))
1048 II = std::next(x: BasicBlock::reverse_iterator(NewInst));
1049 }
1050 }
1051
1052 // If we have a TT matmul or a TT add, lift the transpose. We may be able
1053 // to fold into consuming multiply or add.
1054 for (BasicBlock &BB : Func) {
1055 for (Instruction &I : llvm::make_early_inc_range(Range&: BB)) {
1056 Changed |= liftTranspose(I);
1057 }
1058 }
1059 return Changed;
1060 }
1061
1062 bool Visit() {
1063 SmallVector<Instruction *, 32> WorkList;
1064
1065 // Initially only the shape of matrix intrinsics is known.
1066 // Initialize the work list with ops carrying shape information.
1067 for (BasicBlock &BB : Func)
1068 for (Instruction &Inst : BB) {
1069 IntrinsicInst *II = dyn_cast<IntrinsicInst>(Val: &Inst);
1070 if (!II)
1071 continue;
1072
1073 switch (II->getIntrinsicID()) {
1074 case Intrinsic::matrix_multiply:
1075 case Intrinsic::matrix_transpose:
1076 case Intrinsic::matrix_column_major_load:
1077 case Intrinsic::matrix_column_major_store:
1078 WorkList.push_back(Elt: &Inst);
1079 break;
1080 default:
1081 break;
1082 }
1083 }
1084
1085 // Avoid unnecessary work if there are no matrix intrinsics in the function.
1086 if (WorkList.empty())
1087 return false;
1088
1089 if (AM) {
1090 ORE = &AM->getResult<OptimizationRemarkEmitterAnalysis>(IR&: Func);
1091 AA = &AM->getResult<AAManager>(IR&: Func);
1092 DT = &AM->getResult<DominatorTreeAnalysis>(IR&: Func);
1093 LI = &AM->getResult<LoopAnalysis>(IR&: Func);
1094 }
1095
1096 // Propagate shapes until nothing changes any longer.
1097 while (!WorkList.empty()) {
1098 WorkList = propagateShapeForward(WorkList);
1099 WorkList = propagateShapeBackward(WorkList);
1100 }
1101
1102 bool Changed = false;
1103 if (!isMinimal()) {
1104 Changed |= optimizeTransposes();
1105 if (Opts.matrix_print_after_transpose_opt) {
1106 dbgs() << "Dump after matrix transpose optimization:\n";
1107 Func.print(OS&: dbgs());
1108 }
1109 }
1110
1111 SmallVector<CallInst *, 16> MaybeFusableInsts;
1112 SmallVector<Instruction *, 16> MatrixInsts;
1113 SmallVector<IntrinsicInst *, 16> LifetimeEnds;
1114
1115 // First, collect all instructions with shape information and candidates for
1116 // fusion (currently only matrix multiplies).
1117 ReversePostOrderTraversal<Function *> RPOT(&Func);
1118 for (auto *BB : RPOT)
1119 for (Instruction &I : *BB) {
1120 if (match(V: &I, P: m_Intrinsic<Intrinsic::lifetime_end>()))
1121 LifetimeEnds.push_back(Elt: cast<IntrinsicInst>(Val: &I));
1122 if (!ShapeMap.contains(Val: &I))
1123 continue;
1124 if (match(V: &I, P: m_Intrinsic<Intrinsic::matrix_multiply>()))
1125 MaybeFusableInsts.push_back(Elt: cast<CallInst>(Val: &I));
1126 MatrixInsts.push_back(Elt: &I);
1127 }
1128
1129 // Second, try to lower any dot products
1130 SmallPtrSet<Instruction *, 16> FusedInsts;
1131 for (CallInst *CI : MaybeFusableInsts)
1132 lowerDotProduct(MatMul: CI, FusedInsts, FMF: getFastMathFlags(Inst: CI));
1133
1134 // Third, try to fuse candidates.
1135 for (CallInst *CI : MaybeFusableInsts)
1136 if (!FusedInsts.contains(Ptr: CI))
1137 LowerMatrixMultiplyFused(MatMul: CI, FusedInsts, LifetimeEnds);
1138
1139 Changed |= !FusedInsts.empty();
1140
1141 // Fourth, pre-process all the PHINode's. The incoming values will be
1142 // assigned later in VisitPHI.
1143 for (Instruction *Inst : MatrixInsts) {
1144 if (FusedInsts.count(Ptr: Inst))
1145 continue;
1146
1147 auto *PHI = dyn_cast<PHINode>(Val: Inst);
1148 if (!PHI)
1149 continue;
1150
1151 const ShapeInfo &SI = ShapeMap.at(Val: Inst);
1152 auto *EltTy = cast<FixedVectorType>(Val: PHI->getType())->getElementType();
1153 MatrixTy PhiM(SI.NumRows, SI.NumColumns, EltTy);
1154
1155 IRBuilder<> Builder(Inst);
1156 for (unsigned VI = 0, VE = PhiM.getNumVectors(); VI != VE; ++VI)
1157 PhiM.setVector(i: VI, V: Builder.CreatePHI(Ty: PhiM.getVectorTy(),
1158 NumReservedValues: PHI->getNumIncomingValues(),
1159 Name: PHI->getName()));
1160 assert(!Inst2ColumnMatrix.contains(PHI) && "map already contains phi?");
1161 Inst2ColumnMatrix[PHI] = PhiM;
1162 }
1163
1164 // Fifth, lower remaining instructions with shape information.
1165 for (Instruction *Inst : MatrixInsts) {
1166 if (FusedInsts.count(Ptr: Inst))
1167 continue;
1168
1169 const ShapeInfo &SI = ShapeMap.at(Val: Inst);
1170
1171 Value *Op1;
1172 Value *Op2;
1173 MatrixTy Result;
1174 IRBuilder<> Builder(Inst);
1175 if (auto *BinOp = dyn_cast<BinaryOperator>(Val: Inst))
1176 Result = VisitBinaryOperator(Inst: BinOp, SI, Builder);
1177 else if (auto *Cast = dyn_cast<CastInst>(Val: Inst))
1178 Result = VisitCastInstruction(Inst: Cast, Shape: SI, Builder);
1179 else if (auto *UnOp = dyn_cast<UnaryOperator>(Val: Inst))
1180 Result = VisitUnaryOperator(Inst: UnOp, SI, Builder);
1181 else if (auto *Intr = dyn_cast<IntrinsicInst>(Val: Inst))
1182 Result = VisitIntrinsicInst(Inst: Intr, SI, Builder);
1183 else if (auto *Select = dyn_cast<SelectInst>(Val: Inst))
1184 Result = VisitSelectInst(Inst: Select, Shape: SI, Builder);
1185 else if (match(V: Inst, P: m_Load(Op: m_Value(V&: Op1))))
1186 Result = VisitLoad(Inst: cast<LoadInst>(Val: Inst), SI, Ptr: Op1, Builder);
1187 else if (match(V: Inst, P: m_Store(ValueOp: m_Value(V&: Op1), PointerOp: m_Value(V&: Op2))))
1188 Result = VisitStore(Inst: cast<StoreInst>(Val: Inst), SI, StoredVal: Op1, Ptr: Op2, Builder);
1189 else if (auto *PHI = dyn_cast<PHINode>(Val: Inst))
1190 Result = VisitPHI(Inst: PHI, SI, Builder);
1191 else
1192 continue;
1193
1194 finalizeLowering(Inst, Matrix: Result, Builder);
1195 Changed = true;
1196 }
1197
1198 if (ORE) {
1199 RemarkGenerator RemarkGen(Inst2ColumnMatrix, *ORE, Func);
1200 RemarkGen.emitRemarks();
1201 }
1202
1203 // Delete the instructions backwards, as it has a reduced likelihood of
1204 // having to update as many def-use and use-def chains.
1205 //
1206 // Because we add to ToRemove during fusion we can't guarantee that defs
1207 // are before uses. Change uses to poison temporarily as these should get
1208 // removed as well.
1209 //
1210 // For verification, we keep track of where we changed uses to poison in
1211 // PoisonedInsts and then check that we in fact remove them.
1212 SmallPtrSet<Instruction *, 16> PoisonedInsts;
1213 for (auto *Inst : reverse(C&: ToRemove)) {
1214 for (Use &U : llvm::make_early_inc_range(Range: Inst->uses())) {
1215 if (auto *Poisoned = dyn_cast<Instruction>(Val: U.getUser()))
1216 PoisonedInsts.insert(Ptr: Poisoned);
1217 U.set(PoisonValue::get(T: Inst->getType()));
1218 }
1219 Inst->eraseFromParent();
1220 PoisonedInsts.erase(Ptr: Inst);
1221 }
1222 if (!PoisonedInsts.empty()) {
1223 // If we didn't remove all poisoned instructions, it's a hard error.
1224 dbgs() << "Poisoned but present instructions:\n";
1225 for (auto *I : PoisonedInsts)
1226 dbgs() << *I << "\n";
1227 llvm_unreachable("Poisoned but instruction not removed");
1228 }
1229
1230 return Changed;
1231 }
1232
1233 /// Replace intrinsic calls.
1234 MatrixTy VisitIntrinsicInst(IntrinsicInst *Inst, const ShapeInfo &SI,
1235 IRBuilder<> &Builder) {
1236 assert(Inst->getCalledFunction() &&
1237 Inst->getCalledFunction()->isIntrinsic());
1238
1239 switch (Inst->getCalledFunction()->getIntrinsicID()) {
1240 case Intrinsic::matrix_multiply:
1241 return LowerMultiply(MatMul: Inst, Builder);
1242 case Intrinsic::matrix_transpose:
1243 return LowerTranspose(Inst, Builder);
1244 case Intrinsic::matrix_column_major_load:
1245 return LowerColumnMajorLoad(Inst, Builder);
1246 case Intrinsic::matrix_column_major_store:
1247 return LowerColumnMajorStore(Inst, Builder);
1248 case Intrinsic::abs:
1249 case Intrinsic::fabs: {
1250 MatrixTy Result;
1251 MatrixTy M = getMatrix(MatrixVal: Inst->getOperand(i_nocapture: 0), SI, Builder);
1252 Builder.setFastMathFlags(getFastMathFlags(Inst));
1253
1254 for (auto *Vector : M.vectors()) {
1255 switch (Inst->getIntrinsicID()) {
1256 case Intrinsic::abs:
1257 Result.addVector(V: Builder.CreateBinaryIntrinsic(ID: Intrinsic::abs, LHS: Vector,
1258 RHS: Inst->getOperand(i_nocapture: 1)));
1259 continue;
1260 case Intrinsic::fabs:
1261 Result.addVector(
1262 V: Builder.CreateUnaryIntrinsic(ID: Inst->getIntrinsicID(), Op: Vector));
1263 continue;
1264 default:
1265 llvm_unreachable("unexpected intrinsic");
1266 }
1267 }
1268
1269 return Result.addNumComputeOps(N: getNumOps(VT: Result.getVectorTy()) *
1270 Result.getNumVectors());
1271 }
1272 default:
1273 break;
1274 }
1275 llvm_unreachable(
1276 "only intrinsics supporting shape info should be seen here");
1277 }
1278
1279 /// Compute the alignment for a column/row \p Idx with \p Stride between them.
1280 /// The address at \p Idx == 0 has alignment \p A. If \p Stride is a
1281 /// ConstantInt, reduce the initial alignment based on the byte offset. For
1282 /// non-ConstantInt strides, return the common alignment of the initial
1283 /// alignment and the element size in bytes.
1284 Align getAlignForIndex(unsigned Idx, Value *Stride, Type *ElementTy,
1285 MaybeAlign A) const {
1286 Align InitialAlign = DL.getValueOrABITypeAlignment(Alignment: A, Ty: ElementTy);
1287 if (Idx == 0)
1288 return InitialAlign;
1289
1290 TypeSize ElementSizeInBits = DL.getTypeSizeInBits(Ty: ElementTy);
1291 if (auto *ConstStride = dyn_cast<ConstantInt>(Val: Stride)) {
1292 uint64_t StrideInBytes =
1293 ConstStride->getZExtValue() * ElementSizeInBits / 8;
1294 return commonAlignment(A: InitialAlign, Offset: Idx * StrideInBytes);
1295 }
1296 return commonAlignment(A: InitialAlign, Offset: ElementSizeInBits / 8);
1297 }
1298
1299 IntegerType *getIndexType(Value *Ptr) const {
1300 return cast<IntegerType>(Val: DL.getIndexType(PtrTy: Ptr->getType()));
1301 }
1302
1303 Value *getIndex(Value *Ptr, uint64_t V) const {
1304 return ConstantInt::get(Ty: getIndexType(Ptr), V);
1305 }
1306
1307 Value *castToIndexType(Value *Ptr, Value *V, IRBuilder<> &Builder) const {
1308 assert(isa<IntegerType>(V->getType()) &&
1309 "Attempted to cast non-integral type to integer index");
1310 // In case the data layout's index type differs in width from the type of
1311 // the value we're given, truncate or zero extend to the appropriate width.
1312 // We zero extend here as indices are unsigned.
1313 return Builder.CreateZExtOrTrunc(V, DestTy: getIndexType(Ptr),
1314 Name: V->getName() + ".cast");
1315 }
1316
1317 /// Load a matrix with \p Shape starting at \p Ptr and using \p Stride between
1318 /// vectors.
1319 MatrixTy loadMatrix(Type *Ty, Value *Ptr, MaybeAlign MAlign, Value *Stride,
1320 bool IsVolatile, ShapeInfo Shape, IRBuilder<> &Builder) {
1321 auto *VType = cast<FixedVectorType>(Val: Ty);
1322 Type *EltTy = VType->getElementType();
1323 Type *VecTy = FixedVectorType::get(ElementType: EltTy, NumElts: Shape.getStride());
1324 Value *EltPtr = Ptr;
1325 MatrixTy Result;
1326 Stride = castToIndexType(Ptr, V: Stride, Builder);
1327 for (unsigned I = 0, E = Shape.getNumVectors(); I < E; ++I) {
1328 Value *GEP = computeVectorAddr(
1329 BasePtr: EltPtr, VecIdx: Builder.getIntN(N: Stride->getType()->getScalarSizeInBits(), C: I),
1330 Stride, NumElements: Shape.getStride(), EltType: EltTy, Builder);
1331 Value *Vector = Builder.CreateAlignedLoad(
1332 Ty: VecTy, Ptr: GEP, Align: getAlignForIndex(Idx: I, Stride, ElementTy: EltTy, A: MAlign),
1333 isVolatile: IsVolatile, Name: "col.load");
1334
1335 Result.addVector(V: Vector);
1336 }
1337 return Result.addNumLoads(N: getNumOps(VT: Result.getVectorTy()) *
1338 Result.getNumVectors());
1339 }
1340
1341 /// Loads a sub-matrix with shape \p ResultShape from a \p R x \p C matrix,
1342 /// starting at \p MatrixPtr[I][J].
1343 MatrixTy loadMatrix(Value *MatrixPtr, MaybeAlign Align, bool IsVolatile,
1344 ShapeInfo MatrixShape, Value *I, Value *J,
1345 ShapeInfo ResultShape, Type *EltTy,
1346 IRBuilder<> &Builder) {
1347 Value *Offset = Builder.CreateAdd(
1348 LHS: Builder.CreateMul(LHS: J, RHS: getIndex(Ptr: MatrixPtr, V: MatrixShape.getStride())), RHS: I);
1349
1350 Value *TileStart = Builder.CreateInBoundsGEP(Ty: EltTy, Ptr: MatrixPtr, IdxList: Offset);
1351 auto *TileTy = FixedVectorType::get(ElementType: EltTy, NumElts: ResultShape.NumRows *
1352 ResultShape.NumColumns);
1353
1354 return loadMatrix(Ty: TileTy, Ptr: TileStart, MAlign: Align,
1355 Stride: getIndex(Ptr: MatrixPtr, V: MatrixShape.getStride()), IsVolatile,
1356 Shape: ResultShape, Builder);
1357 }
1358
1359 /// Lower a load instruction with shape information.
1360 MatrixTy LowerLoad(Instruction *Inst, Value *Ptr, MaybeAlign Align,
1361 Value *Stride, bool IsVolatile, ShapeInfo Shape,
1362 IRBuilder<> &Builder) {
1363 return loadMatrix(Ty: Inst->getType(), Ptr, MAlign: Align, Stride, IsVolatile, Shape,
1364 Builder);
1365 }
1366
1367 /// Lowers llvm.matrix.column.major.load.
1368 ///
1369 /// The intrinsic loads a matrix from memory using a stride between columns.
1370 MatrixTy LowerColumnMajorLoad(CallInst *Inst, IRBuilder<> &Builder) {
1371 assert(Opts.matrix_default_layout == MatrixLayoutTy::ColumnMajor &&
1372 "Intrinsic only supports column-major layout!");
1373 Value *Ptr = Inst->getArgOperand(i: 0);
1374 Value *Stride = Inst->getArgOperand(i: 1);
1375 return LowerLoad(Inst, Ptr, Align: Inst->getParamAlign(ArgNo: 0), Stride,
1376 IsVolatile: cast<ConstantInt>(Val: Inst->getArgOperand(i: 2))->isOne(),
1377 Shape: {Inst->getArgOperand(i: 3), Inst->getArgOperand(i: 4)}, Builder);
1378 }
1379
1380 /// Stores a sub-matrix \p StoreVal into the \p R x \p C matrix starting at \p
1381 /// MatrixPtr[I][J].
1382 void storeMatrix(const MatrixTy &StoreVal, Value *MatrixPtr,
1383 MaybeAlign MAlign, bool IsVolatile, ShapeInfo MatrixShape,
1384 Value *I, Value *J, Type *EltTy, IRBuilder<> &Builder) {
1385 Value *Offset = Builder.CreateAdd(
1386 LHS: Builder.CreateMul(LHS: J, RHS: getIndex(Ptr: MatrixPtr, V: MatrixShape.getStride())), RHS: I);
1387
1388 Value *TileStart = Builder.CreateInBoundsGEP(Ty: EltTy, Ptr: MatrixPtr, IdxList: Offset);
1389 auto *TileTy = FixedVectorType::get(ElementType: EltTy, NumElts: StoreVal.getNumRows() *
1390 StoreVal.getNumColumns());
1391
1392 storeMatrix(Ty: TileTy, StoreVal, Ptr: TileStart, MAlign,
1393 Stride: getIndex(Ptr: MatrixPtr, V: MatrixShape.getStride()), IsVolatile,
1394 Builder);
1395 }
1396
1397 /// Store matrix \p StoreVal starting at \p Ptr and using \p Stride between
1398 /// vectors.
1399 MatrixTy storeMatrix(Type *Ty, MatrixTy StoreVal, Value *Ptr,
1400 MaybeAlign MAlign, Value *Stride, bool IsVolatile,
1401 IRBuilder<> &Builder) {
1402 auto *VType = cast<FixedVectorType>(Val: Ty);
1403 Value *EltPtr = Ptr;
1404 Stride = castToIndexType(Ptr, V: Stride, Builder);
1405 for (auto Vec : enumerate(First: StoreVal.vectors())) {
1406 Value *GEP = computeVectorAddr(
1407 BasePtr: EltPtr,
1408 VecIdx: Builder.getIntN(N: Stride->getType()->getScalarSizeInBits(),
1409 C: Vec.index()),
1410 Stride, NumElements: StoreVal.getStride(), EltType: VType->getElementType(), Builder);
1411 Builder.CreateAlignedStore(Val: Vec.value(), Ptr: GEP,
1412 Align: getAlignForIndex(Idx: Vec.index(), Stride,
1413 ElementTy: VType->getElementType(),
1414 A: MAlign),
1415 isVolatile: IsVolatile);
1416 }
1417 return MatrixTy().addNumStores(N: getNumOps(VT: StoreVal.getVectorTy()) *
1418 StoreVal.getNumVectors());
1419 }
1420
1421 /// Lower a store instruction with shape information.
1422 MatrixTy LowerStore(Instruction *Inst, Value *Matrix, Value *Ptr,
1423 MaybeAlign A, Value *Stride, bool IsVolatile,
1424 ShapeInfo Shape, IRBuilder<> &Builder) {
1425 auto StoreVal = getMatrix(MatrixVal: Matrix, SI: Shape, Builder);
1426 return storeMatrix(Ty: Matrix->getType(), StoreVal, Ptr, MAlign: A, Stride, IsVolatile,
1427 Builder);
1428 }
1429
1430 /// Lowers llvm.matrix.column.major.store.
1431 ///
1432 /// The intrinsic store a matrix back memory using a stride between columns.
1433 MatrixTy LowerColumnMajorStore(CallInst *Inst, IRBuilder<> &Builder) {
1434 assert(Opts.matrix_default_layout == MatrixLayoutTy::ColumnMajor &&
1435 "Intrinsic only supports column-major layout!");
1436 Value *Matrix = Inst->getArgOperand(i: 0);
1437 Value *Ptr = Inst->getArgOperand(i: 1);
1438 Value *Stride = Inst->getArgOperand(i: 2);
1439 return LowerStore(Inst, Matrix, Ptr, A: Inst->getParamAlign(ArgNo: 1), Stride,
1440 IsVolatile: cast<ConstantInt>(Val: Inst->getArgOperand(i: 3))->isOne(),
1441 Shape: {Inst->getArgOperand(i: 4), Inst->getArgOperand(i: 5)},
1442 Builder);
1443 }
1444
1445 // Set elements I..I+NumElts-1 to Block
1446 Value *insertVector(Value *Col, unsigned I, Value *Block,
1447 IRBuilder<> &Builder) {
1448
1449 // First, bring Block to the same size as Col
1450 unsigned BlockNumElts =
1451 cast<FixedVectorType>(Val: Block->getType())->getNumElements();
1452 unsigned NumElts = cast<FixedVectorType>(Val: Col->getType())->getNumElements();
1453 assert(NumElts >= BlockNumElts && "Too few elements for current block");
1454
1455 Block = Builder.CreateShuffleVector(
1456 V: Block, Mask: createSequentialMask(Start: 0, NumInts: BlockNumElts, NumUndefs: NumElts - BlockNumElts));
1457
1458 // If Col is 7 long and I is 2 and BlockNumElts is 2 the mask is: 0, 1, 7,
1459 // 8, 4, 5, 6
1460 SmallVector<int, 16> Mask;
1461 unsigned i;
1462 for (i = 0; i < I; i++)
1463 Mask.push_back(Elt: i);
1464
1465 unsigned VecNumElts =
1466 cast<FixedVectorType>(Val: Col->getType())->getNumElements();
1467 for (; i < I + BlockNumElts; i++)
1468 Mask.push_back(Elt: i - I + VecNumElts);
1469
1470 for (; i < VecNumElts; i++)
1471 Mask.push_back(Elt: i);
1472
1473 return Builder.CreateShuffleVector(V1: Col, V2: Block, Mask);
1474 }
1475
1476 Value *createMulAdd(Value *Sum, Value *A, Value *B, bool UseFPOp,
1477 IRBuilder<> &Builder, bool AllowContraction,
1478 unsigned &NumComputeOps) {
1479 NumComputeOps += getNumOps(VT: A->getType());
1480 if (!Sum)
1481 return UseFPOp ? Builder.CreateFMul(L: A, R: B) : Builder.CreateMul(LHS: A, RHS: B);
1482
1483 if (UseFPOp) {
1484 if (AllowContraction) {
1485 // Use fmuladd for floating point operations and let the backend decide
1486 // if that's profitable.
1487 return Builder.CreateIntrinsic(ID: Intrinsic::fmuladd, OverloadTypes: A->getType(),
1488 Args: {A, B, Sum});
1489 }
1490 NumComputeOps += getNumOps(VT: A->getType());
1491 Value *Mul = Builder.CreateFMul(L: A, R: B);
1492 return Builder.CreateFAdd(L: Sum, R: Mul);
1493 }
1494
1495 NumComputeOps += getNumOps(VT: A->getType());
1496 Value *Mul = Builder.CreateMul(LHS: A, RHS: B);
1497 return Builder.CreateAdd(LHS: Sum, RHS: Mul);
1498 }
1499
1500 /// Cache \p Matrix as result of \p Inst and update the uses of \p Inst. For
1501 /// users with shape information, there's nothing to do: they will use the
1502 /// cached value when they are lowered. For other users, \p Matrix is
1503 /// flattened and the uses are updated to use it. Also marks \p Inst for
1504 /// deletion.
1505 void finalizeLowering(Instruction *Inst, MatrixTy Matrix,
1506 IRBuilder<> &Builder) {
1507 auto inserted = Inst2ColumnMatrix.insert(KV: std::make_pair(x&: Inst, y&: Matrix));
1508 (void)inserted;
1509 assert((inserted.second || isa<PHINode>(Inst)) &&
1510 "multiple matrix lowering mapping");
1511
1512 ToRemove.push_back(Elt: Inst);
1513 Value *Flattened = nullptr;
1514 for (Use &U : llvm::make_early_inc_range(Range: Inst->uses())) {
1515 if (ShapeMap.contains(Val: U.getUser()))
1516 continue;
1517
1518 if (!Flattened) {
1519 Flattened = Matrix.embedInVector(Builder);
1520 LLVM_DEBUG(
1521 if (Instruction *User = dyn_cast<Instruction>(U.getUser())) dbgs()
1522 << "flattening a " << Matrix.shape() << " matrix:\n"
1523 << *Inst
1524 << "\nbecause we do not have a shape-aware lowering for its "
1525 "user:\n"
1526 << *User << '\n';);
1527 FlattenedMatrices++;
1528 }
1529 U.set(Flattened);
1530 }
1531 }
1532
1533 /// Special case for MatMul lowering. Prevents scalar loads of row-major
1534 /// vectors Lowers to vector reduction add instead of sequential add if
1535 /// reassocation is enabled.
1536 void lowerDotProduct(CallInst *MatMul,
1537 SmallPtrSet<Instruction *, 16> &FusedInsts,
1538 FastMathFlags FMF) {
1539 if (FusedInsts.contains(Ptr: MatMul) ||
1540 Opts.matrix_default_layout != MatrixLayoutTy::ColumnMajor)
1541 return;
1542 ShapeInfo LShape(MatMul->getArgOperand(i: 2), MatMul->getArgOperand(i: 3));
1543 ShapeInfo RShape(MatMul->getArgOperand(i: 3), MatMul->getArgOperand(i: 4));
1544
1545 if (LShape.NumRows != 1 || RShape.NumColumns != 1) // not a dot product
1546 return;
1547
1548 Value *LHS = MatMul->getArgOperand(i: 0);
1549 Value *RHS = MatMul->getArgOperand(i: 1);
1550
1551 Type *ElementType = cast<FixedVectorType>(Val: LHS->getType())->getElementType();
1552 bool IsIntVec = ElementType->isIntegerTy();
1553
1554 TTI::TargetCostKind CostKind = TTI::TCK_RecipThroughput;
1555
1556 // Floating point reductions require reassocation.
1557 if (!IsIntVec && !FMF.allowReassoc())
1558 return;
1559
1560 auto CanBeFlattened = [](Value *Op) {
1561 if (match(V: Op, P: m_BinOp()))
1562 return true;
1563 return match(
1564 V: Op, P: m_OneUse(SubPattern: m_CombineOr(
1565 Ps: m_Load(Op: m_Value()),
1566 Ps: m_CombineOr(Ps: m_Intrinsic<Intrinsic::matrix_transpose>(),
1567 Ps: m_Intrinsic<Intrinsic::matrix_column_major_load>(
1568 Ops: m_Value(), Ops: m_One())))));
1569 };
1570 // Returns the cost benefit of using \p Op with the dot product lowering. If
1571 // the returned cost is < 0, the argument is cheaper to use in the
1572 // dot-product lowering.
1573 auto GetCostForArg = [this, &CanBeFlattened, CostKind](Value *Op,
1574 unsigned N) {
1575 if (!ShapeMap.contains(Val: Op))
1576 return InstructionCost::getInvalid();
1577
1578 if (!isa<Instruction>(Val: Op))
1579 return InstructionCost(0);
1580
1581 FixedVectorType *VecTy = cast<FixedVectorType>(Val: Op->getType());
1582 Type *EltTy = VecTy->getElementType();
1583
1584 if (!CanBeFlattened(Op)) {
1585 InstructionCost EmbedCost(0);
1586 // Roughly estimate the cost for embedding the columns into a vector.
1587 for (unsigned I = 1; I < N; ++I)
1588 EmbedCost +=
1589 TTI.getShuffleCost(Kind: TTI::SK_Splice, DstTy: FixedVectorType::get(ElementType: EltTy, NumElts: 1),
1590 SrcTy: FixedVectorType::get(ElementType: EltTy, NumElts: 1), CostKind);
1591 return EmbedCost;
1592 }
1593
1594 if (match(V: Op, P: m_BinOp()) && ShapeMap.contains(Val: Op)) {
1595 InstructionCost OriginalCost =
1596 TTI.getArithmeticInstrCost(Opcode: cast<Instruction>(Val: Op)->getOpcode(),
1597 Ty: EltTy, CostKind) *
1598 N;
1599 InstructionCost NewCost = TTI.getArithmeticInstrCost(
1600 Opcode: cast<Instruction>(Val: Op)->getOpcode(), Ty: VecTy, CostKind);
1601 return NewCost - OriginalCost;
1602 }
1603
1604 if (match(V: Op, P: m_Intrinsic<Intrinsic::matrix_transpose>())) {
1605 // The transpose can be skipped for the dot product lowering, roughly
1606 // estimate the savings as the cost of embedding the columns in a
1607 // vector.
1608 InstructionCost EmbedCost(0);
1609 for (unsigned I = 1; I < N; ++I)
1610 EmbedCost -=
1611 TTI.getShuffleCost(Kind: TTI::SK_Splice, DstTy: FixedVectorType::get(ElementType: EltTy, NumElts: 1),
1612 SrcTy: FixedVectorType::get(ElementType: EltTy, NumElts: 1), CostKind);
1613 return EmbedCost;
1614 }
1615
1616 // Costs for loads.
1617 if (N == 1)
1618 return InstructionCost(0);
1619
1620 return TTI.getMemoryOpCost(Opcode: Instruction::Load, Src: VecTy, Alignment: Align(1), AddressSpace: 0,
1621 CostKind) -
1622 N * TTI.getMemoryOpCost(Opcode: Instruction::Load, Src: EltTy, Alignment: Align(1), AddressSpace: 0,
1623 CostKind);
1624 };
1625
1626 // Iterate over LHS and operations feeding LHS and check if it is profitable
1627 // to flatten the visited ops. For each op, we compute the difference
1628 // between the flattened and matrix versions.
1629 SmallPtrSet<Value *, 4> Seen;
1630 SmallVector<Value *> WorkList;
1631 SmallVector<Value *> ToFlatten;
1632 WorkList.push_back(Elt: LHS);
1633 InstructionCost LHSCost(0);
1634 while (!WorkList.empty()) {
1635 Value *Op = WorkList.pop_back_val();
1636 if (!Seen.insert(Ptr: Op).second)
1637 continue;
1638
1639 InstructionCost OpCost = GetCostForArg(Op, LShape.NumColumns);
1640 if (OpCost + LHSCost >= LHSCost)
1641 continue;
1642
1643 LHSCost += OpCost;
1644 ToFlatten.push_back(Elt: Op);
1645 if (auto *I = dyn_cast<Instruction>(Val: Op))
1646 WorkList.append(in_start: I->op_begin(), in_end: I->op_end());
1647 }
1648
1649 // We compare the costs of a vector.reduce.add to sequential add.
1650 int AddOpCode = IsIntVec ? Instruction::Add : Instruction::FAdd;
1651 int MulOpCode = IsIntVec ? Instruction::Mul : Instruction::FMul;
1652 InstructionCost ReductionCost =
1653 TTI.getArithmeticReductionCost(
1654 Opcode: AddOpCode, Ty: cast<FixedVectorType>(Val: LHS->getType()),
1655 FMF: IsIntVec ? std::nullopt : std::optional(FMF), CostKind) +
1656 TTI.getArithmeticInstrCost(Opcode: MulOpCode, Ty: LHS->getType(), CostKind);
1657 InstructionCost SequentialAddCost =
1658 TTI.getArithmeticInstrCost(Opcode: AddOpCode, Ty: ElementType, CostKind) *
1659 (LShape.NumColumns - 1) +
1660 TTI.getArithmeticInstrCost(Opcode: MulOpCode, Ty: ElementType, CostKind) *
1661 (LShape.NumColumns);
1662 if ((LHSCost + ReductionCost - SequentialAddCost) > InstructionCost(0))
1663 return;
1664
1665 FusedInsts.insert(Ptr: MatMul);
1666 IRBuilder<> Builder(MatMul);
1667 auto FlattenArg = [&Builder, &FusedInsts, &CanBeFlattened,
1668 this](Value *Op) {
1669 // Matmul must be the only user of loads because we don't use LowerLoad
1670 // for row vectors (LowerLoad results in scalar loads and shufflevectors
1671 // instead of single vector load).
1672 if (!CanBeFlattened(Op))
1673 return;
1674
1675 if (match(V: Op, P: m_BinOp())) {
1676 auto It = ShapeMap.find(Val: Op);
1677 if (It != ShapeMap.end()) {
1678 It->second = It->second.t();
1679 return;
1680 }
1681 }
1682
1683 FusedInsts.insert(Ptr: cast<Instruction>(Val: Op));
1684 // If vector uses the builtin load, lower to a LoadInst
1685 Value *Arg;
1686 if (match(V: Op, P: m_Intrinsic<Intrinsic::matrix_column_major_load>(
1687 Ops: m_Value(V&: Arg)))) {
1688 auto *MatLoad = cast<IntrinsicInst>(Val: Op);
1689 bool IsVolatile = cast<ConstantInt>(Val: MatLoad->getArgOperand(i: 2))->isOne();
1690 // Preserve the volatile flag and alignment of the original load.
1691 Align Alignment = getAlignForIndex(
1692 Idx: 0, Stride: MatLoad->getArgOperand(i: 1),
1693 ElementTy: cast<FixedVectorType>(Val: Op->getType())->getElementType(),
1694 A: MatLoad->getParamAlign(ArgNo: 0));
1695 auto *NewLoad = Builder.CreateAlignedLoad(Ty: Op->getType(), Ptr: Arg, Align: Alignment,
1696 isVolatile: IsVolatile);
1697 Op->replaceAllUsesWith(V: NewLoad);
1698 eraseFromParentAndRemoveFromShapeMap(Inst: cast<Instruction>(Val: Op));
1699 return;
1700 } else if (match(V: Op, P: m_Intrinsic<Intrinsic::matrix_transpose>(
1701 Ops: m_Value(V&: Arg)))) {
1702 ToRemove.push_back(Elt: cast<Instruction>(Val: Op));
1703 Op->replaceAllUsesWith(V: Arg);
1704 return;
1705 }
1706 };
1707
1708 for (auto *V : ToFlatten)
1709 FlattenArg(V);
1710
1711 LHS = MatMul->getArgOperand(i: 0);
1712
1713 // Insert mul/fmul and llvm.vector.reduce.fadd
1714 Value *Mul =
1715 IsIntVec ? Builder.CreateMul(LHS, RHS) : Builder.CreateFMul(L: LHS, R: RHS);
1716
1717 Value *Result;
1718 if (IsIntVec)
1719 Result = Builder.CreateAddReduce(Src: Mul);
1720 else {
1721 Result = Builder.CreateFAddReduce(
1722 Acc: ConstantFP::get(
1723 Ty: cast<FixedVectorType>(Val: LHS->getType())->getElementType(), V: 0.0),
1724 Src: Mul);
1725 cast<Instruction>(Val: Result)->setFastMathFlags(FMF);
1726 }
1727
1728 // pack scalar back into a matrix and then replace matmul inst
1729 Result = Builder.CreateInsertElement(Vec: PoisonValue::get(T: MatMul->getType()),
1730 NewElt: Result, Idx: uint64_t(0));
1731 MatMul->replaceAllUsesWith(V: Result);
1732 FusedInsts.insert(Ptr: MatMul);
1733 ToRemove.push_back(Elt: MatMul);
1734 }
1735
1736 /// Given \p Remainder iterations of the the matmul inner loop,
1737 /// potentially lower \p Blocksize that is used for the underlying
1738 /// vector.
1739 unsigned capBlockSize(unsigned BlockSize, unsigned Remainder, Type *EltType) {
1740 if (BlockSize <= Remainder)
1741 return BlockSize;
1742
1743 // If the remainder is also a legal type just use it.
1744 auto *VecTy = FixedVectorType::get(ElementType: EltType, NumElts: Remainder);
1745 if (TTI.isTypeLegal(Ty: VecTy))
1746 return Remainder;
1747
1748 // Similarly, if the vector is small enough that we don't want
1749 // to split further.
1750 if (VecTy->getPrimitiveSizeInBits() <=
1751 Opts.matrix_split_matmul_remainder_over_threshold)
1752 return Remainder;
1753
1754 // Gradually lower the vectorization factor to cover the
1755 // remainder.
1756 do {
1757 BlockSize /= 2;
1758 } while (BlockSize > Remainder);
1759 return BlockSize;
1760 }
1761
1762 /// Compute \p Result += \p A * \p B for input matrices with left-associating
1763 /// addition.
1764 ///
1765 /// We can fold a transpose into the operand that is used to extract scalars.
1766 /// This is the first operands with row-major and the second with
1767 /// column-major. If \p IsScalarMatrixTransposed we assume the appropriate
1768 /// operand is transposed.
1769 void emitMatrixMultiply(MatrixTy &Result, const MatrixTy &A,
1770 const MatrixTy &B, IRBuilder<> &Builder, bool IsTiled,
1771 bool IsScalarMatrixTransposed, FastMathFlags FMF) {
1772 const unsigned VF = std::max<unsigned>(
1773 a: TTI.getRegisterBitWidth(K: TargetTransformInfo::RGK_FixedWidthVector)
1774 .getFixedValue() /
1775 Result.getElementType()->getPrimitiveSizeInBits().getFixedValue(),
1776 b: 1U);
1777 unsigned R = Result.getNumRows();
1778 unsigned C = Result.getNumColumns();
1779 unsigned M = A.getNumColumns();
1780
1781 bool IsFP = Result.getElementType()->isFloatingPointTy();
1782 assert(A.isColumnMajor() == B.isColumnMajor() &&
1783 Result.isColumnMajor() == A.isColumnMajor() &&
1784 "operands must agree on matrix layout");
1785 unsigned NumComputeOps = 0;
1786
1787 Builder.setFastMathFlags(FMF);
1788
1789 if (A.isColumnMajor()) {
1790 // Multiply columns from the first operand with scalars from the second
1791 // operand. Then move along the K axes and accumulate the columns. With
1792 // this the adds can be vectorized without reassociation.
1793 for (unsigned J = 0; J < C; ++J) {
1794 unsigned BlockSize = VF;
1795 // If Result is zero, we don't need to accumulate in the K==0 iteration.
1796 bool isSumZero = isa<ConstantAggregateZero>(Val: Result.getColumn(i: J));
1797
1798 for (unsigned I = 0; I < R; I += BlockSize) {
1799 // Lower block size to make sure we stay within bounds.
1800 BlockSize = capBlockSize(BlockSize, Remainder: R - I, EltType: Result.getElementType());
1801 Value *Sum = IsTiled ? Result.extractVector(I, J, NumElts: BlockSize, Builder)
1802 : nullptr;
1803 for (unsigned K = 0; K < M; ++K) {
1804 Value *L = A.extractVector(I, J: K, NumElts: BlockSize, Builder);
1805 Value *RH = Builder.CreateExtractElement(
1806 Vec: B.getColumn(i: IsScalarMatrixTransposed ? K : J),
1807 Idx: IsScalarMatrixTransposed ? J : K);
1808 Value *Splat = Builder.CreateVectorSplat(NumElts: BlockSize, V: RH, Name: "splat");
1809 Sum =
1810 createMulAdd(Sum: isSumZero && K == 0 ? nullptr : Sum, A: L, B: Splat,
1811 UseFPOp: IsFP, Builder, AllowContraction: FMF.allowContract(), NumComputeOps);
1812 }
1813 Result.setVector(i: J,
1814 V: insertVector(Col: Result.getVector(i: J), I, Block: Sum, Builder));
1815 }
1816 }
1817 } else {
1818 // Multiply rows from the second operand with scalars from the first
1819 // operand. Then move along the K axes and accumulate the rows. With this
1820 // the adds can be vectorized without reassociation.
1821 for (unsigned I = 0; I < R; ++I) {
1822 unsigned BlockSize = VF;
1823 bool isSumZero = isa<ConstantAggregateZero>(Val: Result.getRow(i: I));
1824 for (unsigned J = 0; J < C; J += BlockSize) {
1825 // Lower the vectorization factor to cover the remainder.
1826 BlockSize = capBlockSize(BlockSize, Remainder: C - J, EltType: Result.getElementType());
1827
1828 Value *Sum = nullptr;
1829 for (unsigned K = 0; K < M; ++K) {
1830 Value *R = B.extractVector(I: K, J, NumElts: BlockSize, Builder);
1831 Value *LH = Builder.CreateExtractElement(
1832 Vec: A.getVector(i: IsScalarMatrixTransposed ? K : I),
1833 Idx: IsScalarMatrixTransposed ? I : K);
1834 Value *Splat = Builder.CreateVectorSplat(NumElts: BlockSize, V: LH, Name: "splat");
1835 Sum =
1836 createMulAdd(Sum: isSumZero && K == 0 ? nullptr : Sum, A: Splat, B: R,
1837 UseFPOp: IsFP, Builder, AllowContraction: FMF.allowContract(), NumComputeOps);
1838 }
1839 Result.setVector(i: I,
1840 V: insertVector(Col: Result.getVector(i: I), I: J, Block: Sum, Builder));
1841 }
1842 }
1843 }
1844 Result.addNumComputeOps(N: NumComputeOps);
1845 }
1846
1847 /// Ensure that the memory in \p Load does not alias \p Store by potentially
1848 /// copying it to a new location. This new or otherwise the original location
1849 /// is returned.
1850 std::pair<Value *, AllocaInst *>
1851 getNonAliasingPointer(LoadInst *Load, StoreInst *Store, CallInst *MatMul) {
1852 MemoryLocation StoreLoc = MemoryLocation::get(SI: Store);
1853 MemoryLocation LoadLoc = MemoryLocation::get(LI: Load);
1854
1855 // If we can statically determine noalias we're good.
1856 if (AA->isNoAlias(LocA: LoadLoc, LocB: StoreLoc))
1857 return {Load->getPointerOperand(), nullptr};
1858
1859 // If the pointers are in different address spaces, we cannot compare them
1860 // at runtime. Conservatively copy the load operand to a new buffer.
1861 IRBuilder<> AllocaBuilder(&Func.getEntryBlock().front());
1862 if (Load->getPointerAddressSpace() != Store->getPointerAddressSpace()) {
1863 auto *VT = cast<FixedVectorType>(Val: Load->getType());
1864 auto *ArrayTy =
1865 ArrayType::get(ElementType: VT->getElementType(), NumElements: VT->getNumElements());
1866 AllocaInst *Alloca =
1867 AllocaBuilder.CreateAlloca(Ty: ArrayTy, AddrSpace: Load->getPointerAddressSpace());
1868 IRBuilder<> Builder(MatMul);
1869 Builder.CreateLifetimeStart(Ptr: Alloca);
1870 Builder.CreateMemCpy(Dst: Alloca, DstAlign: Alloca->getAlign(),
1871 Src: Load->getPointerOperand(), SrcAlign: Load->getAlign(),
1872 Size: LoadLoc.Size.getValue());
1873 return {Alloca, Alloca};
1874 }
1875
1876 // Create code to check if the memory locations of the Load and Store
1877 // overlap and if they do, copy Load's operand to a new buffer.
1878
1879 // First, create new blocks for 2n part of the check and the copy.
1880 BasicBlock *Check0 = MatMul->getParent();
1881 // FIXME: Use lazy DTU and update SplitBlock to accept a DTU instead of a
1882 // DT. Manually collect dominator tree updates, to avoid unnecessary work,
1883 // as we adjust Check0 and Check1's branches.
1884 SmallVector<DominatorTree::UpdateType, 4> DTUpdates;
1885 for (BasicBlock *Succ : successors(BB: Check0))
1886 DTUpdates.push_back(Elt: {DT->Delete, Check0, Succ});
1887
1888 BasicBlock *Check1 =
1889 SplitBlock(Old: MatMul->getParent(), SplitPt: MatMul, DTU: (DomTreeUpdater *)nullptr, LI,
1890 MSSAU: nullptr, BBName: "alias_cont");
1891 BasicBlock *Copy =
1892 SplitBlock(Old: MatMul->getParent(), SplitPt: MatMul, DTU: (DomTreeUpdater *)nullptr, LI,
1893 MSSAU: nullptr, BBName: "copy");
1894 BasicBlock *Fusion =
1895 SplitBlock(Old: MatMul->getParent(), SplitPt: MatMul, DTU: (DomTreeUpdater *)nullptr, LI,
1896 MSSAU: nullptr, BBName: "no_alias");
1897
1898 // Check if the loaded memory location begins before the end of the store
1899 // location. If the condition holds, they might overlap, otherwise they are
1900 // guaranteed to not overlap.
1901 IRBuilder<> Builder(MatMul);
1902 Check0->getTerminator()->eraseFromParent();
1903 Builder.SetInsertPoint(Check0);
1904 Type *AddrTy = DL.getAddressType(PtrTy: Store->getPointerOperand()->getType());
1905 Value *StoreBegin = Store->getPointerOperand();
1906 Value *StoreEnd = Builder.CreatePtrAdd(
1907 Ptr: StoreBegin, Offset: ConstantInt::get(Ty: AddrTy, V: StoreLoc.Size.getValue()),
1908 Name: "store.end",
1909 NW: GEPNoWrapFlags::inBounds() | GEPNoWrapFlags::noUnsignedWrap());
1910 Value *LoadBegin = Load->getPointerOperand();
1911 CondBrInst *BR1 = Builder.CreateCondBr(
1912 Cond: Builder.CreateICmpULT(LHS: LoadBegin, RHS: StoreEnd), True: Check1, False: Fusion);
1913 setExplicitlyUnknownBranchWeightsIfProfiled(I&: *BR1, DEBUG_TYPE);
1914
1915 // Check if the store begins before the end of the load location. If the
1916 // condition holds, they alias, otherwise they are guaranteed to not
1917 // overlap.
1918 Check1->getTerminator()->eraseFromParent();
1919 Builder.SetInsertPoint(Check1->begin());
1920
1921 auto *VT = cast<FixedVectorType>(Val: Load->getType());
1922 // Use an array type for the alloca, to avoid potentially huge alignment
1923 // requirements for large vector types.
1924 auto *ArrayTy = ArrayType::get(ElementType: VT->getElementType(), NumElements: VT->getNumElements());
1925 AllocaInst *Alloca =
1926 AllocaBuilder.CreateAlloca(Ty: ArrayTy, AddrSpace: Load->getPointerAddressSpace());
1927 Builder.CreateLifetimeStart(Ptr: Alloca);
1928
1929 Value *LoadEnd = Builder.CreatePtrAdd(
1930 Ptr: LoadBegin, Offset: ConstantInt::get(Ty: AddrTy, V: LoadLoc.Size.getValue()),
1931 Name: "load.end",
1932 NW: GEPNoWrapFlags::inBounds() | GEPNoWrapFlags::noUnsignedWrap());
1933 CondBrInst *BR2 = Builder.CreateCondBr(
1934 Cond: Builder.CreateICmpULT(LHS: StoreBegin, RHS: LoadEnd), True: Copy, False: Fusion);
1935 setExplicitlyUnknownBranchWeightsIfProfiled(I&: *BR2, DEBUG_TYPE);
1936
1937 // Copy load operand to new alloca.
1938 Builder.SetInsertPoint(Copy->begin());
1939 Builder.CreateMemCpy(Dst: Alloca, DstAlign: Alloca->getAlign(), Src: Load->getPointerOperand(),
1940 SrcAlign: Load->getAlign(), Size: LoadLoc.Size.getValue());
1941 Builder.SetInsertPoint(Fusion->begin());
1942 PHINode *PHI = Builder.CreatePHI(Ty: Load->getPointerOperandType(), NumReservedValues: 3);
1943 PHI->addIncoming(V: Load->getPointerOperand(), BB: Check0);
1944 PHI->addIncoming(V: Load->getPointerOperand(), BB: Check1);
1945 PHI->addIncoming(V: Alloca, BB: Copy);
1946
1947 // Adjust DT.
1948 DTUpdates.push_back(Elt: {DT->Insert, Check0, Check1});
1949 DTUpdates.push_back(Elt: {DT->Insert, Check0, Fusion});
1950 DTUpdates.push_back(Elt: {DT->Insert, Check1, Copy});
1951 DTUpdates.push_back(Elt: {DT->Insert, Check1, Fusion});
1952 DT->applyUpdates(Updates: DTUpdates);
1953 return {PHI, Alloca};
1954 }
1955
1956 bool isFusionProfitable(CallInst *MatMul) {
1957 if (Opts.force_fuse_matrix)
1958 return true;
1959
1960 ShapeInfo LShape(MatMul->getArgOperand(i: 2), MatMul->getArgOperand(i: 3));
1961 ShapeInfo RShape(MatMul->getArgOperand(i: 3), MatMul->getArgOperand(i: 4));
1962
1963 const unsigned R = LShape.NumRows;
1964 const unsigned C = RShape.NumColumns;
1965 const unsigned M = LShape.NumColumns;
1966 auto *EltType = cast<FixedVectorType>(Val: MatMul->getType())->getElementType();
1967
1968 const unsigned VF = std::max<unsigned>(
1969 a: TTI.getRegisterBitWidth(K: TargetTransformInfo::RGK_FixedWidthVector)
1970 .getFixedValue() /
1971 EltType->getPrimitiveSizeInBits().getFixedValue(),
1972 b: 1U);
1973
1974 // Cost model for tiling
1975 //
1976 // For tiling to be beneficial, we need reuse either along the R or
1977 // the C axis. We vectorize along the R axis so that means at least
1978 // 3 elements.
1979 // TODO: Also consider cost of copying if operands alias.
1980 if (R <= VF && C == 1)
1981 return false;
1982 // Then we need enough elements to exceed the number of vector
1983 // registers we have. Note that this is an oversimplification since
1984 // fusing also takes some extra loads which may exceed the number of
1985 // reloads necessary.
1986 unsigned Op0Regs = (R + VF - 1) / VF * M;
1987 unsigned Op1Regs = (M + VF - 1) / VF * C;
1988 return Op0Regs + Op1Regs >
1989 TTI.getNumberOfRegisters(ClassID: TTI.getRegisterClassForType(Vector: true));
1990 }
1991
1992 MatrixTy getZeroMatrix(Type *EltType, unsigned R, unsigned C) {
1993 MatrixTy Res;
1994 auto *ColumType = FixedVectorType::get(ElementType: EltType, NumElts: R);
1995 for (unsigned I = 0; I < C; ++I)
1996 Res.addVector(V: ConstantAggregateZero::get(Ty: ColumType));
1997 return Res;
1998 }
1999
2000 void createTiledLoops(CallInst *MatMul, Value *LPtr, ShapeInfo LShape,
2001 Value *RPtr, ShapeInfo RShape, StoreInst *Store) {
2002 auto *EltType = cast<FixedVectorType>(Val: MatMul->getType())->getElementType();
2003
2004 // Create the main tiling loop nest.
2005 TileInfo TI(LShape.NumRows, RShape.NumColumns, LShape.NumColumns,
2006 Opts.fuse_matrix_tile_size);
2007 DomTreeUpdater DTU(DT, DomTreeUpdater::UpdateStrategy::Lazy);
2008 Instruction *InsertI = cast<Instruction>(Val: MatMul);
2009 BasicBlock *Start = InsertI->getParent();
2010 BasicBlock *End =
2011 SplitBlock(Old: InsertI->getParent(), SplitPt: InsertI, DT, LI, MSSAU: nullptr, BBName: "continue");
2012 IRBuilder<> Builder(MatMul);
2013 BasicBlock *InnerBody = TI.CreateTiledLoops(Start, End, B&: Builder, DTU, LI&: *LI);
2014
2015 Type *TileVecTy = FixedVectorType::get(ElementType: MatMul->getType()->getScalarType(),
2016 NumElts: Opts.fuse_matrix_tile_size);
2017 MatrixTy TileResult;
2018 // Insert in the inner loop header.
2019 Builder.SetInsertPoint(TI.KLoop.Header->getTerminator());
2020 // Create PHI nodes for the result columns to accumulate across iterations.
2021 SmallVector<PHINode *, 4> ColumnPhis;
2022 for (unsigned I = 0; I < Opts.fuse_matrix_tile_size; I++) {
2023 auto *Phi = Builder.CreatePHI(Ty: TileVecTy, NumReservedValues: 2, Name: "result.vec." + Twine(I));
2024 Phi->addIncoming(V: ConstantAggregateZero::get(Ty: TileVecTy),
2025 BB: TI.RowLoop.Header->getSingleSuccessor());
2026 TileResult.addVector(V: Phi);
2027 ColumnPhis.push_back(Elt: Phi);
2028 }
2029
2030 // Insert in the inner loop body, which computes
2031 // Res += Load(CurrentRow, K) * Load(K, CurrentColumn)
2032 Builder.SetInsertPoint(InnerBody->getTerminator());
2033 // Load tiles of the operands.
2034 MatrixTy A =
2035 loadMatrix(MatrixPtr: LPtr, Align: {}, IsVolatile: false, MatrixShape: LShape, I: TI.RowLoop.Index, J: TI.KLoop.Index,
2036 ResultShape: {Opts.fuse_matrix_tile_size, Opts.fuse_matrix_tile_size},
2037 EltTy: EltType, Builder);
2038 MatrixTy B =
2039 loadMatrix(MatrixPtr: RPtr, Align: {}, IsVolatile: false, MatrixShape: RShape, I: TI.KLoop.Index, J: TI.ColumnLoop.Index,
2040 ResultShape: {Opts.fuse_matrix_tile_size, Opts.fuse_matrix_tile_size},
2041 EltTy: EltType, Builder);
2042 emitMatrixMultiply(Result&: TileResult, A, B, Builder, IsTiled: true, IsScalarMatrixTransposed: false,
2043 FMF: getFastMathFlags(Inst: MatMul));
2044 // Store result after the inner loop is done.
2045 Builder.SetInsertPoint(TI.RowLoop.Latch->getTerminator());
2046 storeMatrix(StoreVal: TileResult, MatrixPtr: Store->getPointerOperand(), MAlign: Store->getAlign(),
2047 IsVolatile: Store->isVolatile(), MatrixShape: {LShape.NumRows, RShape.NumColumns},
2048 I: TI.RowLoop.Index, J: TI.ColumnLoop.Index, EltTy: EltType, Builder);
2049
2050 for (unsigned I = 0; I < TileResult.getNumVectors(); I++)
2051 ColumnPhis[I]->addIncoming(V: TileResult.getVector(i: I), BB: TI.KLoop.Latch);
2052
2053 // Force unrolling of a few iterations of the inner loop, to make sure there
2054 // is enough work per iteration.
2055 // FIXME: The unroller should make this decision directly instead, but
2056 // currently the cost-model is not up to the task.
2057 unsigned InnerLoopUnrollCount =
2058 std::min(a: 10u, b: LShape.NumColumns / Opts.fuse_matrix_tile_size);
2059 addStringMetadataToLoop(TheLoop: LI->getLoopFor(BB: TI.KLoop.Header),
2060 MDString: "llvm.loop.unroll.count", V: InnerLoopUnrollCount);
2061 }
2062
2063 void emitSIMDTiling(CallInst *MatMul, LoadInst *LoadOp0, LoadInst *LoadOp1,
2064 StoreInst *Store,
2065 SmallPtrSetImpl<Instruction *> &FusedInsts) {
2066 assert(Opts.matrix_default_layout == MatrixLayoutTy::ColumnMajor &&
2067 "Tiling only supported for column-major matrixes at the moment!");
2068 if (!isFusionProfitable(MatMul))
2069 return;
2070
2071 ShapeInfo LShape(MatMul->getArgOperand(i: 2), MatMul->getArgOperand(i: 3));
2072 ShapeInfo RShape(MatMul->getArgOperand(i: 3), MatMul->getArgOperand(i: 4));
2073
2074 const unsigned R = LShape.NumRows;
2075 const unsigned C = RShape.NumColumns;
2076 const unsigned M = LShape.NumColumns;
2077 auto *EltType = cast<FixedVectorType>(Val: MatMul->getType())->getElementType();
2078
2079 auto [APtr, AAlloca] = getNonAliasingPointer(Load: LoadOp0, Store, MatMul);
2080 auto [BPtr, BAlloca] = getNonAliasingPointer(Load: LoadOp1, Store, MatMul);
2081 Value *CPtr = Store->getPointerOperand();
2082
2083 // Use loop-based tiling when the number of expected operations exceeds
2084 // threshold.
2085 unsigned NumOps = getNumNativeVectorOps(EltType, R, M, C);
2086 bool UseLoops = (NumOps > Opts.fuse_matrix_loops_threshold) &&
2087 R % Opts.fuse_matrix_tile_size == 0 &&
2088 C % Opts.fuse_matrix_tile_size == 0;
2089 if (UseLoops)
2090 createTiledLoops(MatMul, LPtr: APtr, LShape, RPtr: BPtr, RShape, Store);
2091 else {
2092 IRBuilder<> Builder(Store);
2093 for (unsigned J = 0; J < C; J += Opts.fuse_matrix_tile_size)
2094 for (unsigned I = 0; I < R; I += Opts.fuse_matrix_tile_size) {
2095 const unsigned TileR = std::min(a: R - I, b: Opts.fuse_matrix_tile_size);
2096 const unsigned TileC = std::min(a: C - J, b: Opts.fuse_matrix_tile_size);
2097 MatrixTy Res = getZeroMatrix(EltType, R: TileR, C: TileC);
2098
2099 for (unsigned K = 0; K < M; K += Opts.fuse_matrix_tile_size) {
2100 const unsigned TileM = std::min(a: M - K, b: Opts.fuse_matrix_tile_size);
2101 MatrixTy A =
2102 loadMatrix(MatrixPtr: APtr, Align: LoadOp0->getAlign(), IsVolatile: LoadOp0->isVolatile(),
2103 MatrixShape: LShape, I: getIndex(Ptr: APtr, V: I), J: getIndex(Ptr: APtr, V: K),
2104 ResultShape: {TileR, TileM}, EltTy: EltType, Builder);
2105 MatrixTy B =
2106 loadMatrix(MatrixPtr: BPtr, Align: LoadOp1->getAlign(), IsVolatile: LoadOp1->isVolatile(),
2107 MatrixShape: RShape, I: getIndex(Ptr: BPtr, V: K), J: getIndex(Ptr: BPtr, V: J),
2108 ResultShape: {TileM, TileC}, EltTy: EltType, Builder);
2109 emitMatrixMultiply(Result&: Res, A, B, Builder, IsTiled: true, IsScalarMatrixTransposed: false,
2110 FMF: getFastMathFlags(Inst: MatMul));
2111 }
2112 storeMatrix(StoreVal: Res, MatrixPtr: CPtr, MAlign: Store->getAlign(), IsVolatile: Store->isVolatile(), MatrixShape: {R, M},
2113 I: getIndex(Ptr: CPtr, V: I), J: getIndex(Ptr: CPtr, V: J), EltTy: EltType, Builder);
2114 }
2115 }
2116
2117 // End the lifetime of the allocas used for alias-safe copies.
2118 {
2119 IRBuilder<> Builder(Store);
2120 if (AAlloca)
2121 Builder.CreateLifetimeEnd(Ptr: AAlloca);
2122 if (BAlloca)
2123 Builder.CreateLifetimeEnd(Ptr: BAlloca);
2124 }
2125
2126 // Mark eliminated instructions as fused and remove them.
2127 FusedInsts.insert(Ptr: Store);
2128 FusedInsts.insert(Ptr: MatMul);
2129 eraseFromParentAndRemoveFromShapeMap(Inst: Store);
2130 eraseFromParentAndRemoveFromShapeMap(Inst: MatMul);
2131 if (LoadOp0->use_empty()) {
2132 FusedInsts.insert(Ptr: LoadOp0);
2133 eraseFromParentAndRemoveFromShapeMap(Inst: LoadOp0);
2134 }
2135 if (LoadOp1 != LoadOp0 && LoadOp1->use_empty()) {
2136 FusedInsts.insert(Ptr: LoadOp1);
2137 eraseFromParentAndRemoveFromShapeMap(Inst: LoadOp1);
2138 }
2139 }
2140
2141 /// Try to lower matrix multiply chains by fusing operations.
2142 ///
2143 /// Call finalizeLowering on lowered instructions. Instructions that are
2144 /// completely eliminated by fusion are added to \p FusedInsts.
2145 void
2146 LowerMatrixMultiplyFused(CallInst *MatMul,
2147 SmallPtrSetImpl<Instruction *> &FusedInsts,
2148 SmallVector<IntrinsicInst *, 16> &LifetimeEnds) {
2149 if (!Opts.fuse_matrix || !DT || Opts.fuse_matrix_tile_size == 0)
2150 return;
2151
2152 assert(AA && LI && "Analyses should be available");
2153
2154 Value *A = MatMul->getArgOperand(i: 0);
2155 Value *B = MatMul->getArgOperand(i: 1);
2156
2157 // We can fold the transpose into the operand that is used to fetch scalars.
2158 Value *T;
2159 if (Opts.matrix_default_layout == MatrixLayoutTy::ColumnMajor
2160 ? match(V: B, P: m_Intrinsic<Intrinsic::matrix_transpose>(Ops: m_Value(V&: T)))
2161 : match(V: A, P: m_Intrinsic<Intrinsic::matrix_transpose>(Ops: m_Value(V&: T)))) {
2162 IRBuilder<> Builder(MatMul);
2163 auto *EltType =
2164 cast<FixedVectorType>(Val: MatMul->getType())->getElementType();
2165 ShapeInfo LShape(MatMul->getArgOperand(i: 2), MatMul->getArgOperand(i: 3));
2166 ShapeInfo RShape(MatMul->getArgOperand(i: 3), MatMul->getArgOperand(i: 4));
2167 const unsigned R = LShape.NumRows;
2168 const unsigned M = LShape.NumColumns;
2169 const unsigned C = RShape.NumColumns;
2170
2171 MatrixTy MA;
2172 MatrixTy MB;
2173
2174 Value *Transpose;
2175 if (Opts.matrix_default_layout == MatrixLayoutTy::ColumnMajor) {
2176 MA = getMatrix(MatrixVal: A, SI: ShapeInfo(R, M), Builder);
2177 MB = getMatrix(MatrixVal: T, SI: ShapeInfo(C, M), Builder);
2178 Transpose = B;
2179 } else {
2180 MA = getMatrix(MatrixVal: T, SI: ShapeInfo(R, M), Builder);
2181 MB = getMatrix(MatrixVal: B, SI: ShapeInfo(C, M), Builder);
2182 Transpose = A;
2183 }
2184
2185 // Initialize the output
2186 MatrixTy Result(R, C, EltType);
2187
2188 emitMatrixMultiply(Result, A: MA, B: MB, Builder, IsTiled: false, IsScalarMatrixTransposed: true,
2189 FMF: getFastMathFlags(Inst: MatMul));
2190
2191 FusedInsts.insert(Ptr: MatMul);
2192 if (Transpose->hasOneUse()) {
2193 FusedInsts.insert(Ptr: cast<Instruction>(Val: Transpose));
2194 ToRemove.push_back(Elt: cast<Instruction>(Val: Transpose));
2195 // TODO: add a fake entry for the folded instruction so that this is
2196 // included in the expression in the remark.
2197 Inst2ColumnMatrix[Transpose] = MatrixTy(M, C, EltType);
2198 }
2199 finalizeLowering(Inst: MatMul, Matrix: Result, Builder);
2200 return;
2201 }
2202
2203 if (!MatMul->hasOneUse() ||
2204 Opts.matrix_default_layout != MatrixLayoutTy::ColumnMajor)
2205 return;
2206
2207 // Lower {ld, ld} -> matmul -> st chains. No need to call finalizeLowering
2208 // since the single store user will be lowered as part of this.
2209 auto *LoadOp0 = dyn_cast<LoadInst>(Val: A);
2210 auto *LoadOp1 = dyn_cast<LoadInst>(Val: B);
2211 auto *Store = dyn_cast<StoreInst>(Val: *MatMul->user_begin());
2212 if (LoadOp0 && LoadOp1 && Store) {
2213 // The store address must dominate the MatMul instruction, otherwise
2214 // we create invalid IR.
2215 SetVector<Value *> WorkList;
2216 WorkList.insert(X: Store->getOperand(i_nocapture: 1));
2217 SmallVector<Instruction *> ToHoist;
2218 for (unsigned I = 0; I != WorkList.size(); ++I) {
2219 Value *Current = WorkList[I];
2220 auto *CurrI = dyn_cast<Instruction>(Val: Current);
2221 if (!CurrI)
2222 continue;
2223 if (isa<PHINode>(Val: CurrI))
2224 return;
2225 if (DT->dominates(Def: CurrI, User: MatMul))
2226 continue;
2227 if (CurrI->mayHaveSideEffects() || CurrI->mayReadFromMemory())
2228 return;
2229 ToHoist.push_back(Elt: CurrI);
2230 WorkList.insert_range(R: CurrI->operands());
2231 }
2232
2233 sort(C&: ToHoist, Comp: [this](Instruction *A, Instruction *B) {
2234 return DT->dominates(Def: A, User: B);
2235 });
2236 for (Instruction *I : ToHoist)
2237 I->moveBefore(InsertPos: MatMul->getIterator());
2238
2239 // Deal with lifetime.end calls that might be between Load0/Load1 and the
2240 // store. To avoid introducing loads to dead objects (i.e. after the
2241 // lifetime has been termined by @llvm.lifetime.end), either sink them
2242 // after the store if in the same block, or remove the lifetime.end marker
2243 // otherwise. This might pessimize further optimizations, by extending the
2244 // lifetime of the object until the function returns, but should be
2245 // conservatively correct.
2246 MemoryLocation Load0Loc = MemoryLocation::get(LI: LoadOp0);
2247 MemoryLocation Load1Loc = MemoryLocation::get(LI: LoadOp1);
2248 BasicBlock *StoreParent = Store->getParent();
2249 bool FusableOpsInSameBlock = LoadOp0->getParent() == StoreParent &&
2250 LoadOp1->getParent() == StoreParent;
2251 for (unsigned Idx = 0; Idx != LifetimeEnds.size();) {
2252 IntrinsicInst *End = LifetimeEnds[Idx];
2253 llvm::scope_exit Inc([&Idx]() { Idx++; });
2254 // If the lifetime.end is guaranteed to be before the loads or after the
2255 // store, it won't interfere with fusion.
2256 if (DT->dominates(Def: End, User: LoadOp0) && DT->dominates(Def: End, User: LoadOp1))
2257 continue;
2258 if (DT->dominates(Def: Store, User: End))
2259 continue;
2260 // If all fusable ops are in the same block and the lifetime.end is in a
2261 // different block, it won't interfere with fusion.
2262 if (FusableOpsInSameBlock && End->getParent() != StoreParent)
2263 continue;
2264
2265 // If the loads don't alias the lifetime.end, it won't interfere with
2266 // fusion.
2267 MemoryLocation EndLoc = MemoryLocation::getForArgument(Call: End, ArgIdx: 0, TLI: nullptr);
2268 if (!EndLoc.Ptr)
2269 continue;
2270 if (AA->isNoAlias(LocA: Load0Loc, LocB: EndLoc) && AA->isNoAlias(LocA: Load1Loc, LocB: EndLoc))
2271 continue;
2272
2273 // If both lifetime.end and the store are in the same block, extend the
2274 // lifetime until after the store, so the new lifetime covers the loads
2275 // we introduce later.
2276 if (End->getParent() == StoreParent) {
2277 End->moveAfter(MovePos: Store);
2278 continue;
2279 }
2280
2281 // Otherwise remove the conflicting lifetime.end marker.
2282 ToRemove.push_back(Elt: End);
2283 std::swap(a&: LifetimeEnds[Idx], b&: LifetimeEnds.back());
2284 LifetimeEnds.pop_back();
2285 Inc.release();
2286 }
2287
2288 emitSIMDTiling(MatMul, LoadOp0, LoadOp1, Store, FusedInsts);
2289 return;
2290 }
2291 }
2292
2293 /// Lowers llvm.matrix.multiply.
2294 MatrixTy LowerMultiply(CallInst *MatMul, IRBuilder<> &Builder) {
2295 auto *EltType = cast<FixedVectorType>(Val: MatMul->getType())->getElementType();
2296 ShapeInfo LShape(MatMul->getArgOperand(i: 2), MatMul->getArgOperand(i: 3));
2297 ShapeInfo RShape(MatMul->getArgOperand(i: 3), MatMul->getArgOperand(i: 4));
2298
2299 const MatrixTy &Lhs = getMatrix(MatrixVal: MatMul->getArgOperand(i: 0), SI: LShape, Builder);
2300 const MatrixTy &Rhs = getMatrix(MatrixVal: MatMul->getArgOperand(i: 1), SI: RShape, Builder);
2301 assert(Lhs.getElementType() == Rhs.getElementType() &&
2302 "Matrix multiply argument element types do not match.");
2303
2304 const unsigned R = LShape.NumRows;
2305 const unsigned C = RShape.NumColumns;
2306 assert(LShape.NumColumns == RShape.NumRows);
2307
2308 // Initialize the output
2309 MatrixTy Result(R, C, EltType);
2310 assert(Lhs.getElementType() == Result.getElementType() &&
2311 "Matrix multiply result element type does not match arguments.");
2312
2313 emitMatrixMultiply(Result, A: Lhs, B: Rhs, Builder, IsTiled: false, IsScalarMatrixTransposed: false,
2314 FMF: getFastMathFlags(Inst: MatMul));
2315 return Result;
2316 }
2317
2318 /// Lowers llvm.matrix.transpose.
2319 MatrixTy LowerTranspose(CallInst *Inst, IRBuilder<> &Builder) {
2320 MatrixTy Result;
2321 Value *InputVal = Inst->getArgOperand(i: 0);
2322 FixedVectorType *VectorTy = cast<FixedVectorType>(Val: InputVal->getType());
2323 ShapeInfo ArgShape(Inst->getArgOperand(i: 1), Inst->getArgOperand(i: 2));
2324 MatrixTy InputMatrix = getMatrix(MatrixVal: InputVal, SI: ArgShape, Builder);
2325
2326 const unsigned NewNumVecs =
2327 InputMatrix.isColumnMajor() ? ArgShape.NumRows : ArgShape.NumColumns;
2328 const unsigned NewNumElts =
2329 InputMatrix.isColumnMajor() ? ArgShape.NumColumns : ArgShape.NumRows;
2330
2331 for (unsigned I = 0; I < NewNumVecs; ++I) {
2332 // Build a single result vector. First initialize it.
2333 Value *ResultVector = PoisonValue::get(
2334 T: FixedVectorType::get(ElementType: VectorTy->getElementType(), NumElts: NewNumElts));
2335 // Go through the old elements and insert it into the resulting vector.
2336 for (auto J : enumerate(First: InputMatrix.vectors())) {
2337 Value *Elt = Builder.CreateExtractElement(Vec: J.value(), Idx: I);
2338 // Row and column indices are transposed.
2339 ResultVector =
2340 Builder.CreateInsertElement(Vec: ResultVector, NewElt: Elt, Idx: J.index());
2341 }
2342 Result.addVector(V: ResultVector);
2343 }
2344
2345 // TODO: Improve estimate of operations needed for transposes. Currently we
2346 // just count the insertelement/extractelement instructions, but do not
2347 // account for later simplifications/combines.
2348 return Result.addNumComputeOps(N: 2 * ArgShape.NumRows * ArgShape.NumColumns)
2349 .addNumExposedTransposes(N: 1);
2350 }
2351
2352 /// Lower load instructions.
2353 MatrixTy VisitLoad(LoadInst *Inst, const ShapeInfo &SI, Value *Ptr,
2354 IRBuilder<> &Builder) {
2355 return LowerLoad(Inst, Ptr, Align: Inst->getAlign(), Stride: getIndex(Ptr, V: SI.getStride()),
2356 IsVolatile: Inst->isVolatile(), Shape: SI, Builder);
2357 }
2358
2359 MatrixTy VisitStore(StoreInst *Inst, const ShapeInfo &SI, Value *StoredVal,
2360 Value *Ptr, IRBuilder<> &Builder) {
2361 return LowerStore(Inst, Matrix: StoredVal, Ptr, A: Inst->getAlign(),
2362 Stride: getIndex(Ptr, V: SI.getStride()), IsVolatile: Inst->isVolatile(), Shape: SI,
2363 Builder);
2364 }
2365
2366 MatrixTy VisitPHI(PHINode *Inst, const ShapeInfo &SI, IRBuilder<> &Builder) {
2367 auto BlockIP = Inst->getParent()->getFirstInsertionPt();
2368 Builder.SetInsertPoint(BlockIP);
2369 MatrixTy PhiM = getMatrix(MatrixVal: Inst, SI, Builder);
2370
2371 // Cache the reshaped columns per incoming block, so that a block listed
2372 // more than once contributes identical incoming values to the new PHIs.
2373 SmallDenseMap<BasicBlock *, MatrixTy> ReshapedIncoming;
2374 for (auto [IncomingV, IncomingB] :
2375 llvm::zip_equal(t: Inst->incoming_values(), u: Inst->blocks())) {
2376 // getMatrix() may insert some instructions to help with reshaping. The
2377 // safest place for those is just before the terminator of the incoming
2378 // block. If there's a valid insert point before the def, even better.
2379 Builder.SetInsertPoint(IncomingB->getTerminator());
2380 if (auto *IncomingInst = dyn_cast<Instruction>(Val&: IncomingV))
2381 if (auto MaybeIP = IncomingInst->getInsertionPointAfterDef())
2382 Builder.SetInsertPoint(*MaybeIP);
2383
2384 auto [It, Inserted] = ReshapedIncoming.try_emplace(Key: IncomingB);
2385 if (Inserted)
2386 It->second = getMatrix(MatrixVal: IncomingV, SI, Builder);
2387 const MatrixTy &OpM = It->second;
2388
2389 for (unsigned VI = 0, VE = PhiM.getNumVectors(); VI != VE; ++VI) {
2390 PHINode *NewPHI = cast<PHINode>(Val: PhiM.getVector(i: VI));
2391 NewPHI->addIncoming(V: OpM.getVector(i: VI), BB: IncomingB);
2392 }
2393 }
2394
2395 // finalizeLowering() may also insert instructions in some cases. The safe
2396 // place for those is at the end of the initial block of PHIs.
2397 Builder.SetInsertPoint(BlockIP);
2398 return PhiM;
2399 }
2400
2401 /// Lower binary operators.
2402 MatrixTy VisitBinaryOperator(BinaryOperator *Inst, const ShapeInfo &SI,
2403 IRBuilder<> &Builder) {
2404 Value *Lhs = Inst->getOperand(i_nocapture: 0);
2405 Value *Rhs = Inst->getOperand(i_nocapture: 1);
2406
2407 MatrixTy Result;
2408 MatrixTy A = getMatrix(MatrixVal: Lhs, SI, Builder);
2409 MatrixTy B = getMatrix(MatrixVal: Rhs, SI, Builder);
2410 assert(A.isColumnMajor() == B.isColumnMajor() &&
2411 Result.isColumnMajor() == A.isColumnMajor() &&
2412 "operands must agree on matrix layout");
2413
2414 Builder.setFastMathFlags(getFastMathFlags(Inst));
2415
2416 for (auto [AV, BV] : llvm::zip_equal(t: A.vectors(), u: B.vectors()))
2417 Result.addVector(V: Builder.CreateBinOp(Opc: Inst->getOpcode(), LHS: AV, RHS: BV));
2418
2419 return Result.addNumComputeOps(N: getNumOps(VT: Result.getVectorTy()) *
2420 Result.getNumVectors());
2421 }
2422
2423 /// Lower unary operators.
2424 MatrixTy VisitUnaryOperator(UnaryOperator *Inst, const ShapeInfo &SI,
2425 IRBuilder<> &Builder) {
2426 Value *Op = Inst->getOperand(i_nocapture: 0);
2427
2428 MatrixTy Result;
2429 MatrixTy M = getMatrix(MatrixVal: Op, SI, Builder);
2430
2431 Builder.setFastMathFlags(getFastMathFlags(Inst));
2432
2433 // Helper to perform unary op on vectors.
2434 auto BuildVectorOp = [&Builder, Inst](Value *Op) {
2435 switch (Inst->getOpcode()) {
2436 case Instruction::FNeg:
2437 return Builder.CreateFNeg(V: Op);
2438 default:
2439 llvm_unreachable("Unsupported unary operator for matrix");
2440 }
2441 };
2442
2443 for (auto *Vector : M.vectors())
2444 Result.addVector(V: BuildVectorOp(Vector));
2445
2446 return Result.addNumComputeOps(N: getNumOps(VT: Result.getVectorTy()) *
2447 Result.getNumVectors());
2448 }
2449
2450 /// Lower cast instructions.
2451 MatrixTy VisitCastInstruction(CastInst *Inst, const ShapeInfo &Shape,
2452 IRBuilder<> &Builder) {
2453 Value *Op = Inst->getOperand(i_nocapture: 0);
2454
2455 MatrixTy Result;
2456 MatrixTy M = getMatrix(MatrixVal: Op, SI: Shape, Builder);
2457
2458 Builder.setFastMathFlags(getFastMathFlags(Inst));
2459
2460 auto *OrigVTy = cast<VectorType>(Val: Inst->getType());
2461 auto *NewVTy = VectorType::get(ElementType: OrigVTy->getElementType(),
2462 EC: ElementCount::getFixed(MinVal: M.getStride()));
2463
2464 for (auto *Vector : M.vectors())
2465 Result.addVector(V: Builder.CreateCast(Op: Inst->getOpcode(), V: Vector, DestTy: NewVTy));
2466
2467 return Result.addNumComputeOps(N: getNumOps(VT: Result.getVectorTy()) *
2468 Result.getNumVectors());
2469 }
2470
2471 /// Lower selects.
2472 MatrixTy VisitSelectInst(SelectInst *Inst, const ShapeInfo &Shape,
2473 IRBuilder<> &Builder) {
2474 Value *Cond = Inst->getOperand(i_nocapture: 0);
2475 Value *OpA = Inst->getOperand(i_nocapture: 1);
2476 Value *OpB = Inst->getOperand(i_nocapture: 2);
2477
2478 MatrixTy Result;
2479 MatrixTy A = getMatrix(MatrixVal: OpA, SI: Shape, Builder);
2480 MatrixTy B = getMatrix(MatrixVal: OpB, SI: Shape, Builder);
2481
2482 SmallVector<Value*> CondV;
2483 Instruction *MDFrom = nullptr;
2484 if (isa<FixedVectorType>(Val: Cond->getType())) {
2485 MatrixTy C = getMatrix(MatrixVal: Cond, SI: Shape, Builder);
2486 llvm::copy(Range: C.vectors(), Out: std::back_inserter(x&: CondV));
2487 } else {
2488 CondV.resize(N: A.getNumVectors());
2489 llvm::fill(Range&: CondV, Value&: Cond);
2490 if (!ProfcheckDisableMetadataFixes)
2491 MDFrom = Inst;
2492 }
2493
2494 for (auto [CV, AV, BV] : llvm::zip_equal(t&: CondV, u: A.vectors(), args: B.vectors())) {
2495 assert(!(isa<VectorType>(CV->getType()) && static_cast<bool>(MDFrom)) &&
2496 "If we have a vector conditional, we should be propagating "
2497 "profile information.");
2498 Result.addVector(V: Builder.CreateSelect(C: CV, True: AV, False: BV, Name: "", MDFrom));
2499 }
2500
2501 return Result.addNumComputeOps(N: getNumOps(VT: Result.getVectorTy()) *
2502 Result.getNumVectors());
2503 }
2504
2505 /// Helper to linearize a matrix expression tree into a string. Currently
2506 /// matrix expressions are linarized by starting at an expression leaf and
2507 /// linearizing bottom up.
2508 struct ExprLinearizer {
2509 unsigned LengthToBreak = 100;
2510 std::string Str;
2511 raw_string_ostream Stream;
2512 unsigned LineLength = 0;
2513 const DataLayout &DL;
2514
2515 /// Mapping from instructions to matrixes. It is used to identify
2516 /// matrix instructions.
2517 const MapVector<Value *, MatrixTy> &Inst2Matrix;
2518
2519 /// Mapping from values to the leaves of all expressions that the value is
2520 /// part of.
2521 const DenseMap<Value *, SmallPtrSet<Value *, 2>> &Shared;
2522
2523 /// Set of matrix expressions in the scope of a given DISubprogram.
2524 const SmallSetVector<Value *, 32> &ExprsInSubprogram;
2525
2526 /// Leaf node of the expression to linearize.
2527 Value *Leaf;
2528
2529 /// Used to keep track of sub-expressions that get reused while linearizing
2530 /// the expression. Re-used sub-expressions are marked as (reused).
2531 SmallPtrSet<Value *, 8> ReusedExprs;
2532
2533 ExprLinearizer(const DataLayout &DL,
2534 const MapVector<Value *, MatrixTy> &Inst2Matrix,
2535 const DenseMap<Value *, SmallPtrSet<Value *, 2>> &Shared,
2536 const SmallSetVector<Value *, 32> &ExprsInSubprogram,
2537 Value *Leaf)
2538 : Stream(Str), DL(DL), Inst2Matrix(Inst2Matrix), Shared(Shared),
2539 ExprsInSubprogram(ExprsInSubprogram), Leaf(Leaf) {}
2540
2541 void indent(unsigned N) {
2542 LineLength += N;
2543 for (unsigned i = 0; i < N; i++)
2544 Stream << " ";
2545 }
2546
2547 void lineBreak() {
2548 Stream << "\n";
2549 LineLength = 0;
2550 }
2551
2552 void maybeIndent(unsigned Indent) {
2553 if (LineLength >= LengthToBreak)
2554 lineBreak();
2555
2556 if (LineLength == 0)
2557 indent(N: Indent);
2558 }
2559
2560 void write(StringRef S) {
2561 LineLength += S.size();
2562 Stream << S;
2563 }
2564
2565 Value *getUnderlyingObjectThroughLoads(Value *V) {
2566 if (Value *Ptr = getPointerOperand(V))
2567 return getUnderlyingObjectThroughLoads(V: Ptr);
2568 else if (V->getType()->isPointerTy())
2569 return getUnderlyingObject(V);
2570 return V;
2571 }
2572
2573 /// Returns true if \p V is a matrix value in the given subprogram.
2574 bool isMatrix(Value *V) const { return ExprsInSubprogram.count(key: V); }
2575
2576 /// If \p V is a matrix value, print its shape as NumRows x NumColumns to
2577 /// \p SS.
2578 void prettyPrintMatrixType(Value *V, raw_string_ostream &SS) {
2579 auto M = Inst2Matrix.find(Key: V);
2580 if (M == Inst2Matrix.end())
2581 SS << "unknown";
2582 else {
2583 SS << M->second.getNumRows();
2584 SS << "x";
2585 SS << M->second.getNumColumns();
2586 }
2587 }
2588
2589 /// Write the called function name. Handles calls to llvm.matrix.*
2590 /// specially: we write the name, followed by the dimensions of the input
2591 /// matrixes, followed by the scalar type name.
2592 void writeFnName(CallInst *CI) {
2593 if (!CI->getCalledFunction())
2594 write(S: "<no called fn>");
2595 else {
2596 StringRef Name = CI->getCalledFunction()->getName();
2597 if (!Name.starts_with(Prefix: "llvm.matrix")) {
2598 write(S: Name);
2599 return;
2600 }
2601 auto *II = cast<IntrinsicInst>(Val: CI);
2602 write(S: Intrinsic::getBaseName(id: II->getIntrinsicID())
2603 .drop_front(N: StringRef("llvm.matrix.").size()));
2604 write(S: ".");
2605 std::string Tmp;
2606 raw_string_ostream SS(Tmp);
2607
2608 switch (II->getIntrinsicID()) {
2609 case Intrinsic::matrix_multiply:
2610 prettyPrintMatrixType(V: II->getOperand(i_nocapture: 0), SS);
2611 SS << ".";
2612 prettyPrintMatrixType(V: II->getOperand(i_nocapture: 1), SS);
2613 SS << "." << *II->getType()->getScalarType();
2614 break;
2615 case Intrinsic::matrix_transpose:
2616 prettyPrintMatrixType(V: II->getOperand(i_nocapture: 0), SS);
2617 SS << "." << *II->getType()->getScalarType();
2618 break;
2619 case Intrinsic::matrix_column_major_load:
2620 prettyPrintMatrixType(V: II, SS);
2621 SS << "." << *II->getType()->getScalarType();
2622 break;
2623 case Intrinsic::matrix_column_major_store:
2624 prettyPrintMatrixType(V: II->getOperand(i_nocapture: 0), SS);
2625 SS << "." << *II->getOperand(i_nocapture: 0)->getType()->getScalarType();
2626 break;
2627 default:
2628 llvm_unreachable("Unhandled case");
2629 }
2630 write(S: Tmp);
2631 }
2632 }
2633
2634 unsigned getNumShapeArgs(CallInst *CI) const {
2635 if (IntrinsicInst *II = dyn_cast<IntrinsicInst>(Val: CI)) {
2636 switch (II->getIntrinsicID()) {
2637 case Intrinsic::matrix_multiply:
2638 return 3;
2639 case Intrinsic::matrix_transpose:
2640 return 2;
2641 case Intrinsic::matrix_column_major_load:
2642 case Intrinsic::matrix_column_major_store:
2643 return 3;
2644 default:
2645 return 0;
2646 }
2647 }
2648 return 0;
2649 }
2650
2651 /// Special printing for values: for pointers, we print if they refer to an
2652 /// (function) external address or a stack address, for other values we
2653 /// either print the constant or "scalar"/"matrix" for other values.
2654 void write(Value *V) {
2655 V = getUnderlyingObjectThroughLoads(V);
2656 if (V->getType()->isPointerTy()) {
2657 if (isa<AllocaInst>(Val: V)) {
2658 Stream << "stack addr";
2659 LineLength += StringRef("stack addr").size();
2660 } else {
2661 Stream << "addr";
2662 LineLength += StringRef("addr").size();
2663 }
2664 if (!V->getName().empty()) {
2665 Stream << " %" << V->getName() << "";
2666 LineLength += V->getName().size() + 2;
2667 }
2668 return;
2669 }
2670
2671 std::string Tmp;
2672 raw_string_ostream TmpStream(Tmp);
2673
2674 if (auto *CI = dyn_cast<ConstantInt>(Val: V))
2675 TmpStream << CI->getValue();
2676 else if (isa<Constant>(Val: V))
2677 TmpStream << "constant";
2678 else {
2679 if (isMatrix(V))
2680 TmpStream << "matrix";
2681 else
2682 TmpStream << "scalar";
2683 }
2684 Tmp = std::string(StringRef(Tmp).trim());
2685 LineLength += Tmp.size();
2686 Stream << Tmp;
2687 }
2688
2689 /// Linearize expression \p Expr starting at an indentation of \p Indent.
2690 /// Expressions that are re-used multiple times are prefixed with (reused)
2691 /// at the re-used root instruction.
2692 void linearizeExpr(Value *Expr, unsigned Indent, bool ParentReused,
2693 bool ParentShared) {
2694 auto *I = cast<Instruction>(Val: Expr);
2695 maybeIndent(Indent);
2696 SmallVector<Value *, 8> Ops;
2697
2698 // Is Expr shared with other expression leaves?
2699 bool ExprShared = false;
2700
2701 // Deal with shared subtrees. Mark them as shared, if required.
2702 if (!ParentShared) {
2703 auto SI = Shared.find(Val: Expr);
2704 assert(SI != Shared.end() && SI->second.count(Leaf));
2705
2706 for (Value *S : SI->second) {
2707 if (S == Leaf)
2708 continue;
2709 DebugLoc DL = cast<Instruction>(Val: S)->getDebugLoc();
2710 write(S: "shared with remark at line " + std::to_string(val: DL.getLine()) +
2711 " column " + std::to_string(val: DL.getCol()) + " (");
2712 }
2713 ExprShared = SI->second.size() > 1;
2714 }
2715
2716 bool Reused = !ReusedExprs.insert(Ptr: Expr).second;
2717 if (Reused && !ParentReused)
2718 write(S: "(reused) ");
2719
2720 if (auto *CI = dyn_cast<CallInst>(Val: I)) {
2721 writeFnName(CI);
2722
2723 Ops.append(in_start: CI->arg_begin(), in_end: CI->arg_end() - getNumShapeArgs(CI));
2724 } else if (isa<BitCastInst>(Val: Expr)) {
2725 // Special case bitcasts, which are used to materialize matrixes from
2726 // non-matrix ops.
2727 write(S: "matrix");
2728 return;
2729 } else {
2730 Ops.append(in_start: I->value_op_begin(), in_end: I->value_op_end());
2731 write(S: I->getOpcodeName());
2732 }
2733
2734 write(S: "(");
2735
2736 unsigned NumOpsToBreak = 1;
2737 if (match(V: Expr, P: m_Intrinsic<Intrinsic::matrix_column_major_load>()))
2738 NumOpsToBreak = 2;
2739
2740 for (Value *Op : Ops) {
2741 if (Ops.size() > NumOpsToBreak)
2742 lineBreak();
2743
2744 maybeIndent(Indent: Indent + 1);
2745 if (isMatrix(V: Op))
2746 linearizeExpr(Expr: Op, Indent: Indent + 1, ParentReused: Reused, ParentShared: ExprShared);
2747 else
2748 write(V: Op);
2749 if (Op != Ops.back())
2750 write(S: ", ");
2751 }
2752
2753 write(S: ")");
2754 }
2755
2756 const std::string &getResult() {
2757 return Str;
2758 }
2759 };
2760
2761 /// Generate remarks for matrix operations in a function. To generate remarks
2762 /// for matrix expressions, the following approach is used:
2763 /// 1. Use the inlined-at debug information to group matrix operations to the
2764 /// DISubprograms they are contained in.
2765 /// 2. Collect leaves of matrix expressions (done in
2766 /// RemarkGenerator::getExpressionLeaves) for each subprogram - expression
2767 // mapping. Leaves are lowered matrix instructions without other matrix
2768 // users (like stores) in the current subprogram.
2769 /// 3. For each leaf, create a remark containing a linearizied version of the
2770 /// matrix expression. The expression is linearized by a recursive
2771 /// bottom-up traversal of the matrix operands, starting at a leaf. Note
2772 /// that multiple leaves can share sub-expressions. Shared subexpressions
2773 /// are explicitly marked as shared().
2774 struct RemarkGenerator {
2775 const MapVector<Value *, MatrixTy> &Inst2Matrix;
2776 OptimizationRemarkEmitter &ORE;
2777 Function &Func;
2778 const DataLayout &DL;
2779
2780 RemarkGenerator(const MapVector<Value *, MatrixTy> &Inst2Matrix,
2781 OptimizationRemarkEmitter &ORE, Function &Func)
2782 : Inst2Matrix(Inst2Matrix), ORE(ORE), Func(Func),
2783 DL(Func.getDataLayout()) {}
2784
2785 /// Return all leaves of the expressions in \p ExprsInSubprogram. Those are
2786 /// instructions in Inst2Matrix returning void or without any users in
2787 /// \p ExprsInSubprogram. Currently that should only include stores.
2788 SmallVector<Value *, 4>
2789 getExpressionLeaves(const SmallSetVector<Value *, 32> &ExprsInSubprogram) {
2790 SmallVector<Value *, 4> Leaves;
2791 for (auto *Expr : ExprsInSubprogram)
2792 if (Expr->getType()->isVoidTy() ||
2793 !any_of(Range: Expr->users(), P: [&ExprsInSubprogram](User *U) {
2794 return ExprsInSubprogram.count(key: U);
2795 }))
2796 Leaves.push_back(Elt: Expr);
2797 return Leaves;
2798 }
2799
2800 /// Recursively traverse expression \p V starting at \p Leaf and add \p Leaf
2801 /// to all visited expressions in \p Shared. Limit the matrix operations to
2802 /// the ones in \p ExprsInSubprogram.
2803 void collectSharedInfo(Value *Leaf, Value *V,
2804 const SmallSetVector<Value *, 32> &ExprsInSubprogram,
2805 DenseMap<Value *, SmallPtrSet<Value *, 2>> &Shared) {
2806
2807 if (!ExprsInSubprogram.count(key: V))
2808 return;
2809
2810 Shared[V].insert(Ptr: Leaf);
2811
2812 for (Value *Op : cast<Instruction>(Val: V)->operand_values())
2813 collectSharedInfo(Leaf, V: Op, ExprsInSubprogram, Shared);
2814 }
2815
2816 /// Calculate the number of exclusive and shared op counts for expression
2817 /// starting at \p V. Expressions used multiple times are counted once.
2818 /// Limit the matrix operations to the ones in \p ExprsInSubprogram.
2819 std::pair<OpInfoTy, OpInfoTy>
2820 sumOpInfos(Value *Root, SmallPtrSetImpl<Value *> &ReusedExprs,
2821 const SmallSetVector<Value *, 32> &ExprsInSubprogram,
2822 DenseMap<Value *, SmallPtrSet<Value *, 2>> &Shared) const {
2823 if (!ExprsInSubprogram.count(key: Root))
2824 return {};
2825
2826 // Already counted this expression. Stop.
2827 if (!ReusedExprs.insert(Ptr: Root).second)
2828 return {};
2829
2830 OpInfoTy SharedCount;
2831 OpInfoTy Count;
2832
2833 auto I = Shared.find(Val: Root);
2834 auto CM = Inst2Matrix.find(Key: Root);
2835 if (I->second.size() == 1)
2836 Count = CM->second.getOpInfo();
2837 else
2838 SharedCount = CM->second.getOpInfo();
2839
2840 for (Value *Op : cast<Instruction>(Val: Root)->operand_values()) {
2841 auto C = sumOpInfos(Root: Op, ReusedExprs, ExprsInSubprogram, Shared);
2842 Count += C.first;
2843 SharedCount += C.second;
2844 }
2845 return {Count, SharedCount};
2846 }
2847
2848 void emitRemarks() {
2849 if (!ORE.allowExtraAnalysis(DEBUG_TYPE))
2850 return;
2851
2852 // Map matrix operations to their containting subprograms, by traversing
2853 // the inlinedAt chain. If the function does not have a DISubprogram, we
2854 // only map them to the containing function.
2855 MapVector<DISubprogram *, SmallVector<Value *, 8>> Subprog2Exprs;
2856 for (const auto &KV : Inst2Matrix) {
2857 if (Func.getSubprogram()) {
2858 auto *I = cast<Instruction>(Val: KV.first);
2859 DILocation *Context = I->getDebugLoc();
2860 while (Context) {
2861 Subprog2Exprs[getSubprogram(Scope: Context->getScope())].push_back(
2862 Elt: KV.first);
2863 Context = DebugLoc(Context).getInlinedAt();
2864 }
2865 } else {
2866 Subprog2Exprs[nullptr].push_back(Elt: KV.first);
2867 }
2868 }
2869 for (auto &KV : Subprog2Exprs) {
2870 SmallSetVector<Value *, 32> ExprsInSubprogram(KV.second.begin(),
2871 KV.second.end());
2872 auto Leaves = getExpressionLeaves(ExprsInSubprogram);
2873
2874 DenseMap<Value *, SmallPtrSet<Value *, 2>> Shared;
2875 for (Value *Leaf : Leaves)
2876 collectSharedInfo(Leaf, V: Leaf, ExprsInSubprogram, Shared);
2877
2878 // Generate remarks for each leaf.
2879 for (auto *L : Leaves) {
2880
2881 DebugLoc Loc = cast<Instruction>(Val: L)->getDebugLoc();
2882 DILocation *Context = cast<Instruction>(Val: L)->getDebugLoc();
2883 while (Context) {
2884 if (getSubprogram(Scope: Context->getScope()) == KV.first) {
2885 Loc = Context;
2886 break;
2887 }
2888 Context = DebugLoc(Context).getInlinedAt();
2889 }
2890
2891 SmallPtrSet<Value *, 8> ReusedExprs;
2892 OpInfoTy Counts, SharedCounts;
2893 std::tie(args&: Counts, args&: SharedCounts) =
2894 sumOpInfos(Root: L, ReusedExprs, ExprsInSubprogram, Shared);
2895
2896 OptimizationRemark Rem(DEBUG_TYPE, "matrix-lowered", Loc,
2897 cast<Instruction>(Val: L)->getParent());
2898
2899 Rem << "Lowered with ";
2900 Rem << ore::NV("NumStores", Counts.NumStores) << " stores, "
2901 << ore::NV("NumLoads", Counts.NumLoads) << " loads, "
2902 << ore::NV("NumComputeOps", Counts.NumComputeOps)
2903 << " compute ops, "
2904 << ore::NV("NumExposedTransposes", Counts.NumExposedTransposes)
2905 << " exposed transposes";
2906
2907 if (SharedCounts.NumStores > 0 || SharedCounts.NumLoads > 0 ||
2908 SharedCounts.NumComputeOps > 0) {
2909 Rem << ",\nadditionally "
2910 << ore::NV("NumStores", SharedCounts.NumStores) << " stores, "
2911 << ore::NV("NumLoads", SharedCounts.NumLoads) << " loads, "
2912 << ore::NV("NumFPOps", SharedCounts.NumComputeOps)
2913 << " compute ops"
2914 << " are shared with other expressions";
2915 }
2916
2917 Rem << ("\n" + linearize(L, Shared, ExprsInSubprogram, DL));
2918 ORE.emit(OptDiag&: Rem);
2919 }
2920 }
2921 }
2922
2923 std::string
2924 linearize(Value *L,
2925 const DenseMap<Value *, SmallPtrSet<Value *, 2>> &Shared,
2926 const SmallSetVector<Value *, 32> &ExprsInSubprogram,
2927 const DataLayout &DL) {
2928 ExprLinearizer Lin(DL, Inst2Matrix, Shared, ExprsInSubprogram, L);
2929 Lin.linearizeExpr(Expr: L, Indent: 0, ParentReused: false, ParentShared: false);
2930 return Lin.getResult();
2931 }
2932 };
2933};
2934} // namespace
2935
2936PreservedAnalyses LowerMatrixIntrinsicsPass::run(Function &F,
2937 FunctionAnalysisManager &AM) {
2938 auto &TTI = AM.getResult<TargetIRAnalysis>(IR&: F);
2939
2940 LowerMatrixIntrinsics LMT(F, TTI, Minimal ? nullptr : &AM);
2941 if (LMT.Visit()) {
2942 PreservedAnalyses PA;
2943 if (!Minimal) {
2944 PA.preserve<LoopAnalysis>();
2945 PA.preserve<DominatorTreeAnalysis>();
2946 }
2947 return PA;
2948 }
2949 return PreservedAnalyses::all();
2950}
2951
2952void LowerMatrixIntrinsicsPass::printPipeline(
2953 raw_ostream &OS, function_ref<StringRef(StringRef)> MapClassName2PassName) {
2954 static_cast<PassInfoMixin<LowerMatrixIntrinsicsPass> *>(this)->printPipeline(
2955 OS, MapClassName2PassName);
2956 OS << '<';
2957 if (Minimal)
2958 OS << "minimal";
2959 OS << '>';
2960}
2961