@@ -653,8 +653,9 @@ struct RewriteAvgPoolAsConv2D
653653 auto filterTy = cast<RankedTensorType>(poolOp.getInputs ()[1 ].getType ());
654654 auto outputTy = cast<RankedTensorType>(poolOp.getResultTypes ()[0 ]);
655655
656- auto kernelShape =
657- SmallVector<int64_t >{filterTy.getDimSize (0 ), filterTy.getDimSize (1 )};
656+ auto c = inputTy.getDimSize (1 );
657+ auto kernelShape = SmallVector<int64_t >{c, c, filterTy.getDimSize (0 ),
658+ filterTy.getDimSize (1 )};
658659 auto kernelTy =
659660 RankedTensorType::get (kernelShape, filterTy.getElementType ());
660661 TypedAttr kernelVals = rewriter.getOneAttr (kernelTy);
@@ -681,55 +682,17 @@ struct RewriteAvgPoolAsConv2D
681682 }
682683 auto kernel =
683684 arith::ConstantOp::create (rewriter, poolOp.getLoc (), kernelVals);
684-
685- // Rewrite the 2D avg pool output shape is N x C x H' x W'. Apply the kernel
686- // to each channel separately, and insert into the output.
687- auto outputVal = poolOp.getOutputs ()[0 ];
688- // The filter must ensure each output channel i is only the sum of values
689- // from input channel i. So the filter uses an identity matrix when f ==
690- // c and 0 otherwise.
691- RankedTensorType twoDOutputType =
692- RankedTensorType::get ({outputTy.getDimSize (2 ), outputTy.getDimSize (3 )},
693- outputTy.getElementType ());
694- RankedTensorType twoDInputType =
695- RankedTensorType::get ({inputTy.getDimSize (2 ), inputTy.getDimSize (3 )},
696- inputTy.getElementType ());
697- Value convOutput = tensor::EmptyOp::create (rewriter, poolOp.getLoc (),
698- twoDOutputType.getShape (),
699- twoDOutputType.getElementType ());
700- for (int n = 0 ; n < inputTy.getDimSize (0 ); ++n) {
701- for (int c = 0 ; c < inputTy.getDimSize (1 ); ++c) {
702- // Compute the 2-D constant convolution.
703- SmallVector<OpFoldResult> offsets = {
704- rewriter.getIndexAttr (n), rewriter.getIndexAttr (c),
705- rewriter.getIndexAttr (0 ), rewriter.getIndexAttr (0 )};
706- SmallVector<OpFoldResult> inputSizes = {
707- rewriter.getIndexAttr (1 ), rewriter.getIndexAttr (1 ),
708- rewriter.getIndexAttr (inputTy.getDimSize (2 )),
709- rewriter.getIndexAttr (inputTy.getDimSize (3 ))};
710- SmallVector<OpFoldResult> strides (4 , rewriter.getIndexAttr (1 ));
711- auto extractInputOp = tensor::ExtractSliceOp::create (
712- rewriter, poolOp.getLoc (), twoDInputType, poolOp.getInputs ()[0 ],
713- offsets, inputSizes, strides);
714-
715- auto convOp = linalg::Conv2DOp::create (
716- rewriter, poolOp.getLoc (), twoDOutputType,
717- ValueRange{extractInputOp, kernel}, ValueRange{convOutput});
718- // Insert into the outputVal.
719- SmallVector<OpFoldResult> outputSizes = {
720- rewriter.getIndexAttr (1 ), rewriter.getIndexAttr (1 ),
721- rewriter.getIndexAttr (outputTy.getDimSize (2 )),
722- rewriter.getIndexAttr (outputTy.getDimSize (3 ))};
723- outputVal = tensor::InsertSliceOp::create (
724- rewriter, poolOp.getLoc (), convOp.getResult (0 ), outputVal, offsets,
725- outputSizes, strides);
726- }
727- }
685+ Value conv = linalg::Conv2DNchwFchwOp::create (
686+ rewriter, poolOp.getLoc (), outputTy,
687+ ValueRange{poolOp.getInputs ()[0 ], kernel},
688+ ValueRange{poolOp.getOutputs ()[0 ]}, poolOp.getStrides (),
689+ poolOp.getDilations ())
690+ .getResult (0 );
728691
729692 if (avgPoolOutput) {
730- rewriter.replaceAllUsesWith (avgPoolOutput, outputVal );
693+ rewriter.replaceAllUsesWith (avgPoolOutput, conv );
731694 } else {
732- rewriter.replaceOp (poolOp, outputVal );
695+ rewriter.replaceOp (poolOp, conv );
733696 }
734697 return success ();
735698 }
@@ -746,6 +709,14 @@ struct LowerConv2DNchwFchw
746709
747710 LogicalResult matchAndRewrite (mlir::linalg::Conv2DNchwFchwOp convOp,
748711 PatternRewriter& rewriter) const override {
712+ // Fails is strides > 1.
713+ if (!llvm::all_of (convOp.getStrides (), [](const APInt& element) {
714+ return element.getSExtValue () == 1 ;
715+ })) {
716+ return rewriter.notifyMatchFailure (convOp,
717+ " expected all ones for strides" );
718+ }
719+
749720 Location loc = convOp.getLoc ();
750721 Value image = convOp.getInputs ()[0 ];
751722 Value filter = convOp.getInputs ()[1 ];
0 commit comments