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})