-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathinferencedriver.py
More file actions
116 lines (101 loc) · 3.8 KB
/
Copy pathinferencedriver.py
File metadata and controls
116 lines (101 loc) · 3.8 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
import numpy as np
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from timeit import default_timer as timer
from probpy import ProbPy
import copy
class InferenceDriver:
def __init__(self, model):
self.pp = ProbPy()
self.model = model
self.samples = []
self.lls = []
def init_model(self):
# prime the database
self.model(self.pp)
self.pp.accept_proposed_trace()
def burn_in(self, steps):
start = timer()
for i in range(steps):
self.inference_step()
print("Burn in: %.2fs" % (timer() - start))
def run_inference(self, interval, samples):
self.num_samples = samples
total_start = timer()
for s in range(samples):
if s%10 == 0:
print("sample %d" % (s))
for i in range(interval):
self.inference_step()
self.samples.append(copy.deepcopy(self.pp.table.trace))
print("Total inference: %.2fs" % (timer() - total_start))
def inference_step(self):
# score the current trace
ll = self.pp.score_current_trace()
self.lls.append(ll)
# start new trace (copy of the old trace)
self.pp.propose_new_trace()
# pick a random ERP
label, entry = self.pp.pick_random_erp()
# propose a new value
if entry["erp"] == "choice":
value, F, R = self.pp.choice_proposal_kernal(entry["value"],
entry["parameters"]["elements"], entry["parameters"]["p"])
else:
value, F, R = self.pp.simple_proposal_kernal(entry["value"])
# value, F, R = self.pp.sample_erp(entry["erp"], entry["parameters"]) # sample kernal
self.pp.store_new_erp(label, value, entry["erp"], entry["parameters"])
# re-run the model
self.model(self.pp)
# score the new trace
ll_prime, ll_fresh, ll_stale = self.pp.score_proposed_trace()
# calculate MH acceptance ratio
threshold = ll_prime - ll + R - F + ll_stale - ll_fresh
# accept or reject
if np.log(np.random.rand()) < threshold:
self.pp.accept_proposed_trace()
def condition(self, label, value):
self.pp.table.condition(label, value)
def prior(self, label, value):
self.pp.table.prior(label, value)
def return_traces(self):
return self.samples
def return_values(self, keys):
values = {}
val_cnt = {}
for s in self.samples:
for key, item in s.items():
if key.split("-")[0] in keys:
if key in values:
#values[key] = float(values[key] + item["value"]) / 2.0
values[key] = float(values[key] + item["value"])
val_cnt[key] += 1
else:
values[key] = item["value"]
val_cnt[key] = 1.0
#return values
return {k: v / val_cnt[k] for k, v in values.items()}
def return_string_values(self, keys):
values = {}
for s in self.samples:
for key, item in s.items():
if key.split("-")[0] in keys:
if key in values:
values[key].append(item["value"])
else:
values[key] = [item["value"]]
return values
def return_plt_data(self, keys):
data = {}
for k in keys:
data[k] = []
for s in self.samples:
values_dict = {}
for key, item in s.items():
if k + '-0' == key:
data[k].append(item['value'])
return data
def graph_ll(self):
plt.plot(range(len(self.lls)), self.lls)
plt.savefig("ll_figure.png")