21#include <mlir/Analysis/DataFlow/DeadCodeAnalysis.h>
22#include <mlir/Dialect/SCF/IR/SCF.h>
24#include <llvm/ADT/EquivalenceClasses.h>
25#include <llvm/ADT/TypeSwitch.h>
43std::optional<UnreducedInterval> mergeUnreducedIntervals(
44 const std::optional<UnreducedInterval> &lhs,
const std::optional<UnreducedInterval> &rhs
46 if (!lhs.has_value() || !rhs.has_value()) {
49 return lhs->doUnion(*rhs);
53std::optional<UnreducedInterval>
55 if (!lhs.hasUnreducedInterval() || !rhs.hasUnreducedInterval()) {
58 return fn(lhs.getUnreducedInterval(), rhs.getUnreducedInterval());
63 if (expr.getInterval() != newInterval) {
64 refined = refined.dropUnreducedInterval();
69bool isInMaybeSkippedScfRegion(Operation *op) {
70 for (Operation *parent = op->getParentOp(); parent !=
nullptr; parent = parent->getParentOp()) {
71 if (llvm::isa<FuncDefOp>(parent)) {
78 if (llvm::isa<scf::ForOp, scf::IfOp, scf::WhileOp>(parent)) {
85std::optional<UnreducedInterval> getBooleanUnreducedInterval(
const Interval &interval) {
86 return interval.isBoolean() ? std::optional<UnreducedInterval>(interval.firstUnreduced())
90FailureOr<std::vector<SourceRef>>
92 std::vector<SourceRef> refs;
93 for (
const auto &[prefix, vals] : translations) {
94 if (!ref.isValidPrefix(prefix)) {
99 auto suffix = ref.getSuffix(prefix);
100 ensure(succeeded(suffix),
"prefix checked before SourceRef suffix extraction");
102 std::vector<SourceRefIndex> arraySuffix, remainingSuffix;
103 bool suffixIsPastArray =
false;
105 if (!suffixIsPastArray && arraySuffix.size() < vals.getNumArrayDims() &&
106 (idx.isIndex() || idx.isIndexRange())) {
107 arraySuffix.push_back(idx);
110 suffixIsPastArray =
true;
111 remainingSuffix.push_back(idx);
114 auto resolvedValsRes = vals.extract(arraySuffix);
115 ensure(succeeded(resolvedValsRes),
"could not resolve translated SourceRef array child");
116 SourceRefSet folded = resolvedValsRes->first.foldToScalar();
117 if (remainingSuffix.empty()) {
118 refs.insert(refs.end(), folded.begin(), folded.end());
122 for (
const SourceRef &baseRef : folded) {
123 auto translatedRef = mlir::FailureOr<SourceRef>(baseRef);
125 if (failed(translatedRef)) {
128 translatedRef = translatedRef->createChild(idx);
130 if (succeeded(translatedRef)) {
131 refs.push_back(*translatedRef);
135 for (
const SourceRef &replacement : vals.getScalarValue()) {
136 auto translated = ref.translate(prefix, replacement);
137 if (succeeded(translated)) {
138 refs.push_back(*translated);
150bool isDirectSourceRefValue(Value value) {
151 if (llvm::isa<BlockArgument>(value)) {
155 Operation *definingOp = value.getDefiningOp();
156 return llvm::isa_and_present<MemberReadOp, ReadArrayOp, polymorphic::ConstReadOp>(definingOp);
159std::optional<SourceRefLatticeValue>
160getIdentitySourceRefState(DataFlowSolver &solver, Value value) {
161 if (isDirectSourceRefValue(value)) {
163 if (val.isScalar()) {
169 auto createArray = llvm::dyn_cast_if_present<CreateArrayOp>(value.getDefiningOp());
175 for (
auto [idx, element] : llvm::enumerate(createArray.getElements())) {
176 std::optional<SourceRefLatticeValue> elementVal = getIdentitySourceRefState(solver, element);
177 if (!elementVal.has_value()) {
180 (void)arrayVal.getElemFlatIdx(idx).setValue(*elementVal);
185llvm::EquivalenceClasses<SourceRef>
186collectDirectEqualityRefs(DataFlowSolver &solver,
FuncDefOp fn) {
187 llvm::EquivalenceClasses<SourceRef> eqRefs;
189 Operation *op = eqOp.getOperation();
194 Value lhs = eqOp.getLhs();
195 Value rhs = eqOp.getRhs();
196 if (!isDirectSourceRefValue(lhs) || !isDirectSourceRefValue(rhs)) {
202 if (!lhsState.isScalar() || !rhsState.isScalar() || !lhsState.isSingleValue() ||
203 !rhsState.isSingleValue()) {
207 const SourceRef &lhsRef = lhsState.getSingleValue();
208 const SourceRef &rhsRef = rhsState.getSingleValue();
209 if (lhsRef.isConstant() || rhsRef.isConstant()) {
212 eqRefs.unionSets(lhsRef, rhsRef);
222 const llvm::SMTSolverRef &solver, Operation *op,
const ExpressionValue &val,
223 StringRef suffix =
""
228 DynamicAPInt invVal = field.
inv(iv.
lhs());
236 if (!suffix.empty()) {
237 symName += suffix.str();
239 llvm::SMTExprRef invSym = field.
createSymbol(solver, symName.c_str());
240 llvm::SMTExprRef one = solver->mkBitvector(APSInt::get(1), field.
bitWidth());
242 llvm::SMTExprRef mult = solver->mkBVMul(val.
getExpr(), invSym);
243 llvm::SMTExprRef
mod = solver->mkBVURem(mult, prime);
244 llvm::SMTExprRef constraint = solver->mkEqual(
mod, one);
245 solver->addConstraint(constraint);
250 if (expr ==
nullptr && rhs.expr ==
nullptr) {
251 return i == rhs.i && unreduced == rhs.unreduced;
253 if (expr ==
nullptr || rhs.expr ==
nullptr) {
256 return i == rhs.i && unreduced == rhs.unreduced && *expr == *rhs.expr;
261 llvm::SMTExprRef zero = solver->mkBitvector(mlir::APSInt::get(0), bitwidth);
262 llvm::SMTExprRef one = solver->mkBitvector(mlir::APSInt::get(1), bitwidth);
263 llvm::SMTExprRef boolToFeltConv = solver->mkIte(expr.
getExpr(), one, zero);
283 llvm::SMTExprRef resultExpr =
285 std::optional<UnreducedInterval> resultUnreduced;
293 resultUnreduced = mergeUnreducedIntervals(
297 return ExpressionValue(resultExpr, resultInterval, std::move(resultUnreduced));
304 const auto *exprEq = solver->mkEqual(lhs.expr, rhs.expr);
311 res.i = lhs.i + rhs.i;
312 res.expr = solver->mkBVAdd(lhs.expr, rhs.expr);
320 res.i = lhs.i - rhs.i;
321 res.expr = solver->mkBVSub(lhs.expr, rhs.expr);
329 res.i = lhs.i * rhs.i;
330 res.expr = solver->mkBVMul(lhs.expr, rhs.expr);
339 auto divRes =
feltDiv(lhs.i, rhs.i);
340 if (failed(divRes)) {
347 "non-degenerate felt.div divisors are not tracked precisely, and the divisor may "
348 "contain zero. Range of division result will be treated as unbounded."
353 "non-degenerate felt.div divisors are not tracked precisely because precise field "
354 "division over intervals would require enumerating divisor inverses. Range of "
355 "division result will be treated as unbounded."
361 "divisor is zero, leading to a divide-by-zero error. Range of division result will "
362 "be treated as unbounded."
371 res.expr = solver->mkBVMul(lhs.expr, invExpr);
376 const llvm::SMTSolverRef &solver, Operation *op,
const ExpressionValue &lhs,
381 if (failed(divRes)) {
383 "divisor is not restricted to non-zero values, leading to potential divide-by-zero error."
384 " Range of division result will be treated as unbounded."
391 res.expr = solver->mkBVUDiv(lhs.expr, rhs.expr);
396 const llvm::SMTSolverRef &solver, Operation *op,
const ExpressionValue &lhs,
401 if (failed(divRes)) {
403 "divisor is not restricted to non-zero values, leading to potential divide-by-zero error."
404 " Range of division result will be treated as unbounded."
411 res.expr = solver->mkBVSDiv(lhs.expr, rhs.expr);
418 res.i = lhs.i % rhs.i;
419 res.expr = solver->mkBVURem(lhs.expr, rhs.expr);
434 res.i = lhs.i & rhs.i;
435 res.expr = solver->mkBVAnd(lhs.expr, rhs.expr);
442 res.i = lhs.i | rhs.i;
443 res.expr = solver->mkBVOr(lhs.expr, rhs.expr);
450 return boolXor(solver, lhs, rhs);
454 res.i = lhs.i ^ rhs.i;
455 res.expr = solver->mkBVXor(lhs.expr, rhs.expr);
463 res.i = lhs.i << rhs.i;
464 res.expr = solver->mkBVShl(lhs.expr, rhs.expr);
472 res.i = lhs.i >> rhs.i;
473 res.expr = solver->mkBVLshr(lhs.expr, rhs.expr);
485 case FeltCmpPredicate::EQ:
486 res.expr = solver->mkEqual(lhs.expr, rhs.expr);
493 case FeltCmpPredicate::NE:
494 res.expr = solver->mkNot(solver->mkEqual(lhs.expr, rhs.expr));
501 case FeltCmpPredicate::LT:
502 res.expr = solver->mkBVUlt(lhs.expr, rhs.expr);
510 case FeltCmpPredicate::LE:
511 res.expr = solver->mkBVUle(lhs.expr, rhs.expr);
519 case FeltCmpPredicate::GT:
520 res.expr = solver->mkBVUgt(lhs.expr, rhs.expr);
528 case FeltCmpPredicate::GE:
529 res.expr = solver->mkBVUge(lhs.expr, rhs.expr);
546 res.expr = solver->mkAnd(lhs.expr, rhs.expr);
554 res.i =
boolOr(lhs.i, rhs.i);
555 res.expr = solver->mkOr(lhs.expr, rhs.expr);
565 res.expr = solver->mkAnd(
566 solver->mkOr(lhs.expr, rhs.expr), solver->mkNot(solver->mkAnd(lhs.expr, rhs.expr))
575 res.expr = solver->mkBVNeg(val.expr);
585 res.expr = solver->mkBVNot(val.expr);
592 res.expr = solver->mkNot(val.expr);
602 res.expr = TypeSwitch<Operation *, llvm::SMTExprRef>(op)
605 }).Default([](Operation *unsupported) {
606 llvm::report_fatal_error(
607 "no fallback provided for " + mlir::Twine(unsupported->getName().getStringRef())
612 if (llvm::isa<InvFeltOp>(op)) {
626 os <<
"<null expression>";
629 os <<
" ( interval: " << i <<
" )";
630 if (unreduced.has_value()) {
631 os <<
" ( unreduced: " << *unreduced <<
" )";
640 return ChangeResult::NoChange;
646 return ChangeResult::NoChange;
650 os <<
"IntervalAnalysisLattice { " << val <<
" }";
655 return ChangeResult::NoChange;
658 return ChangeResult::Change;
667 if (!constraints.contains(e)) {
668 constraints.insert(e);
669 return ChangeResult::Change;
671 return ChangeResult::NoChange;
680std::vector<SourceRefIndex> IntervalDataFlowAnalysis::getArrayAccessIndices(
681 Operation * , ArrayAccessOpInterface arrayAccessOp
683 std::vector<SourceRefIndex> indices;
684 ArrayType arrayType = arrayAccessOp.getArrRefType();
685 size_t numIndices = arrayAccessOp.getIndices().size();
686 indices.reserve(numIndices);
688 for (
size_t i = 0; i < numIndices; ++i) {
689 Value idxOperand = arrayAccessOp.getIndices()[i];
690 SourceRefLatticeValue idxVals = getSourceRefState(idxOperand);
693 if (idxVals.isSingleValue() && idxVals.getSingleValue().isConstant()) {
694 indices.emplace_back(*idxVals.getSingleValue().getConstantValue());
696 auto lower = APInt::getZero(64);
697 APInt upper(64, arrayType.getDimSize(i));
698 indices.emplace_back(lower, upper);
705mlir::FailureOr<SourceRef> IntervalDataFlowAnalysis::getArrayAccessRef(
708 std::vector<SourceRefIndex> indices = getArrayAccessIndices(baseOp, arrayAccessOp);
709 Value arrayVal = arrayAccessOp.getArrRef();
710 if (
auto blockArg = llvm::dyn_cast<BlockArgument>(arrayVal)) {
711 return SourceRef(blockArg, std::move(indices));
713 if (
auto result = llvm::dyn_cast<OpResult>(arrayVal)) {
714 return SourceRef(result, std::move(indices));
720 if (
auto it = writeResults.find(ref); it != writeResults.end()) {
721 return it->second.getInterval();
724 if (ref.isConstantInt()) {
725 auto constVal = ref.getConstantValue();
726 if (succeeded(constVal)) {
731 if (ref.isRooted() && ref.getPath().empty()) {
732 auto rootVal = ref.getRoot();
733 if (succeeded(rootVal) && !llvm::isa<ArrayType, StructType, pod::PodType>(rootVal->getType())) {
735 if (rootExpr.getExpr() !=
nullptr) {
736 return rootExpr.getInterval();
741 return getDefaultIntervalForType(ref.getType());
744std::optional<UnreducedInterval>
745IntervalDataFlowAnalysis::getDefaultUnreducedIntervalForType(mlir::Type ty)
const {
746 if (!trackUnreducedIntervals) {
749 if (isBooleanType(ty)) {
750 return UnreducedInterval(0, 1);
752 return UnreducedInterval(field.get().zero(), field.get().maxVal());
755std::optional<UnreducedInterval>
756IntervalDataFlowAnalysis::getRefUnreducedInterval(
const SourceRef &ref) {
757 if (!trackUnreducedIntervals) {
761 if (
auto it = writeResults.find(ref); it != writeResults.end()) {
762 return it->second.getOptionalUnreducedInterval();
765 if (ref.isConstantInt()) {
766 auto constVal = ref.getConstantValue();
767 if (succeeded(constVal)) {
768 return UnreducedInterval(*constVal, *constVal);
772 if (ref.isRooted() && ref.getPath().empty()) {
773 auto rootVal = ref.getRoot();
774 if (succeeded(rootVal) && !llvm::isa<ArrayType, StructType, pod::PodType>(rootVal->getType())) {
776 if (rootExpr.hasUnreducedInterval()) {
777 return rootExpr.getUnreducedInterval();
782 return getRefInterval(ref).firstUnreduced();
786 if (
auto it = writeResults.find(ref); it != writeResults.end()) {
789 return createUnknownValue(val)
794void IntervalDataFlowAnalysis::recordRefWrite(
797 auto joinStoredWrite = [
this, &writtenRef](
798 const ExpressionValue &old,
const ExpressionValue &next
799 ) -> ExpressionValue {
800 Interval combinedWrite = old.getInterval().join(next.getInterval());
801 auto combinedUnreduced = mergeUnreducedIntervals(
802 old.getOptionalUnreducedInterval(), next.getOptionalUnreducedInterval()
804 if (old.getExpr() !=
nullptr && next.getExpr() !=
nullptr &&
805 *old.getExpr() == *next.getExpr()) {
806 return old.withInterval(combinedWrite).withOptionalUnreducedInterval(combinedUnreduced);
809 return ExpressionValue(
814 if (
auto it = writeResults.find(writtenRef); it != writeResults.end()) {
815 it->second = joinStoredWrite(it->second, writeVal);
816 }
else if (mayBeSkipped) {
817 ExpressionValue noWrite(
819 getRefUnreducedInterval(writtenRef)
821 writeResults[writtenRef] = joinStoredWrite(noWrite, writeVal);
823 writeResults[writtenRef] = writeVal;
826 const ExpressionValue &readerUpdate = mayBeSkipped ? writeResults[writtenRef] : writeVal;
827 for (Lattice *readerLattice : readResults[writtenRef]) {
828 ExpressionValue prior = readerLattice->getValue().getScalarValue();
830 ExpressionValue newVal = prior.withInterval(
intersection);
831 propagateIfChanged(readerLattice, readerLattice->setValue(newVal));
836 Operation *op, ArrayRef<const Lattice *> operands, ArrayRef<Lattice *> results
845 if (operands.empty() && results.empty()) {
850 llvm::SmallVector<LatticeValue> operandVals;
851 llvm::SmallVector<std::optional<SourceRef>> operandRefs;
852 auto resolveRefStateValue =
854 ensure(refSet.isScalar(),
"should have ruled out array values already");
856 if (refSet.getScalarValue().empty()) {
863 " is empty; defining operation is unsupported by SourceRef analysis"
869 if (!refSet.isSingleValue()) {
871 std::optional<UnreducedInterval> joinedUnreduced = std::nullopt;
872 bool sawFirst =
false;
873 for (
const SourceRef &ref : refSet.getScalarValue()) {
874 joinedInterval = joinedInterval.
join(getRefInterval(ref));
875 auto refUnreduced = getRefUnreducedInterval(ref);
877 joinedUnreduced = refUnreduced;
880 joinedUnreduced = mergeUnreducedIntervals(joinedUnreduced, refUnreduced);
886 return LatticeValue(anyVal);
889 return LatticeValue(getRefValue(refSet.getSingleValue(), value));
891 for (
unsigned opNum = 0; opNum < op->getNumOperands(); ++opNum) {
892 Value val = op->getOperand(opNum);
897 operandRefs.push_back(std::nullopt);
900 auto priorState = operands[opNum]->getValue();
901 if (priorState.getScalarValue().getExpr() !=
nullptr) {
902 operandVals.push_back(priorState);
906 if (
auto readArr = llvm::dyn_cast_if_present<ReadArrayOp>(val.getDefiningOp())) {
907 auto arrayRef = getArrayAccessRef(op, readArr);
908 if (succeeded(arrayRef)) {
909 if (
auto it = writeResults.find(*arrayRef); it != writeResults.end()) {
910 operandVals.emplace_back(it->second);
912 (void)operandLattice->
setValue(it->second);
920 Type valTy = val.getType();
921 if (llvm::isa<ArrayType, StructType, pod::PodType>(valTy)) {
923 operandVals.emplace_back(anyVal);
927 auto resolvedValue = resolveRefStateValue(val, refSet);
928 if (!resolvedValue.has_value()) {
933 operandVals.push_back(*resolvedValue);
939 (void)operandLattice->
setValue(operandVals[opNum]);
942 if (isReadOp(op) && op->getNumResults() == 1) {
943 Value resultVal = op->getResult(0);
944 if (!llvm::isa<ArrayType, StructType, pod::PodType>(resultVal.getType())) {
945 auto resolvedValue = resolveRefStateValue(resultVal, getSourceRefState(resultVal));
946 if (resolvedValue.has_value()) {
947 propagateIfChanged(results[0], results[0]->setValue(*resolvedValue));
955 llvm::DynamicAPInt constVal = getConst(op);
956 llvm::SMTExprRef expr;
957 if (isBoolConstOp(op)) {
958 expr = createConstBoolExpr(constVal != 0);
960 expr = createConstBitvectorExpr(constVal);
964 if (trackUnreducedIntervals) {
967 propagateIfChanged(results[0], results[0]->setValue(latticeVal));
968 }
else if (isArithmeticOp(op)) {
970 if (operands.size() == 2) {
971 result = performBinaryArithmetic(op, operandVals[0], operandVals[1]);
973 result = performUnaryArithmetic(op, operandVals[0]);
977 const ExpressionValue &prior = results[0]->getValue().getScalarValue();
981 propagateIfChanged(results[0], results[0]->setValue(result));
982 }
else if (
auto selectOp = llvm::dyn_cast<arith::SelectOp>(op)) {
984 smtSolver, operandVals[0].getScalarValue(), operandVals[1].getScalarValue(),
985 operandVals[2].getScalarValue()
987 const ExpressionValue &prior = results[0]->getValue().getScalarValue();
991 propagateIfChanged(results[0], results[0]->setValue(result));
992 }
else if (
EmitEqualityOp emitEq = llvm::dyn_cast<EmitEqualityOp>(op)) {
993 Value lhsVal = emitEq.getLhs(), rhsVal = emitEq.getRhs();
999 auto res = getGeneralizedDecompInterval(op, lhsVal, rhsVal);
1000 if (succeeded(res)) {
1001 for (Value signalVal : res->first) {
1002 applyInterval(emitEq, signalVal, res->second);
1010 applyInterval(emitEq, lhsVal, constrainInterval);
1011 applyInterval(emitEq, rhsVal, constrainInterval);
1012 }
else if (
auto assertOp = llvm::dyn_cast<AssertOp>(op)) {
1015 Value cond = assertOp.getCondition();
1018 auto assertExpr = operandVals[0].getScalarValue();
1021 }
else if (
auto writePod = llvm::dyn_cast<pod::WritePodOp>(op)) {
1022 const bool maySkipWrite = isInMaybeSkippedScfRegion(op);
1026 ensure(succeeded(recordRefsRes),
"could not create SourceRef child for pod write");
1028 Type valueTy = writePod.
getValue().getType();
1029 if (!llvm::isa<ArrayType, StructType, pod::PodType>(valueTy)) {
1032 recordRefWrite(recordRef, writeVal, maySkipWrite);
1034 }
else if (operandRefs[1].has_value()) {
1035 llvm::SmallVector<std::pair<SourceRef, ExpressionValue>> remappedWrites;
1037 for (
const auto &[writtenRef, writtenVal] : writeResults) {
1038 if (writtenRef.isValidPrefix(*operandRefs[1])) {
1039 auto translated = writtenRef.translate(*operandRefs[1], recordRef);
1040 ensure(succeeded(translated),
"could not translate aggregate pod write");
1041 remappedWrites.emplace_back(*translated, writtenVal);
1045 for (
const auto &[translatedRef, translatedVal] : remappedWrites) {
1046 recordRefWrite(translatedRef, translatedVal, maySkipWrite);
1050 }
else if (
auto writem = llvm::dyn_cast<MemberWriteOp>(op)) {
1051 const bool maySkipWrite = isInMaybeSkippedScfRegion(op);
1054 auto cmp = writem.getComponent();
1058 auto memberDefRes = writem.getMemberDefOp(tables);
1059 if (succeeded(memberDefRes)) {
1062 ensure(succeeded(memberRefRes),
"could not create SourceRef child for member write");
1063 const SourceRef &memberRef = *memberRefRes;
1064 Type memberTy = writem.getVal().
getType();
1065 if (!llvm::isa<ArrayType, StructType, pod::PodType>(memberTy)) {
1067 recordRefWrite(memberRef, writeVal, maySkipWrite);
1070 std::optional<SourceRef> rhsPrefix;
1071 if (operandRefs[1].has_value() && operandRefs[1]->isRooted()) {
1072 rhsPrefix = operandRefs[1];
1073 }
else if (
auto blockArg = llvm::dyn_cast<BlockArgument>(writem.getVal())) {
1075 }
else if (
auto result = llvm::dyn_cast<OpResult>(writem.getVal())) {
1079 if (rhsPrefix.has_value()) {
1080 llvm::SmallVector<std::pair<SourceRef, ExpressionValue>> remappedWrites;
1081 for (
const auto &[writtenRef, writtenVal] : writeResults) {
1082 if (!writtenRef.isValidPrefix(*rhsPrefix)) {
1086 auto translatedRef = writtenRef.translate(*rhsPrefix, memberRef);
1087 ensure(succeeded(translatedRef),
"could not translate composite member write");
1088 remappedWrites.emplace_back(*translatedRef, writtenVal);
1091 for (
const auto &[translatedRef, translatedVal] : remappedWrites) {
1092 recordRefWrite(translatedRef, translatedVal, maySkipWrite);
1098 }
else if (
auto writeArr = llvm::dyn_cast<WriteArrayOp>(op)) {
1099 const bool maySkipWrite = isInMaybeSkippedScfRegion(op);
1101 auto arrayRef = getArrayAccessRef(op, writeArr);
1102 if (succeeded(arrayRef)) {
1103 recordRefWrite(*arrayRef, writeVal, maySkipWrite);
1108 std::vector<SourceRefIndex> indices = getArrayAccessIndices(op, writeArr);
1109 auto targetRefsRes = arrayVals.
extract(indices);
1110 ensure(succeeded(targetRefsRes),
"could not create SourceRef child for array write");
1111 auto [targetRefs, _] = *targetRefsRes;
1112 ensure(targetRefs.isScalar(),
"array write must resolve to scalar references");
1113 for (
const SourceRef &ref : targetRefs.getScalarValue()) {
1114 recordRefWrite(ref, writeVal, maySkipWrite);
1117 }
else if (
auto createArray = llvm::dyn_cast<CreateArrayOp>(op)) {
1118 const auto &elements = createArray.getElements();
1119 ArrayType arrayTy = createArray.getType();
1122 if (!elements.empty() && !llvm::isa<ArrayType, StructType, pod::PodType>(elemTy)) {
1123 ensure(arrayTy.hasStaticShape(),
"array.new with explicit elements must have static shape");
1125 std::cmp_equal(elements.size(), arrayTy.getNumElements()),
1126 "array.new explicit initializer length must match array shape"
1130 auto arrayRes = llvm::cast<OpResult>(createArray->getResult(0));
1131 for (
unsigned i = 0; i < elements.size(); ++i) {
1132 auto maybeIndices = indexGen.
delinearize(i, op->getContext());
1133 ensure(maybeIndices.has_value(),
"could not delinearize array.new element index");
1136 path.reserve(maybeIndices->size());
1137 for (Attribute attr : *maybeIndices) {
1138 auto idxAttr = llvm::dyn_cast<IntegerAttr>(attr);
1139 ensure(idxAttr !=
nullptr,
"array.new delinearize should produce integer attributes");
1140 path.emplace_back(idxAttr.getValue());
1143 recordRefWrite(
SourceRef(arrayRes, std::move(path)), operandVals[i].getScalarValue());
1145 }
else if (!elements.empty()) {
1146 ensure(arrayTy.hasStaticShape(),
"aggregate array.new initializer requires static shape");
1148 SourceRef arrayRoot(llvm::cast<OpResult>(createArray.getResult()));
1149 llvm::SmallVector<std::pair<SourceRef, ExpressionValue>> remappedWrites;
1150 for (
auto [i, element] : llvm::enumerate(elements)) {
1155 auto maybeIndices = indexGen.
delinearize(i, op->getContext());
1156 ensure(maybeIndices.has_value(),
"could not delinearize aggregate array.new index");
1158 for (Attribute attr : *maybeIndices) {
1161 ensure(succeeded(child),
"could not create aggregate array element SourceRef");
1162 elementTarget = *child;
1164 for (
const auto &[writtenRef, writtenVal] : writeResults) {
1166 auto translated = writtenRef.translate(elementRefs.
getSingleValue(), elementTarget);
1167 ensure(succeeded(translated),
"could not translate aggregate array initializer");
1168 remappedWrites.emplace_back(*translated, writtenVal);
1172 for (
const auto &[translatedRef, translatedVal] : remappedWrites) {
1173 recordRefWrite(translatedRef, translatedVal);
1176 }
else if (
auto newPod = llvm::dyn_cast<pod::NewPodOp>(op)) {
1177 SourceRef podRoot(llvm::cast<OpResult>(newPod.getResult()));
1178 for (
auto [idx, record] : llvm::enumerate(newPod.getInitializedRecordValues())) {
1181 ensure(succeeded(recordRef),
"could not create SourceRef child for pod initializer");
1182 if (!llvm::isa<ArrayType, StructType, pod::PodType>(record.value.getType())) {
1183 recordRefWrite(*recordRef, operandVals[idx].getScalarValue());
1191 llvm::SmallVector<std::pair<SourceRef, ExpressionValue>> remappedWrites;
1192 for (
const auto &[writtenRef, writtenVal] : writeResults) {
1194 auto translated = writtenRef.translate(sourceRefs.
getSingleValue(), *recordRef);
1195 ensure(succeeded(translated),
"could not translate aggregate pod initializer");
1196 remappedWrites.emplace_back(*translated, writtenVal);
1199 for (
const auto &[translatedRef, translatedVal] : remappedWrites) {
1200 recordRefWrite(translatedRef, translatedVal);
1203 }
else if (isa<IntToFeltOp, FeltToIndexOp>(op)) {
1210 expr =
boolToFelt(smtSolver, expr, field.get().bitWidth());
1212 propagateIfChanged(results[0], results[0]->setValue(expr));
1213 }
else if (
auto yieldOp = dyn_cast<scf::YieldOp>(op)) {
1216 Operation *parent = op->getParentOp();
1217 ensure(parent,
"yield operation must have parent operation");
1219 for (
unsigned idx = 0; idx < yieldOp.getResults().size(); ++idx) {
1220 Value parentRes = parent->getResult(idx);
1226 if (
auto loopOp = llvm::dyn_cast<LoopLikeOpInterface>(parent)) {
1230 if (exprVal.
getExpr() !=
nullptr) {
1242 propagateIfChanged(resLattice, resLattice->
setValue(newResVal));
1252 && !isDefinitionOp(op)
1254 && !llvm::isa<CreateArrayOp, CreateStructOp, pod::NewPodOp, NonDetOp>(op)
1256 op->emitWarning(
"unhandled operation, analysis may be incomplete").report();
1263 auto it = refSymbols.find(r);
1264 if (it != refSymbols.end()) {
1267 llvm::SMTExprRef sym = createSymbol(r);
1268 refSymbols[r] = sym;
1272llvm::SMTExprRef IntervalDataFlowAnalysis::createSymbol(mlir::Type ty,
const char *name)
const {
1273 if (isBooleanType(ty)) {
1274 return smtSolver->mkSymbol(name, smtSolver->getBoolSort());
1276 return field.get().createSymbol(smtSolver, name);
1279llvm::SMTExprRef IntervalDataFlowAnalysis::createSymbol(
const SourceRef &r)
const {
1281 return createSymbol(r.getType(), name.c_str());
1284llvm::SMTExprRef IntervalDataFlowAnalysis::createSymbol(Value v)
const {
1286 return createSymbol(v.getType(), name.c_str());
1289llvm::DynamicAPInt IntervalDataFlowAnalysis::getConst(Operation *op)
const {
1290 ensure(isConstOp(op),
"op is not a const op");
1294 llvm::DynamicAPInt fieldConst = TypeSwitch<Operation *, llvm::DynamicAPInt>(op)
1295 .Case<FeltConstantOp>([&](
auto feltConst) {
1296 llvm::APSInt constOpVal(feltConst.getValue());
1297 return field.get().reduce(constOpVal);
1299 .Case<arith::ConstantIndexOp>([&](
auto indexConst) {
1300 return DynamicAPInt(indexConst.value());
1302 .Case<arith::ConstantIntOp>([&](
auto intConst) {
1303 auto valAttr = dyn_cast<IntegerAttr>(intConst.getValue());
1304 ensure(valAttr !=
nullptr,
"arith::ConstantIntOp must have an IntegerAttr as its value");
1307 .Default([](
auto *illegalOp) {
1309 debug::Appender(err) <<
"unhandled getConst case: " << *illegalOp;
1310 llvm::report_fatal_error(Twine(err));
1311 return llvm::DynamicAPInt();
1318 Operation *op,
const LatticeValue &a,
const LatticeValue &b
1320 ensure(isArithmeticOp(op),
"is not arithmetic op");
1322 auto lhs = a.getScalarValue(), rhs = b.getScalarValue();
1323 ensure(lhs.getExpr(),
"cannot perform arithmetic over null lhs smt expr");
1324 ensure(rhs.getExpr(),
"cannot perform arithmetic over null rhs smt expr");
1327 auto res = TypeSwitch<Operation *, ExpressionValue>(op)
1328 .Case<AddFeltOp>([&](
auto) {
return add(smtSolver, lhs, rhs); })
1329 .Case<SubFeltOp>([&](
auto) {
return sub(smtSolver, lhs, rhs); })
1330 .Case<MulFeltOp>([&](
auto) {
return mul(smtSolver, lhs, rhs); })
1331 .Case<DivFeltOp>([&](
auto) {
return div(smtSolver, op, lhs, rhs); })
1332 .Case<UnsignedIntDivFeltOp>([&](
auto) {
return uintDiv(smtSolver, op, lhs, rhs); })
1333 .Case<SignedIntDivFeltOp>([&](
auto) {
return sintDiv(smtSolver, op, lhs, rhs); })
1334 .Case<UnsignedModFeltOp>([&](
auto) {
return mod(smtSolver, lhs, rhs); })
1335 .Case<SignedModFeltOp>([&](
auto) {
return sintMod(smtSolver, lhs, rhs); })
1336 .Case<AndFeltOp>([&](
auto) {
return bitAnd(smtSolver, lhs, rhs); })
1337 .Case<OrFeltOp>([&](
auto) {
return bitOr(smtSolver, lhs, rhs); })
1338 .Case<XorFeltOp, arith::XOrIOp>([&](
auto) {
return bitXor(smtSolver, lhs, rhs); })
1339 .Case<ShlFeltOp>([&](
auto) {
return shiftLeft(smtSolver, lhs, rhs); })
1340 .Case<ShrFeltOp>([&](
auto) {
return shiftRight(smtSolver, lhs, rhs); })
1341 .Case<CmpOp>([&](
auto cmpOp) {
return cmp(smtSolver, cmpOp, lhs, rhs); })
1342 .Case<AndBoolOp>([&](
auto) {
return boolAnd(smtSolver, lhs, rhs); })
1343 .Case<OrBoolOp>([&](
auto) {
return boolOr(smtSolver, lhs, rhs); })
1344 .Case<XorBoolOp>([&](
auto) {
return boolXor(smtSolver, lhs, rhs); })
1345 .Default([&](
auto *unsupported) {
1348 "unsupported binary arithmetic operation"
1351 return ExpressionValue();
1355 ensure(res.getExpr(),
"arithmetic produced null smt expr");
1360IntervalDataFlowAnalysis::performUnaryArithmetic(Operation *op,
const LatticeValue &a) {
1361 ensure(isArithmeticOp(op),
"is not arithmetic op");
1363 auto val = a.getScalarValue();
1364 ensure(val.getExpr(),
"cannot perform arithmetic over null smt expr");
1366 auto res = TypeSwitch<Operation *, ExpressionValue>(op)
1367 .Case<NegFeltOp>([&](
auto) {
return neg(smtSolver, val); })
1368 .Case<NotFeltOp>([&](
auto) {
return notOp(smtSolver, val); })
1369 .Case<NotBoolOp>([&](
auto) {
return boolNot(smtSolver, val); })
1371 .Case<InvFeltOp>([&](
auto inv) {
1373 }).Default([&](
auto *unsupported) {
1376 "unsupported unary arithmetic operation, defaulting to over-approximated interval"
1382 ensure(res.getExpr(),
"arithmetic produced null smt expr");
1386void IntervalDataFlowAnalysis::applyInterval(Operation *valUser, Value val,
Interval newInterval) {
1388 ExpressionValue oldLatticeVal = valLattice->getValue().getScalarValue();
1391 ExpressionValue newLatticeVal = refineReducedInterval(oldLatticeVal,
intersection);
1392 ChangeResult changed = valLattice->setValue(newLatticeVal);
1394 if (
auto blockArg = llvm::dyn_cast<BlockArgument>(val)) {
1395 auto fnOp = dyn_cast<FuncDefOp>(blockArg.getOwner()->getParentOp());
1398 if (propagateInputConstraints && fnOp && fnOp.isStructConstrain() &&
1399 blockArg.getArgNumber() > 0 && !newInterval.isEntire()) {
1400 auto structOp = fnOp->getParentOfType<StructDefOp>();
1401 FuncDefOp computeFn = structOp.getComputeFuncOp();
1402 BlockArgument computeArg = computeFn.getArgument(blockArg.getArgNumber() - 1);
1405 SourceRef ref(computeArg);
1406 ExpressionValue newArgVal(
1408 trackUnreducedIntervals ? std::optional<UnreducedInterval>(newInterval.firstUnreduced())
1411 propagateIfChanged(computeEntryLattice, computeEntryLattice->setValue(newArgVal));
1416 Operation *definingOp = val.getDefiningOp();
1418 propagateIfChanged(valLattice, changed);
1422 const Field &f = field.get();
1430 auto cmpCase = [&](CmpOp cmpOp) {
1436 newInterval.isBoolean() || newInterval.isEmpty(),
1437 "new interval for CmpOp is not boolean or empty"
1439 if (!newInterval.isDegenerate()) {
1444 bool cmpTrue = newInterval.rhs() == f.one();
1446 Value lhs = cmpOp.getLhs(), rhs = cmpOp.getRhs();
1448 ExpressionValue lhsExpr = lhsLat->getValue().getScalarValue(),
1449 rhsExpr = rhsLat->getValue().getScalarValue();
1451 Interval newLhsInterval, newRhsInterval;
1452 const Interval &lhsInterval = lhsExpr.getInterval();
1453 const Interval &rhsInterval = rhsExpr.getInterval();
1457 auto eqCase = [&]() {
1458 return (pred == FeltCmpPredicate::EQ && cmpTrue) ||
1459 (pred == FeltCmpPredicate::NE && !cmpTrue);
1461 auto neCase = [&]() {
1462 return (pred == FeltCmpPredicate::NE && cmpTrue) ||
1463 (pred == FeltCmpPredicate::EQ && !cmpTrue);
1465 auto ltCase = [&]() {
1466 return (pred == FeltCmpPredicate::LT && cmpTrue) ||
1467 (pred == FeltCmpPredicate::GE && !cmpTrue);
1469 auto leCase = [&]() {
1470 return (pred == FeltCmpPredicate::LE && cmpTrue) ||
1471 (pred == FeltCmpPredicate::GT && !cmpTrue);
1473 auto gtCase = [&]() {
1474 return (pred == FeltCmpPredicate::GT && cmpTrue) ||
1475 (pred == FeltCmpPredicate::LE && !cmpTrue);
1477 auto geCase = [&]() {
1478 return (pred == FeltCmpPredicate::GE && cmpTrue) ||
1479 (pred == FeltCmpPredicate::LT && !cmpTrue);
1484 newLhsInterval = newRhsInterval = lhsInterval.intersect(rhsInterval);
1485 }
else if (neCase()) {
1486 if (lhsInterval.isDegenerate() && rhsInterval.isDegenerate() && lhsInterval == rhsInterval) {
1490 }
else if (lhsInterval.isDegenerate()) {
1492 newLhsInterval = lhsInterval;
1493 newRhsInterval = rhsInterval.difference(lhsInterval);
1494 }
else if (rhsInterval.isDegenerate()) {
1496 newLhsInterval = lhsInterval.difference(rhsInterval);
1497 newRhsInterval = rhsInterval;
1500 newLhsInterval = lhsInterval;
1501 newRhsInterval = rhsInterval;
1503 }
else if (ltCase()) {
1504 newLhsInterval = lhsInterval.toUnreduced().computeLTPart(rhsInterval.toUnreduced()).reduce(f);
1505 newRhsInterval = rhsInterval.toUnreduced().computeGEPart(lhsInterval.toUnreduced()).reduce(f);
1506 }
else if (leCase()) {
1507 newLhsInterval = lhsInterval.toUnreduced().computeLEPart(rhsInterval.toUnreduced()).reduce(f);
1508 newRhsInterval = rhsInterval.toUnreduced().computeGTPart(lhsInterval.toUnreduced()).reduce(f);
1509 }
else if (gtCase()) {
1510 newLhsInterval = lhsInterval.toUnreduced().computeGTPart(rhsInterval.toUnreduced()).reduce(f);
1511 newRhsInterval = rhsInterval.toUnreduced().computeLEPart(lhsInterval.toUnreduced()).reduce(f);
1512 }
else if (geCase()) {
1513 newLhsInterval = lhsInterval.toUnreduced().computeGEPart(rhsInterval.toUnreduced()).reduce(f);
1514 newRhsInterval = rhsInterval.toUnreduced().computeLTPart(lhsInterval.toUnreduced()).reduce(f);
1516 cmpOp->emitWarning(
"unhandled cmp predicate").report();
1521 applyInterval(cmpOp, lhs, newLhsInterval);
1522 applyInterval(cmpOp, rhs, newRhsInterval);
1530 auto mulCase = [&](MulFeltOp mulOp) {
1532 auto constCase = [&](FeltConstantOp constOperand, Value multiplicand) {
1534 APInt constVal = constOperand.getValue();
1535 if (constVal.isZero()) {
1540 applyInterval(mulOp, multiplicand, updatedInterval);
1543 Value lhs = mulOp.getLhs(), rhs = mulOp.getRhs();
1545 auto lhsConstOp = dyn_cast_if_present<FeltConstantOp>(lhs.getDefiningOp());
1546 auto rhsConstOp = dyn_cast_if_present<FeltConstantOp>(rhs.getDefiningOp());
1548 if (lhsConstOp && rhsConstOp) {
1550 }
else if (lhsConstOp) {
1551 constCase(lhsConstOp, rhs);
1553 }
else if (rhsConstOp) {
1554 constCase(rhsConstOp, lhs);
1560 if (newInterval.intersect(zeroInt).isNotEmpty()) {
1566 ExpressionValue lhsExpr = lhsLat->getValue().getScalarValue(),
1567 rhsExpr = rhsLat->getValue().getScalarValue();
1568 Interval newLhsInterval = lhsExpr.getInterval().difference(zeroInt);
1569 Interval newRhsInterval = rhsExpr.getInterval().difference(zeroInt);
1570 applyInterval(mulOp, lhs, newLhsInterval);
1571 applyInterval(mulOp, rhs, newRhsInterval);
1574 auto addCase = [&](AddFeltOp addOp) {
1575 Value lhs = addOp.getLhs(), rhs = addOp.getRhs();
1577 ExpressionValue lhsVal = lhsLat->getValue().getScalarValue();
1578 ExpressionValue rhsVal = rhsLat->getValue().getScalarValue();
1580 const Interval &currLhsInt = lhsVal.getInterval(), &currRhsInt = rhsVal.getInterval();
1582 Interval derivedLhsInt = newInterval - currRhsInt;
1583 Interval derivedRhsInt = newInterval - currLhsInt;
1585 Interval finalLhsInt = currLhsInt.intersect(derivedLhsInt);
1586 Interval finalRhsInt = currRhsInt.intersect(derivedRhsInt);
1588 applyInterval(addOp, lhs, finalLhsInt);
1589 applyInterval(addOp, rhs, finalRhsInt);
1592 auto subCase = [&](SubFeltOp subOp) {
1593 Value lhs = subOp.getLhs(), rhs = subOp.getRhs();
1595 ExpressionValue lhsVal = lhsLat->getValue().getScalarValue();
1596 ExpressionValue rhsVal = rhsLat->getValue().getScalarValue();
1598 const Interval &currLhsInt = lhsVal.getInterval(), &currRhsInt = rhsVal.getInterval();
1600 Interval derivedLhsInt = newInterval + currRhsInt;
1601 Interval derivedRhsInt = currLhsInt - newInterval;
1603 Interval finalLhsInt = currLhsInt.intersect(derivedLhsInt);
1604 Interval finalRhsInt = currRhsInt.intersect(derivedRhsInt);
1606 applyInterval(subOp, lhs, finalLhsInt);
1607 applyInterval(subOp, rhs, finalRhsInt);
1610 auto selectCase = [&](arith::SelectOp selectOp) {
1611 Value cond = selectOp.getCondition();
1612 Value trueVal = selectOp.getTrueValue();
1613 Value falseVal = selectOp.getFalseValue();
1619 const Interval &condInterval = condExpr.getInterval();
1620 if (condInterval.isDegenerate() && condInterval.rhs() == f.one()) {
1621 applyInterval(selectOp, trueVal, newInterval);
1624 if (condInterval.isDegenerate() && condInterval.rhs() == f.zero()) {
1625 applyInterval(selectOp, falseVal, newInterval);
1629 Interval trueOverlap = trueExpr.getInterval().intersect(newInterval);
1630 Interval falseOverlap = falseExpr.getInterval().intersect(newInterval);
1631 bool truePossible = trueOverlap.isNotEmpty();
1632 bool falsePossible = falseOverlap.isNotEmpty();
1634 if (truePossible && !falsePossible) {
1636 applyInterval(selectOp, trueVal, newInterval);
1639 if (!truePossible && falsePossible) {
1641 applyInterval(selectOp, falseVal, newInterval);
1644 if (!truePossible && !falsePossible) {
1649 auto readmCase = [&](MemberReadOp) {
1650 SourceRefLatticeValue sourceRefVal = getSourceRefState(val);
1652 if (sourceRefVal.isSingleValue()) {
1653 const SourceRef &ref = sourceRefVal.getSingleValue();
1654 readResults[ref].insert(valLattice);
1657 for (Lattice *l : readResults[ref]) {
1658 if (l != valLattice) {
1659 propagateIfChanged(l, l->setValue(newLatticeVal));
1665 auto readArrCase = [&](ReadArrayOp) {
1666 auto arrayRef = getArrayAccessRef(valUser, llvm::cast<ReadArrayOp>(definingOp));
1667 if (succeeded(arrayRef)) {
1668 readResults[*arrayRef].insert(valLattice);
1670 for (Lattice *l : readResults[*arrayRef]) {
1671 if (l != valLattice) {
1672 propagateIfChanged(l, l->setValue(newLatticeVal));
1677 SourceRefLatticeValue sourceRefVal = getSourceRefState(val);
1679 if (sourceRefVal.isSingleValue()) {
1680 const SourceRef &ref = sourceRefVal.getSingleValue();
1681 readResults[ref].insert(valLattice);
1684 for (Lattice *l : readResults[ref]) {
1685 if (l != valLattice) {
1686 propagateIfChanged(l, l->setValue(newLatticeVal));
1693 auto castCase = [&](Operation *op) { applyInterval(op, op->getOperand(0), newInterval); };
1699 TypeSwitch<Operation *>(definingOp)
1700 .Case<CmpOp>([&](
auto op) { cmpCase(op); })
1701 .Case<AddFeltOp>([&](
auto op) {
return addCase(op); })
1702 .Case<SubFeltOp>([&](
auto op) {
return subCase(op); })
1703 .Case<MulFeltOp>([&](
auto op) { mulCase(op); })
1704 .Case<arith::SelectOp>([&](
auto op) { selectCase(op); })
1705 .Case<MemberReadOp>([&](
auto op){ readmCase(op); })
1706 .Case<ReadArrayOp>([&](
auto op){ readArrCase(op); })
1707 .Case<IntToFeltOp, FeltToIndexOp>([&](
auto op) { castCase(op); })
1708 .Default([&](Operation *) { });
1712 propagateIfChanged(valLattice, changed);
1715FailureOr<std::pair<DenseSet<Value>,
Interval>>
1716IntervalDataFlowAnalysis::getGeneralizedDecompInterval(
1717 Operation * , Value lhs, Value rhs
1719 auto isZeroConst = [
this](Value v) {
1720 Operation *op = v.getDefiningOp();
1724 if (!isConstOp(op)) {
1727 return getConst(op) == field.get().zero();
1729 bool lhsIsZero = isZeroConst(lhs), rhsIsZero = isZeroConst(rhs);
1730 Value exprTree =
nullptr;
1731 if (lhsIsZero && !rhsIsZero) {
1733 }
else if (!lhsIsZero && rhsIsZero) {
1740 std::optional<SourceRef> signalRef = std::nullopt;
1741 DenseSet<Value> signalVals;
1742 SmallVector<DynamicAPInt> consts;
1743 SmallVector<Value> frontier {exprTree};
1744 while (!frontier.empty()) {
1745 Value v = frontier.back();
1746 frontier.pop_back();
1747 Operation *op = v.getDefiningOp();
1751 auto handleRefValue = [
this, &signalRef, &signalVal, &signalVals]() {
1752 SourceRefLatticeValue refSet = getSourceRefState(signalVal);
1753 if (!refSet.isScalar() || !refSet.isSingleValue()) {
1756 SourceRef r = refSet.getSingleValue();
1757 if (signalRef.has_value() && signalRef.value() != r) {
1759 }
else if (!signalRef.has_value()) {
1762 signalVals.insert(signalVal);
1767 if (op && matchPattern(op, subPattern)) {
1768 if (failed(handleRefValue())) {
1771 auto constInt = APSInt(c.getValue());
1772 consts.push_back(field.get().reduce(constInt));
1774 }
else if (
m_RefValue(&signalVal).match(v)) {
1775 if (failed(handleRefValue())) {
1778 consts.push_back(field.get().zero());
1784 if (op && matchPattern(op, mulPattern)) {
1785 frontier.push_back(a);
1786 frontier.push_back(b);
1795 std::sort(consts.begin(), consts.end());
1796 Interval iv = UnreducedInterval(consts.front(), consts.back()).reduce(field.get());
1797 return std::make_pair(std::move(signalVals), iv);
1805 SymbolTableCollection tables;
1807 auto computeIntervalsImpl =
1808 [&solver, &am, &ctx, &tables,
this](
1809 FuncDefOp fn, llvm::MapVector<SourceRef, Interval> &memberRanges,
1810 llvm::MapVector<SourceRef, UnreducedInterval> &memberUnreducedRanges,
1811 llvm::SetVector<ExpressionValue> &
1813 auto setUnreducedRange =
1815 memberUnreducedRanges.erase(ref);
1816 memberUnreducedRanges.insert({ref, interval});
1829 searchSet.insert(ref);
1833 auto mergeInterval = [&memberRanges, &memberUnreducedRanges](
1835 std::optional<UnreducedInterval> unreducedInterval = std::nullopt
1837 auto *existing = memberRanges.find(ref);
1838 if (existing != memberRanges.end()) {
1840 bool intervalChanged = mergedInterval != existing->second;
1841 existing->second = mergedInterval;
1843 if (unreducedInterval.has_value()) {
1844 auto *existingUnreduced = memberUnreducedRanges.find(ref);
1845 if (existingUnreduced != memberUnreducedRanges.end()) {
1846 existingUnreduced->second = existingUnreduced->second.intersect(*unreducedInterval);
1848 memberUnreducedRanges.insert({ref, *unreducedInterval});
1850 }
else if (intervalChanged) {
1851 memberUnreducedRanges.erase(ref);
1856 memberRanges[ref] = interval;
1857 if (unreducedInterval.has_value()) {
1858 memberUnreducedRanges.insert({ref, *unreducedInterval});
1863 for (BlockArgument arg : fn.getArguments()) {
1865 if (searchSet.erase(ref)) {
1869 if (!expr.getExpr()) {
1872 expr = expr.withUnreducedInterval(expr.getInterval().firstUnreduced());
1875 memberRanges[ref] = expr.getInterval();
1876 if (expr.hasUnreducedInterval()) {
1877 setUnreducedRange(ref, expr.getUnreducedInterval());
1879 assert(memberRanges[ref].getField() == ctx.
getField() &&
"bad interval defaults");
1887 if (!lattices.empty() && searchSet.erase(ref)) {
1889 std::optional<UnreducedInterval> joinedUnreduced = std::nullopt;
1890 bool sawFirst =
false;
1902 memberRanges[ref] = joinedInterval;
1903 if (joinedUnreduced.has_value()) {
1904 setUnreducedRange(ref, *joinedUnreduced);
1906 assert(memberRanges[ref].getField() == ctx.
getField() &&
"bad interval defaults");
1911 if (searchSet.erase(ref)) {
1912 memberRanges[ref] = val.getInterval();
1913 if (val.hasUnreducedInterval()) {
1914 setUnreducedRange(ref, val.getUnreducedInterval());
1916 assert(memberRanges[ref].getField() == ctx.
getField() &&
"bad interval defaults");
1923 if (fn.isStructConstrain()) {
1924 auto mergeChildConstrainIntervals = [&](
CallOp fnCall) {
1939 auto calledStruct = calledFn->getParentOfType<
StructDefOp>();
1940 if (calledStruct == structDef) {
1945 if (childAnalysis.inProgress(ctx)) {
1948 if (!childAnalysis.constructed(ctx)) {
1950 succeeded(childAnalysis.runAnalysis(solver, am, ctx)),
1951 "could not construct interval analysis for child struct"
1958 llvm::MapVector<SourceRef, Interval> callOperandIntervals;
1959 for (
unsigned i = 0; i < calledFn.getNumArguments(); i++) {
1960 SourceRef prefix(calledFn.getArgument(i));
1961 Value operand = fnCall.getOperand(i);
1962 std::optional<SourceRefLatticeValue> identityVal =
1963 getIdentitySourceRefState(solver, operand);
1964 if (identityVal.has_value()) {
1965 identityTranslations.push_back({prefix, *identityVal});
1968 if (!llvm::isa<ArrayType, StructType, pod::PodType>(operand.getType())) {
1971 if (lattice !=
nullptr) {
1973 callOperandIntervals[prefix] = expr.
getInterval();
1978 const StructIntervals &childIntervals = childAnalysis.getResult(ctx);
1981 for (
const auto &[childRef, childInterval] : constrainIntervals) {
1982 auto translatedRefs = translateRef(childRef, identityTranslations);
1983 if (failed(translatedRefs)) {
1987 std::optional<UnreducedInterval> childUnreduced = std::nullopt;
1988 if (
const auto *childUnreducedIt = constrainUnreducedIntervals.find(childRef);
1989 childUnreducedIt != constrainUnreducedIntervals.end()) {
1990 childUnreduced = childUnreducedIt->second;
1994 for (
const SourceRef &translatedRef : *translatedRefs) {
1995 uniqueTranslatedRefs.insert(translatedRef);
1997 if (uniqueTranslatedRefs.size() != 1) {
2001 const SourceRef &translatedRef = *uniqueTranslatedRefs.begin();
2002 if (functionRefs.contains(translatedRef)) {
2003 mergeInterval(translatedRef, childInterval, childUnreduced);
2004 searchSet.erase(translatedRef);
2010 llvm::EquivalenceClasses<SourceRef> directEqRefs =
2011 collectDirectEqualityRefs(solver, calledFn);
2012 for (
auto leaderIt = directEqRefs.begin(); leaderIt != directEqRefs.end(); ++leaderIt) {
2013 if (!leaderIt->isLeader()) {
2017 llvm::MapVector<SourceRef, Interval> translatedEqRefs;
2019 bool hasInterval =
false;
2020 bool ambiguousTranslation =
false;
2022 for (
auto memberIt = directEqRefs.member_begin(leaderIt);
2023 memberIt != directEqRefs.member_end(); ++memberIt) {
2025 if (
const auto *childIntervalIt = constrainIntervals.find(*memberIt);
2026 childIntervalIt != constrainIntervals.end()) {
2027 memberInterval = memberInterval.
intersect(childIntervalIt->second);
2029 if (
auto *callOperandIt = callOperandIntervals.find(*memberIt);
2030 callOperandIt != callOperandIntervals.end()) {
2031 memberInterval = memberInterval.
intersect(callOperandIt->second);
2032 contextualInterval = contextualInterval.
intersect(memberInterval);
2036 auto translatedRefs = translateRef(*memberIt, identityTranslations);
2037 if (failed(translatedRefs)) {
2042 for (
const SourceRef &translatedRef : *translatedRefs) {
2043 uniqueTranslatedRefs.insert(translatedRef);
2045 if (uniqueTranslatedRefs.size() != 1) {
2046 ambiguousTranslation =
true;
2050 const SourceRef &translatedRef = *uniqueTranslatedRefs.begin();
2051 if (!functionRefs.contains(translatedRef)) {
2055 if (
auto *parentIntervalIt = memberRanges.find(translatedRef);
2056 parentIntervalIt != memberRanges.end()) {
2057 memberInterval = memberInterval.
intersect(parentIntervalIt->second);
2060 translatedEqRefs[translatedRef] = memberInterval;
2061 contextualInterval = contextualInterval.
intersect(memberInterval);
2065 if (ambiguousTranslation || !hasInterval || translatedEqRefs.empty()) {
2069 for (
const auto &[translatedRef, _] : translatedEqRefs) {
2070 mergeInterval(translatedRef, contextualInterval);
2071 searchSet.erase(translatedRef);
2076 fn.walk(mergeChildConstrainIntervals);
2080 for (
const auto &ref : searchSet) {
2083 setUnreducedRange(ref, memberRanges[ref].firstUnreduced());
2092 llvm::SmallVector<std::pair<SourceRef, Interval>> sortedRanges;
2093 sortedRanges.reserve(memberRanges.size());
2094 for (
const auto &[ref, interval] : memberRanges) {
2095 sortedRanges.emplace_back(ref, interval);
2097 llvm::sort(sortedRanges, [](
const auto &a,
const auto &b) {
return a.first < b.first; });
2098 llvm::SmallVector<std::pair<SourceRef, UnreducedInterval>> sortedUnreducedRanges;
2099 sortedUnreducedRanges.reserve(memberUnreducedRanges.size());
2100 for (
const auto &[ref, interval] : memberUnreducedRanges) {
2101 sortedUnreducedRanges.emplace_back(ref, interval);
2103 llvm::sort(sortedUnreducedRanges, [](
const auto &a,
const auto &b) {
2104 return a.first < b.first;
2106 memberRanges.clear();
2107 memberUnreducedRanges.clear();
2108 for (
auto &[ref, interval] : sortedRanges) {
2109 memberRanges[ref] = interval;
2111 for (
auto &[ref, interval] : sortedUnreducedRanges) {
2112 memberUnreducedRanges.insert({ref, interval});
2116 if (
auto computeFn = structDef.getComputeFuncOp()) {
2117 computeIntervalsImpl(
2118 computeFn, computeMemberRanges, computeMemberUnreducedRanges, computeSolverConstraints
2121 if (
auto constrainFn = structDef.getConstrainFuncOp()) {
2122 computeIntervalsImpl(
2123 constrainFn, constrainMemberRanges, constrainMemberUnreducedRanges,
2124 constrainSolverConstraints
2132 mlir::raw_ostream &os,
bool withConstraints,
bool printCompute,
bool printUnreduced
2134 auto writeIntervals =
2135 [&os, &withConstraints, &printUnreduced](
2136 const char *fnName,
const llvm::MapVector<SourceRef, Interval> &memberRanges,
2137 const llvm::MapVector<SourceRef, UnreducedInterval> &memberUnreducedRanges,
2138 const llvm::SetVector<ExpressionValue> &solverConstraints,
bool printName
2143 os.indent(indent) << fnName <<
" {";
2147 if (memberRanges.empty()) {
2152 for (
const auto &[ref, interval] : memberRanges) {
2154 os.indent(indent) << ref <<
" in " << interval;
2155 if (printUnreduced) {
2156 const auto *unreducedIt = memberUnreducedRanges.find(ref);
2157 if (unreducedIt != memberUnreducedRanges.end()) {
2158 os <<
" ( " << unreducedIt->second <<
" )";
2163 if (withConstraints) {
2165 os.indent(indent) <<
"Solver Constraints { ";
2166 if (solverConstraints.empty()) {
2169 for (
const auto &e : solverConstraints) {
2171 os.indent(indent + 4);
2172 e.getExpr()->print(os);
2175 os.indent(indent) <<
'}';
2181 os.indent(indent - 4) <<
'}';
2185 os <<
"StructIntervals { ";
2186 if (constrainMemberRanges.empty() && (!printCompute || computeMemberRanges.empty())) {
2194 computeSolverConstraints, printCompute
2199 constrainSolverConstraints, printCompute
Tracks a solver expression and an interval range for that expression.
ExpressionValue withUnreducedInterval(const UnreducedInterval &newUnreducedInterval) const
ExpressionValue withExpression(const llvm::SMTExprRef &newExpr) const
Return the current expression with a new SMT expression.
const Interval & getInterval() const
const std::optional< UnreducedInterval > & getOptionalUnreducedInterval() const
ExpressionValue withOptionalUnreducedInterval(std::optional< UnreducedInterval > newUnreducedInterval) const
ExpressionValue withInterval(const Interval &newInterval) const
Return the current expression with a new interval.
void print(mlir::raw_ostream &os) const
bool operator==(const ExpressionValue &rhs) const
llvm::SMTExprRef getExpr() const
bool isBoolSort(const llvm::SMTSolverRef &solver) const
bool hasUnreducedInterval() const
const Field & getField() const
const UnreducedInterval & getUnreducedInterval() const
Information about the prime finite field used for the interval analysis.
llvm::DynamicAPInt zero() const
Returns 0 at the bitwidth of the field.
llvm::DynamicAPInt prime() const
For the prime field p, returns p.
llvm::DynamicAPInt one() const
Returns 1 at the bitwidth of the field.
llvm::DynamicAPInt inv(const llvm::DynamicAPInt &i) const
Returns the multiplicative inverse of i in prime field p.
unsigned bitWidth() const
llvm::SMTExprRef createSymbol(const llvm::SMTSolverRef &solver, const char *name) const
Create a SMT solver symbol with the current field's bitwidth.
llvm::DynamicAPInt maxVal() const
Returns p - 1, which is the max value possible in a prime field described by p.
const LatticeValue & getValue() const
mlir::ChangeResult setValue(const LatticeValue &val)
IntervalAnalysisLatticeValue LatticeValue
mlir::ChangeResult meet(const AbstractSparseLattice &other) override
void print(mlir::raw_ostream &os) const override
mlir::ChangeResult join(const AbstractSparseLattice &other) override
mlir::ChangeResult addSolverConstraint(const ExpressionValue &e)
mlir::LogicalResult visitOperation(mlir::Operation *op, mlir::ArrayRef< const Lattice * > operands, mlir::ArrayRef< Lattice * > results) override
Visit an operation with the lattices of its operands.
llvm::SMTExprRef getOrCreateSymbol(const SourceRef &r)
Either return the existing SMT expression that corresponds to the SourceRef, or create one.
const llvm::DenseMap< SourceRef, ExpressionValue > & getWriteResults() const
const llvm::DenseMap< SourceRef, llvm::DenseSet< Lattice * > > & getReadResults() const
Intervals over a finite field.
static Interval True(const Field &f)
llvm::DynamicAPInt rhs() const
Interval intersect(const Interval &rhs) const
Intersect.
UnreducedInterval toUnreduced() const
Convert to an UnreducedInterval.
static Interval Boolean(const Field &f)
UnreducedInterval firstUnreduced() const
Get the first side of the interval for TypeF intervals, otherwise just get the full interval as an Un...
static Interval Entire(const Field &f)
bool isDegenerate() const
static Interval False(const Field &f)
llvm::DynamicAPInt lhs() const
Interval join(const Interval &rhs) const
Union.
static SourceRefLatticeValue getValueState(mlir::DataFlowSolver &solver, mlir::Value val)
Defines an index into an LLZK object.
A value at a given point of the SourceRefLattice.
mlir::FailureOr< std::pair< SourceRefLatticeValue, mlir::ChangeResult > > referencePodRecord(mlir::StringAttr recordName) const
Add the given pod recordName to the SourceRefs contained within this value.
const SourceRef & getSingleValue() const
mlir::FailureOr< std::pair< SourceRefLatticeValue, mlir::ChangeResult > > extract(const std::vector< SourceRefIndex > &indices) const
Perform an array.extract or array.read operation, depending on how many indices are provided.
A reference to a "source", which is the base value from which other SSA values are derived.
mlir::FailureOr< SourceRef > createChild(const SourceRefIndex &r) const
std::vector< SourceRefIndex > Path
static std::vector< SourceRef > getAllSourceRefs(mlir::SymbolTableCollection &tables, mlir::ModuleOp mod, const SourceRef &root)
Produce all possible SourceRefs that are present starting from the given root.
mlir::Type getType() const
const llvm::MapVector< SourceRef, Interval > & getConstrainIntervals() const
const llvm::MapVector< SourceRef, UnreducedInterval > & getConstrainUnreducedIntervals() const
void print(mlir::raw_ostream &os, bool withConstraints=false, bool printCompute=false, bool printUnreduced=false) const
mlir::LogicalResult computeIntervals(mlir::DataFlowSolver &solver, mlir::AnalysisManager &am, const IntervalAnalysisContext &ctx)
An inclusive interval [a, b] where a and b are arbitrary integers not necessarily bound to a given fi...
UnreducedInterval computeLTPart(const UnreducedInterval &rhs) const
Return the part of the interval that is guaranteed to be less than the rhs's max value.
UnreducedInterval computeGEPart(const UnreducedInterval &rhs) const
Return the part of the interval that is greater than or equal to the rhs's lower bound.
UnreducedInterval computeGTPart(const UnreducedInterval &rhs) const
Return the part of the interval that is greater than the rhs's lower bound.
Interval reduce(const Field &field) const
Reduce the interval to an interval in the given field.
UnreducedInterval computeLEPart(const UnreducedInterval &rhs) const
Return the part of the interval that is less than or equal to the rhs's upper bound.
Helper for converting between linear and multi-dimensional indexing with checks to ensure indices are...
static ArrayIndexGen from(ArrayType)
Construct new ArrayIndexGen. Will assert if hasStaticShape() is false.
std::optional< llvm::SmallVector< mlir::Value > > delinearize(int64_t, mlir::Location, mlir::OpBuilder &) const
::mlir::Type getElementType() const
::llzk::boolean::FeltCmpPredicate getPredicate()
std::variant< ScalarTy, ArrayTy > & getValue()
bool isSingleValue() const
const ScalarTy & getScalarValue() const
IntervalAnalysisLattice * getLatticeElement(mlir::Value value) override
bool isStructConstrain()
Return true iff the function is within a StructDefOp and named FUNC_NAME_CONSTRAIN.
bool isOperationLive(DataFlowSolver &solver, Operation *op)
ExpressionValue boolNot(const llvm::SMTSolverRef &solver, const ExpressionValue &val)
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.
ExpressionValue sintMod(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
RefValueCapture m_RefValue()
ExpressionValue intersection(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
FailureOr< Interval > signedIntDiv(const Interval &lhs, const Interval &rhs)
Computes signed integer division with possibly non-Degenerate divisors.
std::vector< std::pair< SourceRef, SourceRefLatticeValue > > SourceRefRemappings
ExpressionValue mod(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue shiftLeft(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue fallbackUnaryOp(const llvm::SMTSolverRef &solver, Operation *op, const ExpressionValue &val)
constexpr char FUNC_NAME_CONSTRAIN[]
Interval signedMod(const Interval &lhs, const Interval &rhs)
Computes signed integer remainder with possibly non-Degenerate divisors.
void ensure(bool condition, const llvm::Twine &errMsg)
ExpressionValue boolXor(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue cmp(const llvm::SMTSolverRef &solver, CmpOp op, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue neg(const llvm::SMTSolverRef &solver, const ExpressionValue &val)
DynamicAPInt toDynamicAPInt(StringRef str)
llvm::SMTExprRef createFieldInverseExpr(const llvm::SMTSolverRef &solver, Operation *op, const ExpressionValue &val, StringRef suffix="")
ExpressionValue sintDiv(const llvm::SMTSolverRef &solver, Operation *op, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue boolAnd(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
FailureOr< Interval > unsignedIntDiv(const Interval &lhs, const Interval &rhs)
Computes unsigned integer division with possibly non-Degenerate divisors.
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)
std::string buildStringViaPrint(const T &base, Args &&...args)
Generate a string by calling base.print(llvm::raw_ostream &) on a stream backed by the returned strin...
ExpressionValue boolToFelt(const llvm::SMTSolverRef &solver, const ExpressionValue &expr, unsigned bitwidth)
mlir::FailureOr< SymbolLookupResult< T > > resolveCallableSilently(mlir::SymbolTableCollection &symbolTable, mlir::CallOpInterface call)
Resolve a callable without emitting a diagnostic for missing top-level symbols.
ConstantCapture m_Constant()
std::string buildStringViaInsertionOp(Args &&...args)
Generate a string by using the insertion operator (<<) to append all args to a stream backed by the r...
ExpressionValue bitOr(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue uintDiv(const llvm::SMTSolverRef &solver, Operation *op, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue bitAnd(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue shiftRight(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
APSInt toAPSInt(const DynamicAPInt &i)
ExpressionValue sub(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
auto m_CommutativeOp(LhsMatcher lhs, RhsMatcher rhs)
ExpressionValue bitXor(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue notOp(const llvm::SMTSolverRef &solver, const ExpressionValue &val)
FailureOr< Interval > feltDiv(const Interval &lhs, const Interval &rhs)
Computes finite-field division by multiplying the dividend by the multiplicative inverse of the divis...
ExpressionValue boolOr(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
ExpressionValue selectValue(const llvm::SMTSolverRef &solver, const ExpressionValue &cond, const ExpressionValue &trueVal, const ExpressionValue &falseVal)
Parameters and shared objects to pass to child analyses.
const Field & getField() const
IntervalDataFlowAnalysis * intervalDFA
bool doTrackUnreducedIntervals() const