#![cfg_attr(docsrs, feature(doc_cfg))]
pub mod proc_int_data;
use proc_int_data::*;
pub mod cycle_stage;
use cycle_stage::*;
mod circuit_config;
use circuit_config::*;
use gatenative::cpu_build_exec::*;
use gatenative::cpu_data_transform::*;
use gatenative::opencl_build_exec::*;
use gatenative::opencl_data_transform::*;
use gatenative::*;
use gatesim::*;
use infmachine_config::*;
pub use gatenative;
use gatenative::gatesim;
pub use infmachine_config;
use std::fmt::Debug;
use std::hash::Hash;
use std::marker::PhantomData;
use std::str::FromStr;
#[repr(u8)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DataAccess {
Nothing,
ReadOnly,
WriteOnly,
ReadWrite,
}
#[repr(u8)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DataPartMove {
Nothing,
Forward,
Backward,
}
#[repr(u8)]
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DataKind {
MemAddress,
TempBuffer,
ProcId,
}
pub(crate) const GLOBAL_STATE_STOP: u32 = 1;
pub(crate) const GLOBAL_STATE_ILLEGAL: u32 = 2;
#[derive(thiserror::Error, Debug)]
pub enum InfParMachineError<'a, DR, DW, D, E, CS>
where
DR: DataReader,
DW: DataWriter,
D: DataHolder<'a, DR, DW>,
E: Executor<'a, DR, DW, D>,
E::ErrorType: Debug,
CS: CycleStage<'a, DR, DW, D, E>,
CS::ErrorType: Debug,
{
#[error("CycleStage: {0:?}")]
CycleStage(CS::ErrorType),
#[error("Executor: {0:?}")]
Executor(E::ErrorType),
}
pub struct InfParMachine<'a, DR, DW, D, E, CS, IDT, ODT>
where
DR: DataReader,
DW: DataWriter,
D: DataHolder<'a, DR, DW>,
E: Executor<'a, DR, DW, D> + DataTransforms<'a, DR, DW, D, IDT, ODT>,
<E as Executor<'a, DR, DW, D>>::ErrorType: Debug,
<E as DataTransforms<'a, DR, DW, D, IDT, ODT>>::ErrorType: Debug,
CS: CycleStage<'a, DR, DW, D, E>,
CS::ErrorType: Debug,
IDT: DataTransformer<'a, DR, DW, D>,
IDT::ErrorType: Debug,
ODT: DataTransformer<'a, DR, DW, D>,
ODT::ErrorType: Debug,
{
exec_word_len: u32,
config: InfParMachineConfig,
env_config: InfParEnvConfig,
int_states_tx: ODT,
executor: E,
cycle_stage: CS,
int_states: D,
int_state_len: u32,
pidc: ProcIntDataConfig,
proc_ints: D,
memory: D,
memory_last_mask: u32,
cycle_no: u64,
cycle_stage_no: u32,
state: u32,
pdr: PhantomData<&'a DR>,
pdw: PhantomData<&'a DW>,
podt: PhantomData<&'a IDT>,
}
impl<'a, DR, DW, D, E, CS, IDT, ODT> InfParMachine<'a, DR, DW, D, E, CS, IDT, ODT>
where
DR: DataReader,
DW: DataWriter,
D: DataHolder<'a, DR, DW> + RangedData,
E: Executor<'a, DR, DW, D> + DataTransforms<'a, DR, DW, D, IDT, ODT>,
<E as Executor<'a, DR, DW, D>>::ErrorType: Debug,
<E as DataTransforms<'a, DR, DW, D, IDT, ODT>>::ErrorType: Debug,
CS: CycleStage<'a, DR, DW, D, E>,
CS::ErrorType: Debug,
IDT: DataTransformer<'a, DR, DW, D>,
IDT::ErrorType: Debug,
ODT: DataTransformer<'a, DR, DW, D>,
ODT::ErrorType: Debug,
{
pub fn new<B, T>(
builder: B,
config: InfParMachineConfig,
env_config: InfParEnvConfig,
circuit: Circuit<T>,
) -> Result<Self, B::ErrorType>
where
B: Builder<'a, DR, DW, D, E>,
B::ErrorType: Debug,
T: Clone + Copy + Ord + PartialEq + Eq + Hash,
T: Default + TryFrom<usize>,
<T as TryFrom<usize>>::Error: Debug,
usize: TryFrom<T>,
<usize as TryFrom<T>>::Error: Debug,
{
config.valid().unwrap();
env_config.valid().unwrap();
assert!(env_config.flat_memory, "Need flat memory model");
assert!(env_config.max_mem_size.is_some());
assert_ne!(env_config.max_mem_size.unwrap(), 0);
let exec_word_len = builder.word_len();
let circuit_input_len = usize::try_from(circuit.input_len()).unwrap();
let state_len =
circuit_input_len - (1 << config.cell_len_bits) - config.data_part_len as usize - 1;
assert_eq!(state_len, config.state_len as usize);
assert_eq!(circuit.outputs().len(), circuit_input_len + 8);
let memory_size: usize = env_config.max_mem_size.unwrap().try_into().unwrap();
let proc_num_usize: usize = env_config.proc_num.try_into().unwrap();
let pidc = ProcIntDataConfig::new(config, env_config);
let mut exec = build_circuit(builder, circuit, &pidc, proc_num_usize)?;
let cycle_stage = CS::new_from_executor(&exec, config, env_config);
let int_states = exec.new_data_input_elems(proc_num_usize);
let proc_ints = exec.new_data(64 + pidc.len() * proc_num_usize);
let memory = exec.new_data((memory_size + 3) >> 2);
let int_states_tx = exec
.output_transformer((state_len + 31) & !31, &(0..state_len).collect::<Vec<_>>())
.unwrap();
Ok(Self {
exec_word_len,
config,
env_config,
int_states_tx,
executor: exec,
cycle_stage,
int_states,
int_state_len: u32::try_from(state_len).unwrap(),
pidc,
proc_ints,
memory,
memory_last_mask: if (memory_size & 3) != 0 {
(1u32 << ((memory_size & 3) * 8)) - 1u32
} else {
u32::MAX
},
cycle_no: 0,
cycle_stage_no: 0,
state: 0,
pdr: PhantomData,
pdw: PhantomData,
podt: PhantomData,
})
}
pub fn new_from_data<B, T>(builder: B, data: InfParMachineData<T>) -> Result<Self, B::ErrorType>
where
B: Builder<'a, DR, DW, D, E>,
B::ErrorType: Debug,
T: Clone + Copy + Ord + Debug + PartialEq + Eq + Hash + std::ops::Add<Output = T>,
T: FromStr + From<u8>,
<T as FromStr>::Err: Debug,
T: Default + TryFrom<usize>,
<T as TryFrom<usize>>::Error: Debug,
usize: TryFrom<T>,
<usize as TryFrom<T>>::Error: Debug,
{
Self::new(builder, data.config, data.env_config, data.circuit)
}
pub fn initialize(&mut self) {
self.initialize_state();
self.memory.fill(0);
}
pub fn initialize_state(&mut self) {
self.int_states.fill(0);
self.proc_ints.fill(0);
self.cycle_no = 0;
self.cycle_stage_no = 0;
self.state = 0;
}
pub fn execute(&mut self) -> Result<u64, InfParMachineError<'a, DR, DW, D, E, CS>> {
while self.state == 0 {
self.execute_cycle()?;
}
Ok(self.cycle_no)
}
pub fn execute_cycles(
&mut self,
cycles: u64,
) -> Result<u64, InfParMachineError<'a, DR, DW, D, E, CS>> {
for _ in 0..cycles {
if self.state != 0 {
break;
}
self.execute_cycle()?;
}
Ok(self.cycle_no)
}
pub fn execute_cycle(&mut self) -> Result<u64, InfParMachineError<'a, DR, DW, D, E, CS>> {
if self.state == 0 {
for _ in self.cycle_stage_no..4 {
self.execute_cycle_stage()?;
if (self.state & GLOBAL_STATE_ILLEGAL) != 0 {
break;
}
}
}
Ok(self.cycle_no)
}
pub fn execute_cycle_stage(
&mut self,
) -> Result<(u64, u32), InfParMachineError<'a, DR, DW, D, E, CS>> {
if (self.state & GLOBAL_STATE_ILLEGAL) == 0
&& ((self.state & GLOBAL_STATE_STOP) == 0 || self.cycle_stage_no != 0)
{
match self.cycle_stage_no {
0 => {
self.executor
.execute_buffer_single(&mut self.int_states, 0, &mut self.proc_ints)
.map_err(|e| InfParMachineError::Executor(e))?;
}
1 => {
self.cycle_stage
.execute_read(&mut self.proc_ints, &self.memory)
.map_err(|e| InfParMachineError::CycleStage(e))?;
}
2 => {
self.cycle_stage
.execute_clear(&mut self.proc_ints, &mut self.memory)
.map_err(|e| InfParMachineError::CycleStage(e))?;
}
3 => {
self.cycle_stage
.execute_write(&mut self.proc_ints, &mut self.memory)
.map_err(|e| InfParMachineError::CycleStage(e))?;
}
_ => {
panic!("Unexpected!");
}
}
self.cycle_stage_no = if self.cycle_stage_no == 3 {
self.cycle_no += 1;
0
} else {
self.cycle_stage_no + 1
};
let old_len = self.proc_ints.len();
self.proc_ints.set_range(0..1);
self.state = self.proc_ints.process(|d| d[0]);
self.proc_ints.set_range(0..old_len);
}
Ok((self.cycle_no, self.cycle_stage_no))
}
pub fn proc_num(&self) -> u64 {
self.env_config.proc_num
}
pub fn memory_size(&self) -> u64 {
self.env_config.max_mem_size.unwrap()
}
pub fn config(&self) -> &InfParMachineConfig {
&self.config
}
pub fn env_config(&self) -> &InfParEnvConfig {
&self.env_config
}
pub fn cycle_no(&self) -> u64 {
self.cycle_no
}
pub fn cycle_stage_no(&self) -> u32 {
self.cycle_stage_no
}
pub fn state_len(&self) -> u32 {
self.int_state_len
}
pub fn read_memory(&mut self, start: u64, end: u64) -> Vec<u32> {
let start = usize::try_from(start).unwrap();
let end = usize::try_from(end).unwrap();
let mem_length = self.memory.len();
assert!(start <= mem_length);
assert!(end <= mem_length);
self.memory.set_range(start..end);
let mut out = vec![0; end - start];
self.memory.process(|d| {
out.copy_from_slice(d);
if end == mem_length {
*out.last_mut().unwrap() &= self.memory_last_mask;
}
});
self.memory.set_range(0..mem_length);
out
}
pub fn write_memory(&mut self, start: u64, src: &[u32]) {
let start = usize::try_from(start).unwrap();
let mem_length = self.memory.len();
let end = start + src.len();
assert!(start <= mem_length);
assert!(end <= mem_length);
self.memory.set_range(start..end);
self.memory.process_mut(|d| {
d.copy_from_slice(src);
if end == mem_length {
d[end - start - 1] &= self.memory_last_mask;
}
});
self.memory.set_range(0..mem_length);
}
pub fn read_memory_bytes(&mut self, start: u64, end: u64) -> Vec<u8> {
let start = usize::try_from(start).unwrap();
let end = usize::try_from(end).unwrap();
assert_eq!(start & 3, 0);
assert_eq!(end & 3, 0);
let start = start >> 2;
let end = end >> 2;
let mem_length = self.memory.len();
assert!(start <= mem_length);
assert!(end <= mem_length);
self.memory.set_range(start..end);
let mut out = vec![0u8; (end << 2) - (start << 2)];
self.memory.process(|d| {
for (i, v) in d.iter().enumerate() {
let vb = v.to_ne_bytes();
out[i << 2] = vb[0];
out[(i << 2) + 1] = vb[1];
out[(i << 2) + 2] = vb[2];
out[(i << 2) + 3] = vb[3];
}
if end == mem_length {
let i = end - start - 1;
let vb = self.memory_last_mask.to_ne_bytes();
out[i << 2] &= vb[0];
out[(i << 2) + 1] &= vb[1];
out[(i << 2) + 2] &= vb[2];
out[(i << 2) + 3] &= vb[3];
}
});
self.memory.set_range(0..mem_length);
out
}
pub fn write_memory_bytes(&mut self, start: u64, src: &[u8]) {
let start = usize::try_from(start).unwrap();
let end = start + src.len();
assert_eq!(start & 3, 0);
assert_eq!(end & 3, 0);
let start = start >> 2;
let end = end >> 2;
let mem_length = self.memory.len();
assert!(start <= mem_length);
assert!(end <= mem_length);
self.memory.set_range(start..end);
self.memory.process_mut(|d| {
for (i, v) in d.iter_mut().enumerate() {
let s = [
src[i << 2],
src[(i << 2) + 1],
src[(i << 2) + 2],
src[(i << 2) + 3],
];
*v = u32::from_ne_bytes(s);
}
if end == mem_length {
d[end - start - 1] &= self.memory_last_mask;
}
});
self.memory.set_range(0..mem_length);
}
pub fn states(&mut self, start: u64, end: u64) -> (Vec<u32>, u32) {
let start: usize = start.try_into().unwrap();
let end: usize = end.try_into().unwrap();
let exec_word_len_usize = self.exec_word_len as usize;
let state_len_in_dwords = ((self.int_state_len + 31) >> 5) as usize;
let start_w: usize = (start / exec_word_len_usize).try_into().unwrap();
let end_w: usize = ((end + exec_word_len_usize - 1) / exec_word_len_usize)
.try_into()
.unwrap();
let int_all_states_len = self.int_states.len();
let start_w2 = start_w * (self.int_state_len as usize) * (exec_word_len_usize >> 5);
let end_w2 = end_w * (self.int_state_len as usize) * (exec_word_len_usize >> 5);
self.int_states.set_range(start_w2..end_w2);
let data = self.int_states_tx.transform(&self.int_states).unwrap();
let out = data.process(|d| {
let mut out = vec![0; ((end - start) as usize) * state_len_in_dwords];
out.copy_from_slice(
&d[((start as usize) - start_w * exec_word_len_usize) * state_len_in_dwords
..((end as usize) - start_w * exec_word_len_usize) * state_len_in_dwords],
);
out
});
self.int_states.set_range(0..int_all_states_len);
(out, state_len_in_dwords as u32)
}
pub fn proc_int_data(&mut self, start: u64, end: u64) -> ProcIntDataReader {
assert!(start <= self.env_config.proc_num);
assert!(end <= self.env_config.proc_num);
let start: usize = start.try_into().unwrap();
let end: usize = end.try_into().unwrap();
let entry_len = self.pidc.len();
let proc_ints_len = self.proc_ints.len();
self.proc_ints
.set_range(64 + start * entry_len..64 + end * entry_len);
let pidr = self
.proc_ints
.process(|d| ProcIntDataReader::new_from_config(self.pidc, d.to_vec()));
self.proc_ints.set_range(0..proc_ints_len);
pidr
}
pub fn stopped(&self) -> bool {
(self.state & GLOBAL_STATE_STOP) != 0
}
pub fn illegal_state(&self) -> bool {
(self.state & GLOBAL_STATE_ILLEGAL) != 0
}
}
pub type CPUInfParMachine<'a> = InfParMachine<
'a,
CPUDataReader<'a>,
CPUDataWriter<'a>,
CPUDataHolder,
CPUExecutor,
CPUCycleStage,
CPUDataInputTransformer,
CPUDataOutputTransformer,
>;
pub type OpenCLInfParMachine<'a> = InfParMachine<
'a,
OpenCLDataReader<'a>,
OpenCLDataWriter<'a>,
OpenCLDataHolder,
OpenCLExecutor,
OpenCLCycleStage,
OpenCLDataInputTransformer,
OpenCLDataOutputTransformer,
>;
#[cfg(test)]
mod tests {
use super::*;
use gatenative::clang_writer::*;
#[test]
fn test_infparmachine() {
let builder = CPUBuilder::new_with_cpu_ext_and_clang_config(
CPUExtension::NoExtension,
&CLANG_WRITER_U64,
None,
);
let config = InfParMachineConfig {
state_len: 44,
data_part_len: 32,
cell_len_bits: 5,
};
let env_config = InfParEnvConfig {
proc_num: 64 * 1024,
flat_memory: true,
max_temp_buffer_len: 96,
max_mem_size: Some(400 * 1024),
};
let pic = ProcIntDataConfig::new(config, env_config);
println!("EntryLen: {}", pic.len());
let mem_cell_data_part_len = (1 << config.cell_len_bits) + config.data_part_len;
let circuit_input_len = config.state_len + mem_cell_data_part_len + 1;
let circuit = Circuit::<u32>::new(
circuit_input_len,
[Gate::new_nor(0, circuit_input_len - 1)],
std::iter::once((circuit_input_len, true)).chain(
(1..circuit_input_len - 1)
.chain(0..9)
.map(|i| (i, i >= config.state_len && i < circuit_input_len - 1)),
),
)
.unwrap();
let mut machine =
CPUInfParMachine::<'_>::new(builder, config, env_config, circuit).unwrap();
assert_eq!(machine.proc_num(), 64 * 1024);
{
let mut mem = machine.memory.get_mut();
let mem = mem.get_mut();
for i in 0..100 {
mem[4421 + i] = 532566 + 6 * u32::try_from(i).unwrap();
}
for i in 0..200 {
mem[46969 + i] = 0xddaa00 + 7 * u32::try_from(i).unwrap();
}
mem[100 * 1024 - 1] = 0xaa22cc77;
}
assert_eq!(
(0..100).map(|i| 532566 + 6 * i).collect::<Vec<_>>(),
machine.read_memory(4421, 4421 + 100),
);
assert_eq!(
(0..200).map(|i| 0xddaa00 + 7 * i).collect::<Vec<_>>(),
machine.read_memory(46969, 46969 + 200),
);
assert_eq!(
vec![0, 0xaa22cc77],
machine.read_memory(100 * 1024 - 2, 100 * 1024),
);
assert_eq!(
(0..100)
.map(|i| (532566u32 + 6 * i).to_ne_bytes())
.flatten()
.collect::<Vec<_>>(),
machine.read_memory_bytes(4 * 4421, 4 * (4421 + 100)),
);
assert_eq!(
(0..200)
.map(|i| (0xddaa00u32 + 7 * i).to_ne_bytes())
.flatten()
.collect::<Vec<_>>(),
machine.read_memory_bytes(4 * 46969, 4 * (46969 + 200)),
);
assert_eq!(
0xaa22cc77u32.to_ne_bytes().to_vec(),
machine.read_memory_bytes(4 * (100 * 1024 - 1), 4 * 100 * 1024),
);
machine.write_memory(4500, &(0..200).map(|i| 244211 + 10 * i).collect::<Vec<_>>());
{
let mem = machine.memory.get();
let mem = mem.get();
assert_eq!(
(0..200).map(|i| 244211 + 10 * i).collect::<Vec<_>>(),
&mem[4500..4500 + 200]
);
}
machine.write_memory_bytes(
4 * 4500,
&(0..200)
.map(|i| (244277u32 + 9 * i).to_ne_bytes())
.flatten()
.collect::<Vec<_>>(),
);
{
let mem = machine.memory.get();
let mem = mem.get();
assert_eq!(
(0..200).map(|i| 244277 + 9 * i).collect::<Vec<_>>(),
&mem[4500..4500 + 200]
);
}
let mut it = machine
.executor
.input_transformer(
((config.state_len + 31) & !31) as usize,
&(0..config.state_len)
.map(|i| i as usize)
.collect::<Vec<_>>(),
)
.unwrap();
let states_data = machine.executor.new_data_from_vec(
(0..4096)
.map(|i| [0x677da3b + 5 * i, (0x545a44 + 3 * i) & 0xfff])
.flatten()
.collect::<Vec<_>>(),
);
let states_data = it.transform(&states_data).unwrap();
{
let mut int_states = machine.int_states.get_mut();
let int_states = int_states.get_mut();
let states_data = states_data.get();
let states_data = states_data.get();
int_states[14 * 704 * 2..(14 + 4) * 704 * 2].copy_from_slice(&states_data);
}
let (states, state_len) = machine.states(14 * 1024, 18 * 1024);
assert_eq!(2, state_len);
assert_eq!(2 * 4096, states.len());
for i in 0..4096 {
let j = u32::try_from(i).unwrap();
assert_eq!(
&states[2 * i..2 * i + 2],
[0x677da3b + 5 * j, (0x545a44 + 3 * j) & 0xfff],
"{}",
i
);
}
let (states, state_len) = machine.states(14 * 1024 + 21, 18 * 1024 - 7);
assert_eq!(2, state_len);
assert_eq!(2 * (4096 - 21 - 7), states.len());
for i in 21..4096 - 7 {
let j = u32::try_from(i).unwrap();
assert_eq!(
&states[2 * (i - 21)..2 * (i - 21) + 2],
[0x677da3b + 5 * j, (0x545a44 + 3 * j) & 0xfff],
"{}",
i
);
}
let (states, state_len) = machine.states(14 * 1024 + 20, 18 * 1024 - 32 + 11);
assert_eq!(2, state_len);
assert_eq!(2 * (4096 - 20 - 32 + 11), states.len());
for i in 20..4096 - 32 + 11 {
let j = u32::try_from(i).unwrap();
assert_eq!(
&states[2 * (i - 20)..2 * (i - 20) + 2],
[0x677da3b + 5 * j, (0x545a44 + 3 * j) & 0xfff],
"{}",
i
);
}
let (states, state_len) = machine.states(14 * 1024 + 47, 18 * 1024 - 64 + 13);
assert_eq!(2, state_len);
assert_eq!(2 * (4096 - 47 - 64 + 13), states.len());
for i in 47..4096 - 64 + 13 {
let j = u32::try_from(i).unwrap();
assert_eq!(
&states[2 * (i - 47)..2 * (i - 47) + 2],
[0x677da3b + 5 * j, (0x545a44 + 3 * j) & 0xfff],
"{}",
i
);
}
{
let mut pidata = machine.proc_ints.get_mut();
let pidata = pidata.get_mut();
let entry_len = pic.len();
for i in 0..10 {
let j = u32::try_from(i).unwrap();
pidata[64 + (120 + i) * entry_len] = 0xada355 + j;
pidata[64 + (120 + i) * entry_len + 1] = 0x2a78 + j;
}
}
let pidr = machine.proc_int_data(120, 130);
for i in 0..10u64 {
let j = u32::try_from(i).unwrap();
assert_eq!(0xada355 + u64::from(j), pic.mem_address(&pidr[i]));
assert_eq!(vec![0x2a78 + j, 0, 0], pic.temp_buffer(&pidr[i]));
}
let builder = CPUBuilder::new_with_cpu_ext_and_clang_config(
CPUExtension::NoExtension,
&CLANG_WRITER_U64,
None,
);
let env_config = InfParEnvConfig {
proc_num: 64 * 1024,
flat_memory: true,
max_temp_buffer_len: 96,
max_mem_size: Some(400 * 1024 - 3),
};
let pic = ProcIntDataConfig::new(config, env_config);
println!("EntryLen: {}", pic.len());
let mem_cell_data_part_len = (1 << config.cell_len_bits) + config.data_part_len;
let circuit_input_len = config.state_len + mem_cell_data_part_len + 1;
let circuit = Circuit::<u32>::new(
circuit_input_len,
[Gate::new_nor(0, circuit_input_len - 1)],
std::iter::once((circuit_input_len, true)).chain(
(1..circuit_input_len - 1)
.chain(0..9)
.map(|i| (i, i >= config.state_len && i < circuit_input_len - 1)),
),
)
.unwrap();
let mut machine =
CPUInfParMachine::<'_>::new(builder, config, env_config, circuit).unwrap();
{
let mut mem = machine.memory.get_mut();
let mem = mem.get_mut();
mem[100 * 1024 - 1] = 0xaa22cc77;
}
assert_eq!(
vec![0, 0x77],
machine.read_memory(100 * 1024 - 2, 100 * 1024),
);
assert_eq!(
0x77u32.to_ne_bytes().to_vec(),
machine.read_memory_bytes(4 * (100 * 1024 - 1), 4 * 100 * 1024),
);
machine.write_memory(100 * 1024 - 2, &[0, 0xab348a]);
{
let mem = machine.memory.get();
let mem = mem.get();
assert_eq!(vec![0, 0x8a], &mem[100 * 1024 - 2..]);
}
machine.write_memory_bytes(4 * (100 * 1024 - 1), &[0x5e, 0x11, 0x22, 0x33]);
{
let mem = machine.memory.get();
let mem = mem.get();
assert_eq!(vec![0, 0x5e], &mem[100 * 1024 - 2..]);
}
}
}