LLZK
3.0.0
An open-source IR for Zero Knowledge (ZK) circuits
Toggle main menu visibility
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
15
#include "
llzk/Dialect/Array/IR/Ops.h
"
16
#include "
llzk/Dialect/Array/IR/Types.h
"
17
#include "
llzk/Dialect/Bool/IR/Ops.h
"
18
#include "
llzk/Dialect/Bool/IR/Utils.h
"
19
#include "
llzk/Dialect/Bool/Transforms/TransformationPasses.h
"
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>
25
26
// Include the generated base pass class definitions.
27
namespace
llzk::boolean
{
28
#define GEN_PASS_DEF_LOWERBOOLQUANTIFIERSPASS
29
#include "
llzk/Dialect/Bool/Transforms/TransformationPasses.h.inc
"
30
}
// namespace llzk::boolean
31
32
using namespace
mlir
;
33
using namespace
llzk
;
34
using namespace
llzk::array
;
35
using namespace
llzk::boolean
;
36
37
namespace
{
38
40
static
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.
52
template
<
typename
QuantifierOp,
typename
CombineOp>
53
static
LogicalResult
54
lowerQuantifier(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
91
class
LowerForAllOp :
public
OpRewritePattern<ForAllOp> {
92
public
:
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
100
class
LowerExistsOp :
public
OpRewritePattern<ExistsOp> {
101
public
:
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
109
class
PassImpl :
public
llzk::boolean::impl::LowerBoolQuantifiersPassBase
<PassImpl> {
110
using
Base = LowerBoolQuantifiersPassBase<PassImpl>;
111
112
public
:
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
Ops.h
Types.h
Ops.h
TransformationPasses.h.inc
TransformationPasses.h
Utils.h
llzk::array::ArrayLengthOp
Definition
Ops.h.inc:120
llzk::array::ArrayType
Definition
Types.h.inc:23
llzk::array::ArrayType::getDimensionSizes
::llvm::ArrayRef<::mlir::Attribute > getDimensionSizes() const
Definition
Types.cpp.inc:203
llzk::array::ExtractArrayOp
Definition
Ops.h.inc:566
llzk::array::ReadArrayOp
Definition
Ops.h.inc:876
llzk::boolean::ExistsOp
Definition
Ops.h.inc:139
llzk::boolean::ForAllOp
Definition
Ops.h.inc:291
llzk::boolean::impl::LowerBoolQuantifiersPassBase
Definition
LowerBoolQuantifiersPass.cpp:25
llzk::array
Definition
Ops.cpp:43
llzk::boolean
Definition
Ops.cpp:29
llzk::boolean::getQuantifierOpDomainIterType
mlir::Type getQuantifierOpDomainIterType(llzk::array::ArrayType arr)
Extracts the type used for a quantifier op block argument.
Definition
Utils.h:20
llzk::cast
Definition
Ops.cpp:56
llzk
Definition
AnalysisPassEnums.cpp:19
mlir
Definition
ValueModel.h:30
lib
Dialect
Bool
Transforms
LowerBoolQuantifiersPass.cpp
Generated by
1.17.0
Copyright 2025 Veridise Inc. under the Apache License v2.0. Copyright 2026 Project LLZK under the Apache License v2.0.