use crate::{AnnIR, AnnOpRef, AnnValRef, Annotation, Dialect, IR, OpRef, ValId, ValMap, ValRef};
use std::{cell::RefCell, marker::PhantomData, rc::Rc};
use zhc_utils::{
iter::{CollectInSmallVec, MultiZip},
small::SmallVec,
};
pub struct EagerTranslator<ID: Dialect, OD: Dialect> {
output: IR<OD>,
valmap: ValMap<ValId>,
phantom: PhantomData<ID>,
}
impl<ID: Dialect, OD: Dialect> EagerTranslator<ID, OD> {
pub fn translate_val<'a>(&self, old: ValId) -> 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 register_translation(&mut self, old: ValId, new: ValId) {
assert!(
self.valmap.insert(old, new).is_none(),
"Tried to register a translation twice for {old}"
);
}
}
pub fn eager_translate<'a, ID: Dialect, OD: Dialect>(
ir: &'a IR<ID>,
driver: impl Fn(OpRef<'a, ID>, &mut EagerTranslator<ID, OD>),
) -> IR<OD> {
let output = IR::empty();
let valmap = ir.empty_valmap();
let mut translator = EagerTranslator {
output,
valmap,
phantom: PhantomData,
};
for op in ir.walk_ops_linear() {
driver(op, &mut translator)
}
translator.output
}
pub struct AnnEagerTranslator<ID: Dialect, OpAnn: Annotation, ValAnn: Annotation, OD: Dialect> {
output: IR<OD>,
valmap: ValMap<ValId>,
phantom: PhantomData<(ID, OpAnn, ValAnn)>,
}
impl<ID: Dialect, OpAnn: Annotation, ValAnn: Annotation, OD: Dialect>
AnnEagerTranslator<ID, OpAnn, ValAnn, OD>
{
pub fn translate_val(&self, old: ValId) -> ValId {
self.valmap.get(&old).unwrap().clone()
}
pub fn direct_translation<'a, 'b>(
&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 add_op(&mut self, instr: OD::InstructionSet, args: SmallVec<ValId>) -> SmallVec<ValId> {
self.output.add_op(instr, args).1
}
pub fn register_translation(&mut self, old: ValId, new: ValId) {
assert!(
self.valmap.insert(old, new).is_none(),
"Tried to register a translation twice for {old}"
);
}
}
pub fn eager_translate_ann<
'a,
'b,
ID: Dialect,
OpAnn: Annotation,
ValAnn: Annotation,
OD: Dialect,
>(
ir: &'b AnnIR<'a, ID, OpAnn, ValAnn>,
driver: impl Fn(AnnOpRef<'a, 'b, ID, OpAnn, ValAnn>, &mut AnnEagerTranslator<ID, OpAnn, ValAnn, OD>),
) -> IR<OD> {
let output = IR::empty();
let valmap = ir.empty_valmap();
let mut translator = AnnEagerTranslator {
output,
valmap,
phantom: PhantomData,
};
for op in ir.walk_ops_linear() {
driver(op, &mut translator)
}
translator.output
}
pub struct LazyTranslator<'a, ID: Dialect, OD: Dialect> {
driver: Rc<dyn Fn(OpRef<'a, ID>, &LazyTranslator<'a, ID, OD>) + 'a>,
output: Rc<RefCell<IR<OD>>>,
valmap: Rc<RefCell<ValMap<ValId>>>,
}
impl<'a, ID: Dialect, OD: Dialect> LazyTranslator<'a, ID, OD> {
pub fn translate_val(&self, valref: ValRef<'a, ID>) -> ValId {
if !self.valmap.borrow().contains_key(&*valref) {
(self.driver)(valref.get_origin().opref, self);
}
self.valmap.borrow().get(&*valref).unwrap().clone()
}
fn ignite(&self, opref: OpRef<'a, ID>) {
assert!(
opref.is_effect(),
"Tried to ignite translation on a non-effect op."
);
for arg in opref.get_args_iter() {
self.translate_val(arg);
}
(self.driver)(opref, self)
}
pub fn add_op(&self, instr: OD::InstructionSet, args: SmallVec<ValId>) -> SmallVec<ValId> {
self.output.borrow_mut().add_op(instr, args).1
}
pub fn register_translation(&self, old: ValId, new: ValId) {
assert!(
self.valmap.borrow_mut().insert(old, new).is_none(),
"Tried to register a translation twice for {old}"
);
}
}
pub fn lazy_translate<'a, ID: Dialect, OD: Dialect>(
ir: &'a IR<ID>,
driver: impl Fn(OpRef<'a, ID>, &LazyTranslator<'a, ID, OD>) + 'a,
) -> IR<OD> {
let output = Rc::new(RefCell::new(IR::empty()));
let valmap = Rc::new(RefCell::new(ir.empty_valmap()));
let driver = Rc::new(driver);
let translator = LazyTranslator {
driver,
output,
valmap,
};
for effect in ir.walk_ops_linear().filter(|op| op.is_effect()) {
translator.ignite(effect)
}
RefCell::into_inner(Rc::try_unwrap(translator.output).unwrap())
}
pub struct AnnLazyTranslator<
'a,
'b,
ID: Dialect,
OpAnn: Annotation,
ValAnn: Annotation,
OD: Dialect,
> {
driver: Rc<
dyn Fn(
AnnOpRef<'a, 'b, ID, OpAnn, ValAnn>,
&AnnLazyTranslator<'a, 'b, ID, OpAnn, ValAnn, OD>,
) + 'a,
>,
output: Rc<RefCell<IR<OD>>>,
valmap: Rc<RefCell<ValMap<ValId>>>,
}
impl<'a, 'b, ID: Dialect, OpAnn: Annotation, ValAnn: Annotation, OD: Dialect>
AnnLazyTranslator<'a, 'b, ID, OpAnn, ValAnn, OD>
{
pub fn translate_val(&self, valref: AnnValRef<'a, 'b, ID, OpAnn, ValAnn>) -> ValId {
if !self.valmap.borrow().contains_key(&*valref) {
(self.driver)(valref.get_origin().opref, self);
}
self.valmap.borrow().get(&*valref).unwrap().clone()
}
fn ignite(&self, opref: AnnOpRef<'a, 'b, ID, OpAnn, ValAnn>) {
assert!(
opref.is_effect(),
"Tried to ignite translation on a non-effect op."
);
for arg in opref.get_args_iter() {
self.translate_val(arg);
}
(self.driver)(opref, self)
}
pub fn push_new_op(&self, instr: OD::InstructionSet, args: SmallVec<ValId>) -> SmallVec<ValId> {
self.output.borrow_mut().add_op(instr, args).1
}
pub fn register_translation(&self, old: ValId, new: ValId) {
assert!(
self.valmap.borrow_mut().insert(old, new).is_none(),
"Tried to register a translation twice for {old}"
);
}
}
pub fn lazy_translate_ann<
'a,
'b,
ID: Dialect,
OpAnn: Annotation,
ValAnn: Annotation,
OD: Dialect,
>(
ir: &'b AnnIR<'a, ID, OpAnn, ValAnn>,
driver: impl Fn(
AnnOpRef<'a, 'b, ID, OpAnn, ValAnn>,
&AnnLazyTranslator<'a, 'b, ID, OpAnn, ValAnn, OD>,
) + 'a,
) -> IR<OD> {
let output = Rc::new(RefCell::new(IR::empty()));
let valmap = Rc::new(RefCell::new(ir.empty_valmap()));
let driver = Rc::new(driver);
let translator = AnnLazyTranslator {
driver,
output,
valmap,
};
for effect in ir.walk_ops_linear().filter(|op| op.is_effect()) {
translator.ignite(effect)
}
RefCell::into_inner(Rc::try_unwrap(translator.output).unwrap())
}