use std::collections::HashMap;
use cranelift_jit::JITModule;
use crate::ast::PolydatNode;
use crate::kernel::ProvMask;
pub(super) struct JitCore {
pub(super) buffer: Vec<u64>,
pub(super) coord_count: usize,
pub(super) output_map: HashMap<String, usize>,
pub(super) _module: JITModule,
pub(super) _nodes: Vec<Box<dyn PolydatNode>>,
}
pub(super) fn compute_jit_slot_provenance(
coord_count: usize,
buffer_len: usize,
step_output_slots: &[Vec<usize>],
input_dependents: &[Vec<usize>],
) -> Vec<ProvMask> {
let step_count = step_output_slots.len();
let mut step_prov: Vec<ProvMask> =
(0..step_count).map(|_| ProvMask::empty()).collect();
for (input_slot, deps) in input_dependents.iter().enumerate() {
for &step_idx in deps {
if step_idx < step_count {
step_prov[step_idx].set(input_slot);
}
}
}
let mut slot_prov: Vec<ProvMask> =
(0..buffer_len).map(|_| ProvMask::empty()).collect();
for (i, slot) in slot_prov.iter_mut().enumerate().take(coord_count) {
slot.set(i);
}
for (i, outs) in step_output_slots.iter().enumerate() {
for &slot in outs {
if slot < slot_prov.len() {
slot_prov[slot] = step_prov[i].clone();
}
}
}
slot_prov
}
macro_rules! jit_accessors {
() => {
pub fn coord_count(&self) -> usize { self.core.coord_count }
pub fn resolve_output(&self, name: &str) -> Option<usize> {
self.core.output_map.get(name).copied()
}
#[inline]
pub fn get(&self, name: &str) -> u64 {
self.core.buffer[self.core.output_map[name]]
}
#[inline]
pub fn get_slot(&self, slot: usize) -> u64 {
self.core.buffer[slot]
}
};
}
pub struct JitKernelRaw {
pub(super) core: JitCore,
pub(super) code_fn: unsafe fn(*const u64, *mut u64),
}
impl JitKernelRaw {
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.core.buffer[..self.core.coord_count.min(coords.len())]
.copy_from_slice(&coords[..self.core.coord_count.min(coords.len())]);
let code_fn = self.code_fn;
let buf_ptr_const = self.core.buffer.as_ptr();
let buf_ptr_mut = self.core.buffer.as_mut_ptr();
super::codegen::invoke_with_catch(move || {
unsafe { (code_fn)(buf_ptr_const, buf_ptr_mut); }
});
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.eval(coords);
self.core.buffer[slot]
}
pub fn into_parts(self) -> (unsafe fn(*const u64, *mut u64), JITModule) {
(self.code_fn, self.core._module)
}
jit_accessors!();
}
pub struct JitKernelPush {
pub(super) core: JitCore,
pub(super) code_fn_prov: unsafe fn(*const u64, *mut u64, *mut u8),
pub(super) node_clean: Vec<u8>,
pub(super) input_dependents: Vec<Vec<usize>>,
}
impl JitKernelPush {
#[inline]
fn set_inputs(&mut self, coords: &[u64]) {
for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
if i < self.input_dependents.len() {
for &step_idx in &self.input_dependents[i] {
self.node_clean[step_idx] = 0;
}
}
}
}
}
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.set_inputs(coords);
let code_fn = self.code_fn_prov;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
let clean_mut = self.node_clean.as_mut_ptr();
super::codegen::invoke_with_catch(move || {
unsafe { (code_fn)(buf_const, buf_mut, clean_mut); }
});
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.eval(coords);
self.core.buffer[slot]
}
jit_accessors!();
}
pub struct JitKernelPull {
pub(super) core: JitCore,
pub(super) code_fn: unsafe fn(*const u64, *mut u64),
pub(super) slot_provenance: Vec<ProvMask>,
pub(super) changed_mask: ProvMask,
}
impl JitKernelPull {
#[inline]
fn set_inputs(&mut self, coords: &[u64]) {
self.changed_mask.clear();
for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
self.changed_mask.set(i);
}
}
}
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.set_inputs(coords);
let code_fn = self.code_fn;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
super::codegen::invoke_with_catch(move || {
unsafe { (code_fn)(buf_const, buf_mut); }
});
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.set_inputs(coords);
if slot < self.slot_provenance.len()
&& !self.slot_provenance[slot].intersects(&self.changed_mask) {
return self.core.buffer[slot];
}
let code_fn = self.code_fn;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
super::codegen::invoke_with_catch(move || {
unsafe { (code_fn)(buf_const, buf_mut); }
});
self.core.buffer[slot]
}
jit_accessors!();
}
pub struct JitKernelPushPull {
pub(super) core: JitCore,
pub(super) code_fn_prov: unsafe fn(*const u64, *mut u64, *mut u8),
pub(super) node_clean: Vec<u8>,
pub(super) input_dependents: Vec<Vec<usize>>,
pub(super) slot_provenance: Vec<ProvMask>,
pub(super) changed_mask: ProvMask,
}
impl JitKernelPushPull {
#[inline]
fn set_inputs(&mut self, coords: &[u64]) {
self.changed_mask.clear();
for (i, &c) in coords.iter().enumerate().take(self.core.coord_count) {
if self.core.buffer[i] != c {
self.core.buffer[i] = c;
self.changed_mask.set(i);
if i < self.input_dependents.len() {
for &step_idx in &self.input_dependents[i] {
self.node_clean[step_idx] = 0;
}
}
}
}
}
#[inline]
pub fn eval(&mut self, coords: &[u64]) {
self.set_inputs(coords);
let code_fn = self.code_fn_prov;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
let clean_mut = self.node_clean.as_mut_ptr();
super::codegen::invoke_with_catch(move || {
unsafe { (code_fn)(buf_const, buf_mut, clean_mut); }
});
}
#[inline]
pub fn eval_for_slot(&mut self, coords: &[u64], slot: usize) -> u64 {
self.set_inputs(coords);
if slot < self.slot_provenance.len()
&& !self.slot_provenance[slot].intersects(&self.changed_mask) {
return self.core.buffer[slot];
}
let code_fn = self.code_fn_prov;
let buf_const = self.core.buffer.as_ptr();
let buf_mut = self.core.buffer.as_mut_ptr();
let clean_mut = self.node_clean.as_mut_ptr();
super::codegen::invoke_with_catch(move || {
unsafe { (code_fn)(buf_const, buf_mut, clean_mut); }
});
self.core.buffer[slot]
}
jit_accessors!();
}