@@ -2,7 +2,7 @@ use std::{fmt, hash::Hash};
22
33use serdes:: ExpSerde ;
44
5- use crate :: { field:: FieldArith , hints, utils:: error:: Error } ;
5+ use crate :: { field:: FieldArith , frontend :: HintCaller , hints, utils:: error:: Error } ;
66
77use 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 {
0 commit comments