use hashbrown::HashMap;
use crate::ast::Ast;
use crate::lexer::Lexer;
use crate::error::Error;
use cranelift::prelude::*;
use cranelift_module::{DataContext, Linkage, Module};
use cranelift_simplejit::{SimpleJITBackend, SimpleJITBuilder};
use libm::pow;
use std::mem;
use std::slice;
const POW: &str = "pow";
pub struct JIT {
builder_context: FunctionBuilderContext,
ctx: codegen::Context,
data_ctx: DataContext,
module: Module<SimpleJITBackend>,
required_parameters: HashMap<String, usize>,
}
impl Default for JIT {
#[must_use]
fn default() -> Self {
if cfg!(windows) {
unimplemented!();
}
let mut builder = SimpleJITBuilder::new(cranelift_module::default_libcall_names());
let _s = builder.symbol(POW, pow as *const u8);
let module = Module::new(builder);
Self {
builder_context: FunctionBuilderContext::new(),
ctx: module.make_context(),
data_ctx: DataContext::new(),
module,
required_parameters: HashMap::new(),
}
}
}
impl JIT {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn compile(
&mut self,
input: &str,
) -> Result<Box<dyn Fn(HashMap<String, &[f64]>, usize) -> Result<Vec<f64>, Error>>, Error> {
let mut lexer = Lexer::new(input);
let ast = Ast::from_tokens(&mut lexer.parse().map_err(|e| Error::ParseError(e.to_string()))?, "")
.map_err(|e| Error::ParseError(e.to_string()))?;
self.translate(&ast)?;
let id = self
.module
.declare_function(&input, Linkage::Export, &self.ctx.func.signature)
.map_err(|e| Error::ParseError(e.to_string()))?;
self.module
.define_function(id, &mut self.ctx)
.map_err(|e| Error::ParseError(e.to_string()))?;
self.module.clear_context(&mut self.ctx);
self.module.finalize_definitions();
let code = self.module.get_finalized_function(id);
Ok(Self::dynamic_param_fn(
code,
&self.required_parameters,
))
}
fn dynamic_param_fn(
function: *const u8,
required_parameters: &HashMap<String, usize>,
) -> Box<dyn Fn(HashMap<String, &[f64]>, usize) -> Result<Vec<f64>, Error>> {
let keys = required_parameters.keys();
let mut sorted_keys = vec![];
for k in keys {
sorted_keys.push(k.to_owned());
}
sorted_keys.sort_unstable();
match required_parameters.len() {
0 => {
let function = unsafe { mem::transmute::<_, fn() -> f64>(function) };
Box::new(
move |_params: HashMap<String, &[f64]>, number_of_evaluations: usize| {
let mut results: Vec<f64> = Vec::with_capacity(number_of_evaluations);
for _i in 0..number_of_evaluations {
results.push(function());
}
Ok(results)
},
)
}
1 => {
let function = unsafe { mem::transmute::<_, fn(f64) -> f64>(function) };
Box::new(
move |params: HashMap<String, &[f64]>, number_of_evaluations: usize| {
let mut results: Vec<f64> = Vec::with_capacity(number_of_evaluations);
let param_1 = params.get(&sorted_keys[0]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[0])))?;
for k in &sorted_keys {
if number_of_evaluations > params[k].len() {
return Err(Error::NameError(format!("Missing data for parameter: {}", k)));
}
}
for i in 0..number_of_evaluations {
results.push(function(param_1[i]));
}
Ok(results)
},
)
}
2 => {
let function = unsafe { mem::transmute::<_, fn(f64, f64) -> f64>(function) };
Box::new(
move |params: HashMap<String, &[f64]>, number_of_evaluations: usize| {
let mut results: Vec<f64> = Vec::with_capacity(number_of_evaluations);
let param_1 = params.get(&sorted_keys[0]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[0])))?;
let param_2 = params.get(&sorted_keys[1]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[1])))?;
for k in &sorted_keys {
if number_of_evaluations > params[k].len() {
return Err(Error::NameError(format!("Missing data for parameter: {}", k)));
}
}
for i in 0..number_of_evaluations {
results.push(function(param_1[i], param_2[i]));
}
Ok(results)
},
)
}
3 => {
let function = unsafe { mem::transmute::<_, fn(f64, f64, f64) -> f64>(function) };
Box::new(
move |params: HashMap<String, &[f64]>, number_of_evaluations: usize| {
let mut results: Vec<f64> = Vec::with_capacity(number_of_evaluations);
let param_1 = params.get(&sorted_keys[0]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[0])))?;
let param_2 = params.get(&sorted_keys[1]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[1])))?;
let param_3 = params.get(&sorted_keys[2]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[2])))?;
for k in &sorted_keys {
if number_of_evaluations > params[k].len() {
return Err(Error::NameError(format!("Missing data for parameter: {}", k)));
}
}
for i in 0..number_of_evaluations {
results.push(function(param_1[i], param_2[i], param_3[i]));
}
Ok(results)
},
)
}
4 => {
let function =
unsafe { mem::transmute::<_, fn(f64, f64, f64, f64) -> f64>(function) };
Box::new(
move |params: HashMap<String, &[f64]>, number_of_evaluations: usize| {
let mut results: Vec<f64> = Vec::with_capacity(number_of_evaluations);
let param_1 = params.get(&sorted_keys[0]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[0])))?;
let param_2 = params.get(&sorted_keys[1]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[1])))?;
let param_3 = params.get(&sorted_keys[2]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[2])))?;
let param_4 = params.get(&sorted_keys[3]).ok_or_else(|| Error::NameError(format!("Missing parameter: {}", &sorted_keys[3])))?;
for k in &sorted_keys {
if number_of_evaluations > params[k].len() {
return Err(Error::NameError(format!("Missing data for parameter: {}", k)));
}
}
for i in 0..number_of_evaluations {
results.push(function(param_1[i], param_2[i], param_3[i], param_4[i]));
}
Ok(results)
},
)
}
_ => panic!(),
}
}
pub fn create_data(&mut self, name: &str, contents: Vec<u8>) -> Result<&[u8], Error> {
self.data_ctx.define(contents.into_boxed_slice());
let id = self
.module
.declare_data(name, Linkage::Export, true, None)
.map_err(|e| Error::NameError(e.to_string()))?;
self.module
.define_data(id, &self.data_ctx)
.map_err(|e| Error::ParseError(e.to_string()))?;
self.data_ctx.clear();
self.module.finalize_definitions();
let buffer = self.module.get_finalized_data(id);
Ok(unsafe { slice::from_raw_parts(buffer.0, buffer.1) })
}
fn get_parameters<'a>(ast: &'a Ast, context: &mut Vec<&'a str>) {
match ast {
Ast::Variable(name) => {
context.push(name);
}
Ast::Value(_) => {}
Ast::Function(_, ref arg) => {
Self::get_parameters(arg, context);
}
Ast::Add(ref left, ref right)
| Ast::Sub(ref left, ref right)
| Ast::Mul(ref left, ref right)
| Ast::Div(ref left, ref right)
| Ast::Exp(ref left, ref right) => {
Self::get_parameters(left, context);
Self::get_parameters(right, context);
}
}
}
fn translate(&mut self, ast: &Ast) -> Result<(), Error> {
let mut parameter_names = vec![];
Self::get_parameters(ast, &mut parameter_names);
parameter_names.sort_unstable();
parameter_names.dedup();
for (i, param_name) in parameter_names.iter().enumerate() {
self.required_parameters.insert((*param_name).to_owned(), i);
}
for _p in ¶meter_names {
self.ctx
.func
.signature
.params
.push(AbiParam::new(types::F64));
}
self.ctx
.func
.signature
.returns
.push(AbiParam::new(types::F64));
let mut builder = FunctionBuilder::new(&mut self.ctx.func, &mut self.builder_context);
let entry_ebb = builder.create_ebb();
builder.append_ebb_params_for_function_params(entry_ebb);
builder.switch_to_block(entry_ebb);
builder.seal_block(entry_ebb);
let variables = declare_variables(&mut builder, ¶meter_names, &ast, entry_ebb);
let mut trans = FunctionTranslator {
builder,
variables,
module: &mut self.module,
};
let expression_val = trans.translate_expr(&ast);
let return_variable = trans.variables[".the_return"];
trans.builder.def_var(return_variable, expression_val);
let return_value = trans.builder.use_var(return_variable);
trans.builder.ins().return_(&[return_value]);
trans.builder.finalize();
Ok(())
}
}
struct FunctionTranslator<'a> {
builder: FunctionBuilder<'a>,
variables: HashMap<String, Variable>,
module: &'a mut Module<SimpleJITBackend>,
}
impl<'a> FunctionTranslator<'a> {
fn translate_expr(&mut self, ast: &Ast) -> Value {
match *ast {
Ast::Value(val) => self.builder.ins().f64const(Ieee64::with_float(val)),
Ast::Variable(ref name) => {
let variable = self.variables.get(name).expect("variable not defined");
self.builder.use_var(*variable)
}
Ast::Add(ref left, ref right) => {
let lhs = self.translate_expr(left);
let rhs = self.translate_expr(right);
self.builder.ins().fadd(lhs, rhs)
}
Ast::Sub(ref left, ref right) => {
let lhs = self.translate_expr(left);
let rhs = self.translate_expr(right);
self.builder.ins().fsub(lhs, rhs)
}
Ast::Mul(ref left, ref right) => {
let lhs = self.translate_expr(left);
let rhs = self.translate_expr(right);
self.builder.ins().fmul(lhs, rhs)
}
Ast::Div(ref left, ref right) => {
let lhs = self.translate_expr(left);
let rhs = self.translate_expr(right);
self.builder.ins().fdiv(lhs, rhs)
}
Ast::Exp(ref left, ref right) => {
let lhs = self.translate_expr(left);
let rhs = self.translate_expr(right);
self.translate_call(POW, &[lhs, rhs])
}
_ => self.builder.ins().f64const(Ieee64::with_float(0.0)),
}
}
fn translate_call(&mut self, name: &str, args: &[Value]) -> Value {
let mut sig = self.module.make_signature();
for _arg in args {
sig.params.push(AbiParam::new(types::F64));
}
sig.returns.push(AbiParam::new(types::F64));
let callee = self
.module
.declare_function(name, Linkage::Import, &sig)
.expect("problem declaring function");
let local_callee = self
.module
.declare_func_in_func(callee, &mut self.builder.func);
let call = self.builder.ins().call(local_callee, &args);
self.builder.inst_results(call)[0]
}
fn translate_global_data_addr(&mut self, name: String) -> Value {
let sym = self
.module
.declare_data(&name, Linkage::Export, true, None)
.expect("problem declaring data object");
let local_id = self
.module
.declare_data_in_func(sym, &mut self.builder.func);
let pointer = self.module.target_config().pointer_type();
self.builder.ins().symbol_value(pointer, local_id)
}
}
fn declare_variables(
builder: &mut FunctionBuilder,
params: &[&str],
stmts: &Ast,
entry_ebb: Ebb,
) -> HashMap<String, Variable> {
let mut variables = HashMap::new();
let mut index = 0;
for (i, name) in params.iter().enumerate() {
let value = builder.ebb_params(entry_ebb)[i];
let var = declare_variable(builder, &mut variables, &mut index, name);
builder.def_var(var, value);
}
let zero = builder.ins().f64const(Ieee64::with_float(0.0));
let return_variable = declare_variable(builder, &mut variables, &mut index, ".the_return");
builder.def_var(return_variable, zero);
variables
}
fn declare_variable(
builder: &mut FunctionBuilder,
variables: &mut HashMap<String, Variable>,
index: &mut usize,
name: &str,
) -> Variable {
let var = Variable::new(*index);
if !variables.contains_key(name) {
variables.insert(name.into(), var);
builder.declare_var(var, types::F64);
*index += 1;
}
var
}
#[cfg(test)]
mod tests {
use super::HashMap;
use std::process;
use std::time::Instant;
#[test]
fn bench() {
let watch = Instant::now();
let mut jit = super::JIT::new();
let foo_code = "(var1 + var2 * 3) / (2 + 3) - something";
let compiled_formula = jit.compile(&foo_code).unwrap_or_else(|msg| {
dbg!(msg);
process::exit(1);
});
let mut dict: HashMap<String, Vec<f64>> = HashMap::with_capacity(3);
let capacity = 5_000_000;
let iterations = 5_000_000;
dict.insert("var1".to_owned(), Vec::with_capacity(capacity));
dict.insert("var2".to_owned(), Vec::with_capacity(capacity));
dict.insert("something".to_owned(), Vec::with_capacity(capacity));
for i in 1..=iterations {
dict.get_mut("var1").unwrap().push(10.0 + f64::from(i));
dict.get_mut("var2").unwrap().push(20.0 + f64::from(i));
dict.get_mut("something").unwrap().push(30.0 + f64::from(i));
}
let dict2: HashMap<String, &[f64]> = dict
.iter_mut()
.map(|v| (v.0.to_owned(), v.1.as_slice()))
.collect();
let watch = watch.elapsed();
let watch2 = Instant::now();
let results = compiled_formula(dict2, capacity);
let watch2 = watch2.elapsed();
match results {
Ok(results) => println!("{}", results[0]),
Err(msg) => println!("{}", msg)
}
println!("{}", watch.as_millis());
println!("{}", watch2.as_millis());
}
}