Added VWDictionary tests and doc. Fixed LSH not working (fix from https://github.com/flann-lib/flann/pull/472

This commit is contained in:
matlabbe
2025-12-22 21:07:25 -08:00
parent c815c12fc7
commit f7cb330e31
9 changed files with 1238 additions and 59 deletions
+8 -1
View File
@@ -3647,7 +3647,14 @@ void DBDriverSqlite3::loadQuery(VWDictionary & dictionary, bool lastStateOnly) c
if(dataSize>4 && data)
{
UDEBUG("A flann index was saved in the database (size=%ld).", dataSize);
dictionary.deserializeIndex((const unsigned char*)data, dataSize);
if(!dictionary.deserializeIndex((const unsigned char*)data, dataSize))
{
UERROR("Failed to deserialize dictionary's index! See previous logs for reason.");
}
else
{
UINFO("Sucessfully loaded dictionary's index.");
}
}
else {
UDEBUG("No flann index was saved in the database.");
+13 -3
View File
@@ -304,6 +304,7 @@ void FlannIndex::buildIndex(
break;
case FLANN_INDEX_LSH:
UASSERT(features.type() == CV_8UC1);
UASSERT_MSG(features.cols >= 8, "LSH requires a minimum of 8 dimensions to provide valid results.");
params = rtflann::LshIndexParams(12, 20, 2);
break;
default:
@@ -719,10 +720,11 @@ void FlannIndex::knnSearch(
UERROR("Flann index not yet created!");
return;
}
indices.create(query.rows, knn, sizeof(size_t)==8?CV_64F:CV_32S);
dists.create(query.rows, knn, featuresType_ == CV_8UC1?CV_32S:CV_32F);
rtflann::Matrix<size_t> indicesF((size_t*)indices.data, indices.rows, indices.cols);
dists = cv::Mat(query.rows, knn, featuresType_ == CV_8UC1?CV_32S:CV_32F, cv::Scalar(-1));
std::vector<size_t> indicesBuffer(query.rows * knn, std::numeric_limits<size_t>::max());
rtflann::Matrix<size_t> indicesF((size_t*)indicesBuffer.data(), query.rows, knn);
rtflann::SearchParams params = rtflann::SearchParams(checks, eps, sorted);
@@ -749,6 +751,14 @@ void FlannIndex::knnSearch(
((rtflann::Index<rtflann::L2<float> >*)index_)->knnSearch(queryF, indicesF, distsF, knn, params);
}
}
indices.create(query.rows, knn, CV_32S);
int * ptr = indices.ptr<int>();
for(size_t i=0 ; i<indicesBuffer.size(); i+=2)
{
ptr[i] = indicesBuffer[i] == std::numeric_limits<size_t>::max()?-1:(int)indicesBuffer[i];
ptr[i+1] = indicesBuffer[i+1] == std::numeric_limits<size_t>::max()?-1:(int)indicesBuffer[i+1];
}
}
void FlannIndex::radiusSearch(
+27 -33
View File
@@ -661,14 +661,17 @@ void VWDictionary::update()
ULOGGER_DEBUG("_mapIndexId.size() = %d, words.size()=%d, _dim=%d",_mapIndexId.size(), _visualWords.size(), dim);
ULOGGER_DEBUG("copying data = %f s", timer.ticks());
_flannIndex->buildIndex(
_strategy == kNNFlannNaive ? FlannIndex::FLANN_INDEX_LINEAR:
_strategy == kNNFlannLSH ? FlannIndex::FLANN_INDEX_LSH:
FlannIndex::FLANN_INDEX_KDTREE, // kNNFlannKdTree
_dataTree,
useDistanceL1_,
_incrementalDictionary&&_incrementalFlann?_rebalancingFactor:1);
ULOGGER_DEBUG("Time to create kd tree = %f s", timer.ticks());
if(_strategy < kNNBruteForce)
{
_flannIndex->buildIndex(
_strategy == kNNFlannNaive ? FlannIndex::FLANN_INDEX_LINEAR:
_strategy == kNNFlannLSH ? FlannIndex::FLANN_INDEX_LSH:
FlannIndex::FLANN_INDEX_KDTREE, // kNNFlannKdTree
_dataTree,
useDistanceL1_,
_incrementalDictionary&&_incrementalFlann?_rebalancingFactor:1);
ULOGGER_DEBUG("Time to create kd tree = %f s", timer.ticks());
}
}
}
UDEBUG("Dictionary updated! (size=%d added=%d removed=%d)",
@@ -697,38 +700,38 @@ std::vector<unsigned char> VWDictionary::serializeIndex() const
return _flannIndex->serializeIndex(_serializeWithChecksum);
}
void VWDictionary::deserializeIndex(const std::vector<unsigned char> & data)
bool VWDictionary::deserializeIndex(const std::vector<unsigned char> & data)
{
deserializeIndex(data.data(), data.size());
return deserializeIndex(data.data(), data.size());
}
void VWDictionary::deserializeIndex(const unsigned char * data, size_t size)
bool VWDictionary::deserializeIndex(const unsigned char * data, size_t size)
{
if(data== NULL || size == 0)
{
UWARN("Trying to deserialize empty data, aborting.");
return;
return false;
}
UDEBUG("Loading flann index... (data size=%ld bytes)", size);
if(_strategy >= kNNBruteForce) {
//ignore
return;
return false;
}
if(_flannIndex->isBuilt()) {
UERROR("Flann index is already built, cannot deserialize data!");
return;
return false;
}
if(_visualWords.empty()) {
UERROR("Descriptors should be added before deserializing flann index! See VWDictionary::addWord()");
return;
return false;
}
if(!(_removedIndexedWords.empty() && _visualWords.size() == _notIndexedWords.size())) {
UERROR("State of dictionary not as expected before deserializing. (removed words=%ld, words=%ld, not indexed=%ld)",
_removedIndexedWords.size(), _visualWords.size(), _notIndexedWords.size());
return;
return false;
}
std::map<int, int> mapIndexId;
@@ -818,9 +821,11 @@ void VWDictionary::deserializeIndex(const unsigned char * data, size_t size)
else {
UWARN("Failed deserializing flann index data (error: %s), the index will be rebuilt on next update.", errorMsg.c_str());
_flannIndex->release(); // reset to initial state
return false;
}
ULOGGER_DEBUG("Time to load flann index = %f s", timer.ticks());
return true;
}
void VWDictionary::clear(bool printWarningsIfNotEmpty)
@@ -1075,14 +1080,9 @@ std::list<int> VWDictionary::addNewWords(
for(int j=0; j<dists.cols; ++j)
{
float d = dists.at<float>(i,j);
int index;
if (sizeof(size_t) == 8)
{
index = *((size_t*)&results.at<double>(i, j));
}
else
{
index = *((size_t*)&results.at<int>(i, j));
int index = results.at<int>(i, j);
if(index<0) {
continue;
}
int id = uValue(_mapIndexId, index);
if(d >= 0.0f && id != 0)
@@ -1440,15 +1440,9 @@ std::vector<int> VWDictionary::findNN(const cv::Mat & queryIn) const
for(int j=0; j<dists.cols; ++j)
{
float d = dists.at<float>(i,j);
int index;
if (sizeof(size_t) == 8)
{
index = *((size_t*)&results.at<double>(i, j));
}
else
{
index = *((size_t*)&results.at<int>(i, j));
int index = results.at<int>(i, j);
if(index < 0) {
continue;
}
int id = uValue(_mapIndexId, index);
if(d >= 0.0f && id != 0)
+4 -3
View File
@@ -107,7 +107,7 @@ public:
capacity_(capacity_)
{
// reserving capacity to prevent memory re-allocations
dist_index_.resize(capacity_, DistIndex(std::numeric_limits<DistanceType>::max(),-1));
dist_index_.resize(capacity_, DistIndex(std::numeric_limits<DistanceType>::max(),std::numeric_limits<size_t>::max()));
clear();
}
@@ -210,7 +210,7 @@ public:
KNNResultSet(int capacity) : capacity_(capacity)
{
// reserving capacity to prevent memory re-allocations
dist_index_.resize(capacity_, DistIndex(std::numeric_limits<DistanceType>::max(),-1));
dist_index_.resize(capacity_, DistIndex(std::numeric_limits<DistanceType>::max(),std::numeric_limits<size_t>::max()));
clear();
}
@@ -252,7 +252,8 @@ public:
#endif
{
// Check for duplicate indices
for (size_t j = i - 1; dist_index_[j].dist_ == dist && j--;) {
// https://github.com/flann-lib/flann/pull/472
for (size_t j = i; j-- && dist_index_[j].dist_ == dist;) {
if (dist_index_[j].index_ == index) {
return;
}