Facebook sync (May 2019) + relicense (#838)

Changelog:

- changed license: BSD+Patents -> MIT
- propagates exceptions raised in sub-indexes of IndexShards and IndexReplicas
- support for searching several inverted lists in parallel (parallel_mode != 0)
- better support for PQ codes where nbit != 8 or 16
- IVFSpectralHash implementation: spectral hash codes inside an IVF
- 6-bit per component scalar quantizer (4 and 8 bit were already supported)
- combinations of inverted lists: HStackInvertedLists and VStackInvertedLists
- configurable number of threads for OnDiskInvertedLists prefetching (including 0=no prefetch)
- more test and demo code compatible with Python 3 (print with parentheses)
- refactored benchmark code: data loading is now in a single file
This commit is contained in:
Lucas Hosseini
2019-05-28 16:17:22 +02:00
committed by GitHub
parent 712edb043a
commit a8118acbc5
1592 changed files with 68696 additions and 317410 deletions
+18 -6
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -469,9 +468,22 @@ void ParameterSpace::set_index_parameter (
}
if (DC (IndexShards)) {
// call on all sub-indexes
for (auto & shard_index : ix->shard_indexes) {
set_index_parameter (shard_index, name, val);
}
auto fn =
[this, name, val](int, Index* subIndex) {
set_index_parameter(subIndex, name, val);
};
ix->runOnIndex(fn);
return;
}
if (DC (IndexReplicas)) {
// call on all sub-indexes
auto fn =
[this, name, val](int, Index* subIndex) {
set_index_parameter(subIndex, name, val);
};
ix->runOnIndex(fn);
return;
}
if (DC (IndexRefineFlat)) {
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+40 -5
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -54,6 +53,10 @@ RangeSearchResult::~RangeSearchResult () {
delete [] lims;
}
/***********************************************************************
* BufferList
***********************************************************************/
@@ -148,7 +151,7 @@ void RangeSearchPartialResult::finalize ()
res->do_allocation ();
#pragma omp barrier
set_result ();
copy_result ();
}
@@ -162,7 +165,7 @@ void RangeSearchPartialResult::set_lims ()
}
/// called by range_search after do_allocation
void RangeSearchPartialResult::set_result (bool incremental)
void RangeSearchPartialResult::copy_result (bool incremental)
{
size_t ofs = 0;
for (int i = 0; i < queries.size(); i++) {
@@ -178,6 +181,38 @@ void RangeSearchPartialResult::set_result (bool incremental)
}
}
void RangeSearchPartialResult::merge (std::vector <RangeSearchPartialResult *> &
partial_results, bool do_delete)
{
int npres = partial_results.size();
if (npres == 0) return;
RangeSearchResult *result = partial_results[0]->res;
size_t nx = result->nq;
// count
for (const RangeSearchPartialResult * pres : partial_results) {
if (!pres) continue;
for (const RangeQueryResult &qres : pres->queries) {
result->lims[qres.qno] += qres.nres;
}
}
result->do_allocation ();
for (int j = 0; j < npres; j++) {
if (!partial_results[j]) continue;
partial_results[j]->copy_result (true);
if (do_delete) {
delete partial_results[j];
partial_results[j] = nullptr;
}
}
// reset the limits
for (size_t i = nx; i > 0; i--) {
result->lims [i] = result->lims [i - 1];
}
result->lims [0] = 0;
}
/***********************************************************************
* IDSelectorRange
+29 -13
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -90,10 +89,15 @@ struct IDSelectorBatch: IDSelector {
~IDSelectorBatch() override {}
};
// Below are structures used only by Index implementations
/****************************************************************
* Result structures for range search.
*
* The main constraint here is that we want to support parallel
* queries from different threads in various ways: 1 thread per query,
* several threads per query. We store the actual results in blocks of
* fixed size rather than exponentially increasing memory. At the end,
* we copy the block content to a linear result array.
*****************************************************************/
/** List of temporary buffers used to store results before they are
* copied to the RangeSearchResult object. */
@@ -115,9 +119,10 @@ struct BufferList {
~BufferList ();
// create a new buffer
/// create a new buffer
void append_buffer ();
/// add one result, possibly appending a new buffer if needed
void add (idx_t id, float dis);
/// copy elemnts ofs:ofs+n-1 seen as linear data in the buffers to
@@ -132,10 +137,11 @@ struct RangeSearchPartialResult;
/// result structure for a single query
struct RangeQueryResult {
using idx_t = Index::idx_t;
idx_t qno;
size_t nres;
idx_t qno; //< id of the query
size_t nres; //< nb of results for this query
RangeSearchPartialResult * pres;
/// called by search function to report a new result
void add (float dis, idx_t id);
};
@@ -143,20 +149,30 @@ struct RangeQueryResult {
struct RangeSearchPartialResult: BufferList {
RangeSearchResult * res;
/// eventually the result will be stored in res_in
explicit RangeSearchPartialResult (RangeSearchResult * res_in);
/// query ids + nb of results per query.
std::vector<RangeQueryResult> queries;
/// begin a new result
RangeQueryResult & new_result (idx_t qno);
/*****************************************
* functions used at the end of the search to merge the result
* lists */
void finalize ();
/// called by range_search before do_allocation
void set_lims ();
/// called by range_search after do_allocation
void set_result (bool incremental = false);
void copy_result (bool incremental = false);
/// merge a set of PartialResult's into one RangeSearchResult
/// on ouptut the partialresults are empty!
static void merge (std::vector <RangeSearchPartialResult *> &
partial_results, bool do_delete=true);
};
@@ -212,7 +228,7 @@ struct VectorIOWriter:IOWriter {
* it maintains counters) so the distance functions are not const,
* instanciate one from each thread if needed.
***********************************************************/
struct DistanceComputer {
struct DistanceComputer {
using idx_t = Index::idx_t;
/// called before computing distances
@@ -225,7 +241,7 @@ struct VectorIOWriter:IOWriter {
virtual float symmetric_dis (idx_t i, idx_t j) = 0;
virtual ~DistanceComputer() {}
};
};
/***********************************************************
* Interrupt callback
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+34 -3
View File
@@ -1,14 +1,14 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
// -*- c++ -*-
#include "FaissException.h"
#include <sstream>
namespace faiss {
@@ -32,4 +32,35 @@ FaissException::what() const noexcept {
return msg.c_str();
}
void handleExceptions(
std::vector<std::pair<int, std::exception_ptr>>& exceptions) {
if (exceptions.size() == 1) {
// throw the single received exception directly
std::rethrow_exception(exceptions.front().second);
} else if (exceptions.size() > 1) {
// multiple exceptions; aggregate them and return a single exception
std::stringstream ss;
for (auto& p : exceptions) {
try {
std::rethrow_exception(p.second);
} catch (std::exception& ex) {
if (ex.what()) {
// exception message available
ss << "Exception thrown from index " << p.first << ": "
<< ex.what() << "\n";
} else {
// No message available
ss << "Unknown exception thrown from index " << p.first << "\n";
}
} catch (...) {
ss << "Unknown exception thrown from index " << p.first << "\n";
}
}
throw FaissException(ss.str());
}
}
}
+9 -6
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -13,6 +12,8 @@
#include <exception>
#include <string>
#include <vector>
#include <utility>
namespace faiss {
@@ -32,6 +33,11 @@ class FaissException : public std::exception {
std::string msg;
};
/// Handle multiple exceptions from worker threads, throwing an appropriate
/// exception that aggregates the information
/// The pair int is the thread that generated the exception
void
handleExceptions(std::vector<std::pair<int, std::exception_ptr>>& exceptions);
/** bare-bones unique_ptr
* this one deletes with delete [] */
@@ -60,9 +66,6 @@ struct ScopeDeleter1 {
}
};
}
#endif
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+3 -4
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -19,7 +18,7 @@
#define FAISS_VERSION_MAJOR 1
#define FAISS_VERSION_MINOR 5
#define FAISS_VERSION_PATCH 1
#define FAISS_VERSION_PATCH 2
/**
* @namespace faiss
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+10 -10
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -459,16 +458,17 @@ void search_knn_hamming_heap(const IndexBinaryIVF& ivf,
size_t list_size = ivf.invlists->list_size(key);
InvertedLists::ScopedCodes scodes (ivf.invlists, key);
const Index::idx_t * ids = store_pairs ? nullptr :
ivf.invlists->get_ids (key);
std::unique_ptr<InvertedLists::ScopedIds> sids;
const Index::idx_t * ids = nullptr;
if (!store_pairs) {
sids.reset (new InvertedLists::ScopedIds (ivf.invlists, key));
ids = sids->get();
}
nheap += scanner->scan_codes (list_size, scodes.get(),
ids, simi, idxi, k);
if (ids) {
ivf.invlists->release_ids (ids);
}
nscan += list_size;
if (max_codes && nscan >= max_codes)
break;
@@ -553,7 +553,7 @@ void search_knn_hamming_count(const IndexBinaryIVF& ivf,
csi.update_counter(yj, id);
}
if (ids)
ivf.invlists->release_ids (ids);
ivf.invlists->release_ids (key, ids);
nscan += list_size;
if (max_codes && nscan >= max_codes)
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+309 -111
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -10,7 +9,11 @@
#include "IndexIVF.h"
#include <omp.h>
#include <cstdio>
#include <memory>
#include "utils.h"
#include "hamming.h"
@@ -118,6 +121,7 @@ IndexIVF::IndexIVF (Index * quantizer, size_t d,
code_size (code_size),
nprobe (1),
max_codes (0),
parallel_mode (0),
maintain_direct_map (false)
{
FAISS_THROW_IF_NOT (d == quantizer->d);
@@ -132,7 +136,7 @@ IndexIVF::IndexIVF (Index * quantizer, size_t d,
IndexIVF::IndexIVF ():
invlists (nullptr), own_invlists (false),
code_size (0),
nprobe (1), max_codes (0),
nprobe (1), max_codes (0), parallel_mode (0),
maintain_direct_map (false)
{}
@@ -141,6 +145,60 @@ void IndexIVF::add (idx_t n, const float * x)
add_with_ids (n, x, nullptr);
}
void IndexIVF::add_with_ids (idx_t n, const float * x, const long *xids)
{
// do some blocking to avoid excessive allocs
idx_t bs = 65536;
if (n > bs) {
for (idx_t i0 = 0; i0 < n; i0 += bs) {
idx_t i1 = std::min (n, i0 + bs);
if (verbose) {
printf(" IndexIVF::add_with_ids %ld:%ld\n", i0, i1);
}
add_with_ids (i1 - i0, x + i0 * d,
xids ? xids + i0 : nullptr);
}
return;
}
FAISS_THROW_IF_NOT (is_trained);
std::unique_ptr<idx_t []> idx(new idx_t[n]);
quantizer->assign (n, x, idx.get());
size_t nadd = 0, nminus1 = 0;
for (size_t i = 0; i < n; i++) {
if (idx[i] < 0) nminus1++;
}
std::unique_ptr<uint8_t []> flat_codes(new uint8_t [n * code_size]);
encode_vectors (n, x, idx.get(), flat_codes.get());
#pragma omp parallel reduction(+: nadd)
{
int nt = omp_get_num_threads();
int rank = omp_get_thread_num();
// each thread takes care of a subset of lists
for (size_t i = 0; i < n; i++) {
long list_no = idx [i];
if (list_no >= 0 && list_no % nt == rank) {
long id = xids ? xids[i] : ntotal + i;
invlists->add_entry (list_no, id,
flat_codes.get() + i * code_size);
nadd++;
}
}
}
if (verbose) {
printf(" added %ld / %ld vectors (%ld -1s)\n", nadd, n, nminus1);
}
ntotal += n;
}
void IndexIVF::make_direct_map (bool new_maintain_direct_map)
{
// nothing to do
@@ -214,70 +272,152 @@ void IndexIVF::search_preassigned (idx_t n, const float *x, idx_t k,
{
InvertedListScanner *scanner = get_InvertedListScanner(store_pairs);
ScopeDeleter1<InvertedListScanner> del(scanner);
#pragma omp for
for (size_t i = i0; i < i1; i++) {
// loop over queries
const float * xi = x + i * d;
scanner->set_query (xi);
const long * keysi = keys + i * nprobe;
float * simi = distances + i * k;
long * idxi = labels + i * k;
/*****************************************************
* Depending on parallel_mode, there are two possible ways
* to organize the search. Here we define local functions
* that are in common between the two
******************************************************/
// intialize + reorder a result heap
auto init_result = [&](float *simi, idx_t *idxi) {
if (metric_type == METRIC_INNER_PRODUCT) {
heap_heapify<HeapForIP> (k, simi, idxi);
} else {
heap_heapify<HeapForL2> (k, simi, idxi);
}
};
long nscan = 0;
// loop over probes
for (size_t ik = 0; ik < nprobe; ik++) {
long key = keysi[ik]; /* select the list */
if (key < 0) {
// not enough centroids for multiprobe
continue;
}
FAISS_THROW_IF_NOT_FMT (key < (long) nlist,
"Invalid key=%ld at ik=%ld nlist=%ld\n",
key, ik, nlist);
size_t list_size = invlists->list_size(key);
// don't waste time on empty lists
if (list_size == 0) {
continue;
}
scanner->set_list (key, coarse_dis[i * nprobe + ik]);
nlistv++;
InvertedLists::ScopedCodes scodes (invlists, key);
const Index::idx_t * ids = store_pairs ? nullptr :
invlists->get_ids (key);
nheap += scanner->scan_codes (list_size, scodes.get(),
ids, simi, idxi, k);
if (ids) {
invlists->release_ids (ids);
}
nscan += list_size;
if (max_codes && nscan >= max_codes)
break;
}
ndis += nscan;
auto reorder_result = [&] (float *simi, idx_t *idxi) {
if (metric_type == METRIC_INNER_PRODUCT) {
heap_reorder<HeapForIP> (k, simi, idxi);
} else {
heap_reorder<HeapForL2> (k, simi, idxi);
}
} // parallel for
} // parallel
};
// single list scan using the current scanner (with query
// set porperly) and storing results in simi and idxi
auto scan_one_list = [&] (idx_t key, float coarse_dis_i,
float *simi, idx_t *idxi) {
if (key < 0) {
// not enough centroids for multiprobe
return (size_t)0;
}
FAISS_THROW_IF_NOT_FMT (key < (idx_t) nlist,
"Invalid key=%ld nlist=%ld\n",
key, nlist);
size_t list_size = invlists->list_size(key);
// don't waste time on empty lists
if (list_size == 0) {
return (size_t)0;
}
scanner->set_list (key, coarse_dis_i);
nlistv++;
InvertedLists::ScopedCodes scodes (invlists, key);
std::unique_ptr<InvertedLists::ScopedIds> sids;
const Index::idx_t * ids = nullptr;
if (!store_pairs) {
sids.reset (new InvertedLists::ScopedIds (invlists, key));
ids = sids->get();
}
nheap += scanner->scan_codes (list_size, scodes.get(),
ids, simi, idxi, k);
return list_size;
};
/****************************************************
* Actual loops, depending on parallel_mode
****************************************************/
if (parallel_mode == 0) {
#pragma omp for
for (size_t i = i0; i < i1; i++) {
// loop over queries
scanner->set_query (x + i * d);
float * simi = distances + i * k;
idx_t * idxi = labels + i * k;
init_result (simi, idxi);
long nscan = 0;
// loop over probes
for (size_t ik = 0; ik < nprobe; ik++) {
nscan += scan_one_list (
keys [i * nprobe + ik],
coarse_dis[i * nprobe + ik],
simi, idxi
);
if (max_codes && nscan >= max_codes) {
break;
}
}
ndis += nscan;
reorder_result (simi, idxi);
} // parallel for
} else if (parallel_mode == 1) {
std::vector <idx_t> local_idx (k);
std::vector <float> local_dis (k);
for (size_t i = i0; i < i1; i++) {
scanner->set_query (x + i * d);
init_result (local_dis.data(), local_idx.data());
#pragma omp for schedule(dynamic)
for (size_t ik = 0; ik < nprobe; ik++) {
ndis += scan_one_list (
keys [i * nprobe + ik],
coarse_dis[i * nprobe + ik],
local_dis.data(), local_idx.data()
);
// can't do the test on max_codes
}
// merge thread-local results
float * simi = distances + i * k;
idx_t * idxi = labels + i * k;
#pragma omp single
init_result (simi, idxi);
#pragma omp barrier
#pragma omp critical
{
if (metric_type == METRIC_INNER_PRODUCT) {
heap_addn<HeapForIP>
(k, simi, idxi,
local_dis.data(), local_idx.data(), k);
} else {
heap_addn<HeapForL2>
(k, simi, idxi,
local_dis.data(), local_idx.data(), k);
}
}
#pragma omp barrier
#pragma omp single
reorder_result (simi, idxi);
}
} else {
FAISS_THROW_FMT ("parallel_mode %d not supported\n",
parallel_mode);
}
} // loop over blocks
InterruptCallback::check ();
} // loop over blocks
@@ -294,61 +434,119 @@ void IndexIVF::search_preassigned (idx_t n, const float *x, idx_t k,
void IndexIVF::range_search (idx_t nx, const float *x, float radius,
RangeSearchResult *result) const
{
long * keys = new long [nx * nprobe];
ScopeDeleter<long> del (keys);
float * coarse_dis = new float [nx * nprobe];
ScopeDeleter<float> del2 (coarse_dis);
std::unique_ptr<idx_t[]> keys (new idx_t[nx * nprobe]);
std::unique_ptr<float []> coarse_dis (new float[nx * nprobe]);
double t0 = getmillisecs();
quantizer->search (nx, x, nprobe, coarse_dis, keys);
quantizer->search (nx, x, nprobe, coarse_dis.get (), keys.get ());
indexIVF_stats.quantization_time += getmillisecs() - t0;
t0 = getmillisecs();
invlists->prefetch_lists (keys, nx * nprobe);
invlists->prefetch_lists (keys.get(), nx * nprobe);
range_search_preassigned (nx, x, radius, keys.get (), coarse_dis.get (),
result);
indexIVF_stats.search_time += getmillisecs() - t0;
}
void IndexIVF::range_search_preassigned (
idx_t nx, const float *x, float radius,
const idx_t *keys, const float *coarse_dis,
RangeSearchResult *result) const
{
size_t nlistv = 0, ndis = 0;
bool store_pairs = false;
std::vector<RangeSearchPartialResult *> all_pres (omp_get_max_threads());
#pragma omp parallel reduction(+: nlistv, ndis)
{
RangeSearchPartialResult pres(result);
InvertedListScanner *scanner = get_InvertedListScanner(store_pairs);
ScopeDeleter1<InvertedListScanner> del3(scanner);
std::unique_ptr<InvertedListScanner> scanner
(get_InvertedListScanner(store_pairs));
FAISS_THROW_IF_NOT (scanner.get ());
all_pres[omp_get_thread_num()] = &pres;
// prepare the list scanning function
auto scan_list_func = [&](size_t i, size_t ik, RangeQueryResult &qres) {
idx_t key = keys[i * nprobe + ik]; /* select the list */
if (key < 0) return;
FAISS_THROW_IF_NOT_FMT (
key < (idx_t) nlist,
"Invalid key=%ld at ik=%ld nlist=%ld\n",
key, ik, nlist);
const size_t list_size = invlists->list_size(key);
if (list_size == 0) return;
InvertedLists::ScopedCodes scodes (invlists, key);
InvertedLists::ScopedIds ids (invlists, key);
scanner->set_list (key, coarse_dis[i * nprobe + ik]);
nlistv++;
ndis += list_size;
scanner->scan_codes_range (list_size, scodes.get(),
ids.get(), radius, qres);
};
if (parallel_mode == 0) {
#pragma omp for
for (size_t i = 0; i < nx; i++) {
const float * xi = x + i * d;
scanner->set_query (xi);
const long * keysi = keys + i * nprobe;
for (size_t i = 0; i < nx; i++) {
scanner->set_query (x + i * d);
RangeQueryResult & qres = pres.new_result (i);
RangeQueryResult & qres = pres.new_result (i);
for (size_t ik = 0; ik < nprobe; ik++) {
long key = keysi[ik]; /* select the list */
if (key < 0) continue;
FAISS_THROW_IF_NOT_FMT (key < (long) nlist,
"Invalid key=%ld at ik=%ld nlist=%ld\n",
key, ik, nlist);
const size_t list_size = invlists->list_size(key);
if (list_size == 0) continue;
InvertedLists::ScopedCodes scodes (invlists, key);
InvertedLists::ScopedIds ids (invlists, key);
scanner->set_list (key, coarse_dis[i * nprobe + ik]);
nlistv++;
ndis += list_size;
scanner->scan_codes_range (list_size, scodes.get(),
ids.get(), radius, qres);
for (size_t ik = 0; ik < nprobe; ik++) {
scan_list_func (i, ik, qres);
}
}
}
pres.finalize ();
} else if (parallel_mode == 1) {
for (size_t i = 0; i < nx; i++) {
scanner->set_query (x + i * d);
RangeQueryResult & qres = pres.new_result (i);
#pragma omp for schedule(dynamic)
for (size_t ik = 0; ik < nprobe; ik++) {
scan_list_func (i, ik, qres);
}
}
} else if (parallel_mode == 2) {
std::vector<RangeQueryResult *> all_qres (nx);
RangeQueryResult *qres = nullptr;
#pragma omp for schedule(dynamic)
for (size_t iik = 0; iik < nx * nprobe; iik++) {
size_t i = iik / nprobe;
size_t ik = iik % nprobe;
if (qres == nullptr || qres->qno != i) {
FAISS_ASSERT (!qres || i > qres->qno);
qres = &pres.new_result (i);
scanner->set_query (x + i * d);
}
scan_list_func (i, ik, *qres);
}
} else {
FAISS_THROW_FMT ("parallel_mode %d not supported\n", parallel_mode);
}
if (parallel_mode == 0) {
pres.finalize ();
} else {
#pragma omp barrier
#pragma omp single
RangeSearchPartialResult::merge (all_pres, false);
#pragma omp barrier
}
}
indexIVF_stats.search_time += getmillisecs() - t0;
indexIVF_stats.nq += nx;
indexIVF_stats.nlist += nlistv;
indexIVF_stats.ndis += ndis;
@@ -367,8 +565,8 @@ void IndexIVF::reconstruct (idx_t key, float* recons) const
"direct map is not initialized");
FAISS_THROW_IF_NOT_MSG (key >= 0 && key < direct_map.size(),
"invalid key");
long list_no = direct_map[key] >> 32;
long offset = direct_map[key] & 0xffffffff;
idx_t list_no = direct_map[key] >> 32;
idx_t offset = direct_map[key] & 0xffffffff;
reconstruct_from_offset (list_no, offset, recons);
}
@@ -377,12 +575,12 @@ void IndexIVF::reconstruct_n (idx_t i0, idx_t ni, float* recons) const
{
FAISS_THROW_IF_NOT (ni == 0 || (i0 >= 0 && i0 + ni <= ntotal));
for (long list_no = 0; list_no < nlist; list_no++) {
for (idx_t list_no = 0; list_no < nlist; list_no++) {
size_t list_size = invlists->list_size (list_no);
ScopedIds idlist (invlists, list_no);
for (long offset = 0; offset < list_size; offset++) {
long id = idlist[offset];
for (idx_t offset = 0; offset < list_size; offset++) {
idx_t id = idlist[offset];
if (!(id >= i0 && id < i0 + ni)) {
continue;
}
@@ -398,8 +596,8 @@ void IndexIVF::search_and_reconstruct (idx_t n, const float *x, idx_t k,
float *distances, idx_t *labels,
float *recons) const
{
long * idx = new long [n * nprobe];
ScopeDeleter<long> del (idx);
idx_t * idx = new idx_t [n * nprobe];
ScopeDeleter<idx_t> del (idx);
float * coarse_dis = new float [n * nprobe];
ScopeDeleter<float> del2 (coarse_dis);
@@ -433,8 +631,8 @@ void IndexIVF::search_and_reconstruct (idx_t n, const float *x, idx_t k,
}
void IndexIVF::reconstruct_from_offset(
long /*list_no*/,
long /*offset*/,
idx_t /*list_no*/,
idx_t /*offset*/,
float* /*recons*/) const {
FAISS_THROW_MSG ("reconstruct_from_offset not implemented");
}
@@ -447,16 +645,16 @@ void IndexIVF::reset ()
}
long IndexIVF::remove_ids (const IDSelector & sel)
Index::idx_t IndexIVF::remove_ids (const IDSelector & sel)
{
FAISS_THROW_IF_NOT_MSG (!maintain_direct_map,
"direct map remove not implemented");
std::vector<long> toremove(nlist);
std::vector<idx_t> toremove(nlist);
#pragma omp parallel for
for (long i = 0; i < nlist; i++) {
long l0 = invlists->list_size (i), l = l0, j = 0;
for (idx_t i = 0; i < nlist; i++) {
idx_t l0 = invlists->list_size (i), l = l0, j = 0;
ScopedIds idsi (invlists, i);
while (j < l) {
if (sel.is_member (idsi[j])) {
@@ -472,8 +670,8 @@ long IndexIVF::remove_ids (const IDSelector & sel)
toremove[i] = l0 - l;
}
// this will not run well in parallel on ondisk because of possible shrinks
long nremove = 0;
for (long i = 0; i < nlist; i++) {
idx_t nremove = 0;
for (idx_t i = 0; i < nlist; i++) {
if (toremove[i] > 0) {
nremove += toremove[i];
invlists->resize(
@@ -548,7 +746,7 @@ void IndexIVF::replace_invlists (InvertedLists *il, bool own)
void IndexIVF::copy_subset_to (IndexIVF & other, int subset_type,
long a1, long a2) const
idx_t a1, idx_t a2) const
{
FAISS_THROW_IF_NOT (nlist == other.nlist);
@@ -564,12 +762,12 @@ void IndexIVF::copy_subset_to (IndexIVF & other, int subset_type,
InvertedLists *oivf = other.invlists;
for (long list_no = 0; list_no < nlist; list_no++) {
for (idx_t list_no = 0; list_no < nlist; list_no++) {
size_t n = invlists->list_size (list_no);
ScopedIds ids_in (invlists, list_no);
if (subset_type == 0) {
for (long i = 0; i < n; i++) {
for (idx_t i = 0; i < n; i++) {
idx_t id = ids_in[i];
if (a1 <= id && id < a2) {
oivf->add_entry (list_no,
@@ -579,7 +777,7 @@ void IndexIVF::copy_subset_to (IndexIVF & other, int subset_type,
}
}
} else if (subset_type == 1) {
for (long i = 0; i < n; i++) {
for (idx_t i = 0; i < n; i++) {
idx_t id = ids_in[i];
if (id % a1 == a2) {
oivf->add_entry (list_no,
@@ -596,7 +794,7 @@ void IndexIVF::copy_subset_to (IndexIVF & other, int subset_type,
size_t next_accu_a2 = next_accu_n * a2 / ntotal;
size_t i2 = next_accu_a2 - accu_a2;
for (long i = i1; i < i2; i++) {
for (idx_t i = i1; i < i2; i++) {
oivf->add_entry (list_no,
invlists->get_single_id (list_no, i),
ScopedCodes (invlists, list_no, i).get());
+22 -8
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -98,9 +97,17 @@ struct IndexIVF: Index, Level1Quantizer {
size_t nprobe; ///< number of probes at query time
size_t max_codes; ///< max nb of codes to visit to do a query
/** Parallel mode determines how queries are parallelized with OpenMP
*
* 0 (default): parallelize over queries
* 1: parallelize over over inverted lists
* 2: parallelize over both
*/
int parallel_mode;
/// map for direct access to the elements. Enables reconstruct().
bool maintain_direct_map;
std::vector <long> direct_map;
std::vector <idx_t> direct_map;
/** The Inverted file takes a quantizer (an Index) on input,
* which implements the function mapping a vector to a list
@@ -119,6 +126,9 @@ struct IndexIVF: Index, Level1Quantizer {
/// Calls add_with_ids with NULL ids
void add(idx_t n, const float* x) override;
/// default implementation that calls encode_vectors
void add_with_ids(idx_t n, const float* x, const idx_t* xids) override;
/** Encodes a set of vectors as they would appear in the inverted lists
*
* @param list_nos inverted list ids as returned by the
@@ -166,6 +176,10 @@ struct IndexIVF: Index, Level1Quantizer {
void range_search (idx_t n, const float* x, float radius,
RangeSearchResult* result) const override;
void range_search_preassigned(idx_t nx, const float *x, float radius,
const idx_t *keys, const float *coarse_dis,
RangeSearchResult *result) const;
/// get a scanner for this index (store_pairs means ignore labels)
virtual InvertedListScanner *get_InvertedListScanner (
bool store_pairs=false) const;
@@ -203,13 +217,13 @@ struct IndexIVF: Index, Level1Quantizer {
* the inv list offset is computed by search_preassigned() with
* `store_pairs` set.
*/
virtual void reconstruct_from_offset (long list_no, long offset,
virtual void reconstruct_from_offset (idx_t list_no, idx_t offset,
float* recons) const;
/// Dataset manipulation functions
long remove_ids(const IDSelector& sel) override;
idx_t remove_ids(const IDSelector& sel) override;
/** check that the two indexes are compatible (ie, they are
* trained in the same way and have the same
@@ -229,7 +243,7 @@ struct IndexIVF: Index, Level1Quantizer {
* elements are left before and a2 elements are after
*/
virtual void copy_subset_to (IndexIVF & other, int subset_type,
long a1, long a2) const;
idx_t a1, idx_t a2) const;
~IndexIVF() override;
@@ -249,7 +263,7 @@ struct IndexIVF: Index, Level1Quantizer {
IndexIVF ();
};
class RangeQueryResult;
struct RangeQueryResult;
/** Object that handles a query. The inverted lists to scan are
* provided externally. The object has a lot of state, but
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+132 -92
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -447,6 +446,7 @@ void IndexIVFPQ::precompute_table ()
}
}
}
namespace {
@@ -756,23 +756,65 @@ struct QueryTables {
};
template<class C>
struct KnnSearchResults {
idx_t key;
const idx_t *ids;
// heap params
size_t k;
float * heap_sim;
long * heap_ids;
size_t nup;
inline void add (idx_t j, float dis) {
if (C::cmp (heap_sim[0], dis)) {
heap_pop<C> (k, heap_sim, heap_ids);
long id = ids ? ids[j] : (key << 32 | j);
heap_push<C> (k, heap_sim, heap_ids, dis, id);
nup++;
}
}
};
template<class C>
struct RangeSearchResults {
idx_t key;
const idx_t *ids;
// wrapped result structure
float radius;
RangeQueryResult & rres;
inline void add (idx_t j, float dis) {
if (C::cmp (radius, dis)) {
long id = ids ? ids[j] : (key << 32 | j);
rres.add (dis, id);
}
}
};
/*****************************************************
* Scaning the codes.
* The scanning functions call their favorite precompute_*
* function to precompute the tables they need.
*****************************************************/
template <typename IDType, bool store_pairs, class C, MetricType METRIC_TYPE>
template <typename IDType, MetricType METRIC_TYPE>
struct IVFPQScannerT: QueryTables {
const uint8_t * list_codes;
const IDType * list_ids;
size_t list_size;
explicit IVFPQScannerT (const IndexIVFPQ & ivfpq,
const IVFSearchParameters *params):
IVFPQScannerT (const IndexIVFPQ & ivfpq, const IVFSearchParameters *params):
QueryTables (ivfpq, params)
{
FAISS_THROW_IF_NOT (pq.byte_per_idx == 1);
FAISS_THROW_IF_NOT (pq.nbits == 8);
assert(METRIC_TYPE == metric_type);
}
@@ -795,11 +837,10 @@ struct IVFPQScannerT: QueryTables {
*****************************************************/
/// version of the scan where we use precomputed tables
size_t scan_list_with_table (
size_t ncode, const uint8_t *codes, const idx_t *ids,
size_t k, float * heap_sim, long * heap_ids) const
template<class SearchResultType>
void scan_list_with_table (size_t ncode, const uint8_t *codes,
SearchResultType & res) const
{
int nup = 0;
for (size_t j = 0; j < ncode; j++) {
float dis = dis0;
@@ -810,24 +851,17 @@ struct IVFPQScannerT: QueryTables {
tab += pq.ksub;
}
if (C::cmp (heap_sim[0], dis)) {
heap_pop<C> (k, heap_sim, heap_ids);
long id = store_pairs ? (key << 32 | j) : ids[j];
heap_push<C> (k, heap_sim, heap_ids, dis, id);
nup++;
}
res.add(j, dis);
}
return nup;
}
/// tables are not precomputed, but pointers are provided to the
/// relevant X_c|x_r tables
size_t scan_list_with_pointer (
size_t ncode, const uint8_t *codes, const idx_t *ids,
size_t k, float * heap_sim, long * heap_ids) const
template<class SearchResultType>
void scan_list_with_pointer (size_t ncode, const uint8_t *codes,
SearchResultType & res) const
{
size_t nup = 0;
for (size_t j = 0; j < ncode; j++) {
float dis = dis0;
@@ -838,26 +872,18 @@ struct IVFPQScannerT: QueryTables {
dis += sim_table_ptrs [m][ci] - 2 * tab [ci];
tab += pq.ksub;
}
if (C::cmp (heap_sim[0], dis)) {
heap_pop<C> (k, heap_sim, heap_ids);
long id = store_pairs ? (key << 32 | j) : ids[j];
heap_push<C> (k, heap_sim, heap_ids, dis, id);
nup++;
}
res.add (j, dis);
}
return nup;
}
/// nothing is precomputed: access residuals on-the-fly
size_t scan_on_the_fly_dist (
size_t ncode, const uint8_t *codes, const idx_t *ids,
size_t k, float * heap_sim, long * heap_ids) const
template<class SearchResultType>
void scan_on_the_fly_dist (size_t ncode, const uint8_t *codes,
SearchResultType &res) const
{
const float *dvec;
float dis0 = 0;
size_t nup = 0;
if (by_residual) {
if (METRIC_TYPE == METRIC_INNER_PRODUCT) {
ivfpq.quantizer->reconstruct (key, residual_vec);
@@ -882,25 +908,18 @@ struct IVFPQScannerT: QueryTables {
} else {
dis = fvec_L2sqr (decoded_vec, dvec, d);
}
if (C::cmp (heap_sim[0], dis)) {
heap_pop<C> (k, heap_sim, heap_ids);
long id = store_pairs ? (key << 32 | j) : ids[j];
heap_push<C> (k, heap_sim, heap_ids, dis, id);
nup++;
}
res.add (j, dis);
}
return nup;
}
/*****************************************************
* Scanning codes with polysemous filtering
*****************************************************/
template <class HammingComputer>
size_t scan_list_polysemous_hc (
size_t ncode, const uint8_t *codes, const idx_t *ids,
size_t k, float * heap_sim, long * heap_ids) const
template <class HammingComputer, class SearchResultType>
void scan_list_polysemous_hc (
size_t ncode, const uint8_t *codes,
SearchResultType & res) const
{
int ht = ivfpq.polysemous_ht;
size_t n_hamming_pass = 0, nup = 0;
@@ -923,12 +942,7 @@ struct IVFPQScannerT: QueryTables {
tab += pq.ksub;
}
if (C::cmp (heap_sim[0], dis)) {
heap_pop<C> (k, heap_sim, heap_ids);
long id = store_pairs ? (key << 32 | j) : ids[j];
heap_push<C> (k, heap_sim, heap_ids, dis, id);
nup++;
}
res.add (j, dis);
}
codes += code_size;
}
@@ -936,18 +950,19 @@ struct IVFPQScannerT: QueryTables {
{
indexIVFPQ_stats.n_hamming_pass += n_hamming_pass;
}
return nup;
}
size_t scan_list_polysemous (
size_t ncode, const uint8_t *codes, const idx_t *ids,
size_t k, float * heap_sim, long * heap_ids) const
template<class SearchResultType>
void scan_list_polysemous (
size_t ncode, const uint8_t *codes,
SearchResultType &res) const
{
switch (pq.code_size) {
#define HANDLE_CODE_SIZE(cs) \
case cs: \
return scan_list_polysemous_hc <HammingComputer ## cs> \
(ncode, codes, ids, k, heap_sim, heap_ids); \
scan_list_polysemous_hc \
<HammingComputer ## cs, SearchResultType> \
(ncode, codes, res); \
break
HANDLE_CODE_SIZE(4);
HANDLE_CODE_SIZE(8);
@@ -958,11 +973,13 @@ struct IVFPQScannerT: QueryTables {
#undef HANDLE_CODE_SIZE
default:
if (pq.code_size % 8 == 0)
return scan_list_polysemous_hc <HammingComputerM8>
(ncode, codes, ids, k, heap_sim, heap_ids);
scan_list_polysemous_hc
<HammingComputerM8, SearchResultType>
(ncode, codes, res);
else
return scan_list_polysemous_hc <HammingComputerM4>
(ncode, codes, ids, k, heap_sim, heap_ids);
scan_list_polysemous_hc
<HammingComputerM4, SearchResultType>
(ncode, codes, res);
break;
}
}
@@ -974,17 +991,18 @@ struct IVFPQScannerT: QueryTables {
* gain in runtime is worth the code bloat. C is the comparator < or
* >, it is directly related to METRIC_TYPE. precompute_mode is how
* much we precompute (2 = precompute distance tables, 1 = precompute
* pointers to distances, 0 = compute distances one by one). Currently
* only 2 is supported. */
template<MetricType METRIC_TYPE, bool store_pairs, class C,
int precompute_mode>
* pointers to distances, 0 = compute distances one by one).
* Currently only 2 is supported */
template<MetricType METRIC_TYPE, class C, int precompute_mode>
struct IVFPQScanner:
IVFPQScannerT<Index::idx_t, store_pairs, C, METRIC_TYPE>,
IVFPQScannerT<Index::idx_t, METRIC_TYPE>,
InvertedListScanner
{
bool store_pairs;
IVFPQScanner(const IndexIVFPQ & ivfpq):
IVFPQScannerT<Index::idx_t, store_pairs, C, METRIC_TYPE>(ivfpq, nullptr)
IVFPQScanner(const IndexIVFPQ & ivfpq, bool store_pairs):
IVFPQScannerT<Index::idx_t, METRIC_TYPE>(ivfpq, nullptr),
store_pairs(store_pairs)
{
}
@@ -1014,25 +1032,57 @@ struct IVFPQScanner:
float *heap_sim, idx_t *heap_ids,
size_t k) const override
{
KnnSearchResults<C> res = {
/* key */ this->key,
/* ids */ this->store_pairs ? nullptr : ids,
/* k */ k,
/* heap_sim */ heap_sim,
/* heap_ids */ heap_ids,
/* nup */ 0
};
if (this->polysemous_ht > 0) {
assert(precompute_mode == 2);
this->scan_list_polysemous
(ncode, codes, ids, k, heap_sim, heap_ids);
this->scan_list_polysemous (ncode, codes, res);
} else if (precompute_mode == 2) {
this->scan_list_with_table
(ncode, codes, ids, k, heap_sim, heap_ids);
this->scan_list_with_table (ncode, codes, res);
} else if (precompute_mode == 1) {
this->scan_list_with_pointer
(ncode, codes, ids, k, heap_sim, heap_ids);
this->scan_list_with_pointer (ncode, codes, res);
} else if (precompute_mode == 0) {
this->scan_on_the_fly_dist
(ncode, codes, ids, k, heap_sim, heap_ids);
this->scan_on_the_fly_dist (ncode, codes, res);
} else {
FAISS_THROW_MSG("bad precomp mode");
}
return 0;
return res.nup;
}
void scan_codes_range (size_t ncode,
const uint8_t *codes,
const idx_t *ids,
float radius,
RangeQueryResult & rres) const override
{
RangeSearchResults<C> res = {
/* key */ this->key,
/* ids */ this->store_pairs ? nullptr : ids,
/* radius */ radius,
/* rres */ rres
};
if (this->polysemous_ht > 0) {
assert(precompute_mode == 2);
this->scan_list_polysemous (ncode, codes, res);
} else if (precompute_mode == 2) {
this->scan_list_with_table (ncode, codes, res);
} else if (precompute_mode == 1) {
this->scan_list_with_pointer (ncode, codes, res);
} else if (precompute_mode == 0) {
this->scan_on_the_fly_dist (ncode, codes, res);
} else {
FAISS_THROW_MSG("bad precomp mode");
}
}
};
@@ -1044,21 +1094,11 @@ InvertedListScanner *
IndexIVFPQ::get_InvertedListScanner (bool store_pairs) const
{
if (metric_type == METRIC_INNER_PRODUCT) {
if (store_pairs) {
return new IVFPQScanner<
METRIC_INNER_PRODUCT, true, CMin<float, long>, 2> (*this);
} else {
return new IVFPQScanner<
METRIC_INNER_PRODUCT, false, CMin<float, long>, 2>(*this);
}
return new IVFPQScanner<METRIC_INNER_PRODUCT, CMin<float, long>, 2>
(*this, store_pairs);
} else if (metric_type == METRIC_L2) {
if (store_pairs) {
return new IVFPQScanner<
METRIC_L2, true, CMax<float, long>, 2> (*this);
} else {
return new IVFPQScanner<
METRIC_L2, false, CMax<float, long>, 2>(*this);
}
return new IVFPQScanner<METRIC_L2, CMax<float, long>, 2>
(*this, store_pairs);
}
return nullptr;
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+327
View File
@@ -0,0 +1,327 @@
/**
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
// -*- c++ -*-
#include "IndexIVFSpectralHash.h"
#include <memory>
#include <algorithm>
#include "hamming.h"
#include "utils.h"
#include "FaissAssert.h"
#include "AuxIndexStructures.h"
#include "VectorTransform.h"
namespace faiss {
IndexIVFSpectralHash::IndexIVFSpectralHash (
Index * quantizer, size_t d, size_t nlist,
int nbit, float period):
IndexIVF (quantizer, d, nlist, (nbit + 7) / 8, METRIC_L2),
nbit (nbit), period (period), threshold_type (Thresh_global)
{
FAISS_THROW_IF_NOT (code_size % 4 == 0);
RandomRotationMatrix *rr = new RandomRotationMatrix (d, nbit);
rr->init (1234);
vt = rr;
own_fields = true;
is_trained = false;
}
IndexIVFSpectralHash::IndexIVFSpectralHash():
IndexIVF(), vt(nullptr), own_fields(false),
nbit(0), period(0), threshold_type(Thresh_global)
{}
IndexIVFSpectralHash::~IndexIVFSpectralHash ()
{
if (own_fields) {
delete vt;
}
}
namespace {
float median (size_t n, float *x) {
std::sort(x, x + n);
if (n % 2 == 1) {
return x [n / 2];
} else {
return (x [n / 2 - 1] + x [n / 2]) / 2;
}
}
}
void IndexIVFSpectralHash::train_residual (idx_t n, const float *x)
{
if (!vt->is_trained) {
vt->train (n, x);
}
if (threshold_type == Thresh_global) {
// nothing to do
return;
} else if (threshold_type == Thresh_centroid ||
threshold_type == Thresh_centroid_half) {
// convert all centroids with vt
std::vector<float> centroids (nlist * d);
quantizer->reconstruct_n (0, nlist, centroids.data());
trained.resize(nlist * nbit);
vt->apply_noalloc (nlist, centroids.data(), trained.data());
if (threshold_type == Thresh_centroid_half) {
for (size_t i = 0; i < nlist * nbit; i++) {
trained[i] -= 0.25 * period;
}
}
return;
}
// otherwise train medians
// assign
std::unique_ptr<idx_t []> idx (new idx_t [n]);
quantizer->assign (n, x, idx.get());
std::vector<size_t> sizes(nlist + 1);
for (size_t i = 0; i < n; i++) {
FAISS_THROW_IF_NOT (idx[i] >= 0);
sizes[idx[i]]++;
}
size_t ofs = 0;
for (int j = 0; j < nlist; j++) {
size_t o0 = ofs;
ofs += sizes[j];
sizes[j] = o0;
}
// transform
std::unique_ptr<float []> xt (vt->apply (n, x));
// transpose + reorder
std::unique_ptr<float []> xo (new float[n * nbit]);
for (size_t i = 0; i < n; i++) {
size_t idest = sizes[idx[i]]++;
for (size_t j = 0; j < nbit; j++) {
xo[idest + n * j] = xt[i * nbit + j];
}
}
trained.resize (n * nbit);
// compute medians
#pragma omp for
for (int i = 0; i < nlist; i++) {
size_t i0 = i == 0 ? 0 : sizes[i - 1];
size_t i1 = sizes[i];
for (int j = 0; j < nbit; j++) {
float *xoi = xo.get() + i0 + n * j;
if (i0 == i1) { // nothing to train
trained[i * nbit + j] = 0.0;
} else if (i1 == i0 + 1) {
trained[i * nbit + j] = xoi[0];
} else {
trained[i * nbit + j] = median(i1 - i0, xoi);
}
}
}
}
namespace {
void binarize_with_freq(size_t nbit, float freq,
const float *x, const float *c,
uint8_t *codes)
{
memset (codes, 0, (nbit + 7) / 8);
for (size_t i = 0; i < nbit; i++) {
float xf = (x[i] - c[i]);
int xi = int(floor(xf * freq));
int bit = xi & 1;
codes[i >> 3] |= bit << (i & 7);
}
}
};
void IndexIVFSpectralHash::encode_vectors(idx_t n, const float* x_in,
const idx_t *list_nos,
uint8_t * codes) const
{
FAISS_THROW_IF_NOT (is_trained);
float freq = 2.0 / period;
// transform with vt
std::unique_ptr<float []> x (vt->apply (n, x_in));
#pragma omp parallel
{
std::vector<float> zero (nbit);
// each thread takes care of a subset of lists
#pragma omp for
for (size_t i = 0; i < n; i++) {
long list_no = list_nos [i];
if (list_no >= 0) {
const float *c;
if (threshold_type == Thresh_global) {
c = zero.data();
} else {
c = trained.data() + list_no * nbit;
}
binarize_with_freq (nbit, freq,
x.get() + i * nbit, c,
codes + i * code_size) ;
}
}
}
}
namespace {
template<class HammingComputer>
struct IVFScanner: InvertedListScanner {
// copied from index structure
const IndexIVFSpectralHash *index;
size_t code_size;
size_t nbit;
bool store_pairs;
float period, freq;
std::vector<float> q;
std::vector<float> zero;
std::vector<uint8_t> qcode;
HammingComputer hc;
using idx_t = Index::idx_t;
IVFScanner (const IndexIVFSpectralHash * index,
bool store_pairs):
index (index),
code_size(index->code_size),
nbit(index->nbit),
store_pairs(store_pairs),
period(index->period), freq(2.0 / index->period),
q(nbit), zero(nbit), qcode(code_size),
hc(qcode.data(), code_size)
{
}
void set_query (const float *query) override {
FAISS_THROW_IF_NOT(query);
FAISS_THROW_IF_NOT(q.size() == nbit);
index->vt->apply_noalloc (1, query, q.data());
if (index->threshold_type ==
IndexIVFSpectralHash::Thresh_global) {
binarize_with_freq
(nbit, freq, q.data(), zero.data(), qcode.data());
hc.set (qcode.data(), code_size);
}
}
idx_t list_no;
void set_list (idx_t list_no, float /*coarse_dis*/) override {
this->list_no = list_no;
if (index->threshold_type != IndexIVFSpectralHash::Thresh_global) {
const float *c = index->trained.data() + list_no * nbit;
binarize_with_freq (nbit, freq, q.data(), c, qcode.data());
hc.set (qcode.data(), code_size);
}
}
float distance_to_code (const uint8_t *code) const final {
return hc.hamming (code);
}
size_t scan_codes (size_t list_size,
const uint8_t *codes,
const idx_t *ids,
float *simi, idx_t *idxi,
size_t k) const override
{
size_t nup = 0;
for (size_t j = 0; j < list_size; j++) {
float dis = hc.hamming (codes);
if (dis < simi [0]) {
maxheap_pop (k, simi, idxi);
long id = store_pairs ? (list_no << 32 | j) : ids[j];
maxheap_push (k, simi, idxi, dis, id);
nup++;
}
codes += code_size;
}
return nup;
}
void scan_codes_range (size_t list_size,
const uint8_t *codes,
const idx_t *ids,
float radius,
RangeQueryResult & res) const override
{
for (size_t j = 0; j < list_size; j++) {
float dis = hc.hamming (codes);
if (dis < radius) {
long id = store_pairs ? (list_no << 32 | j) : ids[j];
res.add (dis, id);
}
codes += code_size;
}
}
};
} // anonymous namespace
InvertedListScanner* IndexIVFSpectralHash::get_InvertedListScanner
(bool store_pairs) const
{
switch (code_size) {
#define HANDLE_CODE_SIZE(cs) \
case cs: \
return new IVFScanner<HammingComputer ## cs> (this, store_pairs)
HANDLE_CODE_SIZE(4);
HANDLE_CODE_SIZE(8);
HANDLE_CODE_SIZE(16);
HANDLE_CODE_SIZE(20);
HANDLE_CODE_SIZE(32);
HANDLE_CODE_SIZE(64);
#undef HANDLE_CODE_SIZE
default:
if (code_size % 8 == 0) {
return new IVFScanner<HammingComputerM8>(this, store_pairs);
} else if (code_size % 4 == 0) {
return new IVFScanner<HammingComputerM4>(this, store_pairs);
} else {
FAISS_THROW_MSG("not supported");
}
}
}
} // namespace faiss
+74
View File
@@ -0,0 +1,74 @@
/**
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
// -*- c++ -*-
#ifndef FAISS_INDEX_IVFSH_H
#define FAISS_INDEX_IVFSH_H
#include <vector>
#include "IndexIVF.h"
namespace faiss {
struct VectorTransform;
/** Inverted list that stores binary codes of size nbit. Before the
* binary conversion, the dimension of the vectors is transformed from
* dim d into dim nbit by vt (a random rotation by default).
*
* Each coordinate is subtracted from a value determined by
* threshold_type, and split into intervals of size period. Half of
* the interval is a 0 bit, the other half a 1.
*/
struct IndexIVFSpectralHash: IndexIVF {
VectorTransform *vt; // transformation from d to nbit dim
bool own_fields;
int nbit;
float period;
enum ThresholdType {
Thresh_global,
Thresh_centroid,
Thresh_centroid_half,
Thresh_median
};
ThresholdType threshold_type;
// size nlist * nbit or 0 if Thresh_global
std::vector<float> trained;
IndexIVFSpectralHash (Index * quantizer, size_t d, size_t nlist,
int nbit, float period);
IndexIVFSpectralHash ();
void train_residual(idx_t n, const float* x) override;
void encode_vectors(idx_t n, const float* x,
const idx_t *list_nos,
uint8_t * codes) const override;
InvertedListScanner *get_InvertedListScanner (bool store_pairs)
const override;
~IndexIVFSpectralHash () override;
};
}; // namespace faiss
#endif
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+4 -5
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -280,7 +279,7 @@ static size_t polysemous_inner_loop (
void IndexPQ::search_core_polysemous (idx_t n, const float *x, idx_t k,
float *distances, idx_t *labels) const
{
FAISS_THROW_IF_NOT (pq.byte_per_idx == 1);
FAISS_THROW_IF_NOT (pq.nbits == 8);
// PQ distance tables
float * dis_tables = new float [n * pq.ksub * pq.M];
@@ -415,7 +414,7 @@ void IndexPQ::hamming_distance_histogram (idx_t n, const float *x,
{
FAISS_THROW_IF_NOT (metric_type == METRIC_L2);
FAISS_THROW_IF_NOT (pq.code_size % 8 == 0);
FAISS_THROW_IF_NOT (pq.byte_per_idx == 1);
FAISS_THROW_IF_NOT (pq.nbits == 8);
// Hamming embedding queries
uint8_t * q_codes = new uint8_t [n * pq.code_size];
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+70 -124
View File
@@ -1,177 +1,123 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#include "IndexReplicas.h"
#include "FaissAssert.h"
namespace faiss {
template<class IndexClass>
IndexReplicasTemplate<IndexClass>::IndexReplicasTemplate()
: own_fields(false) {
template <typename IndexT>
IndexReplicasTemplate<IndexT>::IndexReplicasTemplate(bool threaded)
: ThreadedIndex<IndexT>(threaded) {
}
template<class IndexClass>
IndexReplicasTemplate<IndexClass>::~IndexReplicasTemplate() {
if (own_fields) {
for (auto& index : this->indices_)
delete index.first;
}
template <typename IndexT>
IndexReplicasTemplate<IndexT>::IndexReplicasTemplate(idx_t d, bool threaded)
: ThreadedIndex<IndexT>(d, threaded) {
}
template<class IndexClass>
void IndexReplicasTemplate<IndexClass>::addIndex(IndexClass* index) {
// Make sure that the parameters are the same for all prior indices
if (!indices_.empty()) {
auto& existing = indices_.front().first;
template <typename IndexT>
IndexReplicasTemplate<IndexT>::IndexReplicasTemplate(int d, bool threaded)
: ThreadedIndex<IndexT>(d, threaded) {
}
FAISS_THROW_IF_NOT_FMT(index->d == existing->d,
"IndexReplicas::addIndex: dimension mismatch for "
"newly added index; prior index has dim %d, "
"new index has %d",
existing->d, index->d);
template <typename IndexT>
void
IndexReplicasTemplate<IndexT>::onAfterAddIndex(IndexT* index) {
// Make sure that the parameters are the same for all prior indices, unless
// we're the first index to be added
if (this->count() > 0 && this->at(0) != index) {
auto existing = this->at(0);
FAISS_THROW_IF_NOT_FMT(index->ntotal == existing->ntotal,
"IndexReplicas::addIndex: newly added index does "
"IndexReplicas: newly added index does "
"not have same number of vectors as prior index; "
"prior index has %ld vectors, new index has %ld",
existing->ntotal, index->ntotal);
FAISS_THROW_IF_NOT_MSG(index->metric_type == existing->metric_type,
"IndexReplicas::addIndex: newly added index is "
"of different metric type than old index");
FAISS_THROW_IF_NOT_MSG(index->is_trained == existing->is_trained,
"IndexReplicas: newly added index does "
"not have same train status as prior index");
} else {
// Set our parameters
// FIXME: this is a little bit weird
this->d = index->d;
// Set our parameters based on the first index we're adding
// (dimension is handled in ThreadedIndex)
this->ntotal = index->ntotal;
this->verbose = index->verbose;
this->is_trained = index->is_trained;
this->metric_type = index->metric_type;
}
this->indices_.emplace_back(
std::make_pair(index,
std::unique_ptr<WorkerThread>(new WorkerThread)));
}
template<class IndexClass>
void IndexReplicasTemplate<IndexClass>::removeIndex(IndexClass* index) {
for (auto it = this->indices_.begin(); it != indices_.end(); ++it) {
if (it->first == index) {
// This is our index; stop the worker thread before removing it,
// to ensure that it has finished before function exit
it->second->stop();
it->second->waitForThreadExit();
this->indices_.erase(it);
return;
}
}
// could not find our index
FAISS_THROW_MSG("IndexReplicas::removeIndex: index not found");
template <typename IndexT>
void
IndexReplicasTemplate<IndexT>::train(idx_t n, const component_t* x) {
this->runOnIndex([n, x](int, IndexT* index){ index->train(n, x); });
}
template<class IndexClass>
void IndexReplicasTemplate<IndexClass>::runOnIndex(std::function<void(IndexClass*)> f) {
FAISS_THROW_IF_NOT_MSG(!indices_.empty(), "no replicas in index");
std::vector<std::future<bool>> v;
for (auto& index : this->indices_) {
auto indexPtr = index.first;
v.emplace_back(index.second->add([indexPtr, f](){ f(indexPtr); }));
}
// Blocking wait for completion
for (auto& func : v) {
func.get();
}
}
template<class IndexClass>
void IndexReplicasTemplate<IndexClass>::reset() {
runOnIndex([](IndexClass* index){ index->reset(); });
this->ntotal = 0;
}
template<class IndexClass>
void IndexReplicasTemplate<IndexClass>::train(idx_t n, const component_t* x) {
runOnIndex([n, x](IndexClass* index){ index->train(n, x); });
}
template<class IndexClass>
void IndexReplicasTemplate<IndexClass>::add(idx_t n, const component_t* x) {
runOnIndex([n, x](IndexClass* index){ index->add(n, x); });
template <typename IndexT>
void
IndexReplicasTemplate<IndexT>::add(idx_t n, const component_t* x) {
this->runOnIndex([n, x](int, IndexT* index){ index->add(n, x); });
this->ntotal += n;
}
template<class IndexClass>
void IndexReplicasTemplate<IndexClass>::reconstruct(idx_t n, component_t* x) const {
FAISS_THROW_IF_NOT_MSG(!indices_.empty(), "no replicas in index");
indices_[0].first->reconstruct (n, x);
template <typename IndexT>
void
IndexReplicasTemplate<IndexT>::reconstruct(idx_t n, component_t* x) const {
FAISS_THROW_IF_NOT_MSG(this->count() > 0, "no replicas in index");
// Just pass to the first replica
this->at(0)->reconstruct(n, x);
}
template<class IndexClass>
void IndexReplicasTemplate<IndexClass>::search(
idx_t n,
const component_t* x,
idx_t k,
distance_t* distances,
idx_t* labels) const {
FAISS_THROW_IF_NOT_MSG(!indices_.empty(), "no replicas in index");
template <typename IndexT>
void
IndexReplicasTemplate<IndexT>::search(idx_t n,
const component_t* x,
idx_t k,
distance_t* distances,
idx_t* labels) const {
FAISS_THROW_IF_NOT_MSG(this->count() > 0, "no replicas in index");
if (n == 0) {
return;
}
auto dim = indices_.front().first->d;
std::vector<std::future<bool>> v;
auto dim = this->d;
size_t componentsPerVec =
sizeof(component_t) == 1 ? (dim + 7) / 8 : dim;
// Partition the query by the number of indices we have
auto queriesPerIndex =
(faiss::Index::idx_t) (n + indices_.size() - 1) / indices_.size();
FAISS_ASSERT(n / queriesPerIndex <= indices_.size());
faiss::Index::idx_t queriesPerIndex =
(faiss::Index::idx_t) (n + this->count() - 1) /
(faiss::Index::idx_t) this->count();
FAISS_ASSERT(n / queriesPerIndex <= this->count());
for (faiss::Index::idx_t i = 0; i < indices_.size(); ++i) {
auto base = i * queriesPerIndex;
if (base >= n) {
break;
}
auto fn =
[queriesPerIndex, componentsPerVec,
n, x, k, distances, labels](int i, const IndexT* index) {
faiss::Index::idx_t base = (faiss::Index::idx_t) i * queriesPerIndex;
auto numForIndex = std::min(queriesPerIndex, n - base);
size_t components_per_vec = sizeof(component_t) == 1 ? (dim + 7) / 8 : dim;
auto queryStart = x + base * components_per_vec;
auto distancesStart = distances + base * k;
auto labelsStart = labels + base * k;
if (base < n) {
auto numForIndex = std::min(queriesPerIndex, n - base);
auto indexPtr = indices_[i].first;
auto fn =
[indexPtr, numForIndex, queryStart, k, distancesStart, labelsStart]() {
indexPtr->search(numForIndex, queryStart,
k, distancesStart, labelsStart);
};
index->search(numForIndex,
x + base * componentsPerVec,
k,
distances + base * k,
labels + base * k);
}
};
v.emplace_back(indices_[i].second->add(std::move(fn)));
}
// Blocking wait for completion
for (auto& f : v) {
f.get();
}
this->runOnIndex(fn);
}
// explicit instanciations
// explicit instantiations
template struct IndexReplicasTemplate<Index>;
template struct IndexReplicasTemplate<IndexBinary>;
} // namespace
+32 -46
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -10,9 +9,7 @@
#include "Index.h"
#include "IndexBinary.h"
#include "WorkerThread.h"
#include <memory>
#include <vector>
#include "ThreadedIndex.h"
namespace faiss {
@@ -20,71 +17,60 @@ namespace faiss {
/// sending to each Index instance, and joins the results together
/// when done.
/// Each index is managed by a separate CPU thread.
template<class IndexClass>
class IndexReplicasTemplate : public IndexClass {
template <typename IndexT>
class IndexReplicasTemplate : public ThreadedIndex<IndexT> {
public:
using idx_t = typename IndexClass::idx_t;
using component_t = typename IndexClass::component_t;
using distance_t = typename IndexClass::distance_t;
using idx_t = typename IndexT::idx_t;
using component_t = typename IndexT::component_t;
using distance_t = typename IndexT::distance_t;
IndexReplicasTemplate();
~IndexReplicasTemplate() override;
/// The dimension that all sub-indices must share will be the dimension of the
/// first sub-index added
/// @param threaded do we use one thread per sub-index or do queries
/// sequentially?
explicit IndexReplicasTemplate(bool threaded = true);
/// Adds an index that is managed by ourselves.
/// WARNING: once an index is added to this proxy, it becomes unsafe
/// to touch it from any other thread than that on which is managing
/// it, until we are shut down. Use runOnIndex to perform work on it
/// instead.
void addIndex(IndexClass* index);
/// @param d the dimension that all sub-indices must share
/// @param threaded do we use one thread per sub index or do queries
/// sequentially?
explicit IndexReplicasTemplate(idx_t d, bool threaded = true);
/// Remove an index that is managed by ourselves.
/// This will flush all pending work on that index, and then shut
/// down its managing thread, and will remove the index.
void removeIndex(IndexClass* index);
/// int version due to the implicit bool conversion ambiguity of int as
/// dimension
explicit IndexReplicasTemplate(int d, bool threaded = true);
/// Run a function on all indices, in the thread that the index is
/// managed in.
void runOnIndex(std::function<void(IndexClass*)> f);
/// Alias for addIndex()
void add_replica(IndexT* index) { this->addIndex(index); }
/// Alias for removeIndex()
void remove_replica(IndexT* index) { this->removeIndex(index); }
/// faiss::Index API
/// All indices receive the same call
void reset() override;
void train(idx_t n, const component_t* x) override;
/// faiss::Index API
/// All indices receive the same call
virtual void train(idx_t n, const component_t* x) override;
/// faiss::Index API
/// All indices receive the same call
virtual void add(idx_t n, const component_t* x) override;
void add(idx_t n, const component_t* x) override;
/// faiss::Index API
/// Query is partitioned into a slice for each sub-index
/// split by ceil(n / #indices) for our sub-indices
virtual void search(idx_t n,
void search(idx_t n,
const component_t* x,
idx_t k,
distance_t* distances,
idx_t* labels) const override;
/// reconstructs from the first index
virtual void reconstruct(idx_t, component_t *v) const override;
void reconstruct(idx_t, component_t *v) const override;
bool own_fields;
int count() const {return indices_.size(); }
IndexClass* at(int i) {return indices_[i].first; }
const IndexClass* at(int i) const {return indices_[i].first; }
private:
/// Collection of Index instances, with their managing worker thread
mutable std::vector<std::pair<IndexClass*,
std::unique_ptr<WorkerThread> > > indices_;
protected:
/// Called just after an index is added
void onAfterAddIndex(IndexT* index) override;
};
using IndexReplicas = IndexReplicasTemplate<Index>;
using IndexBinaryReplicas = IndexReplicasTemplate<IndexBinary>;
} // namespace
+91 -28
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -133,6 +132,67 @@ struct Codec4bit {
#endif
};
struct Codec6bit {
static void encode_component (float x, uint8_t *code, int i) {
int bits = (int)(x * 63.0);
code += (i >> 2) * 3;
switch(i & 3) {
case 0:
code[0] |= bits;
break;
case 1:
code[0] |= bits << 6;
code[1] |= bits >> 2;
break;
case 2:
code[1] |= bits << 4;
code[2] |= bits >> 4;
break;
case 3:
code[2] |= bits << 2;
break;
}
}
static float decode_component (const uint8_t *code, int i) {
uint8_t bits;
code += (i >> 2) * 3;
switch(i & 3) {
case 0:
bits = code[0] & 0x3f;
break;
case 1:
bits = code[0] >> 6;
bits |= (code[1] & 0xf) << 2;
break;
case 2:
bits = code[1] >> 4;
bits |= (code[2] & 3) << 4;
break;
case 3:
bits = code[2] >> 2;
break;
}
return (bits + 0.5f) / 63.0f;
}
#ifdef USE_AVX
static __m256 decode_8_components (const uint8_t *code, int i) {
return _mm256_set_ps
(decode_component(code, i + 7),
decode_component(code, i + 6),
decode_component(code, i + 5),
decode_component(code, i + 4),
decode_component(code, i + 3),
decode_component(code, i + 2),
decode_component(code, i + 1),
decode_component(code, i + 0));
}
#endif
};
#ifdef USE_AVX
@@ -496,6 +556,8 @@ Quantizer *select_quantizer (
switch(qtype) {
case ScalarQuantizer::QT_8bit:
return new QuantizerTemplate<Codec8bit, false, SIMDWIDTH>(d, trained);
case ScalarQuantizer::QT_6bit:
return new QuantizerTemplate<Codec6bit, false, SIMDWIDTH>(d, trained);
case ScalarQuantizer::QT_4bit:
return new QuantizerTemplate<Codec4bit, false, SIMDWIDTH>(d, trained);
case ScalarQuantizer::QT_8bit_uniform:
@@ -1125,6 +1187,10 @@ SQDistanceComputer *select_distance_computer (
return new DCTemplate<QuantizerTemplate<Codec8bit, false, SIMDWIDTH>,
Sim, SIMDWIDTH>(d, trained);
case ScalarQuantizer::QT_6bit:
return new DCTemplate<QuantizerTemplate<Codec6bit, false, SIMDWIDTH>,
Sim, SIMDWIDTH>(d, trained);
case ScalarQuantizer::QT_4bit:
return new DCTemplate<QuantizerTemplate<Codec4bit, false, SIMDWIDTH>,
Sim, SIMDWIDTH>(d, trained);
@@ -1169,6 +1235,9 @@ ScalarQuantizer::ScalarQuantizer
case QT_4bit_uniform:
code_size = (d + 1) / 2;
break;
case QT_6bit:
code_size = (d * 6 + 7) / 8;
break;
case QT_fp16:
code_size = d * 2;
break;
@@ -1186,6 +1255,7 @@ void ScalarQuantizer::train (size_t n, const float *x)
int bit_per_dim =
qtype == QT_4bit_uniform ? 4 :
qtype == QT_4bit ? 4 :
qtype == QT_6bit ? 6 :
qtype == QT_8bit_uniform ? 8 :
qtype == QT_8bit ? 8 : -1;
@@ -1194,7 +1264,7 @@ void ScalarQuantizer::train (size_t n, const float *x)
train_Uniform (rangestat, rangestat_arg,
n * d, 1 << bit_per_dim, x, trained);
break;
case QT_4bit: case QT_8bit:
case QT_4bit: case QT_8bit: case QT_6bit:
train_NonUniform (rangestat, rangestat_arg,
n, d, 1 << bit_per_dim, x, trained);
break;
@@ -1263,10 +1333,10 @@ ScalarQuantizer::get_distance_computer (MetricType metric) const
namespace {
template<bool store_pairs, class DCClass>
template<class DCClass>
struct IVFSQScannerIP: InvertedListScanner {
DCClass dc;
bool by_residual;
bool store_pairs, by_residual;
size_t code_size;
@@ -1274,8 +1344,10 @@ struct IVFSQScannerIP: InvertedListScanner {
float accu0; /// added to all distances
IVFSQScannerIP(int d, const std::vector<float> & trained,
size_t code_size, bool by_residual=false):
dc(d, trained), by_residual(by_residual),
size_t code_size, bool store_pairs,
bool by_residual):
dc(d, trained), store_pairs(store_pairs),
by_residual(by_residual),
code_size(code_size), list_no(0), accu0(0)
{}
@@ -1336,12 +1408,12 @@ struct IVFSQScannerIP: InvertedListScanner {
};
template<bool store_pairs, class DCClass>
template<class DCClass>
struct IVFSQScannerL2: InvertedListScanner {
DCClass dc;
bool by_residual;
bool store_pairs, by_residual;
size_t code_size;
const Index *quantizer;
idx_t list_no; /// current inverted list
@@ -1351,8 +1423,8 @@ struct IVFSQScannerL2: InvertedListScanner {
IVFSQScannerL2(int d, const std::vector<float> & trained,
size_t code_size, const Index *quantizer,
bool by_residual):
dc(d, trained), by_residual(by_residual),
bool store_pairs, bool by_residual):
dc(d, trained), store_pairs(store_pairs), by_residual(by_residual),
code_size(code_size), quantizer(quantizer),
list_no (0), x (nullptr), tmp (d)
{
@@ -1429,21 +1501,11 @@ InvertedListScanner* sel2_InvertedListScanner
const Index *quantizer, bool store_pairs, bool r)
{
if (DCClass::Sim::metric_type == METRIC_L2) {
if (store_pairs) {
return new IVFSQScannerL2<true, DCClass>
(sq->d, sq->trained, sq->code_size, quantizer, r);
} else {
return new IVFSQScannerL2<false, DCClass>
(sq->d, sq->trained, sq->code_size, quantizer, r);
}
return new IVFSQScannerL2<DCClass>(sq->d, sq->trained, sq->code_size,
quantizer, store_pairs, r);
} else {
if (store_pairs) {
return new IVFSQScannerIP<true, DCClass>
(sq->d, sq->trained, sq->code_size, r);
} else {
return new IVFSQScannerIP<false, DCClass>
(sq->d, sq->trained, sq->code_size, r);
}
return new IVFSQScannerIP<DCClass>(sq->d, sq->trained, sq->code_size,
store_pairs, r);
}
}
@@ -1479,6 +1541,9 @@ InvertedListScanner* sel1_InvertedListScanner
case ScalarQuantizer::QT_4bit:
return sel12_InvertedListScanner
<Similarity, Codec4bit, false>(sq, quantizer, store_pairs, r);
case ScalarQuantizer::QT_6bit:
return sel12_InvertedListScanner
<Similarity, Codec6bit, false>(sq, quantizer, store_pairs, r);
case ScalarQuantizer::QT_fp16:
return sel2_InvertedListScanner
<DCTemplate<QuantizerFP16<SIMDWIDTH>, Similarity, SIMDWIDTH> >
@@ -1720,8 +1785,6 @@ void IndexIVFScalarQuantizer::encode_vectors(idx_t n, const float* x,
xi = residual.data ();
}
squant->encode_vector (xi, codes + i * code_size);
} else {
memset (codes + i * code_size, 0, code_size);
}
}
}
+3 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -39,6 +38,7 @@ struct ScalarQuantizer {
QT_4bit_uniform,
QT_fp16,
QT_8bit_direct, /// fast indexing of uint8s
QT_6bit, ///< 6 bits per component
};
QuantizerType qtype;
+219 -232
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -43,288 +42,276 @@ void translate_labels (long n, idx_t *labels, long translation)
*/
template <class IndexClass, class C>
void merge_tables (long n, long k, long nshard,
typename IndexClass::distance_t *distances,
idx_t *labels,
const typename IndexClass::distance_t *all_distances,
idx_t *all_labels,
const long *translations)
{
if(k == 0) {
return;
}
using distance_t = typename IndexClass::distance_t;
void
merge_tables(long n, long k, long nshard,
typename IndexClass::distance_t *distances,
idx_t *labels,
const std::vector<typename IndexClass::distance_t>& all_distances,
const std::vector<idx_t>& all_labels,
const std::vector<long>& translations) {
if (k == 0) {
return;
}
using distance_t = typename IndexClass::distance_t;
long stride = n * k;
long stride = n * k;
#pragma omp parallel
{
std::vector<int> buf (2 * nshard);
int * pointer = buf.data();
int * shard_ids = pointer + nshard;
std::vector<distance_t> buf2 (nshard);
distance_t * heap_vals = buf2.data();
{
std::vector<int> buf (2 * nshard);
int * pointer = buf.data();
int * shard_ids = pointer + nshard;
std::vector<distance_t> buf2 (nshard);
distance_t * heap_vals = buf2.data();
#pragma omp for
for (long i = 0; i < n; i++) {
// the heap maps values to the shard where they are
// produced.
const distance_t *D_in = all_distances + i * k;
const idx_t *I_in = all_labels + i * k;
int heap_size = 0;
for (long i = 0; i < n; i++) {
// the heap maps values to the shard where they are
// produced.
const distance_t *D_in = all_distances.data() + i * k;
const idx_t *I_in = all_labels.data() + i * k;
int heap_size = 0;
for (long s = 0; s < nshard; s++) {
pointer[s] = 0;
if (I_in[stride * s] >= 0)
heap_push<C> (++heap_size, heap_vals, shard_ids,
D_in[stride * s], s);
}
distance_t *D = distances + i * k;
idx_t *I = labels + i * k;
for (int j = 0; j < k; j++) {
if (heap_size == 0) {
I[j] = -1;
D[j] = C::neutral();
} else {
// pop best element
int s = shard_ids[0];
int & p = pointer[s];
D[j] = heap_vals[0];
I[j] = I_in[stride * s + p] + translations[s];
heap_pop<C> (heap_size--, heap_vals, shard_ids);
p++;
if (p < k && I_in[stride * s + p] >= 0)
heap_push<C> (++heap_size, heap_vals, shard_ids,
D_in[stride * s + p], s);
}
}
for (long s = 0; s < nshard; s++) {
pointer[s] = 0;
if (I_in[stride * s] >= 0) {
heap_push<C> (++heap_size, heap_vals, shard_ids,
D_in[stride * s], s);
}
}
distance_t *D = distances + i * k;
idx_t *I = labels + i * k;
for (int j = 0; j < k; j++) {
if (heap_size == 0) {
I[j] = -1;
D[j] = C::neutral();
} else {
// pop best element
int s = shard_ids[0];
int & p = pointer[s];
D[j] = heap_vals[0];
I[j] = I_in[stride * s + p] + translations[s];
heap_pop<C> (heap_size--, heap_vals, shard_ids);
p++;
if (p < k && I_in[stride * s + p] >= 0) {
heap_push<C> (++heap_size, heap_vals, shard_ids,
D_in[stride * s + p], s);
}
}
}
}
}
}
template<class IndexClass>
void runOnIndexes(bool threaded,
std::function<void(int no, IndexClass*)> f,
std::vector<IndexClass *> indexes)
{
FAISS_THROW_IF_NOT_MSG(!indexes.empty(), "no shards in index");
if (!threaded) {
for (int no = 0; no < indexes.size(); no++) {
IndexClass *index = indexes[no];
f(no, index);
}
} else {
std::vector<std::unique_ptr<WorkerThread> > threads;
std::vector<std::future<bool>> v;
for (int no = 0; no < indexes.size(); no++) {
IndexClass *index = indexes[no];
threads.emplace_back(new WorkerThread());
WorkerThread *wt = threads.back().get();
v.emplace_back(wt->add([no, index, f](){ f(no, index); }));
}
// Blocking wait for completion
for (auto& func : v) {
func.get();
}
}
};
} // anonymous namespace
template<class IndexClass>
IndexShardsTemplate<IndexClass>::IndexShardsTemplate (idx_t d, bool threaded, bool successive_ids):
IndexClass (d), own_fields (false),
threaded (threaded), successive_ids (successive_ids)
{
template <typename IndexT>
IndexShardsTemplate<IndexT>::IndexShardsTemplate(idx_t d,
bool threaded,
bool successive_ids)
: ThreadedIndex<IndexT>(d, threaded),
successive_ids(successive_ids) {
}
template<class IndexClass>
void IndexShardsTemplate<IndexClass>::add_shard (IndexClass *idx)
{
shard_indexes.push_back (idx);
sync_with_shard_indexes ();
template <typename IndexT>
IndexShardsTemplate<IndexT>::IndexShardsTemplate(int d,
bool threaded,
bool successive_ids)
: ThreadedIndex<IndexT>(d, threaded),
successive_ids(successive_ids) {
}
template<class IndexClass>
void IndexShardsTemplate<IndexClass>::sync_with_shard_indexes ()
{
if (shard_indexes.empty()) return;
IndexClass * index0 = shard_indexes[0];
this->d = index0->d;
this->metric_type = index0->metric_type;
this->is_trained = index0->is_trained;
this->ntotal = index0->ntotal;
for (int i = 1; i < shard_indexes.size(); i++) {
IndexClass * index = shard_indexes[i];
FAISS_THROW_IF_NOT (this->metric_type == index->metric_type);
FAISS_THROW_IF_NOT (this->d == index->d);
this->ntotal += index->ntotal;
}
template <typename IndexT>
IndexShardsTemplate<IndexT>::IndexShardsTemplate(bool threaded,
bool successive_ids)
: ThreadedIndex<IndexT>(threaded),
successive_ids(successive_ids) {
}
template <typename IndexT>
void
IndexShardsTemplate<IndexT>::onAfterAddIndex(IndexT* index /* unused */) {
sync_with_shard_indexes();
}
template <typename IndexT>
void
IndexShardsTemplate<IndexT>::onAfterRemoveIndex(IndexT* index /* unused */) {
sync_with_shard_indexes();
}
template <typename IndexT>
void
IndexShardsTemplate<IndexT>::sync_with_shard_indexes() {
if (!this->count()) {
this->is_trained = false;
this->ntotal = 0;
return;
}
auto firstIndex = this->at(0);
this->metric_type = firstIndex->metric_type;
this->is_trained = firstIndex->is_trained;
this->ntotal = firstIndex->ntotal;
template<class IndexClass>
void IndexShardsTemplate<IndexClass>::train (idx_t n, const component_t *x)
{
auto train_func = [n, x](int no, IndexClass *index)
{
if (index->verbose)
printf ("begin train shard %d on %ld points\n", no, n);
index->train(n, x);
if (index->verbose)
printf ("end train shard %d\n", no);
for (int i = 1; i < this->count(); ++i) {
auto index = this->at(i);
FAISS_THROW_IF_NOT(this->metric_type == index->metric_type);
FAISS_THROW_IF_NOT(this->d == index->d);
this->ntotal += index->ntotal;
}
}
template <typename IndexT>
void
IndexShardsTemplate<IndexT>::train(idx_t n,
const component_t *x) {
auto fn =
[n, x](int no, IndexT *index) {
if (index->verbose) {
printf("begin train shard %d on %ld points\n", no, n);
}
index->train(n, x);
if (index->verbose) {
printf("end train shard %d\n", no);
}
};
runOnIndexes<IndexClass> (threaded, train_func, shard_indexes);
sync_with_shard_indexes ();
this->runOnIndex(fn);
sync_with_shard_indexes();
}
template<class IndexClass>
void IndexShardsTemplate<IndexClass>::add (idx_t n, const component_t *x)
{
add_with_ids (n, x, nullptr);
template <typename IndexT>
void
IndexShardsTemplate<IndexT>::add(idx_t n,
const component_t *x) {
add_with_ids(n, x, nullptr);
}
template<class IndexClass>
void IndexShardsTemplate<IndexClass>::add_with_ids (idx_t n, const component_t * x, const idx_t *xids)
{
template <typename IndexT>
void
IndexShardsTemplate<IndexT>::add_with_ids(idx_t n,
const component_t * x,
const idx_t *xids) {
FAISS_THROW_IF_NOT_MSG(!(successive_ids && xids),
"It makes no sense to pass in ids and "
"request them to be shifted");
FAISS_THROW_IF_NOT_MSG(!(successive_ids && xids),
"It makes no sense to pass in ids and "
"request them to be shifted");
if (successive_ids) {
FAISS_THROW_IF_NOT_MSG(!xids,
"It makes no sense to pass in ids and "
"request them to be shifted");
FAISS_THROW_IF_NOT_MSG(this->ntotal == 0,
"when adding to IndexShards with sucessive_ids, "
"only add() in a single pass is supported");
if (successive_ids) {
FAISS_THROW_IF_NOT_MSG(!xids,
"It makes no sense to pass in ids and "
"request them to be shifted");
FAISS_THROW_IF_NOT_MSG(this->ntotal == 0,
"when adding to IndexShards with sucessive_ids, "
"only add() in a single pass is supported");
}
idx_t nshard = this->count();
const idx_t *ids = xids;
std::vector<idx_t> aids;
if (!ids && !successive_ids) {
aids.resize(n);
for (idx_t i = 0; i < n; i++) {
aids[i] = this->ntotal + i;
}
long nshard = shard_indexes.size();
const idx_t *ids = xids;
ScopeDeleter<idx_t> del;
if (!ids && !successive_ids) {
idx_t *aids = new idx_t[n];
for (idx_t i = 0; i < n; i++)
aids[i] = this->ntotal + i;
ids = aids;
del.set (ids);
}
ids = aids.data();
}
size_t components_per_vec =
sizeof(component_t) == 1 ? (this->d + 7) / 8 : this->d;
size_t components_per_vec =
sizeof(component_t) == 1 ? (this->d + 7) / 8 : this->d;
auto add_func = [n, ids, x, nshard, components_per_vec]
(int no, IndexClass *index) {
auto fn =
[n, ids, x, nshard, components_per_vec](int no, IndexT *index) {
idx_t i0 = (idx_t) no * n / nshard;
idx_t i1 = ((idx_t) no + 1) * n / nshard;
auto x0 = x + i0 * components_per_vec;
idx_t i0 = no * n / nshard;
idx_t i1 = (no + 1) * n / nshard;
if (index->verbose) {
printf ("begin add shard %d on %ld points\n", no, n);
}
auto x0 = x + i0 * components_per_vec;
if (ids) {
index->add_with_ids (i1 - i0, x0, ids + i0);
} else {
index->add (i1 - i0, x0);
}
if (index->verbose) {
printf ("begin add shard %d on %ld points\n", no, n);
}
if (ids) {
index->add_with_ids (i1 - i0, x0, ids + i0);
} else {
index->add (i1 - i0, x0);
}
if (index->verbose) {
printf ("end add shard %d on %ld points\n", no, i1 - i0);
}
if (index->verbose) {
printf ("end add shard %d on %ld points\n", no, i1 - i0);
}
};
runOnIndexes<IndexClass> (threaded, add_func, shard_indexes);
this->runOnIndex(fn);
this->ntotal += n;
// This is safe to do here because the current thread controls execution in
// all threads, and nothing else is happening
this->ntotal += n;
}
template<class IndexClass>
void IndexShardsTemplate<IndexClass>::reset ()
{
for (int i = 0; i < shard_indexes.size(); i++) {
shard_indexes[i]->reset ();
}
sync_with_shard_indexes ();
}
template <typename IndexT>
void
IndexShardsTemplate<IndexT>::search(idx_t n,
const component_t *x,
idx_t k,
distance_t *distances,
idx_t *labels) const {
long nshard = this->count();
template<class IndexClass>
void IndexShardsTemplate<IndexClass>::search (
idx_t n, const component_t *x, idx_t k,
distance_t *distances, idx_t *labels) const
{
long nshard = shard_indexes.size();
distance_t *all_distances = new distance_t [nshard * k * n];
idx_t *all_labels = new idx_t [nshard * k * n];
ScopeDeleter<distance_t> del (all_distances);
ScopeDeleter<idx_t> del2 (all_labels);
std::vector<distance_t> all_distances(nshard * k * n);
std::vector<idx_t> all_labels(nshard * k * n);
auto query_func = [n, k, x, all_distances, all_labels]
(int no, IndexClass *index) {
auto fn =
[n, k, x, &all_distances, &all_labels](int no, const IndexT *index) {
if (index->verbose) {
printf ("begin query shard %d on %ld points\n", no, n);
}
if (index->verbose) {
printf ("begin query shard %d on %ld points\n", no, n);
}
index->search (n, x, k,
all_distances + no * k * n,
all_labels + no * k * n);
if (index->verbose) {
printf ("end query shard %d\n", no);
}
index->search (n, x, k,
all_distances.data() + no * k * n,
all_labels.data() + no * k * n);
if (index->verbose) {
printf ("end query shard %d\n", no);
}
};
runOnIndexes<IndexClass> (threaded, query_func, shard_indexes);
this->runOnIndex(fn);
std::vector<long> translations (nshard, 0);
if (successive_ids) {
translations[0] = 0;
for (int s = 0; s + 1 < nshard; s++)
translations [s + 1] = translations [s] +
shard_indexes [s]->ntotal;
std::vector<long> translations(nshard, 0);
// Because we just called runOnIndex above, it is safe to access the sub-index
// ntotal here
if (successive_ids) {
translations[0] = 0;
for (int s = 0; s + 1 < nshard; s++) {
translations[s + 1] = translations[s] + this->at(s)->ntotal;
}
}
if (this->metric_type == METRIC_L2) {
merge_tables<IndexClass, CMin<distance_t, int> > (
n, k, nshard, distances, labels,
all_distances, all_labels, translations.data ());
} else {
merge_tables<IndexClass, CMax<distance_t, int> > (
n, k, nshard, distances, labels,
all_distances, all_labels, translations.data ());
}
}
template<class IndexClass>
IndexShardsTemplate<IndexClass>::~IndexShardsTemplate ()
{
if (own_fields) {
for (int s = 0; s < shard_indexes.size(); s++)
delete shard_indexes [s];
}
if (this->metric_type == METRIC_L2) {
merge_tables<IndexT, CMin<distance_t, int>>(
n, k, nshard, distances, labels,
all_distances, all_labels, translations);
} else {
merge_tables<IndexT, CMax<distance_t, int>>(
n, k, nshard, distances, labels,
all_distances, all_labels, translations);
}
}
// explicit instanciations
template struct IndexShardsTemplate<Index>;
template struct IndexShardsTemplate<IndexBinary>;
} // namespace faiss
+68 -51
View File
@@ -1,79 +1,96 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
// -*- c++ -*-
#pragma once
#include <vector>
#include "Index.h"
#include "IndexBinary.h"
#include "ThreadedIndex.h"
namespace faiss {
/** Index that concatenates the results from several sub-indexes
*
/**
* Index that concatenates the results from several sub-indexes
*/
template<class IndexClass>
struct IndexShardsTemplate : IndexClass {
template <typename IndexT>
struct IndexShardsTemplate : public ThreadedIndex<IndexT> {
using idx_t = typename IndexT::idx_t;
using component_t = typename IndexT::component_t;
using distance_t = typename IndexT::distance_t;
using idx_t = typename IndexClass::idx_t;
using component_t = typename IndexClass::component_t;
using distance_t = typename IndexClass::distance_t;
/**
* The dimension that all sub-indices must share will be the dimension of the
* first sub-index added
*
* @param threaded do we use one thread per sub_index or do
* queries sequentially?
* @param successive_ids should we shift the returned ids by
* the size of each sub-index or return them
* as they are?
*/
explicit IndexShardsTemplate(bool threaded = false,
bool successive_ids = true);
std::vector<IndexClass*> shard_indexes;
bool own_fields; /// should the sub-indexes be deleted along with this?
bool threaded;
bool successive_ids;
/**
* @param threaded do we use one thread per sub_index or do
* queries sequentially?
* @param successive_ids should we shift the returned ids by
* the size of each sub-index or return them
* as they are?
*/
explicit IndexShardsTemplate(idx_t d,
bool threaded = false,
bool successive_ids = true);
/**
* @param threaded do we use one thread per sub_index or do
* queries sequentially?
* @param successive_ids should we shift the returned ids by
* the size of each sub-index or return them
* as they are?
*/
explicit IndexShardsTemplate (idx_t d, bool threaded = false,
bool successive_ids = true);
/// int version due to the implicit bool conversion ambiguity of int as
/// dimension
explicit IndexShardsTemplate(int d,
bool threaded = false,
bool successive_ids = true);
void add_shard (IndexClass *);
/// Alias for addIndex()
void add_shard(IndexT* index) { this->addIndex(index); }
// update metric_type and ntotal. Call if you changes something in
// the shard indexes.
void sync_with_shard_indexes ();
/// Alias for removeIndex()
void remove_shard(IndexT* index) { this->removeIndex(index); }
IndexClass *at(int i) {return shard_indexes[i]; }
/// supported only for sub-indices that implement add_with_ids
void add(idx_t n, const component_t* x) override;
/// supported only for sub-indices that implement add_with_ids
void add(idx_t n, const component_t* x) override;
/**
* Cases (successive_ids, xids):
* - true, non-NULL ERROR: it makes no sense to pass in ids and
* request them to be shifted
* - true, NULL OK, but should be called only once (calls add()
* on sub-indexes).
* - false, non-NULL OK: will call add_with_ids with passed in xids
* distributed evenly over shards
* - false, NULL OK: will call add_with_ids on each sub-index,
* starting at ntotal
*/
void add_with_ids(idx_t n, const component_t* x, const idx_t* xids) override;
/**
* Cases (successive_ids, xids):
* - true, non-NULL ERROR: it makes no sense to pass in ids and
* request them to be shifted
* - true, NULL OK, but should be called only once (calls add()
* on sub-indexes).
* - false, non-NULL OK: will call add_with_ids with passed in xids
* distributed evenly over shards
* - false, NULL OK: will call add_with_ids on each sub-index,
* starting at ntotal
*/
void add_with_ids(idx_t n, const component_t* x, const idx_t* xids) override;
void search(idx_t n, const component_t* x, idx_t k,
distance_t* distances, idx_t* labels) const override;
void search(
idx_t n, const component_t* x, idx_t k,
distance_t* distances, idx_t* labels) const override;
void train(idx_t n, const component_t* x) override;
void train(idx_t n, const component_t* x) override;
// update metric_type and ntotal. Call if you changes something in
// the shard indexes.
void sync_with_shard_indexes();
void reset() override;
bool successive_ids;
~IndexShardsTemplate() override;
protected:
/// Called just after an index is added
void onAfterAddIndex(IndexT* index) override;
/// Called just after an index is removed
void onAfterRemoveIndex(IndexT* index) override;
};
using IndexShards = IndexShardsTemplate<Index>;
+336 -26
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -41,13 +40,13 @@ InvertedLists::idx_t InvertedLists::get_single_id (
}
void InvertedLists::release_codes (const uint8_t *) const
void InvertedLists::release_codes (size_t, const uint8_t *) const
{}
void InvertedLists::release_ids (const idx_t *) const
void InvertedLists::release_ids (size_t, const idx_t *) const
{}
void InvertedLists::prefetch_lists (const long *, int) const
void InvertedLists::prefetch_lists (const idx_t *, int) const
{}
const uint8_t * InvertedLists::get_single_code (
@@ -78,7 +77,7 @@ void InvertedLists::reset () {
void InvertedLists::merge_from (InvertedLists *oivf, size_t add_id) {
#pragma omp parallel for
for (long i = 0; i < nlist; i++) {
for (idx_t i = 0; i < nlist; i++) {
size_t list_size = oivf->list_size (i);
ScopedIds ids (oivf, i);
if (add_id == 0) {
@@ -124,6 +123,13 @@ void InvertedLists::print_stats () const {
}
}
size_t InvertedLists::compute_ntotal () const {
size_t tot = 0;
for (size_t i = 0; i < nlist; i++) {
tot += list_size(i);
}
return tot;
}
/*****************************************
* ArrayInvertedLists implementation
@@ -189,14 +195,38 @@ void ArrayInvertedLists::update_entries (
ArrayInvertedLists::~ArrayInvertedLists ()
{}
/*****************************************************************
* Meta-inverted list implementations
*****************************************************************/
size_t ReadOnlyInvertedLists::add_entries (
size_t , size_t ,
const idx_t* , const uint8_t *)
{
FAISS_THROW_MSG ("not implemented");
}
void ReadOnlyInvertedLists::update_entries (size_t, size_t , size_t ,
const idx_t *, const uint8_t *)
{
FAISS_THROW_MSG ("not implemented");
}
void ReadOnlyInvertedLists::resize (size_t , size_t )
{
FAISS_THROW_MSG ("not implemented");
}
/*****************************************
* ConcatenatedInvertedLists implementation
* HStackInvertedLists implementation
******************************************/
ConcatenatedInvertedLists::ConcatenatedInvertedLists (
HStackInvertedLists::HStackInvertedLists (
int nil, const InvertedLists **ils_in):
InvertedLists (nil > 0 ? ils_in[0]->nlist : 0,
ReadOnlyInvertedLists (nil > 0 ? ils_in[0]->nlist : 0,
nil > 0 ? ils_in[0]->code_size : 0)
{
FAISS_THROW_IF_NOT (nil > 0);
@@ -207,7 +237,7 @@ ConcatenatedInvertedLists::ConcatenatedInvertedLists (
}
}
size_t ConcatenatedInvertedLists::list_size(size_t list_no) const
size_t HStackInvertedLists::list_size(size_t list_no) const
{
size_t sz = 0;
for (int i = 0; i < ils.size(); i++) {
@@ -217,7 +247,7 @@ size_t ConcatenatedInvertedLists::list_size(size_t list_no) const
return sz;
}
const uint8_t * ConcatenatedInvertedLists::get_codes (size_t list_no) const
const uint8_t * HStackInvertedLists::get_codes (size_t list_no) const
{
uint8_t *codes = new uint8_t [code_size * list_size(list_no)], *c = codes;
@@ -232,7 +262,7 @@ const uint8_t * ConcatenatedInvertedLists::get_codes (size_t list_no) const
return codes;
}
const uint8_t * ConcatenatedInvertedLists::get_single_code (
const uint8_t * HStackInvertedLists::get_single_code (
size_t list_no, size_t offset) const
{
for (int i = 0; i < ils.size(); i++) {
@@ -250,11 +280,11 @@ const uint8_t * ConcatenatedInvertedLists::get_single_code (
}
void ConcatenatedInvertedLists::release_codes (const uint8_t *codes) const {
void HStackInvertedLists::release_codes (size_t, const uint8_t *codes) const {
delete [] codes;
}
const Index::idx_t * ConcatenatedInvertedLists::get_ids (size_t list_no) const
const Index::idx_t * HStackInvertedLists::get_ids (size_t list_no) const
{
idx_t *ids = new idx_t [list_size(list_no)], *c = ids;
@@ -269,7 +299,7 @@ const Index::idx_t * ConcatenatedInvertedLists::get_ids (size_t list_no) const
return ids;
}
Index::idx_t ConcatenatedInvertedLists::get_single_id (
Index::idx_t HStackInvertedLists::get_single_id (
size_t list_no, size_t offset) const
{
@@ -285,28 +315,308 @@ Index::idx_t ConcatenatedInvertedLists::get_single_id (
}
void ConcatenatedInvertedLists::release_ids (const idx_t *ids) const {
void HStackInvertedLists::release_ids (size_t, const idx_t *ids) const {
delete [] ids;
}
size_t ConcatenatedInvertedLists::add_entries (
size_t , size_t ,
const idx_t* , const uint8_t *)
void HStackInvertedLists::prefetch_lists (const idx_t *list_nos, int nlist) const
{
FAISS_THROW_MSG ("not implemented");
for (int i = 0; i < ils.size(); i++) {
const InvertedLists *il = ils[i];
il->prefetch_lists (list_nos, nlist);
}
}
void ConcatenatedInvertedLists::update_entries (size_t, size_t , size_t ,
const idx_t *, const uint8_t *)
/*****************************************
* SliceInvertedLists implementation
******************************************/
namespace {
using idx_t = InvertedLists::idx_t;
idx_t translate_list_no (const SliceInvertedLists *sil,
idx_t list_no) {
FAISS_THROW_IF_NOT (list_no >= 0 && list_no < sil->nlist);
return list_no + sil->i0;
}
};
SliceInvertedLists::SliceInvertedLists (
const InvertedLists *il, idx_t i0, idx_t i1):
ReadOnlyInvertedLists (i1 - i0, il->code_size),
il (il), i0(i0), i1(i1)
{
FAISS_THROW_MSG ("not implemented");
}
void ConcatenatedInvertedLists::resize (size_t , size_t )
size_t SliceInvertedLists::list_size(size_t list_no) const
{
FAISS_THROW_MSG ("not implemented");
return il->list_size (translate_list_no (this, list_no));
}
const uint8_t * SliceInvertedLists::get_codes (size_t list_no) const
{
return il->get_codes (translate_list_no (this, list_no));
}
const uint8_t * SliceInvertedLists::get_single_code (
size_t list_no, size_t offset) const
{
return il->get_single_code (translate_list_no (this, list_no), offset);
}
void SliceInvertedLists::release_codes (
size_t list_no, const uint8_t *codes) const {
return il->release_codes (translate_list_no (this, list_no), codes);
}
const Index::idx_t * SliceInvertedLists::get_ids (size_t list_no) const
{
return il->get_ids (translate_list_no (this, list_no));
}
Index::idx_t SliceInvertedLists::get_single_id (
size_t list_no, size_t offset) const
{
return il->get_single_id (translate_list_no (this, list_no), offset);
}
void SliceInvertedLists::release_ids (size_t list_no, const idx_t *ids) const {
return il->release_ids (translate_list_no (this, list_no), ids);
}
void SliceInvertedLists::prefetch_lists (const idx_t *list_nos, int nlist) const
{
std::vector<idx_t> translated_list_nos;
for (int j = 0; j < nlist; j++) {
idx_t list_no = list_nos[j];
if (list_no < 0) continue;
translated_list_nos.push_back (translate_list_no (this, list_no));
}
il->prefetch_lists (translated_list_nos.data(),
translated_list_nos.size());
}
/*****************************************
* VStackInvertedLists implementation
******************************************/
namespace {
using idx_t = InvertedLists::idx_t;
// find the invlist this number belongs to
int translate_list_no (const VStackInvertedLists *vil,
idx_t list_no) {
FAISS_THROW_IF_NOT (list_no >= 0 && list_no < vil->nlist);
int i0 = 0, i1 = vil->ils.size();
const idx_t *cumsz = vil->cumsz.data();
while (i0 + 1 < i1) {
int imed = (i0 + i1) / 2;
if (list_no >= cumsz[imed]) {
i0 = imed;
} else {
i1 = imed;
}
}
assert(list_no >= cumsz[i0] && list_no < cumsz[i0 + 1]);
return i0;
}
idx_t sum_il_sizes (int nil, const InvertedLists **ils_in) {
idx_t tot = 0;
for (int i = 0; i < nil; i++) {
tot += ils_in[i]->nlist;
}
return tot;
}
};
VStackInvertedLists::VStackInvertedLists (
int nil, const InvertedLists **ils_in):
ReadOnlyInvertedLists (sum_il_sizes(nil, ils_in),
nil > 0 ? ils_in[0]->code_size : 0)
{
FAISS_THROW_IF_NOT (nil > 0);
cumsz.resize (nil + 1);
for (int i = 0; i < nil; i++) {
ils.push_back (ils_in[i]);
FAISS_THROW_IF_NOT (ils_in[i]->code_size == code_size);
cumsz[i + 1] = cumsz[i] + ils_in[i]->nlist;
}
}
size_t VStackInvertedLists::list_size(size_t list_no) const
{
int i = translate_list_no (this, list_no);
list_no -= cumsz[i];
return ils[i]->list_size (list_no);
}
const uint8_t * VStackInvertedLists::get_codes (size_t list_no) const
{
int i = translate_list_no (this, list_no);
list_no -= cumsz[i];
return ils[i]->get_codes (list_no);
}
const uint8_t * VStackInvertedLists::get_single_code (
size_t list_no, size_t offset) const
{
int i = translate_list_no (this, list_no);
list_no -= cumsz[i];
return ils[i]->get_single_code (list_no, offset);
}
void VStackInvertedLists::release_codes (
size_t list_no, const uint8_t *codes) const {
int i = translate_list_no (this, list_no);
list_no -= cumsz[i];
return ils[i]->release_codes (list_no, codes);
}
const Index::idx_t * VStackInvertedLists::get_ids (size_t list_no) const
{
int i = translate_list_no (this, list_no);
list_no -= cumsz[i];
return ils[i]->get_ids (list_no);
}
Index::idx_t VStackInvertedLists::get_single_id (
size_t list_no, size_t offset) const
{
int i = translate_list_no (this, list_no);
list_no -= cumsz[i];
return ils[i]->get_single_id (list_no, offset);
}
void VStackInvertedLists::release_ids (size_t list_no, const idx_t *ids) const {
int i = translate_list_no (this, list_no);
list_no -= cumsz[i];
return ils[i]->release_ids (list_no, ids);
}
void VStackInvertedLists::prefetch_lists (
const idx_t *list_nos, int nlist) const
{
std::vector<int> ilno (nlist, -1);
std::vector<int> n_per_il (ils.size(), 0);
for (int j = 0; j < nlist; j++) {
idx_t list_no = list_nos[j];
if (list_no < 0) continue;
int i = ilno[j] = translate_list_no (this, list_no);
n_per_il[i]++;
}
std::vector<int> cum_n_per_il (ils.size() + 1, 0);
for (int j = 0; j < ils.size(); j++) {
cum_n_per_il[j + 1] = cum_n_per_il[j] + n_per_il[j];
}
std::vector<idx_t> sorted_list_nos (cum_n_per_il.back());
for (int j = 0; j < nlist; j++) {
idx_t list_no = list_nos[j];
if (list_no < 0) continue;
int i = ilno[j];
list_no -= cumsz[i];
sorted_list_nos[cum_n_per_il[i]++] = list_no;
}
int i0 = 0;
for (int j = 0; j < ils.size(); j++) {
int i1 = i0 + n_per_il[j];
if (i1 > i0) {
ils[j]->prefetch_lists (sorted_list_nos.data() + i0,
i1 - i0);
}
i0 = i1;
}
}
/*****************************************
* MaskedInvertedLists implementation
******************************************/
MaskedInvertedLists::MaskedInvertedLists (const InvertedLists *il0,
const InvertedLists *il1):
ReadOnlyInvertedLists (il0->nlist, il0->code_size),
il0 (il0), il1 (il1)
{
FAISS_THROW_IF_NOT (il1->nlist == nlist);
FAISS_THROW_IF_NOT (il1->code_size == code_size);
}
size_t MaskedInvertedLists::list_size(size_t list_no) const
{
size_t sz = il0->list_size(list_no);
return sz ? sz : il1->list_size(list_no);
}
const uint8_t * MaskedInvertedLists::get_codes (size_t list_no) const
{
size_t sz = il0->list_size(list_no);
return (sz ? il0 : il1)->get_codes(list_no);
}
const idx_t * MaskedInvertedLists::get_ids (size_t list_no) const
{
size_t sz = il0->list_size (list_no);
return (sz ? il0 : il1)->get_ids (list_no);
}
void MaskedInvertedLists::release_codes (
size_t list_no, const uint8_t *codes) const
{
size_t sz = il0->list_size (list_no);
(sz ? il0 : il1)->release_codes (list_no, codes);
}
void MaskedInvertedLists::release_ids (size_t list_no, const idx_t *ids) const
{
size_t sz = il0->list_size (list_no);
(sz ? il0 : il1)->release_ids (list_no, ids);
}
idx_t MaskedInvertedLists::get_single_id (size_t list_no, size_t offset) const
{
size_t sz = il0->list_size (list_no);
return (sz ? il0 : il1)->get_single_id (list_no, offset);
}
const uint8_t * MaskedInvertedLists::get_single_code (
size_t list_no, size_t offset) const
{
size_t sz = il0->list_size (list_no);
return (sz ? il0 : il1)->get_single_code (list_no, offset);
}
void MaskedInvertedLists::prefetch_lists (
const idx_t *list_nos, int nlist) const
{
std::vector<idx_t> list0, list1;
for (int i = 0; i < nlist; i++) {
idx_t list_no = list_nos[i];
if (list_no < 0) continue;
size_t sz = il0->list_size(list_no);
(sz ? list0 : list1).push_back (list_no);
}
il0->prefetch_lists (list0.data(), list0.size());
il1->prefetch_lists (list1.data(), list1.size());
}
+129 -31
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -58,10 +57,10 @@ struct InvertedLists {
virtual const idx_t * get_ids (size_t list_no) const = 0;
/// release codes returned by get_codes (default implementation is nop
virtual void release_codes (const uint8_t *codes) const;
virtual void release_codes (size_t list_no, const uint8_t *codes) const;
/// release ids returned by get_ids
virtual void release_ids (const idx_t *ids) const;
virtual void release_ids (size_t list_no, const idx_t *ids) const;
/// @return a single id in an inverted list
virtual idx_t get_single_id (size_t list_no, size_t offset) const;
@@ -73,7 +72,7 @@ struct InvertedLists {
/// prepare the following lists (default does nothing)
/// a list can be -1 hence the signed long
virtual void prefetch_lists (const long *list_nos, int nlist) const;
virtual void prefetch_lists (const idx_t *list_nos, int nlist) const;
/*************************
* writing functions */
@@ -110,6 +109,8 @@ struct InvertedLists {
/// display some stats about the inverted lists
void print_stats () const;
/// sum up list sizes
size_t compute_ntotal () const;
/**************************************
* Scoped inverted lists (for automatic deallocation)
@@ -118,7 +119,7 @@ struct InvertedLists {
*
* uint8_t * codes = invlists->get_codes (10);
* ... use codes
* invlists->release_codes(codes)
* invlists->release_codes(10, codes)
*
* write:
*
@@ -135,9 +136,10 @@ struct InvertedLists {
struct ScopedIds {
const InvertedLists *il;
const idx_t *ids;
size_t list_no;
ScopedIds (const InvertedLists *il, size_t list_no):
il (il), ids (il->get_ids (list_no))
il (il), ids (il->get_ids (list_no)), list_no (list_no)
{}
const idx_t *get() {return ids; }
@@ -147,26 +149,28 @@ struct InvertedLists {
}
~ScopedIds () {
il->release_ids (ids);
il->release_ids (list_no, ids);
}
};
struct ScopedCodes {
const InvertedLists *il;
const uint8_t *codes;
size_t list_no;
ScopedCodes (const InvertedLists *il, size_t list_no):
il (il), codes (il->get_codes (list_no))
il (il), codes (il->get_codes (list_no)), list_no (list_no)
{}
ScopedCodes (const InvertedLists *il, size_t list_no, size_t offset):
il (il), codes (il->get_single_code (list_no, offset))
il (il), codes (il->get_single_code (list_no, offset)),
list_no (list_no)
{}
const uint8_t *get() {return codes; }
~ScopedCodes () {
il->release_codes (codes);
il->release_codes (list_no, codes);
}
};
@@ -197,27 +201,17 @@ struct ArrayInvertedLists: InvertedLists {
virtual ~ArrayInvertedLists ();
};
/*****************************************************************
* Meta-inverted lists
*
* About terminology: the inverted lists are seen as a sparse matrix,
* that can be stacked horizontally, vertically and sliced.
*****************************************************************/
/// inverted lists built as the concatenation of a set of invlists
/// (read-only)
struct ConcatenatedInvertedLists: InvertedLists {
struct ReadOnlyInvertedLists: InvertedLists {
std::vector<const InvertedLists *>ils;
/// build InvertedLists by concatenating nil of them
ConcatenatedInvertedLists (int nil, const InvertedLists **ils);
size_t list_size(size_t list_no) const override;
const uint8_t * get_codes (size_t list_no) const override;
const idx_t * get_ids (size_t list_no) const override;
void release_codes (const uint8_t *codes) const override;
void release_ids (const idx_t *ids) const override;
idx_t get_single_id (size_t list_no, size_t offset) const override;
const uint8_t * get_single_code (
size_t list_no, size_t offset) const override;
ReadOnlyInvertedLists (size_t nlist, size_t code_size):
InvertedLists (nlist, code_size) {}
size_t add_entries (
size_t list_no, size_t n_entry,
@@ -230,6 +224,110 @@ struct ConcatenatedInvertedLists: InvertedLists {
};
/// Horizontal stack of inverted lists
struct HStackInvertedLists: ReadOnlyInvertedLists {
std::vector<const InvertedLists *>ils;
/// build InvertedLists by concatenating nil of them
HStackInvertedLists (int nil, const InvertedLists **ils);
size_t list_size(size_t list_no) const override;
const uint8_t * get_codes (size_t list_no) const override;
const idx_t * get_ids (size_t list_no) const override;
void prefetch_lists (const idx_t *list_nos, int nlist) const override;
void release_codes (size_t list_no, const uint8_t *codes) const override;
void release_ids (size_t list_no, const idx_t *ids) const override;
idx_t get_single_id (size_t list_no, size_t offset) const override;
const uint8_t * get_single_code (
size_t list_no, size_t offset) const override;
};
using ConcatenatedInvertedLists = HStackInvertedLists;
/// vertical slice of indexes in another InvertedLists
struct SliceInvertedLists: ReadOnlyInvertedLists {
const InvertedLists *il;
idx_t i0, i1;
SliceInvertedLists(const InvertedLists *il, idx_t i0, idx_t i1);
size_t list_size(size_t list_no) const override;
const uint8_t * get_codes (size_t list_no) const override;
const idx_t * get_ids (size_t list_no) const override;
void release_codes (size_t list_no, const uint8_t *codes) const override;
void release_ids (size_t list_no, const idx_t *ids) const override;
idx_t get_single_id (size_t list_no, size_t offset) const override;
const uint8_t * get_single_code (
size_t list_no, size_t offset) const override;
void prefetch_lists (const idx_t *list_nos, int nlist) const override;
};
struct VStackInvertedLists: ReadOnlyInvertedLists {
std::vector<const InvertedLists *>ils;
std::vector<idx_t> cumsz;
/// build InvertedLists by concatenating nil of them
VStackInvertedLists (int nil, const InvertedLists **ils);
size_t list_size(size_t list_no) const override;
const uint8_t * get_codes (size_t list_no) const override;
const idx_t * get_ids (size_t list_no) const override;
void release_codes (size_t list_no, const uint8_t *codes) const override;
void release_ids (size_t list_no, const idx_t *ids) const override;
idx_t get_single_id (size_t list_no, size_t offset) const override;
const uint8_t * get_single_code (
size_t list_no, size_t offset) const override;
void prefetch_lists (const idx_t *list_nos, int nlist) const override;
};
/** use the first inverted lists if they are non-empty otherwise use the second
*
* This is useful if il1 has a few inverted lists that are too long,
* and that il0 has replacement lists for those, with empty lists for
* the others. */
struct MaskedInvertedLists: ReadOnlyInvertedLists {
const InvertedLists *il0;
const InvertedLists *il1;
MaskedInvertedLists (const InvertedLists *il0,
const InvertedLists *il1);
size_t list_size(size_t list_no) const override;
const uint8_t * get_codes (size_t list_no) const override;
const idx_t * get_ids (size_t list_no) const override;
void release_codes (size_t list_no, const uint8_t *codes) const override;
void release_ids (size_t list_no, const idx_t *ids) const override;
idx_t get_single_id (size_t list_no, size_t offset) const override;
const uint8_t * get_single_code (
size_t list_no, size_t offset) const override;
void prefetch_lists (const idx_t *list_nos, int nlist) const override;
};
} // namespace faiss
+17 -27
View File
@@ -1,31 +1,21 @@
MIT License
BSD License
Copyright (c) Facebook, Inc. and its affiliates.
For Faiss software
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
Copyright (c) 2016-present, Facebook, Inc. All rights reserved.
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
Redistribution and use in source and binary forms, with or without modification,
are permitted provided that the following conditions are met:
* Redistributions of source code must retain the above copyright notice, this
list of conditions and the following disclaimer.
* Redistributions in binary form must reproduce the above copyright notice,
this list of conditions and the following disclaimer in the documentation
and/or other materials provided with the distribution.
* Neither the name Facebook nor the names of its contributors may be used to
endorse or promote products derived from this software without specific
prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND
ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED
WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE FOR
ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES
(INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES;
LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON
ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS
SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
+3 -4
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
-include makefile.inc
@@ -109,4 +108,4 @@ misc/test_blas: misc/test_blas.cpp
$(CXX) $(CPPFLAGS) $(CXXFLAGS) $(LDFLAGS) -o $@ $^ $(LIBS)
.PHONY: all clean demos install installdirs py test gpu_test uninstall
.PHONY: all clean demos install installdirs py test test_gpu uninstall
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+3 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -15,6 +14,7 @@
#include <unordered_map>
#include "Index.h"
#include "IndexShards.h"
#include "IndexReplicas.h"
namespace faiss {
+122 -54
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -19,6 +18,7 @@
#include <sys/types.h>
#include "FaissAssert.h"
#include "utils.h"
namespace faiss {
@@ -144,12 +144,41 @@ struct OnDiskInvertedLists::OngoingPrefetch {
struct Thread {
pthread_t pth;
const OnDiskInvertedLists *od;
int64_t list_no;
OngoingPrefetch *pf;
bool one_list () {
idx_t list_no = pf->get_next_list();
if(list_no == -1) return false;
const OnDiskInvertedLists *od = pf->od;
od->locks->lock_1 (list_no);
size_t n = od->list_size (list_no);
const Index::idx_t *idx = od->get_ids (list_no);
const uint8_t *codes = od->get_codes (list_no);
int cs = 0;
for (size_t i = 0; i < n;i++) {
cs += idx[i];
}
const idx_t *codes8 = (const idx_t*)codes;
idx_t n8 = n * od->code_size / 8;
for (size_t i = 0; i < n8;i++) {
cs += codes8[i];
}
od->locks->unlock_1(list_no);
global_cs += cs & 1;
return true;
}
};
std::vector<Thread> threads;
pthread_mutex_t list_ids_mutex;
std::vector<idx_t> list_ids;
int cur_list;
// mutex for the list of tasks
pthread_mutex_t mutex;
// pretext to avoid code below to be optimized out
@@ -157,50 +186,57 @@ struct OnDiskInvertedLists::OngoingPrefetch {
const OnDiskInvertedLists *od;
OngoingPrefetch (const OnDiskInvertedLists *od): od (od)
explicit OngoingPrefetch (const OnDiskInvertedLists *od): od (od)
{
pthread_mutex_init (&mutex, nullptr);
pthread_mutex_init (&list_ids_mutex, nullptr);
cur_list = 0;
}
static void* prefetch_list (void * arg) {
Thread *th = static_cast<Thread*>(arg);
th->od->locks->lock_1(th->list_no);
size_t n = th->od->list_size(th->list_no);
const Index::idx_t *idx = th->od->get_ids(th->list_no);
const uint8_t *codes = th->od->get_codes(th->list_no);
int cs = 0;
for (size_t i = 0; i < n;i++) {
cs += idx[i];
}
const long *codes8 = (const long*)codes;
long n8 = n * th->od->code_size / 8;
while (th->one_list()) ;
for (size_t i = 0; i < n8;i++) {
cs += codes8[i];
}
th->od->locks->unlock_1(th->list_no);
global_cs += cs & 1;
return nullptr;
}
void prefetch_lists (const long *list_nos, int n) {
pthread_mutex_lock (&mutex);
for (auto &th: threads) {
if (th.list_no != -1) {
pthread_join (th.pth, nullptr);
}
idx_t get_next_list () {
idx_t list_no = -1;
pthread_mutex_lock (&list_ids_mutex);
if (cur_list >= 0 && cur_list < list_ids.size()) {
list_no = list_ids[cur_list++];
}
threads.resize (n);
for (int i = 0; i < n; i++) {
long list_no = list_nos[i];
Thread & th = threads[i];
if (list_no >= 0 && od->list_size(list_no) > 0) {
th.list_no = list_no;
th.od = od;
pthread_mutex_unlock (&list_ids_mutex);
return list_no;
}
void prefetch_lists (const idx_t *list_nos, int n) {
pthread_mutex_lock (&mutex);
pthread_mutex_lock (&list_ids_mutex);
list_ids.clear ();
pthread_mutex_unlock (&list_ids_mutex);
for (auto &th: threads) {
pthread_join (th.pth, nullptr);
}
threads.resize (0);
cur_list = 0;
int nt = std::min (n, od->prefetch_nthread);
if (nt > 0) {
// prepare tasks
for (int i = 0; i < n; i++) {
idx_t list_no = list_nos[i];
if (list_no >= 0 && od->list_size(list_no) > 0) {
list_ids.push_back (list_no);
}
}
// prepare threads
threads.resize (nt);
for (Thread &th: threads) {
th.pf = this;
pthread_create (&th.pth, nullptr, prefetch_list, &th);
} else {
th.list_no = -1;
}
}
pthread_mutex_unlock (&mutex);
@@ -209,12 +245,11 @@ struct OnDiskInvertedLists::OngoingPrefetch {
~OngoingPrefetch () {
pthread_mutex_lock (&mutex);
for (auto &th: threads) {
if (th.list_no != -1) {
pthread_join (th.pth, nullptr);
}
pthread_join (th.pth, nullptr);
}
pthread_mutex_unlock (&mutex);
pthread_mutex_destroy (&mutex);
pthread_mutex_destroy (&list_ids_mutex);
}
};
@@ -222,7 +257,7 @@ struct OnDiskInvertedLists::OngoingPrefetch {
int OnDiskInvertedLists::OngoingPrefetch::global_cs = 0;
void OnDiskInvertedLists::prefetch_lists (const long *list_nos, int n) const
void OnDiskInvertedLists::prefetch_lists (const idx_t *list_nos, int n) const
{
pf->prefetch_lists (list_nos, n);
}
@@ -260,7 +295,7 @@ void OnDiskInvertedLists::update_totsize (size_t new_size)
// unmap file
if (ptr != nullptr) {
int err = munmap (ptr, totsize);
FAISS_THROW_IF_NOT_FMT (err == 0, "mumap error: %s",
FAISS_THROW_IF_NOT_FMT (err == 0, "munmap error: %s",
strerror(errno));
}
if (totsize == 0) {
@@ -329,7 +364,8 @@ OnDiskInvertedLists::OnDiskInvertedLists (
ptr (nullptr),
read_only (false),
locks (new LockLevels ()),
pf (new OngoingPrefetch (this))
pf (new OngoingPrefetch (this)),
prefetch_nthread (32)
{
lists.resize (nlist);
@@ -337,12 +373,7 @@ OnDiskInvertedLists::OnDiskInvertedLists (
}
OnDiskInvertedLists::OnDiskInvertedLists ():
InvertedLists (0, 0),
totsize (0),
ptr (nullptr),
read_only (false),
locks (new LockLevels ()),
pf (new OngoingPrefetch (this))
OnDiskInvertedLists (0, 0, "")
{
}
@@ -353,9 +384,10 @@ OnDiskInvertedLists::~OnDiskInvertedLists ()
// unmap all lists
if (ptr != nullptr) {
int err = munmap (ptr, totsize);
FAISS_THROW_IF_NOT_FMT (err == 0,
"mumap error: %s",
strerror(errno));
if (err != 0) {
fprintf(stderr, "mumap error: %s",
strerror(errno));
}
}
delete locks;
}
@@ -559,7 +591,8 @@ void OnDiskInvertedLists::free_slot (size_t offset, size_t capacity) {
* Compact form
*****************************************/
size_t OnDiskInvertedLists::merge_from (const InvertedLists **ils, int n_il)
size_t OnDiskInvertedLists::merge_from (const InvertedLists **ils, int n_il,
bool verbose)
{
FAISS_THROW_IF_NOT_MSG (totsize == 0, "works only on an empty InvertedLists");
@@ -585,6 +618,10 @@ size_t OnDiskInvertedLists::merge_from (const InvertedLists **ils, int n_il)
update_totsize (cums);
size_t nmerged = 0;
double t0 = getmillisecs(), last_t = t0;
#pragma omp parallel for
for (size_t j = 0; j < nlist; j++) {
List & l = lists[j];
@@ -593,14 +630,45 @@ size_t OnDiskInvertedLists::merge_from (const InvertedLists **ils, int n_il)
size_t n_entry = il->list_size(j);
l.size += n_entry;
update_entries (j, l.size - n_entry, n_entry,
il->get_ids(j),
il->get_codes(j));
ScopedIds(il, j).get(),
ScopedCodes(il, j).get());
}
assert (l.size == l.capacity);
if (verbose) {
#pragma omp critical
{
nmerged++;
double t1 = getmillisecs();
if (t1 - last_t > 500) {
printf("merged %ld lists in %.3f s\r",
nmerged, (t1 - t0) / 1000.0);
fflush(stdout);
last_t = t1;
}
}
}
}
if(verbose) {
printf("\n");
}
return ntotal;
}
void OnDiskInvertedLists::crop_invlists(size_t l0, size_t l1)
{
FAISS_THROW_IF_NOT(0 <= l0 && l0 <= l1 && l1 <= nlist);
std::vector<List> new_lists (l1 - l0);
memcpy (new_lists.data(), &lists[l0], (l1 - l0) * sizeof(List));
lists.swap(new_lists);
nlist = l1 - l0;
}
} // namespace faiss
+10 -5
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -58,6 +57,7 @@ struct OnDiskInvertedLists: InvertedLists {
List ();
};
// size nlist
std::vector<List> lists;
struct Slot {
@@ -67,6 +67,7 @@ struct OnDiskInvertedLists: InvertedLists {
Slot ();
};
// size whatever space remains
std::list<Slot> slots;
std::string filename;
@@ -92,9 +93,12 @@ struct OnDiskInvertedLists: InvertedLists {
// copy all inverted lists into *this, in compact form (without
// allocating slots)
size_t merge_from (const InvertedLists **ils, int n_il);
size_t merge_from (const InvertedLists **ils, int n_il, bool verbose=false);
void prefetch_lists (const long *list_nos, int nlist) const override;
/// restrict the inverted lists to l0:l1 without touching the mmapped region
void crop_invlists(size_t l0, size_t l1);
void prefetch_lists (const idx_t *list_nos, int nlist) const override;
virtual ~OnDiskInvertedLists ();
@@ -105,6 +109,7 @@ struct OnDiskInvertedLists: InvertedLists {
// encapsulates the threads that are busy prefeteching
struct OngoingPrefetch;
OngoingPrefetch *pf;
int prefetch_nthread;
void do_mmap ();
void update_totsize (size_t new_totsize);
-33
View File
@@ -1,33 +0,0 @@
Additional Grant of Patent Rights Version 2
"Software" means the Faiss software distributed by Facebook, Inc.
Facebook, Inc. ("Facebook") hereby grants to each recipient of the Software
("you") a perpetual, worldwide, royalty-free, non-exclusive, irrevocable
(subject to the termination provision below) license under any Necessary
Claims, to make, have made, use, sell, offer to sell, import, and otherwise
transfer the Software. For avoidance of doubt, no license is granted under
Facebook’s rights in any patent claims that are infringed by (i) modifications
to the Software made by you or any third party or (ii) the Software in
combination with any software or other technology.
The license granted hereunder will terminate, automatically and without notice,
if you (or any of your subsidiaries, corporate affiliates or agents) initiate
directly or indirectly, or take a direct financial interest in, any Patent
Assertion: (i) against Facebook or any of its subsidiaries or corporate
affiliates, (ii) against any party if such Patent Assertion arises in whole or
in part from any software, technology, product or service of Facebook or any of
its subsidiaries or corporate affiliates, or (iii) against any party relating
to the Software. Notwithstanding the foregoing, if Facebook or any of its
subsidiaries or corporate affiliates files a lawsuit alleging patent
infringement against you in the first instance, and you respond by filing a
patent infringement counterclaim in that lawsuit against that party that is
unrelated to the Software, the license granted hereunder will not terminate
under section (i) of this paragraph due to such counterclaim.
A "Necessary Claim" is a claim of a patent owned by Facebook that is
necessarily infringed by the Software standing alone.
A "Patent Assertion" is any lawsuit or other action alleging direct, indirect,
or contributory infringement or inducement to infringe any patent, including a
cross-claim or counterclaim.
+3 -4
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -837,7 +836,7 @@ void PolysemousTraining::optimize_ranking (
pq.compute_codes (x, all_codes.data(), n);
FAISS_THROW_IF_NOT (pq.byte_per_idx == 1);
FAISS_THROW_IF_NOT (pq.nbits == 8);
if (n == 0)
pq.compute_sdc_table ();
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+290 -129
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -14,6 +13,7 @@
#include <cstddef>
#include <cstring>
#include <cstdio>
#include <memory>
#include <algorithm>
@@ -97,7 +97,7 @@ void pq_estimators_from_tables_M4 (const CT * codes,
template <typename CT, class C>
static inline void pq_estimators_from_tables (const ProductQuantizer * pq,
static inline void pq_estimators_from_tables (const ProductQuantizer& pq,
const CT * codes,
size_t ncodes,
const float * dis_table,
@@ -106,24 +106,24 @@ static inline void pq_estimators_from_tables (const ProductQuantizer * pq,
long * heap_ids)
{
if (pq->M == 4) {
if (pq.M == 4) {
pq_estimators_from_tables_M4<CT, C> (codes, ncodes,
dis_table, pq->ksub, k,
heap_dis, heap_ids);
dis_table, pq.ksub, k,
heap_dis, heap_ids);
return;
}
if (pq->M % 4 == 0) {
pq_estimators_from_tables_Mmul4<CT, C> (pq->M, codes, ncodes,
dis_table, pq->ksub, k,
heap_dis, heap_ids);
if (pq.M % 4 == 0) {
pq_estimators_from_tables_Mmul4<CT, C> (pq.M, codes, ncodes,
dis_table, pq.ksub, k,
heap_dis, heap_ids);
return;
}
/* Default is relatively slow */
const size_t M = pq->M;
const size_t ksub = pq->ksub;
const size_t M = pq.M;
const size_t ksub = pq.ksub;
for (size_t j = 0; j < ncodes; j++) {
float dis = 0;
const float * __restrict dt = dis_table;
@@ -138,6 +138,36 @@ static inline void pq_estimators_from_tables (const ProductQuantizer * pq,
}
}
template <class C>
static inline void pq_estimators_from_tables_generic(const ProductQuantizer& pq,
size_t nbits,
const uint8_t *codes,
size_t ncodes,
const float *dis_table,
size_t k,
float *heap_dis,
long *heap_ids)
{
const size_t M = pq.M;
const size_t ksub = pq.ksub;
for (size_t j = 0; j < ncodes; ++j) {
faiss::ProductQuantizer::PQDecoderGeneric decoder(
codes + j * pq.code_size, nbits
);
float dis = 0;
const float * __restrict dt = dis_table;
for (size_t m = 0; m < M; m++) {
uint64_t c = decoder.decode();
dis += dt[c];
dt += ksub;
}
if (C::cmp(heap_dis[0], dis)) {
heap_pop<C>(k, heap_dis, heap_ids);
heap_push<C>(k, heap_dis, heap_ids, dis, j);
}
}
}
/*********************************************
* PQ implementation
@@ -151,27 +181,20 @@ ProductQuantizer::ProductQuantizer (size_t d, size_t M, size_t nbits):
set_derived_values ();
}
ProductQuantizer::ProductQuantizer ():
d(0), M(1), nbits(0), assign_index(nullptr)
{
set_derived_values ();
}
ProductQuantizer::ProductQuantizer ()
: ProductQuantizer(0, 1, 0) {}
void ProductQuantizer::set_derived_values () {
// quite a few derived values
FAISS_THROW_IF_NOT (d % M == 0);
dsub = d / M;
byte_per_idx = (nbits + 7) / 8;
code_size = byte_per_idx * M;
code_size = (nbits * M + 7) / 8;
ksub = 1 << nbits;
centroids.resize (d * ksub);
verbose = false;
train_type = Train_default;
}
void ProductQuantizer::set_params (const float * centroids_, int m)
{
memcpy (get_centroids(m, 0), centroids_,
@@ -304,48 +327,71 @@ void ProductQuantizer::train (int n, const float * x)
}
}
template<class PQEncoder>
void compute_code(const ProductQuantizer& pq, const float *x, uint8_t *code) {
float distances [pq.ksub];
PQEncoder encoder(code, pq.nbits);
for (size_t m = 0; m < pq.M; m++) {
float mindis = 1e20;
uint64_t idxm = 0;
const float * xsub = x + m * pq.dsub;
void ProductQuantizer::compute_code (const float * x, uint8_t * code) const
{
float distances [ksub];
for (size_t m = 0; m < M; m++) {
float mindis = 1e20;
int idxm = -1;
const float * xsub = x + m * dsub;
fvec_L2sqr_ny(distances, xsub, pq.get_centroids(m, 0), pq.dsub, pq.ksub);
fvec_L2sqr_ny (distances, xsub, get_centroids(m, 0), dsub, ksub);
/* Find best centroid */
size_t i;
for (i = 0; i < ksub; i++) {
float dis = distances [i];
if (dis < mindis) {
mindis = dis;
idxm = i;
}
}
switch (byte_per_idx) {
case 1: code[m] = (uint8_t) idxm; break;
case 2: ((uint16_t *) code)[m] = (uint16_t) idxm; break;
}
/* Find best centroid */
for (size_t i = 0; i < pq.ksub; i++) {
float dis = distances[i];
if (dis < mindis) {
mindis = dis;
idxm = i;
}
}
encoder.encode(idxm);
}
}
void ProductQuantizer::compute_code(const float * x, uint8_t * code) const {
switch (nbits) {
case 8:
faiss::compute_code<PQEncoder8>(*this, x, code);
break;
case 16:
faiss::compute_code<PQEncoder16>(*this, x, code);
break;
default:
faiss::compute_code<PQEncoderGeneric>(*this, x, code);
break;
}
}
template<class PQDecoder>
void decode(const ProductQuantizer& pq, const uint8_t *code, float *x)
{
PQDecoder decoder(code, pq.nbits);
for (size_t m = 0; m < pq.M; m++) {
uint64_t c = decoder.decode();
memcpy(x + m * pq.dsub, pq.get_centroids(m, c), sizeof(float) * pq.dsub);
}
}
void ProductQuantizer::decode (const uint8_t *code, float *x) const
{
if (byte_per_idx == 1) {
for (size_t m = 0; m < M; m++) {
memcpy (x + m * dsub, get_centroids(m, code[m]),
sizeof(float) * dsub);
}
} else {
const uint16_t *c = (const uint16_t*) code;
for (size_t m = 0; m < M; m++) {
memcpy (x + m * dsub, get_centroids(m, c[m]),
sizeof(float) * dsub);
}
}
switch (nbits) {
case 8:
faiss::decode<PQDecoder8>(*this, code, x);
break;
case 16:
faiss::decode<PQDecoder16>(*this, code, x);
break;
default:
faiss::decode<PQDecoderGeneric>(*this, code, x);
break;
}
}
@@ -360,23 +406,22 @@ void ProductQuantizer::decode (const uint8_t *code, float *x, size_t n) const
void ProductQuantizer::compute_code_from_distance_table (const float *tab,
uint8_t *code) const
{
for (size_t m = 0; m < M; m++) {
float mindis = 1e20;
int idxm = -1;
PQEncoderGeneric encoder(code, nbits);
for (size_t m = 0; m < M; m++) {
float mindis = 1e20;
uint64_t idxm = 0;
/* Find best centroid */
for (size_t j = 0; j < ksub; j++) {
float dis = *tab++;
if (dis < mindis) {
mindis = dis;
idxm = j;
}
}
switch (byte_per_idx) {
case 1: code[m] = (uint8_t) idxm; break;
case 2: ((uint16_t *) code)[m] = (uint16_t) idxm; break;
}
/* Find best centroid */
for (size_t j = 0; j < ksub; j++) {
float dis = *tab++;
if (dis < mindis) {
mindis = dis;
idxm = j;
}
}
encoder.encode(idxm);
}
}
void ProductQuantizer::compute_codes_with_assign_index (
@@ -406,25 +451,27 @@ void ProductQuantizer::compute_codes_with_assign_index (
assign_index->assign (i1 - i0, xslice, assign);
switch (byte_per_idx) {
case 1:
{
uint8_t *c = codes + code_size * i0 + m;
for (size_t i = i0; i < i1; i++) {
*c = assign[i - i0];
c += M;
}
}
break;
case 2:
{
uint16_t *c = (uint16_t*)(codes + code_size * i0 + m * 2);
for (size_t i = i0; i < i1; i++) {
*c = assign[i - i0];
c += M;
}
}
break;
if (nbits == 8) {
uint8_t *c = codes + code_size * i0 + m;
for (size_t i = i0; i < i1; i++) {
*c = assign[i - i0];
c += M;
}
} else if (nbits == 16) {
uint16_t *c = (uint16_t*)(codes + code_size * i0 + m * 2);
for (size_t i = i0; i < i1; i++) {
*c = assign[i - i0];
c += M;
}
} else {
for (size_t i = i0; i < i1; ++i) {
uint8_t *c = codes + code_size * i + ((m * nbits) / 8);
uint8_t offset = (m * nbits) % 8;
uint64_t ass = assign[i - i0];
PQEncoderGeneric encoder(c, nbits, offset);
encoder.encode(ass);
}
}
}
@@ -436,8 +483,7 @@ void ProductQuantizer::compute_codes (const float * x,
uint8_t * codes,
size_t n) const
{
// process by blocks to avoid using too much RAM
// process by blocks to avoid using too much RAM
size_t bs = 256 * 1024;
if (n > bs) {
for (size_t i0 = 0; i0 < n; i0 += bs) {
@@ -553,9 +599,10 @@ void ProductQuantizer::compute_inner_prod_tables (
}
}
template <typename CT, class C>
template <class C>
static void pq_knn_search_with_tables (
const ProductQuantizer * pq,
const ProductQuantizer& pq,
size_t nbits,
const float *dis_tables,
const uint8_t * codes,
const size_t ncodes,
@@ -563,7 +610,7 @@ static void pq_knn_search_with_tables (
bool init_finalize_heap)
{
size_t k = res->k, nx = res->nh;
size_t ksub = pq->ksub, M = pq->M;
size_t ksub = pq.ksub, M = pq.M;
#pragma omp parallel for
@@ -579,10 +626,30 @@ static void pq_knn_search_with_tables (
heap_heapify<C> (k, heap_dis, heap_ids);
}
pq_estimators_from_tables<CT, C> (pq,
(CT*)codes, ncodes,
dis_table,
k, heap_dis, heap_ids);
switch (nbits) {
case 8:
pq_estimators_from_tables<uint8_t, C> (pq,
codes, ncodes,
dis_table,
k, heap_dis, heap_ids);
break;
case 16:
pq_estimators_from_tables<uint16_t, C> (pq,
(uint16_t*)codes, ncodes,
dis_table,
k, heap_dis, heap_ids);
break;
default:
pq_estimators_from_tables_generic<C> (pq,
nbits,
codes, ncodes,
dis_table,
k, heap_dis, heap_ids);
break;
}
if (init_finalize_heap) {
heap_reorder<C> (k, heap_dis, heap_ids);
}
@@ -597,21 +664,11 @@ void ProductQuantizer::search (const float * __restrict x,
bool init_finalize_heap) const
{
FAISS_THROW_IF_NOT (nx == res->nh);
float * dis_tables = new float [nx * ksub * M];
ScopeDeleter<float> del(dis_tables);
compute_distance_tables (nx, x, dis_tables);
if (byte_per_idx == 1) {
pq_knn_search_with_tables<uint8_t, CMax<float, long> > (
this, dis_tables, codes, ncodes, res, init_finalize_heap);
} else if (byte_per_idx == 2) {
pq_knn_search_with_tables<uint16_t, CMax<float, long> > (
this, dis_tables, codes, ncodes, res, init_finalize_heap);
}
std::unique_ptr<float[]> dis_tables(new float [nx * ksub * M]);
compute_distance_tables (nx, x, dis_tables.get());
pq_knn_search_with_tables<CMax<float, long>> (
*this, nbits, dis_tables.get(), codes, ncodes, res, init_finalize_heap);
}
void ProductQuantizer::search_ip (const float * __restrict x,
@@ -622,20 +679,11 @@ void ProductQuantizer::search_ip (const float * __restrict x,
bool init_finalize_heap) const
{
FAISS_THROW_IF_NOT (nx == res->nh);
float * dis_tables = new float [nx * ksub * M];
ScopeDeleter<float> del(dis_tables);
compute_inner_prod_tables (nx, x, dis_tables);
if (byte_per_idx == 1) {
pq_knn_search_with_tables<uint8_t, CMin<float, long> > (
this, dis_tables, codes, ncodes, res, init_finalize_heap);
} else if (byte_per_idx == 2) {
pq_knn_search_with_tables<uint16_t, CMin<float, long> > (
this, dis_tables, codes, ncodes, res, init_finalize_heap);
}
std::unique_ptr<float[]> dis_tables(new float [nx * ksub * M]);
compute_inner_prod_tables (nx, x, dis_tables.get());
pq_knn_search_with_tables<CMin<float, long> > (
*this, nbits, dis_tables.get(), codes, ncodes, res, init_finalize_heap);
}
@@ -675,7 +723,7 @@ void ProductQuantizer::search_sdc (const uint8_t * qcodes,
bool init_finalize_heap) const
{
FAISS_THROW_IF_NOT (sdc_table.size() == M * ksub * ksub);
FAISS_THROW_IF_NOT (byte_per_idx == 1);
FAISS_THROW_IF_NOT (nbits == 8);
size_t k = res->k;
@@ -712,4 +760,117 @@ void ProductQuantizer::search_sdc (const uint8_t * qcodes,
}
} // namespace faiss
ProductQuantizer::PQEncoderGeneric::PQEncoderGeneric(uint8_t *code, int nbits,
uint8_t offset)
: code(code), offset(offset), nbits(nbits), reg(0) {
assert(nbits <= 64);
if (offset > 0) {
reg = (*code & ((1 << offset) - 1));
}
}
void ProductQuantizer::PQEncoderGeneric::encode(uint64_t x) {
reg |= (uint8_t)(x << offset);
x >>= (8 - offset);
if (offset + nbits >= 8) {
*code++ = reg;
for (int i = 0; i < (nbits - (8 - offset)) / 8; ++i) {
*code++ = (uint8_t)x;
x >>= 8;
}
offset += nbits;
offset &= 7;
reg = (uint8_t)x;
} else {
offset += nbits;
}
}
ProductQuantizer::PQEncoderGeneric::~PQEncoderGeneric() {
if (offset > 0) {
*code = reg;
}
}
ProductQuantizer::PQEncoder8::PQEncoder8(uint8_t *code, int nbits)
: code(code) {
assert(8 == nbits);
}
void ProductQuantizer::PQEncoder8::encode(uint64_t x) {
*code++ = (uint8_t)x;
}
ProductQuantizer::PQEncoder16::PQEncoder16(uint8_t *code, int nbits)
: code((uint16_t *)code) {
assert(16 == nbits);
}
void ProductQuantizer::PQEncoder16::encode(uint64_t x) {
*code++ = (uint16_t)x;
}
ProductQuantizer::PQDecoderGeneric::PQDecoderGeneric(const uint8_t *code,
int nbits)
: code(code),
offset(0),
nbits(nbits),
mask((1ull << nbits) - 1),
reg(0) {
assert(nbits <= 64);
}
uint64_t ProductQuantizer::PQDecoderGeneric::decode() {
if (offset == 0) {
reg = *code;
}
uint64_t c = (reg >> offset);
if (offset + nbits >= 8) {
uint64_t e = 8 - offset;
++code;
for (int i = 0; i < (nbits - (8 - offset)) / 8; ++i) {
c |= ((uint64_t)(*code++) << e);
e += 8;
}
offset += nbits;
offset &= 7;
if (offset > 0) {
reg = *code;
c |= ((uint64_t)reg << e);
}
} else {
offset += nbits;
}
return c & mask;
}
ProductQuantizer::PQDecoder8::PQDecoder8(const uint8_t *code, int nbits)
: code(code) {
assert(8 == nbits);
}
uint64_t ProductQuantizer::PQDecoder8::decode() {
return (uint64_t)(*code++);
}
ProductQuantizer::PQDecoder16::PQDecoder16(const uint8_t *code, int nbits)
: code((uint16_t *)code) {
assert(16 == nbits);
}
uint64_t ProductQuantizer::PQDecoder16::decode() {
return (uint64_t)(*code++);
}
} // namespace faiss
+64 -6
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -31,7 +30,6 @@ struct ProductQuantizer {
// values derived from the above
size_t dsub; ///< dimensionality of each subvector
size_t byte_per_idx; ///< nb bytes per code component (1 or 2)
size_t code_size; ///< byte per indexed vector
size_t ksub; ///< number of centroids for each subquantizer
bool verbose; ///< verbose during training?
@@ -142,7 +140,7 @@ struct ProductQuantizer {
/** perform a search (L2 distance)
* @param x query vectors, size nx * d
* @param nx nb of queries
* @param codes database codes, size ncodes * byte_per_idx
* @param codes database codes, size ncodes * code_size
* @param ncodes nb of nb vectors
* @param res heap array to store results (nh == nx)
* @param init_finalize_heap initialize heap (input) and sort (output)?
@@ -176,10 +174,70 @@ struct ProductQuantizer {
float_maxheap_array_t * res,
bool init_finalize_heap = true) const;
struct PQEncoderGeneric {
uint8_t *code; ///< code for this vector
uint8_t offset;
const int nbits; ///< number of bits per subquantizer index
uint8_t reg;
PQEncoderGeneric(uint8_t *code, int nbits, uint8_t offset = 0);
void encode(uint64_t x);
~PQEncoderGeneric();
};
struct PQEncoder8 {
uint8_t *code;
PQEncoder8(uint8_t *code, int nbits);
void encode(uint64_t x);
};
struct PQEncoder16 {
uint16_t *code;
PQEncoder16(uint8_t *code, int nbits);
void encode(uint64_t x);
};
struct PQDecoderGeneric {
const uint8_t *code;
uint8_t offset;
const int nbits;
const uint64_t mask;
uint8_t reg;
PQDecoderGeneric(const uint8_t *code, int nbits);
uint64_t decode();
};
struct PQDecoder8 {
const uint8_t *code;
PQDecoder8(const uint8_t *code, int nbits);
uint64_t decode();
};
struct PQDecoder16 {
const uint16_t *code;
PQDecoder16(const uint8_t *code, int nbits);
uint64_t decode();
};
};
} // namespace faiss
} // namespace faiss
#endif
+3 -2
View File
@@ -4,6 +4,8 @@ Faiss is a library for efficient similarity search and clustering of dense vecto
## NEWS
*NEW: version 1.5.2 (2018-05-27) the license was relaxed to MIT from BSD+Patents. Read LICENSE for details.*
*NEW: version 1.5.0 (2018-12-19) GPU binary flat index and binary HNSW index*
*NEW: version 1.4.0 (2018-08-30) no more crashes in pure Python code*
@@ -80,5 +82,4 @@ We monitor the [issues page](http://github.com/facebookresearch/faiss/issues) of
## License
Faiss is BSD-licensed. We also provide an additional patent grant.
Faiss is MIT-licensed.
+192
View File
@@ -0,0 +1,192 @@
/**
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#include "FaissAssert.h"
#include <exception>
#include <iostream>
namespace faiss {
template <typename IndexT>
ThreadedIndex<IndexT>::ThreadedIndex(bool threaded)
// 0 is default dimension
: ThreadedIndex(0, threaded) {
}
template <typename IndexT>
ThreadedIndex<IndexT>::ThreadedIndex(int d, bool threaded)
: IndexT(d),
own_fields(false),
isThreaded_(threaded) {
}
template <typename IndexT>
ThreadedIndex<IndexT>::~ThreadedIndex() {
for (auto& p : indices_) {
if (isThreaded_) {
// should have worker thread
FAISS_ASSERT((bool) p.second);
// This will also flush all pending work
p.second->stop();
p.second->waitForThreadExit();
} else {
// should not have worker thread
FAISS_ASSERT(!(bool) p.second);
}
if (own_fields) {
delete p.first;
}
}
}
template <typename IndexT>
void ThreadedIndex<IndexT>::addIndex(IndexT* index) {
// We inherit the dimension from the first index added to us if we don't have
// a set dimension
if (indices_.empty() && this->d == 0) {
this->d = index->d;
}
// The new index must match our set dimension
FAISS_THROW_IF_NOT_FMT(this->d == index->d,
"addIndex: dimension mismatch for "
"newly added index; expecting dim %d, "
"new index has dim %d",
this->d, index->d);
if (!indices_.empty()) {
auto& existing = indices_.front().first;
FAISS_THROW_IF_NOT_MSG(index->metric_type == existing->metric_type,
"addIndex: newly added index is "
"of different metric type than old index");
// Make sure this index is not duplicated
for (auto& p : indices_) {
FAISS_THROW_IF_NOT_MSG(p.first != index,
"addIndex: attempting to add index "
"that is already in the collection");
}
}
indices_.emplace_back(
std::make_pair(
index,
std::unique_ptr<WorkerThread>(isThreaded_ ?
new WorkerThread : nullptr)));
onAfterAddIndex(index);
}
template <typename IndexT>
void ThreadedIndex<IndexT>::removeIndex(IndexT* index) {
for (auto it = indices_.begin(); it != indices_.end(); ++it) {
if (it->first == index) {
// This is our index; stop the worker thread before removing it,
// to ensure that it has finished before function exit
if (isThreaded_) {
// should have worker thread
FAISS_ASSERT((bool) it->second);
it->second->stop();
it->second->waitForThreadExit();
} else {
// should not have worker thread
FAISS_ASSERT(!(bool) it->second);
}
indices_.erase(it);
onAfterRemoveIndex(index);
if (own_fields) {
delete index;
}
return;
}
}
// could not find our index
FAISS_THROW_MSG("IndexReplicas::removeIndex: index not found");
}
template <typename IndexT>
void ThreadedIndex<IndexT>::runOnIndex(std::function<void(int, IndexT*)> f) {
if (isThreaded_) {
std::vector<std::future<bool>> v;
for (int i = 0; i < this->indices_.size(); ++i) {
auto& p = this->indices_[i];
auto indexPtr = p.first;
v.emplace_back(p.second->add([f, i, indexPtr](){ f(i, indexPtr); }));
}
waitAndHandleFutures(v);
} else {
// Multiple exceptions may be thrown; gather them as we encounter them,
// while letting everything else run to completion
std::vector<std::pair<int, std::exception_ptr>> exceptions;
for (int i = 0; i < this->indices_.size(); ++i) {
auto& p = this->indices_[i];
try {
f(i, p.first);
} catch (...) {
exceptions.emplace_back(std::make_pair(i, std::current_exception()));
}
}
handleExceptions(exceptions);
}
}
template <typename IndexT>
void ThreadedIndex<IndexT>::runOnIndex(
std::function<void(int, const IndexT*)> f) const {
const_cast<ThreadedIndex<IndexT>*>(this)->runOnIndex(
[f](int i, IndexT* idx){ f(i, idx); });
}
template <typename IndexT>
void ThreadedIndex<IndexT>::reset() {
runOnIndex([](int, IndexT* index){ index->reset(); });
this->ntotal = 0;
this->is_trained = false;
}
template <typename IndexT>
void
ThreadedIndex<IndexT>::onAfterAddIndex(IndexT* index) {
}
template <typename IndexT>
void
ThreadedIndex<IndexT>::onAfterRemoveIndex(IndexT* index) {
}
template <typename IndexT>
void
ThreadedIndex<IndexT>::waitAndHandleFutures(std::vector<std::future<bool>>& v) {
// Blocking wait for completion for all of the indices, capturing any
// exceptions that are generated
std::vector<std::pair<int, std::exception_ptr>> exceptions;
for (int i = 0; i < v.size(); ++i) {
auto& fut = v[i];
try {
fut.get();
} catch (...) {
exceptions.emplace_back(std::make_pair(i, std::current_exception()));
}
}
handleExceptions(exceptions);
}
} // namespace
+80
View File
@@ -0,0 +1,80 @@
/**
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#pragma once
#include "Index.h"
#include "IndexBinary.h"
#include "WorkerThread.h"
#include <memory>
#include <vector>
namespace faiss {
/// A holder of indices in a collection of threads
/// The interface to this class itself is not thread safe
template <typename IndexT>
class ThreadedIndex : public IndexT {
public:
explicit ThreadedIndex(bool threaded);
explicit ThreadedIndex(int d, bool threaded);
~ThreadedIndex() override;
/// override an index that is managed by ourselves.
/// WARNING: once an index is added, it becomes unsafe to touch it from any
/// other thread than that on which is managing it, until we are shut
/// down. Use runOnIndex to perform work on it instead.
void addIndex(IndexT* index);
/// Remove an index that is managed by ourselves.
/// This will flush all pending work on that index, and then shut
/// down its managing thread, and will remove the index.
void removeIndex(IndexT* index);
/// Run a function on all indices, in the thread that the index is
/// managed in.
/// Function arguments are (index in collection, index pointer)
void runOnIndex(std::function<void(int, IndexT*)> f);
void runOnIndex(std::function<void(int, const IndexT*)> f) const;
/// faiss::Index API
/// All indices receive the same call
void reset() override;
/// Returns the number of sub-indices
int count() const { return indices_.size(); }
/// Returns the i-th sub-index
IndexT* at(int i) { return indices_[i].first; }
/// Returns the i-th sub-index (const version)
const IndexT* at(int i) const { return indices_[i].first; }
/// Whether or not we are responsible for deleting our contained indices
bool own_fields;
protected:
/// Called just after an index is added
virtual void onAfterAddIndex(IndexT* index);
/// Called just after an index is removed
virtual void onAfterRemoveIndex(IndexT* index);
protected:
static void waitAndHandleFutures(std::vector<std::future<bool>>& v);
/// Collection of Index instances, with their managing worker thread if any
std::vector<std::pair<IndexT*, std::unique_ptr<WorkerThread>>> indices_;
/// Is this index multi-threaded?
bool isThreaded_;
};
} // namespace
#include "ThreadedIndex-inl.h"
+5 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -222,6 +221,7 @@ void RandomRotationMatrix::init (int seed)
float_randn(q, d_out * d_in, seed);
matrix_qr(d_in, d_out, q);
} else {
// use tight-frame transformation
A.resize (d_out * d_out);
float *q = A.data();
float_randn(q, d_out * d_out, seed);
@@ -867,6 +867,7 @@ IndexPreTransform::IndexPreTransform (
index (index), own_fields (false)
{
is_trained = index->is_trained;
ntotal = index->ntotal;
}
@@ -877,6 +878,7 @@ IndexPreTransform::IndexPreTransform (
index (index), own_fields (false)
{
is_trained = index->is_trained;
ntotal = index->ntotal;
prepend_transform (ltrans);
}
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+21 -7
View File
@@ -1,17 +1,32 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
#include "WorkerThread.h"
#include "FaissAssert.h"
#include <exception>
namespace faiss {
namespace {
// Captures any exceptions thrown by the lambda and returns them via the promise
void runCallback(std::function<void()>& fn,
std::promise<bool>& promise) {
try {
fn();
promise.set_value(true);
} catch (...) {
promise.set_exception(std::current_exception());
}
}
} // namespace
WorkerThread::WorkerThread() :
wantStop_(false) {
startThread();
@@ -70,9 +85,9 @@ WorkerThread::threadMain() {
// Call all pending tasks
FAISS_ASSERT(wantStop_);
// flush all pending operations
for (auto& f : queue_) {
f.first();
f.second.set_value(true);
runCallback(f.first, f.second);
}
}
@@ -96,8 +111,7 @@ WorkerThread::threadLoop() {
queue_.pop_front();
}
data.first();
data.second.set_value(true);
runCallback(data.first, data.second);
}
}
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
+8 -8
View File
@@ -1,11 +1,11 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
from __future__ import print_function
import os
import numpy as np
import faiss
@@ -41,14 +41,14 @@ aa('--eval_freq', default=100, type=int)
args = parser.parse_args()
print "args:", args
print("args:", args)
os.system('echo -n "nb processors "; '
'cat /proc/cpuinfo | grep ^processor | wc -l; '
'cat /proc/cpuinfo | grep ^"model name" | tail -1')
ngpu = faiss.get_num_gpus()
print "nb GPUs:", ngpu
print("nb GPUs:", ngpu)
######################################################
# Load dataset
@@ -71,7 +71,7 @@ xb = xb[:args.nb]
d = xb.shape[1]
if args.pcadim != -1:
print "training PCA: %d -> %d" % (d, args.pcadim)
print("training PCA: %d -> %d" % (d, args.pcadim))
pca = faiss.PCAMatrix(d, args.pcadim)
pca.train(sanitize(xt_pca))
xt = pca.apply_py(sanitize(xt))
@@ -87,7 +87,7 @@ if args.pcadim != -1:
index = faiss.IndexFlatL2(d)
if ngpu > 0:
print "moving index to GPU"
print("moving index to GPU")
index = faiss.index_cpu_to_all_gpus(index)
@@ -115,4 +115,4 @@ for iter0 in range(0, args.niter, args.eval_freq):
error = ((xb - centroids[I.ravel()]) ** 2).sum()
print "iter1=%d quantization error on test: %.4f" % (iter1, error)
print("iter1=%d quantization error on test: %.4f" % (iter1, error))
+19 -19
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#! /usr/bin/env python2
@@ -10,6 +9,7 @@
Common functions to load datasets and compute their ground-truth
"""
from __future__ import print_function
import time
import numpy as np
import faiss
@@ -101,7 +101,7 @@ class ResultHeap:
def compute_GT_sliced(xb, xq, k):
print "compute GT"
print("compute GT")
t0 = time.time()
nb, d = xb.shape
nq, d = xq.shape
@@ -120,24 +120,24 @@ def compute_GT_sliced(xb, xq, k):
D, I = db_gt.search(xqs, k)
rh.add_batch_result(D, I, i0)
db_gt.reset()
print "\r %d/%d, %.3f s" % (i0, nb, time.time() - t0),
print("\r %d/%d, %.3f s" % (i0, nb, time.time() - t0), end=' ')
sys.stdout.flush()
print
print()
rh.finalize()
gt_I = rh.I
print "GT time: %.3f s" % (time.time() - t0)
print("GT time: %.3f s" % (time.time() - t0))
return gt_I
def do_compute_gt(xb, xq, k):
print "computing GT"
print("computing GT")
nb, d = xb.shape
index = faiss.index_cpu_to_all_gpus(faiss.IndexFlatL2(d))
if nb < 100 * 1000:
print " add"
print(" add")
index.add(np.ascontiguousarray(xb, dtype='float32'))
print " search"
print(" search")
D, I = index.search(np.ascontiguousarray(xq, dtype='float32'), k)
else:
I = compute_GT_sliced(xb, xq, k)
@@ -147,7 +147,7 @@ def do_compute_gt(xb, xq, k):
def load_data(dataset='deep1M', compute_gt=False):
print "load data", dataset
print("load data", dataset)
if dataset == 'sift1M':
basedir = simdir + 'sift1M/'
@@ -189,7 +189,7 @@ def load_data(dataset='deep1M', compute_gt=False):
gt_fname = basedir + "%s_groundtruth.ivecs" % dataset
if compute_gt:
gt = do_compute_gt(xb, xq, 100)
print "store", gt_fname
print("store", gt_fname)
ivecs_write(gt_fname, gt)
gt = ivecs_read(gt_fname)
@@ -197,8 +197,8 @@ def load_data(dataset='deep1M', compute_gt=False):
else:
assert False
print "dataset %s sizes: B %s Q %s T %s" % (
dataset, xb.shape, xq.shape, xt.shape)
print("dataset %s sizes: B %s Q %s T %s" % (
dataset, xb.shape, xq.shape, xt.shape))
return xt, xb, xq, gt
@@ -213,7 +213,7 @@ def evaluate_DI(D, I, gt):
rank = 1
while rank <= k:
recall = (I[:, :rank] == gt[:, :1]).sum() / float(nq)
print "R@%d: %.4f" % (rank, recall),
print("R@%d: %.4f" % (rank, recall), end=' ')
rank *= 10
@@ -222,13 +222,13 @@ def evaluate(xq, gt, index, k=100, endl=True):
D, I = index.search(xq, k)
t1 = time.time()
nq = xq.shape[0]
print "\t %8.4f ms per query, " % (
(t1 - t0) * 1000.0 / nq),
print("\t %8.4f ms per query, " % (
(t1 - t0) * 1000.0 / nq), end=' ')
rank = 1
while rank <= k:
recall = (I[:, :rank] == gt[:, :1]).sum() / float(nq)
print "R@%d: %.4f" % (rank, recall),
print("R@%d: %.4f" % (rank, recall), end=' ')
rank *= 10
if endl:
print
print()
return D, I
+2 -3
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#! /usr/bin/env python2
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# @nolint
+68 -78
View File
@@ -1,11 +1,11 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#! /usr/bin/env python2
from __future__ import print_function
import numpy as np
import time
import os
@@ -14,7 +14,7 @@ import faiss
import re
from multiprocessing.dummy import Pool as ThreadPool
from datasets import ivecs_read
####################################################################
# Parse command line
@@ -22,7 +22,7 @@ from multiprocessing.dummy import Pool as ThreadPool
def usage():
print >>sys.stderr, """
print("""
Usage: bench_gpu_1bn.py dataset indextype [options]
@@ -61,7 +61,7 @@ indextype: any index type supported by index_factory that runs on GPU.
%d will be replaced with the nprobe
-oD xx%d.npy output the search result distances to this file
"""
""", file=sys.stderr)
sys.exit(1)
@@ -108,29 +108,19 @@ while args:
elif not dbname: dbname = a
elif not index_key: index_key = a
else:
print >> sys.stderr, "argument %s unknown" % a
print("argument %s unknown" % a, file=sys.stderr)
sys.exit(1)
cacheroot = '/tmp/bench_gpu_1bn'
if not os.path.isdir(cacheroot):
print "%s does not exist, creating it" % cacheroot
print("%s does not exist, creating it" % cacheroot)
os.mkdir(cacheroot)
#################################################################
# Small Utility Functions
#################################################################
def ivecs_read(fname):
a = np.fromfile(fname, dtype='int32')
d = a[0]
return a.reshape(-1, d + 1)[:, 1:].copy()
def fvecs_read(fname):
return ivecs_read(fname).view('float32')
# we mem-map the biggest files to avoid having them in memory all at
# once
@@ -203,7 +193,7 @@ def eval_intersection_measure(gt_I, I):
# Prepare dataset
#################################################################
print "Preparing dataset", dbname
print("Preparing dataset", dbname)
if dbname.startswith('SIFT'):
# SIFT1M to SIFT1000M
@@ -226,7 +216,7 @@ elif dbname == 'Deep1B':
gt_I = ivecs_read('deep1b/deep1B_groundtruth.ivecs')
else:
print >> sys.stderr, 'unknown dataset', dbname
print('unknown dataset', dbname, file=sys.stderr)
sys.exit(1)
@@ -242,9 +232,9 @@ if knngraph:
gt_I = None
print "sizes: B %s Q %s T %s gt %s" % (
print("sizes: B %s Q %s T %s gt %s" % (
xb.shape, xq.shape, xt.shape,
gt_I.shape if gt_I is not None else None)
gt_I.shape if gt_I is not None else None))
@@ -299,17 +289,17 @@ if not use_cache:
cent_cachefile = None
index_cachefile = None
print "cachefiles:"
print preproc_cachefile
print cent_cachefile
print index_cachefile
print("cachefiles:")
print(preproc_cachefile)
print(cent_cachefile)
print(index_cachefile)
#################################################################
# Wake up GPUs
#################################################################
print "preparing resources for %d GPUs" % ngpu
print("preparing resources for %d GPUs" % ngpu)
gpu_resources = []
@@ -338,7 +328,7 @@ def make_vres_vdev(i0=0, i1=-1):
def compute_GT():
print "compute GT"
print("compute GT")
t0 = time.time()
gt_I = np.zeros((nq_gt, gt_sl), dtype='int64')
@@ -367,23 +357,23 @@ def compute_GT():
heaps.addn_with_ids(
gt_sl, faiss.swig_ptr(D), faiss.swig_ptr(I), gt_sl)
db_gt_gpu.reset()
print "\r %d/%d, %.3f s" % (i0, n, time.time() - t0),
print
print("\r %d/%d, %.3f s" % (i0, n, time.time() - t0), end=' ')
print()
heaps.reorder()
print "GT time: %.3f s" % (time.time() - t0)
print("GT time: %.3f s" % (time.time() - t0))
return gt_I
if knngraph:
if gt_cachefile and os.path.exists(gt_cachefile):
print "load GT", gt_cachefile
print("load GT", gt_cachefile)
gt_I = np.load(gt_cachefile)
else:
gt_I = compute_GT()
if gt_cachefile:
print "store GT", gt_cachefile
print("store GT", gt_cachefile)
np.save(gt_cachefile, gt_I)
#################################################################
@@ -392,7 +382,7 @@ if knngraph:
def train_preprocessor():
print "train preproc", preproc_str
print("train preproc", preproc_str)
d = xt.shape[1]
t0 = time.time()
if preproc_str.startswith('OPQ'):
@@ -406,7 +396,7 @@ def train_preprocessor():
else:
assert False
preproc.train(sanitize(xt[:1000000]))
print "preproc train done in %.3f s" % (time.time() - t0)
print("preproc train done in %.3f s" % (time.time() - t0))
return preproc
@@ -415,10 +405,10 @@ def get_preprocessor():
if not preproc_cachefile or not os.path.exists(preproc_cachefile):
preproc = train_preprocessor()
if preproc_cachefile:
print "store", preproc_cachefile
print("store", preproc_cachefile)
faiss.write_VectorTransform(preproc, preproc_cachefile)
else:
print "load", preproc_cachefile
print("load", preproc_cachefile)
preproc = faiss.read_VectorTransform(preproc_cachefile)
else:
d = xb.shape[1]
@@ -438,11 +428,11 @@ def train_coarse_quantizer(x, k, preproc):
# clus.niter = 2
clus.max_points_per_centroid = 10000000
print "apply preproc on shape", x.shape, 'k=', k
print("apply preproc on shape", x.shape, 'k=', k)
t0 = time.time()
x = preproc.apply_py(sanitize(x))
print " preproc %.3f s output shape %s" % (
time.time() - t0, x.shape)
print(" preproc %.3f s output shape %s" % (
time.time() - t0, x.shape))
vres, vdev = make_vres_vdev()
index = faiss.index_cpu_to_gpu_multiple(
@@ -457,16 +447,16 @@ def train_coarse_quantizer(x, k, preproc):
def prepare_coarse_quantizer(preproc):
if cent_cachefile and os.path.exists(cent_cachefile):
print "load centroids", cent_cachefile
print("load centroids", cent_cachefile)
centroids = np.load(cent_cachefile)
else:
nt = max(1000000, 256 * ncent)
print "train coarse quantizer..."
print("train coarse quantizer...")
t0 = time.time()
centroids = train_coarse_quantizer(xt[:nt], ncent, preproc)
print "Coarse train time: %.3f s" % (time.time() - t0)
print("Coarse train time: %.3f s" % (time.time() - t0))
if cent_cachefile:
print "store centroids", cent_cachefile
print("store centroids", cent_cachefile)
np.save(cent_cachefile, centroids)
coarse_quantizer = faiss.IndexFlatL2(preproc.d_out)
@@ -485,13 +475,13 @@ def prepare_trained_index(preproc):
coarse_quantizer = prepare_coarse_quantizer(preproc)
d = preproc.d_out
if pqflat_str == 'Flat':
print "making an IVFFlat index"
print("making an IVFFlat index")
idx_model = faiss.IndexIVFFlat(coarse_quantizer, d, ncent,
faiss.METRIC_L2)
else:
m = int(pqflat_str[2:])
assert m < 56 or use_float16, "PQ%d will work only with -float16" % m
print "making an IVFPQ index, m = ", m
print("making an IVFPQ index, m = ", m)
idx_model = faiss.IndexIVFPQ(coarse_quantizer, d, ncent, m, 8)
coarse_quantizer.this.disown()
@@ -499,10 +489,10 @@ def prepare_trained_index(preproc):
# finish training on CPU
t0 = time.time()
print "Training vector codes"
print("Training vector codes")
x = preproc.apply_py(sanitize(xt[:1000000]))
idx_model.train(x)
print " done %.3f s" % (time.time() - t0)
print(" done %.3f s" % (time.time() - t0))
return idx_model
@@ -526,43 +516,43 @@ def compute_populated_index(preproc):
gpu_index = faiss.index_cpu_to_gpu_multiple(
vres, vdev, indexall, co)
print "add..."
print("add...")
t0 = time.time()
nb = xb.shape[0]
for i0, xs in dataset_iterator(xb, preproc, add_batch_size):
i1 = i0 + xs.shape[0]
gpu_index.add_with_ids(xs, np.arange(i0, i1))
if max_add > 0 and gpu_index.ntotal > max_add:
print "Flush indexes to CPU"
print("Flush indexes to CPU")
for i in range(ngpu):
index_src_gpu = faiss.downcast_index(gpu_index.at(i))
index_src = faiss.index_gpu_to_cpu(index_src_gpu)
print " index %d size %d" % (i, index_src.ntotal)
print(" index %d size %d" % (i, index_src.ntotal))
index_src.copy_subset_to(indexall, 0, 0, nb)
index_src_gpu.reset()
index_src_gpu.reserveMemory(max_add)
gpu_index.sync_with_shard_indexes()
print '\r%d/%d (%.3f s) ' % (
i0, nb, time.time() - t0),
print('\r%d/%d (%.3f s) ' % (
i0, nb, time.time() - t0), end=' ')
sys.stdout.flush()
print "Add time: %.3f s" % (time.time() - t0)
print("Add time: %.3f s" % (time.time() - t0))
print "Aggregate indexes to CPU"
print("Aggregate indexes to CPU")
t0 = time.time()
if hasattr(gpu_index, 'at'):
# it is a sharded index
for i in range(ngpu):
index_src = faiss.index_gpu_to_cpu(gpu_index.at(i))
print " index %d size %d" % (i, index_src.ntotal)
print(" index %d size %d" % (i, index_src.ntotal))
index_src.copy_subset_to(indexall, 0, 0, nb)
else:
# simple index
index_src = faiss.index_gpu_to_cpu(gpu_index)
index_src.copy_subset_to(indexall, 0, 0, nb)
print " done in %.3f s" % (time.time() - t0)
print(" done in %.3f s" % (time.time() - t0))
if max_add > 0:
# it does not contain all the vectors
@@ -591,7 +581,7 @@ def compute_populated_index_2(preproc):
stage2 = rate_limited_imap(quantize, stage1)
print "add..."
print("add...")
t0 = time.time()
nb = xb.shape[0]
@@ -606,10 +596,10 @@ def compute_populated_index_2(preproc):
else:
assert False
print '\r%d/%d (%.3f s) ' % (
i0, nb, time.time() - t0),
print('\r%d/%d (%.3f s) ' % (
i0, nb, time.time() - t0), end=' ')
sys.stdout.flush()
print "Add time: %.3f s" % (time.time() - t0)
print("Add time: %.3f s" % (time.time() - t0))
return None, indexall
@@ -623,10 +613,10 @@ def get_populated_index(preproc):
else:
gpu_index, indexall = compute_populated_index_2(preproc)
if index_cachefile:
print "store", index_cachefile
print("store", index_cachefile)
faiss.write_index(indexall, index_cachefile)
else:
print "load", index_cachefile
print("load", index_cachefile)
indexall = faiss.read_index(index_cachefile)
gpu_index = None
@@ -638,11 +628,11 @@ def get_populated_index(preproc):
co.verbose = True
co.shard = True # the replicas will be made "manually"
t0 = time.time()
print "CPU index contains %d vectors, move to GPU" % indexall.ntotal
print("CPU index contains %d vectors, move to GPU" % indexall.ntotal)
if replicas == 1:
if not gpu_index:
print "copying loaded index to GPUs"
print("copying loaded index to GPUs")
vres, vdev = make_vres_vdev()
index = faiss.index_cpu_to_gpu_multiple(
vres, vdev, indexall, co)
@@ -652,7 +642,7 @@ def get_populated_index(preproc):
else:
del gpu_index # We override the GPU index
print "Copy CPU index to %d sharded GPU indexes" % replicas
print("Copy CPU index to %d sharded GPU indexes" % replicas)
index = faiss.IndexReplicas()
@@ -661,7 +651,7 @@ def get_populated_index(preproc):
gpu1 = ngpu * (i + 1) / replicas
vres, vdev = make_vres_vdev(gpu0, gpu1)
print " dispatch to GPUs %d:%d" % (gpu0, gpu1)
print(" dispatch to GPUs %d:%d" % (gpu0, gpu1))
index1 = faiss.index_cpu_to_gpu_multiple(
vres, vdev, indexall, co)
@@ -669,7 +659,7 @@ def get_populated_index(preproc):
index.addIndex(index1)
index.own_fields = True
del indexall
print "move to GPU done in %.3f s" % (time.time() - t0)
print("move to GPU done in %.3f s" % (time.time() - t0))
return index
@@ -685,7 +675,7 @@ def eval_dataset(index, preproc):
ps.initialize(index)
nq_gt = gt_I.shape[0]
print "search..."
print("search...")
sl = query_batch_size
nq = xq.shape[0]
for nprobe in nprobes:
@@ -701,8 +691,8 @@ def eval_dataset(index, preproc):
inter_res = ''
for i0, xs in dataset_iterator(xq, preproc, sl):
print '\r%d/%d (%.3f s%s) ' % (
i0, nq, time.time() - t0, inter_res),
print('\r%d/%d (%.3f s%s) ' % (
i0, nq, time.time() - t0, inter_res), end=' ')
sys.stdout.flush()
i1 = i0 + xs.shape[0]
@@ -719,24 +709,24 @@ def eval_dataset(index, preproc):
t1 = time.time()
if knngraph:
ires = eval_intersection_measure(gt_I[:, :nnn], I[:nq_gt])
print " probe=%-3d: %.3f s rank-%d intersection results: %.4f" % (
nprobe, t1 - t0, nnn, ires)
print(" probe=%-3d: %.3f s rank-%d intersection results: %.4f" % (
nprobe, t1 - t0, nnn, ires))
else:
print " probe=%-3d: %.3f s" % (nprobe, t1 - t0),
print(" probe=%-3d: %.3f s" % (nprobe, t1 - t0), end=' ')
gtc = gt_I[:, :1]
nq = xq.shape[0]
for rank in 1, 10, 100:
if rank > nnn: continue
nok = (I[:, :rank] == gtc).sum()
print "1-R@%d: %.4f" % (rank, nok / float(nq)),
print
print("1-R@%d: %.4f" % (rank, nok / float(nq)), end=' ')
print()
if I_fname:
I_fname_i = I_fname % I
print "storing", I_fname_i
print("storing", I_fname_i)
np.save(I, I_fname_i)
if D_fname:
D_fname_i = I_fname % I
print "storing", D_fname_i
print("storing", D_fname_i)
np.save(D, D_fname_i)
+19 -50
View File
@@ -1,46 +1,25 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
from __future__ import print_function
import os
import time
import numpy as np
import pdb
import faiss
#################################################################
# I/O functions
#################################################################
def ivecs_read(fname):
a = np.fromfile(fname, dtype='int32')
d = a[0]
return a.reshape(-1, d + 1)[:, 1:].copy()
def fvecs_read(fname):
return ivecs_read(fname).view('float32')
from datasets import load_sift1M, evaluate
#################################################################
# Main program
#################################################################
print "load data"
xt = fvecs_read("sift1M/sift_learn.fvecs")
xb = fvecs_read("sift1M/sift_base.fvecs")
xq = fvecs_read("sift1M/sift_query.fvecs")
print("load data")
xb, xq, xt, gt = load_sift1M()
nq, d = xq.shape
print "load GT"
gt = ivecs_read("sift1M/sift_groundtruth.ivecs")
# we need only a StandardGpuResources per GPU
res = faiss.StandardGpuResources()
@@ -49,40 +28,36 @@ res = faiss.StandardGpuResources()
# Exact search experiment
#################################################################
print "============ Exact search"
print("============ Exact search")
flat_config = faiss.GpuIndexFlatConfig()
flat_config.device = 0
index = faiss.GpuIndexFlatL2(res, d, flat_config)
print "add vectors to index"
print("add vectors to index")
index.add(xb)
print "warmup"
print("warmup")
index.search(xq, 123)
print "benchmark"
print("benchmark")
for lk in range(11):
k = 1 << lk
t0 = time.time()
D, I = index.search(xq, k)
t1 = time.time()
t, r = evaluate(index, xq, gt, k)
# the recall should be 1 at all times
recall_at_1 = (I[:, :1] == gt[:, :1]).sum() / float(nq)
print "k=%d %.3f s, R@1 %.4f" % (
k, t1 - t0, recall_at_1)
print("k=%d %.3f ms, R@1 %.4f" % (k, t, r[1]))
#################################################################
# Approximate search experiment
#################################################################
print "============ Approximate search"
print("============ Approximate search")
index = faiss.index_factory(d, "IVF4096,PQ64")
@@ -97,29 +72,23 @@ co.useFloat16 = True
index = faiss.index_cpu_to_gpu(res, 0, index, co)
print "train"
print("train")
index.train(xt)
print "add vectors to index"
print("add vectors to index")
index.add(xb)
print "warmup"
print("warmup")
index.search(xq, 123)
print "benchmark"
print("benchmark")
for lnprobe in range(10):
nprobe = 1 << lnprobe
index.setNumProbes(nprobe)
t0 = time.time()
D, I = index.search(xq, 100)
t1 = time.time()
t, r = evaluate(index, xq, gt, 100)
print "nprobe=%4d %.3f s recalls=" % (nprobe, t1 - t0),
for rank in 1, 10, 100:
n_ok = (I[:, :rank] == gt[:, :1]).sum()
print "%.4f" % (n_ok / float(nq)),
print
print("nprobe=%4d %.3f ms recalls= %.4f %.4f %.4f" % (nprobe, t, r[1], r[10], r[100]))
+31 -57
View File
@@ -1,51 +1,25 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
from __future__ import print_function
import time
import sys
import numpy as np
import faiss
from datasets import load_sift1M
#################################################################
# Small I/O functions
#################################################################
def ivecs_read(fname):
a = np.fromfile(fname, dtype='int32')
d = a[0]
return a.reshape(-1, d + 1)[:, 1:].copy()
def fvecs_read(fname):
return ivecs_read(fname).view('float32')
#################################################################
# Main program
#################################################################
print "load data"
xt = fvecs_read("sift1M/sift_learn.fvecs")
xb = fvecs_read("sift1M/sift_base.fvecs")
xq = fvecs_read("sift1M/sift_query.fvecs")
nq, d = xq.shape
print "load GT"
gt = ivecs_read("sift1M/sift_groundtruth.ivecs")
k = int(sys.argv[1])
todo = sys.argv[1:]
print("load data")
xb, xq, xt, gt = load_sift1M()
nq, d = xq.shape
if todo == []:
todo = 'hnsw hnsw_sq ivf ivf_hnsw_quantizer kmeans kmeans_hnsw'.split()
@@ -60,13 +34,13 @@ def evaluate(index):
missing_rate = (I == -1).sum() / float(k * nq)
recall_at_1 = (I == gt[:, :1]).sum() / float(nq)
print "\t %7.3f ms per query, R@1 %.4f, missing rate %.4f" % (
(t1 - t0) * 1000.0 / nq, recall_at_1, missing_rate)
print("\t %7.3f ms per query, R@1 %.4f, missing rate %.4f" % (
(t1 - t0) * 1000.0 / nq, recall_at_1, missing_rate))
if 'hnsw' in todo:
print "Testing HNSW Flat"
print("Testing HNSW Flat")
index = faiss.IndexHNSWFlat(d, 32)
@@ -76,27 +50,27 @@ if 'hnsw' in todo:
# construct
index.hnsw.efConstruction = 40
print "add"
print("add")
# to see progress
index.verbose = True
index.add(xb)
print "search"
print("search")
for efSearch in 16, 32, 64, 128, 256:
for bounded_queue in [True, False]:
print "efSearch", efSearch, "bounded queue", bounded_queue,
print("efSearch", efSearch, "bounded queue", bounded_queue, end=' ')
index.hnsw.search_bounded_queue = bounded_queue
index.hnsw.efSearch = efSearch
evaluate(index)
if 'hnsw_sq' in todo:
print "Testing HNSW with a scalar quantizer"
print("Testing HNSW with a scalar quantizer")
# also set M so that the vectors and links both use 128 bytes per
# entry (total 256 bytes)
index = faiss.IndexHNSWSQ(d, faiss.ScalarQuantizer.QT_8bit, 16)
print "training"
print("training")
# training for the scalar quantizer
index.train(xt)
@@ -104,20 +78,20 @@ if 'hnsw_sq' in todo:
# construct
index.hnsw.efConstruction = 40
print "add"
print("add")
# to see progress
index.verbose = True
index.add(xb)
print "search"
print("search")
for efSearch in 16, 32, 64, 128, 256:
print "efSearch", efSearch,
print("efSearch", efSearch, end=' ')
index.hnsw.efSearch = efSearch
evaluate(index)
if 'ivf' in todo:
print "Testing IVF Flat (baseline)"
print("Testing IVF Flat (baseline)")
quantizer = faiss.IndexFlatL2(d)
index = faiss.IndexIVFFlat(quantizer, d, 16384)
index.cp.min_points_per_centroid = 5 # quiet warning
@@ -125,21 +99,21 @@ if 'ivf' in todo:
# to see progress
index.verbose = True
print "training"
print("training")
index.train(xt)
print "add"
print("add")
index.add(xb)
print "search"
print("search")
for nprobe in 1, 4, 16, 64, 256:
print "nprobe", nprobe,
print("nprobe", nprobe, end=' ')
index.nprobe = nprobe
evaluate(index)
if 'ivf_hnsw_quantizer' in todo:
print "Testing IVF Flat with HNSW quantizer"
print("Testing IVF Flat with HNSW quantizer")
quantizer = faiss.IndexHNSWFlat(d, 32)
index = faiss.IndexIVFFlat(quantizer, d, 16384)
index.cp.min_points_per_centroid = 5 # quiet warning
@@ -148,23 +122,23 @@ if 'ivf_hnsw_quantizer' in todo:
# to see progress
index.verbose = True
print "training"
print("training")
index.train(xt)
print "add"
print("add")
index.add(xb)
print "search"
print("search")
quantizer.hnsw.efSearch = 64
for nprobe in 1, 4, 16, 64, 256:
print "nprobe", nprobe,
print("nprobe", nprobe, end=' ')
index.nprobe = nprobe
evaluate(index)
# Bonus: 2 kmeans tests
if 'kmeans' in todo:
print "Performing kmeans on sift1M database vectors (baseline)"
print("Performing kmeans on sift1M database vectors (baseline)")
clus = faiss.Clustering(d, 16384)
clus.verbose = True
clus.niter = 10
@@ -173,7 +147,7 @@ if 'kmeans' in todo:
if 'kmeans_hnsw' in todo:
print "Performing kmeans on sift1M using HNSW assignment"
print("Performing kmeans on sift1M using HNSW assignment")
clus = faiss.Clustering(d, 16384)
clus.verbose = True
clus.niter = 10
+22
View File
@@ -0,0 +1,22 @@
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from __future__ import print_function
import faiss
from datasets import load_sift1M, evaluate
xb, xq, xt, gt = load_sift1M()
nq, d = xq.shape
k = 32
for nbits in 4, 6, 8, 10, 12:
index = faiss.IndexPQ(d, 8, nbits)
index.train(xt)
index.add(xb)
t, r = evaluate(index, xq, gt, k)
print("\t %7.3f ms per query, R@1 %.4f" % (t, r[1]))
del index
+27 -39
View File
@@ -1,11 +1,11 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
from __future__ import print_function
import os
import sys
import time
@@ -13,20 +13,8 @@ import numpy as np
import re
import faiss
from multiprocessing.dummy import Pool as ThreadPool
from datasets import ivecs_read
#################################################################
# I/O functions
#################################################################
def ivecs_read(fname):
a = np.fromfile(fname, dtype='int32')
d = a[0]
return a.reshape(-1, d + 1)[:, 1:].copy()
def fvecs_read(fname):
return ivecs_read(fname).view('float32')
# we mem-map the biggest files to avoid having them in memory all at
# once
@@ -57,7 +45,7 @@ parametersets = sys.argv[3:]
tmpdir = '/tmp/bench_polysemous'
if not os.path.isdir(tmpdir):
print "%s does not exist, creating it" % tmpdir
print("%s does not exist, creating it" % tmpdir)
os.mkdir(tmpdir)
@@ -66,7 +54,7 @@ if not os.path.isdir(tmpdir):
#################################################################
print "Preparing dataset", dbname
print("Preparing dataset", dbname)
if dbname.startswith('SIFT'):
# SIFT1M to SIFT1000M
@@ -89,12 +77,12 @@ elif dbname == 'Deep1B':
gt = ivecs_read('deep1b/deep1B_groundtruth.ivecs')
else:
print >> sys.stderr, 'unknown dataset', dbname
print('unknown dataset', dbname, file=sys.stderr)
sys.exit(1)
print "sizes: B %s Q %s T %s gt %s" % (
xb.shape, xq.shape, xt.shape, gt.shape)
print("sizes: B %s Q %s T %s gt %s" % (
xb.shape, xq.shape, xt.shape, gt.shape))
nq, d = xq.shape
nb, d = xb.shape
@@ -132,7 +120,7 @@ def get_trained_index():
n_train = choose_train_size(index_key)
xtsub = xt[:n_train]
print "Keeping %d train vectors" % xtsub.shape[0]
print("Keeping %d train vectors" % xtsub.shape[0])
# make sure the data is actually in RAM and in float
xtsub = xtsub.astype('float32').copy()
index.verbose = True
@@ -140,11 +128,11 @@ def get_trained_index():
t0 = time.time()
index.train(xtsub)
index.verbose = False
print "train done in %.3f s" % (time.time() - t0)
print "storing", filename
print("train done in %.3f s" % (time.time() - t0))
print("storing", filename)
faiss.write_index(index, filename)
else:
print "loading", filename
print("loading", filename)
index = faiss.read_index(filename)
return index
@@ -187,16 +175,16 @@ def get_populated_index():
t0 = time.time()
for xs in matrix_slice_iterator(xb, 100000):
i1 = i0 + xs.shape[0]
print '\radd %d:%d, %.3f s' % (i0, i1, time.time() - t0),
print('\radd %d:%d, %.3f s' % (i0, i1, time.time() - t0), end=' ')
sys.stdout.flush()
index.add(xs)
i0 = i1
print
print "Add done in %.3f s" % (time.time() - t0)
print "storing", filename
print()
print("Add done in %.3f s" % (time.time() - t0))
print("storing", filename)
faiss.write_index(index, filename)
else:
print "loading", filename
print("loading", filename)
index = faiss.read_index(filename)
return index
@@ -229,28 +217,28 @@ if parametersets == ['autotune'] or parametersets == ['autotuneMT']:
crit.set_groundtruth(None, gt.astype('int64'))
# then we let Faiss find the optimal parameters by itself
print "exploring operating points"
print("exploring operating points")
t0 = time.time()
op = ps.explore(index, xq, crit)
print "Done in %.3f s, available OPs:" % (time.time() - t0)
print("Done in %.3f s, available OPs:" % (time.time() - t0))
# opv is a C++ vector, so it cannot be accessed like a Python array
opv = op.optimal_pts
print "%-40s 1-R@1 time" % "Parameters"
print("%-40s 1-R@1 time" % "Parameters")
for i in range(opv.size()):
opt = opv.at(i)
print "%-40s %.4f %7.3f" % (opt.key, opt.perf, opt.t)
print("%-40s %.4f %7.3f" % (opt.key, opt.perf, opt.t))
else:
# we do queries in a single thread
faiss.omp_set_num_threads(1)
print ' ' * len(parametersets[0]), '\t', 'R@1 R@10 R@100 time %pass'
print(' ' * len(parametersets[0]), '\t', 'R@1 R@10 R@100 time %pass')
for param in parametersets:
print param, '\t',
print(param, '\t', end=' ')
sys.stdout.flush()
ps.set_index_parameters(index, param)
t0 = time.time()
@@ -259,6 +247,6 @@ else:
t1 = time.time()
for rank in 1, 10, 100:
n_ok = (I[:, :rank] == gt[:, :1]).sum()
print "%.4f" % (n_ok / float(nq)),
print "%8.3f " % ((t1 - t0) * 1000.0 / nq),
print "%5.2f" % (ivfpq_stats.n_hamming_pass * 100.0 / ivfpq_stats.ncode)
print("%.4f" % (n_ok / float(nq)), end=' ')
print("%8.3f " % ((t1 - t0) * 1000.0 / nq), end=' ')
print("%5.2f" % (ivfpq_stats.n_hamming_pass * 100.0 / ivfpq_stats.ncode))
+12 -47
View File
@@ -1,58 +1,32 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
from __future__ import print_function
import time
import numpy as np
import faiss
#################################################################
# Small I/O functions
#################################################################
from datasets import load_sift1M, evaluate
def ivecs_read(fname):
a = np.fromfile(fname, dtype='int32')
d = a[0]
return a.reshape(-1, d + 1)[:, 1:].copy()
def fvecs_read(fname):
return ivecs_read(fname).view('float32')
#################################################################
# Main program
#################################################################
print "load data"
xt = fvecs_read("sift1M/sift_learn.fvecs")
xb = fvecs_read("sift1M/sift_base.fvecs")
xq = fvecs_read("sift1M/sift_query.fvecs")
print("load data")
xb, xq, xt, gt = load_sift1M()
nq, d = xq.shape
print "load GT"
gt = ivecs_read("sift1M/sift_groundtruth.ivecs")
# index with 16 subquantizers, 8 bit each
index = faiss.IndexPQ(d, 16, 8)
index.do_polysemous_training = True
index.verbose = True
print "train"
print("train")
index.train(xt)
print "add vectors to index"
print("add vectors to index")
index.add(xb)
@@ -60,22 +34,13 @@ nt = 1
faiss.omp_set_num_threads(1)
def evaluate():
t0 = time.time()
D, I = index.search(xq, 1)
t1 = time.time()
recall_at_1 = (I == gt[:, :1]).sum() / float(nq)
print "\t %7.3f ms per query, R@1 %.4f" % (
(t1 - t0) * 1000.0 / nq * nt, recall_at_1)
print "PQ baseline",
print("PQ baseline", end=' ')
index.search_type = faiss.IndexPQ.ST_PQ
evaluate()
for ht in 64, 62, 58, 54, 50, 46, 42, 38, 34, 30:
print "Polysemous", ht,
print("Polysemous", ht, end=' ')
index.search_type = faiss.IndexPQ.ST_polysemous
index.polysemous_ht = ht
evaluate()
t, r = evaluate(index, xq, gt, 1)
print("\t %7.3f ms per query, R@1 %.4f" % (t, r[1]))
+17 -46
View File
@@ -1,51 +1,22 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
from __future__ import print_function
import time
import numpy as np
import faiss
#################################################################
# I/O functions
#################################################################
from datasets import load_sift1M
def ivecs_read(fname):
a = np.fromfile(fname, dtype='int32')
d = a[0]
return a.reshape(-1, d + 1)[:, 1:].copy()
def fvecs_read(fname):
return ivecs_read(fname).view('float32')
#################################################################
# Main program
#################################################################
print "load data"
xt = fvecs_read("sift1M/sift_learn.fvecs")
xb = fvecs_read("sift1M/sift_base.fvecs")
xq = fvecs_read("sift1M/sift_query.fvecs")
# xq = xq[:1000]
# xb = xb[:100000]
print("load data")
xb, xq, xt, gt = load_sift1M()
nq, d = xq.shape
print "load GT"
gt = ivecs_read("sift1M/sift_groundtruth.ivecs")
# gt = gt[:1000]
ncent = 256
variants = [(name, getattr(faiss.ScalarQuantizer, name))
@@ -58,7 +29,7 @@ quantizer = faiss.IndexFlatL2(d)
if False:
for name, qtype in [('flat', 0)] + variants:
print "============== test", name
print("============== test", name)
t0 = time.time()
if name == 'flat':
@@ -69,23 +40,23 @@ if False:
qtype, faiss.METRIC_L2)
index.nprobe = 16
print "[%.3f s] train" % (time.time() - t0)
print("[%.3f s] train" % (time.time() - t0))
index.train(xt)
print "[%.3f s] add" % (time.time() - t0)
print("[%.3f s] add" % (time.time() - t0))
index.add(xb)
print "[%.3f s] search" % (time.time() - t0)
print("[%.3f s] search" % (time.time() - t0))
D, I = index.search(xq, 100)
print "[%.3f s] eval" % (time.time() - t0)
print("[%.3f s] eval" % (time.time() - t0))
for rank in 1, 10, 100:
n_ok = (I[:, :rank] == gt[:, :1]).sum()
print "%.4f" % (n_ok / float(nq)),
print
print("%.4f" % (n_ok / float(nq)), end=' ')
print()
if True:
for name, qtype in variants:
print "============== test", name
print("============== test", name)
for rsname, vals in [('RS_minmax',
[-0.4, -0.2, -0.1, -0.05, 0.0, 0.1, 0.5]),
@@ -93,7 +64,7 @@ if True:
('RS_quantiles', [0.02, 0.05, 0.1, 0.15]),
('RS_optim', [0.0])]:
for val in vals:
print "%-15s %5g " % (rsname, val),
print("%-15s %5g " % (rsname, val), end=' ')
index = faiss.IndexIVFScalarQuantizer(quantizer, d, ncent,
qtype, faiss.METRIC_L2)
index.nprobe = 16
@@ -110,5 +81,5 @@ if True:
for rank in 1, 10, 100:
n_ok = (I[:, :rank] == gt[:, :1]).sum()
print "%.4f" % (n_ok / float(nq)),
print " %.3f s" % (t1 - t0)
print("%.4f" % (n_ok / float(nq)), end=' ')
print(" %.3f s" % (t1 - t0))
+8 -8
View File
@@ -1,11 +1,11 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#! /usr/bin/env python2
from __future__ import print_function
import numpy as np
import faiss
import time
@@ -27,9 +27,9 @@ np.random.seed(1234)
faiss.omp_set_num_threads(1)
print 'xd=%d yd=%d' % (xd, yd)
print('xd=%d yd=%d' % (xd, yd))
print 'Running inner products test..'
print('Running inner products test..')
for d in 3, 4, 12, 36, 64:
x = faiss.rand(xd * d).reshape(xd, d)
@@ -54,10 +54,10 @@ for d in 3, 4, 12, 36, 64:
num += abs(distances[xi, yi] - np.dot(x[xi], y[yi]))
denom += abs(distances[xi, yi])
print 'd=%d t=%.3f s diff=%g' % (d, t1 - t0, num / denom)
print('d=%d t=%.3f s diff=%g' % (d, t1 - t0, num / denom))
print 'Running L2sqr test..'
print('Running L2sqr test..')
for d in 3, 4, 12, 36, 64:
x = faiss.rand(xd * d).reshape(xd, d)
@@ -82,4 +82,4 @@ for d in 3, 4, 12, 36, 64:
num += abs(distances[xi, yi] - np.sum((x[xi] - y[yi]) ** 2))
denom += abs(distances[xi, yi])
print 'd=%d t=%.3f s diff=%g' % (d, t1 - t0, num / denom)
print('d=%d t=%.3f s diff=%g' % (d, t1 - t0, num / denom))
+45
View File
@@ -0,0 +1,45 @@
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
from __future__ import print_function
import sys
import time
import numpy as np
def ivecs_read(fname):
a = np.fromfile(fname, dtype='int32')
d = a[0]
return a.reshape(-1, d + 1)[:, 1:].copy()
def fvecs_read(fname):
return ivecs_read(fname).view('float32')
def load_sift1M():
print("Loading sift1M...", end='', file=sys.stderr)
xt = fvecs_read("sift1M/sift_learn.fvecs")
xb = fvecs_read("sift1M/sift_base.fvecs")
xq = fvecs_read("sift1M/sift_query.fvecs")
gt = ivecs_read("sift1M/sift_groundtruth.ivecs")
print("done", file=sys.stderr)
return xb, xq, xt, gt
def evaluate(index, xq, gt, k):
nq = xq.shape[0]
t0 = time.time()
D, I = index.search(xq, k) # noqa: E741
t1 = time.time()
recalls = {}
i = 1
while i <= k:
recalls[i] = (I[:, :i] == gt[:, :1]).sum() / float(nq)
i *= 10
return (t1 - t0) * 1000.0 / nq, recalls
+2 -3
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#! /usr/bin/env python2
+3
View File
@@ -1,3 +1,6 @@
README for the link & code implementation
=========================================
+2 -3
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
+2 -3
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#! /usr/bin/env python2
+2 -3
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#! /usr/bin/env python2
+2 -2
View File
@@ -35,7 +35,7 @@ build:
about:
home: https://github.com/facebookresearch/faiss
license: BSD 3-Clause
license_family: BSD
license: MIT
license_family: MIT
license_file: LICENSE
summary: A library for efficient similarity search and clustering of dense vectors.
+2 -2
View File
@@ -30,7 +30,7 @@ build:
about:
home: https://github.com/facebookresearch/faiss
license: BSD 3-Clause
license_family: BSD
license: MIT
license_family: MIT
license_file: LICENSE
summary: A library for efficient similarity search and clustering of dense vectors.
Vendored
+6 -6
View File
@@ -9,9 +9,9 @@
# This configure script is free software; the Free Software Foundation
# gives unlimited permission to copy, distribute and modify it.
#
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# This source code is licensed under the BSD+Patents license found in the
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
## -------------------- ##
## M4sh Initialization. ##
@@ -1447,9 +1447,9 @@ Copyright (C) 2012 Free Software Foundation, Inc.
This configure script is free software; the Free Software Foundation
gives unlimited permission to copy, distribute and modify it.
Copyright (c) 2015-present, Facebook, Inc.
All rights reserved.
This source code is licensed under the BSD+Patents license found in the
Copyright (c) Facebook, Inc. and its affiliates.
This source code is licensed under the MIT license found in the
LICENSE file in the root directory of this source tree.
_ACEOF
exit
+3 -3
View File
@@ -3,9 +3,9 @@
AC_PREREQ([2.69])
AC_INIT([faiss], [1.0])
AC_COPYRIGHT([Copyright (c) 2015-present, Facebook, Inc.
All rights reserved.
This source code is licensed under the BSD+Patents license found in the
AC_COPYRIGHT([Copyright (c) Facebook, Inc. and its affiliates.
This source code is licensed under the MIT license found in the
LICENSE file in the root directory of this source tree.])
AC_CONFIG_SRCDIR([Index.h])
AC_CONFIG_AUX_DIR([build-aux])
+2 -3
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
-include ../makefile.inc
+11 -11
View File
@@ -1,11 +1,11 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
from __future__ import print_function
import os
import time
import numpy as np
@@ -54,7 +54,7 @@ def plot_OperatingPoints(ops, nq, **kwargs):
t0 = time.time()
print "load data"
print("load data")
xt = fvecs_read("sift1M/sift_learn.fvecs")
xb = fvecs_read("sift1M/sift_base.fvecs")
@@ -62,13 +62,13 @@ xq = fvecs_read("sift1M/sift_query.fvecs")
d = xt.shape[1]
print "load GT"
print("load GT")
gt = ivecs_read("sift1M/sift_groundtruth.ivecs")
gt = gt.astype('int64')
k = gt.shape[1]
print "prepare criterion"
print("prepare criterion")
# criterion = 1-recall at 1
crit = faiss.OneRecallAtRCriterion(xq.shape[0], 1)
@@ -122,7 +122,7 @@ op = faiss.OperatingPoints()
for index_key in keys_to_test:
print "============ key", index_key
print("============ key", index_key)
# make the index described by the key
index = faiss.index_factory(d, index_key)
@@ -137,17 +137,17 @@ for index_key in keys_to_test:
params.initialize(index)
print "[%.3f s] train & add" % (time.time() - t0)
print("[%.3f s] train & add" % (time.time() - t0))
index.train(xt)
index.add(xb)
print "[%.3f s] explore op points" % (time.time() - t0)
print("[%.3f s] explore op points" % (time.time() - t0))
# find operating points for this index
opi = params.explore(index, xq, crit)
print "[%.3f s] result operating points:" % (time.time() - t0)
print("[%.3f s] result operating points:" % (time.time() - t0))
opi.display()
# update best operating points so far
@@ -170,6 +170,6 @@ for index_key in keys_to_test:
fig.savefig('tmp/demo_auto_tune.png')
print "[%.3f s] final result:" % (time.time() - t0)
print("[%.3f s] final result:" % (time.time() - t0))
op.display()
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+3 -4
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
@@ -97,7 +96,7 @@ int main ()
index.add (nb, database.data());
printf ("[%.3f s] imbalance factor: %g\n",
elapsed() - t0, index.invlists->imbalance_factor());
elapsed() - t0, index.invlists->imbalance_factor ());
// remember a few elements from the database as queries
int i0 = 1234;
+2 -3
View File
@@ -1,7 +1,6 @@
# Copyright (c) 2015-present, Facebook, Inc.
# All rights reserved.
# Copyright (c) Facebook, Inc. and its affiliates.
#
# This source code is licensed under the BSD+Patents license found in the
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
#!/usr/bin/env python2
+2 -3
View File
@@ -1,8 +1,7 @@
/**
* Copyright (c) 2015-present, Facebook, Inc.
* All rights reserved.
* Copyright (c) Facebook, Inc. and its affiliates.
*
* This source code is licensed under the BSD+Patents license found in the
* This source code is licensed under the MIT license found in the
* LICENSE file in the root directory of this source tree.
*/
+1215 -1213
View File
File diff suppressed because it is too large Load Diff

Some files were not shown because too many files have changed in this diff Show More