LLZK 3.0.0
An open-source IR for Zero Knowledge (ZK) circuits
Loading...
Searching...
No Matches
Ops.cpp
Go to the documentation of this file.
1//===-- Ops.cpp - Boolean operation implementations ----------------*- C++ -*-===//
2//
3// Part of the LLZK Project, under the Apache License v2.0.
4// See LICENSE.txt for license information.
5// Copyright 2025 Veridise Inc.
6// SPDX-License-Identifier: Apache-2.0
7//
8//===----------------------------------------------------------------------===//
9
11
16
17#include <mlir/IR/BuiltinAttributes.h>
18#include <mlir/IR/OpImplementation.h>
19#include <mlir/Support/LLVM.h>
20
21#include <cassert>
22
23// TableGen'd implementation files
24#define GET_OP_CLASSES
26
27using namespace mlir;
28
29namespace llzk::boolean {
30
31//===------------------------------------------------------------------===//
32// AssertOp
33//===------------------------------------------------------------------===//
34
35// This side effect models "program termination". Based on
36// https://github.com/llvm/llvm-project/blob/f325e4b2d836d6e65a4d0cf3efc6b0996ccf3765/mlir/lib/Dialect/ControlFlow/IR/ControlFlowOps.cpp#L92-L97
38 SmallVectorImpl<SideEffects::EffectInstance<MemoryEffects::Effect>> &effects
39) {
40 effects.emplace_back(MemoryEffects::Write::get());
41}
42
43//===------------------------------------------------------------------===//
44// Fold helpers
45//===------------------------------------------------------------------===//
46
47namespace {
48
51static FailureOr<bool> getBoolValue(Attribute attr) {
52 auto ia = llvm::dyn_cast_or_null<IntegerAttr>(attr);
53 if (!ia || !ia.getType().isInteger(1)) {
54 return failure();
55 }
56 return ia.getValue().getBoolValue();
57}
58
60static IntegerAttr makeBoolAttr(MLIRContext *ctx, bool val) {
61 auto i1Ty = IntegerType::get(ctx, 1);
62 return IntegerAttr::get(i1Ty, val ? 1 : 0);
63}
64
65} // namespace
66
67//===------------------------------------------------------------------===//
68// AndBoolOp
69//===------------------------------------------------------------------===//
70
71OpFoldResult AndBoolOp::fold(FoldAdaptor adaptor) {
72 auto lhs = getBoolValue(adaptor.getLhs());
73 auto rhs = getBoolValue(adaptor.getRhs());
74 if (failed(lhs) || failed(rhs)) {
75 return {};
76 }
77 return makeBoolAttr(getContext(), *lhs && *rhs);
78}
79
80//===------------------------------------------------------------------===//
81// OrBoolOp
82//===------------------------------------------------------------------===//
83
84OpFoldResult OrBoolOp::fold(FoldAdaptor adaptor) {
85 auto lhs = getBoolValue(adaptor.getLhs());
86 auto rhs = getBoolValue(adaptor.getRhs());
87 if (failed(lhs) || failed(rhs)) {
88 return {};
89 }
90 return makeBoolAttr(getContext(), *lhs || *rhs);
91}
92
93//===------------------------------------------------------------------===//
94// XorBoolOp
95//===------------------------------------------------------------------===//
96
97OpFoldResult XorBoolOp::fold(FoldAdaptor adaptor) {
98 auto lhs = getBoolValue(adaptor.getLhs());
99 auto rhs = getBoolValue(adaptor.getRhs());
100 if (failed(lhs) || failed(rhs)) {
101 return {};
102 }
103 return makeBoolAttr(getContext(), *lhs != *rhs);
104}
105
106//===------------------------------------------------------------------===//
107// NotBoolOp
108//===------------------------------------------------------------------===//
109
110OpFoldResult NotBoolOp::fold(FoldAdaptor adaptor) {
111 auto val = getBoolValue(adaptor.getOperand());
112 if (failed(val)) {
113 return {};
114 }
115 return makeBoolAttr(getContext(), !*val);
116}
117
118//===------------------------------------------------------------------===//
119// CmpOp
120//===------------------------------------------------------------------===//
121
122inline static bool eval(FeltCmpPredicate pred, const llvm::APInt &lval, const llvm::APInt &rval) {
123 switch (pred) {
125 return lval == rval;
127 return lval != rval;
129 return lval.ult(rval);
131 return lval.ule(rval);
133 return lval.ugt(rval);
135 return lval.uge(rval);
136 }
137 llvm_unreachable("invalid FeltCmpPredicate");
138}
139
140OpFoldResult CmpOp::fold(FoldAdaptor adaptor) {
141 auto lhsAttr = llvm::dyn_cast_or_null<felt::FeltConstAttr>(adaptor.getLhs());
142 auto rhsAttr = llvm::dyn_cast_or_null<felt::FeltConstAttr>(adaptor.getRhs());
143 if (!lhsAttr || !rhsAttr) {
144 return {};
145 }
146
147 // Normalize to a common bit width for unsigned comparison.
148 llvm::APInt lval = lhsAttr.getValue();
149 llvm::APInt rval = rhsAttr.getValue();
150 unsigned w = std::max(lval.getBitWidth(), rval.getBitWidth());
151 if (lval.getBitWidth() < w) {
152 lval = lval.zext(w);
153 }
154 if (rval.getBitWidth() < w) {
155 rval = rval.zext(w);
156 }
157 return makeBoolAttr(getContext(), eval(getPredicate(), lval, rval));
158}
159
160//===------------------------------------------------------------------===//
161// Quantifier ops common impl
162//===------------------------------------------------------------------===//
163
164namespace {
169template <typename Op> LogicalResult verifyQuantOp(Op op) {
170 auto *block = op.getBody();
171 if (!block || block->getNumArguments() != 1) {
172 return op->emitOpError() << "must have one block argument";
173 }
174 auto argType = block->getArgument(0).getType();
175 auto eltType = getQuantifierOpDomainIterType(op.getSort().getType());
176 if (argType != eltType) {
177 return op->emitOpError() << "expects element type " << argType << " but sort has element type "
178 << eltType;
179 }
180
181 auto termOp = block->getTerminator();
182 if (!llvm::dyn_cast_if_present<YieldOp>(termOp)) {
183 return op->emitOpError() << "expects 'bool.yield' terminator op";
184 }
185 return success();
186}
187
196ParseResult parseQuantOp(OpAsmParser &parser, OperationState &result) {
197 OpAsmParser::Argument arg;
198 if (parser.parseArgument(arg)) {
199 return failure();
200 }
201 assert(!arg.type);
202 if (succeeded(parser.parseOptionalColon())) {
203 if (parser.parseType(arg.type)) {
204 return failure();
205 }
206 }
207 if (parser.parseKeyword("in")) {
208 return failure();
209 }
210 OpAsmParser::UnresolvedOperand sortOperand;
211 array::ArrayType sortType;
212 if (parser.parseOperand(sortOperand)) {
213 return failure();
214 }
215 if (parser.parseColonType(sortType)) {
216 return failure();
217 }
218 if (parser.resolveOperand(sortOperand, sortType, result.operands)) {
219 return failure();
220 }
221
222 if (!arg.type) {
223 arg.type = getQuantifierOpDomainIterType(sortType);
224 assert(arg.type && "argument type must be inferred from the array element type");
225 }
226
227 auto *body = result.addRegion();
228 SMLoc loc = parser.getCurrentLocation();
229 if (parser.parseRegion(
230 *body, {arg},
231 /*enableNameShadowing=*/false
232 )) {
233 return failure();
234 }
235
236 if (body->empty()) {
237 return parser.emitError(loc, "expected non-empty body");
238 }
239 if (parser.parseOptionalAttrDictWithKeyword(result.attributes)) {
240 return failure();
241 }
242
243 result.types = {parser.getBuilder().getI1Type()};
244 return success();
245}
246
248template <typename Op> void printQuantOp(OpAsmPrinter &p, Op op) {
249 p << ' ';
250 p.printRegionArgument(op.getBody()->getArgument(0));
251 p << " in ";
252 p.printOperand(op.getSort());
253 p << " : " << op.getSort().getType();
254 p << ' ';
255 p.printRegion(op.getRegion(), /*printEntryBlockArgs=*/false);
256 p << ' ';
257 p.printOptionalAttrDictWithKeyword(op->getAttrs());
258}
259} // namespace
260
261//===------------------------------------------------------------------===//
262// ForAllOp
263//===------------------------------------------------------------------===//
264
265LogicalResult ForAllOp::verify() { return verifyQuantOp(*this); }
266
267ParseResult ForAllOp::parse(OpAsmParser &parser, OperationState &result) {
268 return parseQuantOp(parser, result);
269}
270
271void ForAllOp::print(OpAsmPrinter &p) { printQuantOp(p, *this); }
272
273//===------------------------------------------------------------------===//
274// ExistsOp
275//===------------------------------------------------------------------===//
276
277LogicalResult ExistsOp::verify() { return verifyQuantOp(*this); }
278
279ParseResult ExistsOp::parse(OpAsmParser &parser, OperationState &result) {
280 return parseQuantOp(parser, result);
281}
282
283void ExistsOp::print(OpAsmPrinter &p) { printQuantOp(p, *this); }
284
285} // namespace llzk::boolean
::mlir::OpFoldResult fold(FoldAdaptor adaptor)
Definition Ops.cpp:71
GenericAdaptor<::llvm::ArrayRef<::mlir::Attribute > > FoldAdaptor
Definition Ops.h.inc:446
void getEffects(::llvm::SmallVectorImpl<::mlir::SideEffects::EffectInstance<::mlir::MemoryEffects::Effect > > &effects)
Definition Ops.cpp:37
GenericAdaptor<::llvm::ArrayRef<::mlir::Attribute > > FoldAdaptor
Definition Ops.h.inc:870
::llzk::boolean::FeltCmpPredicate getPredicate()
Definition Ops.cpp.inc:873
::mlir::OpFoldResult fold(FoldAdaptor adaptor)
Definition Ops.cpp:140
::llvm::LogicalResult verify()
Definition Ops.cpp:277
void print(::mlir::OpAsmPrinter &p)
::mlir::ParseResult parse(::mlir::OpAsmParser &parser, ::mlir::OperationState &result)
::mlir::ParseResult parse(::mlir::OpAsmParser &parser, ::mlir::OperationState &result)
Definition Ops.cpp:267
::llvm::LogicalResult verify()
Definition Ops.cpp:265
void print(::mlir::OpAsmPrinter &p)
Definition Ops.cpp:271
GenericAdaptor<::llvm::ArrayRef<::mlir::Attribute > > FoldAdaptor
Definition Ops.h.inc:1089
::mlir::OpFoldResult fold(FoldAdaptor adaptor)
Definition Ops.cpp:110
GenericAdaptor<::llvm::ArrayRef<::mlir::Attribute > > FoldAdaptor
Definition Ops.h.inc:1258
::mlir::OpFoldResult fold(FoldAdaptor adaptor)
Definition Ops.cpp:84
GenericAdaptor<::llvm::ArrayRef<::mlir::Attribute > > FoldAdaptor
Definition Ops.h.inc:1436
::mlir::OpFoldResult fold(FoldAdaptor adaptor)
Definition Ops.cpp:97
mlir::Type getQuantifierOpDomainIterType(llzk::array::ArrayType arr)
Extracts the type used for a quantifier op block argument.
Definition Utils.h:20