23#include <mlir/IR/BuiltinOps.h>
24#include <mlir/IR/Dominance.h>
26#include <llvm/ADT/DenseMap.h>
27#include <llvm/ADT/DenseMapInfo.h>
28#include <llvm/ADT/DenseSet.h>
29#include <llvm/ADT/STLExtras.h>
30#include <llvm/ADT/ScopeExit.h>
31#include <llvm/ADT/SmallVector.h>
32#include <llvm/ADT/StringMap.h>
33#include <llvm/Support/Debug.h>
34#include <llvm/Support/raw_ostream.h>
43#define GEN_PASS_DEF_POLYLOWERINGPASS
55#define DEBUG_TYPE "llzk-poly-lowering-pass"
56#define AUXILIARY_MEMBER_PREFIX "__llzk_poly_lowering_pass_aux_member_"
61 std::string auxMemberName;
67struct MutableContainmentElement {
72enum class AuxAssignmentVisitState : uint8_t {
78class DegreeComputationError :
public std::runtime_error {
80 DegreeComputationError(Location errorLoc,
const std::string &message)
81 : std::runtime_error(message), loc(errorLoc) {}
83 Location getLoc()
const {
return loc; }
90 using Base = PolyLoweringPassBase<PassImpl>;
93 unsigned auxCounter = 0;
95 static void collectStructDefs(ModuleOp modOp, SmallVectorImpl<StructDefOp> &structDefs) {
97 structDefs.push_back(structDef);
98 return WalkResult::skip();
103 static void addAuxDependency(
104 unsigned dep,
unsigned owner, DenseSet<unsigned> &seenDeps, SmallVectorImpl<unsigned> &deps
109 if (seenDeps.insert(dep).second) {
115 static void collectAuxDependencies(
116 Value val,
unsigned owner,
const DenseMap<Value, unsigned> &auxValueToIndex,
117 const llvm::StringMap<unsigned> &auxNameToIndex, DenseSet<Value> &visitedValues,
118 DenseSet<unsigned> &seenDeps, SmallVectorImpl<unsigned> &deps
122 if (!val || !visitedValues.insert(val).second) {
126 if (
auto it = auxValueToIndex.find(val); it != auxValueToIndex.end()) {
127 addAuxDependency(it->second, owner, seenDeps, deps);
130 if (Operation *defOp = val.getDefiningOp()) {
131 if (
auto readOp = llvm::dyn_cast<MemberReadOp>(defOp)) {
132 auto it = auxNameToIndex.find(readOp.getMemberName());
133 if (it != auxNameToIndex.end()) {
134 addAuxDependency(it->second, owner, seenDeps, deps);
138 for (Value operand : defOp->getOperands()) {
139 collectAuxDependencies(
140 operand, owner, auxValueToIndex, auxNameToIndex, visitedValues, seenDeps, deps
147 static LogicalResult visitAuxAssignment(
148 unsigned idx, ArrayRef<SmallVector<unsigned>> deps,
149 SmallVectorImpl<AuxAssignmentVisitState> &visitState, SmallVectorImpl<unsigned> &ordered,
150 ArrayRef<AuxAssignment> auxAssignments
152 if (visitState[idx] == AuxAssignmentVisitState::Done) {
155 if (visitState[idx] == AuxAssignmentVisitState::Visiting) {
156 return emitError(auxAssignments[idx].computedValue.getLoc())
157 <<
"poly lowering generated cyclic auxiliary dependency involving @"
158 << auxAssignments[idx].auxMemberName;
161 visitState[idx] = AuxAssignmentVisitState::Visiting;
163 for (
unsigned dep : deps[idx]) {
164 if (failed(visitAuxAssignment(dep, deps, visitState, ordered, auxAssignments))) {
168 visitState[idx] = AuxAssignmentVisitState::Done;
169 ordered.push_back(idx);
175 orderAuxAssignments(ArrayRef<AuxAssignment> auxAssignments, SmallVectorImpl<unsigned> &ordered) {
176 DenseMap<Value, unsigned> auxValueToIndex;
177 llvm::StringMap<unsigned> auxNameToIndex;
178 auxValueToIndex.reserve(auxAssignments.size());
179 for (
auto [idx, assign] : llvm::enumerate(auxAssignments)) {
180 if (assign.auxValue) {
181 auxValueToIndex[assign.auxValue] = idx;
183 auxNameToIndex[assign.auxMemberName] = idx;
186 SmallVector<SmallVector<unsigned>> deps(auxAssignments.size());
187 for (
auto [idx, assign] : llvm::enumerate(auxAssignments)) {
188 DenseSet<Value> visitedValues;
189 DenseSet<unsigned> seenDeps;
190 collectAuxDependencies(
191 assign.computedValue, idx, auxValueToIndex, auxNameToIndex, visitedValues, seenDeps,
196 SmallVector<AuxAssignmentVisitState> visitState(
197 auxAssignments.size(), AuxAssignmentVisitState::Unvisited
199 for (
unsigned idx = 0, e = auxAssignments.size(); idx < e; ++idx) {
200 if (failed(visitAuxAssignment(idx, deps, visitState, ordered, auxAssignments))) {
208 unsigned getDegree(Value val, DenseMap<Value, unsigned> &memo) {
209 if (
auto it = memo.find(val); it != memo.end()) {
213 if (llvm::isa<BlockArgument>(val)) {
214 return memo[val] = 1;
216 if (Operation *defOp = val.getDefiningOp()) {
217 if (llvm::isa<FeltConstantOp>(defOp)) {
218 return memo[val] = 0;
220 if (llvm::isa<NonDetOp, MemberReadOp>(defOp)) {
221 return memo[val] = 1;
223 if (
auto add = llvm::dyn_cast<AddFeltOp>(defOp)) {
224 return memo[val] = std::max(getDegree(
add.getLhs(), memo), getDegree(
add.getRhs(), memo));
226 if (
auto sub = llvm::dyn_cast<SubFeltOp>(defOp)) {
227 return memo[val] = std::max(getDegree(
sub.getLhs(), memo), getDegree(
sub.getRhs(), memo));
229 if (
auto mul = llvm::dyn_cast<MulFeltOp>(defOp)) {
230 return memo[val] = getDegree(
mul.getLhs(), memo) + getDegree(
mul.getRhs(), memo);
232 if (
auto div = llvm::dyn_cast<DivFeltOp>(defOp)) {
233 return memo[val] = getDegree(
div.getLhs(), memo) + getDegree(
div.getRhs(), memo);
235 if (
auto neg = llvm::dyn_cast<NegFeltOp>(defOp)) {
236 return memo[val] = getDegree(
neg.getOperand(), memo);
238 if (
auto call = llvm::dyn_cast<CallOp>(defOp)) {
240 llvm::raw_string_ostream(message)
242 <<
"' in degree computation. Try running '-llzk-inline-free-functions' first.";
243 throw DegreeComputationError(val.getLoc(), message);
248 llvm::raw_string_ostream(message) <<
"Unhandled value in degree computation: " << val;
249 throw DegreeComputationError(val.getLoc(), message);
252 Value lowerExpression(
254 DominanceInfo &dominanceInfo, DenseMap<Value, unsigned> °reeMemo,
255 DenseMap<Value, Value> &rewrites, SmallVector<AuxAssignment> &auxAssignments
257 auto rewriteIt = rewrites.find(val);
258 if (rewriteIt != rewrites.end() && dominanceInfo.properlyDominates(rewriteIt->second, useOp)) {
259 return rewriteIt->second;
262 auto cacheIdentityRewriteIfAbsent = [&rewrites, &val]() {
264 if (!rewrites.contains(val)) {
269 unsigned degree = getDegree(val, degreeMemo);
270 if (degree <= maxDegree) {
274 cacheIdentityRewriteIfAbsent();
279 auto lowerBinaryRoot = [&](
auto op) -> Value {
280 Value lhs = lowerExpression(
281 op.getLhs(), structDef, constrainFunc, op.getOperation(), dominanceInfo, degreeMemo,
282 rewrites, auxAssignments
284 Value rhs = lowerExpression(
285 op.getRhs(), structDef, constrainFunc, op.getOperation(), dominanceInfo, degreeMemo,
286 rewrites, auxAssignments
289 if (lhs != op.getLhs()) {
290 op.getLhsMutable().set(lhs);
292 if (rhs != op.getRhs()) {
293 op.getRhsMutable().set(rhs);
295 degreeMemo[val] = std::max(getDegree(lhs, degreeMemo), getDegree(rhs, degreeMemo));
296 cacheIdentityRewriteIfAbsent();
300 Operation *defOp = val.getDefiningOp();
301 if (
auto addOp = llvm::dyn_cast_if_present<AddFeltOp>(defOp)) {
302 return lowerBinaryRoot(addOp);
303 }
else if (
auto subOp = llvm::dyn_cast_if_present<SubFeltOp>(defOp)) {
304 return lowerBinaryRoot(subOp);
305 }
else if (
auto negOp = llvm::dyn_cast_if_present<NegFeltOp>(defOp)) {
306 Value operand = lowerExpression(
307 negOp.getOperand(), structDef, constrainFunc, negOp.getOperation(), dominanceInfo,
308 degreeMemo, rewrites, auxAssignments
311 if (operand != negOp.getOperand()) {
312 negOp.getOperandMutable().set(operand);
314 degreeMemo[val] = getDegree(operand, degreeMemo);
315 cacheIdentityRewriteIfAbsent();
317 }
else if (
auto mulOp = llvm::dyn_cast_if_present<MulFeltOp>(defOp)) {
319 Value lhs = lowerExpression(
320 mulOp.getLhs(), structDef, constrainFunc, mulOp.getOperation(), dominanceInfo, degreeMemo,
321 rewrites, auxAssignments
323 Value rhs = lowerExpression(
324 mulOp.getRhs(), structDef, constrainFunc, mulOp.getOperation(), dominanceInfo, degreeMemo,
325 rewrites, auxAssignments
328 unsigned lhsDeg = getDegree(lhs, degreeMemo);
329 unsigned rhsDeg = getDegree(rhs, degreeMemo);
331 OpBuilder builder(mulOp.getOperation()->getBlock(), ++Block::iterator(mulOp));
333 bool eraseMul = lhsDeg + rhsDeg > maxDegree;
335 if (lhs == rhs && eraseMul) {
340 lhs.getLoc(), lhs.getType(), selfVal, auxMember.getNameAttr()
342 auxAssignments.push_back({auxName, lhs, auxVal});
343 Location loc = builder.getFusedLoc({auxVal.getLoc(), lhs.getLoc()});
347 degreeMemo[auxVal] = 1;
348 rewrites[lhs] = auxVal;
349 rewrites[rhs] = auxVal;
360 while (lhsDeg + rhsDeg > maxDegree) {
361 Value &toFactor = (lhsDeg >= rhsDeg) ? lhs : rhs;
369 toFactor.getLoc(), toFactor.getType(), selfVal, auxMember.getNameAttr()
373 Location loc = builder.getFusedLoc({auxVal.getLoc(), toFactor.getLoc()});
374 auto eqOp = builder.create<
EmitEqualityOp>(loc, auxVal, toFactor);
375 auxAssignments.push_back({auxName, toFactor, auxVal});
377 rewrites[toFactor] = auxVal;
378 degreeMemo[auxVal] = 1;
386 lhsDeg = getDegree(lhs, degreeMemo);
387 rhsDeg = getDegree(rhs, degreeMemo);
391 auto mulVal = builder.create<
MulFeltOp>(lhs.getLoc(), lhs.getType(), lhs, rhs);
393 mulOp->replaceAllUsesWith(mulVal);
398 degreeMemo[mulVal] = lhsDeg + rhsDeg;
399 rewrites[val] = mulVal;
405 cacheIdentityRewriteIfAbsent();
409 Value materializeCallArgument(
411 DominanceInfo &dominanceInfo, DenseMap<Value, unsigned> °reeMemo,
412 DenseMap<Value, Value> &rewrites, SmallVector<AuxAssignment> &auxAssignments
414 Value loweredVal = lowerExpression(
415 val, structDef, constrainFunc, callOp.getOperation(), dominanceInfo, degreeMemo, rewrites,
418 DenseMap<Value, unsigned> checkMemo;
419 if (getDegree(loweredVal, checkMemo) <= 1) {
428 OpBuilder builder(callOp);
431 loweredVal.getLoc(), loweredVal.getType(), selfVal, auxMember.getNameAttr()
434 Location loc = builder.getFusedLoc({auxVal.getLoc(), loweredVal.getLoc()});
436 auxAssignments.push_back({auxName, loweredVal, auxVal});
438 degreeMemo[auxVal] = 1;
439 rewrites[loweredVal] = auxVal;
440 rewrites[val] = auxVal;
444 LogicalResult checkEqualityDegrees(
FuncDefOp constrainFunc) {
445 auto res = constrainFunc.walk([
this](
EmitEqualityOp eqOp) -> WalkResult {
447 Value rhs = eqOp.
getRhs();
448 if (llvm::isa<FeltType>(lhs.getType()) && llvm::isa<FeltType>(rhs.getType())) {
449 DenseMap<Value, unsigned> checkMemo;
450 unsigned lhsDegree = getDegree(lhs, checkMemo);
451 unsigned rhsDegree = getDegree(rhs, checkMemo);
454 return eqOp.emitOpError().append(
455 "poly lowering postcondition failed: equality operand degree exceeds max-degree ",
456 maxDegree.getValue(),
" (lhs degree ", lhsDegree,
", rhs degree ", rhsDegree,
')'
460 return WalkResult::advance();
462 return failure(res.wasInterrupted());
465 LogicalResult checkStructConstrainCallArguments(
FuncDefOp constrainFunc) {
466 auto res = constrainFunc.walk([
this](
CallOp callOp) -> WalkResult {
469 if (!llvm::isa<FeltType>(arg.getType())) {
472 DenseMap<Value, unsigned> checkMemo;
473 unsigned argDegree = getDegree(arg, checkMemo);
475 return callOp.emitOpError()
476 <<
"poly lowering postcondition failed: "
477 "struct constrain call argument degree exceeds 1 (argument degree "
482 return WalkResult::advance();
484 return failure(res.wasInterrupted());
488 static bool isFeltArray(Type type) {
489 if (
auto arrayType = llvm::dyn_cast<ArrayType>(type)) {
490 return llvm::isa<FeltType>(arrayType.getElementType());
496 return containOp.emitOpError()
497 <<
"poly lowering cannot resolve containment RHS row write history: " <<
detail;
501 template <
typename IndexRange,
typename PrefixRange>
502 static bool indexStartsWith(
const IndexRange &index,
const PrefixRange &prefix) {
503 auto indexIt = index.begin();
504 for (Attribute attr : prefix) {
505 if (indexIt == index.end() || *indexIt != attr) {
514 template <
typename LhsRange,
typename RhsRange>
515 static bool prefixesCanOverlap(
const LhsRange &lhs,
const RhsRange &rhs) {
516 auto lhsIt = lhs.begin();
517 auto rhsIt = rhs.begin();
518 while (lhsIt != lhs.end() && rhsIt != rhs.end()) {
519 if (*lhsIt != *rhsIt) {
529 static ArrayAttr dropIndexPrefix(MLIRContext *ctx, ArrayAttr index,
size_t prefixSize) {
530 SmallVector<Attribute> attrs;
532 for (Attribute attr : index) {
533 if (idx++ >= prefixSize) {
534 attrs.push_back(attr);
537 return ArrayAttr::get(ctx, attrs);
541 template <
typename PrefixRange>
542 static ArrayAttr appendIndex(MLIRContext *ctx,
const PrefixRange &prefix, ArrayAttr suffix) {
543 SmallVector<Attribute> attrs;
544 for (Attribute attr : prefix) {
545 attrs.push_back(attr);
547 for (Attribute attr : suffix) {
548 attrs.push_back(attr);
550 return ArrayAttr::get(ctx, attrs);
554 static inline ArrayAttr getStaticAccessIndex(Operation *op) {
555 return llvm::cast<ArrayAccessOpInterface>(op).indexOperandsToAttributeArray();
560 static std::optional<SmallVector<ArrayAttr>>
561 getViewIndices(ArrayType arrayType, ArrayRef<Attribute> viewPrefix) {
567 SmallVector<ArrayAttr> viewIndices;
568 MLIRContext *ctx = arrayType.getContext();
569 for (ArrayAttr index : *allIndices) {
570 if (indexStartsWith(index, viewPrefix)) {
571 viewIndices.push_back(dropIndexPrefix(ctx, index, viewPrefix.size()));
580 static LogicalResult collectMutableContainmentElements(
581 Value arrayValue, Operation *boundaryOp, ArrayRef<Attribute> viewPrefix,
582 EmitContainmentOp containOp, DenseSet<Value> &activeArrays,
583 SmallVectorImpl<MutableContainmentElement> &elements
585 auto arrayType = llvm::dyn_cast<ArrayType>(arrayValue.getType());
586 if (!arrayType || !llvm::isa<FeltType>(arrayType.
getElementType())) {
590 DenseMap<Attribute, OpOperand *> finalElements;
591 if (failed(collectMutableContainmentElementMap(
592 arrayValue, boundaryOp, viewPrefix, containOp, activeArrays, finalElements
597 std::optional<SmallVector<ArrayAttr>> viewIndices = getViewIndices(arrayType, viewPrefix);
599 if (finalElements.empty()) {
602 return emitAmbiguousContainmentRhs(containOp,
"array shape is not static");
605 for (ArrayAttr relativeIndex : *viewIndices) {
606 auto elementIt = finalElements.find(relativeIndex);
607 if (elementIt != finalElements.end()) {
608 elements.push_back(MutableContainmentElement {relativeIndex, elementIt->second});
617 static std::optional<std::pair<Value, FlatSymbolRefAttr>> resolveStructReadSource(Value v) {
618 if (
auto readOp = v.getDefiningOp<MemberReadOp>()) {
619 return std::make_pair(readOp.getComponent(), readOp.getMemberNameAttr());
629 static bool mayAliasArraySource(Value a, Value b) {
633 auto srcA = resolveStructReadSource(a);
637 auto srcB = resolveStructReadSource(b);
641 return srcA->first == srcB->first && srcA->second == srcB->second;
647 static LogicalResult collectMutableContainmentElementMap(
648 Value arrayValue, Operation *boundaryOp, ArrayRef<Attribute> viewPrefix,
649 EmitContainmentOp containOp, DenseSet<Value> &activeArrays,
650 DenseMap<Attribute, OpOperand *> &finalElements
654 auto arrayType = llvm::dyn_cast<ArrayType>(arrayValue.getType());
655 if (!arrayType || !llvm::isa<FeltType>(arrayType.
getElementType())) {
658 if (!boundaryOp || !boundaryOp->getBlock()) {
659 return emitAmbiguousContainmentRhs(containOp,
"missing observation block");
661 if (!activeArrays.insert(arrayValue).second) {
662 return emitAmbiguousContainmentRhs(containOp,
"cyclic array update");
664 auto cleanup = llvm::make_scope_exit([&]() { activeArrays.erase(arrayValue); });
666 MLIRContext *ctx = arrayType.getContext();
668 if (
auto arrayOp = arrayValue.getDefiningOp<CreateArrayOp>()) {
669 MutableOperandRange elementOperands = arrayOp.getElementsMutable();
670 if (!elementOperands.empty()) {
673 return emitAmbiguousContainmentRhs(containOp,
"array.new shape is not static");
675 assert(allIndices->size() == elementOperands.size() &&
"array.new verifier mismatch");
677 auto *indexIt = allIndices->begin();
678 for (OpOperand &elementOperand : elementOperands) {
679 ArrayAttr index = *indexIt++;
680 if (indexStartsWith(index, viewPrefix)) {
681 finalElements[dropIndexPrefix(ctx, index, viewPrefix.size())] = &elementOperand;
685 }
else if (
auto extractOp = arrayValue.getDefiningOp<ExtractArrayOp>()) {
686 ArrayAttr extractIndex = getStaticAccessIndex(extractOp.getOperation());
688 return emitAmbiguousContainmentRhs(containOp,
"array.extract index is not static");
691 SmallVector<Attribute> sourcePrefix;
692 for (Attribute attr : extractIndex) {
693 sourcePrefix.push_back(attr);
695 for (Attribute attr : viewPrefix) {
696 sourcePrefix.push_back(attr);
699 if (failed(collectMutableContainmentElementMap(
700 extractOp.getArrRef(), extractOp.getOperation(), sourcePrefix, containOp,
701 activeArrays, finalElements
707 for (Operation &op : *boundaryOp->getBlock()) {
708 if (&op == boundaryOp) {
712 if (
auto writeOp = llvm::dyn_cast<WriteArrayOp>(&op)) {
713 if (!mayAliasArraySource(writeOp.getArrRef(), arrayValue)) {
717 ArrayAttr writeIndex = getStaticAccessIndex(writeOp.getOperation());
719 return emitAmbiguousContainmentRhs(containOp,
"array.write index is not static");
721 if (indexStartsWith(writeIndex, viewPrefix)) {
722 finalElements[dropIndexPrefix(ctx, writeIndex, viewPrefix.size())] =
723 &writeOp.getRvalueMutable();
728 if (
auto insertOp = llvm::dyn_cast<InsertArrayOp>(&op)) {
729 if (!mayAliasArraySource(insertOp.getArrRef(), arrayValue)) {
733 ArrayAttr insertIndex = getStaticAccessIndex(insertOp.getOperation());
735 return emitAmbiguousContainmentRhs(containOp,
"array.insert index is not static");
737 if (!prefixesCanOverlap(insertIndex, viewPrefix)) {
741 auto rvalueType = llvm::dyn_cast<ArrayType>(insertOp.getRvalue().getType());
742 if (!rvalueType || !llvm::isa<FeltType>(rvalueType.getElementType())) {
746 std::optional<SmallVector<ArrayAttr>> rvalueIndices = rvalueType.getSubelementIndices();
747 if (!rvalueIndices) {
748 return emitAmbiguousContainmentRhs(containOp,
"array.insert rvalue shape is not static");
751 SmallVector<MutableContainmentElement> insertedElements;
752 if (failed(collectMutableContainmentElements(
753 insertOp.getRvalue(), insertOp.getOperation(), ArrayRef<Attribute> {}, containOp,
754 activeArrays, insertedElements
759 DenseMap<Attribute, OpOperand *> insertedElementMap;
760 for (MutableContainmentElement element : insertedElements) {
761 insertedElementMap[element.index] = element.operand;
764 for (ArrayAttr rvalueIndex : *rvalueIndices) {
765 ArrayAttr targetIndex = appendIndex(ctx, insertIndex, rvalueIndex);
766 if (!indexStartsWith(targetIndex, viewPrefix)) {
770 ArrayAttr relativeIndex = dropIndexPrefix(ctx, targetIndex, viewPrefix.size());
771 auto elementIt = insertedElementMap.find(rvalueIndex);
772 if (elementIt == insertedElementMap.end()) {
773 finalElements.erase(relativeIndex);
776 finalElements[relativeIndex] = elementIt->second;
786 LogicalResult lowerContainmentRhsFeltOperand(
787 OpOperand &operand, StructDefOp structDef, FuncDefOp constrainFunc,
788 DominanceInfo &dominanceInfo, DenseMap<Value, unsigned> °reeMemo,
789 DenseMap<Value, Value> &rewrites, SmallVector<AuxAssignment> &auxAssignments
791 Value value = operand.get();
792 if (!llvm::isa<FeltType>(value.getType())) {
796 unsigned degree = getDegree(value, degreeMemo);
797 if (degree > maxDegree) {
798 operand.set(lowerExpression(
799 value, structDef, constrainFunc, operand.getOwner(), dominanceInfo, degreeMemo, rewrites,
808 LogicalResult lowerContainmentRhsValue(
809 OpOperand &operand, StructDefOp structDef, FuncDefOp constrainFunc,
810 DominanceInfo &dominanceInfo, DenseMap<Value, unsigned> °reeMemo,
811 DenseMap<Value, Value> &rewrites, SmallVector<AuxAssignment> &auxAssignments,
812 EmitContainmentOp containOp
814 Value value = operand.get();
815 if (llvm::isa<FeltType>(value.getType())) {
816 return lowerContainmentRhsFeltOperand(
817 operand, structDef, constrainFunc, dominanceInfo, degreeMemo, rewrites, auxAssignments
821 if (!isFeltArray(value.getType())) {
825 DenseSet<Value> activeArrays;
826 SmallVector<MutableContainmentElement> elements;
827 if (failed(collectMutableContainmentElements(
828 value, containOp.getOperation(), ArrayRef<Attribute> {}, containOp, activeArrays,
834 for (MutableContainmentElement element : elements) {
835 if (failed(lowerContainmentRhsFeltOperand(
836 *element.operand, structDef, constrainFunc, dominanceInfo, degreeMemo, rewrites,
848 LogicalResult checkContainmentRhsFeltValue(
849 Value value, EmitContainmentOp containOp, DenseMap<Value, unsigned> &checkMemo
851 if (!llvm::isa<FeltType>(value.getType())) {
855 unsigned valueDegree = getDegree(value, checkMemo);
856 if (valueDegree <= maxDegree) {
860 return containOp.emitOpError()
861 <<
"poly lowering postcondition failed: "
862 "containment RHS element degree exceeds max-degree "
863 << maxDegree.getValue() <<
" (element degree " << valueDegree <<
')';
868 LogicalResult checkContainmentRhsValue(
869 Value value, EmitContainmentOp containOp, DenseMap<Value, unsigned> &checkMemo
871 if (llvm::isa<FeltType>(value.getType())) {
872 return checkContainmentRhsFeltValue(value, containOp, checkMemo);
875 if (!isFeltArray(value.getType())) {
879 DenseSet<Value> activeArrays;
880 SmallVector<MutableContainmentElement> elements;
881 auto res = collectMutableContainmentElements(
882 value, containOp.getOperation(), {}, containOp, activeArrays, elements
888 for (MutableContainmentElement element : elements) {
889 if (failed(checkContainmentRhsFeltValue(element.operand->get(), containOp, checkMemo))) {
898 LogicalResult checkContainmentRhsDegrees(FuncDefOp constrainFunc) {
899 auto res = constrainFunc.walk([&](EmitContainmentOp containOp) -> WalkResult {
900 DenseMap<Value, unsigned> memo;
901 return checkContainmentRhsValue(containOp.
getRhs(), containOp, memo);
903 return failure(res.wasInterrupted());
906 LogicalResult lowerInConstrain(
907 StructDefOp structDef, FuncDefOp constrainFunc, SmallVector<AuxAssignment> &auxAssignments
910 DenseMap<Value, unsigned> degreeMemo;
911 DenseMap<Value, Value> rewrites;
912 DominanceInfo dominanceInfo(constrainFunc);
915 constrainFunc.walk([&](EmitEqualityOp constraintOp) {
916 if (!llvm::isa<FeltType>(constraintOp.
getLhs().getType()) ||
917 !llvm::isa<FeltType>(constraintOp.
getRhs().getType())) {
923 unsigned degreeLhs = getDegree(lhsOperand.get(), degreeMemo);
924 unsigned degreeRhs = getDegree(rhsOperand.get(), degreeMemo);
926 if (degreeLhs > maxDegree) {
927 Value loweredExpr = lowerExpression(
928 lhsOperand.get(), structDef, constrainFunc, constraintOp.getOperation(), dominanceInfo,
929 degreeMemo, rewrites, auxAssignments
931 lhsOperand.set(loweredExpr);
933 if (degreeRhs > maxDegree) {
934 Value loweredExpr = lowerExpression(
935 rhsOperand.get(), structDef, constrainFunc, constraintOp.getOperation(), dominanceInfo,
936 degreeMemo, rewrites, auxAssignments
938 rhsOperand.set(loweredExpr);
943 auto res = constrainFunc.walk([&](EmitContainmentOp containOp) -> WalkResult {
944 return lowerContainmentRhsValue(
945 containOp.
getRhsMutable(), structDef, constrainFunc, dominanceInfo, degreeMemo, rewrites,
946 auxAssignments, containOp
949 if (res.wasInterrupted()) {
954 constrainFunc.walk([&](CallOp callOp) {
956 SmallVector<Value> newOperands = llvm::to_vector(callOp.
getArgOperands());
957 bool modified =
false;
959 for (Value &arg : newOperands) {
960 if (!llvm::isa<FeltType>(arg.getType())) {
964 DenseMap<Value, unsigned> callMemo;
965 if (getDegree(arg, callMemo) > 1) {
966 arg = materializeCallArgument(
967 arg, structDef, constrainFunc, callOp, dominanceInfo, degreeMemo, rewrites,
975 OpBuilder builder(callOp);
976 builder.create<CallOp>(
977 callOp.getLoc(), callOp.getResultTypes(), callOp.
getCallee(),
992 rebuildInCompute(FuncDefOp computeFunc,
const SmallVector<AuxAssignment> &auxAssignments) {
993 DenseMap<Value, Value> rebuildMemo;
994 Block &computeBlock = computeFunc.
getBody().front();
995 OpBuilder builder(&computeBlock, computeBlock.getTerminator()->getIterator());
998 SmallVector<unsigned> orderedAuxAssignments;
999 orderedAuxAssignments.reserve(auxAssignments.size());
1000 if (failed(orderAuxAssignments(auxAssignments, orderedAuxAssignments))) {
1004 for (
unsigned assignIdx : orderedAuxAssignments) {
1005 const auto &assign = auxAssignments[assignIdx];
1011 builder.create<MemberWriteOp>(
1012 assign.computedValue.getLoc(), selfVal, builder.getStringAttr(assign.auxMemberName),
1015 if (assign.auxValue) {
1018 rebuildMemo[assign.auxValue] = rebuiltExpr;
1024 void runOnOperation()
override {
1025 ModuleOp moduleOp = getOperation();
1028 if (maxDegree < 2) {
1029 moduleOp.emitError()
1030 .append(
"Invalid max degree: ", maxDegree.getValue(),
". Must be >= 2.")
1032 signalPassFailure();
1036 auto moduleRes = moduleOp.walk([
this](StructDefOp structDef) -> WalkResult {
1039 return WalkResult::interrupt();
1043 if (!constrainFunc) {
1044 return structDef.emitOpError() <<
'"' << structDef.getName() <<
"\" doesn't have a \"@"
1049 return WalkResult::interrupt();
1054 return structDef.emitOpError() <<
'"' << structDef.getName() <<
"\" doesn't have a \"@"
1059 return WalkResult::interrupt();
1062 SmallVector<AuxAssignment> auxAssignments;
1063 if (failed(lowerInConstrain(structDef, constrainFunc, auxAssignments))) {
1064 return WalkResult::interrupt();
1067 if (failed(checkEqualityDegrees(constrainFunc))) {
1068 return WalkResult::interrupt();
1071 if (failed(checkContainmentRhsDegrees(constrainFunc))) {
1072 return WalkResult::interrupt();
1075 if (failed(checkStructConstrainCallArguments(constrainFunc))) {
1076 return WalkResult::interrupt();
1079 if (failed(rebuildInCompute(computeFunc, auxAssignments))) {
1080 return WalkResult::interrupt();
1083 }
catch (
const DegreeComputationError &err) {
1084 mlir::emitError(err.getLoc()) << err.what();
1085 return WalkResult::interrupt();
1088 return WalkResult::advance();
1091 if (moduleRes.wasInterrupted()) {
1092 signalPassFailure();
Shared utility function implementations for LLZK lowering passes.
#define AUXILIARY_MEMBER_PREFIX
std::optional<::llvm::SmallVector<::mlir::ArrayAttr > > getSubelementIndices() const
Return a list of all valid indices for this ArrayType.
::mlir::Type getElementType() const
::llzk::function::FuncDefOp getConstrainFuncOp()
Gets the FuncDefOp that defines the constrain function in this structure, if present,...
::llzk::function::FuncDefOp getComputeFuncOp()
Gets the FuncDefOp that defines the compute function in this structure, if present,...
::mlir::TypedValue<::mlir::Type > getRhs()
::mlir::OpOperand & getRhsMutable()
::mlir::OpOperand & getRhsMutable()
::mlir::TypedValue<::mlir::Type > getLhs()
::mlir::OpOperand & getLhsMutable()
::mlir::TypedValue<::mlir::Type > getRhs()
bool calleeIsStructConstrain()
Return true iff the callee function name is FUNC_NAME_CONSTRAIN within a StructDefOp.
::mlir::SymbolRefAttr getCallee()
::llvm::ArrayRef< int32_t > getNumDimsPerMap()
::mlir::Operation::operand_range getArgOperands()
static constexpr ::llvm::StringLiteral getOperationName()
::mlir::OperandRangeRange getMapOperands()
static ::llvm::SmallVector<::mlir::ValueRange > toVectorOfValueRange(::mlir::OperandRangeRange)
Allocate consecutive storage of the ValueRange instances in the parameter so it can be passed to the ...
::mlir::Value getSelfValueFromCompute()
Return the "self" value (i.e.
::mlir::Value getSelfValueFromConstrain()
Return the "self" value (i.e.
::mlir::Region & getBody()
::mlir::Pass::Option< unsigned > maxDegree
ExpressionValue add(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
constexpr char FUNC_NAME_COMPUTE[]
Symbol name for the witness generation (and resp.
Value rebuildExprInCompute(Value val, FuncDefOp computeFunc, OpBuilder &builder, DenseMap< Value, Value > &memo)
void replaceSubsequentUsesWith(Value oldVal, Value newVal, Operation *afterOp)
constexpr char FUNC_NAME_CONSTRAIN[]
ExpressionValue neg(const llvm::SMTSolverRef &solver, const ExpressionValue &val)
MemberDefOp addAuxMember(StructDefOp structDef, StringRef name, Type type)
ExpressionValue div(const llvm::SMTSolverRef &solver, Operation *op, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue mul(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
LogicalResult checkFuncBodyIsStraightLine(FuncDefOp func, StringRef passName)
ExpressionValue sub(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
LogicalResult checkForAuxMemberConflicts(StructDefOp structDef, StringRef prefix)