mirror of
https://github.com/facebookresearch/faiss.git
synced 2026-10-11 22:50:00 +00:00
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:
+18
-6
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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
@@ -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];
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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);
|
||||
|
||||
@@ -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.
|
||||
@@ -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 ();
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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,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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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))
|
||||
|
||||
@@ -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]))
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
@@ -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,3 +1,6 @@
|
||||
|
||||
|
||||
|
||||
README for the link & code implementation
|
||||
=========================================
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
#! /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.
|
||||
|
||||
#! /usr/bin/env python2
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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.
|
||||
*/
|
||||
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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,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.
|
||||
*/
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user