use std::collections::HashMap;
use shape_vm::bytecode::{BuiltinFunction, BytecodeProgram, OpCode, Operand};
use crate::loop_analysis::LoopInfo;
#[derive(Debug, Clone)]
pub struct HoistableCall {
pub call_idx: usize,
pub arg_count: usize,
pub first_arg_idx: usize,
}
#[derive(Debug, Clone, Default)]
pub struct LicmPlan {
pub hoistable_calls_by_loop: HashMap<usize, Vec<HoistableCall>>,
}
fn is_pure_builtin(builtin: &BuiltinFunction) -> bool {
matches!(
builtin,
BuiltinFunction::Sin
| BuiltinFunction::Cos
| BuiltinFunction::Tan
| BuiltinFunction::Asin
| BuiltinFunction::Acos
| BuiltinFunction::Atan
| BuiltinFunction::Sqrt
| BuiltinFunction::Abs
| BuiltinFunction::Floor
| BuiltinFunction::Ceil
| BuiltinFunction::Round
| BuiltinFunction::Exp
| BuiltinFunction::Ln
| BuiltinFunction::Log
| BuiltinFunction::Pow
| BuiltinFunction::Sign
| BuiltinFunction::Hypot
)
}
fn is_pure_method_name(name: &str) -> bool {
matches!(name, "row" | "col" | "transpose" | "shape" | "len")
}
fn is_invariant_value_producer(
instr_idx: usize,
program: &BytecodeProgram,
info: &LoopInfo,
) -> bool {
let instr = &program.instructions[instr_idx];
match instr.opcode {
OpCode::PushConst | OpCode::PushNull => true,
OpCode::LoadLocal | OpCode::LoadLocalTrusted => {
if let Some(Operand::Local(slot)) = &instr.operand {
info.invariant_locals.contains(slot)
} else {
false
}
}
OpCode::LoadModuleBinding => {
if let Some(Operand::ModuleBinding(slot)) = &instr.operand {
info.invariant_module_bindings.contains(slot)
} else {
false
}
}
_ => false,
}
}
fn analyze_loop_calls(
program: &BytecodeProgram,
info: &LoopInfo,
) -> Vec<HoistableCall> {
let mut hoistable = Vec::new();
let mut nested_depth = 0usize;
let mut i = info.header_idx + 1;
while i < info.end_idx {
let instr = &program.instructions[i];
match instr.opcode {
OpCode::LoopStart => {
nested_depth += 1;
i += 1;
continue;
}
OpCode::LoopEnd if nested_depth > 0 => {
nested_depth -= 1;
i += 1;
continue;
}
_ => {}
}
if nested_depth > 0 {
i += 1;
continue;
}
if instr.opcode == OpCode::BuiltinCall {
if let Some(Operand::Builtin(builtin)) = &instr.operand {
if is_pure_builtin(builtin) {
if let Some(call) =
try_hoist_builtin_call(program, info, i)
{
hoistable.push(call);
}
}
}
}
if instr.opcode == OpCode::CallMethod {
match &instr.operand {
Some(Operand::TypedMethodCall { string_id, arg_count: _, .. }) => {
let str_idx = *string_id as usize;
if let Some(method_name) = program.strings.get(str_idx) {
if is_pure_method_name(method_name) {
if let Some(call) =
try_hoist_method_call(program, info, i)
{
hoistable.push(call);
}
}
}
}
Some(Operand::TypedMethodCall {
string_id,
..
}) => {
let str_idx = *string_id as usize;
if let Some(method_name) = program.strings.get(str_idx) {
if is_pure_method_name(method_name) {
if let Some(call) =
try_hoist_method_call(program, info, i)
{
hoistable.push(call);
}
}
}
}
_ => {}
}
}
i += 1;
}
hoistable
}
fn try_hoist_builtin_call(
program: &BytecodeProgram,
info: &LoopInfo,
call_idx: usize,
) -> Option<HoistableCall> {
if call_idx == 0 {
return None;
}
let argc_instr = &program.instructions[call_idx - 1];
if argc_instr.opcode != OpCode::PushConst {
return None;
}
let arg_count = read_const_int(program, &argc_instr.operand)?;
if arg_count > 8 {
return None; }
let first_arg_idx = (call_idx - 1).checked_sub(arg_count)?;
if first_arg_idx <= info.header_idx {
return None; }
for j in first_arg_idx..(call_idx - 1) {
if !is_invariant_value_producer(j, program, info) {
return None;
}
}
Some(HoistableCall {
call_idx,
arg_count,
first_arg_idx,
})
}
fn try_hoist_method_call(
program: &BytecodeProgram,
info: &LoopInfo,
call_idx: usize,
) -> Option<HoistableCall> {
let operand_arg_count = match &program.instructions[call_idx].operand {
Some(Operand::TypedMethodCall { arg_count, .. }) => *arg_count as usize,
Some(Operand::TypedMethodCall { arg_count, .. }) => *arg_count as usize,
_ => return None,
};
if operand_arg_count > 8 {
return None; }
if call_idx == 0 {
return None;
}
let argc_instr = &program.instructions[call_idx - 1];
if argc_instr.opcode != OpCode::PushConst {
return None;
}
let total_pushes = 1 + operand_arg_count;
let first_arg_idx = (call_idx - 1).checked_sub(total_pushes)?;
if first_arg_idx <= info.header_idx {
return None;
}
for j in first_arg_idx..(call_idx - 1) {
if !is_invariant_value_producer(j, program, info) {
return None;
}
}
Some(HoistableCall {
call_idx,
arg_count: total_pushes, first_arg_idx,
})
}
fn read_const_int(
program: &BytecodeProgram,
operand: &Option<Operand>,
) -> Option<usize> {
let Some(Operand::Const(const_idx)) = operand else {
return None;
};
match program.constants.get(*const_idx as usize) {
Some(shape_vm::bytecode::Constant::Int(v)) => {
if *v >= 0 {
Some(*v as usize)
} else {
None
}
}
Some(shape_vm::bytecode::Constant::UInt(v)) => Some(*v as usize),
Some(shape_vm::bytecode::Constant::Number(v)) if *v >= 0.0 && *v == (*v as usize) as f64 => {
Some(*v as usize)
}
_ => None,
}
}
pub fn analyze_licm(
program: &BytecodeProgram,
loop_info: &HashMap<usize, LoopInfo>,
) -> LicmPlan {
let mut plan = LicmPlan::default();
for (header, info) in loop_info {
let calls = analyze_loop_calls(program, info);
if !calls.is_empty() {
plan.hoistable_calls_by_loop.insert(*header, calls);
}
}
plan
}
#[cfg(test)]
mod tests {
use super::*;
use shape_vm::bytecode::*;
fn make_instr(opcode: OpCode, operand: Option<Operand>) -> Instruction {
Instruction { opcode, operand }
}
fn make_program(instrs: Vec<Instruction>, constants: Vec<Constant>) -> BytecodeProgram {
BytecodeProgram {
instructions: instrs,
constants,
strings: vec![],
functions: vec![],
debug_info: DebugInfo::default(),
data_schema: None,
module_binding_names: vec![],
top_level_locals_count: 0,
top_level_local_storage_hints: vec![],
type_schema_registry: Default::default(),
module_binding_storage_hints: vec![],
function_local_storage_hints: vec![],
compiled_annotations: Default::default(),
trait_method_symbols: Default::default(),
expanded_function_defs: Default::default(),
string_index: Default::default(),
foreign_functions: vec![],
native_struct_layouts: vec![],
content_addressed: None,
function_blob_hashes: vec![],
top_level_frame: None,
..Default::default()
}
}
fn make_program_with_strings(
instrs: Vec<Instruction>,
constants: Vec<Constant>,
strings: Vec<String>,
) -> BytecodeProgram {
let mut p = make_program(instrs, constants);
p.strings = strings;
p
}
#[test]
fn test_pure_builtin_hoistable_single_arg() {
let instrs = vec![
make_instr(OpCode::LoopStart, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), make_instr(OpCode::LtInt, None), make_instr(OpCode::JumpIfFalse, Some(Operand::Offset(8))), make_instr(OpCode::LoadLocal, Some(Operand::Local(2))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::BuiltinCall, Some(Operand::Builtin(BuiltinFunction::Sin))), make_instr(OpCode::StoreLocal, Some(Operand::Local(3))), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::AddInt, None), make_instr(OpCode::StoreLocal, Some(Operand::Local(0))), make_instr(OpCode::LoopEnd, None), ];
let program = make_program(instrs, vec![Constant::Int(1)]);
let loop_info = crate::loop_analysis::analyze_loops(&program);
let plan = analyze_licm(&program, &loop_info);
assert!(plan.hoistable_calls_by_loop.contains_key(&0));
let calls = &plan.hoistable_calls_by_loop[&0];
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].call_idx, 7);
assert_eq!(calls[0].arg_count, 1);
assert_eq!(calls[0].first_arg_idx, 5);
}
#[test]
fn test_non_invariant_arg_not_hoisted() {
let instrs = vec![
make_instr(OpCode::LoopStart, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), make_instr(OpCode::LtInt, None), make_instr(OpCode::JumpIfFalse, Some(Operand::Offset(8))), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::BuiltinCall, Some(Operand::Builtin(BuiltinFunction::Sin))), make_instr(OpCode::StoreLocal, Some(Operand::Local(3))), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::AddInt, None), make_instr(OpCode::StoreLocal, Some(Operand::Local(0))), make_instr(OpCode::LoopEnd, None), ];
let program = make_program(instrs, vec![Constant::Int(1)]);
let loop_info = crate::loop_analysis::analyze_loops(&program);
let plan = analyze_licm(&program, &loop_info);
assert!(
plan.hoistable_calls_by_loop.get(&0).map_or(true, |c| c.is_empty()),
"sin(i) should not be hoisted when i is the IV"
);
}
#[test]
fn test_impure_builtin_not_hoisted() {
let instrs = vec![
make_instr(OpCode::LoopStart, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), make_instr(OpCode::LtInt, None), make_instr(OpCode::JumpIfFalse, Some(Operand::Offset(8))), make_instr(OpCode::LoadLocal, Some(Operand::Local(2))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::BuiltinCall, Some(Operand::Builtin(BuiltinFunction::Print))), make_instr(OpCode::StoreLocal, Some(Operand::Local(3))), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::AddInt, None), make_instr(OpCode::StoreLocal, Some(Operand::Local(0))), make_instr(OpCode::LoopEnd, None), ];
let program = make_program(instrs, vec![Constant::Int(1)]);
let loop_info = crate::loop_analysis::analyze_loops(&program);
let plan = analyze_licm(&program, &loop_info);
assert!(
plan.hoistable_calls_by_loop.get(&0).map_or(true, |c| c.is_empty()),
"print() should not be hoisted (impure)"
);
}
#[test]
fn test_pure_method_call_hoistable() {
use shape_value::StringId;
let instrs = vec![
make_instr(OpCode::LoopStart, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), make_instr(OpCode::LtInt, None), make_instr(OpCode::JumpIfFalse, Some(Operand::Offset(8))), make_instr(OpCode::LoadLocal, Some(Operand::Local(2))), make_instr(OpCode::PushConst, Some(Operand::Const(1))), make_instr(
OpCode::CallMethod,
Some(Operand::TypedMethodCall {
method_id: 0,
arg_count: 0,
string_id: 0,
receiver_type_tag: 0xFF,
}),
), make_instr(OpCode::StoreLocal, Some(Operand::Local(3))), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::AddInt, None), make_instr(OpCode::StoreLocal, Some(Operand::Local(0))), make_instr(OpCode::LoopEnd, None), ];
let program = make_program_with_strings(
instrs,
vec![Constant::Int(1), Constant::Int(0)],
vec!["shape".to_string()],
);
let loop_info = crate::loop_analysis::analyze_loops(&program);
let plan = analyze_licm(&program, &loop_info);
assert!(plan.hoistable_calls_by_loop.contains_key(&0));
let calls = &plan.hoistable_calls_by_loop[&0];
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].call_idx, 7);
}
#[test]
fn test_nested_loop_ignores_inner() {
let instrs = vec![
make_instr(OpCode::LoopStart, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), make_instr(OpCode::LtInt, None), make_instr(OpCode::JumpIfFalse, Some(Operand::Offset(15))), make_instr(OpCode::LoadLocal, Some(Operand::Local(2))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::BuiltinCall, Some(Operand::Builtin(BuiltinFunction::Sin))), make_instr(OpCode::StoreLocal, Some(Operand::Local(3))), make_instr(OpCode::LoopStart, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(2))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::BuiltinCall, Some(Operand::Builtin(BuiltinFunction::Cos))), make_instr(OpCode::Pop, None), make_instr(OpCode::LoopEnd, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::AddInt, None), make_instr(OpCode::StoreLocal, Some(Operand::Local(0))), make_instr(OpCode::LoopEnd, None), ];
let program = make_program(instrs, vec![Constant::Int(1)]);
let loop_info = crate::loop_analysis::analyze_loops(&program);
let plan = analyze_licm(&program, &loop_info);
if let Some(calls) = plan.hoistable_calls_by_loop.get(&0) {
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].call_idx, 7, "should be the sin() call in outer body");
}
}
#[test]
fn test_constant_arg_hoistable() {
let instrs = vec![
make_instr(OpCode::LoopStart, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::LoadLocal, Some(Operand::Local(1))), make_instr(OpCode::LtInt, None), make_instr(OpCode::JumpIfFalse, Some(Operand::Offset(8))), make_instr(OpCode::PushConst, Some(Operand::Const(1))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::BuiltinCall, Some(Operand::Builtin(BuiltinFunction::Sin))), make_instr(OpCode::Pop, None), make_instr(OpCode::LoadLocal, Some(Operand::Local(0))), make_instr(OpCode::PushConst, Some(Operand::Const(0))), make_instr(OpCode::AddInt, None), make_instr(OpCode::StoreLocal, Some(Operand::Local(0))), make_instr(OpCode::LoopEnd, None), ];
let program = make_program(
instrs,
vec![Constant::Int(1), Constant::Number(3.14)],
);
let loop_info = crate::loop_analysis::analyze_loops(&program);
let plan = analyze_licm(&program, &loop_info);
assert!(plan.hoistable_calls_by_loop.contains_key(&0));
let calls = &plan.hoistable_calls_by_loop[&0];
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].call_idx, 7);
}
}