diff --git a/AnnService/BalancedDataPartition.vcxproj b/AnnService/BalancedDataPartition.vcxproj
index 6ce001f60..7751af309 100644
--- a/AnnService/BalancedDataPartition.vcxproj
+++ b/AnnService/BalancedDataPartition.vcxproj
@@ -148,5 +148,12 @@
+
+
+
+ This project references NuGet package(s) that are missing on this computer. Use NuGet Package Restore to download them. For more information, see http://go.microsoft.com/fwlink/?LinkID=322105. The missing file is {0}.
+
+
+
\ No newline at end of file
diff --git a/AnnService/CMakeLists.txt b/AnnService/CMakeLists.txt
index 323f3fe2b..de5f588e1 100644
--- a/AnnService/CMakeLists.txt
+++ b/AnnService/CMakeLists.txt
@@ -38,9 +38,9 @@ if(${CMAKE_CXX_COMPILER_ID} STREQUAL "GNU")
endif()
add_library (SPTAGLib SHARED ${SRC_FILES} ${HDR_FILES})
-target_link_libraries (SPTAGLib DistanceUtils libzstd_shared)
+target_link_libraries (SPTAGLib DistanceUtils libzstd_shared ${NUMA_LIBRARY})
add_library (SPTAGLibStatic STATIC ${SRC_FILES} ${HDR_FILES})
-target_link_libraries (SPTAGLibStatic DistanceUtils libzstd_static)
+target_link_libraries (SPTAGLibStatic DistanceUtils libzstd_static ${NUMA_LIBRARY_STATIC})
if(${CMAKE_CXX_COMPILER_ID} STREQUAL "GNU")
target_compile_options(SPTAGLibStatic PRIVATE -fPIC)
endif()
diff --git a/AnnService/inc/Core/BKT/Index.h b/AnnService/inc/Core/BKT/Index.h
index f3131343c..500799c47 100644
--- a/AnnService/inc/Core/BKT/Index.h
+++ b/AnnService/inc/Core/BKT/Index.h
@@ -72,7 +72,6 @@ namespace SPTAG
std::shared_timed_mutex m_dataDeleteLock;
COMMON::Labelset m_deletedID;
- std::unique_ptr> m_workSpacePool;
Helper::ThreadPool m_threadPool;
int m_iNumberOfThreads;
@@ -85,6 +84,7 @@ namespace SPTAG
int m_iNumberOfInitialDynamicPivots;
int m_iNumberOfOtherDynamicPivots;
int m_iHashTableExp;
+ std::unique_ptr> m_workSpaceFactory;
public:
Index()
@@ -98,6 +98,7 @@ namespace SPTAG
m_pSamples.SetName("Vector");
m_fComputeDistance = std::function(COMMON::DistanceCalcSelector(m_iDistCalcMethod));
m_iBaseSquare = (m_iDistCalcMethod == DistCalcMethod::Cosine) ? COMMON::Utils::GetBase() * COMMON::Utils::GetBase() : 1;
+ m_workSpaceFactory = std::make_unique>();
}
~Index() {}
@@ -166,6 +167,24 @@ namespace SPTAG
ErrorCode RefineIndex(const std::vector>& p_indexStreams, IAbortOperation* p_abort);
ErrorCode RefineIndex(std::shared_ptr& p_newIndex);
+ ErrorCode SetWorkSpaceFactory(std::unique_ptr> up_workSpaceFactory)
+ {
+ SPTAG::COMMON::IWorkSpaceFactory* raw_generic_ptr = up_workSpaceFactory.release();
+ if (!raw_generic_ptr) return ErrorCode::Fail;
+
+
+ SPTAG::COMMON::IWorkSpaceFactory* raw_specialized_ptr = dynamic_cast*>(raw_generic_ptr);
+ if (!raw_specialized_ptr)
+ {
+ delete raw_generic_ptr;
+ return ErrorCode::Fail;
+ }
+ else
+ {
+ m_workSpaceFactory = std::unique_ptr>(raw_specialized_ptr);
+ return ErrorCode::Success;
+ }
+ }
private:
void SearchIndex(COMMON::QueryResultSet &p_query, COMMON::WorkSpace &p_space, bool p_searchDeleted, bool p_searchDuplicated) const;
diff --git a/AnnService/inc/Core/Common.h b/AnnService/inc/Core/Common.h
index 34ea786d8..686524ec8 100644
--- a/AnnService/inc/Core/Common.h
+++ b/AnnService/inc/Core/Common.h
@@ -127,7 +127,7 @@ extern std::shared_ptr GetLogger();
#define LOG(l, ...) GetLogger()->Logging("SPTAG", l, __FILE__, __LINE__, __FUNCTION__, __VA_ARGS__)
-class MyException : public std::exception
+class MyException : public std::exception
{
private:
std::string Exp;
@@ -249,6 +249,26 @@ enum class QuantizerType : std::uint8_t
};
static_assert(static_cast(QuantizerType::Undefined) != 0, "Empty QuantizerType!");
+enum class NumaStrategy : std::uint8_t
+{
+#define DefineNumaStrategy(Name) Name,
+#include "DefinitionList.h"
+#undef DefineNumaStrategy
+
+ Undefined
+};
+static_assert(static_cast(NumaStrategy::Undefined) != 0, "Empty NumaStrategy!");
+
+enum class OrderStrategy : std::uint8_t
+{
+#define DefineOrderStrategy(Name) Name,
+#include "DefinitionList.h"
+#undef DefineOrderStrategy
+
+ Undefined
+};
+static_assert(static_cast(OrderStrategy::Undefined) != 0, "Empty OrderStrategy!");
+
} // namespace SPTAG
#endif // _SPTAG_CORE_COMMONDEFS_H_
diff --git a/AnnService/inc/Core/Common/Heap.h b/AnnService/inc/Core/Common/Heap.h
index a5d544b1d..5365c53ca 100644
--- a/AnnService/inc/Core/Common/Heap.h
+++ b/AnnService/inc/Core/Common/Heap.h
@@ -27,7 +27,12 @@ namespace SPTAG
~Heap() {}
inline int size() { return count; }
inline bool empty() { return count == 0; }
- inline void clear() { count = 0; }
+ inline void clear(int size)
+ {
+ if (size > length) Resize(size);
+ count = 0;
+ }
+
inline T& Top() { if (count == 0) return heap[0]; else return heap[1]; }
// Insert a new element in the heap.
diff --git a/AnnService/inc/Core/Common/WorkSpace.h b/AnnService/inc/Core/Common/WorkSpace.h
index 2689095f9..52e9e32dc 100644
--- a/AnnService/inc/Core/Common/WorkSpace.h
+++ b/AnnService/inc/Core/Common/WorkSpace.h
@@ -14,6 +14,32 @@ namespace SPTAG
{
namespace COMMON
{
+ template
+ class IWorkSpaceFactory
+ {
+ public:
+ virtual std::unique_ptr GetWorkSpace() = 0;
+ virtual void ReturnWorkSpace(std::unique_ptr ws) = 0;
+ };
+
+ template
+ class ThreadLocalWorkSpaceFactory : public IWorkSpaceFactory
+ {
+ public:
+ static thread_local std::unique_ptr m_workspace;
+
+ virtual std::unique_ptr< WorkSpaceType> GetWorkSpace() override
+ {
+ return std::move(m_workspace);
+ }
+
+ virtual void ReturnWorkSpace(std::unique_ptr ws) override
+ {
+ m_workspace = std::move(ws);
+ }
+
+ };
+
class OptHashPosVector
{
protected:
@@ -198,8 +224,10 @@ namespace SPTAG
}
};
+ class IWorkSpace {};
+
// Variables for each single NN search
- struct WorkSpace
+ struct WorkSpace : public IWorkSpace
{
WorkSpace() {}
@@ -232,8 +260,8 @@ namespace SPTAG
void Reset(int maxCheck, int resultNum)
{
nodeCheckStatus.clear();
- m_SPTQueue.clear();
- m_NGQueue.clear();
+ m_SPTQueue.clear(maxCheck * 10);
+ m_NGQueue.clear(maxCheck * 30);
m_Results.clear(max(maxCheck / 16, resultNum));
m_iNumOfContinuousNoBetterPropagation = 0;
diff --git a/AnnService/inc/Core/DefinitionList.h b/AnnService/inc/Core/DefinitionList.h
index 812cc3cd0..ad5612777 100644
--- a/AnnService/inc/Core/DefinitionList.h
+++ b/AnnService/inc/Core/DefinitionList.h
@@ -125,4 +125,18 @@ DefineTruthFileType(XVEC)
// row(int32_t), column(int32_t), data...
DefineTruthFileType(DEFAULT)
-#endif // DefineTruthFileType
\ No newline at end of file
+#endif // DefineTruthFileType
+
+#ifdef DefineNumaStrategy
+
+DefineNumaStrategy(LOCAL)
+DefineNumaStrategy(SCATTER)
+
+#endif // DefineNumaStrategy
+
+#ifdef DefineOrderStrategy
+
+DefineOrderStrategy(ASC)
+DefineOrderStrategy(DESC)
+
+#endif // DefineOrderStrategy
\ No newline at end of file
diff --git a/AnnService/inc/Core/KDT/Index.h b/AnnService/inc/Core/KDT/Index.h
index 2eab91df5..1d95d2744 100644
--- a/AnnService/inc/Core/KDT/Index.h
+++ b/AnnService/inc/Core/KDT/Index.h
@@ -70,7 +70,6 @@ namespace SPTAG
std::shared_timed_mutex m_dataDeleteLock;
COMMON::Labelset m_deletedID;
- std::unique_ptr> m_workSpacePool;
Helper::ThreadPool m_threadPool;
int m_iNumberOfThreads;
@@ -83,6 +82,7 @@ namespace SPTAG
int m_iNumberOfInitialDynamicPivots;
int m_iNumberOfOtherDynamicPivots;
int m_iHashTableExp;
+ std::unique_ptr> m_workSpaceFactory;
public:
Index()
@@ -96,6 +96,7 @@ namespace SPTAG
m_pSamples.SetName("Vector");
m_fComputeDistance = std::function(COMMON::DistanceCalcSelector(m_iDistCalcMethod));
m_iBaseSquare = (m_iDistCalcMethod == DistCalcMethod::Cosine) ? COMMON::Utils::GetBase() * COMMON::Utils::GetBase() : 1;
+ m_workSpaceFactory = std::make_unique>();
}
~Index() {}
@@ -164,6 +165,24 @@ namespace SPTAG
ErrorCode RefineIndex(const std::vector>& p_indexStreams, IAbortOperation* p_abort);
ErrorCode RefineIndex(std::shared_ptr& p_newIndex);
+ ErrorCode SetWorkSpaceFactory(std::unique_ptr> up_workSpaceFactory)
+ {
+ SPTAG::COMMON::IWorkSpaceFactory* raw_generic_ptr = up_workSpaceFactory.release();
+ if (!raw_generic_ptr) return ErrorCode::Fail;
+
+
+ SPTAG::COMMON::IWorkSpaceFactory* raw_specialized_ptr = dynamic_cast*>(raw_generic_ptr);
+ if (!raw_specialized_ptr)
+ {
+ delete raw_generic_ptr;
+ return ErrorCode::Fail;
+ }
+ else
+ {
+ m_workSpaceFactory = std::unique_ptr>(raw_specialized_ptr);
+ return ErrorCode::Success;
+ }
+ }
private:
template
diff --git a/AnnService/inc/Core/SPANN/IExtraSearcher.h b/AnnService/inc/Core/SPANN/IExtraSearcher.h
index 8db3f0f5e..274e30b27 100644
--- a/AnnService/inc/Core/SPANN/IExtraSearcher.h
+++ b/AnnService/inc/Core/SPANN/IExtraSearcher.h
@@ -101,15 +101,14 @@ namespace SPTAG {
std::size_t m_pageBufferSize;
};
- struct ExtraWorkSpace
+ struct ExtraWorkSpace : public SPTAG::COMMON::IWorkSpace
{
ExtraWorkSpace() {}
- ~ExtraWorkSpace() {}
+ ~ExtraWorkSpace() { g_spaceCount--; }
ExtraWorkSpace(ExtraWorkSpace& other) {
Initialize(other.m_deduper.MaxCheck(), other.m_deduper.HashTableExponent(), (int)other.m_pageBuffers.size(), (int)(other.m_pageBuffers[0].GetPageSize()), other.m_enableDataCompression);
- m_spaceID = g_spaceCount++;
}
void Initialize(int p_maxCheck, int p_hashExp, int p_internalResultNum, int p_maxPages, bool enableDataCompression) {
@@ -128,6 +127,7 @@ namespace SPTAG {
if (enableDataCompression) {
m_decompressBuffer.ReservePageBuffer(p_maxPages);
}
+ m_spaceID = g_spaceCount++;
}
void Initialize(va_list& arg) {
@@ -139,6 +139,28 @@ namespace SPTAG {
Initialize(maxCheck, hashExp, internalResultNum, maxPages, enableDataCompression);
}
+ void Clear(int p_internalResultNum, int p_maxPages, bool enableDataCompression) {
+ if (p_internalResultNum > m_pageBuffers.size()) {
+ m_postingIDs.reserve(p_internalResultNum);
+ m_processIocp.reset(p_internalResultNum);
+ m_pageBuffers.resize(p_internalResultNum);
+ for (int pi = 0; pi < p_internalResultNum; pi++) {
+ m_pageBuffers[pi].ReservePageBuffer(p_maxPages);
+ }
+ m_diskRequests.resize(p_internalResultNum);
+ for (int pi = 0; pi < p_internalResultNum; pi++) {
+ m_diskRequests[pi].m_extension = m_processIocp.handle();
+ }
+ } else if (p_maxPages > m_pageBuffers[0].GetPageSize()) {
+ for (int pi = 0; pi < m_pageBuffers.size(); pi++) m_pageBuffers[pi].ReservePageBuffer(p_maxPages);
+ }
+
+ m_enableDataCompression = enableDataCompression;
+ if (enableDataCompression) {
+ m_decompressBuffer.ReservePageBuffer(p_maxPages);
+ }
+ }
+
static void Reset() { g_spaceCount = 0; }
std::vector m_postingIDs;
diff --git a/AnnService/inc/Core/SPANN/Index.h b/AnnService/inc/Core/SPANN/Index.h
index 1a4f61098..224138758 100644
--- a/AnnService/inc/Core/SPANN/Index.h
+++ b/AnnService/inc/Core/SPANN/Index.h
@@ -47,16 +47,17 @@ namespace SPTAG
std::unordered_map m_headParameters;
std::shared_ptr m_extraSearcher;
- std::unique_ptr> m_workSpacePool;
Options m_options;
std::function m_fComputeDistance;
int m_iBaseSquare;
+ std::unique_ptr> m_workSpaceFactory;
public:
Index()
{
+ m_workSpaceFactory = std::make_unique>();
m_fComputeDistance = std::function(COMMON::DistanceCalcSelector(m_options.m_distCalcMethod));
m_iBaseSquare = (m_options.m_distCalcMethod == DistCalcMethod::Cosine) ? COMMON::Utils::GetBase() * COMMON::Utils::GetBase() : 1;
}
@@ -138,6 +139,33 @@ namespace SPTAG
ErrorCode DeleteIndex(const SizeType& p_id) { return ErrorCode::Undefined; }
ErrorCode RefineIndex(const std::vector>& p_indexStreams, IAbortOperation* p_abort) { return ErrorCode::Undefined; }
ErrorCode RefineIndex(std::shared_ptr& p_newIndex) { return ErrorCode::Undefined; }
+ ErrorCode SetWorkSpaceFactory(std::unique_ptr> up_workSpaceFactory)
+ {
+ SPTAG::COMMON::IWorkSpaceFactory* raw_generic_ptr = up_workSpaceFactory.release();
+ if (!raw_generic_ptr) return ErrorCode::Fail;
+
+
+ SPTAG::COMMON::IWorkSpaceFactory* raw_specialized_ptr = dynamic_cast*>(raw_generic_ptr);
+ if (!raw_specialized_ptr)
+ {
+ // If it is of type SPTAG::COMMON::WorkSpace, we should pass on to child index
+ if (!m_index)
+ {
+ delete raw_generic_ptr;
+ return ErrorCode::Fail;
+ }
+ else
+ {
+ return m_index->SetWorkSpaceFactory(std::unique_ptr>(raw_generic_ptr));
+ }
+
+ }
+ else
+ {
+ m_workSpaceFactory = std::unique_ptr>(raw_specialized_ptr);
+ return ErrorCode::Success;
+ }
+ }
private:
bool CheckHeadIndexType();
void SelectHeadAdjustOptions(int p_vectorCount);
diff --git a/AnnService/inc/Core/VectorIndex.h b/AnnService/inc/Core/VectorIndex.h
index 19896251c..7f62f2499 100644
--- a/AnnService/inc/Core/VectorIndex.h
+++ b/AnnService/inc/Core/VectorIndex.h
@@ -11,6 +11,7 @@
#include "inc/Helper/SimpleIniReader.h"
#include
#include "inc/Core/Common/IQuantizer.h"
+#include "inc/Core/Common/WorkSpace.h"
namespace SPTAG
{
@@ -143,6 +144,8 @@ class VectorIndex
virtual ErrorCode RefineIndex(const std::vector>& p_indexStreams, IAbortOperation* p_abort) = 0;
+ virtual ErrorCode SetWorkSpaceFactory(std::unique_ptr> up_workSpaceFactory) = 0;
+
inline bool HasMetaMapping() const { return nullptr != m_pMetaToVec; }
inline SizeType GetMetaMapping(std::string& meta) const;
diff --git a/AnnService/inc/Helper/AsyncFileReader.h b/AnnService/inc/Helper/AsyncFileReader.h
index b4be27727..0ad1b812c 100644
--- a/AnnService/inc/Helper/AsyncFileReader.h
+++ b/AnnService/inc/Helper/AsyncFileReader.h
@@ -16,23 +16,25 @@
#include
#define ASYNC_READ 1
+#define BATCH_READ 1
#ifdef _MSC_VER
#include
#include
-#define BATCH_READ 1
#else
-#define BATCH_READ 1
#include
#include
#include
+#ifdef NUMA
+#include
+#endif
#endif
namespace SPTAG
{
namespace Helper
{
-
+ void SetThreadAffinity(int threadID, std::thread& thread, NumaStrategy socketStrategy = NumaStrategy::LOCAL, OrderStrategy idStrategy = OrderStrategy::ASC);
#ifdef _MSC_VER
namespace DiskUtils
{
@@ -168,7 +170,7 @@ namespace SPTAG
m_fileIocp.Reset(::CreateIoCompletionPort(m_fileHandle.GetHandle(), NULL, NULL, iocpThreads));
for (int i = 0; i < iocpThreads; ++i)
{
- m_fileIocpThreads.emplace_back(std::thread(std::bind(&AsyncFileIO::ListionIOCP, this)));
+ m_fileIocpThreads.emplace_back(std::thread(std::bind(&AsyncFileIO::ListionIOCP, this, i)));
}
return m_fileIocp.IsValid();
}
@@ -375,8 +377,10 @@ namespace SPTAG
ExitProcess(dw);
}
- void ListionIOCP()
+ void ListionIOCP(int i)
{
+ SetThreadAffinity(i, m_fileIocpThreads[i], NumaStrategy::SCATTER, OrderStrategy::DESC); // avoid IO threads overlap with search threads
+
DWORD cBytes;
ULONG_PTR key;
OVERLAPPED* ol;
diff --git a/AnnService/inc/SSDServing/SSDIndex.h b/AnnService/inc/SSDServing/SSDIndex.h
index 2858be7c2..8a26b8b7d 100644
--- a/AnnService/inc/SSDServing/SSDIndex.h
+++ b/AnnService/inc/SSDServing/SSDIndex.h
@@ -115,37 +115,39 @@ namespace SPTAG {
Utils::StopW sw;
- auto func = [&]()
- {
- Utils::StopW threadws;
- size_t index = 0;
- while (true)
+ for (int i = 0; i < p_numThreads; i++) { threads.emplace_back([&, i]()
{
- index = queriesSent.fetch_add(1);
- if (index < numQueries)
+ NumaStrategy ns = (p_index->GetDiskIndex() != nullptr) ? NumaStrategy::SCATTER : NumaStrategy::LOCAL; // Only for SPANN, we need to avoid IO threads overlap with search threads.
+ Helper::SetThreadAffinity(i, threads[i], ns, OrderStrategy::ASC);
+
+ Utils::StopW threadws;
+ size_t index = 0;
+ while (true)
{
- if ((index & ((1 << 14) - 1)) == 0)
+ index = queriesSent.fetch_add(1);
+ if (index < numQueries)
{
- LOG(Helper::LogLevel::LL_Info, "Sent %.2lf%%...\n", index * 100.0 / numQueries);
- }
+ if ((index & ((1 << 14) - 1)) == 0)
+ {
+ LOG(Helper::LogLevel::LL_Info, "Sent %.2lf%%...\n", index * 100.0 / numQueries);
+ }
- double startTime = threadws.getElapsedMs();
- p_index->GetMemoryIndex()->SearchIndex(p_results[index]);
- double endTime = threadws.getElapsedMs();
- p_index->SearchDiskIndex(p_results[index], &(p_stats[index]));
- double exEndTime = threadws.getElapsedMs();
+ double startTime = threadws.getElapsedMs();
+ p_index->GetMemoryIndex()->SearchIndex(p_results[index]);
+ double endTime = threadws.getElapsedMs();
+ p_index->SearchDiskIndex(p_results[index], &(p_stats[index]));
+ double exEndTime = threadws.getElapsedMs();
- p_stats[index].m_exLatency = exEndTime - endTime;
- p_stats[index].m_totalLatency = p_stats[index].m_totalSearchLatency = exEndTime - startTime;
- }
- else
- {
- return;
+ p_stats[index].m_exLatency = exEndTime - endTime;
+ p_stats[index].m_totalLatency = p_stats[index].m_totalSearchLatency = exEndTime - startTime;
+ }
+ else
+ {
+ return;
+ }
}
- }
- };
-
- for (int i = 0; i < p_numThreads; i++) { threads.emplace_back(func); }
+ });
+ }
for (auto& thread : threads) { thread.join(); }
double sendingCost = sw.getElapsedSec();
diff --git a/AnnService/src/Core/BKT/BKTIndex.cpp b/AnnService/src/Core/BKT/BKTIndex.cpp
index f9ef95f6f..a2676959d 100644
--- a/AnnService/src/Core/BKT/BKTIndex.cpp
+++ b/AnnService/src/Core/BKT/BKTIndex.cpp
@@ -10,8 +10,12 @@
namespace SPTAG
{
+ template
+ thread_local std::unique_ptr COMMON::ThreadLocalWorkSpaceFactory::m_workspace;
+
namespace BKT
{
+
template
ErrorCode Index::LoadConfig(Helper::IniReader& p_reader)
{
@@ -67,8 +71,6 @@ namespace SPTAG
else if (m_deletedID.Load((char*)p_indexBlobs[3].Data(), m_iDataBlockSize, m_iDataCapacity) != ErrorCode::Success) return ErrorCode::FailedParseValue;
omp_set_num_threads(m_iNumberOfThreads);
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
m_threadPool.init();
return ErrorCode::Success;
}
@@ -86,8 +88,6 @@ namespace SPTAG
else if ((ret = m_deletedID.Load(p_indexStreams[3], m_iDataBlockSize, m_iDataCapacity)) != ErrorCode::Success) return ret;
omp_set_num_threads(m_iNumberOfThreads);
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
m_threadPool.init();
return ret;
}
@@ -95,9 +95,12 @@ namespace SPTAG
template
ErrorCode Index::SaveConfig(std::shared_ptr p_configOut)
{
- auto workSpace = m_workSpacePool->Rent();
- m_iHashTableExp = workSpace->HashTableExponent();
- m_workSpacePool->Return(workSpace);
+ auto workspace = m_workSpaceFactory->GetWorkSpace();
+ if (workspace)
+ {
+ m_iHashTableExp = workspace->HashTableExponent();
+ }
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workspace));
#define DefineBKTParameter(VarName, VarType, DefaultValue, RepresentStr) \
IOSTRING(p_configOut, WriteString, (RepresentStr + std::string("=") + GetParameter(RepresentStr) + std::string("\n")).c_str());
@@ -276,12 +279,16 @@ namespace SPTAG
{
if (!m_bReady) return ErrorCode::EmptyIndex;
- auto workSpace = m_workSpacePool->Rent();
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new COMMON::WorkSpace());
+ workSpace->Initialize(max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
+ }
workSpace->Reset(m_iMaxCheck, p_query.GetResultNum());
SearchIndex(*((COMMON::QueryResultSet*)&p_query), *workSpace, p_searchDeleted, true);
- m_workSpacePool->Return(workSpace);
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
if (p_query.WithMeta() && nullptr != m_pMetadata)
{
@@ -297,19 +304,26 @@ namespace SPTAG
template
ErrorCode Index::RefineSearchIndex(QueryResult &p_query, bool p_searchDeleted) const
{
- auto workSpace = m_workSpacePool->Rent();
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new COMMON::WorkSpace());
+ workSpace->Initialize(max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
+ }
workSpace->Reset(m_pGraph.m_iMaxCheckForRefineGraph, p_query.GetResultNum());
-
SearchIndex(*((COMMON::QueryResultSet*)&p_query), *workSpace, p_searchDeleted, false);
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
- m_workSpacePool->Return(workSpace);
return ErrorCode::Success;
}
template
ErrorCode Index::SearchTree(QueryResult& p_query) const
{
- auto workSpace = m_workSpacePool->Rent();
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new COMMON::WorkSpace());
+ workSpace->Initialize(max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
+ }
workSpace->Reset(m_pGraph.m_iMaxCheckForRefineGraph, p_query.GetResultNum());
COMMON::QueryResultSet* p_results = (COMMON::QueryResultSet*)&p_query;
@@ -322,7 +336,8 @@ namespace SPTAG
res[i].VID = cell.node;
res[i].Dist = cell.distance;
}
- m_workSpacePool->Return(workSpace);
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
+
return ErrorCode::Success;
}
#pragma endregion
@@ -346,8 +361,6 @@ namespace SPTAG
}
}
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
m_threadPool.init();
auto t1 = std::chrono::high_resolution_clock::now();
@@ -400,8 +413,6 @@ namespace SPTAG
LOG(Helper::LogLevel::LL_Info, "Refine... from %d -> %d\n", GetNumSamples(), newR);
if (newR == 0) return ErrorCode::EmptyIndex;
- ptr->m_workSpacePool.reset(new COMMON::WorkSpacePool());
- ptr->m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
ptr->m_threadPool.init();
ErrorCode ret = ErrorCode::Success;
@@ -569,8 +580,6 @@ namespace SPTAG
Index::UpdateIndex()
{
omp_set_num_threads(m_iNumberOfThreads);
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
return ErrorCode::Success;
}
diff --git a/AnnService/src/Core/KDT/KDTIndex.cpp b/AnnService/src/Core/KDT/KDTIndex.cpp
index 1431e266c..c97fdbce7 100644
--- a/AnnService/src/Core/KDT/KDTIndex.cpp
+++ b/AnnService/src/Core/KDT/KDTIndex.cpp
@@ -10,8 +10,11 @@
namespace SPTAG
{
+ template
+ thread_local std::unique_ptr COMMON::ThreadLocalWorkSpaceFactory::m_workspace;
namespace KDT
{
+
template
ErrorCode Index::LoadConfig(Helper::IniReader& p_reader)
{
@@ -66,8 +69,6 @@ namespace SPTAG
else if (m_deletedID.Load((char*)p_indexBlobs[3].Data(), m_iDataBlockSize, m_iDataCapacity) != ErrorCode::Success) return ErrorCode::FailedParseValue;
omp_set_num_threads(m_iNumberOfThreads);
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
m_threadPool.init();
return ErrorCode::Success;
}
@@ -85,8 +86,6 @@ namespace SPTAG
else if ((ret = m_deletedID.Load(p_indexStreams[3], m_iDataBlockSize, m_iDataCapacity)) != ErrorCode::Success) return ret;
omp_set_num_threads(m_iNumberOfThreads);
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
m_threadPool.init();
return ret;
}
@@ -94,9 +93,12 @@ namespace SPTAG
template
ErrorCode Index::SaveConfig(std::shared_ptr p_configOut)
{
- auto workSpace = m_workSpacePool->Rent();
- m_iHashTableExp = workSpace->HashTableExponent();
- m_workSpacePool->Return(workSpace);
+ auto workspace = m_workSpaceFactory->GetWorkSpace();
+ if (workspace)
+ {
+ m_iHashTableExp = workspace->HashTableExponent();
+ }
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workspace));
#define DefineKDTParameter(VarName, VarType, DefaultValue, RepresentStr) \
IOSTRING(p_configOut, WriteString, (RepresentStr + std::string("=") + GetParameter(RepresentStr) + std::string("\n")).c_str());
@@ -183,7 +185,11 @@ namespace SPTAG
{
if (!m_bReady) return ErrorCode::EmptyIndex;
- auto workSpace = m_workSpacePool->Rent();
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new COMMON::WorkSpace());
+ workSpace->Initialize(max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
+ }
workSpace->Reset(m_iMaxCheck, p_query.GetResultNum());
COMMON::QueryResultSet* p_results = (COMMON::QueryResultSet*) & p_query;
@@ -211,9 +217,9 @@ case VectorValueType::Name: \
SearchIndex(*p_results, *workSpace, p_searchDeleted);
}
- m_workSpacePool->Return(workSpace);
-
- if (p_query.WithMeta() && nullptr != m_pMetadata)
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
+
+ if (p_query.WithMeta() && nullptr != m_pMetadata)
{
for (int i = 0; i < p_query.GetResultNum(); ++i)
{
@@ -221,13 +227,18 @@ case VectorValueType::Name: \
p_query.SetMetadata(i, (result < 0) ? ByteArray::c_empty : m_pMetadata->GetMetadataCopy(result));
}
}
+
return ErrorCode::Success;
}
template
ErrorCode Index::RefineSearchIndex(QueryResult &p_query, bool p_searchDeleted) const
{
- auto workSpace = m_workSpacePool->Rent();
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new COMMON::WorkSpace());
+ workSpace->Initialize(max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
+ }
workSpace->Reset(m_pGraph.m_iMaxCheckForRefineGraph, p_query.GetResultNum());
COMMON::QueryResultSet* p_results = (COMMON::QueryResultSet*) & p_query;
@@ -254,15 +265,19 @@ case VectorValueType::Name: \
{
SearchIndex(*p_results, *workSpace, p_searchDeleted);
}
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
- m_workSpacePool->Return(workSpace);
return ErrorCode::Success;
}
template
ErrorCode Index::SearchTree(QueryResult& p_query) const
{
- auto workSpace = m_workSpacePool->Rent();
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new COMMON::WorkSpace());
+ workSpace->Initialize(max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
+ }
workSpace->Reset(m_pGraph.m_iMaxCheckForRefineGraph, p_query.GetResultNum());
COMMON::QueryResultSet* p_results = (COMMON::QueryResultSet*)&p_query;
@@ -299,7 +314,8 @@ case VectorValueType::Name: \
res[i].VID = cell.node;
res[i].Dist = cell.distance;
}
- m_workSpacePool->Return(workSpace);
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
+
return ErrorCode::Success;
}
#pragma endregion
@@ -323,8 +339,6 @@ case VectorValueType::Name: \
}
}
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
m_threadPool.init();
auto t1 = std::chrono::high_resolution_clock::now();
@@ -377,8 +391,6 @@ case VectorValueType::Name: \
LOG(Helper::LogLevel::LL_Info, "Refine... from %d -> %d\n", GetNumSamples(), newR);
if (newR == 0) return ErrorCode::EmptyIndex;
- ptr->m_workSpacePool.reset(new COMMON::WorkSpacePool());
- ptr->m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
ptr->m_threadPool.init();
ErrorCode ret = ErrorCode::Success;
@@ -555,8 +567,6 @@ case VectorValueType::Name: \
Index::UpdateIndex()
{
omp_set_num_threads(m_iNumberOfThreads);
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_iNumberOfThreads, max(m_iMaxCheck, m_pGraph.m_iMaxCheckForRefineGraph), m_iHashTableExp);
return ErrorCode::Success;
}
diff --git a/AnnService/src/Core/SPANN/SPANNIndex.cpp b/AnnService/src/Core/SPANN/SPANNIndex.cpp
index 5b7e7c844..f017b3b41 100644
--- a/AnnService/src/Core/SPANN/SPANNIndex.cpp
+++ b/AnnService/src/Core/SPANN/SPANNIndex.cpp
@@ -12,6 +12,8 @@
namespace SPTAG
{
+ template
+ thread_local std::unique_ptr COMMON::ThreadLocalWorkSpaceFactory::m_workspace;
namespace SPANN
{
std::atomic_int ExtraWorkSpace::g_spaceCount(0);
@@ -101,8 +103,6 @@ namespace SPTAG
m_vectorTranslateMap.reset((std::uint64_t*)(p_indexBlobs.back().Data()), [=](std::uint64_t* ptr) {});
omp_set_num_threads(m_options.m_iSSDNumberOfThreads);
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_options.m_iSSDNumberOfThreads, m_options.m_maxCheck, m_options.m_hashExp, m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, int(m_options.m_enableDataCompression));
return ErrorCode::Success;
}
@@ -133,8 +133,6 @@ namespace SPTAG
IOBINARY(p_indexStreams[m_index->GetIndexFiles()->size()], ReadBinary, sizeof(std::uint64_t) * m_index->GetNumSamples(), reinterpret_cast(m_vectorTranslateMap.get()));
omp_set_num_threads(m_options.m_iSSDNumberOfThreads);
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_options.m_iSSDNumberOfThreads, m_options.m_maxCheck, m_options.m_hashExp, m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, int(m_options.m_enableDataCompression));
return ErrorCode::Success;
}
@@ -203,9 +201,15 @@ namespace SPTAG
m_index->SearchIndex(*p_queryResults);
- std::shared_ptr workSpace = nullptr;
if (m_extraSearcher != nullptr) {
- workSpace = m_workSpacePool->Rent();
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new ExtraWorkSpace());
+ workSpace->Initialize(m_options.m_maxCheck, m_options.m_hashExp, m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, m_options.m_enableDataCompression);
+ }
+ else {
+ workSpace->Clear(m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, m_options.m_enableDataCompression);
+ }
workSpace->m_deduper.clear();
workSpace->m_postingIDs.clear();
@@ -223,7 +227,7 @@ namespace SPTAG
}
// Don't do disk reads for irrelevant pages
- if (workSpace->m_postingIDs.size() >= m_options.m_searchInternalResultNum ||
+ if (workSpace->m_postingIDs.size() >= m_options.m_searchInternalResultNum ||
(limitDist > 0.1 && res->Dist > limitDist) ||
!m_extraSearcher->CheckValidPosting(postingID))
continue;
@@ -232,7 +236,7 @@ namespace SPTAG
p_queryResults->Reverse();
m_extraSearcher->SearchIndex(workSpace.get(), *p_queryResults, m_index, nullptr);
- m_workSpacePool->Return(workSpace);
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
p_queryResults->SortResult();
}
@@ -258,7 +262,15 @@ namespace SPTAG
if (nullptr == m_extraSearcher) return ErrorCode::EmptyIndex;
COMMON::QueryResultSet* p_queryResults = (COMMON::QueryResultSet*) & p_query;
- std::shared_ptr workSpace = m_workSpacePool->Rent();
+
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new ExtraWorkSpace());
+ workSpace->Initialize(m_options.m_maxCheck, m_options.m_hashExp, m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, m_options.m_enableDataCompression);
+ }
+ else {
+ workSpace->Clear(m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, m_options.m_enableDataCompression);
+ }
workSpace->m_deduper.clear();
workSpace->m_postingIDs.clear();
@@ -294,7 +306,7 @@ namespace SPTAG
p_queryResults->Reverse();
m_extraSearcher->SearchIndex(workSpace.get(), *p_queryResults, m_index, p_stats);
- m_workSpacePool->Return(workSpace);
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
p_queryResults->SortResult();
return ErrorCode::Success;
}
@@ -321,29 +333,34 @@ namespace SPTAG
}
newResults.Reverse();
- auto auto_ws = m_workSpacePool->Rent();
- auto_ws->m_deduper.clear();
+ auto workSpace = m_workSpaceFactory->GetWorkSpace();
+ if (!workSpace) {
+ workSpace.reset(new ExtraWorkSpace());
+ workSpace->Initialize(m_options.m_maxCheck, m_options.m_hashExp, m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, m_options.m_enableDataCompression);
+ }
+ else {
+ workSpace->Clear(m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, m_options.m_enableDataCompression);
+ }
+ workSpace->m_deduper.clear();
int partitions = (p_internalResultNum + p_subInternalResultNum - 1) / p_subInternalResultNum;
float limitDist = p_query.GetResult(0)->Dist * m_options.m_maxDistRatio;
for (SizeType p = 0; p < partitions; p++) {
int subInternalResultNum = min(p_subInternalResultNum, p_internalResultNum - p_subInternalResultNum * p);
- auto_ws->m_postingIDs.clear();
+ workSpace->m_postingIDs.clear();
for (int i = p * p_subInternalResultNum; i < p * p_subInternalResultNum + subInternalResultNum; i++)
{
auto res = p_query.GetResult(i);
if (res->VID == -1 || (limitDist > 0.1 && res->Dist > limitDist)) break;
if (!m_extraSearcher->CheckValidPosting(res->VID)) continue;
- auto_ws->m_postingIDs.emplace_back(res->VID);
+ workSpace->m_postingIDs.emplace_back(res->VID);
}
- m_extraSearcher->SearchIndex(auto_ws.get(), newResults, m_index, p_stats, truth, found);
+ m_extraSearcher->SearchIndex(workSpace.get(), newResults, m_index, p_stats, truth, found);
}
-
- m_workSpacePool->Return(auto_ws);
-
+ m_workSpaceFactory->ReturnWorkSpace(std::move(workSpace));
newResults.SortResult();
std::copy(newResults.GetResults(), newResults.GetResults() + newResults.GetResultNum(), p_query.GetResults());
return ErrorCode::Success;
@@ -776,10 +793,6 @@ namespace SPTAG
}
}
- if (m_extraSearcher != nullptr) {
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_options.m_iSSDNumberOfThreads, m_options.m_maxCheck, m_options.m_hashExp, m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, int(m_options.m_enableDataCompression));
- }
m_bReady = true;
return ErrorCode::Success;
}
@@ -845,8 +858,6 @@ namespace SPTAG
//m_index->SetParameter("MaxCheck", std::to_string(m_options.m_maxCheck));
//m_index->SetParameter("HashTableExponent", std::to_string(m_options.m_hashExp));
m_index->UpdateIndex();
- m_workSpacePool.reset(new COMMON::WorkSpacePool());
- m_workSpacePool->Init(m_options.m_iSSDNumberOfThreads, m_options.m_maxCheck, m_options.m_hashExp, m_options.m_searchInternalResultNum, max(m_options.m_postingPageLimit, m_options.m_searchPostingPageLimit + 1) << PageSizeEx, int(m_options.m_enableDataCompression));
return ErrorCode::Success;
}
diff --git a/AnnService/src/Helper/AsyncFileReader.cpp b/AnnService/src/Helper/AsyncFileReader.cpp
index cbbc88d63..e5f4fdb82 100644
--- a/AnnService/src/Helper/AsyncFileReader.cpp
+++ b/AnnService/src/Helper/AsyncFileReader.cpp
@@ -6,6 +6,49 @@
namespace SPTAG {
namespace Helper {
#ifndef _MSC_VER
+ void SetThreadAffinity(int threadID, std::thread& thread, NumaStrategy socketStrategy, OrderStrategy idStrategy)
+ {
+#ifdef NUMA
+ int numGroups = numa_num_task_nodes();
+ int numCpus = numa_num_task_cpus() / numGroups;
+
+ int group = threadID / numCpus;
+ int cpuid = threadID % numCpus;
+ if (socketStrategy == NumaStrategy::SCATTER) {
+ group = threadID % numGroups;
+ cpuid = (threadID / numGroups) % numCpus;
+ }
+
+ struct bitmask* cpumask = numa_allocate_cpumask();
+ if (!numa_node_to_cpus(group, cpumask)) {
+ unsigned int nodecpu = 0;
+ for (unsigned int i = 0; i < cpumask->size; i++) {
+ if (numa_bitmask_isbitset(cpumask, i)) {
+ if (cpuid == nodecpu) {
+ cpu_set_t cpuset;
+ CPU_ZERO(&cpuset);
+ CPU_SET(i, &cpuset);
+ int rc = pthread_setaffinity_np(thread.native_handle(), sizeof(cpu_set_t), &cpuset);
+ if (rc != 0) {
+ LOG(Helper::LogLevel::LL_Error, "Error calling pthread_setaffinity_np for thread %d: %d\n", threadID, rc);
+ }
+ break;
+ }
+ nodecpu++;
+ }
+ }
+ }
+#else
+ cpu_set_t cpuset;
+ CPU_ZERO(&cpuset);
+ CPU_SET(threadID, &cpuset);
+ int rc = pthread_setaffinity_np(thread.native_handle(), sizeof(cpu_set_t), &cpuset);
+ if (rc != 0) {
+ LOG(Helper::LogLevel::LL_Error, "Error calling pthread_setaffinity_np for thread %d: %d\n", threadID, rc);
+ }
+#endif
+ }
+
struct timespec AIOTimeout {0, 30000};
void BatchReadFileAsync(std::vector>& handlers, AsyncReadRequest* readRequests, int num)
{
@@ -80,6 +123,58 @@ namespace SPTAG {
}
}
#else
+ ULONGLONG GetCpuMasks(WORD group, DWORD numCpus)
+ {
+ ULONGLONG masks = 0, mask = 1;
+ for (DWORD i = 0; i < numCpus; ++i)
+ {
+ masks |= mask;
+ mask <<= 1;
+ }
+
+ return masks;
+ }
+
+ void SetThreadAffinity(int threadID, std::thread& thread, NumaStrategy socketStrategy, OrderStrategy idStrategy)
+ {
+ WORD numGroups = GetActiveProcessorGroupCount();
+ DWORD numCpus = GetActiveProcessorCount(0);
+
+ GROUP_AFFINITY ga;
+ memset(&ga, 0, sizeof(ga));
+ PROCESSOR_NUMBER pn;
+ memset(&pn, 0, sizeof(pn));
+
+ WORD group = (WORD)(threadID / numCpus);
+ pn.Number = (BYTE)(threadID % numCpus);
+ if (socketStrategy == NumaStrategy::SCATTER) {
+ group = (WORD)(threadID % numGroups);
+ pn.Number = (BYTE)((threadID / numGroups) % numCpus);
+ }
+
+ ga.Group = group;
+ ga.Mask = GetCpuMasks(group, numCpus);
+ BOOL res = SetThreadGroupAffinity(GetCurrentThread(), &ga, NULL);
+ if (!res)
+ {
+ LOG(Helper::LogLevel::LL_Error, "Failed SetThreadGroupAffinity for group %d and mask %I64x for thread %d.\n", ga.Group, ga.Mask, threadID);
+ return;
+ }
+ pn.Group = group;
+ if (idStrategy == OrderStrategy::DESC) {
+ pn.Number = (BYTE)(numCpus - 1 - pn.Number);
+ }
+ res = SetThreadIdealProcessorEx(GetCurrentThread(), &pn, NULL);
+ if (!res)
+ {
+ LOG(Helper::LogLevel::LL_Error, "Unable to set ideal processor for thread %d.\n", threadID);
+ return;
+ }
+
+ //LOG(Helper::LogLevel::LL_Info, "numGroup:%d numCPUs:%d threadID:%d group:%d cpuid:%d\n", (int)(numGroups), (int)numCpus, threadID, (int)(group), (int)(pn.Number));
+ YieldProcessor();
+ }
+
void BatchReadFileAsync(std::vector>& handlers, AsyncReadRequest* readRequests, int num)
{
if (handlers.size() == 1) {
diff --git a/AnnService/src/IndexSearcher/main.cpp b/AnnService/src/IndexSearcher/main.cpp
index 8497f831f..29eca8a72 100644
--- a/AnnService/src/IndexSearcher/main.cpp
+++ b/AnnService/src/IndexSearcher/main.cpp
@@ -5,6 +5,7 @@
#include "inc/Helper/SimpleIniReader.h"
#include "inc/Helper/CommonHelper.h"
#include "inc/Helper/StringConvert.h"
+#include "inc/Helper/AsyncFileReader.h"
#include "inc/Core/Common/CommonUtils.h"
#include "inc/Core/Common/TruthSet.h"
#include "inc/Core/Common/QueryResultSet.h"
@@ -187,29 +188,32 @@ int Process(std::shared_ptr options, VectorIndex& index)
std::atomic_size_t queriesSent(0);
std::vector threads;
- auto func = [&]()
- {
- size_t qid = 0;
- while (true)
- {
- qid = queriesSent.fetch_add(1);
- if (qid < numQuerys)
- {
- auto t1 = std::chrono::high_resolution_clock::now();
- index.SearchIndex(results[qid]);
- auto t2 = std::chrono::high_resolution_clock::now();
- latencies[qid] = (float)(std::chrono::duration_cast(t2 - t1).count() / 1000000.0);
- }
- else
- {
- return;
- }
- }
- };
auto batchstart = std::chrono::high_resolution_clock::now();
- for (std::uint32_t i = 0; i < options->m_threadNum; i++) { threads.emplace_back(func); }
+ for (std::uint32_t i = 0; i < options->m_threadNum; i++) {
+ threads.emplace_back([&, i] {
+ NumaStrategy ns = (index.GetIndexAlgoType() == IndexAlgoType::SPANN)? NumaStrategy::SCATTER: NumaStrategy::LOCAL; // Only for SPANN, we need to avoid IO threads overlap with search threads.
+ Helper::SetThreadAffinity(i, threads[i], ns, OrderStrategy::ASC);
+
+ size_t qid = 0;
+ while (true)
+ {
+ qid = queriesSent.fetch_add(1);
+ if (qid < numQuerys)
+ {
+ auto t1 = std::chrono::high_resolution_clock::now();
+ index.SearchIndex(results[qid]);
+ auto t2 = std::chrono::high_resolution_clock::now();
+ latencies[qid] = (float)(std::chrono::duration_cast(t2 - t1).count() / 1000000.0);
+ }
+ else
+ {
+ return;
+ }
+ }
+ });
+ }
for (auto& thread : threads) { thread.join(); }
auto batchend = std::chrono::high_resolution_clock::now();
diff --git a/CMakeLists.txt b/CMakeLists.txt
index f1b97a451..a49fa402e 100644
--- a/CMakeLists.txt
+++ b/CMakeLists.txt
@@ -26,6 +26,31 @@ if(${CMAKE_CXX_COMPILER_ID} STREQUAL "GNU")
set (CMAKE_CXX_FLAGS "-Wall -Wunreachable-code -Wno-reorder -Wno-sign-compare -Wno-unknown-pragmas -Wcast-align -lm -lrt -std=c++14 -fopenmp")
set (CMAKE_CXX_FLAGS_RELEASE "-DNDEBUG -O3 -march=native")
set (CMAKE_CXX_FLAGS_DEBUG "-g -DDEBUG")
+
+
+ find_path(NUMA_INCLUDE_DIR NAME numa.h
+ HINTS $ENV{HOME}/local/include /opt/local/include /usr/local/include /usr/include)
+
+ find_library(NUMA_LIBRARY NAME libnuma.so
+ HINTS $ENV{HOME}/local/lib64 $ENV{HOME}/local/lib /usr/local/lib64 /usr/local/lib /opt/local/lib64 /opt/local/lib /usr/lib64 /usr/lib)
+
+ find_library(NUMA_LIBRARY_STATIC NAME libnuma.a
+ HINTS $ENV{HOME}/local/lib64 $ENV{HOME}/local/lib /usr/local/lib64 /usr/local/lib /opt/local/lib64 /opt/local/lib /usr/lib64 /usr/lib)
+
+ if (NUMA_INCLUDE_DIR AND NUMA_LIBRARY AND NUMA_LIBRARY_STATIC)
+ set (CMAKE_C_FLAGS "${CMAKE_C_FLAGS} -lnuma")
+ set (CMAKE_CXX_FLAGS "${CMAKE_CXX_FLAGS} -lnuma")
+
+ include_directories (${NUMA_INCLUDE_DIR})
+ message (STATUS "Found numa library: inc=${NUMA_INCLUDE_DIR}, lib=${NUMA_LIBRARY}, staticlib=${NUMA_LIBRARY_STATIC}")
+ add_definitions(-DNUMA)
+ else ()
+ set (NUMA_LIBRARY "")
+ set (NUMA_LIBRARY_STATIC "")
+ message (STATUS "WARNING: Numa library not found.")
+ message (STATUS "Try: 'sudo yum install numactl numactl-devel' (or sudo apt-get install libnuma libnuma-dev)")
+ endif ()
+
elseif(WIN32)
if(NOT MSVC14)
message(FATAL_ERROR "On Windows, only MSVC version 14 are supported!")
diff --git a/Dockerfile b/Dockerfile
index 00ec1f842..5f7758925 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -4,7 +4,7 @@ WORKDIR /app
ENV DEBIAN_FRONTEND=noninteractive
RUN apt-get update && apt-get -y install wget build-essential \
- swig cmake git \
+ swig cmake git libnuma libnuma-dev \
libboost-filesystem-dev libboost-test-dev libboost-serialization-dev libboost-regex-dev libboost-serialization-dev libboost-regex-dev libboost-thread-dev libboost-system-dev
ENV PYTHONPATH=/app/Release
diff --git a/Dockerfile.cuda b/Dockerfile.cuda
index bfd265104..3428a2bf8 100644
--- a/Dockerfile.cuda
+++ b/Dockerfile.cuda
@@ -4,7 +4,7 @@ WORKDIR /app
ENV DEBIAN_FRONTEND=noninteractive
RUN apt-get update && apt-get -y install build-essential \
- swig cmake git \
+ swig cmake git libnuma libnuma-dev \
libboost-filesystem-dev libboost-test-dev libboost-serialization-dev libboost-regex-dev libboost-serialization-dev libboost-regex-dev libboost-thread-dev libboost-system-dev
ENV PYTHONPATH=/app/Release
diff --git a/GPUSupport/CMakeLists.txt b/GPUSupport/CMakeLists.txt
index 36845381e..391ca5ef8 100644
--- a/GPUSupport/CMakeLists.txt
+++ b/GPUSupport/CMakeLists.txt
@@ -54,11 +54,11 @@ if (CUDA_FOUND)
endif()
CUDA_ADD_LIBRARY(GPUSPTAGLib SHARED ${GPU_SRC_FILES} ${GPU_HDR_FILES})
- target_link_libraries(GPUSPTAGLib DistanceUtils ${Boost_LIBRARIES} ${CUDA_LIBRARIES} libzstd_shared)
+ target_link_libraries(GPUSPTAGLib DistanceUtils ${Boost_LIBRARIES} ${CUDA_LIBRARIES} libzstd_shared ${NUMA_LIBRARY})
target_compile_definitions(GPUSPTAGLib PRIVATE ${Definition})
CUDA_ADD_LIBRARY(GPUSPTAGLibStatic STATIC ${GPU_SRC_FILES} ${GPU_HDR_FILES})
- target_link_libraries(GPUSPTAGLibStatic DistanceUtils ${Boost_LIBRARIES} ${CUDA_LIBRARIES} libzstd_static)
+ target_link_libraries(GPUSPTAGLibStatic DistanceUtils ${Boost_LIBRARIES} ${CUDA_LIBRARIES} libzstd_static ${NUMA_LIBRARY_STATIC})
target_compile_definitions(GPUSPTAGLibStatic PRIVATE ${Definition})
add_dependencies(GPUSPTAGLibStatic GPUSPTAGLib)
diff --git a/Wrappers/CMakeLists.txt b/Wrappers/CMakeLists.txt
index 1dcfa6d0d..3637f013b 100644
--- a/Wrappers/CMakeLists.txt
+++ b/Wrappers/CMakeLists.txt
@@ -54,9 +54,9 @@ if (Python_FOUND)
add_library (_SPTAG SHARED ${CORE_SRC_FILES} ${CORE_HDR_FILES})
set_target_properties(_SPTAG PROPERTIES PREFIX "" SUFFIX ${PY_SUFFIX})
if (CUDA_FOUND)
- target_link_libraries(_SPTAG GPUSPTAGLibStatic ${Python_LIBRARIES})
+ target_link_libraries(_SPTAG GPUSPTAGLib ${Python_LIBRARIES})
else()
- target_link_libraries(_SPTAG SPTAGLibStatic ${Python_LIBRARIES})
+ target_link_libraries(_SPTAG SPTAGLib ${Python_LIBRARIES})
endif()
add_custom_command(TARGET _SPTAG POST_BUILD COMMAND ${CMAKE_COMMAND} -E copy ${PROJECT_SOURCE_DIR}/Wrappers/inc/SPTAG.py ${EXECUTABLE_OUTPUT_PATH})
@@ -64,7 +64,7 @@ if (Python_FOUND)
file(GLOB CLIENT_SRC_FILES ${PROJECT_SOURCE_DIR}/Wrappers/src/ClientInterface.cpp ${PROJECT_SOURCE_DIR}/AnnService/src/Socket/*.cpp ${PROJECT_SOURCE_DIR}/AnnService/src/Client/*.cpp ${PROJECT_SOURCE_DIR}/Wrappers/inc/ClientInterface_pwrap.cpp)
add_library (_SPTAGClient SHARED ${CLIENT_SRC_FILES} ${CLIENT_HDR_FILES})
set_target_properties(_SPTAGClient PROPERTIES PREFIX "" SUFFIX ${PY_SUFFIX})
- target_link_libraries(_SPTAGClient SPTAGLibStatic ${Python_LIBRARIES} ${Boost_LIBRARIES})
+ target_link_libraries(_SPTAGClient SPTAGLib ${Python_LIBRARIES} ${Boost_LIBRARIES})
add_custom_command(TARGET _SPTAGClient POST_BUILD COMMAND ${CMAKE_COMMAND} -E copy ${PROJECT_SOURCE_DIR}/Wrappers/inc/SPTAGClient.py ${EXECUTABLE_OUTPUT_PATH})
install(TARGETS _SPTAG _SPTAGClient
@@ -99,13 +99,13 @@ if (JNI_FOUND)
file(GLOB CORE_SRC_FILES ${PROJECT_SOURCE_DIR}/Wrappers/src/CoreInterface.cpp ${PROJECT_SOURCE_DIR}/Wrappers/inc/CoreInterface_jwrap.cpp)
add_library (JAVASPTAG SHARED ${CORE_SRC_FILES} ${CORE_HDR_FILES})
set_target_properties(JAVASPTAG PROPERTIES SUFFIX ${JAVA_SUFFIX})
- target_link_libraries(JAVASPTAG SPTAGLibStatic ${JNI_LIBRARIES})
+ target_link_libraries(JAVASPTAG SPTAGLib ${JNI_LIBRARIES})
file(GLOB CLIENT_HDR_FILES ${PROJECT_SOURCE_DIR}/Wrappers/inc/ClientInterface.h ${PROJECT_SOURCE_DIR}/AnnService/inc/Socket/*.h ${PROJECT_SOURCE_DIR}/AnnService/inc/Client/*.h)
file(GLOB CLIENT_SRC_FILES ${PROJECT_SOURCE_DIR}/Wrappers/src/ClientInterface.cpp ${PROJECT_SOURCE_DIR}/AnnService/src/Socket/*.cpp ${PROJECT_SOURCE_DIR}/AnnService/src/Client/*.cpp ${PROJECT_SOURCE_DIR}/Wrappers/inc/ClientInterface_jwrap.cpp)
add_library (JAVASPTAGClient SHARED ${CLIENT_SRC_FILES} ${CLIENT_HDR_FILES})
set_target_properties(JAVASPTAGClient PROPERTIES SUFFIX ${JAVA_SUFFIX})
- target_link_libraries(JAVASPTAGClient SPTAGLibStatic ${JNI_LIBRARIES} ${Boost_LIBRARIES})
+ target_link_libraries(JAVASPTAGClient SPTAGLib ${JNI_LIBRARIES} ${Boost_LIBRARIES})
file(GLOB JAVA_FILES ${PROJECT_SOURCE_DIR}/Wrappers/inc/*.java)
foreach(JAVA_FILE ${JAVA_FILES})
@@ -164,13 +164,13 @@ if (DOTNET_FOUND)
file(GLOB CORE_SRC_FILES ${PROJECT_SOURCE_DIR}/Wrappers/src/CoreInterface.cpp ${PROJECT_SOURCE_DIR}/Wrappers/inc/CoreInterface_cwrap.cpp)
add_library (CSHARPSPTAG SHARED ${CORE_SRC_FILES} ${CORE_HDR_FILES})
set_target_properties(CSHARPSPTAG PROPERTIES SUFFIX ${CSHARP_SUFFIX})
- target_link_libraries(CSHARPSPTAG SPTAGLibStatic)
+ target_link_libraries(CSHARPSPTAG SPTAGLib)
file(GLOB CLIENT_HDR_FILES ${PROJECT_SOURCE_DIR}/Wrappers/inc/ClientInterface.h ${PROJECT_SOURCE_DIR}/AnnService/inc/Socket/*.h ${PROJECT_SOURCE_DIR}/AnnService/inc/Client/*.h)
file(GLOB CLIENT_SRC_FILES ${PROJECT_SOURCE_DIR}/Wrappers/src/ClientInterface.cpp ${PROJECT_SOURCE_DIR}/AnnService/src/Socket/*.cpp ${PROJECT_SOURCE_DIR}/AnnService/src/Client/*.cpp ${PROJECT_SOURCE_DIR}/Wrappers/inc/ClientInterface_cwrap.cpp)
add_library (CSHARPSPTAGClient SHARED ${CLIENT_SRC_FILES} ${CLIENT_HDR_FILES})
set_target_properties(CSHARPSPTAGClient PROPERTIES SUFFIX ${CSHARP_SUFFIX})
- target_link_libraries(CSHARPSPTAGClient SPTAGLibStatic ${Boost_LIBRARIES})
+ target_link_libraries(CSHARPSPTAGClient SPTAGLib ${Boost_LIBRARIES})
file(GLOB CSHARP_FILES ${PROJECT_SOURCE_DIR}/Wrappers/inc/*.cs)
foreach(CSHARP_FILE ${CSHARP_FILES})