Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion zirgen/Dialect/R1CS/IR/Ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
#include "zirgen/Dialect/R1CS/IR/R1CS.h"

#include "mlir/IR/Builders.h"
//#include "mlir/IR/PatternMatch.h"
// #include "mlir/IR/PatternMatch.h"

#define GET_OP_CLASSES
#include "zirgen/Dialect/R1CS/IR/Ops.cpp.inc"
Expand Down
14 changes: 7 additions & 7 deletions zirgen/Dialect/R1CS/IR/Types.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,13 +14,13 @@

#include "mlir/IR/Types.h"
#include "mlir/IR/Builders.h"
//#include "mlir/IR/BuiltinTypes.h"
//#include "mlir/IR/Diagnostics.h"
//#include "mlir/IR/DialectImplementation.h"
//#include "mlir/IR/PatternMatch.h"
//#include "mlir/IR/StorageUniquerSupport.h"
//#include "llvm/ADT/StringExtras.h"
//#include "llvm/ADT/TypeSwitch.h"
// #include "mlir/IR/BuiltinTypes.h"
// #include "mlir/IR/Diagnostics.h"
// #include "mlir/IR/DialectImplementation.h"
// #include "mlir/IR/PatternMatch.h"
// #include "mlir/IR/StorageUniquerSupport.h"
// #include "llvm/ADT/StringExtras.h"
// #include "llvm/ADT/TypeSwitch.h"

