1 //===- Tiling.cpp - Implementation of linalg Tiling -----------------------===// 2 // 3 // Copyright 2019 The MLIR Authors. 4 // 5 // Licensed under the Apache License, Version 2.0 (the "License"); 6 // you may not use this file except in compliance with the License. 7 // You may obtain a copy of the License at 8 // 9 // http://www.apache.org/licenses/LICENSE-2.0 10 // 11 // Unless required by applicable law or agreed to in writing, software 12 // distributed under the License is distributed on an "AS IS" BASIS, 13 // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. 14 // See the License for the specific language governing permissions and 15 // limitations under the License. 16 // ============================================================================= 17 // 18 // This file implements the linalg dialect Tiling pass. 19 // 20 //===----------------------------------------------------------------------===// 21 22 #include "mlir/Dialect/Linalg/IR/LinalgOps.h" 23 #include "mlir/Dialect/Linalg/IR/LinalgTypes.h" 24 #include "mlir/Dialect/Linalg/Passes.h" 25 #include "mlir/Dialect/Linalg/Utils/Intrinsics.h" 26 #include "mlir/Dialect/Linalg/Utils/Utils.h" 27 #include "mlir/Dialect/LoopOps/LoopOps.h" 28 #include "mlir/EDSC/Helpers.h" 29 #include "mlir/IR/AffineExpr.h" 30 #include "mlir/IR/AffineExprVisitor.h" 31 #include "mlir/IR/AffineMap.h" 32 #include "mlir/IR/OpImplementation.h" 33 #include "mlir/Pass/Pass.h" 34 #include "mlir/Support/LLVM.h" 35 #include "mlir/Support/STLExtras.h" 36 #include "mlir/Transforms/FoldUtils.h" 37 38 #include "llvm/Support/CommandLine.h" 39 40 using namespace mlir; 41 using namespace mlir::edsc; 42 using namespace mlir::edsc::intrinsics; 43 using namespace mlir::linalg; 44 using namespace mlir::linalg::intrinsics; 45 using namespace mlir::loop; 46 47 #define DEBUG_TYPE "linalg-tiling" 48 49 static llvm::cl::OptionCategory clOptionsCategory(DEBUG_TYPE " options"); 50 static llvm::cl::list<unsigned> 51 clTileSizes("linalg-tile-sizes", 52 llvm::cl::desc("Tile sizes by which to tile linalg operations"), 53 llvm::cl::ZeroOrMore, llvm::cl::MiscFlags::CommaSeparated, 54 llvm::cl::cat(clOptionsCategory)); 55 56 static bool isZero(Value *v) { 57 return isa_and_nonnull<ConstantIndexOp>(v->getDefiningOp()) && 58 cast<ConstantIndexOp>(v->getDefiningOp()).getValue() == 0; 59 } 60 61 // Creates a number of ranges equal to the number of non-zero in `tileSizes`. 62 // One for each loop of the LinalgOp that is tiled. The `tileSizes` argument has 63 // one entry per surrounding loop. It uses zero as the convention that a 64 // particular loop is not tiled. This convention simplifies implementations by 65 // avoiding affine map manipulations. 66 // The returned ranges correspond to the loop ranges, in the proper order, that 67 // are tiled and for which new loops will be created. 68 static SmallVector<SubViewOp::Range, 4> 69 makeTiledLoopRanges(OpBuilder &b, Location loc, AffineMap map, 70 ArrayRef<Value *> allViewSizes, 71 ArrayRef<Value *> allTileSizes, OperationFolder *folder) { 72 assert(allTileSizes.size() == map.getNumResults()); 73 // Apply `map` to get view sizes in loop order. 74 auto viewSizes = applyMapToValues(b, loc, map, allViewSizes, folder); 75 SmallVector<Value *, 4> tileSizes(allTileSizes.begin(), allTileSizes.end()); 76 77 // Traverse the tile sizes, which are in loop order, erase zeros everywhere. 78 for (int idx = tileSizes.size() - 1; idx >= 0; --idx) { 79 if (isZero(tileSizes[idx])) { 80 viewSizes.erase(viewSizes.begin() + idx); 81 tileSizes.erase(tileSizes.begin() + idx); 82 } 83 } 84 85 // Create a new range with the applied tile sizes. 86 SmallVector<SubViewOp::Range, 4> res; 87 for (unsigned idx = 0, e = tileSizes.size(); idx < e; ++idx) { 88 res.push_back(SubViewOp::Range{constant_index(folder, 0), viewSizes[idx], 89 tileSizes[idx]}); 90 } 91 return res; 92 } 93 94 namespace { 95 // Helper visitor to determine whether an AffineExpr is tiled. 96 // This is achieved by traversing every AffineDimExpr with position `pos` and 97 // checking whether the corresponding `tileSizes[pos]` is non-zero. 98 // This also enforces only positive coefficients occur in multiplications. 99 // 100 // Example: 101 // `d0 + 2 * d1 + d3` is tiled by [0, 0, 0, 2] but not by [0, 0, 2, 0] 102 // 103 struct TileCheck : public AffineExprVisitor<TileCheck> { 104 TileCheck(ArrayRef<Value *> tileSizes) 105 : isTiled(false), tileSizes(tileSizes) {} 106 107 void visitDimExpr(AffineDimExpr expr) { 108 isTiled |= !isZero(tileSizes[expr.getPosition()]); 109 } 110 void visitAffineBinaryOpExpr(AffineBinaryOpExpr expr) { 111 visit(expr.getLHS()); 112 visit(expr.getRHS()); 113 if (expr.getKind() == mlir::AffineExprKind::Mul) 114 assert(expr.getRHS().cast<AffineConstantExpr>().getValue() > 0 && 115 "nonpositive multiplying coefficient"); 116 } 117 bool isTiled; 118 ArrayRef<Value *> tileSizes; 119 }; 120 } // namespace 121 122 static bool isTiled(AffineExpr expr, ArrayRef<Value *> tileSizes) { 123 if (!expr) 124 return false; 125 TileCheck t(tileSizes); 126 t.visit(expr); 127 return t.isTiled; 128 } 129 130 // Checks whether the view with index `viewIndex` within `linalgOp` varies with 131 // respect to a non-zero `tileSize`. 132 static bool isTiled(AffineMap map, ArrayRef<Value *> tileSizes) { 133 if (!map) 134 return false; 135 for (unsigned r = 0; r < map.getNumResults(); ++r) 136 if (isTiled(map.getResult(r), tileSizes)) 137 return true; 138 return false; 139 } 140 141 static SmallVector<Value *, 4> 142 makeTiledViews(OpBuilder &b, Location loc, LinalgOp linalgOp, 143 ArrayRef<Value *> ivs, ArrayRef<Value *> tileSizes, 144 ArrayRef<Value *> viewSizes, OperationFolder *folder) { 145 assert(ivs.size() == static_cast<size_t>(llvm::count_if( 146 llvm::make_range(tileSizes.begin(), tileSizes.end()), 147 [](Value *v) { return !isZero(v); })) && 148 "expected as many ivs as non-zero sizes"); 149 150 using edsc::intrinsics::select; 151 using edsc::op::operator+; 152 using edsc::op::operator<; 153 154 // Construct (potentially temporary) mins and maxes on which to apply maps 155 // that define tile subviews. 156 SmallVector<Value *, 8> lbs, subViewSizes; 157 for (unsigned idx = 0, idxIvs = 0, e = tileSizes.size(); idx < e; ++idx) { 158 bool isTiled = !isZero(tileSizes[idx]); 159 lbs.push_back(isTiled ? ivs[idxIvs++] : (Value *)constant_index(folder, 0)); 160 subViewSizes.push_back(isTiled ? tileSizes[idx] : viewSizes[idx]); 161 } 162 163 auto *op = linalgOp.getOperation(); 164 165 SmallVector<Value *, 4> res; 166 res.reserve(op->getNumOperands()); 167 auto viewIteratorBegin = linalgOp.getInputsAndOutputs().begin(); 168 for (unsigned viewIndex = 0; viewIndex < linalgOp.getNumInputsAndOutputs(); 169 ++viewIndex) { 170 Value *view = *(viewIteratorBegin + viewIndex); 171 unsigned rank = view->getType().cast<MemRefType>().getRank(); 172 auto map = loopToOperandRangesMaps(linalgOp)[viewIndex]; 173 // If the view is not tiled, we can use it as is. 174 if (!isTiled(map, tileSizes)) { 175 res.push_back(view); 176 continue; 177 } 178 179 // Construct a new subview for the tile. 180 SmallVector<Value *, 4> offsets, sizes, strides; 181 offsets.reserve(rank); 182 sizes.reserve(rank); 183 strides.reserve(rank); 184 for (unsigned r = 0; r < rank; ++r) { 185 if (!isTiled(map.getSubMap({r}), tileSizes)) { 186 offsets.push_back(constant_index(folder, 0)); 187 sizes.push_back(dim(view, r)); 188 strides.push_back(constant_index(folder, 1)); 189 continue; 190 } 191 192 // Tiling creates a new slice at the proper index, the slice step is 1 193 // (i.e. the slice view does not subsample, stepping occurs in the loop). 194 auto m = map.getSubMap({r}); 195 auto *offset = applyMapToValues(b, loc, m, lbs, folder).front(); 196 offsets.push_back(offset); 197 auto *size = applyMapToValues(b, loc, m, subViewSizes, folder).front(); 198 sizes.push_back(size); 199 strides.push_back(constant_index(folder, 1)); 200 } 201 // TODO(b/144419024) Atm std.subview is not guaranteed in-bounds. Depending 202 // on the semantics we attach to it, we may need to use min(size, dim) here 203 // and canonicalize later. 204 res.push_back(b.create<SubViewOp>(loc, view, offsets, sizes, strides)); 205 } 206 207 // Traverse the mins/maxes and erase those that don't have uses left. 208 // This is a special type of folding that we only apply when `folder` is 209 // defined. 210 if (folder) 211 for (auto *v : llvm::concat<Value *>(lbs, subViewSizes)) 212 if (v->use_empty()) 213 v->getDefiningOp()->erase(); 214 215 return res; 216 } 217 218 llvm::Optional<TiledLinalgOp> 219 mlir::linalg::tileLinalgOp(OpBuilder &b, LinalgOp op, 220 ArrayRef<Value *> tileSizes, 221 OperationFolder *folder) { 222 // 1. Enforce the convention that "tiling by zero" skips tiling a particular 223 // dimension. This convention is significantly simpler to handle instead of 224 // adjusting affine maps to account for missing dimensions. 225 assert(op.getNumParallelLoops() + op.getNumReductionLoops() + 226 op.getNumWindowLoops() == 227 tileSizes.size() && 228 "expected matching number of tile sizes and loops"); 229 OpBuilder::InsertionGuard g(b); 230 b.setInsertionPoint(op); 231 ScopedContext scope(b, op.getLoc()); 232 // 2. Build the tiled loop ranges. 233 auto viewSizes = getViewSizes(op); 234 // The flattened loopToOperandRangesMaps is expected to be an invertible 235 // permutation map (asserted in the inverse calculation). 236 auto viewSizesToLoopsMap = 237 inversePermutation(concatAffineMaps(loopToOperandRangesMaps(op))); 238 assert(viewSizesToLoopsMap && "expected invertible map"); 239 auto loopRanges = 240 makeTiledLoopRanges(b, scope.getLocation(), viewSizesToLoopsMap, 241 viewSizes, tileSizes, folder); 242 243 // 3. Create the tiled loops. 244 LinalgOp res = op; 245 SmallVector<IndexHandle, 4> ivs(loopRanges.size()); 246 auto pivs = makeIndexHandlePointers(ivs); 247 LoopNestRangeBuilder(pivs, loopRanges)([&] { 248 auto b = ScopedContext::getBuilder(); 249 auto loc = ScopedContext::getLocation(); 250 SmallVector<Value *, 4> ivValues(ivs.begin(), ivs.end()); 251 auto views = 252 makeTiledViews(b, loc, op, ivValues, tileSizes, viewSizes, folder); 253 auto operands = getAssumedNonViewOperands(op); 254 views.append(operands.begin(), operands.end()); 255 res = op.clone(b, loc, views); 256 }); 257 258 // 4. Gather the newly created loops and return them with the new op. 259 SmallVector<ForOp, 8> loops; 260 loops.reserve(ivs.size()); 261 for (auto iv : ivs) 262 loops.push_back(loop::getForInductionVarOwner(iv)); 263 264 return TiledLinalgOp{res, loops}; 265 } 266 267 llvm::Optional<TiledLinalgOp> 268 mlir::linalg::tileLinalgOp(OpBuilder &b, LinalgOp op, 269 ArrayRef<int64_t> tileSizes, 270 OperationFolder *folder) { 271 if (tileSizes.empty()) 272 return llvm::None; 273 274 // The following uses the convention that "tiling by zero" skips tiling a 275 // particular dimension. This convention is significantly simpler to handle 276 // instead of adjusting affine maps to account for missing dimensions. 277 auto nLoops = op.getNumParallelLoops() + op.getNumReductionLoops() + 278 op.getNumWindowLoops(); 279 tileSizes = tileSizes.take_front(nLoops); 280 // If only 0 tilings are left, then return. 281 if (llvm::all_of(tileSizes, [](int64_t v) { return v == 0; })) 282 return llvm::None; 283 284 // Create a builder for tile size constants. 285 OpBuilder::InsertionGuard g(b); 286 b.setInsertionPoint(op); 287 ScopedContext scope(b, op.getLoc()); 288 289 // Materialize concrete tile size values to pass the generic tiling function. 290 SmallVector<Value *, 8> tileSizeValues; 291 tileSizeValues.reserve(tileSizes.size()); 292 for (auto ts : tileSizes) 293 tileSizeValues.push_back(constant_index(folder, ts)); 294 // Pad tile sizes with zero values to enforce our convention. 295 if (tileSizeValues.size() < nLoops) { 296 for (unsigned i = tileSizeValues.size(); i < nLoops; ++i) 297 tileSizeValues.push_back(constant_index(folder, 0)); 298 } 299 300 return tileLinalgOp(b, op, tileSizeValues, folder); 301 } 302 303 static void tileLinalgOps(FuncOp f, ArrayRef<int64_t> tileSizes) { 304 OpBuilder b(f); 305 OperationFolder folder(f.getContext()); 306 f.walk([tileSizes, &b, &folder](LinalgOp op) { 307 auto opLoopsPair = tileLinalgOp(b, op, tileSizes, &folder); 308 // If tiling occurred successfully, erase old op. 309 if (opLoopsPair) 310 op.erase(); 311 }); 312 f.walk([](LinalgOp op) { 313 if (!op.getOperation()->hasNoSideEffect()) 314 return; 315 if (op.getOperation()->use_empty()) 316 op.erase(); 317 }); 318 } 319 320 namespace { 321 struct LinalgTilingPass : public FunctionPass<LinalgTilingPass> { 322 LinalgTilingPass() = default; 323 LinalgTilingPass(ArrayRef<int64_t> sizes); 324 325 void runOnFunction() override { tileLinalgOps(getFunction(), tileSizes); } 326 327 SmallVector<int64_t, 8> tileSizes; 328 }; 329 } // namespace 330 331 LinalgTilingPass::LinalgTilingPass(ArrayRef<int64_t> sizes) { 332 this->tileSizes.assign(sizes.begin(), sizes.end()); 333 } 334 335 std::unique_ptr<OpPassBase<FuncOp>> 336 mlir::linalg::createLinalgTilingPass(ArrayRef<int64_t> tileSizes) { 337 return std::make_unique<LinalgTilingPass>(tileSizes); 338 } 339 340 static PassRegistration<LinalgTilingPass> 341 pass("linalg-tile", "Tile operations in the linalg dialect", [] { 342 auto pass = std::make_unique<LinalgTilingPass>(); 343 pass->tileSizes.assign(clTileSizes.begin(), clTileSizes.end()); 344 return pass; 345 }); 346