Skip to content

Commit adc0358

Browse files
committed
Tests passing again.
1 parent 88f76ae commit adc0358

20 files changed

Lines changed: 148 additions & 582 deletions

File tree

compiler/include/graphalg/GraphAlgOps.td

Lines changed: 1 addition & 78 deletions
Original file line numberDiff line numberDiff line change
@@ -467,87 +467,10 @@ def ForOp : Core_Op<"for", [
467467
}];
468468
}
469469

470-
// Not core according to spec, but we don't want to unroll in the general case.
471-
def ForConstOp : Core_Op<"for_const", [
472-
Pure,
473-
AllTypesMatch<["rangeBegin", "rangeEnd"]>,
474-
DeclareOpInterfaceMethods<RegionBranchOpInterface, ["getEntrySuccessorOperands"]>]> {
475-
let summary = "For loop with constant bounds";
476-
477-
let description = [{
478-
A loop iterating over the integer range starting at `rangeBegin`
479-
(inclusive) and ending at `rangeEnd` (exclusive).
480-
The `body` region is executed once for every value in the integer range
481-
(that value is passed as the first block argument).
482-
At the first iteration of the loop, the other block arguments take the
483-
values of `initArgs`. For subsequent iterations, results from the
484-
previous iteration (produced by `YieldOp`) are taken instead.
485-
The `until` region, if present, is executed after `body`, and produces a
486-
single boolean scalar indicating whether the loop should terminate
487-
early.
488-
489-
In more imperative terms, `initArgs` can be seen as the set of variables
490-
that are updated in the loop body.
491-
Within the loop body, those variables can be accessed through the block
492-
arguments, and their updated values are set through `YieldOp`.
493-
Finally, `results` represents the new state of those variables after the
494-
loop terminates.
495-
}];
496-
497-
let arguments = (ins
498-
Variadic<Matrix>:$initArgs,
499-
I64Scalar:$rangeBegin,
500-
I64Scalar:$rangeEnd);
501-
502-
let results = (outs Variadic<Matrix>:$results);
503-
504-
let regions = (region SizedRegion<1>:$body, MaxSizedRegion<1>:$until);
505-
506-
let assemblyFormat = [{
507-
`range` `(`
508-
$rangeBegin `,`
509-
$rangeEnd
510-
`)` `:` type($rangeEnd)
511-
`init` `(` $initArgs `)` `:` type($initArgs) `->` type($results) attr-dict
512-
`body` $body
513-
`until` $until
514-
}];
515-
516-
let hasRegionVerifier = 1;
517-
}
518-
519-
def ForDimOp : Core_Op<"for_dim", [
520-
Pure,
521-
DeclareOpInterfaceMethods<RegionBranchOpInterface, ["getEntrySuccessorOperands"]>]> {
522-
let summary = "For loop over a matrix dimension";
523-
524-
let description = [{
525-
A loop iterating over the half-open range [0..`dim`).
526-
527-
This op is otherwise equivalent to `ForConstOp`.
528-
}];
529-
530-
let arguments = (ins Variadic<Matrix>:$initArgs, DimAttr:$dim);
531-
532-
let results = (outs Variadic<Matrix>:$results);
533-
534-
let regions = (region SizedRegion<1>:$body, MaxSizedRegion<1>:$until);
535-
536-
let assemblyFormat = [{
537-
`range` `(` custom<BareAttr>($dim) `)`
538-
`init` `(` $initArgs `)` `:` type($initArgs) `->` type($results) attr-dict
539-
`body` $body
540-
`until` $until
541-
}];
542-
543-
let hasRegionVerifier = 1;
544-
let hasCanonicalizer = 1;
545-
}
546-
547470
def YieldOp : Core_Op<"yield", [
548471
Pure,
549472
Terminator,
550-
ParentOneOf<["ForOp", "ForConstOp", "ForDimOp"]>,
473+
HasParent<"ForOp">,
551474
DeclareOpInterfaceMethods<RegionBranchTerminatorOpInterface>]> {
552475
let summary = "Yield from a loop body";
553476

compiler/src/garel/GraphAlgToRel.cpp

Lines changed: 26 additions & 67 deletions
Original file line numberDiff line numberDiff line change
@@ -320,6 +320,11 @@ static mlir::Value createDimRead(mlir::Location loc, graphalg::DimAttr dim,
320320
*/
321321
static 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-
991951
template <>
992952
mlir::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

compiler/src/graphalg/GraphAlgCanonicalize.cpp

Lines changed: 2 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -272,38 +272,9 @@ ForOp::fold(FoldAdaptor adaptor,
272272
return mlir::success();
273273
}
274274

275-
return mlir::failure();
276-
}
277-
278-
static mlir::LogicalResult forDimConst(ForDimOp op,
279-
mlir::PatternRewriter &rewriter) {
280-
if (!op.getDim().isConcrete()) {
281-
return mlir::failure();
282-
}
283-
284-
// The number of iterations is known, so we can replace with a ForConstOp.
285-
auto end = op.getDim().getConcreteDim();
286-
287-
// Range from 0 to dim.
288-
auto intType =
289-
MatrixType::scalarOf(SemiringTypes::forInt(rewriter.getContext()));
290-
auto beginOp = rewriter.create<ConstantMatrixOp>(
291-
op->getLoc(), intType, rewriter.getI64IntegerAttr(0));
292-
auto endOp = rewriter.create<ConstantMatrixOp>(
293-
op->getLoc(), intType, rewriter.getI64IntegerAttr(end));
275+
// TODO: Fold if iters=0
294276

295-
auto forConstOp = rewriter.create<ForConstOp>(
296-
op->getLoc(), op->getResultTypes(), op.getInitArgs(), beginOp, endOp);
297-
rewriter.inlineRegionBefore(op.getBody(), forConstOp.getBody(),
298-
forConstOp.getBody().begin());
299-
rewriter.replaceOp(op, forConstOp);
300-
301-
return mlir::success();
302-
}
303-
304-
void ForDimOp::getCanonicalizationPatterns(mlir::RewritePatternSet &patterns,
305-
mlir::MLIRContext *context) {
306-
patterns.add(forDimConst);
277+
return mlir::failure();
307278
}
308279

309280
mlir::OpFoldResult PickAnyOp::fold(FoldAdaptor adaptor) {

compiler/src/graphalg/GraphAlgLoopAggregate.cpp

Lines changed: 7 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -28,19 +28,10 @@ class GraphAlgLoopAggregate
2828
// If the loop body ends with an aggregation op, ensure the init arg is also an
2929
// aggregation. Later on in the AvantGraph query pipeline, this will signal to
3030
// the optimizer that the iteration state can be kept as an aggregate table.
31-
static void addInitReduce(mlir::Operation *op, mlir::IRRewriter &rewriter) {
32-
mlir::Block *body;
33-
llvm::SmallVector<mlir::Value> newInitArgs;
34-
if (auto constOp = llvm::dyn_cast<ForConstOp>(op)) {
35-
body = &constOp.getBody().front();
36-
newInitArgs = constOp.getInitArgs();
37-
} else {
38-
auto dimOp = llvm::cast<ForDimOp>(op);
39-
body = &dimOp.getBody().front();
40-
newInitArgs = dimOp.getInitArgs();
41-
}
31+
static void addInitReduce(ForOp op, mlir::IRRewriter &rewriter) {
32+
llvm::SmallVector<mlir::Value> newInitArgs(op.getInitArgs());
4233

43-
auto yieldOp = llvm::cast<YieldOp>(body->getTerminator());
34+
auto yieldOp = llvm::cast<YieldOp>(op.getBody().front().getTerminator());
4435
for (auto [i, iterResult] : llvm::enumerate(yieldOp.getInputs())) {
4536
auto iterLastOp = iterResult.getDefiningOp();
4637
if (llvm::isa_and_present<PickAnyOp, DeferredReduceOp>(iterLastOp)) {
@@ -50,20 +41,13 @@ static void addInitReduce(mlir::Operation *op, mlir::IRRewriter &rewriter) {
5041
}
5142
}
5243

53-
rewriter.modifyOpInPlace(op, [&]() {
54-
if (auto constOp = llvm::dyn_cast<ForConstOp>(op)) {
55-
constOp.getInitArgsMutable().assign(newInitArgs);
56-
} else {
57-
auto dimOp = llvm::cast<ForDimOp>(op);
58-
dimOp.getInitArgsMutable().assign(newInitArgs);
59-
}
60-
});
44+
rewriter.modifyOpInPlace(
45+
op, [&]() { op.getInitArgsMutable().assign(newInitArgs); });
6146
}
6247

6348
void GraphAlgLoopAggregate::runOnOperation() {
64-
llvm::SmallVector<mlir::Operation *> loopOps;
65-
getOperation()->walk([&](ForConstOp op) { loopOps.emplace_back(op); });
66-
getOperation()->walk([&](ForDimOp op) { loopOps.emplace_back(op); });
49+
llvm::SmallVector<ForOp> loopOps;
50+
getOperation()->walk([&](ForOp op) { loopOps.emplace_back(op); });
6751

6852
mlir::IRRewriter rewriter(&getContext());
6953
for (auto op : loopOps) {

0 commit comments

Comments
 (0)