use crate::graph::{ComputationGraph, NodeId};
use crate::{CompiledKernel, ExecutionStats, JitError, JitResult, TensorRef};
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Instant;
#[derive(Clone)]
pub struct JitRuntime {
cache: Arc<Mutex<KernelCache>>,
stats: Arc<Mutex<ExecutionStats>>,
config: RuntimeConfig,
}
impl JitRuntime {
pub fn new(config: crate::JitConfig) -> Self {
Self {
cache: Arc::new(Mutex::new(KernelCache::new())),
stats: Arc::new(Mutex::new(ExecutionStats::default())),
config: RuntimeConfig::from_jit_config(config),
}
}
pub fn execute(
&self,
graph: &ComputationGraph,
kernels: &[CompiledKernel],
inputs: &[TensorRef],
) -> JitResult<Vec<TensorRef>> {
let start_time = Instant::now();
let mut context = ExecutionContext::new(graph, inputs)?;
for kernel in kernels {
self.execute_kernel(&mut context, kernel)?;
}
self.update_stats(start_time.elapsed().as_micros() as u64, kernels.len());
context.get_outputs()
}
fn execute_kernel(
&self,
context: &mut ExecutionContext,
kernel: &CompiledKernel,
) -> JitResult<()> {
let cache_hit = if self.config.enable_caching {
self.cache
.lock()
.expect("lock should not be poisoned")
.get(&kernel.id)
.is_some()
} else {
false
};
if cache_hit {
let mut cache = self.cache.lock().expect("lock should not be poisoned");
if let Some(exec_fn) = cache.get(&kernel.id) {
exec_fn(context, kernel)?;
}
} else {
let exec_fn = self.compile_kernel(kernel)?;
exec_fn(context, kernel)?;
if self.config.enable_caching {
self.cache
.lock()
.expect("cache lock should not be poisoned")
.insert(kernel.id.clone(), exec_fn);
}
}
Ok(())
}
fn compile_kernel(&self, _kernel: &CompiledKernel) -> JitResult<ExecutableFn> {
Ok(Box::new(move |context, kernel| {
interpreter_execute(context, kernel)
}))
}
fn update_stats(&self, elapsed_us: u64, kernel_count: usize) {
let mut stats = self.stats.lock().expect("lock should not be poisoned");
stats.total_time_us += elapsed_us;
stats.kernel_launches += kernel_count;
let cache = self.cache.lock().expect("lock should not be poisoned");
stats.cache_hit_rate = cache.hit_rate();
}
pub fn stats(&self) -> ExecutionStats {
self.stats
.lock()
.expect("lock should not be poisoned")
.clone()
}
pub fn clear_cache(&self) {
self.cache
.lock()
.expect("lock should not be poisoned")
.clear();
}
}
#[derive(Debug, Clone)]
struct RuntimeConfig {
enable_caching: bool,
#[allow(dead_code)]
enable_profiling: bool,
#[allow(dead_code)]
max_cache_size: usize,
}
impl RuntimeConfig {
fn from_jit_config(config: crate::JitConfig) -> Self {
Self {
enable_caching: config.enable_caching,
enable_profiling: config.enable_profiling,
max_cache_size: 1000, }
}
}
struct KernelCache {
cache: HashMap<String, ExecutableFn>,
hits: usize,
misses: usize,
max_size: usize,
}
impl KernelCache {
fn new() -> Self {
Self {
cache: HashMap::new(),
hits: 0,
misses: 0,
max_size: 1000,
}
}
fn get(&mut self, key: &str) -> Option<&ExecutableFn> {
if self.cache.contains_key(key) {
self.hits += 1;
self.cache.get(key)
} else {
self.misses += 1;
None
}
}
fn insert(&mut self, key: String, value: ExecutableFn) {
if self.cache.len() >= self.max_size {
if let Some(first_key) = self.cache.keys().next().cloned() {
self.cache.remove(&first_key);
}
}
self.cache.insert(key, value);
}
fn clear(&mut self) {
self.cache.clear();
self.hits = 0;
self.misses = 0;
}
fn hit_rate(&self) -> f32 {
let total = self.hits + self.misses;
if total > 0 {
self.hits as f32 / total as f32
} else {
0.0
}
}
}
type ExecutableFn =
Box<dyn Fn(&mut ExecutionContext, &CompiledKernel) -> JitResult<()> + Send + Sync>;
pub struct ExecutionContext {
#[allow(dead_code)]
inputs: Vec<TensorRef>,
intermediates: HashMap<NodeId, TensorRef>,
output_ids: Vec<NodeId>,
}
impl ExecutionContext {
fn new(graph: &ComputationGraph, inputs: &[TensorRef]) -> JitResult<Self> {
if inputs.len() != graph.inputs.len() {
return Err(JitError::RuntimeError(format!(
"Expected {} inputs, got {}",
graph.inputs.len(),
inputs.len()
)));
}
let mut intermediates = HashMap::new();
for (i, &node_id) in graph.inputs.iter().enumerate() {
intermediates.insert(node_id, inputs[i].clone());
}
Ok(Self {
inputs: inputs.to_vec(),
intermediates,
output_ids: graph.outputs.clone(),
})
}
pub fn get_tensor(&self, node_id: NodeId) -> Option<&TensorRef> {
self.intermediates.get(&node_id)
}
pub fn set_tensor(&mut self, node_id: NodeId, tensor: TensorRef) {
self.intermediates.insert(node_id, tensor);
}
fn get_outputs(&self) -> JitResult<Vec<TensorRef>> {
let mut outputs = Vec::new();
for &output_id in &self.output_ids {
let tensor = self.intermediates.get(&output_id).ok_or_else(|| {
JitError::RuntimeError(format!("Output node {:?} not computed", output_id))
})?;
outputs.push(tensor.clone());
}
Ok(outputs)
}
}
fn interpreter_execute(context: &mut ExecutionContext, kernel: &CompiledKernel) -> JitResult<()> {
if kernel.source_nodes.is_empty() {
let missing_outputs: Vec<_> = context
.output_ids
.iter()
.filter(|&&id| !context.intermediates.contains_key(&id))
.copied()
.collect();
for &output_id in &missing_outputs {
let input_data = if let Some(input_tensor) = context.intermediates.values().next() {
input_tensor.data.clone()
} else {
vec![1.0; 10] };
let output_data: Vec<f32> = input_data
.iter()
.map(|&x| if x > 0.0 { x } else { 0.0 })
.collect();
let output_tensor = crate::TensorRef { data: output_data };
context.set_tensor(output_id, output_tensor);
}
} else {
for &node_id in &kernel.source_nodes {
let input_data = if let Some(input_tensor) = context.intermediates.values().next() {
input_tensor.data.clone()
} else {
vec![1.0; 10] };
let output_data: Vec<f32> = input_data
.iter()
.map(|&x| if x > 0.0 { x } else { 0.0 })
.collect();
let output_tensor = crate::TensorRef { data: output_data };
context.set_tensor(node_id, output_tensor);
}
}
Ok(())
}
pub struct MemoryPool {
pools: HashMap<usize, Vec<Vec<u8>>>,
}
impl MemoryPool {
pub fn new() -> Self {
Self {
pools: HashMap::new(),
}
}
pub fn allocate(&mut self, size: usize) -> Vec<u8> {
let pool_size = size.next_power_of_two();
if let Some(pool) = self.pools.get_mut(&pool_size) {
if let Some(mut buffer) = pool.pop() {
buffer.resize(size, 0);
return buffer;
}
}
vec![0u8; size]
}
pub fn release(&mut self, mut buffer: Vec<u8>) {
let pool_size = buffer.capacity().next_power_of_two();
buffer.clear();
self.pools.entry(pool_size).or_default().push(buffer);
}
}
impl Default for MemoryPool {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::{ComputationGraph, Node};
#[test]
fn test_kernel_cache() {
let mut cache = KernelCache::new();
cache.max_size = 2;
let fn1: ExecutableFn = Box::new(|_, _| Ok(()));
cache.insert("kernel1".to_string(), fn1);
assert!(cache.get("kernel1").is_some());
assert!(cache.get("kernel2").is_none());
assert_eq!(cache.hits, 1);
assert_eq!(cache.misses, 1);
assert_eq!(cache.hit_rate(), 0.5);
}
#[test]
fn test_memory_pool() {
let mut pool = MemoryPool::new();
let buf1 = pool.allocate(100);
assert_eq!(buf1.len(), 100);
pool.release(buf1);
let buf2 = pool.allocate(100);
assert_eq!(buf2.len(), 100);
}
#[test]
fn test_execution_context() {
let mut graph = ComputationGraph::new();
let input_node = graph.add_node(
Node::new(crate::graph::Operation::Input, "input".to_string())
.with_output_shapes(vec![Some(crate::graph::shape_from_slice(&[10]))])
.with_dtypes(vec![torsh_core::DType::F32])
.with_device(torsh_core::DeviceType::Cpu),
);
graph.add_input(input_node);
let inputs = vec![crate::TensorRef {
data: vec![1.0; 10],
}];
let context = ExecutionContext::new(&graph, &inputs);
assert!(context.is_ok());
}
}