@@ -125,20 +125,6 @@ class DimConversionPattern : public mlir::ConversionPattern {
125125 mlir::ConversionPatternRewriter &rewriter) const override ;
126126};
127127
128- /* * Template for rewrites (without type conversion). */
129- template <typename T>
130- class DimOpRewritePattern : public mlir ::OpRewritePattern<T> {
131- private:
132- const DimMapper &_mapper;
133-
134- mlir::LogicalResult
135- matchAndRewrite (T op, mlir::PatternRewriter &rewriter) const override ;
136-
137- public:
138- DimOpRewritePattern (const DimMapper &mapper, mlir::MLIRContext *ctx)
139- : mlir::OpRewritePattern<T>(ctx), _mapper(mapper) {}
140- };
141-
142128} // namespace
143129
144130mlir::FailureOr<DimMapper>
@@ -309,31 +295,33 @@ mlir::LogicalResult DimConversionPattern::matchAndRewrite(
309295 return mlir::success ();
310296}
311297
312- template <>
313- mlir::LogicalResult DimOpRewritePattern<CastDimOp>::matchAndRewrite(
314- CastDimOp op, mlir::PatternRewriter &rewriter) const {
315- auto dim = _mapper.convertAttr (op.getInput ());
316- if (!dim) {
317- return mlir::failure ();
298+ static mlir::LogicalResult updateDim (ForOp op, DimMapper &mapper) {
299+ if (!op.getIters () || !op.getIters ()->isAbstract ()) {
300+ // No update needed
301+ return mlir::success ();
318302 }
319303
320- // The folder on CastDimOp should turn this into a constant.
321- auto newOp = rewriter.createOrFold <CastDimOp>(op->getLoc (), dim);
322- rewriter.replaceOp (op, newOp);
304+ auto dim = mapper.convertAttr (*op.getIters ());
305+ if (!dim) {
306+ return op.emitOpError (" no mapping for " ) << *op.getIters ();
307+ }
323308
309+ op.setItersAttr (dim);
324310 return mlir::success ();
325311}
326312
327- template <>
328- mlir::LogicalResult DimOpRewritePattern<ForDimOp>::matchAndRewrite(
329- ForDimOp op, mlir::PatternRewriter &rewriter) const {
330- auto dim = _mapper.convertAttr (op.getDim ());
331- if (!dim) {
332- return mlir::failure ();
313+ static mlir::LogicalResult updateDim (CastDimOp op, DimMapper &mapper) {
314+ if (!op.getInput ().isAbstract ()) {
315+ // No update needed
316+ return mlir::success ();
333317 }
334318
335- rewriter.modifyOpInPlace (op, [&]() { op.setDimAttr (dim); });
319+ auto dim = mapper.convertAttr (op.getInput ());
320+ if (!dim) {
321+ return op.emitOpError (" no mapping for " ) << op.getInput ();
322+ }
336323
324+ op.setInputAttr (dim);
337325 return mlir::success ();
338326}
339327
@@ -356,6 +344,19 @@ void GraphAlgSetDimensions::runOnOperation() {
356344 return signalPassFailure ();
357345 }
358346
347+ // Update direct references to dimensions.
348+ bool failedDirectUpdate = false ;
349+ func->walk ([&](ForOp op) {
350+ if (mlir::failed (updateDim (op, *dimMapper))) {
351+ failedDirectUpdate = true ;
352+ }
353+ });
354+ func->walk ([&](CastDimOp op) {
355+ if (mlir::failed (updateDim (op, *dimMapper))) {
356+ failedDirectUpdate = true ;
357+ }
358+ });
359+
359360 mlir::ConversionTarget target (getContext ());
360361 target.addDynamicallyLegalDialect <GraphAlgDialect>(
361362 doesNotUseAbstractDimensions);
@@ -374,12 +375,6 @@ void GraphAlgSetDimensions::runOnOperation() {
374375 // Convert all result types and block argument types.
375376 patterns.add <DimConversionPattern>(typeConverter, &getContext ());
376377
377- // Convert ops that have a special dependency on DimAttr.
378- patterns.add <DimOpRewritePattern<CastDimOp>, DimOpRewritePattern<ForDimOp>>(
379- *dimMapper, &getContext ());
380- // Use the canonicalization pattern to rewrite ForDimOp into ForConstOp.
381- ForDimOp::getCanonicalizationPatterns (patterns, &getContext ());
382-
383378 if (mlir::failed (
384379 mlir::applyPartialConversion (func, target, std::move (patterns)))) {
385380 return signalPassFailure ();
0 commit comments