Skip to content

Commit 6cdd36b

Browse files
committed
draft
1 parent e2f40b5 commit 6cdd36b

16 files changed

Lines changed: 189 additions & 39 deletions

File tree

expander_compiler/src/builder/hint_normalize.rs

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -260,9 +260,10 @@ impl<'a, C: Config> InsnTransformAndExecute<'a, C, IrcIn<C>, IrcOut<C>> for Buil
260260
}
261261
}
262262
}
263-
CustomGate { gate_type, inputs } => ir::hint_normalized::Instruction::CustomGate {
263+
CustomGate { gate_type, inputs , num_outputs} => ir::hint_normalized::Instruction::CustomGate {
264264
gate_type: *gate_type,
265265
inputs: inputs.clone(),
266+
num_outputs: *num_outputs,
266267
},
267268
ToBinary { x, num_bits } => {
268269
let bits = self.push_insn_multi_out(InsnOut::Hint {

expander_compiler/src/circuit/ir/hint_normalized/mod.rs

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@ pub enum Instruction<C: Config> {
4242
CustomGate {
4343
gate_type: usize,
4444
inputs: Vec<usize>,
45+
num_outputs: usize,
4546
},
4647
}
4748

@@ -210,6 +211,7 @@ impl<C: Config> Instruction<C> {
210211
public_inputs: &[CircuitField<C>],
211212
hint_caller: &impl HintCaller<CircuitField<C>>,
212213
) -> EvalResult<C> {
214+
println!("instruction eval safe");
213215
if let Instruction::ConstantLike(coef) = self {
214216
return match coef {
215217
Coef::Constant(c) => EvalResult::Value(*c),
@@ -232,6 +234,7 @@ impl<C: Config> Instruction<C> {
232234
};
233235
}
234236
if let Instruction::CustomGate { .. } = self {
237+
println!("There's a custom gate");
235238
return EvalResult::Error(Error::UserError(
236239
"CustomGate currently unsupported".to_string(),
237240
));
@@ -499,9 +502,11 @@ impl<C: Config> RootCircuit<C> {
499502
public_inputs: &[CircuitField<C>],
500503
hint_caller: &impl HintCaller<CircuitField<C>>,
501504
) -> Result<Vec<CircuitField<C>>, Error> {
505+
println!("root circuit eval sub safe");
502506
let mut values = vec![CircuitField::<C>::zero(); 1];
503507
values.extend(inputs);
504508
for insn in circuit.instructions.iter() {
509+
println!("there's insn");
505510
match insn.eval_safe(&values, public_inputs, hint_caller) {
506511
EvalResult::Value(v) => {
507512
values.push(v);
@@ -525,8 +530,10 @@ impl<C: Config> RootCircuit<C> {
525530
}
526531
let mut res = Vec::new();
527532
for &o in circuit.outputs.iter() {
533+
println!("circuit output {:?}", o);
528534
res.push(values[o]);
529535
}
536+
println!("circuit done");
530537
Ok(res)
531538
}
532539

expander_compiler/src/circuit/ir/hint_normalized/serde.rs

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -41,10 +41,11 @@ impl<C: Config> ExpSerde for Instruction<C> {
4141
inputs.serialize_into(&mut writer)?;
4242
num_outputs.serialize_into(&mut writer)?;
4343
}
44-
Instruction::CustomGate { gate_type, inputs } => {
44+
Instruction::CustomGate { gate_type, inputs, num_outputs} => {
4545
6u8.serialize_into(&mut writer)?;
4646
gate_type.serialize_into(&mut writer)?;
4747
inputs.serialize_into(&mut writer)?;
48+
num_outputs.serialize_into(&mut writer)?;
4849
}
4950
};
5051
Ok(())
@@ -71,6 +72,7 @@ impl<C: Config> ExpSerde for Instruction<C> {
7172
6 => Instruction::CustomGate {
7273
gate_type: usize::deserialize_from(&mut reader)?,
7374
inputs: Vec::<usize>::deserialize_from(&mut reader)?,
75+
num_outputs: usize::deserialize_from(&mut reader)?,
7476
},
7577
_ => {
7678
return Err(IoError::new(

expander_compiler/src/circuit/ir/source/mod.rs

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -59,6 +59,7 @@ pub enum Instruction<C: Config> {
5959
CustomGate {
6060
gate_type: usize,
6161
inputs: Vec<usize>,
62+
num_outputs: usize,
6263
},
6364
ToBinary {
6465
x: usize,
@@ -256,9 +257,10 @@ impl<C: Config> common::Instruction<C> for Instruction<C> {
256257
if_true: f(*if_true),
257258
if_false: f(*if_false),
258259
},
259-
Instruction::CustomGate { gate_type, inputs } => Instruction::CustomGate {
260+
Instruction::CustomGate { gate_type, inputs , num_outputs} => Instruction::CustomGate {
260261
gate_type: *gate_type,
261262
inputs: inputs.iter().map(|i| f(*i)).collect(),
263+
num_outputs: *num_outputs,
262264
},
263265
Instruction::ToBinary { x, num_bits } => Instruction::ToBinary {
264266
x: f(*x),
@@ -397,7 +399,8 @@ impl<C: Config> common::Instruction<C> for Instruction<C> {
397399
} else {
398400
values[*if_true]
399401
}),
400-
Instruction::CustomGate { gate_type, inputs } => {
402+
Instruction::CustomGate { gate_type, inputs , num_outputs} => {
403+
// TODO: impl custom gate
401404
let outputs =
402405
hints::stub_impl(*gate_type, &inputs.iter().map(|i| values[*i]).collect(), 1);
403406
EvalResult::Values(outputs)

expander_compiler/src/circuit/ir/source/serde.rs

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -77,10 +77,11 @@ impl<C: Config> ExpSerde for Instruction<C> {
7777
if_true.serialize_into(&mut writer)?;
7878
if_false.serialize_into(&mut writer)?;
7979
}
80-
Instruction::CustomGate { gate_type, inputs } => {
80+
Instruction::CustomGate { gate_type, inputs , num_outputs} => {
8181
12u8.serialize_into(&mut writer)?;
8282
gate_type.serialize_into(&mut writer)?;
8383
inputs.serialize_into(&mut writer)?;
84+
num_outputs.serialize_into(&mut writer)?;
8485
}
8586
Instruction::ToBinary { x, num_bits } => {
8687
13u8.serialize_into(&mut writer)?;
@@ -169,6 +170,7 @@ impl<C: Config> ExpSerde for Instruction<C> {
169170
12 => Instruction::CustomGate {
170171
gate_type: usize::deserialize_from(&mut reader)?,
171172
inputs: Vec::<usize>::deserialize_from(&mut reader)?,
173+
num_outputs: usize::deserialize_from(&mut reader)?,
172174
},
173175
13 => Instruction::ToBinary {
174176
x: usize::deserialize_from(&mut reader)?,

expander_compiler/src/circuit/ir/source/tests.rs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -67,6 +67,7 @@ impl<C: Config> RandomInstruction for Instruction<C> {
6767
inputs: (0..num_inputs)
6868
.map(|_| rnd.next_u64() as usize % num_vars + 1)
6969
.collect(),
70+
num_outputs,
7071
}
7172
}
7273
} else if prob1 < 0.74 {

expander_compiler/src/circuit/layered/mod.rs

Lines changed: 24 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@ use std::{fmt, hash::Hash};
22

33
use serdes::ExpSerde;
44

5-
use crate::{field::FieldArith, hints, utils::error::Error};
5+
use crate::{field::FieldArith, frontend::HintCaller, hints, utils::error::Error};
66

77
use super::config::{CircuitField, Config};
88

@@ -733,10 +733,12 @@ impl<C: Config, I: InputType> Circuit<C, I> {
733733
&self,
734734
inputs: Vec<CircuitField<C>>,
735735
public_inputs: &[CircuitField<C>],
736+
customgate_caller: &impl HintCaller<CircuitField<C>>,
736737
) -> (Vec<CircuitField<C>>, bool) {
737738
if inputs.len() != self.input_size() {
738739
panic!("input length mismatch");
739740
}
741+
println!("inputs {:?}", inputs);
740742
let mut cur = vec![inputs];
741743
for id in self.layer_ids.iter() {
742744
let mut next = vec![CircuitField::<C>::zero(); self.segments[*id].num_outputs];
@@ -749,6 +751,7 @@ impl<C: Config, I: InputType> Circuit<C, I> {
749751
&inputs,
750752
&mut next,
751753
public_inputs,
754+
customgate_caller,
752755
);
753756
cur.push(next);
754757
}
@@ -772,6 +775,7 @@ impl<C: Config, I: InputType> Circuit<C, I> {
772775
cur: &[&[CircuitField<C>]],
773776
nxt: &mut [CircuitField<C>],
774777
public_inputs: &[CircuitField<C>],
778+
customgate_caller: &impl HintCaller<CircuitField<C>>,
775779
) {
776780
for m in seg.gate_muls.iter() {
777781
nxt[m.output] += cur[m.inputs[0].layer()][m.inputs[0].offset()]
@@ -790,9 +794,14 @@ impl<C: Config, I: InputType> Circuit<C, I> {
790794
for input in cu.inputs.iter() {
791795
inputs.push(cur[input.layer()][input.offset()]);
792796
}
793-
let outputs = hints::stub_impl(cu.gate_type, &inputs, 1);
794-
for (i, output) in outputs.iter().enumerate() {
795-
nxt[cu.output + i] += *output * cu.coef.get_value_with_public_inputs(public_inputs);
797+
// let outputs = hints::stub_impl(cu.gate_type, &inputs, 1);
798+
// TODO: custom outputs num setable
799+
// let outputs = customgate_caller.call_custom_gate(cu.gate_type, &inputs, cu.output);
800+
let outputs = customgate_caller.call_custom_gate(cu.gate_type, &inputs, 1);
801+
if let Ok(outputs) = outputs {
802+
for (i, output) in outputs.iter().enumerate() {
803+
nxt[cu.output + i] += *output * cu.coef.get_value_with_public_inputs(public_inputs);
804+
}
796805
}
797806
}
798807
for (sub_id, allocs) in seg.child_segs.iter() {
@@ -810,6 +819,7 @@ impl<C: Config, I: InputType> Circuit<C, I> {
810819
&inputs,
811820
&mut nxt[a.output_offset..a.output_offset + subc.num_outputs],
812821
public_inputs,
822+
customgate_caller,
813823
);
814824
}
815825
}
@@ -819,6 +829,7 @@ impl<C: Config, I: InputType> Circuit<C, I> {
819829
&self,
820830
inputs: Vec<SF>,
821831
public_inputs: &[SF],
832+
customgate_caller: &impl HintCaller<CircuitField<C>>,
822833
) -> (Vec<SF>, Vec<bool>) {
823834
if inputs.len() != self.input_size() {
824835
panic!("input length mismatch");
@@ -835,6 +846,7 @@ impl<C: Config, I: InputType> Circuit<C, I> {
835846
&inputs,
836847
&mut next,
837848
public_inputs,
849+
customgate_caller,
838850
);
839851
cur.push(next);
840852
}
@@ -860,6 +872,7 @@ impl<C: Config, I: InputType> Circuit<C, I> {
860872
cur: &[&[SF]],
861873
nxt: &mut [SF],
862874
public_inputs: &[SF],
875+
customgate_caller: &impl HintCaller<CircuitField<C>>,
863876
) {
864877
for m in seg.gate_muls.iter() {
865878
nxt[m.output] += cur[m.inputs[0].layer()][m.inputs[0].offset()]
@@ -884,10 +897,14 @@ impl<C: Config, I: InputType> Circuit<C, I> {
884897
let mut outputs = Vec::with_capacity(SF::PACK_SIZE);
885898
for x in inputs.iter() {
886899
// TODO: better handle custom gates
900+
/*
887901
if cu.gate_type == 12348 {
888-
outputs.push(vec![tmp_mul(&x)]);
902+
outputs.push(vec![inner_product(&x)]);
889903
} else {
890904
outputs.push(hints::stub_impl(cu.gate_type, x, 1));
905+
} */
906+
if let Ok(out) = customgate_caller.call_custom_gate(cu.gate_type, x, cu.output) {
907+
outputs.push(out);
891908
}
892909
}
893910
for i in 0..outputs[0].len() {
@@ -915,6 +932,7 @@ impl<C: Config, I: InputType> Circuit<C, I> {
915932
&inputs,
916933
&mut nxt[a.output_offset..a.output_offset + subc.num_outputs],
917934
public_inputs,
935+
customgate_caller,
918936
);
919937
}
920938
}
@@ -1004,7 +1022,7 @@ impl<C: Config, I: InputType> fmt::Display for Circuit<C, I> {
10041022
}
10051023
}
10061024

1007-
fn tmp_mul<F: crate::field::Field>(a: &[F]) -> F {
1025+
fn inner_product<F: crate::field::Field>(a: &[F]) -> F {
10081026
let mut sum = F::ZERO;
10091027
let n = a.len() / 2;
10101028
for i in 0..n {

expander_compiler/src/circuit/layered/witness.rs

Lines changed: 31 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ use serdes::{ExpSerde, SerdeResult};
66

77
use super::{Circuit, InputType};
88
use crate::circuit::config::{CircuitField, Config, SIMDField};
9+
use crate::frontend::{EmptyHintCaller, HintCaller};
910

1011
#[derive(Clone, Debug)]
1112
pub enum WitnessValues<C: Config> {
@@ -199,16 +200,18 @@ impl<C: Config, I: InputType> Circuit<C, I> {
199200
&self,
200201
witness: &Witness<C>,
201202
need_output: bool,
203+
customgate_caller: &impl HintCaller<CircuitField<C>>,
202204
) -> (Vec<bool>, Vec<Vec<CircuitField<C>>>) {
203205
if witness.num_witnesses == 0 {
204206
panic!("expected at least 1 witness")
205207
}
206208
let mut outputs = Vec::new();
207209
let mut constraints = Vec::new();
210+
println!("use simd ? {}", use_simd::<C>(witness.num_witnesses));
208211
if use_simd::<C>(witness.num_witnesses) {
209212
for (inputs, public_inputs) in witness.iter_simd() {
210213
let (out, constraint_result) =
211-
self.eval_with_public_inputs_simd(inputs, &public_inputs);
214+
self.eval_with_public_inputs_simd(inputs, &public_inputs, customgate_caller);
212215
if need_output {
213216
let n = outputs.len();
214217
for _ in 0..SIMDField::<C>::PACK_SIZE {
@@ -226,21 +229,43 @@ impl<C: Config, I: InputType> Circuit<C, I> {
226229
constraints.truncate(witness.num_witnesses);
227230
} else {
228231
for (inputs, public_inputs) in witness.iter_scalar() {
229-
let (out, constraint_result) = self.eval_with_public_inputs(inputs, &public_inputs);
232+
let (out, constraint_result) = self.eval_with_public_inputs(inputs, &public_inputs, customgate_caller);
230233
outputs.push(out);
231234
constraints.push(constraint_result);
232235
}
233236
}
234237
(constraints, outputs)
235238
}
236239

237-
pub fn run(&self, witness: &Witness<C>) -> Vec<bool> {
238-
let (constraints, _) = self.run_inner(witness, false);
240+
pub fn run(
241+
&self,
242+
witness: &Witness<C>,
243+
) -> Vec<bool> {
244+
let (constraints, _) = self.run_inner(witness, false, &EmptyHintCaller);
239245
constraints
240246
}
241247

242-
pub fn run_with_output(&self, witness: &Witness<C>) -> (Vec<bool>, Vec<Vec<CircuitField<C>>>) {
243-
self.run_inner(witness, true)
248+
pub fn run_with_output(
249+
&self,
250+
witness: &Witness<C>,
251+
customgate_caller: &impl HintCaller<CircuitField<C>>,
252+
) -> (Vec<bool>, Vec<Vec<CircuitField<C>>>) {
253+
self.run_inner(witness, true, customgate_caller)
254+
}
255+
256+
pub fn run_with_options(
257+
&self,
258+
witness: &Witness<C>,
259+
need_output: bool,
260+
customgate_caller: &impl HintCaller<CircuitField<C>>,
261+
) -> (Vec<bool>, Option<Vec<Vec<CircuitField<C>>>>) {
262+
let rst = self.run_inner(witness, need_output, customgate_caller);
263+
if need_output {
264+
(rst.0, Some(rst.1))
265+
}
266+
else {
267+
(rst.0, None)
268+
}
244269
}
245270
}
246271

expander_compiler/src/frontend/api.rs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -84,7 +84,7 @@ pub trait BasicAPI<C: Config> {
8484
inputs: &[Variable],
8585
num_outputs: usize,
8686
) -> Vec<Variable>;
87-
fn custom_gate(&mut self, gate_type: usize, inputs: &[Variable]) -> Variable;
87+
fn custom_gate(&mut self, gate_type: usize, inputs: &[Variable], num_outputs: usize) -> Variable;
8888
fn constant(&mut self, x: impl ToVariableOrValue<CircuitField<C>>) -> Variable;
8989
// try to get the value of a compile-time constant variable
9090
// this function has different behavior in normal and debug mode, in debug mode it always returns Some(value)

expander_compiler/src/frontend/builder.rs

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -445,11 +445,12 @@ impl<C: Config> BasicAPI<C> for Builder<C> {
445445
(0..num_outputs).map(|_| self.new_var()).collect()
446446
}
447447

448-
fn custom_gate(&mut self, gate_type: usize, inputs: &[Variable]) -> Variable {
448+
fn custom_gate(&mut self, gate_type: usize, inputs: &[Variable], num_outputs: usize) -> Variable {
449449
ensure_variables_valid(inputs);
450450
self.instructions.push(SourceInstruction::CustomGate {
451451
gate_type,
452452
inputs: inputs.iter().map(|v| v.id).collect(),
453+
num_outputs,
453454
});
454455
self.new_var()
455456
}
@@ -689,8 +690,8 @@ impl<C: Config> BasicAPI<C> for RootBuilder<C> {
689690
self.last_builder().new_hint(hint_key, inputs, num_outputs)
690691
}
691692

692-
fn custom_gate(&mut self, gate_type: usize, inputs: &[Variable]) -> Variable {
693-
self.last_builder().custom_gate(gate_type, inputs)
693+
fn custom_gate(&mut self, gate_type: usize, inputs: &[Variable], num_outputs: usize) -> Variable {
694+
self.last_builder().custom_gate(gate_type, inputs, num_outputs)
694695
}
695696

696697
fn constant(&mut self, x: impl ToVariableOrValue<CircuitField<C>>) -> Variable {

0 commit comments

Comments
 (0)