use crate::{AnnIR, AnnOpRef, Annotation, AsValId, Dialect, IR, OpId, OpRef, ValId, ValMap};
use std::marker::PhantomData;
use zhc_utils::{
iter::{CollectInSmallVec, MultiZip},
small::SmallVec,
};
pub enum Order {
Linear,
Topological,
Custom(Vec<OpId>),
}
pub struct Translator<ID: Dialect, OD: Dialect> {
output: IR<OD>,
valmap: ValMap<ValId>,
phantom: PhantomData<ID>,
}
impl<ID: Dialect, OD: Dialect> Translator<ID, OD> {
pub fn translate_val(&self, old: impl AsValId) -> ValId {
self.valmap.get(old).unwrap().clone()
}
pub fn add_op(&mut self, instr: OD::InstructionSet, args: SmallVec<ValId>) -> SmallVec<ValId> {
self.output.add_op(instr, args).1
}
pub fn has_translation(&self, old: impl AsValId) -> bool {
self.valmap.contains_key(old)
}
pub fn register_translation(&mut self, old: impl AsValId, new: impl AsValId) {
let old = old.val_id();
let new = new.val_id();
assert!(
self.valmap.insert(old, new).is_none(),
"Tried to register a translation twice for {old}"
);
}
pub fn direct_translation<'a, 'b, OpAnn: Annotation, ValAnn: Annotation>(
&mut self,
op: AnnOpRef<'a, 'b, ID, OpAnn, ValAnn>,
instr: OD::InstructionSet,
) {
let new_args = op
.get_arg_valids()
.iter()
.map(|v| self.translate_val(*v))
.cosvec();
let new_rets = self.add_op(instr, new_args);
assert_eq!(new_rets.len(), op.get_return_arity());
(new_rets.into_iter(), op.get_return_valids().iter())
.mzip()
.for_each(|(new, old)| self.register_translation(*old, new));
}
}
pub fn translate<'a, ID: Dialect, OD: Dialect>(
ir: &'a IR<ID>,
order: Order,
driver: impl Fn(OpRef<'a, ID>, &mut Translator<ID, OD>),
) -> IR<OD> {
let output = IR::empty();
let valmap = ir.empty_valmap();
let mut translator = Translator {
output,
valmap,
phantom: PhantomData,
};
match order {
Order::Linear => {
for op in ir.walk_ops_linear() {
driver(op, &mut translator);
}
}
Order::Topological => {
for op in ir.walk_ops_topological() {
driver(op, &mut translator);
}
}
Order::Custom(ids) => {
for op in ir.walk_ops_with(ids.into_iter()) {
driver(op, &mut translator);
}
}
}
translator.output
}
pub fn translate_ann<'a, 'b, ID: Dialect, OpAnn: Annotation, ValAnn: Annotation, OD: Dialect>(
ir: &'b AnnIR<'a, ID, OpAnn, ValAnn>,
order: Order,
driver: impl Fn(AnnOpRef<'a, 'b, ID, OpAnn, ValAnn>, &mut Translator<ID, OD>),
) -> IR<OD> {
let output = IR::empty();
let valmap = ir.empty_valmap();
let mut translator = Translator {
output,
valmap,
phantom: PhantomData,
};
match order {
Order::Linear => {
for op in ir.walk_ops_linear() {
driver(op, &mut translator);
}
}
Order::Topological => {
for op in ir.walk_ops_topological() {
driver(op, &mut translator);
}
}
Order::Custom(ids) => {
for op in ir.walk_ops_with(ids.into_iter()) {
driver(op, &mut translator);
}
}
}
translator.output
}