use crate::{
graph::{ComputationGraph, ConstantValue, Operation},
JitError, JitResult, NodeId,
};
use petgraph::visit::EdgeRef;
use std::collections::HashMap;
use std::fmt::Write;
#[derive(Debug, Clone)]
pub struct LlvmBackend {
module: String,
symbol_table: HashMap<NodeId, String>,
next_id: usize,
module_name: String,
target_triple: String,
data_layout: String,
}
impl Default for LlvmBackend {
fn default() -> Self {
Self::new("torsh_module")
}
}
impl LlvmBackend {
pub fn new(module_name: &str) -> Self {
Self {
module: String::new(),
symbol_table: HashMap::new(),
next_id: 0,
module_name: module_name.to_string(),
target_triple: "x86_64-unknown-linux-gnu".to_string(),
data_layout: "e-m:e-p270:32:32-p271:32:32-p272:64:64-i64:64-f80:128-n8:16:32:64-S128"
.to_string(),
}
}
pub fn set_target_triple(&mut self, triple: String) {
self.target_triple = triple;
}
pub fn set_data_layout(&mut self, layout: String) {
self.data_layout = layout;
}
pub fn generate(&mut self, graph: &ComputationGraph) -> JitResult<String> {
self.reset();
self.generate_module_header()?;
self.generate_global_declarations()?;
self.generate_main_function(graph)?;
self.generate_helper_functions()?;
Ok(self.module.clone())
}
fn reset(&mut self) {
self.module.clear();
self.symbol_table.clear();
self.next_id = 0;
}
fn generate_module_header(&mut self) -> JitResult<()> {
writeln!(self.module, "; Generated LLVM IR module from ToRSh JIT")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "; Module: {}", self.module_name)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "").map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "target triple = \"{}\"", self.target_triple)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "target datalayout = \"{}\"", self.data_layout)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "").map_err(|e| JitError::CodeGenError(e.to_string()))?;
Ok(())
}
fn generate_global_declarations(&mut self) -> JitResult<()> {
writeln!(self.module, "; External function declarations")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "declare float @expf(float)")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "declare float @logf(float)")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "declare float @tanhf(float)")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "declare float @sqrtf(float)")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
"declare <4 x float> @llvm.exp.v4f32(<4 x float>)"
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
"declare <4 x float> @llvm.log.v4f32(<4 x float>)"
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
"declare <4 x float> @llvm.sqrt.v4f32(<4 x float>)"
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "declare i8* @malloc(i64)")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "declare void @free(i8*)")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
"declare void @llvm.memcpy.p0i8.p0i8.i64(i8*, i8*, i64, i1)"
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "").map_err(|e| JitError::CodeGenError(e.to_string()))?;
Ok(())
}
fn generate_main_function(&mut self, graph: &ComputationGraph) -> JitResult<()> {
let input_count = self.count_input_nodes(graph);
let output_count = self.count_output_nodes(graph);
write!(self.module, "define ").map_err(|e| JitError::CodeGenError(e.to_string()))?;
if output_count == 1 {
write!(self.module, "float* ").map_err(|e| JitError::CodeGenError(e.to_string()))?;
} else {
write!(self.module, "{{").map_err(|e| JitError::CodeGenError(e.to_string()))?;
for i in 0..output_count {
if i > 0 {
write!(self.module, ", ").map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
write!(self.module, "float*").map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
write!(self.module, "}} ").map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
write!(self.module, "@main(").map_err(|e| JitError::CodeGenError(e.to_string()))?;
for i in 0..input_count {
if i > 0 {
write!(self.module, ", ").map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
write!(self.module, "float* %input{}, i64 %size{}", i, i)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
writeln!(self.module, ") {{").map_err(|e| JitError::CodeGenError(e.to_string()))?;
self.generate_function_body(graph)?;
writeln!(self.module, "}}").map_err(|e| JitError::CodeGenError(e.to_string()))?;
Ok(())
}
fn generate_function_body(&mut self, graph: &ComputationGraph) -> JitResult<()> {
writeln!(self.module, "entry:").map_err(|e| JitError::CodeGenError(e.to_string()))?;
self.generate_llvm_constants(graph)?;
let topo_order = self.topological_sort(graph)?;
for node_id in topo_order {
if let Some(node) = graph.node(node_id) {
if !matches!(node.op, Operation::Constant(_)) {
self.generate_llvm_operation(graph, node_id, node)?;
}
}
}
self.generate_return_statement(graph)?;
Ok(())
}
fn generate_llvm_constants(&mut self, graph: &ComputationGraph) -> JitResult<()> {
for (node_id, node) in graph.nodes() {
if let Operation::Constant(ref const_info) = node.op {
let var_name = self.get_or_create_symbol(node_id);
match &const_info.value {
ConstantValue::Scalar(val) => {
writeln!(self.module, " %{}_ptr = alloca float, align 4", var_name)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store float {}, float* %{}_ptr, align 4",
val, var_name
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
ConstantValue::IntScalar(val) => {
writeln!(self.module, " %{}_ptr = alloca i64, align 8", var_name)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store i64 {}, i64* %{}_ptr, align 8",
val, var_name
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
ConstantValue::Tensor {
shape: _,
data,
dtype: _,
} => {
let size = data.len();
writeln!(self.module, " %{}_size = add i64 0, {}", var_name, size)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_bytes = mul i64 %{}_size, 4",
var_name, var_name
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_raw = call i8* @malloc(i64 %{}_bytes)",
var_name, var_name
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_ptr = bitcast i8* %{}_raw to float*",
var_name, var_name
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
for (i, &value) in data.iter().enumerate() {
writeln!(
self.module,
" %{}_elem{}_ptr = getelementptr inbounds float, float* %{}_ptr, i64 {}",
var_name, i, var_name, i
).map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store float {}, float* %{}_elem{}_ptr, align 4",
value, var_name, i
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
}
ConstantValue::Bool(val) => {
writeln!(self.module, " %{}_ptr = alloca i1, align 1", var_name)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store i1 {}, i1* %{}_ptr, align 1",
if *val { "true" } else { "false" },
var_name
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
ConstantValue::Int(val) => {
writeln!(self.module, " %{}_ptr = alloca i64, align 8", var_name)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store i64 {}, i64* %{}_ptr, align 8",
val, var_name
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
ConstantValue::Float(val) => {
writeln!(self.module, " %{}_ptr = alloca double, align 8", var_name)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store double {}, double* %{}_ptr, align 8",
val, var_name
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
_ => {
writeln!(self.module, " ; Unhandled constant type for {}", var_name)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
}
}
}
Ok(())
}
fn generate_llvm_operation(
&mut self,
graph: &ComputationGraph,
node_id: NodeId,
node: &crate::Node,
) -> JitResult<()> {
let output_var = self.get_or_create_symbol(node_id);
let inputs = self.get_input_symbols(graph, node_id);
match &node.op {
Operation::Add => {
if inputs.len() >= 2 {
self.generate_elementwise_binary_op(&output_var, &inputs, "fadd")?;
}
}
Operation::Sub => {
if inputs.len() >= 2 {
self.generate_elementwise_binary_op(&output_var, &inputs, "fsub")?;
}
}
Operation::Mul => {
if inputs.len() >= 2 {
self.generate_elementwise_binary_op(&output_var, &inputs, "fmul")?;
}
}
Operation::Div => {
if inputs.len() >= 2 {
self.generate_elementwise_binary_op(&output_var, &inputs, "fdiv")?;
}
}
Operation::MatMul => {
if inputs.len() >= 2 {
self.generate_matmul(&output_var, &inputs)?;
}
}
Operation::Relu => {
if !inputs.is_empty() {
self.generate_relu(&output_var, &inputs[0])?;
}
}
Operation::Sigmoid => {
if !inputs.is_empty() {
self.generate_sigmoid(&output_var, &inputs[0])?;
}
}
Operation::Tanh => {
if !inputs.is_empty() {
self.generate_tanh(&output_var, &inputs[0])?;
}
}
Operation::Exp => {
if !inputs.is_empty() {
self.generate_unary_math_op(&output_var, &inputs[0], "exp")?;
}
}
Operation::Log => {
if !inputs.is_empty() {
self.generate_unary_math_op(&output_var, &inputs[0], "log")?;
}
}
Operation::Neg => {
if !inputs.is_empty() {
self.generate_neg(&output_var, &inputs[0])?;
}
}
_ => {
writeln!(self.module, " ; Unsupported operation: {:?}", node.op)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
if !inputs.is_empty() {
writeln!(
self.module,
" %{}_ptr = alloca float*, align 8",
output_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store float* %{}_ptr, float** %{}_ptr",
inputs[0], output_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
}
}
Ok(())
}
fn generate_elementwise_binary_op(
&mut self,
output_var: &str,
inputs: &[String],
op: &str,
) -> JitResult<()> {
let size_var = format!("{}_size", output_var);
let loop_var = format!("{}_loop", output_var);
let cond_var = format!("{}_cond", output_var);
let next_var = format!("{}_next", output_var);
writeln!(
self.module,
" %{} = add i64 0, 1024 ; placeholder size",
size_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_bytes = mul i64 %{}, 4",
output_var, size_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_raw = call i8* @malloc(i64 %{}_bytes)",
output_var, output_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_ptr = bitcast i8* %{}_raw to float*",
output_var, output_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, " br label %{}_head", loop_var)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "{}:", format!("{}_head", loop_var))
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_i = phi i64 [ 0, %entry ], [ %{}, %{}_body ]",
loop_var, next_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{} = icmp ult i64 %{}_i, %{}",
cond_var, loop_var, size_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" br i1 %{}, label %{}_body, label %{}_end",
cond_var, loop_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "{}:", format!("{}_body", loop_var))
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_a_ptr = getelementptr inbounds float, float* %{}_ptr, i64 %{}_i",
loop_var, inputs[0], loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_a = load float, float* %{}_a_ptr, align 4",
loop_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_b_ptr = getelementptr inbounds float, float* %{}_ptr, i64 %{}_i",
loop_var, inputs[1], loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_b = load float, float* %{}_b_ptr, align 4",
loop_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_result = {} float %{}_a, %{}_b",
loop_var, op, loop_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_out_ptr = getelementptr inbounds float, float* %{}_ptr, i64 %{}_i",
loop_var, output_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store float %{}_result, float* %{}_out_ptr, align 4",
loop_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, " %{} = add i64 %{}_i, 1", next_var, loop_var)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, " br label %{}_head", loop_var)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "{}:", format!("{}_end", loop_var))
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
Ok(())
}
fn generate_matmul(&mut self, output_var: &str, inputs: &[String]) -> JitResult<()> {
writeln!(
self.module,
" ; Matrix multiplication: {} = {} * {}",
output_var, inputs[0], inputs[1]
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_ptr = call float* @matmul_impl(float* %{}_ptr, float* %{}_ptr)",
output_var, inputs[0], inputs[1]
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
Ok(())
}
fn generate_relu(&mut self, output_var: &str, input_var: &str) -> JitResult<()> {
self.generate_elementwise_unary_op(output_var, input_var, |module, i, input, output| {
writeln!(module, " %{}_zero = add float 0.0, 0.0", i)?;
writeln!(
module,
" %{}_cmp = fcmp ogt float %{}, %{}_zero",
i, input, i
)?;
writeln!(
module,
" %{} = select i1 %{}_cmp, float %{}, float %{}_zero",
output, i, input, i
)?;
Ok(())
})
}
fn generate_sigmoid(&mut self, output_var: &str, input_var: &str) -> JitResult<()> {
self.generate_elementwise_unary_op(output_var, input_var, |module, i, input, output| {
writeln!(module, " %{}_neg = fsub float 0.0, %{}", i, input)?;
writeln!(module, " %{}_exp = call float @expf(float %{}_neg)", i, i)?;
writeln!(module, " %{}_one = add float 1.0, 0.0", i)?;
writeln!(module, " %{}_sum = fadd float %{}_one, %{}_exp", i, i, i)?;
writeln!(module, " %{} = fdiv float %{}_one, %{}_sum", output, i, i)?;
Ok(())
})
}
fn generate_tanh(&mut self, output_var: &str, input_var: &str) -> JitResult<()> {
self.generate_elementwise_unary_op(output_var, input_var, |module, i, input, output| {
writeln!(
module,
" %{} = call float @tanhf(float %{})",
output, input
)?;
Ok(())
})
}
fn generate_unary_math_op(
&mut self,
output_var: &str,
input_var: &str,
op: &str,
) -> JitResult<()> {
self.generate_elementwise_unary_op(output_var, input_var, |module, i, input, output| {
writeln!(
module,
" %{} = call float @{}f(float %{})",
output, op, input
)?;
Ok(())
})
}
fn generate_neg(&mut self, output_var: &str, input_var: &str) -> JitResult<()> {
self.generate_elementwise_unary_op(output_var, input_var, |module, i, input, output| {
writeln!(module, " %{} = fsub float 0.0, %{}", output, input)?;
Ok(())
})
}
fn generate_elementwise_unary_op<F>(
&mut self,
output_var: &str,
input_var: &str,
body_gen: F,
) -> JitResult<()>
where
F: Fn(&mut String, &str, &str, &str) -> Result<(), std::fmt::Error>,
{
let size_var = format!("{}_size", output_var);
let loop_var = format!("{}_loop", output_var);
let cond_var = format!("{}_cond", output_var);
let next_var = format!("{}_next", output_var);
writeln!(
self.module,
" %{} = add i64 0, 1024 ; placeholder size",
size_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_bytes = mul i64 %{}, 4",
output_var, size_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_raw = call i8* @malloc(i64 %{}_bytes)",
output_var, output_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_ptr = bitcast i8* %{}_raw to float*",
output_var, output_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, " br label %{}_head", loop_var)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "{}:", format!("{}_head", loop_var))
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_i = phi i64 [ 0, %entry ], [ %{}, %{}_body ]",
loop_var, next_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{} = icmp ult i64 %{}_i, %{}",
cond_var, loop_var, size_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" br i1 %{}, label %{}_body, label %{}_end",
cond_var, loop_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "{}:", format!("{}_body", loop_var))
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_in_ptr = getelementptr inbounds float, float* %{}_ptr, i64 %{}_i",
loop_var, input_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_in = load float, float* %{}_in_ptr, align 4",
loop_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
body_gen(
&mut self.module,
&loop_var,
&format!("{}_in", loop_var),
&format!("{}_result", loop_var),
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" %{}_out_ptr = getelementptr inbounds float, float* %{}_ptr, i64 %{}_i",
loop_var, output_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store float %{}_result, float* %{}_out_ptr, align 4",
loop_var, loop_var
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, " %{} = add i64 %{}_i, 1", next_var, loop_var)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, " br label %{}_head", loop_var)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "{}:", format!("{}_end", loop_var))
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
Ok(())
}
fn generate_return_statement(&mut self, graph: &ComputationGraph) -> JitResult<()> {
let mut output_vars = Vec::new();
for (node_id, _) in graph.nodes() {
let has_outgoing = graph
.edges_directed(node_id, petgraph::Direction::Outgoing)
.next()
.is_some();
if !has_outgoing {
if let Some(var_name) = self.symbol_table.get(&node_id) {
output_vars.push(format!("{}_ptr", var_name));
}
}
}
if output_vars.is_empty() {
writeln!(self.module, " ret void")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
} else if output_vars.len() == 1 {
writeln!(self.module, " ret float* %{}", output_vars[0])
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
} else {
writeln!(
self.module,
" %result = alloca {{{}}}*, align 8",
output_vars
.iter()
.map(|_| "float*")
.collect::<Vec<_>>()
.join(", ")
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
for (i, var) in output_vars.iter().enumerate() {
writeln!(
self.module,
" %result_ptr{} = getelementptr inbounds {{{}}}, {{{}}}* %result, i32 0, i32 {}",
i,
output_vars.iter().map(|_| "float*").collect::<Vec<_>>().join(", "),
output_vars.iter().map(|_| "float*").collect::<Vec<_>>().join(", "),
i
).map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" store float* %{}, float** %result_ptr{}, align 8",
var, i
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
writeln!(
self.module,
" %result_val = load {{{}}}, {{{}}}* %result, align 8",
output_vars
.iter()
.map(|_| "float*")
.collect::<Vec<_>>()
.join(", "),
output_vars
.iter()
.map(|_| "float*")
.collect::<Vec<_>>()
.join(", ")
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(
self.module,
" ret {{{}}} %result_val",
output_vars
.iter()
.map(|_| "float*")
.collect::<Vec<_>>()
.join(", ")
)
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
}
Ok(())
}
fn generate_helper_functions(&mut self) -> JitResult<()> {
writeln!(self.module, "").map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "; Helper function declarations")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
writeln!(self.module, "declare float* @matmul_impl(float*, float*)")
.map_err(|e| JitError::CodeGenError(e.to_string()))?;
Ok(())
}
fn get_or_create_symbol(&mut self, node_id: NodeId) -> String {
if let Some(symbol) = self.symbol_table.get(&node_id) {
symbol.clone()
} else {
let symbol = format!("v{}", self.next_id);
self.next_id += 1;
self.symbol_table.insert(node_id, symbol.clone());
symbol
}
}
fn get_input_symbols(&mut self, graph: &ComputationGraph, node_id: NodeId) -> Vec<String> {
let mut inputs = Vec::new();
for edge in graph.edges_directed(node_id, petgraph::Direction::Incoming) {
let src_id = edge.source();
let symbol = self.get_or_create_symbol(src_id);
inputs.push(symbol);
}
inputs
}
fn topological_sort(&self, graph: &ComputationGraph) -> JitResult<Vec<NodeId>> {
use petgraph::algo::toposort;
toposort(&graph.graph, None)
.map_err(|_| JitError::GraphError("Cyclic graph detected".to_string()))
}
fn count_input_nodes(&self, graph: &ComputationGraph) -> usize {
graph
.nodes()
.filter(|(node_id, _)| {
graph
.edges_directed(*node_id, petgraph::Direction::Incoming)
.next()
.is_none()
})
.count()
}
fn count_output_nodes(&self, graph: &ComputationGraph) -> usize {
graph
.nodes()
.filter(|(node_id, _)| {
graph
.edges_directed(*node_id, petgraph::Direction::Outgoing)
.next()
.is_none()
})
.count()
}
}
#[derive(Debug, Clone)]
pub struct LlvmOptimizer {
opt_level: u8,
target_specific: bool,
vectorize: bool,
loop_opt: bool,
}
impl Default for LlvmOptimizer {
fn default() -> Self {
Self::new(2)
}
}
impl LlvmOptimizer {
pub fn new(opt_level: u8) -> Self {
Self {
opt_level,
target_specific: opt_level >= 2,
vectorize: opt_level >= 2,
loop_opt: opt_level >= 1,
}
}
pub fn enable_target_specific(&mut self, enabled: bool) {
self.target_specific = enabled;
}
pub fn enable_vectorization(&mut self, enabled: bool) {
self.vectorize = enabled;
}
pub fn enable_loop_optimization(&mut self, enabled: bool) {
self.loop_opt = enabled;
}
pub fn optimize(&self, llvm_ir: &str) -> JitResult<String> {
let mut optimized = llvm_ir.to_string();
if self.opt_level >= 1 {
optimized = self.apply_basic_optimizations(optimized)?;
}
if self.opt_level >= 2 {
optimized = self.apply_advanced_optimizations(optimized)?;
}
if self.opt_level >= 3 {
optimized = self.apply_aggressive_optimizations(optimized)?;
}
Ok(optimized)
}
fn apply_basic_optimizations(&self, ir: String) -> JitResult<String> {
let mut result = ir;
result = result.replace("store float %unused,", "; removed dead store:");
result = result.replace("fadd float %x, 0.0", "; simplified: %x");
result = result.replace("fmul float %x, 1.0", "; simplified: %x");
result = result.replace("fmul float %x, 0.0", "; simplified: 0.0");
Ok(result)
}
fn apply_advanced_optimizations(&self, ir: String) -> JitResult<String> {
let mut result = ir;
if self.vectorize {
result = result.replace(
"for.body:",
"for.body: ; vectorizable loop\n !llvm.loop !{!llvm.loop.vectorize.enable, i1 true}"
);
}
if self.loop_opt {
result = result.replace(
"for.inc:",
"for.inc: ; unroll loop\n !llvm.loop !{!llvm.loop.unroll.count, i32 4}",
);
}
Ok(result)
}
fn apply_aggressive_optimizations(&self, ir: String) -> JitResult<String> {
let mut result = ir;
result = result.replace("define ", "define alwaysinline ");
result = result.replace("fadd float", "fadd fast float");
result = result.replace("fsub float", "fsub fast float");
result = result.replace("fmul float", "fmul fast float");
result = result.replace("fdiv float", "fdiv fast float");
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::{ComputationGraph, ConstantInfo, Operation};
use torsh_core::{DType, DeviceType, Shape};
#[test]
fn test_llvm_backend_creation() {
let backend = LlvmBackend::new("test_module");
assert_eq!(backend.module_name, "test_module");
assert!(backend.symbol_table.is_empty());
}
#[test]
fn test_llvm_optimizer() {
let optimizer = LlvmOptimizer::new(2);
assert_eq!(optimizer.opt_level, 2);
assert!(optimizer.target_specific);
assert!(optimizer.vectorize);
}
#[test]
fn test_simple_llvm_generation() {
let mut backend = LlvmBackend::new("test");
let mut graph = ComputationGraph::new();
let node = crate::Node::new(
Operation::Constant(ConstantInfo {
value: ConstantValue::Scalar(42.0),
}),
"const1".to_string(),
)
.with_output_shapes(vec![Some(Shape::new(vec![1]))])
.with_dtypes(vec![DType::F32])
.with_device(DeviceType::Cpu);
graph.add_node(node);
let result = backend.generate(&graph);
assert!(result.is_ok());
let llvm_ir = result.unwrap();
assert!(llvm_ir.contains("target triple"));
assert!(llvm_ir.contains("define"));
assert!(llvm_ir.contains("42"));
}
}