()
| 167 | |
| 168 | |
| 169 | def 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] |
no test coverage detected