(X_train, ncodebooks, codebook_bits=4, niters=20,
initial_kmeans_iters=1, R_sz=16, **sink)
| 341 | |
| 342 | @_memory.cache # opq with block diagonal rotations |
| 343 | def learn_bopq(X_train, ncodebooks, codebook_bits=4, niters=20, |
| 344 | initial_kmeans_iters=1, R_sz=16, **sink): |
| 345 | |
| 346 | t0 = time.time() |
| 347 | |
| 348 | X = X_train.astype(np.float32) |
| 349 | N, D = X.shape |
| 350 | ncentroids = int(2**codebook_bits) |
| 351 | subvect_len = D // ncodebooks |
| 352 | |
| 353 | assert D % subvect_len == 0 # equal number of dims for each codebook |
| 354 | |
| 355 | # compute number of rotations and subspaces associated with each |
| 356 | nrots = int(D / R_sz) |
| 357 | rot_starts = R_sz * np.arange(nrots) |
| 358 | rot_ends = rot_starts + R_sz |
| 359 | |
| 360 | # X_rotated, R = opq_initialize(X_train, ncodebooks=ncodebooks, init=init) |
| 361 | X_rotated = X # hardcode identity init # TODO allow others |
| 362 | rotations = [np.eye(R_sz) for i in range(nrots)] |
| 363 | |
| 364 | # initialize codebooks by running kmeans on each rotated dim; this way, |
| 365 | # setting niters=0 corresponds to normal PQ |
| 366 | codebooks, assignments = learn_pq(X_rotated, ncentroids=ncentroids, |
| 367 | nsubvects=ncodebooks, |
| 368 | subvect_len=subvect_len, |
| 369 | max_kmeans_iters=1) |
| 370 | |
| 371 | for it in np.arange(niters): |
| 372 | # compute reconstruction errors |
| 373 | X_hat = reconstruct_X_pq(assignments, codebooks) |
| 374 | # err = compute_reconstruction_error(X_rotated, X_hat, subvect_len=subvect_len) |
| 375 | err = compute_reconstruction_error(X_rotated, X_hat) |
| 376 | print "---- BOPQ {} {}x{}b iter {}: mse / variance = {:.5f}".format( |
| 377 | R_sz, ncodebooks, codebook_bits, it, err) |
| 378 | |
| 379 | rotations = [] |
| 380 | for i in range(nrots): |
| 381 | start, end = rot_starts[i], rot_ends[i] |
| 382 | |
| 383 | X_sub = X[:, start:end] |
| 384 | X_hat_sub = X_hat[:, start:end] |
| 385 | |
| 386 | # update rotation matrix based on reconstruction errors |
| 387 | U, s, V = np.linalg.svd(np.dot(X_hat_sub.T, X_sub), full_matrices=False) |
| 388 | R = np.dot(U, V) |
| 389 | rotations.append(R) |
| 390 | |
| 391 | X_rotated[:, start:end] = np.dot(X_sub, R.T) |
| 392 | |
| 393 | # update assignments and codebooks based on new rotations |
| 394 | assignments = _encode_X_pq(X_rotated, codebooks) |
| 395 | codebooks = _update_centroids_opq(X_rotated, assignments, ncentroids) |
| 396 | |
| 397 | X_hat = reconstruct_X_pq(assignments, codebooks) |
| 398 | err = compute_reconstruction_error(X_rotated, X_hat) |
| 399 | t = time.time() - t0 |
| 400 | print "---- BOPQ {} {}x{}b final mse / variance = {:.5f} ({:.3f}s)".format( |
nothing calls this directly
no test coverage detected