use rlx_ir::{Graph, NodeId, Op};
use std::collections::{HashMap, HashSet};
use std::ops::Range;
use crate::mps_graph_lower::{MpsGraphPlan, try_lower_with_constants};
pub enum HybridStep {
SubGraph {
plan: MpsGraphPlan,
boundary_parent_ids: HashMap<String, NodeId>,
output_parent_ids: Vec<(NodeId, NodeId)>,
thunk_skip: Range<usize>,
},
Thunks(Range<usize>),
}
pub struct ExtractedSubgraph {
pub graph: Graph,
pub boundaries: HashMap<String, NodeId>,
pub output_parent_ids: Vec<(NodeId, NodeId)>,
}
pub fn extract_subgraph(full: &Graph, segment_nodes: &[NodeId]) -> ExtractedSubgraph {
let seg_set: HashSet<NodeId> = segment_nodes.iter().copied().collect();
let mut boundary_parent: HashMap<String, NodeId> = HashMap::new();
for &nid in segment_nodes {
for &inp in &full.node(nid).inputs {
if !seg_set.contains(&inp) {
boundary_parent
.entry(format!("__boundary_{}", inp.0))
.or_insert(inp);
}
}
}
let mut sub = Graph::new(format!("{}_hybrid", full.name));
let mut map: HashMap<NodeId, NodeId> = HashMap::new();
let mut boundary_names: Vec<String> = boundary_parent.keys().cloned().collect();
boundary_names.sort();
for name in &boundary_names {
let parent_id = boundary_parent[name];
let bn = full.node(parent_id);
let new_id = match &bn.op {
Op::Input { name: n } => sub.input(n.clone(), bn.shape.clone()),
Op::Param { name: n } => sub.param(n.clone(), bn.shape.clone()),
Op::Constant { data } => sub.add_node(
Op::Constant { data: data.clone() },
vec![],
bn.shape.clone(),
),
_ => sub.input(name.clone(), bn.shape.clone()),
};
map.insert(parent_id, new_id);
}
for &nid in segment_nodes {
if map.contains_key(&nid) {
continue;
}
let n = full.node(nid);
let new_inputs: Vec<NodeId> = n
.inputs
.iter()
.map(|&i| *map.get(&i).expect("dependency mapped"))
.collect();
let new_id = sub.add_node(n.op.clone(), new_inputs, n.shape.clone());
map.insert(nid, new_id);
}
let graph_outputs: HashSet<NodeId> = full.outputs.iter().copied().collect();
let mut outs = Vec::new();
let mut output_parent_ids = Vec::new();
for &nid in segment_nodes {
let used_outside = full.users(nid).iter().any(|u| !seg_set.contains(u));
if used_outside || graph_outputs.contains(&nid) {
let sub_out = *map.get(&nid).unwrap();
outs.push(sub_out);
output_parent_ids.push((sub_out, nid));
}
}
if outs.is_empty() {
if let Some(&last) = segment_nodes.last() {
let sub_out = *map.get(&last).unwrap();
outs.push(sub_out);
output_parent_ids.push((sub_out, last));
}
}
sub.set_outputs(outs);
ExtractedSubgraph {
graph: sub,
boundaries: boundary_parent,
output_parent_ids,
}
}
fn can_lower_dequant_in_mps(
_graph: &Graph,
_node_id: NodeId,
_params_as_constants: Option<&HashMap<String, Vec<u8>>>,
) -> bool {
false
}
pub fn build_hybrid_plan(
graph: &Graph,
params_as_constants: Option<&HashMap<String, Vec<u8>>>,
) -> Option<Vec<HybridStep>> {
let mut steps: Vec<HybridStep> = Vec::new();
let mut pending: Vec<NodeId> = Vec::new();
let mut pending_idxs: Vec<usize> = Vec::new();
let flush_mps =
|pending: &mut Vec<NodeId>, pending_idxs: &mut Vec<usize>, steps: &mut Vec<HybridStep>| {
if pending.is_empty() {
return;
}
let start = *pending_idxs.iter().min().expect("pending idx");
let end = *pending_idxs.iter().max().expect("pending idx") + 1;
let extracted = extract_subgraph(graph, pending);
match try_lower_with_constants(&extracted.graph, params_as_constants) {
Some(plan) => {
steps.push(HybridStep::SubGraph {
plan,
boundary_parent_ids: extracted.boundaries,
output_parent_ids: extracted.output_parent_ids,
thunk_skip: start..end,
});
}
None => {
steps.push(HybridStep::Thunks(start..end));
}
}
pending.clear();
pending_idxs.clear();
};
for (thunk_idx, node) in graph.nodes().iter().enumerate() {
let id = node.id;
let op = &node.op;
if matches!(
op,
Op::Input { .. } | Op::Param { .. } | Op::Constant { .. }
) {
continue;
}
if matches!(
op,
Op::GatedDeltaNet { .. } | Op::Lstm { .. } | Op::Attention { .. }
) || (matches!(op, Op::DequantMatMul { .. })
&& !can_lower_dequant_in_mps(graph, id, params_as_constants))
{
flush_mps(&mut pending, &mut pending_idxs, &mut steps);
steps.push(HybridStep::Thunks(thunk_idx..thunk_idx + 1));
} else {
pending.push(id);
pending_idxs.push(thunk_idx);
}
}
flush_mps(&mut pending, &mut pending_idxs, &mut steps);
if steps.iter().all(|s| matches!(s, HybridStep::Thunks(_))) {
return None;
}
Some(steps)
}
pub fn hybrid_has_mps(steps: &[HybridStep]) -> bool {
steps
.iter()
.any(|s| matches!(s, HybridStep::SubGraph { .. }))
}