use crate::Builtin;
use crate::draws;
use crate::ir::*;
use crate::liveness::SlotSet;
use probl_syntax::Span;
#[derive(Clone, Copy, Debug)]
pub struct Update<'a> {
pub slot: SlotId,
pub likelihood: Likelihood<'a>,
}
#[derive(Clone, Copy, Debug)]
pub enum Likelihood<'a> {
Binomial { value: &'a Expr, trials: &'a Expr },
Bernoulli { value: &'a Expr },
Poisson { value: &'a Expr },
Normal { value: &'a Expr, sd: &'a Expr },
}
impl<'a> Update<'a> {
pub fn others(&self) -> Vec<&'a Expr> {
match self.likelihood {
Likelihood::Binomial { value, trials } => vec![value, trials],
Likelihood::Bernoulli { value } => vec![value],
Likelihood::Poisson { value } => vec![value],
Likelihood::Normal { value, sd } => vec![value, sd],
}
}
}
pub fn update<'a>(value: &'a Expr, from: Option<&'a Expr>) -> Option<Update<'a>> {
let builtin = |e: &'a Expr| match &e.kind {
ExprKind::Builtin { func, args, named } if named.is_empty() => Some((*func, args.as_slice())),
_ => None,
};
let slot = |e: &Expr| match e.kind {
ExprKind::Slot(s) => Some(s),
_ => None,
};
let probability_slot = |e: &'a Expr| match builtin(e)? {
(Builtin::Prob, [x]) => slot(x),
_ => None,
};
let parameter_slot = |e: &'a Expr| slot(e).or_else(|| probability_slot(e));
let mut distribution = builtin(from?)?;
if let (Builtin::BooleanLaw, [inner]) = distribution {
if let Some((Builtin::Bernoulli, args)) = builtin(inner) {
distribution = (Builtin::Bernoulli, args);
}
}
let (slot, likelihood) = match distribution {
(Builtin::Binomial, [trials, x]) => (parameter_slot(x)?, Likelihood::Binomial { value, trials }),
(Builtin::Bernoulli | Builtin::ScoreLaw, [x]) => (parameter_slot(x)?, Likelihood::Bernoulli { value }),
(Builtin::BooleanLaw, [x]) => (probability_slot(x)?, Likelihood::Bernoulli { value }),
(Builtin::Poisson, [x]) => (slot(x)?, Likelihood::Poisson { value }),
(Builtin::Normal, [x, sd]) => (slot(x)?, Likelihood::Normal { value, sd }),
_ => return None,
};
let update = Update { slot, likelihood };
let mut elsewhere = false;
for e in update.others() {
draws::expr(e, &mut |s| elsewhere |= s == slot);
}
(!elsewhere).then_some(update)
}
#[derive(Debug, Default)]
pub struct Conjugacy {
pub variables: Vec<Variable>,
pub delays: Vec<Option<u32>>,
pub draws_first: Vec<Vec<SlotId>>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct Variable {
pub function: FnId,
pub slot: SlotId,
pub name: String,
pub span: Span,
}
pub fn analyze(program: &Program) -> Conjugacy {
let n = program.stmt_count as usize;
let mut result = Conjugacy {
variables: Vec::new(),
delays: vec![None; n],
draws_first: vec![Vec::new(); n],
};
for (f, func) in program.functions.iter().enumerate() {
let mut observed = SlotSet::with_capacity(func.n_slots());
let mut drawn = SlotSet::with_capacity(func.n_slots());
each_stmt(&func.body, &mut |s| match &s.kind {
StmtKind::Observe { value, from } => {
if let Some(u) = update(value, from.as_ref()) {
observed.insert(u.slot);
}
}
StmtKind::Draw { place, .. } if place.path.is_empty() => drawn.insert(place.slot),
_ => {}
});
let mut index = vec![None; func.n_slots()];
for slot in observed.iter().filter(|&s| drawn.contains(s)) {
index[slot as usize] = Some(result.variables.len() as u32);
let info = &func.slots[slot as usize];
result.variables.push(Variable {
function: f as FnId,
slot,
name: info.name.clone(),
span: info.span,
});
}
if index.iter().all(Option::is_none) {
continue;
}
each_stmt(&func.body, &mut |s| {
if let StmtKind::Draw { place, .. } = &s.kind {
if place.path.is_empty() {
result.delays[s.id as usize] = index[place.slot as usize];
}
}
let mut reads: Vec<SlotId> = Vec::new();
own_reads(s, &mut |slot| {
if index[slot as usize].is_some() && !reads.contains(&slot) {
reads.push(slot);
}
});
result.draws_first[s.id as usize] = reads;
});
}
result
}
fn each_stmt(b: &Block, f: &mut impl FnMut(&Stmt)) {
for s in &b.stmts {
f(s);
match &s.kind {
StmtKind::If { then, otherwise, .. } => {
each_stmt(then, f);
each_stmt(otherwise, f);
}
StmtKind::Chance { arms, otherwise, .. } => {
for (_, body) in arms {
each_stmt(body, f);
}
if let Some(body) = otherwise {
each_stmt(body, f);
}
}
StmtKind::Loop { body, .. } => each_stmt(body, f),
StmtKind::Try { body, catches } => {
each_stmt(body, f);
for c in catches {
each_stmt(&c.body, f);
}
}
_ => {}
}
}
}
fn own_reads(s: &Stmt, f: &mut impl FnMut(SlotId)) {
match &s.kind {
StmtKind::Set { place, value } => {
self::place(place, f);
draws::value_reads(value, f);
}
StmtKind::Draw { place, dist } => {
self::place(place, f);
draws::value_reads(dist, f);
}
StmtKind::Take { place, bag } => {
self::place(place, f);
f(bag.slot);
self::place(bag, f);
}
StmtKind::Call { dest, callee, args } => {
self::place(dest, f);
match callee {
Callee::Fn { capture_args, .. } => capture_args.iter().for_each(|&s| f(s)),
Callee::Value(e, _) => draws::value_reads(e, f),
}
args.iter().for_each(|a| draws::value_reads(a, f));
}
StmtKind::If { cond, .. } => draws::value_reads(cond, f),
StmtKind::Chance { arms, .. } => arms.iter().for_each(|(weight, _)| draws::value_reads(weight, f)),
StmtKind::Return(e) => draws::value_reads(e, f),
StmtKind::Observe { value, from } => match update(value, from.as_ref()) {
Some(u) => u.others().into_iter().for_each(|e| draws::value_reads(e, f)),
None => {
draws::value_reads(value, f);
if let Some(d) = from {
draws::value_reads(d, f);
}
}
},
StmtKind::Report { value, key, .. } => {
draws::value_reads(value, f);
if let Some(k) = key {
draws::value_reads(k, f);
}
}
StmtKind::Loop { .. } | StmtKind::Try { .. } | StmtKind::Break | StmtKind::Continue | StmtKind::Fail { .. } => {
}
StmtKind::Check { .. } => {}
}
}
fn place(p: &Place, f: &mut impl FnMut(SlotId)) {
if p.path.is_empty() {
return;
}
f(p.slot);
for elem in &p.path {
if let PathElem::Index(e) = elem {
draws::value_reads(e, f);
}
}
}