@@ -320,6 +320,11 @@ static mlir::Value createDimRead(mlir::Location loc, graphalg::DimAttr dim,
320320 */
321321static mlir::Value createDimInt (mlir::Location loc, graphalg::DimAttr dim,
322322 mlir::OpBuilder &builder) {
323+ if (dim.isConcrete ()) {
324+ return builder.create <ConstantOp>(
325+ loc, builder.getI64IntegerAttr (dim.getConcreteDim ()));
326+ }
327+
323328 auto input = createDimRead (loc, dim, builder);
324329 llvm::ArrayRef<ColumnIdx> groupBy;
325330 std::array<AggregatorAttr, 1 > aggregators{
@@ -889,71 +894,48 @@ mlir::LogicalResult OpConversion<graphalg::DiagOp>::matchAndRewrite(
889894 return mlir::success ();
890895}
891896
892- // Sharing logic between ForConstOp and ForDimOp
893- static mlir::LogicalResult
894- convertFor (mlir::Operation *op, mlir::ValueRange adaptorInitArgs,
895- mlir::Value rangeBegin, mlir::Value rangeEnd, mlir::Region &body,
896- mlir::Region &until, const mlir::TypeConverter *typeConverter,
897- mlir::ConversionPatternRewriter &rewriter) {
898- llvm::SmallVector<mlir::Value> initArgs{rangeBegin};
899- initArgs.append (adaptorInitArgs.begin (), adaptorInitArgs.end ());
897+ template <>
898+ mlir::LogicalResult OpConversion<graphalg::ForOp>::matchAndRewrite(
899+ graphalg::ForOp op, OpAdaptor adaptor,
900+ mlir::ConversionPatternRewriter &rewriter) const {
901+ auto begin =
902+ rewriter.create <ConstantOp>(op.getLoc (), rewriter.getI64IntegerAttr (0 ));
903+ auto iters = createDimInt (op.getLoc (), *op.getIters (), rewriter);
904+ llvm::SmallVector<mlir::Value> initArgs{begin};
905+ initArgs.append (adaptor.getInitArgs ().begin (), adaptor.getInitArgs ().end ());
900906
901- auto blockSignature = typeConverter->convertBlockSignature (&body.front ());
907+ auto blockSignature =
908+ typeConverter->convertBlockSignature (&op.getBody ().front ());
902909 if (!blockSignature) {
903910 return op->emitOpError (" Failed to convert iter args" );
904911 }
905912
906- mlir::Value iters;
907- if (isConstantZeroI64 (rangeBegin)) {
908- iters = rangeEnd;
909- } else {
910- // Subtract rangeBegin from rangeEnd
911- auto joinOp = rewriter.create <JoinOp>(
912- op->getLoc (), mlir::ValueRange{rangeBegin, rangeEnd},
913- rewriter.getAttr <JoinPredicatesAttr>(
914- llvm::ArrayRef<JoinPredicateAttr>{}));
915- auto projOp = rewriter.create <ProjectOp>(
916- op->getLoc (), getI64RelationType (op->getContext ()), joinOp);
917-
918- auto &block = projOp.createProjectionsBlock ();
919- mlir::OpBuilder::InsertionGuard guard{rewriter};
920- rewriter.setInsertionPointToStart (&block);
921-
922- auto begin =
923- rewriter.create <ExtractOp>(op->getLoc (), 0 , block.getArgument (0 ));
924- auto end =
925- rewriter.create <ExtractOp>(op->getLoc (), 1 , block.getArgument (0 ));
926- auto res = rewriter.create <mlir::arith::SubIOp>(op->getLoc (), end, begin);
927- rewriter.create <ProjectReturnOp>(op->getLoc (), mlir::ValueRange{res});
928-
929- iters = projOp;
930- }
931-
932913 // The relational version of this op can only have a single output value.
933914 // For loops with multiple results, duplicate.
934915 llvm::SmallVector<mlir::Value> resultValues;
935916 for (auto i : llvm::seq (op->getNumResults ())) {
936917 auto result = op->getResult (i);
937918 if (result.use_empty ()) {
938919 // Not used. Take init arg as a dummy value.
939- resultValues.push_back (adaptorInitArgs [i]);
920+ resultValues.push_back (adaptor. getInitArgs () [i]);
940921 continue ;
941922 }
942923
943924 // We are adding the iteration count variable as a first argument, so offset
944925 // the result index accordingly.
945926 std::int64_t resultIdx = i + 1 ;
946- auto resultType = adaptorInitArgs [i].getType ();
927+ auto resultType = adaptor. getInitArgs () [i].getType ();
947928 auto forOp = rewriter.create <ForOp>(op->getLoc (), resultType, initArgs,
948929 iters, resultIdx);
949930 // body block
950- rewriter.cloneRegionBefore (body, forOp.getBody (), forOp.getBody ().begin ());
931+ rewriter.cloneRegionBefore (op.getBody (), forOp.getBody (),
932+ forOp.getBody ().begin ());
951933 rewriter.applySignatureConversion (&forOp.getBody ().front (),
952934 *blockSignature);
953935
954936 // until block
955- if (!until .empty ()) {
956- rewriter.cloneRegionBefore (until , forOp.getUntil (),
937+ if (!op. getUntil () .empty ()) {
938+ rewriter.cloneRegionBefore (op. getUntil () , forOp.getUntil (),
957939 forOp.getUntil ().begin ());
958940 rewriter.applySignatureConversion (&forOp.getUntil ().front (),
959941 *blockSignature);
@@ -966,28 +948,6 @@ convertFor(mlir::Operation *op, mlir::ValueRange adaptorInitArgs,
966948 return mlir::success ();
967949}
968950
969- template <>
970- mlir::LogicalResult OpConversion<graphalg::ForConstOp>::matchAndRewrite(
971- graphalg::ForConstOp op, OpAdaptor adaptor,
972- mlir::ConversionPatternRewriter &rewriter) const {
973- return convertFor (op, adaptor.getInitArgs (), adaptor.getRangeBegin (),
974- adaptor.getRangeEnd (), op.getBody (), op.getUntil (),
975- typeConverter, rewriter);
976- }
977-
978- template <>
979- mlir::LogicalResult OpConversion<graphalg::ForDimOp>::matchAndRewrite(
980- graphalg::ForDimOp op, OpAdaptor adaptor,
981- mlir::ConversionPatternRewriter &rewriter) const {
982- auto ctx = op.getContext ();
983- auto rangeBegin =
984- rewriter.create <ConstantOp>(op.getLoc (), rewriter.getI64IntegerAttr (0 ));
985- auto rangeEnd = createDimInt (op.getLoc (), op.getDim (), rewriter);
986-
987- return convertFor (op, adaptor.getInitArgs (), rangeBegin, rangeEnd,
988- op.getBody (), op.getUntil (), typeConverter, rewriter);
989- }
990-
991951template <>
992952mlir::LogicalResult OpConversion<graphalg::YieldOp>::matchAndRewrite(
993953 graphalg::YieldOp op, OpAdaptor adaptor,
@@ -1481,11 +1441,10 @@ void GraphAlgToRel::runOnOperation() {
14811441 OpConversion<graphalg::TransposeOp>, OpConversion<graphalg::BroadcastOp>,
14821442 OpConversion<graphalg::ConstantMatrixOp>,
14831443 OpConversion<graphalg::DeferredReduceOp>, OpConversion<graphalg::DiagOp>,
1484- OpConversion<graphalg::ForConstOp>, OpConversion<graphalg::ForDimOp>,
1485- OpConversion<graphalg::YieldOp>, OpConversion<graphalg::MatMulJoinOp>,
1486- OpConversion<graphalg::PickAnyOp>, OpConversion<graphalg::TrilOp>,
1487- OpConversion<graphalg::UnionOp>, OpConversion<graphalg::CastDimOp>>(
1488- matrixTypeConverter, &getContext ());
1444+ OpConversion<graphalg::ForOp>, OpConversion<graphalg::YieldOp>,
1445+ OpConversion<graphalg::MatMulJoinOp>, OpConversion<graphalg::PickAnyOp>,
1446+ OpConversion<graphalg::TrilOp>, OpConversion<graphalg::UnionOp>,
1447+ OpConversion<graphalg::CastDimOp>>(matrixTypeConverter, &getContext ());
14891448 patterns.add <ApplyOpConversion>(semiringTypeConverter, matrixTypeConverter,
14901449 &getContext ());
14911450
0 commit comments