From 77d8b323d4f6e05ca97d9cbef43ac85fd8040d61 Mon Sep 17 00:00:00 2001 From: Cathy Yeh Date: Mon, 13 Nov 2017 14:42:52 -0800 Subject: copy scripts from lgs branch --- beliefs/utils/random_variables.py | 21 +++++++++++++++++++++ 1 file changed, 21 insertions(+) create mode 100644 beliefs/utils/random_variables.py (limited to 'beliefs/utils/random_variables.py') diff --git a/beliefs/utils/random_variables.py b/beliefs/utils/random_variables.py new file mode 100644 index 0000000..1a0b0f7 --- /dev/null +++ b/beliefs/utils/random_variables.py @@ -0,0 +1,21 @@ + + +def get_reachable_observed_variables_for_inferred_variables(model, observed=set()): + """ + After performing inference on a BayesianModel, get the labels of observed variables + ("reachable observed variables") that influenced the beliefs of variables inferred + to be in a definite state. + + INPUT + model: instance of BayesianModel class or subclass + observed: set of labels (strings) corresponding to vars pinned to definite + state during inference. + RETURNS + dict, of form key - source label (a string), value - a list of strings + """ + if not observed: + return {} + + source_vars = model.get_unobserved_variables_in_definite_state(observed) + + return {var: model.reachable_observed_variables(var, observed) for var in source_vars} -- cgit v1.2.3