35#include <mlir/Conversion/AffineToStandard/AffineToStandard.h>
36#include <mlir/Dialect/Arith/IR/Arith.h>
37#include <mlir/Dialect/ControlFlow/IR/ControlFlowOps.h>
38#include <mlir/Dialect/Func/IR/FuncOps.h>
39#include <mlir/Dialect/LLVMIR/LLVMDialect.h>
40#include <mlir/Dialect/MemRef/IR/MemRef.h>
41#include <mlir/Dialect/SCF/IR/SCF.h>
42#include <mlir/Dialect/Utils/IndexingUtils.h>
43#include <mlir/IR/Builders.h>
44#include <mlir/IR/BuiltinAttributes.h>
45#include <mlir/IR/BuiltinOps.h>
46#include <mlir/IR/SymbolTable.h>
47#include <mlir/Pass/PassManager.h>
48#include <mlir/Transforms/Passes.h>
50#include <llvm/ADT/APInt.h>
51#include <llvm/ADT/STLExtras.h>
52#include <llvm/ADT/SmallString.h>
53#include <llvm/ADT/StringMap.h>
54#include <llvm/ADT/TypeSwitch.h>
55#include <llvm/Support/MathExtras.h>
67 llvm::SmallVector<Value> leaves;
71static FailureOr<std::reference_wrapper<const Field>> getModuleField(ModuleOp moduleOp) {
73 if (failed(
collectFields(moduleOp.getOperation(), fields,
false))) {
74 moduleOp.emitError(
"failed to collect fields for llzk-witgen lowering");
77 if (fields.size() != 1) {
78 moduleOp.emitError(
"llzk-witgen execution-engine lowering requires exactly one field");
81 return *fields.begin();
85static std::string mangleFunctionName(function::FuncDefOp funcOp) {
86 auto symbolRef = funcOp.getFullyQualifiedName(
false);
87 llvm::SmallString<128> result(
"__llzk_witgen_");
88 for (StringRef piece :
getNames(symbolRef)) {
89 if (!result.empty() && result.back() !=
'_') {
92 for (
char c : piece) {
93 result += llvm::isAlnum(c) ? c :
'_';
96 return std::string(result);
100static Value makeIndexConstant(OpBuilder &builder, Location loc, int64_t value) {
101 return builder.create<arith::ConstantIndexOp>(loc, value).getResult();
105static Value makeOneFelt(OpBuilder &builder, Location loc,
const Field &field) {
106 return builder.create<arith::ConstantOp>(
107 loc, IntegerAttr::get(IntegerType::get(builder.getContext(), field.bitWidth()), 1)
112static FailureOr<Type> lowerScalarType(MLIRContext *context, Type type,
const Field &field) {
113 if (isa<felt::FeltType>(type)) {
114 return IntegerType::get(context, field.bitWidth());
116 if (isa<IndexType>(type)) {
119 if (
auto intType = dyn_cast<IntegerType>(type)) {
120 if (intType.getWidth() == 1) {
128static bool isScalarType(Type type) {
129 return isa<felt::FeltType, IndexType>(type) ||
130 (isa<IntegerType>(type) && llvm::cast<IntegerType>(type).getWidth() == 1);
134static LogicalResult flattenTypeLeaves(
135 Type type, SymbolTableCollection &tables, Operation *origin,
const Field &field,
136 SmallVectorImpl<Type> &out, llvm::ArrayRef<int64_t> prefixShape = {},
bool storage = false
138 auto emitScalarLeaf = [&](Type leafType) {
139 auto lowered = lowerScalarType(origin->getContext(), leafType, field);
140 if (failed(lowered)) {
143 if (!storage && prefixShape.empty()) {
144 out.push_back(*lowered);
147 llvm::SmallVector<int64_t> shape(prefixShape.begin(), prefixShape.end());
151 out.push_back(MemRefType::get(shape, *lowered));
155 if (isScalarType(type)) {
156 return emitScalarLeaf(type);
159 if (
auto arrayType = dyn_cast<array::ArrayType>(type)) {
160 llvm::SmallVector<int64_t> newPrefix(prefixShape.begin(), prefixShape.end());
161 newPrefix.append(arrayType.getShape().begin(), arrayType.getShape().end());
162 return flattenTypeLeaves(
163 arrayType.getElementType(), tables, origin, field, out, newPrefix,
true
167 if (
auto podType = dyn_cast<pod::PodType>(type)) {
168 for (pod::RecordAttr record : podType.getRecords()) {
170 flattenTypeLeaves(record.getType(), tables, origin, field, out, prefixShape,
true)
178 if (
auto structType = dyn_cast<component::StructType>(type)) {
179 auto def = structType.getDefinition(tables, origin);
181 origin->emitError(
"could not resolve struct type during witgen lowering");
184 for (component::MemberDefOp member : def->get().getMemberDefs()) {
186 flattenTypeLeaves(member.getType(), tables, origin, field, out, prefixShape,
true)
194 origin->emitError(
"unsupported type in llzk-witgen lowering: ") << type;
200getStridedMemRefType(MLIRContext *context, ArrayRef<int64_t> shape, Type elementType) {
201 SmallVector<int64_t> strides(shape.size(), ShapedType::kDynamic);
202 return MemRefType::get(
203 shape, elementType, StridedLayoutAttr::get(context, ShapedType::kDynamic, strides)
208static LogicalResult flattenABILeafTypes(
209 Type type, SymbolTableCollection &tables, Operation *origin,
const Field &field,
210 SmallVectorImpl<Type> &out,
size_t prefixRank = 0,
bool aggregateStorage =
false
212 auto emitScalarLeaf = [&](Type leafType) {
213 auto lowered = lowerScalarType(origin->getContext(), leafType, field);
214 if (failed(lowered)) {
217 if (!aggregateStorage && prefixRank == 0) {
218 out.push_back(*lowered);
221 SmallVector<int64_t> shape;
222 if (prefixRank == 0) {
225 shape.assign(prefixRank, ShapedType::kDynamic);
227 out.push_back(getStridedMemRefType(origin->getContext(), shape, *lowered));
231 if (isScalarType(type)) {
232 return emitScalarLeaf(type);
235 if (
auto arrayType = dyn_cast<array::ArrayType>(type)) {
236 return flattenABILeafTypes(
237 arrayType.getElementType(), tables, origin, field, out, prefixRank + arrayType.getRank(),
242 if (
auto podType = dyn_cast<pod::PodType>(type)) {
243 for (pod::RecordAttr record : podType.getRecords()) {
245 flattenABILeafTypes(record.getType(), tables, origin, field, out, prefixRank,
true)
253 if (
auto structType = dyn_cast<component::StructType>(type)) {
254 auto def = structType.getDefinition(tables, origin);
256 origin->emitError(
"could not resolve struct type during witgen lowering");
259 for (component::MemberDefOp member : def->get().getMemberDefs()) {
261 flattenABILeafTypes(member.getType(), tables, origin, field, out, prefixRank,
true)
269 origin->emitError(
"unsupported type in llzk-witgen lowering: ") << type;
274static FailureOr<size_t>
275getLeafCount(Type type, SymbolTableCollection &tables, Operation *origin,
const Field &field) {
276 SmallVector<Type> leaves;
277 if (failed(flattenTypeLeaves(type, tables, origin, field, leaves))) {
280 return leaves.size();
284static FailureOr<SmallVector<Type>>
285getLeafTypes(Type type, SymbolTableCollection &tables, Operation *origin,
const Field &field) {
286 SmallVector<Type> leaves;
287 if (failed(flattenTypeLeaves(type, tables, origin, field, leaves))) {
294static FailureOr<SmallVector<Type>>
295getABILeafTypes(Type type, SymbolTableCollection &tables, Operation *origin,
const Field &field) {
296 SmallVector<Type> leaves;
297 if (failed(flattenABILeafTypes(type, tables, origin, field, leaves))) {
304static FailureOr<std::pair<size_t, size_t>> getNamedLeafSpan(
305 Type ownerType, StringRef name, SymbolTableCollection &tables, Operation *origin,
308 if (
auto podType = dyn_cast<pod::PodType>(ownerType)) {
310 for (pod::RecordAttr record : podType.getRecords()) {
311 auto count = getLeafCount(record.getType(), tables, origin, field);
315 if (record.getName().getValue() == name) {
316 return std::pair<size_t, size_t> {running, *count};
322 if (
auto structType = dyn_cast<component::StructType>(ownerType)) {
323 auto def = structType.getDefinition(tables, origin);
325 origin->emitError(
"could not resolve struct type during witgen lowering");
329 for (component::MemberDefOp member : def->get().getMemberDefs()) {
330 auto count = getLeafCount(member.getType(), tables, origin, field);
334 if (member.getSymName() == name) {
335 return std::pair<size_t, size_t> {running, *count};
341 origin->emitError(
"could not resolve aggregate member/record @") << name;
346static FailureOr<Type>
347getNamedSubType(Type ownerType, StringRef name, SymbolTableCollection &tables, Operation *origin) {
348 if (
auto podType = dyn_cast<pod::PodType>(ownerType)) {
349 for (pod::RecordAttr record : podType.getRecords()) {
350 if (record.getName().getValue() == name) {
351 return record.getType();
355 if (
auto structType = dyn_cast<component::StructType>(ownerType)) {
356 auto def = structType.getDefinition(tables, origin);
358 origin->emitError(
"could not resolve struct type during witgen lowering");
361 for (component::MemberDefOp member : def->get().getMemberDefs()) {
362 if (member.getSymName() == name) {
363 return member.getType();
367 origin->emitError(
"could not resolve aggregate member/record @") << name;
372static FailureOr<Value> createZeroMemRef(OpBuilder &builder, Location loc, MemRefType memrefType) {
375 emitError(loc) << llvm::toString(elementCount.takeError());
378 Value alloc = builder.create<memref::AllocOp>(loc, memrefType);
379 auto elementType = memrefType.getElementType();
381 if (isa<IndexType>(elementType)) {
382 zero = builder.create<arith::ConstantIndexOp>(loc, 0);
384 zero = builder.create<arith::ConstantOp>(
385 loc, IntegerAttr::get(llvm::cast<IntegerType>(elementType), 0)
388 auto strides = mlir::computeStrides(memrefType.getShape());
389 for (
size_t flat = 0; flat < *elementCount; ++flat) {
392 emitError(loc) << llvm::toString(flatSigned.takeError());
395 SmallVector<Value> indices;
396 for (int64_t index : mlir::delinearize(*flatSigned, strides)) {
397 indices.push_back(makeIndexConstant(builder, loc, index));
399 builder.create<memref::StoreOp>(loc, zero, alloc, indices);
405static FailureOr<Value> createRandomMemRef(
406 OpBuilder &builder, Location loc, MemRefType memrefType,
const Field &field,
411 emitError(loc) << llvm::toString(elementCount.takeError());
414 Value alloc = builder.create<memref::AllocOp>(loc, memrefType);
415 auto elementType = memrefType.getElementType();
416 auto strides = mlir::computeStrides(memrefType.getShape());
417 for (
size_t flat = 0; flat < *elementCount; ++flat) {
420 emitError(loc) << llvm::toString(flatSigned.takeError());
423 SmallVector<Value> indices;
424 for (int64_t index : mlir::delinearize(*flatSigned, strides)) {
425 indices.push_back(makeIndexConstant(builder, loc, index));
427 if (isa<IndexType>(elementType)) {
429 builder.create<memref::StoreOp>(
430 loc, builder.create<arith::ConstantIndexOp>(loc, value), alloc, indices
434 auto intType = llvm::cast<IntegerType>(elementType);
435 if (intType.getWidth() == 1) {
436 builder.create<memref::StoreOp>(
438 builder.create<arith::ConstantOp>(
446 builder.create<memref::StoreOp>(
448 builder.create<arith::ConstantOp>(
458static FailureOr<LoweredValue> createDefaultValue(
459 OpBuilder &builder, Location loc, Type type, SymbolTableCollection &tables, Operation *origin,
462 LoweredValue lowered {type, {}};
463 auto leafTypes = getLeafTypes(type, tables, origin, field);
464 if (failed(leafTypes)) {
467 for (Type leafType : *leafTypes) {
470 "fail-mode default materialization is unsupported in witgen lowering because it would "
471 "hide uninitialized reads"
476 if (
auto memrefType = dyn_cast<MemRefType>(leafType)) {
477 auto randomMemRef = createRandomMemRef(builder, loc, memrefType, field, rng);
478 if (failed(randomMemRef)) {
481 lowered.leaves.push_back(*randomMemRef);
484 if (isa<IndexType>(leafType)) {
485 lowered.leaves.push_back(
490 auto intType = llvm::cast<IntegerType>(leafType);
491 if (intType.getWidth() == 1) {
492 lowered.leaves.push_back(builder.create<arith::ConstantOp>(
498 lowered.leaves.push_back(builder.create<arith::ConstantOp>(
503 if (
auto memrefType = dyn_cast<MemRefType>(leafType)) {
504 auto zeroMemRef = createZeroMemRef(builder, loc, memrefType);
505 if (failed(zeroMemRef)) {
508 lowered.leaves.push_back(*zeroMemRef);
511 if (isa<IndexType>(leafType)) {
512 lowered.leaves.push_back(builder.create<arith::ConstantIndexOp>(loc, 0));
515 lowered.leaves.push_back(builder.create<arith::ConstantOp>(
516 loc, IntegerAttr::get(llvm::cast<IntegerType>(leafType), 0)
523static Value normalizeWideValue(
524 OpBuilder &builder, Location loc, Value wideValue,
unsigned dstWidth,
const Field &field
526 auto wideType = llvm::cast<IntegerType>(wideValue.getType());
527 Value modulus = builder.create<arith::ConstantOp>(
528 loc, field.getPrimeAttr(builder.getContext(), wideType.getWidth())
530 Value reduced = builder.create<arith::RemUIOp>(loc, wideValue, modulus);
531 return builder.create<arith::TruncIOp>(
532 loc, IntegerType::get(builder.getContext(), dstWidth), reduced
537static Value normalizeSignedWideValue(
538 OpBuilder &builder, Location loc, Value wideValue,
unsigned dstWidth,
const Field &field
540 auto wideType = llvm::cast<IntegerType>(wideValue.getType());
541 Value modulus = builder.create<arith::ConstantOp>(
542 loc, field.getPrimeAttr(builder.getContext(), wideType.getWidth())
544 Value reduced = builder.create<arith::RemSIOp>(loc, wideValue, modulus);
545 Value zero = builder.create<arith::ConstantOp>(loc, IntegerAttr::get(wideType, 0));
546 Value isNegative = builder.create<arith::CmpIOp>(loc, arith::CmpIPredicate::slt, reduced, zero);
547 Value adjusted = builder.create<arith::AddIOp>(loc, reduced, modulus);
548 Value canonical = builder.create<arith::SelectOp>(loc, isNegative, adjusted, reduced);
549 return builder.create<arith::TruncIOp>(
550 loc, IntegerType::get(builder.getContext(), dstWidth), canonical
556lowerFeltToSignedWide(OpBuilder &builder, Location loc, Value operand,
const Field &field) {
557 unsigned width = field.bitWidth();
558 unsigned wideWidth = width + 1;
559 auto feltType = IntegerType::get(builder.getContext(), width);
560 auto wideType = IntegerType::get(builder.getContext(), wideWidth);
561 Value operandWide = builder.create<arith::ExtUIOp>(loc, wideType, operand);
563 builder.create<arith::ConstantOp>(loc, field.getPrimeAttr(builder.getContext(), wideWidth));
564 Value half = builder.create<arith::ConstantOp>(
567 Value isNegative = builder.create<arith::CmpIOp>(loc, arith::CmpIPredicate::uge, operand, half);
568 Value signedOperand = builder.create<arith::SubIOp>(loc, operandWide, prime);
569 return builder.create<arith::SelectOp>(loc, isNegative, signedOperand, operandWide);
573static void assertNonZeroFelt(OpBuilder &builder, Location loc, Value operand, StringRef message) {
574 auto operandType = llvm::cast<IntegerType>(operand.getType());
575 Value zero = builder.create<arith::ConstantOp>(loc, IntegerAttr::get(operandType, 0));
576 Value isNonZero = builder.create<arith::CmpIOp>(loc, arith::CmpIPredicate::ne, operand, zero);
577 builder.create<cf::AssertOp>(loc, isNonZero, message);
582lowerFeltAdd(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
583 unsigned width = field.bitWidth();
584 unsigned wideWidth = width + 1;
585 auto wideType = IntegerType::get(builder.getContext(), wideWidth);
586 Value lhsWide = builder.create<arith::ExtUIOp>(loc, wideType, lhs);
587 Value rhsWide = builder.create<arith::ExtUIOp>(loc, wideType, rhs);
588 Value sum = builder.create<arith::AddIOp>(loc, lhsWide, rhsWide);
589 return normalizeWideValue(builder, loc, sum, width, field);
594lowerFeltSub(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
595 unsigned width = field.bitWidth();
596 unsigned wideWidth = width + 1;
597 auto wideType = IntegerType::get(builder.getContext(), wideWidth);
598 Value lhsWide = builder.create<arith::ExtUIOp>(loc, wideType, lhs);
599 Value rhsWide = builder.create<arith::ExtUIOp>(loc, wideType, rhs);
601 builder.create<arith::ConstantOp>(loc, field.getPrimeAttr(builder.getContext(), wideWidth));
602 Value lhsPlusMod = builder.create<arith::AddIOp>(loc, lhsWide, modulus);
603 Value diff = builder.create<arith::SubIOp>(loc, lhsPlusMod, rhsWide);
604 return normalizeWideValue(builder, loc, diff, width, field);
608static Value lowerFeltNeg(OpBuilder &builder, Location loc, Value operand,
const Field &field) {
609 unsigned width = field.bitWidth();
610 unsigned wideWidth = width + 1;
611 auto wideType = IntegerType::get(builder.getContext(), wideWidth);
612 Value operandWide = builder.create<arith::ExtUIOp>(loc, wideType, operand);
614 builder.create<arith::ConstantOp>(loc, field.getPrimeAttr(builder.getContext(), wideWidth));
615 Value diff = builder.create<arith::SubIOp>(loc, modulus, operandWide);
616 return normalizeWideValue(builder, loc, diff, width, field);
621lowerFeltMul(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
622 unsigned width = field.bitWidth();
623 unsigned wideWidth = width * 2;
624 auto wideType = IntegerType::get(builder.getContext(), wideWidth);
625 Value lhsWide = builder.create<arith::ExtUIOp>(loc, wideType, lhs);
626 Value rhsWide = builder.create<arith::ExtUIOp>(loc, wideType, rhs);
627 Value product = builder.create<arith::MulIOp>(loc, lhsWide, rhsWide);
628 return normalizeWideValue(builder, loc, product, width, field);
632static Value lowerFeltInv(OpBuilder &builder, Location loc, Value operand,
const Field &field) {
634 Value result = makeOneFelt(builder, loc, field);
635 Value base = operand;
636 for (
unsigned bit = 0; bit < exponent.getBitWidth(); ++bit) {
638 result = lowerFeltMul(builder, loc, result, base, field);
640 if (bit + 1 < exponent.getBitWidth()) {
641 base = lowerFeltMul(builder, loc, base, base, field);
649lowerFeltDiv(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
650 return lowerFeltMul(builder, loc, lhs, lowerFeltInv(builder, loc, rhs, field), field);
655lowerFeltPow(OpBuilder &builder, Location loc, Value base, Value exponent,
const Field &field) {
656 auto feltType = IntegerType::get(builder.getContext(), field.bitWidth());
657 Value zero = builder.create<arith::ConstantOp>(loc, IntegerAttr::get(feltType, 0));
658 Value one = builder.create<arith::ConstantOp>(loc, IntegerAttr::get(feltType, 1));
659 Value result = makeOneFelt(builder, loc, field);
660 Value currentBase = base;
661 for (
unsigned bit = 0; bit < field.bitWidth(); ++bit) {
662 Value bitIndex = builder.create<arith::ConstantOp>(
663 loc, IntegerAttr::get(feltType, llvm::APInt(field.bitWidth(), bit))
665 Value shifted = builder.create<arith::ShRUIOp>(loc, exponent, bitIndex);
666 Value masked = builder.create<arith::AndIOp>(loc, shifted, one);
667 Value bitIsSet = builder.create<arith::CmpIOp>(loc, arith::CmpIPredicate::ne, masked, zero);
668 auto ifOp = builder.create<scf::IfOp>(loc, TypeRange {feltType}, bitIsSet,
true);
670 OpBuilder::InsertionGuard guard(builder);
671 builder.setInsertionPointToStart(&ifOp.getThenRegion().front());
672 Value multiplied = lowerFeltMul(builder, loc, result, currentBase, field);
673 builder.create<scf::YieldOp>(loc, multiplied);
676 OpBuilder::InsertionGuard guard(builder);
677 builder.setInsertionPointToStart(&ifOp.getElseRegion().front());
678 builder.create<scf::YieldOp>(loc, result);
680 result = ifOp.getResult(0);
681 if (bit + 1 < field.bitWidth()) {
682 currentBase = lowerFeltMul(builder, loc, currentBase, currentBase, field);
690lowerFeltShl(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
691 auto feltType = IntegerType::get(builder.getContext(), field.bitWidth());
692 Value two = builder.create<arith::ConstantOp>(loc, IntegerAttr::get(feltType, 2));
693 return lowerFeltMul(builder, loc, lhs, lowerFeltPow(builder, loc, two, rhs, field), field);
698lowerFeltOr(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
699 unsigned width = field.bitWidth();
700 auto wideType = IntegerType::get(builder.getContext(), width + 1);
701 Value orValue = builder.create<arith::OrIOp>(loc, lhs, rhs);
702 Value orWide = builder.create<arith::ExtUIOp>(loc, wideType, orValue);
703 return normalizeWideValue(builder, loc, orWide, width, field);
708lowerFeltXor(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
709 unsigned width = field.bitWidth();
710 auto wideType = IntegerType::get(builder.getContext(), width + 1);
711 Value xorValue = builder.create<arith::XOrIOp>(loc, lhs, rhs);
712 Value xorWide = builder.create<arith::ExtUIOp>(loc, wideType, xorValue);
713 return normalizeWideValue(builder, loc, xorWide, width, field);
717static Value lowerFeltUnsignedDiv(OpBuilder &builder, Location loc, Value lhs, Value rhs) {
718 return builder.create<arith::DivUIOp>(loc, lhs, rhs);
723lowerFeltSignedDiv(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
724 unsigned width = field.bitWidth();
725 Value lhsSigned = lowerFeltToSignedWide(builder, loc, lhs, field);
726 Value rhsSigned = lowerFeltToSignedWide(builder, loc, rhs, field);
727 Value quotient = builder.create<arith::DivSIOp>(loc, lhsSigned, rhsSigned);
728 return normalizeSignedWideValue(builder, loc, quotient, width, field);
732static Value lowerFeltUnsignedMod(OpBuilder &builder, Location loc, Value lhs, Value rhs) {
733 return builder.create<arith::RemUIOp>(loc, lhs, rhs);
738lowerFeltSignedMod(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
739 unsigned width = field.bitWidth();
740 Value lhsSigned = lowerFeltToSignedWide(builder, loc, lhs, field);
741 Value rhsSigned = lowerFeltToSignedWide(builder, loc, rhs, field);
742 Value remainder = builder.create<arith::RemSIOp>(loc, lhsSigned, rhsSigned);
743 return normalizeSignedWideValue(builder, loc, remainder, width, field);
748lowerFeltShr(OpBuilder &builder, Location loc, Value lhs, Value rhs,
const Field &field) {
749 auto feltType = IntegerType::get(builder.getContext(), field.bitWidth());
750 Value width = builder.create<arith::ConstantOp>(
751 loc, IntegerAttr::get(feltType, llvm::APInt(field.bitWidth(), field.bitWidth()))
753 Value shiftTooLarge = builder.create<arith::CmpIOp>(loc, arith::CmpIPredicate::uge, rhs, width);
754 Value zero = builder.create<arith::ConstantOp>(loc, IntegerAttr::get(feltType, 0));
755 Value maxValidShift = builder.create<arith::ConstantOp>(
756 loc, IntegerAttr::get(feltType, llvm::APInt(field.bitWidth(), field.bitWidth() - 1))
758 Value clampedShift = builder.create<arith::MinUIOp>(loc, rhs, maxValidShift);
759 Value shifted = builder.create<arith::ShRUIOp>(loc, lhs, clampedShift);
760 return builder.create<arith::SelectOp>(loc, shiftTooLarge, zero, shifted);
764static Value lowerFeltNot(OpBuilder &builder, Location loc, Value operand,
const Field &field) {
765 unsigned width = field.bitWidth();
766 auto feltType = IntegerType::get(builder.getContext(), width);
767 auto wideType = IntegerType::get(builder.getContext(), width + 1);
768 Value maxMask = builder.create<arith::ConstantOp>(
769 loc, IntegerAttr::get(feltType, llvm::APInt::getAllOnes(width))
771 Value complement = builder.create<arith::XOrIOp>(loc, operand, maxMask);
772 Value complementWide = builder.create<arith::ExtUIOp>(loc, wideType, complement);
773 return normalizeWideValue(builder, loc, complementWide, width, field);
777static Value loadStorageScalar(OpBuilder &builder, Location loc, Value storageLeaf) {
778 auto memrefType = llvm::cast<MemRefType>(storageLeaf.getType());
779 SmallVector<Value> indices;
780 indices.reserve(memrefType.getRank());
781 for (int64_t dim = 0; dim < memrefType.getRank(); ++dim) {
782 indices.push_back(makeIndexConstant(builder, loc, 0));
784 return builder.create<memref::LoadOp>(loc, storageLeaf, indices);
788static void storeStorageScalar(OpBuilder &builder, Location loc, Value scalar, Value storageLeaf) {
789 auto memrefType = llvm::cast<MemRefType>(storageLeaf.getType());
790 SmallVector<Value> indices;
791 indices.reserve(memrefType.getRank());
792 for (int64_t dim = 0; dim < memrefType.getRank(); ++dim) {
793 indices.push_back(makeIndexConstant(builder, loc, 0));
795 builder.create<memref::StoreOp>(loc, scalar, storageLeaf, indices);
799static LogicalResult copyIntoStorage(
800 OpBuilder &builder, Location loc, Type sourceType, ArrayRef<Value> destLeaves,
801 ArrayRef<Value> sourceLeaves, SymbolTableCollection &tables, Operation *origin,
804 auto leafTypes = getLeafTypes(sourceType, tables, origin, field);
805 if (failed(leafTypes)) {
808 if (destLeaves.size() != sourceLeaves.size() || destLeaves.size() != leafTypes->size()) {
809 origin->emitError(
"flattened leaf mismatch while copying aggregate storage");
812 for (
auto [leafType, destLeaf, srcLeaf] : llvm::zip(*leafTypes, destLeaves, sourceLeaves)) {
813 if (isa<MemRefType>(leafType)) {
814 builder.create<memref::CopyOp>(loc, srcLeaf, destLeaf);
817 storeStorageScalar(builder, loc, srcLeaf, destLeaf);
823static FailureOr<LoweredValue> readNamedAggregateValue(
824 OpBuilder &builder, Location loc, Type ownerType, StringRef name,
const LoweredValue &owner,
825 SymbolTableCollection &tables, Operation *origin,
const Field &field
827 auto subType = getNamedSubType(ownerType, name, tables, origin);
828 if (failed(subType)) {
831 auto span = getNamedLeafSpan(ownerType, name, tables, origin, field);
835 LoweredValue result {*subType, {}};
836 auto leafTypes = getLeafTypes(*subType, tables, origin, field);
837 if (failed(leafTypes)) {
840 auto leaves = ArrayRef<Value>(owner.leaves).slice(span->first, span->second);
841 for (
auto [leafType, leafValue] : llvm::zip(*leafTypes, leaves)) {
842 if (isa<MemRefType>(leafType)) {
843 result.leaves.push_back(leafValue);
845 result.leaves.push_back(loadStorageScalar(builder, loc, leafValue));
852static LogicalResult writeNamedAggregateValue(
853 OpBuilder &builder, Location loc, Type ownerType, StringRef name, LoweredValue &owner,
854 const LoweredValue &value, SymbolTableCollection &tables, Operation *origin,
const Field &field
856 auto subType = getNamedSubType(ownerType, name, tables, origin);
857 if (failed(subType)) {
860 auto span = getNamedLeafSpan(ownerType, name, tables, origin, field);
864 return copyIntoStorage(
865 builder, loc, *subType, ArrayRef<Value>(owner.leaves).slice(span->first, span->second),
866 value.leaves, tables, origin, field
871static FailureOr<Value>
872createElementSubview(OpBuilder &builder, Location loc, Value
source, ValueRange outerIndices) {
873 auto sourceType = llvm::cast<MemRefType>(
source.getType());
874 SmallVector<OpFoldResult> mixedOffsets;
875 SmallVector<OpFoldResult> mixedSizes;
876 SmallVector<OpFoldResult> mixedStrides;
879 emitError(loc) << llvm::toString(indexedRank.takeError());
882 mixedOffsets.reserve(sourceType.getRank());
883 mixedSizes.reserve(sourceType.getRank());
884 mixedStrides.reserve(sourceType.getRank());
885 for (Value index : outerIndices) {
886 mixedOffsets.push_back(index);
888 for (int64_t dim = *indexedRank; dim < sourceType.getRank(); ++dim) {
889 mixedOffsets.push_back(builder.getIndexAttr(0));
891 for (int64_t dim = 0; dim < *indexedRank; ++dim) {
892 mixedSizes.push_back(builder.getIndexAttr(1));
894 for (int64_t dim = *indexedRank; dim < sourceType.getRank(); ++dim) {
895 mixedSizes.push_back(memref::getMixedSize(builder, loc,
source, dim));
897 for (int64_t dim = 0; dim < sourceType.getRank(); ++dim) {
898 mixedStrides.push_back(builder.getIndexAttr(1));
900 SmallVector<int64_t> desiredShape;
903 emitError(loc) << llvm::toString(reserveSize.takeError());
906 desiredShape.reserve(*reserveSize);
907 for (int64_t dim = *indexedRank; dim < sourceType.getRank(); ++dim) {
910 emitError(loc) << llvm::toString(dimIndex.takeError());
913 if (
auto attr = llvm::dyn_cast<Attribute>(mixedSizes[*dimIndex])) {
914 desiredShape.push_back(llvm::cast<IntegerAttr>(attr).getInt());
916 desiredShape.push_back(ShapedType::kDynamic);
919 if (desiredShape.empty()) {
920 desiredShape.push_back(1);
922 auto resultType = llvm::cast<MemRefType>(memref::SubViewOp::inferRankReducedResultType(
923 desiredShape, sourceType, mixedOffsets, mixedSizes, mixedStrides
925 auto op = builder.create<memref::SubViewOp>(
926 loc, resultType,
source, mixedOffsets, mixedSizes, mixedStrides
928 return success(op.getResult());
932static FailureOr<LoweredValue> readArrayElement(
933 OpBuilder &builder, Location loc, array::ArrayType arrayType,
const LoweredValue &arrayValue,
934 ArrayRef<Value> indices
936 Type elementType = arrayType.getElementType();
937 LoweredValue result {elementType, {}};
938 if (isScalarType(elementType)) {
939 result.leaves.push_back(
940 builder.create<memref::LoadOp>(loc, arrayValue.leaves.front(), indices)
945 for (Value sourceLeaf : arrayValue.leaves) {
946 auto subview = createElementSubview(builder, loc, sourceLeaf, indices);
947 if (failed(subview)) {
950 result.leaves.push_back(*subview);
956static LogicalResult writeArrayElement(
957 OpBuilder &builder, Location loc, array::ArrayType arrayType, LoweredValue &arrayValue,
958 ArrayRef<Value> indices,
const LoweredValue &elementValue
960 Type elementType = arrayType.getElementType();
961 if (isScalarType(elementType)) {
962 builder.create<memref::StoreOp>(
963 loc, elementValue.leaves.front(), arrayValue.leaves.front(), indices
968 for (
auto [destLeaf, srcLeaf] : llvm::zip(arrayValue.leaves, elementValue.leaves)) {
969 auto subview = createElementSubview(builder, loc, destLeaf, indices);
970 if (failed(subview)) {
973 builder.create<memref::CopyOp>(loc, srcLeaf, *subview);
979static LogicalResult appendFlatLeavesToTypes(
980 OpBuilder &builder, Location loc,
const LoweredValue &value, ArrayRef<Type> targetLeafTypes,
981 SmallVectorImpl<Value> &out, Operation *origin
983 if (targetLeafTypes.size() != value.leaves.size()) {
984 origin->emitError(
"flattened leaf mismatch during call lowering");
987 for (
auto [leafValue, leafType] : llvm::zip(value.leaves, targetLeafTypes)) {
988 if (leafValue.getType() == leafType) {
989 out.push_back(leafValue);
992 if (isa<MemRefType>(leafValue.getType()) && isa<MemRefType>(leafType)) {
993 out.push_back(builder.create<memref::CastOp>(loc, leafType, leafValue));
996 origin->emitError(
"lowered leaf type mismatch during call lowering");
1007 ModuleOp
mod, SymbolTableCollection &symbolTables,
const Field &moduleField,
1008 const WitgenOptions &options
1010 : moduleOp(
mod), tables(symbolTables), field(moduleField),
1011 uninitializedBehavior(options.uninitializedBehavior), rng(
makeDefaultValueRng(options)) {}
1014 FailureOr<func::FuncOp> lowerFunction(function::FuncDefOp funcOp) {
1015 if (funcOp.isExternal()) {
1016 funcOp.emitError(
"execution-engine backend does not lower extern functions");
1019 if (!funcOp.getBody().hasOneBlock()) {
1020 funcOp.emitError(
"execution-engine backend only supports single-block functions");
1024 SmallVector<Type> loweredArgTypes;
1025 for (Type argType : funcOp.getArgumentTypes()) {
1027 flattenABILeafTypes(argType, tables, funcOp.getOperation(), field, loweredArgTypes)
1032 SmallVector<Type> loweredResultTypes;
1033 for (Type resultType : funcOp.getResultTypes()) {
1034 if (failed(flattenABILeafTypes(
1035 resultType, tables, funcOp.getOperation(), field, loweredResultTypes
1041 OpBuilder moduleBuilder(moduleOp.getContext());
1042 moduleBuilder.setInsertionPointToEnd(moduleOp.getBody());
1043 auto loweredFunc = moduleBuilder.create<func::FuncOp>(
1044 funcOp.getLoc(), mangleFunctionName(funcOp),
1045 moduleBuilder.getFunctionType(loweredArgTypes, loweredResultTypes)
1047 Block *entry = loweredFunc.addEntryBlock();
1048 OpBuilder builder(entry, entry->begin());
1050 DenseMap<Value, LoweredValue> valueMap;
1051 unsigned cursor = 0;
1052 for (
auto [arg, argType] :
1053 llvm::zip(funcOp.getBody().front().getArguments(), funcOp.getArgumentTypes())) {
1054 auto leafCount = getLeafCount(argType, tables, funcOp.getOperation(), field);
1055 if (failed(leafCount)) {
1056 loweredFunc.erase();
1059 LoweredValue lowered {argType, {}};
1060 lowered.leaves.append(
1061 entry->getArguments().begin() + cursor,
1062 entry->getArguments().begin() + cursor + *leafCount
1064 cursor += *leafCount;
1065 valueMap[arg] = std::move(lowered);
1068 if (failed(lowerBlock(builder, funcOp.getBody().front(), valueMap))) {
1069 loweredFunc.erase();
1077 SymbolTableCollection &tables;
1080 std::mt19937_64 rng;
1083 FailureOr<LoweredValue>
1084 lookup(Value value, DenseMap<Value, LoweredValue> &valueMap, Operation *origin) {
1085 auto it = valueMap.find(value);
1086 if (it == valueMap.end()) {
1087 origin->emitError(
"failed to find lowered SSA value");
1095 lookupScalar(Value value, DenseMap<Value, LoweredValue> &valueMap, Operation *origin) {
1096 auto lowered = lookup(value, valueMap, origin);
1097 if (failed(lowered) || lowered->leaves.size() != 1 ||
1098 isa<MemRefType>(lowered->leaves.front().getType())) {
1099 origin->emitError(
"expected scalar lowered value");
1102 return lowered->leaves.front();
1107 lowerBlock(OpBuilder &builder, Block &block, DenseMap<Value, LoweredValue> &valueMap) {
1108 for (Operation &op : block) {
1109 if (failed(lowerOperation(builder, op, valueMap))) {
1118 lowerFeltCmp(OpBuilder &builder, Location loc, boolean::CmpOp cmpOp, Value lhs, Value rhs) {
1119 arith::CmpIPredicate predicate;
1120 switch (cmpOp.getPredicate()) {
1122 predicate = arith::CmpIPredicate::eq;
1125 predicate = arith::CmpIPredicate::ne;
1128 predicate = arith::CmpIPredicate::ult;
1131 predicate = arith::CmpIPredicate::ule;
1134 predicate = arith::CmpIPredicate::ugt;
1137 predicate = arith::CmpIPredicate::uge;
1140 return builder.create<arith::CmpIOp>(loc, predicate, lhs, rhs).getResult();
1145 lowerOperation(OpBuilder &builder, Operation &op, DenseMap<Value, LoweredValue> &valueMap) {
1146 Location loc = op.getLoc();
1148 auto bind = [&](Value result, LoweredValue lowered) {
1149 valueMap[result] = std::move(lowered);
1153 if (
auto returnOp = dyn_cast<function::ReturnOp>(op)) {
1154 SmallVector<Value> results;
1155 for (Value operand : returnOp.getOperands()) {
1156 auto lowered = lookup(operand, valueMap, returnOp.getOperation());
1157 auto leafTypes = getABILeafTypes(operand.getType(), tables, returnOp.getOperation(), field);
1158 if (failed(lowered) || failed(leafTypes) ||
1159 failed(appendFlatLeavesToTypes(
1160 builder, loc, *lowered, *leafTypes, results, returnOp.getOperation()
1165 builder.create<func::ReturnOp>(loc, results);
1169 if (
auto yieldOp = dyn_cast<scf::YieldOp>(op)) {
1170 SmallVector<Value> results;
1171 for (Value operand : yieldOp.getOperands()) {
1172 auto lowered = lookup(operand, valueMap, yieldOp.getOperation());
1173 auto leafTypes = getABILeafTypes(operand.getType(), tables, yieldOp.getOperation(), field);
1174 if (failed(lowered) || failed(leafTypes) ||
1175 failed(appendFlatLeavesToTypes(
1176 builder, loc, *lowered, *leafTypes, results, yieldOp.getOperation()
1181 builder.create<scf::YieldOp>(loc, results);
1184 if (
auto conditionOp = dyn_cast<scf::ConditionOp>(op)) {
1186 lookupScalar(conditionOp.getCondition(), valueMap, conditionOp.getOperation());
1187 if (failed(condition)) {
1190 SmallVector<Value> results;
1191 for (Value operand : conditionOp.getArgs()) {
1192 auto lowered = lookup(operand, valueMap, conditionOp.getOperation());
1194 getABILeafTypes(operand.getType(), tables, conditionOp.getOperation(), field);
1195 if (failed(lowered) || failed(leafTypes) ||
1196 failed(appendFlatLeavesToTypes(
1197 builder, loc, *lowered, *leafTypes, results, conditionOp.getOperation()
1202 builder.create<scf::ConditionOp>(loc, *condition, results);
1206 if (
auto constantOp = dyn_cast<arith::ConstantOp>(op)) {
1207 Operation *clone = builder.clone(op);
1209 constantOp.getResult(), LoweredValue {constantOp.getType(), {clone->getResult(0)}}
1213 if (
auto feltConst = dyn_cast<felt::FeltConstantOp>(op)) {
1214 auto intType = IntegerType::get(builder.getContext(), field.bitWidth());
1217 auto modVal = constVal % field.prime();
1219 Value lowered = builder.create<arith::ConstantOp>(loc, IntegerAttr::get(intType, intVal));
1220 return bind(feltConst.getResult(), LoweredValue {feltConst.getType(), {lowered}});
1223 if (
auto nondetOp = dyn_cast<llzk::NonDetOp>(op)) {
1224 auto lowered = createDefaultValue(
1225 builder, loc, nondetOp.getType(), tables, nondetOp.getOperation(), field,
1226 uninitializedBehavior, rng
1228 if (failed(lowered)) {
1231 return bind(nondetOp.getResult(), std::move(*lowered));
1234 if (
auto addOp = dyn_cast<felt::AddFeltOp>(op)) {
1235 auto lhs = lookupScalar(addOp.getLhs(), valueMap, addOp.getOperation());
1236 auto rhs = lookupScalar(addOp.getRhs(), valueMap, addOp.getOperation());
1237 if (failed(lhs) || failed(rhs)) {
1242 LoweredValue {addOp.getType(), {lowerFeltAdd(builder, loc, *lhs, *rhs, field)}}
1245 if (
auto powOp = dyn_cast<felt::PowFeltOp>(op)) {
1246 auto lhs = lookupScalar(powOp.getLhs(), valueMap, powOp.getOperation());
1247 auto rhs = lookupScalar(powOp.getRhs(), valueMap, powOp.getOperation());
1248 if (failed(lhs) || failed(rhs)) {
1253 LoweredValue {powOp.getType(), {lowerFeltPow(builder, loc, *lhs, *rhs, field)}}
1256 if (
auto andOp = dyn_cast<felt::AndFeltOp>(op)) {
1257 auto lhs = lookupScalar(andOp.getLhs(), valueMap, andOp.getOperation());
1258 auto rhs = lookupScalar(andOp.getRhs(), valueMap, andOp.getOperation());
1259 if (failed(lhs) || failed(rhs)) {
1264 LoweredValue {andOp.getType(), {builder.create<arith::AndIOp>(loc, *lhs, *rhs)}}
1267 if (
auto orOp = dyn_cast<felt::OrFeltOp>(op)) {
1268 auto lhs = lookupScalar(orOp.getLhs(), valueMap, orOp.getOperation());
1269 auto rhs = lookupScalar(orOp.getRhs(), valueMap, orOp.getOperation());
1270 if (failed(lhs) || failed(rhs)) {
1275 LoweredValue {orOp.getType(), {lowerFeltOr(builder, loc, *lhs, *rhs, field)}}
1278 if (
auto xorOp = dyn_cast<felt::XorFeltOp>(op)) {
1279 auto lhs = lookupScalar(xorOp.getLhs(), valueMap, xorOp.getOperation());
1280 auto rhs = lookupScalar(xorOp.getRhs(), valueMap, xorOp.getOperation());
1281 if (failed(lhs) || failed(rhs)) {
1286 LoweredValue {xorOp.getType(), {lowerFeltXor(builder, loc, *lhs, *rhs, field)}}
1289 if (
auto subOp = dyn_cast<felt::SubFeltOp>(op)) {
1290 auto lhs = lookupScalar(subOp.getLhs(), valueMap, subOp.getOperation());
1291 auto rhs = lookupScalar(subOp.getRhs(), valueMap, subOp.getOperation());
1292 if (failed(lhs) || failed(rhs)) {
1297 LoweredValue {subOp.getType(), {lowerFeltSub(builder, loc, *lhs, *rhs, field)}}
1300 if (
auto mulOp = dyn_cast<felt::MulFeltOp>(op)) {
1301 auto lhs = lookupScalar(mulOp.getLhs(), valueMap, mulOp.getOperation());
1302 auto rhs = lookupScalar(mulOp.getRhs(), valueMap, mulOp.getOperation());
1303 if (failed(lhs) || failed(rhs)) {
1308 LoweredValue {mulOp.getType(), {lowerFeltMul(builder, loc, *lhs, *rhs, field)}}
1311 if (
auto negOp = dyn_cast<felt::NegFeltOp>(op)) {
1312 auto operand = lookupScalar(negOp.getOperand(), valueMap, negOp.getOperation());
1313 if (failed(operand)) {
1318 LoweredValue {negOp.getType(), {lowerFeltNeg(builder, loc, *operand, field)}}
1321 if (
auto invOp = dyn_cast<felt::InvFeltOp>(op)) {
1322 auto operand = lookupScalar(invOp.getOperand(), valueMap, invOp.getOperation());
1323 if (failed(operand)) {
1328 LoweredValue {invOp.getType(), {lowerFeltInv(builder, loc, *operand, field)}}
1331 if (
auto divOp = dyn_cast<felt::DivFeltOp>(op)) {
1332 auto lhs = lookupScalar(divOp.getLhs(), valueMap, divOp.getOperation());
1333 auto rhs = lookupScalar(divOp.getRhs(), valueMap, divOp.getOperation());
1334 if (failed(lhs) || failed(rhs)) {
1339 LoweredValue {divOp.getType(), {lowerFeltDiv(builder, loc, *lhs, *rhs, field)}}
1342 if (
auto uintDivOp = dyn_cast<felt::UnsignedIntDivFeltOp>(op)) {
1343 auto lhs = lookupScalar(uintDivOp.getLhs(), valueMap, uintDivOp.getOperation());
1344 auto rhs = lookupScalar(uintDivOp.getRhs(), valueMap, uintDivOp.getOperation());
1345 if (failed(lhs) || failed(rhs)) {
1348 assertNonZeroFelt(builder, loc, *rhs,
"felt.uintdiv divisor must be non-zero");
1350 uintDivOp.getResult(),
1351 LoweredValue {uintDivOp.getType(), {lowerFeltUnsignedDiv(builder, loc, *lhs, *rhs)}}
1354 if (
auto sintDivOp = dyn_cast<felt::SignedIntDivFeltOp>(op)) {
1355 auto lhs = lookupScalar(sintDivOp.getLhs(), valueMap, sintDivOp.getOperation());
1356 auto rhs = lookupScalar(sintDivOp.getRhs(), valueMap, sintDivOp.getOperation());
1357 if (failed(lhs) || failed(rhs)) {
1360 assertNonZeroFelt(builder, loc, *rhs,
"felt.sintdiv divisor must be non-zero");
1362 sintDivOp.getResult(),
1363 LoweredValue {sintDivOp.getType(), {lowerFeltSignedDiv(builder, loc, *lhs, *rhs, field)}}
1366 if (
auto umodOp = dyn_cast<felt::UnsignedModFeltOp>(op)) {
1367 auto lhs = lookupScalar(umodOp.getLhs(), valueMap, umodOp.getOperation());
1368 auto rhs = lookupScalar(umodOp.getRhs(), valueMap, umodOp.getOperation());
1369 if (failed(lhs) || failed(rhs)) {
1372 assertNonZeroFelt(builder, loc, *rhs,
"felt.umod divisor must be non-zero");
1375 LoweredValue {umodOp.getType(), {lowerFeltUnsignedMod(builder, loc, *lhs, *rhs)}}
1378 if (
auto smodOp = dyn_cast<felt::SignedModFeltOp>(op)) {
1379 auto lhs = lookupScalar(smodOp.getLhs(), valueMap, smodOp.getOperation());
1380 auto rhs = lookupScalar(smodOp.getRhs(), valueMap, smodOp.getOperation());
1381 if (failed(lhs) || failed(rhs)) {
1384 assertNonZeroFelt(builder, loc, *rhs,
"felt.smod divisor must be non-zero");
1387 LoweredValue {smodOp.getType(), {lowerFeltSignedMod(builder, loc, *lhs, *rhs, field)}}
1390 if (
auto shrOp = dyn_cast<felt::ShrFeltOp>(op)) {
1391 auto lhs = lookupScalar(shrOp.getLhs(), valueMap, shrOp.getOperation());
1392 auto rhs = lookupScalar(shrOp.getRhs(), valueMap, shrOp.getOperation());
1393 if (failed(lhs) || failed(rhs)) {
1398 LoweredValue {shrOp.getType(), {lowerFeltShr(builder, loc, *lhs, *rhs, field)}}
1401 if (
auto shlOp = dyn_cast<felt::ShlFeltOp>(op)) {
1402 auto lhs = lookupScalar(shlOp.getLhs(), valueMap, shlOp.getOperation());
1403 auto rhs = lookupScalar(shlOp.getRhs(), valueMap, shlOp.getOperation());
1404 if (failed(lhs) || failed(rhs)) {
1409 LoweredValue {shlOp.getType(), {lowerFeltShl(builder, loc, *lhs, *rhs, field)}}
1412 if (
auto notOp = dyn_cast<felt::NotFeltOp>(op)) {
1413 auto operand = lookupScalar(
notOp.getOperand(), valueMap,
notOp.getOperation());
1414 if (failed(operand)) {
1419 LoweredValue {notOp.getType(), {lowerFeltNot(builder, loc, *operand, field)}}
1423 if (
auto cmpOp = dyn_cast<boolean::CmpOp>(op)) {
1424 auto lhs = lookupScalar(cmpOp.getLhs(), valueMap, cmpOp.getOperation());
1425 auto rhs = lookupScalar(cmpOp.getRhs(), valueMap, cmpOp.getOperation());
1426 if (failed(lhs) || failed(rhs)) {
1429 auto lowered = lowerFeltCmp(builder, loc, cmpOp, *lhs, *rhs);
1430 if (failed(lowered)) {
1433 return bind(cmpOp.getResult(), LoweredValue {cmpOp.getType(), {*lowered}});
1435 if (
auto assertOp = dyn_cast<boolean::AssertOp>(op)) {
1436 auto condition = lookupScalar(assertOp.getCondition(), valueMap, assertOp.getOperation());
1437 if (failed(condition)) {
1440 builder.create<cf::AssertOp>(
1441 loc, *condition, assertOp.getMsg() ? assertOp.getMsg()->str() :
"bool.assert failed"
1445 if (
auto andOp = dyn_cast<boolean::AndBoolOp>(op)) {
1446 auto lhs = lookupScalar(andOp.getLhs(), valueMap, andOp.getOperation());
1447 auto rhs = lookupScalar(andOp.getRhs(), valueMap, andOp.getOperation());
1448 if (failed(lhs) || failed(rhs)) {
1453 LoweredValue {andOp.getType(), {builder.create<arith::AndIOp>(loc, *lhs, *rhs)}}
1456 if (
auto orOp = dyn_cast<boolean::OrBoolOp>(op)) {
1457 auto lhs = lookupScalar(orOp.getLhs(), valueMap, orOp.getOperation());
1458 auto rhs = lookupScalar(orOp.getRhs(), valueMap, orOp.getOperation());
1459 if (failed(lhs) || failed(rhs)) {
1464 LoweredValue {orOp.getType(), {builder.create<arith::OrIOp>(loc, *lhs, *rhs)}}
1467 if (
auto xorOp = dyn_cast<boolean::XorBoolOp>(op)) {
1468 auto lhs = lookupScalar(xorOp.getLhs(), valueMap, xorOp.getOperation());
1469 auto rhs = lookupScalar(xorOp.getRhs(), valueMap, xorOp.getOperation());
1470 if (failed(lhs) || failed(rhs)) {
1475 LoweredValue {xorOp.getType(), {builder.create<arith::XOrIOp>(loc, *lhs, *rhs)}}
1478 if (
auto notOp = dyn_cast<boolean::NotBoolOp>(op)) {
1479 auto operand = lookupScalar(
notOp.getOperand(), valueMap,
notOp.getOperation());
1480 if (failed(operand)) {
1483 Value one = builder.create<arith::ConstantOp>(
1484 loc, IntegerAttr::get(IntegerType::get(builder.getContext(), 1), 1)
1488 LoweredValue {notOp.getType(), {builder.create<arith::XOrIOp>(loc, *operand, one)}}
1492 if (
auto intToFelt = dyn_cast<cast::IntToFeltOp>(op)) {
1493 auto operand = lookupScalar(intToFelt.getValue(), valueMap, intToFelt.getOperation());
1494 if (failed(operand)) {
1497 auto dstType = IntegerType::get(builder.getContext(), field.bitWidth());
1499 if (isa<IndexType>((*operand).getType())) {
1500 lowered = builder.create<arith::IndexCastUIOp>(loc, dstType, *operand);
1502 auto intType = llvm::cast<IntegerType>((*operand).getType());
1503 if (intType.getWidth() < dstType.getWidth()) {
1504 lowered = builder.create<arith::ExtUIOp>(loc, dstType, *operand);
1505 }
else if (intType.getWidth() > dstType.getWidth()) {
1506 lowered = normalizeWideValue(builder, loc, *operand, dstType.getWidth(), field);
1511 return bind(intToFelt.getResult(), LoweredValue {intToFelt.getType(), {lowered}});
1513 if (
auto feltToIndex = dyn_cast<cast::FeltToIndexOp>(op)) {
1514 auto operand = lookupScalar(feltToIndex.getValue(), valueMap, feltToIndex.getOperation());
1515 if (failed(operand)) {
1519 feltToIndex.getResult(),
1521 feltToIndex.getType(),
1522 {builder.create<arith::IndexCastUIOp>(loc, builder.getIndexType(), *operand)}
1527 if (
auto structNewOp = dyn_cast<component::CreateStructOp>(op)) {
1528 auto lowered = createDefaultValue(
1529 builder, loc, structNewOp.getType(), tables, structNewOp.getOperation(), field,
1530 uninitializedBehavior, rng
1532 if (failed(lowered)) {
1535 return bind(structNewOp.getResult(), std::move(*lowered));
1537 if (
auto readMemberOp = dyn_cast<component::MemberReadOp>(op)) {
1538 auto componentValue =
1539 lookup(readMemberOp.getComponent(), valueMap, readMemberOp.getOperation());
1540 if (failed(componentValue)) {
1543 auto lowered = readNamedAggregateValue(
1544 builder, loc, readMemberOp.getComponent().getType(), readMemberOp.getMemberName(),
1545 *componentValue, tables, readMemberOp.getOperation(), field
1547 if (failed(lowered)) {
1550 return bind(readMemberOp.getResult(), std::move(*lowered));
1552 if (
auto writeMemberOp = dyn_cast<component::MemberWriteOp>(op)) {
1553 auto componentValue =
1554 lookup(writeMemberOp.getComponent(), valueMap, writeMemberOp.getOperation());
1555 auto memberValue = lookup(writeMemberOp.getVal(), valueMap, writeMemberOp.getOperation());
1556 if (failed(componentValue) || failed(memberValue)) {
1559 return writeNamedAggregateValue(
1560 builder, loc, writeMemberOp.getComponent().getType(), writeMemberOp.getMemberName(),
1561 valueMap[writeMemberOp.getComponent()], *memberValue, tables,
1562 writeMemberOp.getOperation(), field
1566 if (
auto newPodOp = dyn_cast<pod::NewPodOp>(op)) {
1567 auto lowered = createDefaultValue(
1568 builder, loc, newPodOp.getType(), tables, newPodOp.getOperation(), field,
1569 uninitializedBehavior, rng
1571 if (failed(lowered)) {
1574 for (pod::RecordValue init : newPodOp.getInitializedRecordValues()) {
1575 auto value = lookup(init.value, valueMap, newPodOp.getOperation());
1576 if (failed(value) || failed(writeNamedAggregateValue(
1577 builder, loc, newPodOp.getType(), init.name, *lowered, *value,
1578 tables, newPodOp.getOperation(), field
1583 return bind(newPodOp.getResult(), std::move(*lowered));
1585 if (
auto readPodOp = dyn_cast<pod::ReadPodOp>(op)) {
1586 auto podValue = lookup(readPodOp.getPodRef(), valueMap, readPodOp.getOperation());
1587 if (failed(podValue)) {
1590 auto lowered = readNamedAggregateValue(
1591 builder, loc, readPodOp.getPodRef().getType(), readPodOp.getRecordName(), *podValue,
1592 tables, readPodOp.getOperation(), field
1594 if (failed(lowered)) {
1597 return bind(readPodOp.getResult(), std::move(*lowered));
1599 if (
auto writePodOp = dyn_cast<pod::WritePodOp>(op)) {
1600 auto recordValue = lookup(writePodOp.getValue(), valueMap, writePodOp.getOperation());
1601 if (failed(recordValue)) {
1604 return writeNamedAggregateValue(
1605 builder, loc, writePodOp.getPodRef().getType(), writePodOp.getRecordName(),
1606 valueMap[writePodOp.getPodRef()], *recordValue, tables, writePodOp.getOperation(), field
1610 if (
auto arrayNewOp = dyn_cast<array::CreateArrayOp>(op)) {
1611 auto lowered = createDefaultValue(
1612 builder, loc, arrayNewOp.getType(), tables, arrayNewOp.getOperation(), field,
1613 uninitializedBehavior, rng
1615 if (failed(lowered)) {
1618 if (!arrayNewOp.getElements().empty()) {
1619 auto elementCount = checkedCast<size_t>(arrayNewOp.getType().getNumElements());
1620 if (!elementCount) {
1621 arrayNewOp.emitError() << llvm::toString(elementCount.takeError());
1624 if (arrayNewOp.getElements().size() != *elementCount) {
1625 arrayNewOp.emitError(
"expected one explicit element per array slot in witgen lowering");
1628 auto shape = arrayNewOp.getType().getShape();
1629 for (
auto [flatIndex, operand] : llvm::enumerate(arrayNewOp.getElements())) {
1630 auto elementValue = lookup(operand, valueMap, arrayNewOp.getOperation());
1631 if (failed(elementValue)) {
1634 SmallVector<Value> indices;
1635 auto strides = mlir::computeStrides(shape);
1636 auto flatSigned = checkedCast<int64_t>(flatIndex);
1638 arrayNewOp.emitError() << llvm::toString(flatSigned.takeError());
1641 for (int64_t index : mlir::delinearize(*flatSigned, strides)) {
1642 indices.push_back(makeIndexConstant(builder, loc, index));
1644 if (failed(writeArrayElement(
1645 builder, loc, arrayNewOp.getType(), *lowered, indices, *elementValue
1651 return bind(arrayNewOp.getResult(), std::move(*lowered));
1653 if (
auto readArrayOp = dyn_cast<array::ReadArrayOp>(op)) {
1654 SmallVector<Value> indices;
1655 for (Value indexValue : readArrayOp.getIndices()) {
1656 auto loweredIndex = lookupScalar(indexValue, valueMap, readArrayOp.getOperation());
1657 if (failed(loweredIndex)) {
1660 indices.push_back(*loweredIndex);
1662 auto arrayValue = lookup(readArrayOp.getArrRef(), valueMap, readArrayOp.getOperation());
1663 if (failed(arrayValue)) {
1666 auto lowered = readArrayElement(
1667 builder, loc, llvm::cast<array::ArrayType>(readArrayOp.getArrRef().getType()),
1668 *arrayValue, indices
1670 if (failed(lowered)) {
1673 return bind(readArrayOp.getResult(), std::move(*lowered));
1675 if (
auto writeArrayOp = dyn_cast<array::WriteArrayOp>(op)) {
1676 SmallVector<Value> indices;
1677 for (Value indexValue : writeArrayOp.getIndices()) {
1678 auto loweredIndex = lookupScalar(indexValue, valueMap, writeArrayOp.getOperation());
1679 if (failed(loweredIndex)) {
1682 indices.push_back(*loweredIndex);
1684 auto elementValue = lookup(writeArrayOp.getRvalue(), valueMap, writeArrayOp.getOperation());
1685 if (failed(elementValue)) {
1688 return writeArrayElement(
1689 builder, loc, llvm::cast<array::ArrayType>(writeArrayOp.getArrRef().getType()),
1690 valueMap[writeArrayOp.getArrRef()], indices, *elementValue
1694 if (
auto cmpiOp = dyn_cast<arith::CmpIOp>(op)) {
1695 auto lhs = lookupScalar(cmpiOp.getLhs(), valueMap, cmpiOp.getOperation());
1696 auto rhs = lookupScalar(cmpiOp.getRhs(), valueMap, cmpiOp.getOperation());
1697 if (failed(lhs) || failed(rhs)) {
1704 {builder.create<arith::CmpIOp>(loc, cmpiOp.getPredicate(), *lhs, *rhs)}
1708 if (
auto selectOp = dyn_cast<arith::SelectOp>(op)) {
1709 auto cond = lookupScalar(selectOp.getCondition(), valueMap, selectOp.getOperation());
1710 auto trueValue = lookupScalar(selectOp.getTrueValue(), valueMap, selectOp.getOperation());
1711 auto falseValue = lookupScalar(selectOp.getFalseValue(), valueMap, selectOp.getOperation());
1712 if (failed(cond) || failed(trueValue) || failed(falseValue)) {
1716 selectOp.getResult(),
1719 {builder.create<arith::SelectOp>(loc, *cond, *trueValue, *falseValue)}
1723 if (
auto addiOp = dyn_cast<arith::AddIOp>(op)) {
1724 auto lhs = lookupScalar(addiOp.getLhs(), valueMap, addiOp.getOperation());
1725 auto rhs = lookupScalar(addiOp.getRhs(), valueMap, addiOp.getOperation());
1726 if (failed(lhs) || failed(rhs)) {
1731 LoweredValue {addiOp.getType(), {builder.create<arith::AddIOp>(loc, *lhs, *rhs)}}
1734 if (
auto subiOp = dyn_cast<arith::SubIOp>(op)) {
1735 auto lhs = lookupScalar(subiOp.getLhs(), valueMap, subiOp.getOperation());
1736 auto rhs = lookupScalar(subiOp.getRhs(), valueMap, subiOp.getOperation());
1737 if (failed(lhs) || failed(rhs)) {
1742 LoweredValue {subiOp.getType(), {builder.create<arith::SubIOp>(loc, *lhs, *rhs)}}
1746 if (
auto callOp = dyn_cast<function::CallOp>(op)) {
1747 if (callOp.getTemplateParams() || !callOp.getMapOperands().empty()) {
1748 callOp.emitError(
"execution-engine backend encountered an unflattened function.call");
1751 auto *callable = callOp.resolveCallableInTable(&tables);
1752 auto callee = dyn_cast_or_null<function::FuncDefOp>(callable);
1754 callOp.emitError(
"failed to resolve callee during execution-engine lowering");
1757 SmallVector<Type> resultTypes;
1758 for (Type resultType : callOp.getResultTypes()) {
1760 flattenABILeafTypes(resultType, tables, callOp.getOperation(), field, resultTypes)
1765 SmallVector<Value> flatArgs;
1766 for (Value operand : callOp.getArgOperands()) {
1767 auto lowered = lookup(operand, valueMap, callOp.getOperation());
1768 auto leafTypes = getABILeafTypes(operand.getType(), tables, callOp.getOperation(), field);
1769 if (failed(lowered) || failed(leafTypes) ||
1770 failed(appendFlatLeavesToTypes(
1771 builder, loc, *lowered, *leafTypes, flatArgs, callOp.getOperation()
1777 builder.create<func::CallOp>(loc, mangleFunctionName(callee), resultTypes, flatArgs);
1778 auto loweredCallResults = loweredCall.getResults();
1779 size_t totalResults = loweredCallResults.size();
1781 for (
auto [oldResult, oldType] : llvm::zip(callOp.getResults(), callOp.getResultTypes())) {
1782 auto leafCount = getLeafCount(oldType, tables, callOp.getOperation(), field);
1783 if (failed(leafCount)) {
1786 bool overflow =
false;
1787 size_t nextCursor = llvm::SaturatingAdd(cursor, *leafCount, &overflow);
1788 if (overflow || nextCursor > totalResults) {
1789 callOp.emitError(
"leaf count overflow while lowering function call results");
1792 LoweredValue lowered {oldType, {}};
1793 lowered.leaves.append(
1794 loweredCallResults.begin() +
static_cast<ptrdiff_t
>(cursor),
1795 loweredCallResults.begin() +
static_cast<ptrdiff_t
>(nextCursor)
1797 valueMap[oldResult] = std::move(lowered);
1798 cursor = nextCursor;
1803 if (
auto whileOp = dyn_cast<scf::WhileOp>(op)) {
1804 SmallVector<Value> initArgs;
1805 SmallVector<size_t> beforeLeafCounts;
1806 for (
auto [init, initType] : llvm::zip(whileOp.getInits(), whileOp.getOperandTypes())) {
1807 auto lowered = lookup(init, valueMap, whileOp.getOperation());
1808 auto leafTypes = getABILeafTypes(initType, tables, whileOp.getOperation(), field);
1809 if (failed(lowered) || failed(leafTypes) ||
1810 failed(appendFlatLeavesToTypes(
1811 builder, loc, *lowered, *leafTypes, initArgs, whileOp.getOperation()
1815 auto count = getLeafCount(initType, tables, whileOp.getOperation(), field);
1816 if (failed(count)) {
1819 beforeLeafCounts.push_back(*count);
1822 SmallVector<size_t> resultLeafCounts;
1823 SmallVector<Type> loweredResultTypes;
1824 for (Type resultType : whileOp.getResultTypes()) {
1825 auto leafTypes = getABILeafTypes(resultType, tables, whileOp.getOperation(), field);
1826 auto count = getLeafCount(resultType, tables, whileOp.getOperation(), field);
1827 if (failed(leafTypes) || failed(count)) {
1830 loweredResultTypes.append(leafTypes->begin(), leafTypes->end());
1831 resultLeafCounts.push_back(*count);
1834 auto mapRegionArguments = [&](
auto oldArgs,
auto oldTypes,
auto leafCounts,
auto newArgs,
1835 StringRef overflowMessage,
1836 DenseMap<Value, LoweredValue> ®ionMap) -> LogicalResult {
1837 size_t totalArgs = newArgs.size();
1839 for (
auto [oldArg, oldType, leafCount] : llvm::zip(oldArgs, oldTypes, leafCounts)) {
1840 bool overflow =
false;
1841 size_t nextCursor = llvm::SaturatingAdd(cursor, leafCount, &overflow);
1842 if (overflow || nextCursor > totalArgs) {
1843 whileOp.emitError(overflowMessage);
1846 LoweredValue lowered {oldType, {}};
1847 lowered.leaves.append(
1851 regionMap[oldArg] = std::move(lowered);
1852 cursor = nextCursor;
1857 LogicalResult whileLoweringStatus = success();
1858 auto newWhile = builder.create<scf::WhileOp>(
1859 loc, loweredResultTypes, initArgs,
1860 [&](OpBuilder ®ionBuilder, Location , ValueRange beforeArgs) {
1861 DenseMap<Value, LoweredValue> beforeMap(valueMap.begin(), valueMap.end());
1862 if (failed(mapRegionArguments(
1863 whileOp.getBeforeArguments(), whileOp.getOperandTypes(), beforeLeafCounts,
1864 beforeArgs,
"leaf count overflow while lowering while-loop before-region args",
1867 failed(lowerBlock(regionBuilder, whileOp.getBefore().front(), beforeMap))) {
1868 whileLoweringStatus = failure();
1870 }, [&](OpBuilder ®ionBuilder, Location , ValueRange afterArgs) {
1871 DenseMap<Value, LoweredValue> afterMap(valueMap.begin(), valueMap.end());
1872 if (failed(mapRegionArguments(
1873 whileOp.getAfterArguments(), whileOp.getResultTypes(), resultLeafCounts, afterArgs,
1874 "leaf count overflow while lowering while-loop after-region args", afterMap
1876 failed(lowerBlock(regionBuilder, whileOp.getAfter().front(), afterMap))) {
1877 whileLoweringStatus = failure();
1881 if (failed(whileLoweringStatus)) {
1886 auto newWhileResults = newWhile.getResults();
1887 size_t totalResults = newWhileResults.size();
1889 for (
auto [oldResult, oldType, leafCount] :
1890 llvm::zip(whileOp.getResults(), whileOp.getResultTypes(), resultLeafCounts)) {
1891 bool overflow =
false;
1892 size_t nextCursor = llvm::SaturatingAdd(cursor, leafCount, &overflow);
1893 if (overflow || nextCursor > totalResults) {
1894 whileOp.emitError(
"leaf count overflow while lowering while-loop results");
1897 LoweredValue lowered {oldType, {}};
1898 lowered.leaves.append(
1902 valueMap[oldResult] = std::move(lowered);
1903 cursor = nextCursor;
1908 if (
auto ifOp = dyn_cast<scf::IfOp>(op)) {
1909 auto condition = lookupScalar(ifOp.getCondition(), valueMap, ifOp.getOperation());
1910 if (failed(condition)) {
1914 SmallVector<size_t> resultLeafCounts;
1915 SmallVector<Type> loweredResultTypes;
1916 for (Type resultType : ifOp.getResultTypes()) {
1917 auto leafTypes = getABILeafTypes(resultType, tables, ifOp.getOperation(), field);
1918 auto count = getLeafCount(resultType, tables, ifOp.getOperation(), field);
1919 if (failed(leafTypes) || failed(count)) {
1922 loweredResultTypes.append(leafTypes->begin(), leafTypes->end());
1923 resultLeafCounts.push_back(*count);
1926 auto newIf = builder.create<scf::IfOp>(
1927 loc, loweredResultTypes, *condition,
true, !ifOp.getElseRegion().empty()
1931 OpBuilder thenBuilder = OpBuilder::atBlockBegin(&newIf.getThenRegion().front());
1932 DenseMap<Value, LoweredValue> thenMap(valueMap.begin(), valueMap.end());
1933 if (failed(lowerBlock(thenBuilder, ifOp.getThenRegion().front(), thenMap))) {
1938 if (!ifOp.getElseRegion().empty()) {
1939 OpBuilder elseBuilder = OpBuilder::atBlockBegin(&newIf.getElseRegion().front());
1940 DenseMap<Value, LoweredValue> elseMap(valueMap.begin(), valueMap.end());
1941 if (failed(lowerBlock(elseBuilder, ifOp.getElseRegion().front(), elseMap))) {
1946 auto newIfResults = newIf.getResults();
1947 size_t totalResults = newIfResults.size();
1949 for (
auto [oldResult, oldType, leafCount] :
1950 llvm::zip(ifOp.getResults(), ifOp.getResultTypes(), resultLeafCounts)) {
1951 bool overflow =
false;
1952 size_t nextCursor = llvm::SaturatingAdd(cursor, leafCount, &overflow);
1953 if (overflow || nextCursor > totalResults) {
1954 ifOp.emitError(
"leaf count overflow while lowering if-op results");
1957 LoweredValue lowered {oldType, {}};
1958 lowered.leaves.append(
1962 valueMap[oldResult] = std::move(lowered);
1963 cursor = nextCursor;
1968 if (
auto forOp = dyn_cast<scf::ForOp>(op)) {
1969 auto lb = lookupScalar(forOp.getLowerBound(), valueMap, forOp.getOperation());
1970 auto ub = lookupScalar(forOp.getUpperBound(), valueMap, forOp.getOperation());
1971 auto step = lookupScalar(forOp.getStep(), valueMap, forOp.getOperation());
1972 if (failed(lb) || failed(ub) || failed(step)) {
1976 SmallVector<Value> initArgs;
1977 SmallVector<size_t> initLeafCounts;
1978 for (
auto [init, resultType] : llvm::zip(forOp.getInitArgs(), forOp.getResultTypes())) {
1979 auto lowered = lookup(init, valueMap, forOp.getOperation());
1980 auto leafTypes = getABILeafTypes(resultType, tables, forOp.getOperation(), field);
1981 if (failed(lowered) || failed(leafTypes) ||
1982 failed(appendFlatLeavesToTypes(
1983 builder, loc, *lowered, *leafTypes, initArgs, forOp.getOperation()
1987 auto count = getLeafCount(resultType, tables, forOp.getOperation(), field);
1988 if (failed(count)) {
1991 initLeafCounts.push_back(*count);
1994 auto newFor = builder.create<scf::ForOp>(loc, *lb, *ub, *step, initArgs);
1995 if (Attribute unsignedCmpAttr = forOp->getAttr(
"unsignedCmp")) {
1996 newFor->setAttr(
"unsignedCmp", unsignedCmpAttr);
1998 DenseMap<Value, LoweredValue> bodyMap(valueMap.begin(), valueMap.end());
1999 bodyMap[forOp.getInductionVar()] =
2000 LoweredValue {forOp.getInductionVar().getType(), {newFor.getInductionVar()}};
2002 auto newForIterArgs = newFor.getRegionIterArgs();
2003 size_t totalIterArgs = newForIterArgs.size();
2005 for (
auto [oldIterArg, oldType, leafCount] :
2006 llvm::zip(forOp.getRegionIterArgs(), forOp.getResultTypes(), initLeafCounts)) {
2007 bool overflow =
false;
2008 size_t nextCursor = llvm::SaturatingAdd(cursor, leafCount, &overflow);
2009 if (overflow || nextCursor > totalIterArgs) {
2010 forOp.emitError(
"leaf count overflow while lowering for-loop region iter args");
2013 LoweredValue lowered {oldType, {}};
2014 lowered.leaves.append(
2015 newForIterArgs.begin() +
static_cast<ptrdiff_t
>(cursor),
2016 newForIterArgs.begin() +
static_cast<ptrdiff_t
>(nextCursor)
2018 bodyMap[oldIterArg] = std::move(lowered);
2019 cursor = nextCursor;
2023 newFor.getBody()->clear();
2024 OpBuilder bodyBuilder = OpBuilder::atBlockBegin(newFor.getBody());
2025 if (failed(lowerBlock(bodyBuilder, *forOp.getBody(), bodyMap))) {
2030 auto newForResults = newFor.getResults();
2031 size_t totalForResults = newForResults.size();
2033 for (
auto [oldResult, oldType, leafCount] :
2034 llvm::zip(forOp.getResults(), forOp.getResultTypes(), initLeafCounts)) {
2035 bool overflow =
false;
2036 size_t nextCursor = llvm::SaturatingAdd(cursor, leafCount, &overflow);
2037 if (overflow || nextCursor > totalForResults) {
2038 forOp.emitError(
"leaf count overflow while lowering for-loop results");
2041 LoweredValue lowered {oldType, {}};
2042 lowered.leaves.append(
2043 newForResults.begin() +
static_cast<ptrdiff_t
>(cursor),
2044 newForResults.begin() +
static_cast<ptrdiff_t
>(nextCursor)
2046 valueMap[oldResult] = std::move(lowered);
2047 cursor = nextCursor;
2053 op.emitError(
"unsupported operation in execution-engine lowering: ") << op.getName();
2059class LowerComputeToCorePass :
public PassWrapper<LowerComputeToCorePass, OperationPass<ModuleOp>> {
2061 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LowerComputeToCorePass)
2063 explicit LowerComputeToCorePass(
const WitgenOptions &opts) : options(opts) {}
2066 StringRef getArgument() const final {
return "llzk-lower-compute-to-core"; }
2069 StringRef getDescription() const final {
2070 return "Lower LLZK compute IR to func/arith/cf/scf/memref";
2074 StringRef getName()
const override {
return "LowerComputeToCorePass"; }
2077 void runOnOperation()
override {
2078 ModuleOp moduleOp = getOperation();
2079 auto field = getModuleField(moduleOp);
2080 if (failed(field)) {
2081 signalPassFailure();
2085 SymbolTableCollection tables;
2086 BodyLowerer lowerer(moduleOp, tables, field->get(), options);
2087 auto funcs = walkCollect<function::FuncDefOp>(moduleOp, [](
auto funcOp) {
2088 return !funcOp.nameIsConstrain();
2090 for (function::FuncDefOp funcOp : funcs) {
2091 if (failed(lowerer.lowerFunction(funcOp))) {
2092 signalPassFailure();
2099 WitgenOptions options;
2103class CreateWitgenEntryPass :
public PassWrapper<CreateWitgenEntryPass, OperationPass<ModuleOp>> {
2105 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(CreateWitgenEntryPass)
2108 explicit CreateWitgenEntryPass(OutputScope newOutputScope) : outputScope(newOutputScope) {}
2111 StringRef getArgument() const final {
return "llzk-create-witgen-entry"; }
2114 StringRef getDescription() const final {
2115 return "Create the llzk-witgen execution-engine entry wrapper";
2119 StringRef getName()
const override {
return "CreateWitgenEntryPass"; }
2122 void runOnOperation()
override {
2123 ModuleOp moduleOp = getOperation();
2124 auto field = getModuleField(moduleOp);
2125 if (failed(field)) {
2126 signalPassFailure();
2130 SymbolTableCollection tables;
2132 if (failed(mainDef) || !mainDef.value()) {
2133 moduleOp.emitError(
"module is missing a concrete llzk.main struct");
2134 signalPassFailure();
2137 function::FuncDefOp computeFunc = mainDef->get().getComputeFuncOp();
2139 moduleOp.emitError(
"main struct is missing @compute");
2140 signalPassFailure();
2146 if (failed(outputs)) {
2147 signalPassFailure();
2151 OpBuilder builder(moduleOp.getContext());
2152 builder.setInsertionPointToEnd(moduleOp.getBody());
2154 SmallVector<Type> wrapperArgs;
2155 for (Type argType : computeFunc.getArgumentTypes()) {
2156 SmallVector<Type> loweredLeafTypes;
2157 if (failed(flattenTypeLeaves(
2158 argType, tables, computeFunc.getOperation(), field->get(), loweredLeafTypes, {},
true
2160 signalPassFailure();
2163 if (loweredLeafTypes.size() != 1 || !isa<MemRefType>(loweredLeafTypes.front())) {
2164 computeFunc.emitError(
2165 "execution-engine wrapper only supports felt and array<...xfelt> inputs"
2167 signalPassFailure();
2170 wrapperArgs.push_back(loweredLeafTypes.front());
2172 for (
const OutputBinding &output : *outputs) {
2173 SmallVector<Type> loweredLeafTypes;
2174 if (failed(flattenTypeLeaves(
2175 output.type, tables, computeFunc.getOperation(), field->get(), loweredLeafTypes, {},
2178 signalPassFailure();
2181 if (loweredLeafTypes.size() != 1 || !isa<MemRefType>(loweredLeafTypes.front())) {
2182 computeFunc.emitError(
2183 "execution-engine wrapper only supports felt and array<...xfelt> outputs"
2185 signalPassFailure();
2188 wrapperArgs.push_back(loweredLeafTypes.front());
2191 auto wrapper = builder.create<func::FuncOp>(
2192 computeFunc.getLoc(),
"__llzk_witgen_main",
2193 builder.getFunctionType(wrapperArgs, TypeRange {})
2195 wrapper->setAttr(LLVM::LLVMDialect::getEmitCWrapperAttrName(), builder.getUnitAttr());
2196 Block *entry = wrapper.addEntryBlock();
2197 builder.setInsertionPointToStart(entry);
2199 SmallVector<Type> loweredMainResultTypes;
2200 for (Type resultType : computeFunc.getResultTypes()) {
2201 if (failed(flattenABILeafTypes(
2202 resultType, tables, computeFunc.getOperation(), field->get(), loweredMainResultTypes
2204 signalPassFailure();
2209 SmallVector<Value> mainArgs;
2210 for (
auto [argType, wrapperArg] : llvm::zip(
2211 computeFunc.getArgumentTypes(),
2212 entry->getArguments().take_front(computeFunc.getNumArguments())
2214 if (isScalarType(argType)) {
2215 mainArgs.push_back(loadStorageScalar(builder, computeFunc.getLoc(), wrapperArg));
2218 getABILeafTypes(argType, tables, computeFunc.getOperation(), field->get());
2219 if (failed(abiLeafTypes) || abiLeafTypes->size() != 1 ||
2220 !isa<MemRefType>(abiLeafTypes->front())) {
2221 computeFunc.emitError(
"failed to derive execution-engine ABI type for main input");
2222 signalPassFailure();
2225 if (wrapperArg.getType() == abiLeafTypes->front()) {
2226 mainArgs.push_back(wrapperArg);
2228 mainArgs.push_back(builder.create<memref::CastOp>(
2229 computeFunc.getLoc(), abiLeafTypes->front(), wrapperArg
2234 auto loweredMain = builder.create<func::CallOp>(
2235 computeFunc.getLoc(), mangleFunctionName(computeFunc), loweredMainResultTypes, mainArgs
2238 LoweredValue mainResultValue {
2239 computeFunc.getResultTypes().front(),
2240 llvm::SmallVector<Value>(loweredMain.getResults().begin(), loweredMain.getResults().end())
2243 auto extractOutputSlice = [&](ArrayRef<std::string> path, Type currentType,
2244 ArrayRef<Value> leaves,
2245 auto &self) -> FailureOr<SmallVector<Value>> {
2247 return SmallVector<Value>(leaves.begin(), leaves.end());
2249 if (
auto structType = dyn_cast<component::StructType>(currentType)) {
2250 auto defLookup = structType.getDefinition(tables, computeFunc.getOperation());
2251 if (failed(defLookup)) {
2254 unsigned localCursor = 0;
2255 for (component::MemberDefOp member : defLookup->get().getMemberDefs()) {
2257 getLeafCount(member.getType(), tables, member.getOperation(), field->get());
2258 if (failed(leafCount)) {
2261 ArrayRef<Value> slice = ArrayRef<Value>(leaves).slice(localCursor, *leafCount);
2262 localCursor += *leafCount;
2263 if (member.getSymName() == path.front()) {
2264 return self(path.drop_front(), member.getType(), slice, self);
2267 computeFunc.emitError(
"failed to find struct member while wiring witgen outputs");
2270 if (
auto podType = dyn_cast<pod::PodType>(currentType)) {
2271 unsigned localCursor = 0;
2272 for (pod::RecordAttr record : podType.getRecords()) {
2274 getLeafCount(record.getType(), tables, computeFunc.getOperation(), field->get());
2275 if (failed(leafCount)) {
2278 ArrayRef<Value> slice = ArrayRef<Value>(leaves).slice(localCursor, *leafCount);
2279 localCursor += *leafCount;
2280 if (record.getName().getValue() == path.front()) {
2281 return self(path.drop_front(), record.getType(), slice, self);
2284 computeFunc.emitError(
"failed to find POD record while wiring witgen outputs");
2287 computeFunc.emitError(
"extra witness path components for non-aggregate output");
2291 auto outputArgs = entry->getArguments().drop_front(computeFunc.getNumArguments());
2292 for (
auto [output, outputMemRef] : llvm::zip(*outputs, outputArgs)) {
2293 auto slice = extractOutputSlice(
2294 output.path, mainResultValue.sourceType, mainResultValue.leaves, extractOutputSlice
2296 if (failed(slice) || slice->empty()) {
2297 wrapper.emitError(
"missing selected witness output slice while building witgen entry");
2298 signalPassFailure();
2301 if (isScalarType(output.type)) {
2303 builder, computeFunc.getLoc(),
2304 loadStorageScalar(builder, computeFunc.getLoc(), slice->front()), outputMemRef
2307 builder.create<memref::CopyOp>(computeFunc.getLoc(), slice->front(), outputMemRef);
2310 builder.create<func::ReturnOp>(computeFunc.getLoc());
2313 moduleOp->removeAttr(MAIN_ATTR_NAME);
2315 SmallVector<Operation *> toErase;
2316 for (Operation &op : moduleOp.getBody()->getOperations()) {
2317 if (!isa<func::FuncOp>(op)) {
2318 toErase.push_back(&op);
2321 for (Operation *op : toErase) {
2337 pm.addPass(mlir::createLowerAffinePass());
2340 pm.addPass(mlir::createCanonicalizerPass());
2341 pm.addPass(mlir::createCSEPass());
2345 return std::make_unique<LowerComputeToCorePass>(options);
2349 return std::make_unique<CreateWitgenEntryPass>(outputScope);
This file implements helper methods for constructing DynamicAPInts.
Apache License January AND DISTRIBUTION Definitions License shall mean the terms and conditions for and distribution as defined by Sections through of this document Licensor shall mean the copyright owner or entity authorized by the copyright owner that is granting the License Legal Entity shall mean the union of the acting entity and all other entities that control are controlled by or are under common control with that entity For the purposes of this definition control direct or to cause the direction or management of such whether by contract or including but not limited to software source documentation source
std::unique_ptr<::mlir::Pass > createFlatteningPass()
llvm::Expected< T > checkedCast(U u)
std::mt19937_64 makeDefaultValueRng(const WitgenOptions &options)
Seed an RNG for random/default witness value materialization.
OutputScope
Select the JSON scope emitted by llzk-witgen.
FailureOr< llvm::SmallVector< OutputBinding > > collectOutputBindings(component::StructDefOp mainDef, SymbolTableCollection &tables, Operation *origin, OutputScope scope)
Collect the selected output bindings for the requested scope.
void addWitgenPreparePipeline(OpPassManager &pm, const WitgenOptions &)
UninitializedBehavior
Control how witgen materializes uninitialized/default values.
std::unique_ptr< Pass > createLowerComputeToCorePass(const WitgenOptions &options)
Create the pass that lowers supported LLZK compute IR into core MLIR dialects suitable for LLVM lower...
llvm::DynamicAPInt randomFieldElement(std::mt19937_64 &rng, const Field &field)
Draw a uniformly distributed field element in [0, prime).
bool randomBoolValue(std::mt19937_64 &rng)
Draw a uniformly distributed boolean value.
std::unique_ptr< Pass > createCreateWitgenEntryPass(OutputScope outputScope)
Create the pass that synthesizes the stable llzk-witgen JIT entry wrapper.
llvm::Expected< size_t > getStaticElementCount(ShapedType type, llvm::StringRef context)
int64_t randomIndexValue(std::mt19937_64 &rng)
Draw a uniformly distributed signed index value.
ExpressionValue mod(const llvm::SMTSolverRef &solver, const ExpressionValue &lhs, const ExpressionValue &rhs)
llvm::SmallVector< StringRef > getNames(SymbolRefAttr ref)
DynamicAPInt toDynamicAPInt(StringRef str)
constexpr T checkedCast(U u) noexcept
APInt toExactWidthAPInt(const DynamicAPInt &val, unsigned bitWidth)
FailureOr< SymbolLookupResult< StructDefOp > > getMainInstanceDef(SymbolTableCollection &symbolTable, Operation *lookupFrom)
llvm::SmallSet< FieldRef, 2 > FieldSet
Typealias for a set of Fields.
ExpressionValue notOp(const llvm::SMTSolverRef &solver, const ExpressionValue &val)
mlir::LogicalResult collectFields(mlir::Operation *root, FieldSet &fields, bool silent=true)
Collects all the fields used in a circuit.
Configure one llzk-witgen execution.