LLZK 3.0.0
An open-source IR for Zero Knowledge (ZK) circuits
Loading...
Searching...
No Matches
WitgenLowering.cpp
Go to the documentation of this file.
1//===-- LLZKWitgenLoweringPass.cpp -----------------------------*- C++ -*-===//
2//
3// Part of the LLZK Project, under the Apache License v2.0.
4// See LICENSE.txt for license information.
5// Copyright 2026 Project LLZK
6// SPDX-License-Identifier: Apache-2.0
7//
8//===----------------------------------------------------------------------===//
9
10#include "WitgenLowering.h"
11
12#include "WitgenDriver.h"
13#include "WitgenUtils.h"
14#include "WitnessSelection.h"
15
28#include "llzk/Util/Compare.h"
29#include "llzk/Util/Constants.h"
31#include "llzk/Util/Field.h"
33#include "llzk/Util/Walk.h"
34
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>
49
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>
56
57#include <limits>
58
59using namespace mlir;
60
61namespace llzk::witgen {
62namespace {
63
65struct LoweredValue {
66 Type sourceType;
67 llvm::SmallVector<Value> leaves;
68};
69
71static FailureOr<std::reference_wrapper<const Field>> getModuleField(ModuleOp moduleOp) {
72 FieldSet fields;
73 if (failed(collectFields(moduleOp.getOperation(), fields, false))) {
74 moduleOp.emitError("failed to collect fields for llzk-witgen lowering");
75 return failure();
76 }
77 if (fields.size() != 1) {
78 moduleOp.emitError("llzk-witgen execution-engine lowering requires exactly one field");
79 return failure();
80 }
81 return *fields.begin();
82}
83
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() != '_') {
90 result += "__";
91 }
92 for (char c : piece) {
93 result += llvm::isAlnum(c) ? c : '_';
94 }
95 }
96 return std::string(result);
97}
98
100static Value makeIndexConstant(OpBuilder &builder, Location loc, int64_t value) {
101 return builder.create<arith::ConstantIndexOp>(loc, value).getResult();
102}
103
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)
108 );
109}
110
112static FailureOr<Type> lowerScalarType(MLIRContext *context, Type type, const Field &field) {
113 if (isa<felt::FeltType>(type)) {
114 return IntegerType::get(context, field.bitWidth());
115 }
116 if (isa<IndexType>(type)) {
117 return type;
118 }
119 if (auto intType = dyn_cast<IntegerType>(type)) {
120 if (intType.getWidth() == 1) {
121 return intType;
122 }
123 }
124 return failure();
125}
126
128static bool isScalarType(Type type) {
129 return isa<felt::FeltType, IndexType>(type) ||
130 (isa<IntegerType>(type) && llvm::cast<IntegerType>(type).getWidth() == 1);
131}
132
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
137) {
138 auto emitScalarLeaf = [&](Type leafType) {
139 auto lowered = lowerScalarType(origin->getContext(), leafType, field);
140 if (failed(lowered)) {
141 return failure();
142 }
143 if (!storage && prefixShape.empty()) {
144 out.push_back(*lowered);
145 return success();
146 }
147 llvm::SmallVector<int64_t> shape(prefixShape.begin(), prefixShape.end());
148 if (shape.empty()) {
149 shape.push_back(1);
150 }
151 out.push_back(MemRefType::get(shape, *lowered));
152 return success();
153 };
154
155 if (isScalarType(type)) {
156 return emitScalarLeaf(type);
157 }
158
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
164 );
165 }
166
167 if (auto podType = dyn_cast<pod::PodType>(type)) {
168 for (pod::RecordAttr record : podType.getRecords()) {
169 if (failed(
170 flattenTypeLeaves(record.getType(), tables, origin, field, out, prefixShape, true)
171 )) {
172 return failure();
173 }
174 }
175 return success();
176 }
177
178 if (auto structType = dyn_cast<component::StructType>(type)) {
179 auto def = structType.getDefinition(tables, origin);
180 if (failed(def)) {
181 origin->emitError("could not resolve struct type during witgen lowering");
182 return failure();
183 }
184 for (component::MemberDefOp member : def->get().getMemberDefs()) {
185 if (failed(
186 flattenTypeLeaves(member.getType(), tables, origin, field, out, prefixShape, true)
187 )) {
188 return failure();
189 }
190 }
191 return success();
192 }
193
194 origin->emitError("unsupported type in llzk-witgen lowering: ") << type;
195 return failure();
196}
197
199static MemRefType
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)
204 );
205}
206
208static LogicalResult flattenABILeafTypes(
209 Type type, SymbolTableCollection &tables, Operation *origin, const Field &field,
210 SmallVectorImpl<Type> &out, size_t prefixRank = 0, bool aggregateStorage = false
211) {
212 auto emitScalarLeaf = [&](Type leafType) {
213 auto lowered = lowerScalarType(origin->getContext(), leafType, field);
214 if (failed(lowered)) {
215 return failure();
216 }
217 if (!aggregateStorage && prefixRank == 0) {
218 out.push_back(*lowered);
219 return success();
220 }
221 SmallVector<int64_t> shape;
222 if (prefixRank == 0) {
223 shape.push_back(1);
224 } else {
225 shape.assign(prefixRank, ShapedType::kDynamic);
226 }
227 out.push_back(getStridedMemRefType(origin->getContext(), shape, *lowered));
228 return success();
229 };
230
231 if (isScalarType(type)) {
232 return emitScalarLeaf(type);
233 }
234
235 if (auto arrayType = dyn_cast<array::ArrayType>(type)) {
236 return flattenABILeafTypes(
237 arrayType.getElementType(), tables, origin, field, out, prefixRank + arrayType.getRank(),
238 true
239 );
240 }
241
242 if (auto podType = dyn_cast<pod::PodType>(type)) {
243 for (pod::RecordAttr record : podType.getRecords()) {
244 if (failed(
245 flattenABILeafTypes(record.getType(), tables, origin, field, out, prefixRank, true)
246 )) {
247 return failure();
248 }
249 }
250 return success();
251 }
252
253 if (auto structType = dyn_cast<component::StructType>(type)) {
254 auto def = structType.getDefinition(tables, origin);
255 if (failed(def)) {
256 origin->emitError("could not resolve struct type during witgen lowering");
257 return failure();
258 }
259 for (component::MemberDefOp member : def->get().getMemberDefs()) {
260 if (failed(
261 flattenABILeafTypes(member.getType(), tables, origin, field, out, prefixRank, true)
262 )) {
263 return failure();
264 }
265 }
266 return success();
267 }
268
269 origin->emitError("unsupported type in llzk-witgen lowering: ") << type;
270 return failure();
271}
272
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))) {
278 return failure();
279 }
280 return leaves.size();
281}
282
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))) {
288 return failure();
289 }
290 return leaves;
291}
292
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))) {
298 return failure();
299 }
300 return leaves;
301}
302
304static FailureOr<std::pair<size_t, size_t>> getNamedLeafSpan(
305 Type ownerType, StringRef name, SymbolTableCollection &tables, Operation *origin,
306 const Field &field
307) {
308 if (auto podType = dyn_cast<pod::PodType>(ownerType)) {
309 size_t running = 0;
310 for (pod::RecordAttr record : podType.getRecords()) {
311 auto count = getLeafCount(record.getType(), tables, origin, field);
312 if (failed(count)) {
313 return failure();
314 }
315 if (record.getName().getValue() == name) {
316 return std::pair<size_t, size_t> {running, *count};
317 }
318 running += *count;
319 }
320 }
321
322 if (auto structType = dyn_cast<component::StructType>(ownerType)) {
323 auto def = structType.getDefinition(tables, origin);
324 if (failed(def)) {
325 origin->emitError("could not resolve struct type during witgen lowering");
326 return failure();
327 }
328 size_t running = 0;
329 for (component::MemberDefOp member : def->get().getMemberDefs()) {
330 auto count = getLeafCount(member.getType(), tables, origin, field);
331 if (failed(count)) {
332 return failure();
333 }
334 if (member.getSymName() == name) {
335 return std::pair<size_t, size_t> {running, *count};
336 }
337 running += *count;
338 }
339 }
340
341 origin->emitError("could not resolve aggregate member/record @") << name;
342 return failure();
343}
344
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();
352 }
353 }
354 }
355 if (auto structType = dyn_cast<component::StructType>(ownerType)) {
356 auto def = structType.getDefinition(tables, origin);
357 if (failed(def)) {
358 origin->emitError("could not resolve struct type during witgen lowering");
359 return failure();
360 }
361 for (component::MemberDefOp member : def->get().getMemberDefs()) {
362 if (member.getSymName() == name) {
363 return member.getType();
364 }
365 }
366 }
367 origin->emitError("could not resolve aggregate member/record @") << name;
368 return failure();
369}
370
372static FailureOr<Value> createZeroMemRef(OpBuilder &builder, Location loc, MemRefType memrefType) {
373 auto elementCount = getStaticElementCount(memrefType, "witgen zero memref");
374 if (!elementCount) {
375 emitError(loc) << llvm::toString(elementCount.takeError());
376 return failure();
377 }
378 Value alloc = builder.create<memref::AllocOp>(loc, memrefType);
379 auto elementType = memrefType.getElementType();
380 Value zero;
381 if (isa<IndexType>(elementType)) {
382 zero = builder.create<arith::ConstantIndexOp>(loc, 0);
383 } else {
384 zero = builder.create<arith::ConstantOp>(
385 loc, IntegerAttr::get(llvm::cast<IntegerType>(elementType), 0)
386 );
387 }
388 auto strides = mlir::computeStrides(memrefType.getShape());
389 for (size_t flat = 0; flat < *elementCount; ++flat) {
390 auto flatSigned = checkedCast<int64_t>(flat);
391 if (!flatSigned) {
392 emitError(loc) << llvm::toString(flatSigned.takeError());
393 return failure();
394 }
395 SmallVector<Value> indices;
396 for (int64_t index : mlir::delinearize(*flatSigned, strides)) {
397 indices.push_back(makeIndexConstant(builder, loc, index));
398 }
399 builder.create<memref::StoreOp>(loc, zero, alloc, indices);
400 }
401 return alloc;
402}
403
405static FailureOr<Value> createRandomMemRef(
406 OpBuilder &builder, Location loc, MemRefType memrefType, const Field &field,
407 std::mt19937_64 &rng
408) {
409 auto elementCount = getStaticElementCount(memrefType, "witgen random memref");
410 if (!elementCount) {
411 emitError(loc) << llvm::toString(elementCount.takeError());
412 return failure();
413 }
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) {
418 auto flatSigned = checkedCast<int64_t>(flat);
419 if (!flatSigned) {
420 emitError(loc) << llvm::toString(flatSigned.takeError());
421 return failure();
422 }
423 SmallVector<Value> indices;
424 for (int64_t index : mlir::delinearize(*flatSigned, strides)) {
425 indices.push_back(makeIndexConstant(builder, loc, index));
426 }
427 if (isa<IndexType>(elementType)) {
428 auto value = randomIndexValue(rng);
429 builder.create<memref::StoreOp>(
430 loc, builder.create<arith::ConstantIndexOp>(loc, value), alloc, indices
431 );
432 continue;
433 }
434 auto intType = llvm::cast<IntegerType>(elementType);
435 if (intType.getWidth() == 1) {
436 builder.create<memref::StoreOp>(
437 loc,
438 builder.create<arith::ConstantOp>(
439 loc, IntegerAttr::get(intType, APInt(1, randomBoolValue(rng)))
440 ),
441 alloc, indices
442 );
443 continue;
444 }
445 auto candidate = randomFieldElement(rng, field);
446 builder.create<memref::StoreOp>(
447 loc,
448 builder.create<arith::ConstantOp>(
449 loc, IntegerAttr::get(intType, llzk::toExactWidthAPInt(candidate, intType.getWidth()))
450 ),
451 alloc, indices
452 );
453 }
454 return alloc;
455}
456
458static FailureOr<LoweredValue> createDefaultValue(
459 OpBuilder &builder, Location loc, Type type, SymbolTableCollection &tables, Operation *origin,
460 const Field &field, UninitializedBehavior behavior, std::mt19937_64 &rng
461) {
462 LoweredValue lowered {type, {}};
463 auto leafTypes = getLeafTypes(type, tables, origin, field);
464 if (failed(leafTypes)) {
465 return failure();
466 }
467 for (Type leafType : *leafTypes) {
468 if (behavior == UninitializedBehavior::Fail) {
469 origin->emitError(
470 "fail-mode default materialization is unsupported in witgen lowering because it would "
471 "hide uninitialized reads"
472 );
473 return failure();
474 }
475 if (behavior == UninitializedBehavior::Random) {
476 if (auto memrefType = dyn_cast<MemRefType>(leafType)) {
477 auto randomMemRef = createRandomMemRef(builder, loc, memrefType, field, rng);
478 if (failed(randomMemRef)) {
479 return failure();
480 }
481 lowered.leaves.push_back(*randomMemRef);
482 continue;
483 }
484 if (isa<IndexType>(leafType)) {
485 lowered.leaves.push_back(
486 builder.create<arith::ConstantIndexOp>(loc, randomIndexValue(rng))
487 );
488 continue;
489 }
490 auto intType = llvm::cast<IntegerType>(leafType);
491 if (intType.getWidth() == 1) {
492 lowered.leaves.push_back(builder.create<arith::ConstantOp>(
493 loc, IntegerAttr::get(intType, APInt(1, randomBoolValue(rng)))
494 ));
495 continue;
496 }
497 auto candidate = randomFieldElement(rng, field);
498 lowered.leaves.push_back(builder.create<arith::ConstantOp>(
499 loc, IntegerAttr::get(intType, llzk::toExactWidthAPInt(candidate, intType.getWidth()))
500 ));
501 continue;
502 }
503 if (auto memrefType = dyn_cast<MemRefType>(leafType)) {
504 auto zeroMemRef = createZeroMemRef(builder, loc, memrefType);
505 if (failed(zeroMemRef)) {
506 return failure();
507 }
508 lowered.leaves.push_back(*zeroMemRef);
509 continue;
510 }
511 if (isa<IndexType>(leafType)) {
512 lowered.leaves.push_back(builder.create<arith::ConstantIndexOp>(loc, 0));
513 continue;
514 }
515 lowered.leaves.push_back(builder.create<arith::ConstantOp>(
516 loc, IntegerAttr::get(llvm::cast<IntegerType>(leafType), 0)
517 ));
518 }
519 return lowered;
520}
521
523static Value normalizeWideValue(
524 OpBuilder &builder, Location loc, Value wideValue, unsigned dstWidth, const Field &field
525) {
526 auto wideType = llvm::cast<IntegerType>(wideValue.getType());
527 Value modulus = builder.create<arith::ConstantOp>(
528 loc, field.getPrimeAttr(builder.getContext(), wideType.getWidth())
529 );
530 Value reduced = builder.create<arith::RemUIOp>(loc, wideValue, modulus);
531 return builder.create<arith::TruncIOp>(
532 loc, IntegerType::get(builder.getContext(), dstWidth), reduced
533 );
534}
535
537static Value normalizeSignedWideValue(
538 OpBuilder &builder, Location loc, Value wideValue, unsigned dstWidth, const Field &field
539) {
540 auto wideType = llvm::cast<IntegerType>(wideValue.getType());
541 Value modulus = builder.create<arith::ConstantOp>(
542 loc, field.getPrimeAttr(builder.getContext(), wideType.getWidth())
543 );
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
551 );
552}
553
555static Value
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);
562 Value prime =
563 builder.create<arith::ConstantOp>(loc, field.getPrimeAttr(builder.getContext(), wideWidth));
564 Value half = builder.create<arith::ConstantOp>(
565 loc, IntegerAttr::get(feltType, toExactWidthAPInt(field.half(), width))
566 );
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);
570}
571
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);
578}
579
581static Value
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);
590}
591
593static Value
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);
600 Value modulus =
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);
605}
606
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);
613 Value modulus =
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);
617}
618
620static Value
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);
629}
630
632static Value lowerFeltInv(OpBuilder &builder, Location loc, Value operand, const Field &field) {
633 llvm::APInt exponent = toExactWidthAPInt(field.prime() - 2, field.bitWidth());
634 Value result = makeOneFelt(builder, loc, field);
635 Value base = operand;
636 for (unsigned bit = 0; bit < exponent.getBitWidth(); ++bit) {
637 if (exponent[bit]) {
638 result = lowerFeltMul(builder, loc, result, base, field);
639 }
640 if (bit + 1 < exponent.getBitWidth()) {
641 base = lowerFeltMul(builder, loc, base, base, field);
642 }
643 }
644 return result;
645}
646
648static Value
649lowerFeltDiv(OpBuilder &builder, Location loc, Value lhs, Value rhs, const Field &field) {
650 return lowerFeltMul(builder, loc, lhs, lowerFeltInv(builder, loc, rhs, field), field);
651}
652
654static Value
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))
664 );
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);
669 {
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);
674 }
675 {
676 OpBuilder::InsertionGuard guard(builder);
677 builder.setInsertionPointToStart(&ifOp.getElseRegion().front());
678 builder.create<scf::YieldOp>(loc, result);
679 }
680 result = ifOp.getResult(0);
681 if (bit + 1 < field.bitWidth()) {
682 currentBase = lowerFeltMul(builder, loc, currentBase, currentBase, field);
683 }
684 }
685 return result;
686}
687
689static Value
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);
694}
695
697static Value
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);
704}
705
707static Value
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);
714}
715
717static Value lowerFeltUnsignedDiv(OpBuilder &builder, Location loc, Value lhs, Value rhs) {
718 return builder.create<arith::DivUIOp>(loc, lhs, rhs);
719}
720
722static Value
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);
729}
730
732static Value lowerFeltUnsignedMod(OpBuilder &builder, Location loc, Value lhs, Value rhs) {
733 return builder.create<arith::RemUIOp>(loc, lhs, rhs);
734}
735
737static Value
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);
744}
745
747static Value
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()))
752 );
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))
757 );
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);
761}
762
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))
770 );
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);
774}
775
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));
783 }
784 return builder.create<memref::LoadOp>(loc, storageLeaf, indices);
785}
786
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));
794 }
795 builder.create<memref::StoreOp>(loc, scalar, storageLeaf, indices);
796}
797
799static LogicalResult copyIntoStorage(
800 OpBuilder &builder, Location loc, Type sourceType, ArrayRef<Value> destLeaves,
801 ArrayRef<Value> sourceLeaves, SymbolTableCollection &tables, Operation *origin,
802 const Field &field
803) {
804 auto leafTypes = getLeafTypes(sourceType, tables, origin, field);
805 if (failed(leafTypes)) {
806 return failure();
807 }
808 if (destLeaves.size() != sourceLeaves.size() || destLeaves.size() != leafTypes->size()) {
809 origin->emitError("flattened leaf mismatch while copying aggregate storage");
810 return failure();
811 }
812 for (auto [leafType, destLeaf, srcLeaf] : llvm::zip(*leafTypes, destLeaves, sourceLeaves)) {
813 if (isa<MemRefType>(leafType)) {
814 builder.create<memref::CopyOp>(loc, srcLeaf, destLeaf);
815 continue;
816 }
817 storeStorageScalar(builder, loc, srcLeaf, destLeaf);
818 }
819 return success();
820}
821
823static FailureOr<LoweredValue> readNamedAggregateValue(
824 OpBuilder &builder, Location loc, Type ownerType, StringRef name, const LoweredValue &owner,
825 SymbolTableCollection &tables, Operation *origin, const Field &field
826) {
827 auto subType = getNamedSubType(ownerType, name, tables, origin);
828 if (failed(subType)) {
829 return failure();
830 }
831 auto span = getNamedLeafSpan(ownerType, name, tables, origin, field);
832 if (failed(span)) {
833 return failure();
834 }
835 LoweredValue result {*subType, {}};
836 auto leafTypes = getLeafTypes(*subType, tables, origin, field);
837 if (failed(leafTypes)) {
838 return failure();
839 }
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);
844 } else {
845 result.leaves.push_back(loadStorageScalar(builder, loc, leafValue));
846 }
847 }
848 return result;
849}
850
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
855) {
856 auto subType = getNamedSubType(ownerType, name, tables, origin);
857 if (failed(subType)) {
858 return failure();
859 }
860 auto span = getNamedLeafSpan(ownerType, name, tables, origin, field);
861 if (failed(span)) {
862 return failure();
863 }
864 return copyIntoStorage(
865 builder, loc, *subType, ArrayRef<Value>(owner.leaves).slice(span->first, span->second),
866 value.leaves, tables, origin, field
867 );
868}
869
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;
877 auto indexedRank = checkedCast<int64_t>(outerIndices.size());
878 if (!indexedRank) {
879 emitError(loc) << llvm::toString(indexedRank.takeError());
880 return failure();
881 }
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);
887 }
888 for (int64_t dim = *indexedRank; dim < sourceType.getRank(); ++dim) {
889 mixedOffsets.push_back(builder.getIndexAttr(0));
890 }
891 for (int64_t dim = 0; dim < *indexedRank; ++dim) {
892 mixedSizes.push_back(builder.getIndexAttr(1));
893 }
894 for (int64_t dim = *indexedRank; dim < sourceType.getRank(); ++dim) {
895 mixedSizes.push_back(memref::getMixedSize(builder, loc, source, dim));
896 }
897 for (int64_t dim = 0; dim < sourceType.getRank(); ++dim) {
898 mixedStrides.push_back(builder.getIndexAttr(1));
899 }
900 SmallVector<int64_t> desiredShape;
901 auto reserveSize = checkedCast<size_t>(sourceType.getRank() - *indexedRank);
902 if (!reserveSize) {
903 emitError(loc) << llvm::toString(reserveSize.takeError());
904 return failure();
905 }
906 desiredShape.reserve(*reserveSize);
907 for (int64_t dim = *indexedRank; dim < sourceType.getRank(); ++dim) {
908 auto dimIndex = checkedCast<size_t>(dim);
909 if (!dimIndex) {
910 emitError(loc) << llvm::toString(dimIndex.takeError());
911 return failure();
912 }
913 if (auto attr = llvm::dyn_cast<Attribute>(mixedSizes[*dimIndex])) {
914 desiredShape.push_back(llvm::cast<IntegerAttr>(attr).getInt());
915 } else {
916 desiredShape.push_back(ShapedType::kDynamic);
917 }
918 }
919 if (desiredShape.empty()) {
920 desiredShape.push_back(1);
921 }
922 auto resultType = llvm::cast<MemRefType>(memref::SubViewOp::inferRankReducedResultType(
923 desiredShape, sourceType, mixedOffsets, mixedSizes, mixedStrides
924 ));
925 auto op = builder.create<memref::SubViewOp>(
926 loc, resultType, source, mixedOffsets, mixedSizes, mixedStrides
927 );
928 return success(op.getResult());
929}
930
932static FailureOr<LoweredValue> readArrayElement(
933 OpBuilder &builder, Location loc, array::ArrayType arrayType, const LoweredValue &arrayValue,
934 ArrayRef<Value> indices
935) {
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)
941 );
942 return result;
943 }
944
945 for (Value sourceLeaf : arrayValue.leaves) {
946 auto subview = createElementSubview(builder, loc, sourceLeaf, indices);
947 if (failed(subview)) {
948 return failure();
949 }
950 result.leaves.push_back(*subview);
951 }
952 return result;
953}
954
956static LogicalResult writeArrayElement(
957 OpBuilder &builder, Location loc, array::ArrayType arrayType, LoweredValue &arrayValue,
958 ArrayRef<Value> indices, const LoweredValue &elementValue
959) {
960 Type elementType = arrayType.getElementType();
961 if (isScalarType(elementType)) {
962 builder.create<memref::StoreOp>(
963 loc, elementValue.leaves.front(), arrayValue.leaves.front(), indices
964 );
965 return success();
966 }
967
968 for (auto [destLeaf, srcLeaf] : llvm::zip(arrayValue.leaves, elementValue.leaves)) {
969 auto subview = createElementSubview(builder, loc, destLeaf, indices);
970 if (failed(subview)) {
971 return failure();
972 }
973 builder.create<memref::CopyOp>(loc, srcLeaf, *subview);
974 }
975 return success();
976}
977
979static LogicalResult appendFlatLeavesToTypes(
980 OpBuilder &builder, Location loc, const LoweredValue &value, ArrayRef<Type> targetLeafTypes,
981 SmallVectorImpl<Value> &out, Operation *origin
982) {
983 if (targetLeafTypes.size() != value.leaves.size()) {
984 origin->emitError("flattened leaf mismatch during call lowering");
985 return failure();
986 }
987 for (auto [leafValue, leafType] : llvm::zip(value.leaves, targetLeafTypes)) {
988 if (leafValue.getType() == leafType) {
989 out.push_back(leafValue);
990 continue;
991 }
992 if (isa<MemRefType>(leafValue.getType()) && isa<MemRefType>(leafType)) {
993 out.push_back(builder.create<memref::CastOp>(loc, leafType, leafValue));
994 continue;
995 }
996 origin->emitError("lowered leaf type mismatch during call lowering");
997 return failure();
998 }
999 return success();
1000}
1001
1003class BodyLowerer {
1004public:
1006 BodyLowerer(
1007 ModuleOp mod, SymbolTableCollection &symbolTables, const Field &moduleField,
1008 const WitgenOptions &options
1009 )
1010 : moduleOp(mod), tables(symbolTables), field(moduleField),
1011 uninitializedBehavior(options.uninitializedBehavior), rng(makeDefaultValueRng(options)) {}
1012
1014 FailureOr<func::FuncOp> lowerFunction(function::FuncDefOp funcOp) {
1015 if (funcOp.isExternal()) {
1016 funcOp.emitError("execution-engine backend does not lower extern functions");
1017 return failure();
1018 }
1019 if (!funcOp.getBody().hasOneBlock()) {
1020 funcOp.emitError("execution-engine backend only supports single-block functions");
1021 return failure();
1022 }
1023
1024 SmallVector<Type> loweredArgTypes;
1025 for (Type argType : funcOp.getArgumentTypes()) {
1026 if (failed(
1027 flattenABILeafTypes(argType, tables, funcOp.getOperation(), field, loweredArgTypes)
1028 )) {
1029 return failure();
1030 }
1031 }
1032 SmallVector<Type> loweredResultTypes;
1033 for (Type resultType : funcOp.getResultTypes()) {
1034 if (failed(flattenABILeafTypes(
1035 resultType, tables, funcOp.getOperation(), field, loweredResultTypes
1036 ))) {
1037 return failure();
1038 }
1039 }
1040
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)
1046 );
1047 Block *entry = loweredFunc.addEntryBlock();
1048 OpBuilder builder(entry, entry->begin());
1049
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();
1057 return failure();
1058 }
1059 LoweredValue lowered {argType, {}};
1060 lowered.leaves.append(
1061 entry->getArguments().begin() + cursor,
1062 entry->getArguments().begin() + cursor + *leafCount
1063 );
1064 cursor += *leafCount;
1065 valueMap[arg] = std::move(lowered);
1066 }
1067
1068 if (failed(lowerBlock(builder, funcOp.getBody().front(), valueMap))) {
1069 loweredFunc.erase();
1070 return failure();
1071 }
1072 return loweredFunc;
1073 }
1074
1075private:
1076 ModuleOp moduleOp;
1077 SymbolTableCollection &tables;
1078 const Field &field;
1079 UninitializedBehavior uninitializedBehavior;
1080 std::mt19937_64 rng;
1081
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");
1088 return failure();
1089 }
1090 return it->second;
1091 }
1092
1094 FailureOr<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");
1100 return failure();
1101 }
1102 return lowered->leaves.front();
1103 }
1104
1106 LogicalResult
1107 lowerBlock(OpBuilder &builder, Block &block, DenseMap<Value, LoweredValue> &valueMap) {
1108 for (Operation &op : block) {
1109 if (failed(lowerOperation(builder, op, valueMap))) {
1110 return failure();
1111 }
1112 }
1113 return success();
1114 }
1115
1117 FailureOr<Value>
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;
1123 break;
1125 predicate = arith::CmpIPredicate::ne;
1126 break;
1128 predicate = arith::CmpIPredicate::ult;
1129 break;
1131 predicate = arith::CmpIPredicate::ule;
1132 break;
1134 predicate = arith::CmpIPredicate::ugt;
1135 break;
1137 predicate = arith::CmpIPredicate::uge;
1138 break;
1139 }
1140 return builder.create<arith::CmpIOp>(loc, predicate, lhs, rhs).getResult();
1141 }
1142
1144 LogicalResult
1145 lowerOperation(OpBuilder &builder, Operation &op, DenseMap<Value, LoweredValue> &valueMap) {
1146 Location loc = op.getLoc();
1147
1148 auto bind = [&](Value result, LoweredValue lowered) {
1149 valueMap[result] = std::move(lowered);
1150 return success();
1151 };
1152
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()
1161 ))) {
1162 return failure();
1163 }
1164 }
1165 builder.create<func::ReturnOp>(loc, results);
1166 return success();
1167 }
1168
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()
1177 ))) {
1178 return failure();
1179 }
1180 }
1181 builder.create<scf::YieldOp>(loc, results);
1182 return success();
1183 }
1184 if (auto conditionOp = dyn_cast<scf::ConditionOp>(op)) {
1185 auto condition =
1186 lookupScalar(conditionOp.getCondition(), valueMap, conditionOp.getOperation());
1187 if (failed(condition)) {
1188 return failure();
1189 }
1190 SmallVector<Value> results;
1191 for (Value operand : conditionOp.getArgs()) {
1192 auto lowered = lookup(operand, valueMap, conditionOp.getOperation());
1193 auto leafTypes =
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()
1198 ))) {
1199 return failure();
1200 }
1201 }
1202 builder.create<scf::ConditionOp>(loc, *condition, results);
1203 return success();
1204 }
1205
1206 if (auto constantOp = dyn_cast<arith::ConstantOp>(op)) {
1207 Operation *clone = builder.clone(op);
1208 return bind(
1209 constantOp.getResult(), LoweredValue {constantOp.getType(), {clone->getResult(0)}}
1210 );
1211 }
1212
1213 if (auto feltConst = dyn_cast<felt::FeltConstantOp>(op)) {
1214 auto intType = IntegerType::get(builder.getContext(), field.bitWidth());
1215 // Reduce into the field first, then build an APInt with the exact storage width.
1216 auto constVal = toDynamicAPInt(feltConst.getValue().getValue());
1217 auto modVal = constVal % field.prime();
1218 auto intVal = llzk::toExactWidthAPInt(modVal, field.bitWidth());
1219 Value lowered = builder.create<arith::ConstantOp>(loc, IntegerAttr::get(intType, intVal));
1220 return bind(feltConst.getResult(), LoweredValue {feltConst.getType(), {lowered}});
1221 }
1222
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
1227 );
1228 if (failed(lowered)) {
1229 return failure();
1230 }
1231 return bind(nondetOp.getResult(), std::move(*lowered));
1232 }
1233
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)) {
1238 return failure();
1239 }
1240 return bind(
1241 addOp.getResult(),
1242 LoweredValue {addOp.getType(), {lowerFeltAdd(builder, loc, *lhs, *rhs, field)}}
1243 );
1244 }
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)) {
1249 return failure();
1250 }
1251 return bind(
1252 powOp.getResult(),
1253 LoweredValue {powOp.getType(), {lowerFeltPow(builder, loc, *lhs, *rhs, field)}}
1254 );
1255 }
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)) {
1260 return failure();
1261 }
1262 return bind(
1263 andOp.getResult(),
1264 LoweredValue {andOp.getType(), {builder.create<arith::AndIOp>(loc, *lhs, *rhs)}}
1265 );
1266 }
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)) {
1271 return failure();
1272 }
1273 return bind(
1274 orOp.getResult(),
1275 LoweredValue {orOp.getType(), {lowerFeltOr(builder, loc, *lhs, *rhs, field)}}
1276 );
1277 }
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)) {
1282 return failure();
1283 }
1284 return bind(
1285 xorOp.getResult(),
1286 LoweredValue {xorOp.getType(), {lowerFeltXor(builder, loc, *lhs, *rhs, field)}}
1287 );
1288 }
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)) {
1293 return failure();
1294 }
1295 return bind(
1296 subOp.getResult(),
1297 LoweredValue {subOp.getType(), {lowerFeltSub(builder, loc, *lhs, *rhs, field)}}
1298 );
1299 }
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)) {
1304 return failure();
1305 }
1306 return bind(
1307 mulOp.getResult(),
1308 LoweredValue {mulOp.getType(), {lowerFeltMul(builder, loc, *lhs, *rhs, field)}}
1309 );
1310 }
1311 if (auto negOp = dyn_cast<felt::NegFeltOp>(op)) {
1312 auto operand = lookupScalar(negOp.getOperand(), valueMap, negOp.getOperation());
1313 if (failed(operand)) {
1314 return failure();
1315 }
1316 return bind(
1317 negOp.getResult(),
1318 LoweredValue {negOp.getType(), {lowerFeltNeg(builder, loc, *operand, field)}}
1319 );
1320 }
1321 if (auto invOp = dyn_cast<felt::InvFeltOp>(op)) {
1322 auto operand = lookupScalar(invOp.getOperand(), valueMap, invOp.getOperation());
1323 if (failed(operand)) {
1324 return failure();
1325 }
1326 return bind(
1327 invOp.getResult(),
1328 LoweredValue {invOp.getType(), {lowerFeltInv(builder, loc, *operand, field)}}
1329 );
1330 }
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)) {
1335 return failure();
1336 }
1337 return bind(
1338 divOp.getResult(),
1339 LoweredValue {divOp.getType(), {lowerFeltDiv(builder, loc, *lhs, *rhs, field)}}
1340 );
1341 }
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)) {
1346 return failure();
1347 }
1348 assertNonZeroFelt(builder, loc, *rhs, "felt.uintdiv divisor must be non-zero");
1349 return bind(
1350 uintDivOp.getResult(),
1351 LoweredValue {uintDivOp.getType(), {lowerFeltUnsignedDiv(builder, loc, *lhs, *rhs)}}
1352 );
1353 }
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)) {
1358 return failure();
1359 }
1360 assertNonZeroFelt(builder, loc, *rhs, "felt.sintdiv divisor must be non-zero");
1361 return bind(
1362 sintDivOp.getResult(),
1363 LoweredValue {sintDivOp.getType(), {lowerFeltSignedDiv(builder, loc, *lhs, *rhs, field)}}
1364 );
1365 }
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)) {
1370 return failure();
1371 }
1372 assertNonZeroFelt(builder, loc, *rhs, "felt.umod divisor must be non-zero");
1373 return bind(
1374 umodOp.getResult(),
1375 LoweredValue {umodOp.getType(), {lowerFeltUnsignedMod(builder, loc, *lhs, *rhs)}}
1376 );
1377 }
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)) {
1382 return failure();
1383 }
1384 assertNonZeroFelt(builder, loc, *rhs, "felt.smod divisor must be non-zero");
1385 return bind(
1386 smodOp.getResult(),
1387 LoweredValue {smodOp.getType(), {lowerFeltSignedMod(builder, loc, *lhs, *rhs, field)}}
1388 );
1389 }
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)) {
1394 return failure();
1395 }
1396 return bind(
1397 shrOp.getResult(),
1398 LoweredValue {shrOp.getType(), {lowerFeltShr(builder, loc, *lhs, *rhs, field)}}
1399 );
1400 }
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)) {
1405 return failure();
1406 }
1407 return bind(
1408 shlOp.getResult(),
1409 LoweredValue {shlOp.getType(), {lowerFeltShl(builder, loc, *lhs, *rhs, field)}}
1410 );
1411 }
1412 if (auto notOp = dyn_cast<felt::NotFeltOp>(op)) {
1413 auto operand = lookupScalar(notOp.getOperand(), valueMap, notOp.getOperation());
1414 if (failed(operand)) {
1415 return failure();
1416 }
1417 return bind(
1418 notOp.getResult(),
1419 LoweredValue {notOp.getType(), {lowerFeltNot(builder, loc, *operand, field)}}
1420 );
1421 }
1422
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)) {
1427 return failure();
1428 }
1429 auto lowered = lowerFeltCmp(builder, loc, cmpOp, *lhs, *rhs);
1430 if (failed(lowered)) {
1431 return failure();
1432 }
1433 return bind(cmpOp.getResult(), LoweredValue {cmpOp.getType(), {*lowered}});
1434 }
1435 if (auto assertOp = dyn_cast<boolean::AssertOp>(op)) {
1436 auto condition = lookupScalar(assertOp.getCondition(), valueMap, assertOp.getOperation());
1437 if (failed(condition)) {
1438 return failure();
1439 }
1440 builder.create<cf::AssertOp>(
1441 loc, *condition, assertOp.getMsg() ? assertOp.getMsg()->str() : "bool.assert failed"
1442 );
1443 return success();
1444 }
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)) {
1449 return failure();
1450 }
1451 return bind(
1452 andOp.getResult(),
1453 LoweredValue {andOp.getType(), {builder.create<arith::AndIOp>(loc, *lhs, *rhs)}}
1454 );
1455 }
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)) {
1460 return failure();
1461 }
1462 return bind(
1463 orOp.getResult(),
1464 LoweredValue {orOp.getType(), {builder.create<arith::OrIOp>(loc, *lhs, *rhs)}}
1465 );
1466 }
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)) {
1471 return failure();
1472 }
1473 return bind(
1474 xorOp.getResult(),
1475 LoweredValue {xorOp.getType(), {builder.create<arith::XOrIOp>(loc, *lhs, *rhs)}}
1476 );
1477 }
1478 if (auto notOp = dyn_cast<boolean::NotBoolOp>(op)) {
1479 auto operand = lookupScalar(notOp.getOperand(), valueMap, notOp.getOperation());
1480 if (failed(operand)) {
1481 return failure();
1482 }
1483 Value one = builder.create<arith::ConstantOp>(
1484 loc, IntegerAttr::get(IntegerType::get(builder.getContext(), 1), 1)
1485 );
1486 return bind(
1487 notOp.getResult(),
1488 LoweredValue {notOp.getType(), {builder.create<arith::XOrIOp>(loc, *operand, one)}}
1489 );
1490 }
1491
1492 if (auto intToFelt = dyn_cast<cast::IntToFeltOp>(op)) {
1493 auto operand = lookupScalar(intToFelt.getValue(), valueMap, intToFelt.getOperation());
1494 if (failed(operand)) {
1495 return failure();
1496 }
1497 auto dstType = IntegerType::get(builder.getContext(), field.bitWidth());
1498 Value lowered;
1499 if (isa<IndexType>((*operand).getType())) {
1500 lowered = builder.create<arith::IndexCastUIOp>(loc, dstType, *operand);
1501 } else {
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);
1507 } else {
1508 lowered = *operand;
1509 }
1510 }
1511 return bind(intToFelt.getResult(), LoweredValue {intToFelt.getType(), {lowered}});
1512 }
1513 if (auto feltToIndex = dyn_cast<cast::FeltToIndexOp>(op)) {
1514 auto operand = lookupScalar(feltToIndex.getValue(), valueMap, feltToIndex.getOperation());
1515 if (failed(operand)) {
1516 return failure();
1517 }
1518 return bind(
1519 feltToIndex.getResult(),
1520 LoweredValue {
1521 feltToIndex.getType(),
1522 {builder.create<arith::IndexCastUIOp>(loc, builder.getIndexType(), *operand)}
1523 }
1524 );
1525 }
1526
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
1531 );
1532 if (failed(lowered)) {
1533 return failure();
1534 }
1535 return bind(structNewOp.getResult(), std::move(*lowered));
1536 }
1537 if (auto readMemberOp = dyn_cast<component::MemberReadOp>(op)) {
1538 auto componentValue =
1539 lookup(readMemberOp.getComponent(), valueMap, readMemberOp.getOperation());
1540 if (failed(componentValue)) {
1541 return failure();
1542 }
1543 auto lowered = readNamedAggregateValue(
1544 builder, loc, readMemberOp.getComponent().getType(), readMemberOp.getMemberName(),
1545 *componentValue, tables, readMemberOp.getOperation(), field
1546 );
1547 if (failed(lowered)) {
1548 return failure();
1549 }
1550 return bind(readMemberOp.getResult(), std::move(*lowered));
1551 }
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)) {
1557 return failure();
1558 }
1559 return writeNamedAggregateValue(
1560 builder, loc, writeMemberOp.getComponent().getType(), writeMemberOp.getMemberName(),
1561 valueMap[writeMemberOp.getComponent()], *memberValue, tables,
1562 writeMemberOp.getOperation(), field
1563 );
1564 }
1565
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
1570 );
1571 if (failed(lowered)) {
1572 return failure();
1573 }
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
1579 ))) {
1580 return failure();
1581 }
1582 }
1583 return bind(newPodOp.getResult(), std::move(*lowered));
1584 }
1585 if (auto readPodOp = dyn_cast<pod::ReadPodOp>(op)) {
1586 auto podValue = lookup(readPodOp.getPodRef(), valueMap, readPodOp.getOperation());
1587 if (failed(podValue)) {
1588 return failure();
1589 }
1590 auto lowered = readNamedAggregateValue(
1591 builder, loc, readPodOp.getPodRef().getType(), readPodOp.getRecordName(), *podValue,
1592 tables, readPodOp.getOperation(), field
1593 );
1594 if (failed(lowered)) {
1595 return failure();
1596 }
1597 return bind(readPodOp.getResult(), std::move(*lowered));
1598 }
1599 if (auto writePodOp = dyn_cast<pod::WritePodOp>(op)) {
1600 auto recordValue = lookup(writePodOp.getValue(), valueMap, writePodOp.getOperation());
1601 if (failed(recordValue)) {
1602 return failure();
1603 }
1604 return writeNamedAggregateValue(
1605 builder, loc, writePodOp.getPodRef().getType(), writePodOp.getRecordName(),
1606 valueMap[writePodOp.getPodRef()], *recordValue, tables, writePodOp.getOperation(), field
1607 );
1608 }
1609
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
1614 );
1615 if (failed(lowered)) {
1616 return failure();
1617 }
1618 if (!arrayNewOp.getElements().empty()) {
1619 auto elementCount = checkedCast<size_t>(arrayNewOp.getType().getNumElements());
1620 if (!elementCount) {
1621 arrayNewOp.emitError() << llvm::toString(elementCount.takeError());
1622 return failure();
1623 }
1624 if (arrayNewOp.getElements().size() != *elementCount) {
1625 arrayNewOp.emitError("expected one explicit element per array slot in witgen lowering");
1626 return failure();
1627 }
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)) {
1632 return failure();
1633 }
1634 SmallVector<Value> indices;
1635 auto strides = mlir::computeStrides(shape);
1636 auto flatSigned = checkedCast<int64_t>(flatIndex);
1637 if (!flatSigned) {
1638 arrayNewOp.emitError() << llvm::toString(flatSigned.takeError());
1639 return failure();
1640 }
1641 for (int64_t index : mlir::delinearize(*flatSigned, strides)) {
1642 indices.push_back(makeIndexConstant(builder, loc, index));
1643 }
1644 if (failed(writeArrayElement(
1645 builder, loc, arrayNewOp.getType(), *lowered, indices, *elementValue
1646 ))) {
1647 return failure();
1648 }
1649 }
1650 }
1651 return bind(arrayNewOp.getResult(), std::move(*lowered));
1652 }
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)) {
1658 return failure();
1659 }
1660 indices.push_back(*loweredIndex);
1661 }
1662 auto arrayValue = lookup(readArrayOp.getArrRef(), valueMap, readArrayOp.getOperation());
1663 if (failed(arrayValue)) {
1664 return failure();
1665 }
1666 auto lowered = readArrayElement(
1667 builder, loc, llvm::cast<array::ArrayType>(readArrayOp.getArrRef().getType()),
1668 *arrayValue, indices
1669 );
1670 if (failed(lowered)) {
1671 return failure();
1672 }
1673 return bind(readArrayOp.getResult(), std::move(*lowered));
1674 }
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)) {
1680 return failure();
1681 }
1682 indices.push_back(*loweredIndex);
1683 }
1684 auto elementValue = lookup(writeArrayOp.getRvalue(), valueMap, writeArrayOp.getOperation());
1685 if (failed(elementValue)) {
1686 return failure();
1687 }
1688 return writeArrayElement(
1689 builder, loc, llvm::cast<array::ArrayType>(writeArrayOp.getArrRef().getType()),
1690 valueMap[writeArrayOp.getArrRef()], indices, *elementValue
1691 );
1692 }
1693
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)) {
1698 return failure();
1699 }
1700 return bind(
1701 cmpiOp.getResult(),
1702 LoweredValue {
1703 cmpiOp.getType(),
1704 {builder.create<arith::CmpIOp>(loc, cmpiOp.getPredicate(), *lhs, *rhs)}
1705 }
1706 );
1707 }
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)) {
1713 return failure();
1714 }
1715 return bind(
1716 selectOp.getResult(),
1717 LoweredValue {
1718 selectOp.getType(),
1719 {builder.create<arith::SelectOp>(loc, *cond, *trueValue, *falseValue)}
1720 }
1721 );
1722 }
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)) {
1727 return failure();
1728 }
1729 return bind(
1730 addiOp.getResult(),
1731 LoweredValue {addiOp.getType(), {builder.create<arith::AddIOp>(loc, *lhs, *rhs)}}
1732 );
1733 }
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)) {
1738 return failure();
1739 }
1740 return bind(
1741 subiOp.getResult(),
1742 LoweredValue {subiOp.getType(), {builder.create<arith::SubIOp>(loc, *lhs, *rhs)}}
1743 );
1744 }
1745
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");
1749 return failure();
1750 }
1751 auto *callable = callOp.resolveCallableInTable(&tables);
1752 auto callee = dyn_cast_or_null<function::FuncDefOp>(callable);
1753 if (!callee) {
1754 callOp.emitError("failed to resolve callee during execution-engine lowering");
1755 return failure();
1756 }
1757 SmallVector<Type> resultTypes;
1758 for (Type resultType : callOp.getResultTypes()) {
1759 if (failed(
1760 flattenABILeafTypes(resultType, tables, callOp.getOperation(), field, resultTypes)
1761 )) {
1762 return failure();
1763 }
1764 }
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()
1772 ))) {
1773 return failure();
1774 }
1775 }
1776 auto loweredCall =
1777 builder.create<func::CallOp>(loc, mangleFunctionName(callee), resultTypes, flatArgs);
1778 auto loweredCallResults = loweredCall.getResults();
1779 size_t totalResults = loweredCallResults.size();
1780 size_t cursor = 0;
1781 for (auto [oldResult, oldType] : llvm::zip(callOp.getResults(), callOp.getResultTypes())) {
1782 auto leafCount = getLeafCount(oldType, tables, callOp.getOperation(), field);
1783 if (failed(leafCount)) {
1784 return failure();
1785 }
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");
1790 return failure();
1791 }
1792 LoweredValue lowered {oldType, {}};
1793 lowered.leaves.append(
1794 loweredCallResults.begin() + static_cast<ptrdiff_t>(cursor),
1795 loweredCallResults.begin() + static_cast<ptrdiff_t>(nextCursor)
1796 );
1797 valueMap[oldResult] = std::move(lowered);
1798 cursor = nextCursor;
1799 }
1800 return success();
1801 }
1802
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()
1812 ))) {
1813 return failure();
1814 }
1815 auto count = getLeafCount(initType, tables, whileOp.getOperation(), field);
1816 if (failed(count)) {
1817 return failure();
1818 }
1819 beforeLeafCounts.push_back(*count);
1820 }
1821
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)) {
1828 return failure();
1829 }
1830 loweredResultTypes.append(leafTypes->begin(), leafTypes->end());
1831 resultLeafCounts.push_back(*count);
1832 }
1833
1834 auto mapRegionArguments = [&](auto oldArgs, auto oldTypes, auto leafCounts, auto newArgs,
1835 StringRef overflowMessage,
1836 DenseMap<Value, LoweredValue> &regionMap) -> LogicalResult {
1837 size_t totalArgs = newArgs.size();
1838 size_t cursor = 0;
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);
1844 return failure();
1845 }
1846 LoweredValue lowered {oldType, {}};
1847 lowered.leaves.append(
1848 newArgs.begin() + llzk::checkedCast<ptrdiff_t>(cursor),
1849 newArgs.begin() + llzk::checkedCast<ptrdiff_t>(nextCursor)
1850 );
1851 regionMap[oldArg] = std::move(lowered);
1852 cursor = nextCursor;
1853 }
1854 return success();
1855 };
1856
1857 LogicalResult whileLoweringStatus = success();
1858 auto newWhile = builder.create<scf::WhileOp>(
1859 loc, loweredResultTypes, initArgs,
1860 [&](OpBuilder &regionBuilder, Location /*regionLoc*/, 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",
1865 beforeMap
1866 )) ||
1867 failed(lowerBlock(regionBuilder, whileOp.getBefore().front(), beforeMap))) {
1868 whileLoweringStatus = failure();
1869 }
1870 }, [&](OpBuilder &regionBuilder, Location /*regionLoc*/, 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
1875 )) ||
1876 failed(lowerBlock(regionBuilder, whileOp.getAfter().front(), afterMap))) {
1877 whileLoweringStatus = failure();
1878 }
1879 }
1880 );
1881 if (failed(whileLoweringStatus)) {
1882 newWhile.erase();
1883 return failure();
1884 }
1885
1886 auto newWhileResults = newWhile.getResults();
1887 size_t totalResults = newWhileResults.size();
1888 size_t cursor = 0;
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");
1895 return failure();
1896 }
1897 LoweredValue lowered {oldType, {}};
1898 lowered.leaves.append(
1899 newWhileResults.begin() + llzk::checkedCast<ptrdiff_t>(cursor),
1900 newWhileResults.begin() + llzk::checkedCast<ptrdiff_t>(nextCursor)
1901 );
1902 valueMap[oldResult] = std::move(lowered);
1903 cursor = nextCursor;
1904 }
1905 return success();
1906 }
1907
1908 if (auto ifOp = dyn_cast<scf::IfOp>(op)) {
1909 auto condition = lookupScalar(ifOp.getCondition(), valueMap, ifOp.getOperation());
1910 if (failed(condition)) {
1911 return failure();
1912 }
1913
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)) {
1920 return failure();
1921 }
1922 loweredResultTypes.append(leafTypes->begin(), leafTypes->end());
1923 resultLeafCounts.push_back(*count);
1924 }
1925
1926 auto newIf = builder.create<scf::IfOp>(
1927 loc, loweredResultTypes, *condition, true, !ifOp.getElseRegion().empty()
1928 );
1929
1930 {
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))) {
1934 newIf.erase();
1935 return failure();
1936 }
1937 }
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))) {
1942 newIf.erase();
1943 return failure();
1944 }
1945 }
1946 auto newIfResults = newIf.getResults();
1947 size_t totalResults = newIfResults.size();
1948 size_t cursor = 0;
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");
1955 return failure();
1956 }
1957 LoweredValue lowered {oldType, {}};
1958 lowered.leaves.append(
1959 newIfResults.begin() + llzk::checkedCast<ptrdiff_t>(cursor),
1960 newIfResults.begin() + llzk::checkedCast<ptrdiff_t>(nextCursor)
1961 );
1962 valueMap[oldResult] = std::move(lowered);
1963 cursor = nextCursor;
1964 }
1965 return success();
1966 }
1967
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)) {
1973 return failure();
1974 }
1975
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()
1984 ))) {
1985 return failure();
1986 }
1987 auto count = getLeafCount(resultType, tables, forOp.getOperation(), field);
1988 if (failed(count)) {
1989 return failure();
1990 }
1991 initLeafCounts.push_back(*count);
1992 }
1993
1994 auto newFor = builder.create<scf::ForOp>(loc, *lb, *ub, *step, initArgs);
1995 if (Attribute unsignedCmpAttr = forOp->getAttr("unsignedCmp")) {
1996 newFor->setAttr("unsignedCmp", unsignedCmpAttr);
1997 }
1998 DenseMap<Value, LoweredValue> bodyMap(valueMap.begin(), valueMap.end());
1999 bodyMap[forOp.getInductionVar()] =
2000 LoweredValue {forOp.getInductionVar().getType(), {newFor.getInductionVar()}};
2001 {
2002 auto newForIterArgs = newFor.getRegionIterArgs();
2003 size_t totalIterArgs = newForIterArgs.size();
2004 size_t cursor = 0;
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");
2011 return failure();
2012 }
2013 LoweredValue lowered {oldType, {}};
2014 lowered.leaves.append(
2015 newForIterArgs.begin() + static_cast<ptrdiff_t>(cursor),
2016 newForIterArgs.begin() + static_cast<ptrdiff_t>(nextCursor)
2017 );
2018 bodyMap[oldIterArg] = std::move(lowered);
2019 cursor = nextCursor;
2020 }
2021 }
2022
2023 newFor.getBody()->clear();
2024 OpBuilder bodyBuilder = OpBuilder::atBlockBegin(newFor.getBody());
2025 if (failed(lowerBlock(bodyBuilder, *forOp.getBody(), bodyMap))) {
2026 return failure();
2027 }
2028
2029 {
2030 auto newForResults = newFor.getResults();
2031 size_t totalForResults = newForResults.size();
2032 size_t cursor = 0;
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");
2039 return failure();
2040 }
2041 LoweredValue lowered {oldType, {}};
2042 lowered.leaves.append(
2043 newForResults.begin() + static_cast<ptrdiff_t>(cursor),
2044 newForResults.begin() + static_cast<ptrdiff_t>(nextCursor)
2045 );
2046 valueMap[oldResult] = std::move(lowered);
2047 cursor = nextCursor;
2048 }
2049 }
2050 return success();
2051 }
2052
2053 op.emitError("unsupported operation in execution-engine lowering: ") << op.getName();
2054 return failure();
2055 }
2056};
2057
2059class LowerComputeToCorePass : public PassWrapper<LowerComputeToCorePass, OperationPass<ModuleOp>> {
2060public:
2061 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LowerComputeToCorePass)
2062
2063 explicit LowerComputeToCorePass(const WitgenOptions &opts) : options(opts) {}
2064
2066 StringRef getArgument() const final { return "llzk-lower-compute-to-core"; }
2067
2069 StringRef getDescription() const final {
2070 return "Lower LLZK compute IR to func/arith/cf/scf/memref";
2071 }
2072
2074 StringRef getName() const override { return "LowerComputeToCorePass"; }
2075
2077 void runOnOperation() override {
2078 ModuleOp moduleOp = getOperation();
2079 auto field = getModuleField(moduleOp);
2080 if (failed(field)) {
2081 signalPassFailure();
2082 return;
2083 }
2084
2085 SymbolTableCollection tables;
2086 BodyLowerer lowerer(moduleOp, tables, field->get(), options);
2087 auto funcs = walkCollect<function::FuncDefOp>(moduleOp, [](auto funcOp) {
2088 return !funcOp.nameIsConstrain();
2089 });
2090 for (function::FuncDefOp funcOp : funcs) {
2091 if (failed(lowerer.lowerFunction(funcOp))) {
2092 signalPassFailure();
2093 return;
2094 }
2095 }
2096 }
2097
2098private:
2099 WitgenOptions options;
2100};
2101
2103class CreateWitgenEntryPass : public PassWrapper<CreateWitgenEntryPass, OperationPass<ModuleOp>> {
2104public:
2105 MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(CreateWitgenEntryPass)
2106
2107
2108 explicit CreateWitgenEntryPass(OutputScope newOutputScope) : outputScope(newOutputScope) {}
2109
2111 StringRef getArgument() const final { return "llzk-create-witgen-entry"; }
2112
2114 StringRef getDescription() const final {
2115 return "Create the llzk-witgen execution-engine entry wrapper";
2116 }
2117
2119 StringRef getName() const override { return "CreateWitgenEntryPass"; }
2120
2122 void runOnOperation() override {
2123 ModuleOp moduleOp = getOperation();
2124 auto field = getModuleField(moduleOp);
2125 if (failed(field)) {
2126 signalPassFailure();
2127 return;
2128 }
2129
2130 SymbolTableCollection tables;
2131 auto mainDef = getMainInstanceDef(tables, moduleOp.getOperation());
2132 if (failed(mainDef) || !mainDef.value()) {
2133 moduleOp.emitError("module is missing a concrete llzk.main struct");
2134 signalPassFailure();
2135 return;
2136 }
2137 function::FuncDefOp computeFunc = mainDef->get().getComputeFuncOp();
2138 if (!computeFunc) {
2139 moduleOp.emitError("main struct is missing @compute");
2140 signalPassFailure();
2141 return;
2142 }
2143
2144 auto outputs =
2145 collectOutputBindings(mainDef->get(), tables, computeFunc.getOperation(), outputScope);
2146 if (failed(outputs)) {
2147 signalPassFailure();
2148 return;
2149 }
2150
2151 OpBuilder builder(moduleOp.getContext());
2152 builder.setInsertionPointToEnd(moduleOp.getBody());
2153
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
2159 ))) {
2160 signalPassFailure();
2161 return;
2162 }
2163 if (loweredLeafTypes.size() != 1 || !isa<MemRefType>(loweredLeafTypes.front())) {
2164 computeFunc.emitError(
2165 "execution-engine wrapper only supports felt and array<...xfelt> inputs"
2166 );
2167 signalPassFailure();
2168 return;
2169 }
2170 wrapperArgs.push_back(loweredLeafTypes.front());
2171 }
2172 for (const OutputBinding &output : *outputs) {
2173 SmallVector<Type> loweredLeafTypes;
2174 if (failed(flattenTypeLeaves(
2175 output.type, tables, computeFunc.getOperation(), field->get(), loweredLeafTypes, {},
2176 true
2177 ))) {
2178 signalPassFailure();
2179 return;
2180 }
2181 if (loweredLeafTypes.size() != 1 || !isa<MemRefType>(loweredLeafTypes.front())) {
2182 computeFunc.emitError(
2183 "execution-engine wrapper only supports felt and array<...xfelt> outputs"
2184 );
2185 signalPassFailure();
2186 return;
2187 }
2188 wrapperArgs.push_back(loweredLeafTypes.front());
2189 }
2190
2191 auto wrapper = builder.create<func::FuncOp>(
2192 computeFunc.getLoc(), "__llzk_witgen_main",
2193 builder.getFunctionType(wrapperArgs, TypeRange {})
2194 );
2195 wrapper->setAttr(LLVM::LLVMDialect::getEmitCWrapperAttrName(), builder.getUnitAttr());
2196 Block *entry = wrapper.addEntryBlock();
2197 builder.setInsertionPointToStart(entry);
2198
2199 SmallVector<Type> loweredMainResultTypes;
2200 for (Type resultType : computeFunc.getResultTypes()) {
2201 if (failed(flattenABILeafTypes(
2202 resultType, tables, computeFunc.getOperation(), field->get(), loweredMainResultTypes
2203 ))) {
2204 signalPassFailure();
2205 return;
2206 }
2207 }
2208
2209 SmallVector<Value> mainArgs;
2210 for (auto [argType, wrapperArg] : llvm::zip(
2211 computeFunc.getArgumentTypes(),
2212 entry->getArguments().take_front(computeFunc.getNumArguments())
2213 )) {
2214 if (isScalarType(argType)) {
2215 mainArgs.push_back(loadStorageScalar(builder, computeFunc.getLoc(), wrapperArg));
2216 } else {
2217 auto abiLeafTypes =
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();
2223 return;
2224 }
2225 if (wrapperArg.getType() == abiLeafTypes->front()) {
2226 mainArgs.push_back(wrapperArg);
2227 } else {
2228 mainArgs.push_back(builder.create<memref::CastOp>(
2229 computeFunc.getLoc(), abiLeafTypes->front(), wrapperArg
2230 ));
2231 }
2232 }
2233 }
2234 auto loweredMain = builder.create<func::CallOp>(
2235 computeFunc.getLoc(), mangleFunctionName(computeFunc), loweredMainResultTypes, mainArgs
2236 );
2237
2238 LoweredValue mainResultValue {
2239 computeFunc.getResultTypes().front(),
2240 llvm::SmallVector<Value>(loweredMain.getResults().begin(), loweredMain.getResults().end())
2241 };
2242
2243 auto extractOutputSlice = [&](ArrayRef<std::string> path, Type currentType,
2244 ArrayRef<Value> leaves,
2245 auto &self) -> FailureOr<SmallVector<Value>> {
2246 if (path.empty()) {
2247 return SmallVector<Value>(leaves.begin(), leaves.end());
2248 }
2249 if (auto structType = dyn_cast<component::StructType>(currentType)) {
2250 auto defLookup = structType.getDefinition(tables, computeFunc.getOperation());
2251 if (failed(defLookup)) {
2252 return failure();
2253 }
2254 unsigned localCursor = 0;
2255 for (component::MemberDefOp member : defLookup->get().getMemberDefs()) {
2256 auto leafCount =
2257 getLeafCount(member.getType(), tables, member.getOperation(), field->get());
2258 if (failed(leafCount)) {
2259 return failure();
2260 }
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);
2265 }
2266 }
2267 computeFunc.emitError("failed to find struct member while wiring witgen outputs");
2268 return failure();
2269 }
2270 if (auto podType = dyn_cast<pod::PodType>(currentType)) {
2271 unsigned localCursor = 0;
2272 for (pod::RecordAttr record : podType.getRecords()) {
2273 auto leafCount =
2274 getLeafCount(record.getType(), tables, computeFunc.getOperation(), field->get());
2275 if (failed(leafCount)) {
2276 return failure();
2277 }
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);
2282 }
2283 }
2284 computeFunc.emitError("failed to find POD record while wiring witgen outputs");
2285 return failure();
2286 }
2287 computeFunc.emitError("extra witness path components for non-aggregate output");
2288 return failure();
2289 };
2290
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
2295 );
2296 if (failed(slice) || slice->empty()) {
2297 wrapper.emitError("missing selected witness output slice while building witgen entry");
2298 signalPassFailure();
2299 return;
2300 }
2301 if (isScalarType(output.type)) {
2302 storeStorageScalar(
2303 builder, computeFunc.getLoc(),
2304 loadStorageScalar(builder, computeFunc.getLoc(), slice->front()), outputMemRef
2305 );
2306 } else {
2307 builder.create<memref::CopyOp>(computeFunc.getLoc(), slice->front(), outputMemRef);
2308 }
2309 }
2310 builder.create<func::ReturnOp>(computeFunc.getLoc());
2311
2312 // Remove `llzk.main` attribute because the main struct is deleted below.
2313 moduleOp->removeAttr(MAIN_ATTR_NAME);
2314
2315 SmallVector<Operation *> toErase;
2316 for (Operation &op : moduleOp.getBody()->getOperations()) {
2317 if (!isa<func::FuncOp>(op)) {
2318 toErase.push_back(&op);
2319 }
2320 }
2321 for (Operation *op : toErase) {
2322 op->erase();
2323 }
2324 }
2325
2326private:
2327 OutputScope outputScope;
2328};
2329
2330} // namespace
2331
2332void addWitgenPreparePipeline(OpPassManager &pm, const WitgenOptions &) {
2333 using namespace llzk::polymorphic;
2334 pm.addPass(createFlatteningPass(
2336 ));
2337 pm.addPass(mlir::createLowerAffinePass());
2338 // TODO: simplify lowering with `llzk-inline-structs` and `llzk-pod-to-scalar` when both are
2339 // available and support PODs.
2340 pm.addPass(mlir::createCanonicalizerPass());
2341 pm.addPass(mlir::createCSEPass());
2342}
2343
2344std::unique_ptr<Pass> createLowerComputeToCorePass(const WitgenOptions &options) {
2345 return std::make_unique<LowerComputeToCorePass>(options);
2346}
2347
2348std::unique_ptr<Pass> createCreateWitgenEntryPass(OutputScope outputScope) {
2349 return std::make_unique<CreateWitgenEntryPass>(outputScope);
2350}
2351
2352} // namespace llzk::witgen
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
Definition LICENSE.txt:28
std::unique_ptr<::mlir::Pass > createFlatteningPass()
llvm::Expected< T > checkedCast(U u)
Definition WitgenUtils.h:28
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.
Definition ValueModel.h:55
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
Definition Compare.h:94
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.
Definition Field.h:160
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.
Definition Field.cpp:264
Configure one llzk-witgen execution.