diff --git a/include/svs/orchestrators/dynamic_vamana.h b/include/svs/orchestrators/dynamic_vamana.h index 03d307c4f..d2e99b0aa 100644 --- a/include/svs/orchestrators/dynamic_vamana.h +++ b/include/svs/orchestrators/dynamic_vamana.h @@ -288,6 +288,8 @@ class DynamicVamana : public manager::IndexManager { /// @param distance Distance functor or enum. /// @param threadpool_proto Thread pool or number of threads to use. /// @param graph_allocator Allocator instance to use for the graph. + /// @param logger The logger to use for this index. Defaults to the global SVS logger + /// (``svs::logging::get()``). /// template < manager::QueryTypeDefinition QueryTypes, @@ -301,7 +303,8 @@ class DynamicVamana : public manager::IndexManager { std::span ids, Distance distance, ThreadPoolProto threadpool_proto, - const GraphAllocator& graph_allocator = {} + const GraphAllocator& graph_allocator = {}, + svs::logging::logger_ptr logger = svs::logging::get() ) { auto threadpool = threads::as_threadpool(std::move(threadpool_proto)); auto data = @@ -317,7 +320,8 @@ class DynamicVamana : public manager::IndexManager { ids, std::move(distance_function), std::move(threadpool), - graph_allocator + graph_allocator, + std::move(logger) ); }); } else { @@ -328,12 +332,26 @@ class DynamicVamana : public manager::IndexManager { ids, std::move(distance), std::move(threadpool), - graph_allocator + graph_allocator, + std::move(logger) ); } } // Assembly + /// + /// @brief Load a DynamicVamana index from a previously saved index. + /// + /// @param config_path Path to the directory where the index configuration was saved. + /// @param graph_loader The loader for the graph to use. + /// @param data_loader An acceptable data loader or dataset. + /// @param distance Distance functor or ``svs::DistanceType`` enum. + /// @param threadpool_proto Thread pool or number of threads to use. + /// @param debug_load_from_static Internal/unstable: load files produced by the static + /// index using an identity ID translation. + /// @param logger The logger to use for this index. Defaults to the global SVS logger + /// (``svs::logging::get()``). + /// template < manager::QueryTypeDefinition QueryTypes, typename ConfigProto, @@ -347,7 +365,8 @@ class DynamicVamana : public manager::IndexManager { DataLoader&& data_loader, const Distance& distance, ThreadPoolProto threadpool_proto, - bool debug_load_from_static = false + bool debug_load_from_static = false, + svs::logging::logger_ptr logger = svs::logging::get() ) { return DynamicVamana( AssembleTag(), @@ -358,7 +377,8 @@ class DynamicVamana : public manager::IndexManager { std::forward(data_loader), distance, threads::as_threadpool(std::move(threadpool_proto)), - debug_load_from_static + debug_load_from_static, + std::move(logger) ) ); } diff --git a/include/svs/orchestrators/vamana.h b/include/svs/orchestrators/vamana.h index 3d5058553..6f27d6728 100644 --- a/include/svs/orchestrators/vamana.h +++ b/include/svs/orchestrators/vamana.h @@ -421,6 +421,8 @@ class Vamana : public manager::IndexManager { /// instance or an integer specifying the number of threads to use. In the latter /// case, a new default thread pool will be constructed using ``threadpool_proto`` /// as the number of threads to create. + /// @param logger The logger to use for this index. Defaults to the global SVS logger + /// (``svs::logging::get()``). /// /// The data loader should be any object loadable via ``svs::detail::dispatch_load`` /// returning a Vamana compatible dataset. Concrete examples include: @@ -444,7 +446,8 @@ class Vamana : public manager::IndexManager { const GraphLoaderType& graph_loader, DataLoader&& data_loader, const Distance& distance, - ThreadPoolProto threadpool_proto + ThreadPoolProto threadpool_proto, + svs::logging::logger_ptr logger = svs::logging::get() ) { // If given an `enum` for the distance type, than we need to dispatch over that // enum. @@ -460,7 +463,8 @@ class Vamana : public manager::IndexManager { graph_loader, std::forward(data_loader), distance_function, - std::move(threadpool) + std::move(threadpool), + std::move(logger) ); }); } else { @@ -470,7 +474,8 @@ class Vamana : public manager::IndexManager { graph_loader, std::forward(data_loader), distance, - std::move(threadpool) + std::move(threadpool), + std::move(logger) ); } } @@ -579,6 +584,8 @@ class Vamana : public manager::IndexManager { /// case, a new default thread pool will be constructed using ``threadpool_proto`` /// as the number of threads to create. /// @param graph_allocator The allocator to use for the backing graph. + /// @param logger The logger to use for this index. Defaults to the global SVS logger + /// (``svs::logging::get()``). /// /// The data loader should be any object loadable via ``svs::detail::dispatch_load`` /// returning a Vamana compatible dataset. Concrete examples include: @@ -601,7 +608,8 @@ class Vamana : public manager::IndexManager { DataLoader&& data_loader, Distance distance, ThreadPoolProto threadpool_proto = 1, - const Allocator& graph_allocator = {} + const Allocator& graph_allocator = {}, + svs::logging::logger_ptr logger = svs::logging::get() ) { auto threadpool = threads::as_threadpool(std::move(threadpool_proto)); if constexpr (std::is_same_v, DistanceType>) { @@ -613,7 +621,8 @@ class Vamana : public manager::IndexManager { std::forward(data_loader), std::move(distance_function), std::move(threadpool), - graph_allocator + graph_allocator, + std::move(logger) ); }); } else { @@ -623,7 +632,8 @@ class Vamana : public manager::IndexManager { std::forward(data_loader), distance, std::move(threadpool), - graph_allocator + graph_allocator, + std::move(logger) ); } } diff --git a/tests/svs/orchestrators/dynamic_vamana.cpp b/tests/svs/orchestrators/dynamic_vamana.cpp index 10951e7ca..cb5f89844 100644 --- a/tests/svs/orchestrators/dynamic_vamana.cpp +++ b/tests/svs/orchestrators/dynamic_vamana.cpp @@ -29,16 +29,59 @@ #include "tests/utils/utils.h" #include "tests/utils/vamana_reference.h" +// Logging +#include "spdlog/sinks/callback_sink.h" +#include "svs/core/logging.h" + // Catch2 #include "catch2/catch_approx.hpp" #include "catch2/catch_test_macros.hpp" // STL +#include +#include #include +#include +#include #include namespace { +// Logger capture helpers for the per-index logger tests. +struct CapturingLogger { + std::shared_ptr> messages = + std::make_shared>(); + svs::logging::logger_ptr logger; + + explicit CapturingLogger(const std::string& name) { + auto sink = std::make_shared( + [messages = messages](const spdlog::details::log_msg& msg) { + messages->emplace_back(msg.payload.data(), msg.payload.size()); + } + ); + sink->set_level(spdlog::level::trace); + logger = std::make_shared(name, std::move(sink)); + logger->set_level(spdlog::level::trace); + } + + bool contains(std::string_view needle) const { + return std::any_of(messages->begin(), messages->end(), [&](const auto& m) { + return m.find(needle) != std::string::npos; + }); + } +}; + +// Temporarily replace the global SVS logger; restore it on scope exit. +struct GlobalLoggerGuard { + svs::logging::logger_ptr original = svs::logging::get(); + explicit GlobalLoggerGuard(const svs::logging::logger_ptr& replacement) { + svs::logging::set(replacement); + } + GlobalLoggerGuard(const GlobalLoggerGuard&) = delete; + GlobalLoggerGuard& operator=(const GlobalLoggerGuard&) = delete; + ~GlobalLoggerGuard() { svs::logging::set(original); } +}; + template void test_build(DataLoaderT&& data_loader, DistanceT distance = DistanceT()) { auto expected_result = test_dataset::vamana::expected_build_results( @@ -162,3 +205,110 @@ CATCH_TEST_CASE("DynamicVamana Memory Usage", "[managers][dynamic_vamana]") { const size_t usage_after = index.get_memory_breakdown().total(); CATCH_REQUIRE(usage_after > usage_before); } + +CATCH_TEST_CASE("DynamicVamana Per-Index Logger", "[managers][dynamic_vamana][logging]") { + auto distance = svs::distance::DistanceL2(); + auto expected_result = test_dataset::vamana::expected_build_results( + distance, svsbenchmark::Uncompressed(svs::DataType::float32) + ); + auto build_params = expected_result.build_parameters_.value(); + size_t num_threads = 2; + + auto data = svs::data::SimpleData::load(test_dataset::data_svs_file()); + const size_t n = data.size(); + const size_t half = n / 2; + auto first_data = svs::data::SimpleData(half, data.dimensions()); + for (size_t i = 0; i < half; ++i) { + first_data.set_datum(i, data.get_datum(i)); + } + auto second_data = svs::data::SimpleData(n - half, data.dimensions()); + for (size_t i = 0; i < n - half; ++i) { + second_data.set_datum(i, data.get_datum(half + i)); + } + std::vector first_ids(half); + std::iota(first_ids.begin(), first_ids.end(), 0); + std::vector second_ids(n - half); + std::iota(second_ids.begin(), second_ids.end(), half); + + CATCH_SECTION("Build with custom logger") { + auto global = CapturingLogger("global_logger"); + auto guard = GlobalLoggerGuard(global.logger); + auto custom = CapturingLogger("custom_logger"); + + svs::DynamicVamana index = svs::DynamicVamana::build( + build_params, + first_data, + first_ids, + distance, + num_threads, + svs::data::Blocked>(), + custom.logger + ); + CATCH_REQUIRE(custom.logger.use_count() == 2); + CATCH_REQUIRE(custom.contains("Vamana Build Parameters:")); + CATCH_REQUIRE(custom.contains("Number of syncs:")); + CATCH_REQUIRE(global.messages->empty()); + } + + CATCH_SECTION("Build without logger uses the global logger") { + auto global = CapturingLogger("global_logger"); + auto guard = GlobalLoggerGuard(global.logger); + auto baseline = global.logger.use_count(); + + svs::DynamicVamana index = svs::DynamicVamana::build( + build_params, first_data, first_ids, distance, num_threads + ); + CATCH_REQUIRE(global.logger.use_count() == baseline + 1); + CATCH_REQUIRE(global.contains("Vamana Build Parameters:")); + } + + CATCH_SECTION("Assemble with and without custom logger") { + auto tempdir = svs_test::prepare_temp_directory_v2(); + auto config_dir = tempdir / "config"; + auto graph_dir = tempdir / "graph"; + auto data_dir = tempdir / "data"; + { + svs::DynamicVamana index = svs::DynamicVamana::build( + build_params, first_data, first_ids, distance, num_threads + ); + index.save(config_dir, graph_dir, data_dir); + } + + auto global = CapturingLogger("global_logger"); + auto guard = GlobalLoggerGuard(global.logger); + auto custom = CapturingLogger("custom_logger"); + auto global_baseline = global.logger.use_count(); + + // `debug_load_from_static` precedes the logger and must be passed explicitly. + svs::DynamicVamana with_custom = svs::DynamicVamana::assemble( + config_dir, + svs::GraphLoader(graph_dir), + svs::VectorDataLoader(data_dir), + distance, + num_threads, + false, + custom.logger + ); + CATCH_REQUIRE(custom.logger.use_count() == 2); + CATCH_REQUIRE(global.logger.use_count() == global_baseline); + CATCH_REQUIRE(with_custom.size() == half); + + // Adding points runs the graph builder, which logs through the index's logger. + with_custom.add_points(second_data.cview(), second_ids); + CATCH_REQUIRE(custom.contains("Vamana Build Parameters:")); + CATCH_REQUIRE(global.messages->empty()); + + custom.messages->clear(); + svs::DynamicVamana with_default = svs::DynamicVamana::assemble( + config_dir, + svs::GraphLoader(graph_dir), + svs::VectorDataLoader(data_dir), + distance, + num_threads + ); + CATCH_REQUIRE(global.logger.use_count() == global_baseline + 1); + with_default.add_points(second_data.cview(), second_ids); + CATCH_REQUIRE(global.contains("Vamana Build Parameters:")); + CATCH_REQUIRE(custom.messages->empty()); + } +} diff --git a/tests/svs/orchestrators/vamana.cpp b/tests/svs/orchestrators/vamana.cpp index fdab19de2..0ffd8b479 100644 --- a/tests/svs/orchestrators/vamana.cpp +++ b/tests/svs/orchestrators/vamana.cpp @@ -29,16 +29,59 @@ #include "tests/utils/utils.h" #include "tests/utils/vamana_reference.h" +// Logging +#include "spdlog/sinks/callback_sink.h" +#include "svs/core/logging.h" + // Catch2 #include "catch2/catch_approx.hpp" #include "catch2/catch_test_macros.hpp" // STL +#include +#include #include +#include +#include #include namespace { +// Logger capture helpers for the per-index logger tests. +struct CapturingLogger { + std::shared_ptr> messages = + std::make_shared>(); + svs::logging::logger_ptr logger; + + explicit CapturingLogger(const std::string& name) { + auto sink = std::make_shared( + [messages = messages](const spdlog::details::log_msg& msg) { + messages->emplace_back(msg.payload.data(), msg.payload.size()); + } + ); + sink->set_level(spdlog::level::trace); + logger = std::make_shared(name, std::move(sink)); + logger->set_level(spdlog::level::trace); + } + + bool contains(std::string_view needle) const { + return std::any_of(messages->begin(), messages->end(), [&](const auto& m) { + return m.find(needle) != std::string::npos; + }); + } +}; + +// Temporarily replace the global SVS logger; restore it on scope exit. +struct GlobalLoggerGuard { + svs::logging::logger_ptr original = svs::logging::get(); + explicit GlobalLoggerGuard(const svs::logging::logger_ptr& replacement) { + svs::logging::set(replacement); + } + GlobalLoggerGuard(const GlobalLoggerGuard&) = delete; + GlobalLoggerGuard& operator=(const GlobalLoggerGuard&) = delete; + ~GlobalLoggerGuard() { svs::logging::set(original); } +}; + template void test_build(DataLoaderT&& data_loader, DistanceT distance = DistanceT()) { auto expected_result = test_dataset::vamana::expected_build_results( @@ -143,3 +186,97 @@ CATCH_TEST_CASE("Vamana Memory Usage", "[managers][vamana]") { CATCH_REQUIRE(half_usage > 0); CATCH_REQUIRE(full_usage > half_usage); } + +CATCH_TEST_CASE("Vamana Per-Index Logger", "[managers][vamana][logging]") { + auto distance = svs::distance::DistanceL2(); + auto expected_result = test_dataset::vamana::expected_build_results( + distance, svsbenchmark::Uncompressed(svs::DataType::float32) + ); + auto build_params = expected_result.build_parameters_.value(); + size_t num_threads = 2; + + CATCH_SECTION("Build with custom logger") { + auto global = CapturingLogger("global_logger"); + auto guard = GlobalLoggerGuard(global.logger); + auto custom = CapturingLogger("custom_logger"); + + svs::Vamana index = svs::Vamana::build( + build_params, + svs::data::SimpleData::load(test_dataset::data_svs_file()), + distance, + num_threads, + svs::HugepageAllocator(), + custom.logger + ); + // The index holds its own reference to the custom logger. + CATCH_REQUIRE(custom.logger.use_count() == 2); + CATCH_REQUIRE(custom.contains("Vamana Build Parameters:")); + CATCH_REQUIRE(custom.contains("Number of syncs:")); + CATCH_REQUIRE(global.messages->empty()); + } + + CATCH_SECTION("Build without logger uses the global logger") { + auto global = CapturingLogger("global_logger"); + auto guard = GlobalLoggerGuard(global.logger); + auto baseline = global.logger.use_count(); + + svs::Vamana index = svs::Vamana::build( + build_params, + svs::data::SimpleData::load(test_dataset::data_svs_file()), + distance, + num_threads + ); + CATCH_REQUIRE(global.logger.use_count() == baseline + 1); + CATCH_REQUIRE(global.contains("Vamana Build Parameters:")); + } + + CATCH_SECTION("Assemble with and without custom logger") { + auto tempdir = svs_test::prepare_temp_directory_v2(); + auto config_dir = tempdir / "config"; + auto graph_dir = tempdir / "graph"; + auto data_dir = tempdir / "data"; + { + svs::Vamana index = svs::Vamana::build( + build_params, + svs::data::SimpleData::load(test_dataset::data_svs_file()), + distance, + num_threads + ); + index.save(config_dir, graph_dir, data_dir); + } + + auto global = CapturingLogger("global_logger"); + auto guard = GlobalLoggerGuard(global.logger); + auto custom = CapturingLogger("custom_logger"); + + // Static assembly emits no log messages, so check that the index retains a + // reference to the logger it was given. + auto global_baseline = global.logger.use_count(); + svs::Vamana with_custom = svs::Vamana::assemble( + config_dir, + svs::GraphLoader(graph_dir), + svs::VectorDataLoader(data_dir), + svs::DistanceType::L2, + num_threads, + custom.logger + ); + CATCH_REQUIRE(custom.logger.use_count() == 2); + CATCH_REQUIRE(global.logger.use_count() == global_baseline); + CATCH_REQUIRE( + with_custom.size() == + svs::data::SimpleData::load(test_dataset::data_svs_file()).size() + ); + + svs::Vamana with_default = svs::Vamana::assemble( + config_dir, + svs::GraphLoader(graph_dir), + svs::VectorDataLoader(data_dir), + distance, + num_threads + ); + CATCH_REQUIRE(global.logger.use_count() == global_baseline + 1); + CATCH_REQUIRE(custom.logger.use_count() == 2); + CATCH_REQUIRE(global.messages->empty()); + CATCH_REQUIRE(custom.messages->empty()); + } +}