use std::collections::{HashMap, HashSet};
use crate::graph::Graph;
use crate::node::{Attribute, Node, NodeId};
use crate::shape::Dim;
use crate::value::ValueId;
#[must_use]
pub fn inline_single_trip_scan_bodies(graph: &Graph) -> Graph {
let mut out = graph.clone();
let scan_nodes: Vec<NodeId> = out
.nodes
.iter()
.filter(|(_, node)| is_default_domain_scan(node))
.map(|(nid, _)| nid)
.collect();
for scan_id in scan_nodes {
if let Some(plan) = ScanInlinePlan::extract(&out, scan_id) {
plan.apply(&mut out);
}
}
out
}
fn is_default_domain_scan(node: &Node) -> bool {
node.is_default_domain() && node.op_type == "Scan"
}
struct ScanInlinePlan {
scan_id: NodeId,
body: Graph,
state_inputs: Vec<ValueId>,
scan_inputs: Vec<ValueId>,
state_outputs: Vec<ValueId>,
scan_outputs: Vec<ValueId>,
input_axes: Vec<usize>,
output_axes: Vec<usize>,
}
impl ScanInlinePlan {
fn extract(graph: &Graph, scan_id: NodeId) -> Option<Self> {
let node = graph.try_node(scan_id)?;
let body = graph.subgraphs.get(&(scan_id, "body".to_string()))?.clone();
let num_scan_inputs = usize::try_from(node.attr("num_scan_inputs")?.as_int()?).ok()?;
if num_scan_inputs == 0 || node.inputs.len() < num_scan_inputs {
return None;
}
let num_state = node.inputs.len() - num_scan_inputs;
if num_state == 0 || node.outputs.len() < num_state {
return None;
}
let num_scan_outputs = node.outputs.len() - num_state;
if body.inputs.len() != num_state + num_scan_inputs
|| body.outputs.len() != num_state + num_scan_outputs
{
return None;
}
if body_has_nested_control_flow(&body) {
return None;
}
let state_inputs = collect_present(&node.inputs[..num_state])?;
let scan_inputs = collect_present(&node.inputs[num_state..])?;
let state_outputs = node.outputs[..num_state].to_vec();
let scan_outputs = node.outputs[num_state..].to_vec();
let input_axes_raw = int_list_attr(node, "scan_input_axes", num_scan_inputs)?;
let output_axes_raw = int_list_attr(node, "scan_output_axes", num_scan_outputs)?;
let mut input_axes = Vec::with_capacity(num_scan_inputs);
for (scan_input, &raw_axis) in scan_inputs.iter().zip(&input_axes_raw) {
let rank = graph.value(*scan_input).shape.len();
let axis = normalize_axis(raw_axis, rank)?;
if let Some(Dim::Static(extent)) = graph.value(*scan_input).shape.get(axis)
&& *extent > 1
{
return None;
}
input_axes.push(axis);
}
let mut output_axes = Vec::with_capacity(num_scan_outputs);
for (scan_output, &raw_axis) in scan_outputs.iter().zip(&output_axes_raw) {
let out_rank = graph.value(*scan_output).shape.len().max(1);
output_axes.push(normalize_axis(raw_axis, out_rank)?);
}
let name_index = parent_name_index(graph);
for name in body_capture_names(&body) {
if !name_index.contains_key(&name) {
return None;
}
}
Some(Self {
scan_id,
body,
state_inputs,
scan_inputs,
state_outputs,
scan_outputs,
input_axes,
output_axes,
})
}
fn apply(self, out: &mut Graph) {
let num_state = self.state_inputs.len();
let name_index = parent_name_index(out);
let formal_set: HashSet<ValueId> = self.body.inputs.iter().copied().collect();
let mut remap: HashMap<ValueId, ValueId> = HashMap::new();
for (i, &parent_state_in) in self.state_inputs.iter().enumerate() {
remap.insert(self.body.inputs[i], parent_state_in);
}
for (j, (&parent_scan_in, &axis)) in
self.scan_inputs.iter().zip(&self.input_axes).enumerate()
{
let squeezed = self.emit_squeeze(out, parent_scan_in, axis);
remap.insert(self.body.inputs[num_state + j], squeezed);
}
for (vid, value) in self.body.values.iter() {
let is_capture = value.producer.is_none()
&& !formal_set.contains(&vid)
&& !self.body.initializers.contains_key(&vid);
if is_capture
&& let Some(name) = &value.name
&& let Some(&parent) = name_index.get(name)
{
remap.insert(vid, parent);
}
}
for (&vid, weight) in &self.body.initializers {
let value = self.body.value(vid);
let promoted = out.create_value(value.dtype, value.shape.clone());
if let Some(name) = &value.name {
out.value_mut(promoted).name = Some(name.clone());
}
out.set_initializer(promoted, weight.clone());
remap.insert(vid, promoted);
}
let mut direct_state_out: HashSet<ValueId> = HashSet::new();
let mut identity_state_out: Vec<(ValueId, ValueId)> = Vec::new();
for (k, &parent_state_out) in self.state_outputs.iter().enumerate() {
let body_out = self.body.outputs[k];
let producible = self
.body
.try_value(body_out)
.is_some_and(|v| v.producer.is_some());
if producible && !remap.contains_key(&body_out) && direct_state_out.insert(body_out) {
remap.insert(body_out, parent_state_out);
} else {
identity_state_out.push((body_out, parent_state_out));
}
}
for (vid, value) in self.body.values.iter() {
if remap.contains_key(&vid) || value.producer.is_none() {
continue;
}
let fresh = out.create_value(value.dtype, value.shape.clone());
remap.insert(vid, fresh);
}
let order = self
.body
.topological_order()
.expect("scan body must be acyclic");
for nid in order {
let bn = self.body.node(nid);
let inputs = bn
.inputs
.iter()
.map(|slot| slot.map(|v| remap[&v]))
.collect();
let outputs = bn.outputs.iter().map(|v| remap[v]).collect();
let mut nn = Node::new(NodeId(0), bn.op_type.clone(), inputs, outputs);
nn.name = bn.name.clone();
nn.domain = bn.domain.clone();
nn.version = bn.version;
nn.attributes = bn.attributes.clone();
nn.doc_string = bn.doc_string.clone();
out.insert_node(nn);
}
for (body_out, parent_state_out) in identity_state_out {
let src = remap[&body_out];
out.insert_node(Node::new(
NodeId(0),
"Identity",
vec![Some(src)],
vec![parent_state_out],
));
}
for (m, (&parent_scan_out, &axis)) in
self.scan_outputs.iter().zip(&self.output_axes).enumerate()
{
let src = remap[&self.body.outputs[num_state + m]];
self.emit_unsqueeze(out, src, parent_scan_out, axis);
}
out.remove_node(self.scan_id);
out.subgraphs.retain(|(owner, _), _| *owner != self.scan_id);
}
fn emit_squeeze(&self, out: &mut Graph, parent_scan_in: ValueId, axis: usize) -> ValueId {
let (dtype, mut shape) = {
let v = out.value(parent_scan_in);
(v.dtype, v.shape.clone())
};
if axis < shape.len() {
shape.remove(axis);
}
let squeezed = out.create_value(dtype, shape);
out.insert_node(squeeze_like("Squeeze", parent_scan_in, squeezed, axis));
squeezed
}
fn emit_unsqueeze(&self, out: &mut Graph, src: ValueId, parent_scan_out: ValueId, axis: usize) {
out.insert_node(squeeze_like("Unsqueeze", src, parent_scan_out, axis));
}
}
fn squeeze_like(op: &str, input: ValueId, output: ValueId, axis: usize) -> Node {
let mut node = Node::new(NodeId(0), op, vec![Some(input)], vec![output]);
node.version = Some(1);
node.attributes
.insert("axes".to_string(), Attribute::Ints(vec![axis as i64]));
node
}
fn body_has_nested_control_flow(body: &Graph) -> bool {
if !body.subgraphs.is_empty() {
return true;
}
body.nodes.iter().any(|(_, node)| {
node.attributes
.values()
.any(|attr| matches!(attr, Attribute::Graph(_) | Attribute::Graphs(_)))
})
}
fn body_capture_names(body: &Graph) -> Vec<String> {
let formal_set: HashSet<ValueId> = body.inputs.iter().copied().collect();
body.values
.iter()
.filter_map(|(vid, value)| {
(value.producer.is_none()
&& !formal_set.contains(&vid)
&& !body.initializers.contains_key(&vid))
.then(|| value.name.clone())
.flatten()
})
.collect()
}
fn parent_name_index(graph: &Graph) -> HashMap<String, ValueId> {
let mut index = HashMap::new();
for (vid, value) in graph.values.iter() {
if let Some(name) = &value.name {
index.entry(name.clone()).or_insert(vid);
}
}
index
}
fn collect_present(slots: &[Option<ValueId>]) -> Option<Vec<ValueId>> {
slots.iter().copied().collect()
}
fn int_list_attr(node: &Node, name: &str, expected: usize) -> Option<Vec<i64>> {
match node.attr(name) {
None => Some(vec![0; expected]),
Some(attr) => {
let values = attr.as_ints()?;
(values.len() == expected).then(|| values.to_vec())
}
}
}
fn normalize_axis(axis: i64, rank: usize) -> Option<usize> {
let rank_i = i64::try_from(rank).ok()?;
let resolved = if axis < 0 { axis + rank_i } else { axis };
(0..rank_i).contains(&resolved).then_some(resolved as usize)
}
#[cfg(test)]
mod tests;