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

Function learn_bopq

experiments/python/product_quantize.py:343–403  ·  view source on GitHub ↗
(X_train, ncodebooks, codebook_bits=4, niters=20,
               initial_kmeans_iters=1, R_sz=16, **sink)

Source from the content-addressed store, hash-verified

341
342@_memory.cache # opq with block diagonal rotations
343def 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(

Callers

nothing calls this directly

Calls 8

learn_pqFunction · 0.85
reconstruct_X_pqFunction · 0.85
_update_centroids_opqFunction · 0.85
formatMethod · 0.80
appendMethod · 0.80
_encode_X_pqFunction · 0.70
dotMethod · 0.45

Tested by

no test coverage detected