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
4 changes: 2 additions & 2 deletions zirgen/Dialect/R1CS/IR/Ops.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand All @@ -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
16 changes: 8 additions & 8 deletions zirgen/Dialect/R1CS/IR/Types.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand All @@ -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
9 changes: 4 additions & 5 deletions zirgen/Dialect/Zll/IR/Codegen.h
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand Down 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
6 changes: 3 additions & 3 deletions zirgen/circuit/bigint/elliptic_curve.h
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand Down 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"
14 changes: 10 additions & 4 deletions zirgen/circuit/fib/fib.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand Down 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
10 changes: 7 additions & 3 deletions zirgen/circuit/recursion/micro.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand Down 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
18 changes: 13 additions & 5 deletions zirgen/circuit/recursion/sha.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand Down 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
52 changes: 26 additions & 26 deletions zirgen/circuit/recursion/test/AB.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand Down 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
10 changes: 7 additions & 3 deletions zirgen/circuit/recursion/wom.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
// Copyright 2024 RISC Zero, Inc.
// Copyright 2026 RISC Zero, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
Expand Down 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
2 changes: 1 addition & 1 deletion zirgen/circuit/rv32im/shared/test/defs.bzl
Original file line number Diff line number Diff line change
Expand Up @@ -78,7 +78,7 @@ def compile_riscv_tests():
)
all_bins = all_bins + [test]
native.filegroup(
name = "riscv_test_bins",
name = "riscv_test_bins",
srcs = all_bins,
visibility = ["//visibility:public"],
)
Loading
Loading