MCPcopy Create free account
hub / github.com/dblalock/bolt / test_basic

Function test_basic

tests/test_encoder.py:169–273  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

167
168
169def test_basic():
170 # np.set_printoptions(precision=3)
171 np.set_printoptions(formatter={'float_kind': _fmt_float})
172
173 nqueries = 20
174 X, Q = _load_digits_X_Q(nqueries)
175 # X, _ = load_digits(return_X_y=True)
176 # Q = X[-nqueries:]
177 # X = X[:-nqueries]
178
179 # print "X.shape", X.shape
180 # print "X nbytes", X.nbytes
181
182 # ------------------------------------------------ squared l2
183
184 enc = bolt.Encoder(accuracy='low', reduction=bolt.Reductions.SQUARED_EUCLIDEAN)
185 enc.fit(X)
186
187 l2_corrs = np.empty(nqueries)
188 for i, q in enumerate(Q):
189 l2_true = _dists_sq(X, q).astype(np.int)
190 l2_bolt = enc.transform(q)
191 l2_corrs[i] = corr(l2_true, l2_bolt)[0]
192
193 mean_l2 = np.mean(l2_corrs)
194 std_l2 = np.std(l2_corrs)
195 assert mean_l2 > .95
196 print "squared l2 dist correlation: {} +/- {}".format(mean_l2, std_l2)
197
198 # ------------------------------------------------ dot product
199
200 enc = bolt.Encoder(accuracy='low', reduction=bolt.Reductions.DOT_PRODUCT)
201 enc.fit(X)
202
203 dot_corrs = np.empty(nqueries)
204 for i, q in enumerate(Q):
205 dots_true = np.dot(X, q)
206 dots_bolt = enc.transform(q)
207 dot_corrs[i] = corr(dots_true, dots_bolt)[0]
208
209 mean_dot = np.mean(dot_corrs)
210 std_dot = np.std(dot_corrs)
211 print "dot product correlation: {} +/- {}".format(mean_dot, std_dot)
212
213 # ------------------------------------------------ l2 knn
214
215 enc = bolt.Encoder(accuracy='low', reduction='l2')
216 enc.fit(X)
217
218 k_bolt = 10 # tell bolt to search for true knn
219 k_true = 10 # compute this many true neighbors
220 true_knn = _knn(X, Q, k_true)
221 bolt_knn = [enc.knn(q, k_bolt) for q in Q]
222
223 contained = np.empty((nqueries, k_bolt), dtype=np.bool)
224 for i in range(nqueries):
225 true_neighbors = true_knn[i]
226 bolt_neighbors = bolt_knn[i]

Callers 1

test_encoder.pyFile · 0.85

Calls 11

fitMethod · 0.95
transformMethod · 0.95
knnMethod · 0.95
_load_digits_X_QFunction · 0.85
_dists_sqFunction · 0.85
_knnFunction · 0.85
formatMethod · 0.80
top_k_idxsFunction · 0.70
emptyMethod · 0.45
meanMethod · 0.45
dotMethod · 0.45

Tested by

no test coverage detected