namespace zirgen::R1CS {

Expand Down
7 changes: 3 additions & 4 deletions zirgen/Dialect/Zll/IR/Codegen.h
Original file line number Diff line number Diff line change
Expand Up @@ -526,7 +526,7 @@ struct EmitPart {
// String literals
template <size_t N>
EmitPart(const char (&str)[N])
: emitFunc([str](CodegenEmitter& cg) { *cg.getOutputStream() << str; }){};
: emitFunc([str](CodegenEmitter& cg) { *cg.getOutputStream() << str; }) {};

// References to a generated value.
EmitPart(CodegenValue val) : emitFunc([val](CodegenEmitter& cg) { cg.emitValue(val); }) {}
Expand All @@ -538,7 +538,7 @@ struct EmitPart {
// StringRefs must be explicitly converted so we don't accidentally
// skip canonicalizing identifiers.
explicit EmitPart(llvm::StringRef str)
: emitFunc([str](CodegenEmitter& cg) { *cg.getOutputStream() << str; }){};
: emitFunc([str](CodegenEmitter& cg) { *cg.getOutputStream() << str; }) {};

void emit(CodegenEmitter& cg) { emitFunc(cg); }

Expand All @@ -558,8 +558,7 @@ inline CodegenEmitter& CodegenEmitter::operator<<(EmitPart emitPart) {

template <typename Container, typename UnaryFunctor, typename T>
void CodegenEmitter::interleaveComma(const Container& c, UnaryFunctor each_fn) {
llvm::interleave(
c, *getOutputStream(), [&](const T& elem) { each_fn(elem); }, ", ");
llvm::interleave(c, *getOutputStream(), [&](const T& elem) { each_fn(elem); }, ", ");
}
template <typename Container, typename T> void CodegenEmitter::interleaveComma(const Container& c) {
llvm::interleave(
Expand Down
4 changes: 2 additions & 2 deletions zirgen/circuit/bigint/elliptic_curve.h
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ class WeierstrassCurve {
// y^2 = x^3 + a*x + b (mod p)
public:
WeierstrassCurve(Value prime, Value a_coeff, Value b_coeff)
: _prime(prime), _a_coeff(a_coeff), _b_coeff(b_coeff){};
: _prime(prime), _a_coeff(a_coeff), _b_coeff(b_coeff) {};
const Value& a() const { return _a_coeff; };
const Value& b() const { return _b_coeff; };
const Value& prime() const { return _prime; };
Expand All @@ -44,7 +44,7 @@ class AffinePt {
// A point on a Weierstrass curve expressed in affine coordinates
public:
AffinePt(Value x_coord, Value y_coord, std::shared_ptr<WeierstrassCurve> curve)
: _x(x_coord), _y(y_coord), _curve(curve){};
: _x(x_coord), _y(y_coord), _curve(curve) {};
const Value& x() const { return _x; };
const Value& y() const { return _y; };
const std::shared_ptr<WeierstrassCurve>& curve() const { return _curve; };
Expand Down
12 changes: 6 additions & 6 deletions zirgen/circuit/fib/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -3,18 +3,18 @@ name = "risc0-circuit-fib"
version = "0.1.0"
edition = "2021"

[features]
cuda = []
default = []

[dependencies]
anyhow = "1.0"
log = "0.4"
risc0-zkp = { workspace = true, features = ["default"] }

[dev-dependencies]
env_logger = "0.11"

[build-dependencies]
cc = { version = "1.2", features = ["parallel"] }
glob = "0.3"

[features]
cuda = []
default = []
[dev-dependencies]
env_logger = "0.11"
12 changes: 9 additions & 3 deletions zirgen/circuit/fib/fib.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -31,8 +31,12 @@ int main(int argc, char* argv[]) {
[](Buffer control, Buffer out, Buffer data, Buffer mix, Buffer accum) {
// Normal execution
Register val = data[0];
IF(control[0]) { val = 1; }
IF(control[1]) { val = BACK(1, Val(val)) + BACK(2, Val(val)); }
IF(control[0]) {
val = 1;
}
IF(control[1]) {
val = BACK(1, Val(val)) + BACK(2, Val(val));
}
IF(control[2]) {
// TODO: Fix register equality via BufAccess
out[0] = CaptureVal(val);
Expand All @@ -41,7 +45,9 @@ int main(int argc, char* argv[]) {
barrier(1);
barrier(1);
barrier(1);
IF(control[0] + control[1] + control[2]) { accum[0] = 1; }
IF(control[0] + control[1] + control[2]) {
accum[0] = 1;
}
barrier(1);
});

Expand Down
8 changes: 6 additions & 2 deletions zirgen/circuit/recursion/micro.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,9 @@ void MicroOpImpl::set(MicroInst inst, Val writeAddr, Reg extraPrev, size_t extra
IF(decode->at(size_t(MicroOpcode::INV)) * operands[1]) {
FpExt a = in0->doRead(operands[0]);
in1->doNOP();
NONDET { out->doWrite(writeAddr, inv(a).getElems()); }
NONDET {
out->doWrite(writeAddr, inv(a).getElems());
}
XLOG("INV: %e -> %e", a.getElems(), out->data());
eq(FpExt(Val(1)), FpExt(in0->data()) * FpExt(out->data()));
}
Expand Down Expand Up @@ -87,7 +89,9 @@ void MicroOpImpl::set(MicroInst inst, Val writeAddr, Reg extraPrev, size_t extra
in0->doNOP();
in1->doNOP();
out->doWrite(writeAddr, {0, 0, 0, 0});
NONDET { auto vals = doExtern("readIOPHeader", "", 0, {operands[0], operands[1]}); }
NONDET {
auto vals = doExtern("readIOPHeader", "", 0, {operands[0], operands[1]});
}
}
IF(decode->at(size_t(MicroOpcode::READ_IOP_BODY))) {
in0->doNOP();
Expand Down
16 changes: 12 additions & 4 deletions zirgen/circuit/recursion/sha.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -131,11 +131,15 @@ std::array<Val, kWordSize> toBytes(std::array<Val, 32> in) {

static void setCarry(std::array<Bit, 32> out, ShortVec in, Twit carryLow, Twit carryHigh) {
Val carryLow8 = toBits(out, in[0], 0, 16);
NONDET { carryLow->set(carryLow8 & 3); }
NONDET {
carryLow->set(carryLow8 & 3);
}
Val carryLow1 = (carryLow8 - carryLow) / 4;
eqz(carryLow1 * (1 - carryLow1));
Val carryHigh8 = toBits(out, in[1] + carryLow8, 16, 16);
NONDET { carryHigh->set(carryHigh8 & 3); }
NONDET {
carryHigh->set(carryHigh8 & 3);
}
Val carryHigh1 = (carryHigh8 - carryHigh) / 4;
eqz(carryHigh1 * (1 - carryHigh1));
}
Expand Down Expand Up @@ -210,8 +214,12 @@ void ShaCycleImpl::setLoad(MacroInst inst, Val writeAddr) {
// top4 (unless it's == 4, in which case we set it to 0)
NONDET {
Val top4is4 = isz(top4 - 4);
IF(top4is4) { wCarryLow->set(0); }
IF(1 - top4is4) { wCarryLow->set(top4); }
IF(top4is4) {
wCarryLow->set(0);
}
IF(1 - top4is4) {
wCarryLow->set(top4);
}
}
// XLOG("cur = %u, top4 = %u, bot27 = %u, wCarryLow = %u",
// kBabyBearToMontgomery * io0->data()[0],
Expand Down
50 changes: 25 additions & 25 deletions zirgen/circuit/recursion/test/AB.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -241,31 +241,31 @@ So we verify that the taggedStruct uses below produce the same result
TEST(RECURSION, taggedStruct) {
using namespace llvm;
std::string goal = "566573df13310440113eb4a81e9cb8ab7c1a96aa8ed7450885d06a0ac06b956a";
doAB(
HashType::POSEIDON2,
{{0x0699544a,
0x10740194,
0x5fcfb7ec,
0x24d402b4,
0x2c917c8c,
0x58576ff6,
0x6e6063c6,
0x3fa4a82d}},
[&](Buffer out, ReadIopVal iop) {
auto digest0 = iop.readDigests(1)[0];
auto digest1 = taggedStruct("digest1", {}, {1, 2013265920, 3});
auto digest2 = taggedStruct("digest2", {digest1, digest1}, {2013265920, 5});
auto digest3 =
taggedStruct("digest3", {digest1, digest2, digest1}, {6, 7, 2013265920, 9, 10});
auto digest4 = taggedStruct("digest4", {digest3, digest0, digest2}, {6, 2013265920, 9, 10});
std::vector<Val> bytes;
for (size_t i = 0; i < 32; i++) {
bytes.push_back(hexDigitValue(goal[2 * i]) * 16 + hexDigitValue(goal[2 * i + 1]));
}
auto goal = intoDigest(bytes, Zll::DigestKind::Sha256);
assert_eq(digest4, goal);
out.setDigest(0, digest4, "digest");
});
doAB(HashType::POSEIDON2,
{{0x0699544a,
0x10740194,
0x5fcfb7ec,
0x24d402b4,
0x2c917c8c,
0x58576ff6,
0x6e6063c6,
0x3fa4a82d}},
[&](Buffer out, ReadIopVal iop) {
auto digest0 = iop.readDigests(1)[0];
auto digest1 = taggedStruct("digest1", {}, {1, 2013265920, 3});
auto digest2 = taggedStruct("digest2", {digest1, digest1}, {2013265920, 5});
auto digest3 =
taggedStruct("digest3", {digest1, digest2, digest1}, {6, 7, 2013265920, 9, 10});
auto digest4 =
taggedStruct("digest4", {digest3, digest0, digest2}, {6, 2013265920, 9, 10});
std::vector<Val> bytes;
for (size_t i = 0; i < 32; i++) {
bytes.push_back(hexDigitValue(goal[2 * i]) * 16 + hexDigitValue(goal[2 * i + 1]));
}
auto goal = intoDigest(bytes, Zll::DigestKind::Sha256);
assert_eq(digest4, goal);
out.setDigest(0, digest4, "digest");
});
}

} // namespace zirgen::recursion
8 changes: 6 additions & 2 deletions zirgen/circuit/recursion/wom.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -88,15 +88,19 @@ WomHeaderImpl::WomHeaderImpl()
WomRegImpl::WomRegImpl() : elem(CompContext::allocateFromPool<impl::WomAlloc>("wom")->elem) {}

std::array<Val, kExtSize> WomRegImpl::doRead(Val addr) {
NONDET { elem->setData(doExtern("womRead", "", kExtSize, {addr})); }
NONDET {
elem->setData(doExtern("womRead", "", kExtSize, {addr}));
}
elem->addr->set(addr);
return elem->dataVals();
}

void WomRegImpl::doWrite(Val addr, std::array<Val, kExtSize> data) {
elem->addr->set(addr);
elem->setData(data);
NONDET { doExtern("womWrite", "", 0, elem->toVals()); }
NONDET {
doExtern("womWrite", "", 0, elem->toVals());
}
}

void WomRegImpl::doNOP() {
Expand Down
48 changes: 36 additions & 12 deletions zirgen/circuit/rv32im/v1/edsl/bigint.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,9 @@ void BigIntCycleImpl::set(Top top) {
}
}
}
IF(stageOffset) { ioAddr->set(BACK(1, ioAddr->get())); }
IF(stageOffset) {
ioAddr->set(BACK(1, ioAddr->get()));
}

// In stages 1-3, read an 4 words of input from memory.
// Related to step 4 in the approach description.
Expand Down Expand Up @@ -294,8 +296,12 @@ void BigIntCycleImpl::set(Top top) {
bytes.at(i)->set(0);
}
IF(stage->at(1)) {
IF(1 - stageOffset) { bytes.at(i)->set(q.at(i)); }
IF(stageOffset) { bytes.at(i)->set(q.at(i + BigInt::kBytesSize)); }
IF(1 - stageOffset) {
bytes.at(i)->set(q.at(i));
}
IF(stageOffset) {
bytes.at(i)->set(q.at(i + BigInt::kBytesSize));
}
}
// Stages 2 and 3 constrain the low carries to range [-2^15, 2^15).
// Does so by splitting the value plus 2^15 into two bytes at adjacent indices.
Expand Down Expand Up @@ -337,8 +343,12 @@ void BigIntCycleImpl::set(Top top) {
}
}
IF(stage->at(4)) {
IF(1 - stageOffset) { bytes.at(i)->set(z.at(i)); }
IF(stageOffset) { bytes.at(i)->set(z.at(i + BigInt::kBytesSize)); }
IF(1 - stageOffset) {
bytes.at(i)->set(z.at(i));
}
IF(stageOffset) {
bytes.at(i)->set(z.at(i + BigInt::kBytesSize));
}
}
}
// logBigInt("bytes", bytes);
Expand All @@ -351,8 +361,12 @@ void BigIntCycleImpl::set(Top top) {
c.emplace_back(0);

IF(stage->at(2)) {
IF(1 - stageOffset) { carryHi.at(i)->set(c.at(i + BigInt::kByteWidth)); }
IF(stageOffset) { carryHi.at(i)->set(c.at(i + BigInt::kByteWidth + BigInt::kCarryHiSize)); }
IF(1 - stageOffset) {
carryHi.at(i)->set(c.at(i + BigInt::kByteWidth));
}
IF(stageOffset) {
carryHi.at(i)->set(c.at(i + BigInt::kByteWidth + BigInt::kCarryHiSize));
}
}
IF(stage->at(3)) {
IF(1 - stageOffset) {
Expand All @@ -368,20 +382,30 @@ void BigIntCycleImpl::set(Top top) {
for (size_t i = 0; i < BigInt::kMulBufferSize; i++) {
// At stage 2 copy the denomalized reduction value r into the mulBuffer.
IF(stage->at(2)) {
IF(1 - stageOffset) { mulBuffer.at(i)->set(r.at(i)); }
IF(stageOffset) { mulBuffer.at(i)->set(r.at(i + BigInt::kMulBufferSize)); }
IF(1 - stageOffset) {
mulBuffer.at(i)->set(r.at(i));
}
IF(stageOffset) {
mulBuffer.at(i)->set(r.at(i + BigInt::kMulBufferSize));
}
}
IF(stage->at(4)) {
IF(1 - stageOffset) { mulBuffer.at(i)->set(denormZ.at(i)); }
IF(stageOffset) { mulBuffer.at(i)->set(denormZ.at(i + BigInt::kMulBufferSize)); }
IF(1 - stageOffset) {
mulBuffer.at(i)->set(denormZ.at(i));
}
IF(stageOffset) {
mulBuffer.at(i)->set(denormZ.at(i + BigInt::kMulBufferSize));
}
}
}
}

// At stages 1 and 3, copy the inputs into the multiplier.
for (size_t i = 0; i < BigInt::kMulInSize; i++) {
// At stage 1, copy q from bytes to the first half of the mulBuffer.
IF(stage->at(1)) { mulBuffer.at(i)->set(bytes.at(i)); }
IF(stage->at(1)) {
mulBuffer.at(i)->set(bytes.at(i));
}
// At stage 3, copy x from io to the first half of the mulBuffer.
IF(stage->at(3)) {
mulBuffer.at(i)->set(BACK(2, io.at(i / kWordSize)->data().bytes.at(i % kWordSize)));
Expand Down
4 changes: 3 additions & 1 deletion zirgen/circuit/rv32im/v1/edsl/bigint2.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,9 @@ void BigInt2CycleImpl::set(Top top) {
Val isFirstCycle = BACK(1, body->majorSelect->at(MajorType::kECall));

IF(isFirstCycle) {
NONDET { doExtern("syscallBigInt2Precompute", "", 0, {}); }
NONDET {
doExtern("syscallBigInt2Precompute", "", 0, {});
}
// If first cycle, do special initalization
ECallCycle ecall = body->majorMux->at<MajorType::kECall>();
ECallBigInt2 ecallBigInt2 = ecall->minorMux->at<ECallType::kBigInt2>();
Expand Down
Loading
Loading