LLZK 3.0.0
An open-source IR for Zero Knowledge (ZK) circuits
Loading...
Searching...
No Matches
LowerBoolQuantifiersPass.cpp
Go to the documentation of this file.
1//===-- LowerBoolQuantifiersPass.cpp ---------------------------*- C++ -*-===//
2//
3// Part of the LLZK Project, under the Apache License v2.0.
4// See LICENSE.txt for license information.
5// Copyright 2026 Project LLZK
6// SPDX-License-Identifier: Apache-2.0
7//
8//===----------------------------------------------------------------------===//
13//===----------------------------------------------------------------------===//
14
20
21#include <mlir/Dialect/Arith/IR/Arith.h>
22#include <mlir/Dialect/SCF/IR/SCF.h>
23#include <mlir/IR/IRMapping.h>
24#include <mlir/Transforms/GreedyPatternRewriteDriver.h>
26// Include the generated base pass class definitions.
27namespace llzk::boolean {
28#define GEN_PASS_DEF_LOWERBOOLQUANTIFIERSPASS
30} // namespace llzk::boolean
32using namespace mlir;
33using namespace llzk;
34using namespace llzk::array;
35using namespace llzk::boolean;
36
37namespace {
38
40static Value buildQuantifierIterValue(
41 Location loc, Value sort, ArrayType sortType, Value index, PatternRewriter &rewriter
42) {
43 if (sortType.getDimensionSizes().size() == 1) {
44 return rewriter.create<ReadArrayOp>(loc, sort, ValueRange {index});
45 }
46
47 Type iterType = getQuantifierOpDomainIterType(sortType);
48 return rewriter.create<ExtractArrayOp>(loc, iterType, sort, ValueRange {index});
49}
50
51/// Lower a bool quantifier to an `scf.for` loop over the first dimension of its array sort.
52template <typename QuantifierOp, typename CombineOp>
53static LogicalResult
54lowerQuantifier(QuantifierOp op, PatternRewriter &rewriter, bool initialValue) {
55 PatternRewriter::InsertionGuard guard(rewriter);
56 Location loc = op.getLoc();
57 auto sortType = cast<ArrayType>(op.getSort().getType());
58
59 Value lowerBound = rewriter.create<arith::ConstantIndexOp>(loc, 0);
60 Value upperBound = rewriter.create<ArrayLengthOp>(loc, op.getSort(), lowerBound);
61 Value step = rewriter.create<arith::ConstantIndexOp>(loc, 1);
62 Value init = rewriter.create<arith::ConstantIntOp>(loc, initialValue, rewriter.getI1Type());
63
64 auto loop = rewriter.create<scf::ForOp>(loc, lowerBound, upperBound, step, ValueRange {init});
65 loop->setDiscardableAttrs(op->getDiscardableAttrDictionary());
66
67 Block &loopBody = *loop.getBody();
68 if (!loopBody.empty()) {
69 rewriter.eraseOp(&loopBody.back());
70 }
71
72 rewriter.setInsertionPointToStart(&loopBody);
73 Value iterValue =
74 buildQuantifierIterValue(loc, op.getSort(), sortType, loop.getInductionVar(), rewriter);
75
76 IRMapping mapping;
77 mapping.map(op.getBody()->getArgument(0), iterValue);
78 for (Operation &nestedOp : op.getBody()->without_terminator()) {
79 rewriter.clone(nestedOp, mapping);
80 }
81
82 auto yieldOp = cast<YieldOp>(op.getBody()->getTerminator());
83 Value predicate = mapping.lookupOrDefault(yieldOp.getValue());
84 Value combined = rewriter.create<CombineOp>(loc, loop.getRegionIterArg(0), predicate);
85 rewriter.create<scf::YieldOp>(loc, combined);
86
87 rewriter.replaceOp(op, loop.getResults());
88 return success();
89}
90
91class LowerForAllOp : public OpRewritePattern<ForAllOp> {
92public:
93 using OpRewritePattern<ForAllOp>::OpRewritePattern;
94
95 LogicalResult matchAndRewrite(ForAllOp op, PatternRewriter &rewriter) const override {
96 return lowerQuantifier<ForAllOp, AndBoolOp>(op, rewriter, /*initialValue=*/true);
97 }
98};
99
100class LowerExistsOp : public OpRewritePattern<ExistsOp> {
101public:
102 using OpRewritePattern<ExistsOp>::OpRewritePattern;
103
104 LogicalResult matchAndRewrite(ExistsOp op, PatternRewriter &rewriter) const override {
105 return lowerQuantifier<ExistsOp, OrBoolOp>(op, rewriter, /*initialValue=*/false);
106 }
107};
108
109class PassImpl : public llzk::boolean::impl::LowerBoolQuantifiersPassBase<PassImpl> {
110 using Base = LowerBoolQuantifiersPassBase<PassImpl>;
111
112public:
113 using Base::Base;
114
115 void runOnOperation() override {
116 RewritePatternSet patterns(&getContext());
117 patterns.add<LowerForAllOp, LowerExistsOp>(&getContext());
118
119 if (failed(applyPatternsGreedily(getOperation(), std::move(patterns)))) {
120 signalPassFailure();
121 }
122 }
123};
124
125} // namespace
::llvm::ArrayRef<::mlir::Attribute > getDimensionSizes() const
mlir::Type getQuantifierOpDomainIterType(llzk::array::ArrayType arr)
Extracts the type used for a quantifier op block argument.
Definition Utils.h:20