diff --git a/src/lib.rs b/src/lib.rs index 4ad95f1..546fef3 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -5,3 +5,6 @@ pub mod gadgets; pub mod layers; pub mod model; pub mod utils; + +#[cfg(test)] +pub mod tests; diff --git a/src/tests/mod.rs b/src/tests/mod.rs new file mode 100644 index 0000000..ea5c4e0 --- /dev/null +++ b/src/tests/mod.rs @@ -0,0 +1,2 @@ +pub mod sqrt; +pub mod test_circuit; diff --git a/src/tests/sqrt.rs b/src/tests/sqrt.rs new file mode 100644 index 0000000..3445449 --- /dev/null +++ b/src/tests/sqrt.rs @@ -0,0 +1,61 @@ +use std::{ + collections::{HashMap, HashSet}, + marker::PhantomData, +}; + +use halo2_proofs::{ + circuit::{AssignedCell, Layouter, Value}, + dev::MockProver, + halo2curves::{bn256::Fr, FieldExt}, +}; +use ndarray::{Array, ArrayBase, Dim, IxDyn, IxDynImpl, OwnedRepr}; + +use crate::gadgets::gadget::Gadget; +use crate::{ + gadgets::{ + gadget::{self, GadgetConfig, GadgetType}, + sqrt_big::SqrtBigChip, + }, + tests::test_circuit::{TestCircuit, K}, + utils::loader::TensorMsgpack, +}; + +use super::test_circuit::{TestCircFunc, TestConfig}; + +pub struct TestFunc {} + +impl TestCircFunc for TestFunc { + fn compute_tensor( + layouter: &mut impl Layouter, + config: &TestConfig, + tensors: &Vec>, Dim>>, + constants: &HashMap>, + ) { + let inp = &tensors[0]; + let inp_vec = inp.iter().collect::>(); + + let zero = constants.get(&0).unwrap().clone(); + + let sqrt_chip = SqrtBigChip::::construct(config.gadget_config.clone().into()); + let _result = sqrt_chip.forward(layouter.namespace(|| "test"), &vec![inp_vec], &vec![zero]); + } +} + +#[test] +fn test_sqrt() { + let input_tensor = TensorMsgpack { + idx: 0, + shape: vec![3], + data: vec![5, 10, 100], + }; + let input_tensors: Vec = vec![input_tensor]; + + let gadget = &crate::tests::test_circuit::GADGET_CONFIG; + let cloned_gadget = gadget.lock().unwrap().clone(); + + let circuit = TestCircuit::::new(input_tensors); + + let outp = vec![]; + let prover = MockProver::run(K.try_into().unwrap(), &circuit, vec![outp.clone()]).unwrap(); + assert_eq!(prover.verify(), Ok(())); +} diff --git a/src/tests/test_circuit.rs b/src/tests/test_circuit.rs new file mode 100644 index 0000000..01d04f8 --- /dev/null +++ b/src/tests/test_circuit.rs @@ -0,0 +1,321 @@ +use halo2_proofs::{dev::MockProver, halo2curves::bn256::Fr}; +use std::{ + collections::{HashMap, HashSet}, + marker::PhantomData, + rc::Rc, + sync::Mutex, +}; + +use halo2_proofs::{ + circuit::{AssignedCell, Layouter, SimpleFloorPlanner, Value}, + halo2curves::FieldExt, + plonk::{Advice, Circuit, Column, ConstraintSystem, Error}, +}; +use lazy_static::lazy_static; +use ndarray::{Array, IxDyn}; + +use crate::{ + gadgets::{ + add_pairs::AddPairsChip, + adder::AdderChip, + bias_div_round_relu6::BiasDivRoundRelu6Chip, + dot_prod::DotProductChip, + gadget::{Gadget, GadgetConfig, GadgetType}, + mul_pairs::MulPairsChip, + nonlinear::exp::ExpChip, + nonlinear::{logistic::LogisticGadgetChip, rsqrt::RsqrtGadgetChip}, + sqrt_big::SqrtBigChip, + squared_diff::SquaredDiffGadgetChip, + sub_pairs::SubPairsChip, + var_div::VarDivRoundChip, + }, + layers::{ + dag::{DAGLayerChip, DAGLayerConfig}, + layer::{Layer, LayerConfig, LayerType}, + }, + utils::loader::{load_model_msgpack, ModelMsgpack}, +}; + +// +// Contains a way to unit-test each of the files +// + +// Note this useful idiom: importing names from outer (for mod tests) scope. +use super::*; +use crate::{gadgets::sqrt_big, utils::loader::TensorMsgpack}; +use ndarray::{ArrayBase, Dim, IxDynImpl, OwnedRepr}; + +pub const K: u64 = 19; +pub const GLOBAL_SF: i64 = 2 ^ 16; +pub const NUM_COLS: u64 = 6; + +lazy_static! { + pub static ref GADGET_CONFIG: Mutex = Mutex::new(GadgetConfig::default()); +} + +pub trait TestCircFunc { + fn compute_tensor( + layouter: &mut impl Layouter, + config: &TestConfig, + tensors: &Vec>, Dim>>, + hash_map: &HashMap>, + ); +} + +#[derive(Clone, Debug)] +pub struct TestCircuit> { + pub used_gadgets: HashSet, + pub tensors: HashMap, IxDyn>>, + pub _marker: PhantomData, + pub _func_marker: PhantomData, +} + +impl> TestCircuit { + pub fn new(input_tensors: Vec) -> Self { + let gadget = &GADGET_CONFIG; + let cloned_gadget = gadget.lock().unwrap().clone(); + *gadget.lock().unwrap() = GadgetConfig { + scale_factor: GLOBAL_SF as u64, + shift_min_val: -(GLOBAL_SF * GLOBAL_SF * 1024), + div_outp_min_val: -(1 << (K - 1)), + min_val: -(1 << (K - 1)), + max_val: (1 << (K - 1)) - 10, + num_rows: (1 << K) - 10, + num_cols: NUM_COLS as usize, + ..cloned_gadget + }; + + let to_value = |x: i64| { + let bias = 1 << 31; + let x_pos = x + bias; + Value::known(F::from(x_pos as u64)) - Value::known(F::from(bias as u64)) + }; + + let mut tensors = HashMap::new(); + + for flat in input_tensors { + let value_flat = flat.data.iter().map(|x| to_value(*x)).collect::>(); + let shape = flat.shape.iter().map(|x| *x as usize).collect::>(); + let tensor = Array::from_shape_vec(IxDyn(&shape), value_flat).unwrap(); + tensors.insert(flat.idx, tensor); + } + + let mut used_gadgets = HashSet::new(); + used_gadgets.insert(GadgetType::AddPairs); + used_gadgets.insert(GadgetType::Adder); + used_gadgets.insert(GadgetType::BiasDivRoundRelu6); + used_gadgets.insert(GadgetType::DotProduct); + used_gadgets.insert(GadgetType::Rsqrt); + used_gadgets.insert(GadgetType::Exp); + used_gadgets.insert(GadgetType::Logistic); + + Self { + used_gadgets: used_gadgets, + tensors: tensors, + _marker: PhantomData, + _func_marker: PhantomData, + } + } + + pub fn assign_tensors( + &self, + mut layouter: impl Layouter, + columns: &Vec>, + tensors: &HashMap, IxDyn>>, + ) -> Result, IxDyn>>, Error> { + let tensors = layouter.assign_region( + || "asssignment", + |mut region| { + let mut cell_idx = 0; + let mut assigned_tensors = vec![]; + + for (tensor_idx, tensor) in tensors { + let tensor_idx = *tensor_idx as usize; + let mut flat = vec![]; + for val in tensor.iter() { + let row_idx = cell_idx / columns.len(); + let col_idx = cell_idx % columns.len(); + let cell = region.assign_advice(|| "assignment", columns[col_idx], row_idx, || *val)?; + flat.push(cell); + cell_idx += 1; + } + let tensor = Array::from_shape_vec(tensor.shape(), flat).unwrap(); + // TODO: is there a non-stupid way to do this? + while assigned_tensors.len() <= tensor_idx { + assigned_tensors.push(tensor.clone()); + } + assigned_tensors[tensor_idx] = tensor; + } + + Ok(assigned_tensors) + }, + )?; + + Ok(tensors) + } + + // FIXME: assign to public + pub fn assign_constants( + &self, + mut layouter: impl Layouter, + model_config: &TestConfig, + ) -> Result>, Error> { + let sf = model_config.gadget_config.scale_factor; + let min_val = model_config.gadget_config.min_val; + let max_val = model_config.gadget_config.max_val; + + let constants = layouter.assign_region( + || "constants", + |mut region| { + let mut constants: HashMap> = HashMap::new(); + let zero = region.assign_fixed( + || "zero", + model_config.gadget_config.fixed_columns[0], + 0, + || Value::known(F::zero()), + )?; + let one = region.assign_fixed( + || "one", + model_config.gadget_config.fixed_columns[0], + 1, + || Value::known(F::one()), + )?; + let sf_cell = region.assign_fixed( + || "sf", + model_config.gadget_config.fixed_columns[0], + 2, + || Value::known(F::from(sf)), + )?; + let min_val_cell = region.assign_fixed( + || "min_val", + model_config.gadget_config.fixed_columns[0], + 3, + || Value::known(F::zero() - F::from((-min_val) as u64)), + )?; + // TODO: the table goes from min_val to max_val - 1... fix this + let max_val_cell = region.assign_fixed( + || "max_val", + model_config.gadget_config.fixed_columns[0], + 4, + || Value::known(F::from((max_val - 1) as u64)), + )?; + + constants.insert(0, zero); + constants.insert(1, one); + constants.insert(sf as i64, sf_cell); + constants.insert(min_val, min_val_cell); + constants.insert(max_val, max_val_cell); + Ok(constants) + }, + )?; + Ok(constants) + } +} + +impl> Circuit for TestCircuit { + type Config = TestConfig; + type FloorPlanner = SimpleFloorPlanner; + + fn without_witnesses(&self) -> Self { + todo!() + } + + fn configure(meta: &mut ConstraintSystem) -> Self::Config { + // FIXME: decide which gadgets to make + let mut gadget_config = GADGET_CONFIG.lock().unwrap().clone(); + + let columns = (0..gadget_config.num_cols) + .map(|_| meta.advice_column()) + .collect::>(); + for col in columns.iter() { + meta.enable_equality(*col); + } + gadget_config.columns = columns; + + gadget_config.public_columns = vec![meta.instance_column()]; + meta.enable_equality(gadget_config.public_columns[0]); + + gadget_config.fixed_columns = vec![meta.fixed_column()]; + meta.enable_equality(gadget_config.fixed_columns[0]); + + // FIXME: fix this shit + gadget_config = AddPairsChip::::configure(meta, gadget_config); + gadget_config = AdderChip::::configure(meta, gadget_config); + gadget_config = BiasDivRoundRelu6Chip::::configure(meta, gadget_config); + gadget_config = DotProductChip::::configure(meta, gadget_config); + gadget_config = VarDivRoundChip::::configure(meta, gadget_config); + gadget_config = RsqrtGadgetChip::::configure(meta, gadget_config); + gadget_config = MulPairsChip::::configure(meta, gadget_config); + gadget_config = SubPairsChip::::configure(meta, gadget_config); + gadget_config = ExpChip::::configure(meta, gadget_config); + gadget_config = LogisticGadgetChip::::configure(meta, gadget_config); + gadget_config = SquaredDiffGadgetChip::::configure(meta, gadget_config); + gadget_config = SqrtBigChip::::configure(meta, gadget_config); + + TestConfig { + gadget_config: gadget_config.into(), + _marker: PhantomData, + } + } + + fn synthesize(&self, config: Self::Config, mut layouter: impl Layouter) -> Result<(), Error> { + // Assign tables + let gadget_rc: Rc = config.gadget_config.clone().into(); + + for gadget in self.used_gadgets.iter() { + match gadget { + GadgetType::AddPairs => { + let chip = AddPairsChip::::construct(gadget_rc.clone()); + chip.load_lookups(layouter.namespace(|| "add pairs lookup"))?; + } + GadgetType::Adder => { + let chip = AdderChip::::construct(gadget_rc.clone()); + chip.load_lookups(layouter.namespace(|| "adder lookup"))?; + } + GadgetType::BiasDivRoundRelu6 => { + let chip = BiasDivRoundRelu6Chip::::construct(gadget_rc.clone()); + chip.load_lookups(layouter.namespace(|| "bias div round relu6 lookup"))?; + } + GadgetType::DotProduct => { + let chip = DotProductChip::::construct(gadget_rc.clone()); + chip.load_lookups(layouter.namespace(|| "dot product lookup"))?; + } + GadgetType::VarDivRound => { + let chip = VarDivRoundChip::::construct(gadget_rc.clone()); + chip.load_lookups(layouter.namespace(|| "var div lookup"))?; + } + GadgetType::Rsqrt => { + let chip = RsqrtGadgetChip::::construct(gadget_rc.clone()); + chip.load_lookups(layouter.namespace(|| "rsqrt lookup"))?; + } + GadgetType::Exp => { + let chip = ExpChip::::construct(gadget_rc.clone()); + chip.load_lookups(layouter.namespace(|| "exp lookup"))?; + } + GadgetType::Logistic => { + let chip = LogisticGadgetChip::::construct(gadget_rc.clone()); + chip.load_lookups(layouter.namespace(|| "logistic lookup"))?; + } + _ => panic!("unsupported gadget"), + } + } + + // Assign weights and constants + let tensors = self.assign_tensors( + layouter.namespace(|| "assignment"), + &config.gadget_config.columns, + &self.tensors, + )?; + let constants = self.assign_constants(layouter.namespace(|| "constants"), &config)?; + + G::compute_tensor(&mut layouter, &config, &tensors, &constants); + + Ok(()) + } +} + +#[derive(Clone, Debug)] +pub struct TestConfig { + pub gadget_config: Rc, + pub _marker: PhantomData, +}