Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 25 additions & 5 deletions include/svs/orchestrators/dynamic_vamana.h
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,8 @@ class DynamicVamana : public manager::IndexManager<DynamicVamanaInterface> {
/// @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,
Expand All @@ -301,7 +303,8 @@ class DynamicVamana : public manager::IndexManager<DynamicVamanaInterface> {
std::span<const size_t> 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 =
Expand All @@ -317,7 +320,8 @@ class DynamicVamana : public manager::IndexManager<DynamicVamanaInterface> {
ids,
std::move(distance_function),
std::move(threadpool),
graph_allocator
graph_allocator,
std::move(logger)
);
});
} else {
Expand All @@ -328,12 +332,26 @@ class DynamicVamana : public manager::IndexManager<DynamicVamanaInterface> {
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,
Expand All @@ -347,7 +365,8 @@ class DynamicVamana : public manager::IndexManager<DynamicVamanaInterface> {
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(),
Expand All @@ -358,7 +377,8 @@ class DynamicVamana : public manager::IndexManager<DynamicVamanaInterface> {
std::forward<DataLoader>(data_loader),
distance,
threads::as_threadpool(std::move(threadpool_proto)),
debug_load_from_static
debug_load_from_static,
std::move(logger)
)
);
}
Expand Down
22 changes: 16 additions & 6 deletions include/svs/orchestrators/vamana.h
Original file line number Diff line number Diff line change
Expand Up @@ -421,6 +421,8 @@ class Vamana : public manager::IndexManager<VamanaInterface> {
/// 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:
Expand All @@ -444,7 +446,8 @@ class Vamana : public manager::IndexManager<VamanaInterface> {
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.
Expand All @@ -460,7 +463,8 @@ class Vamana : public manager::IndexManager<VamanaInterface> {
graph_loader,
std::forward<DataLoader>(data_loader),
distance_function,
std::move(threadpool)
std::move(threadpool),
std::move(logger)
);
});
} else {
Expand All @@ -470,7 +474,8 @@ class Vamana : public manager::IndexManager<VamanaInterface> {
graph_loader,
std::forward<DataLoader>(data_loader),
distance,
std::move(threadpool)
std::move(threadpool),
std::move(logger)
);
}
}
Expand Down Expand Up @@ -579,6 +584,8 @@ class Vamana : public manager::IndexManager<VamanaInterface> {
/// 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:
Expand All @@ -601,7 +608,8 @@ class Vamana : public manager::IndexManager<VamanaInterface> {
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<std::decay_t<Distance>, DistanceType>) {
Expand All @@ -613,7 +621,8 @@ class Vamana : public manager::IndexManager<VamanaInterface> {
std::forward<DataLoader>(data_loader),
std::move(distance_function),
std::move(threadpool),
graph_allocator
graph_allocator,
std::move(logger)
);
});
} else {
Expand All @@ -623,7 +632,8 @@ class Vamana : public manager::IndexManager<VamanaInterface> {
std::forward<DataLoader>(data_loader),
distance,
std::move(threadpool),
graph_allocator
graph_allocator,
std::move(logger)
);
}
}
Expand Down
150 changes: 150 additions & 0 deletions tests/svs/orchestrators/dynamic_vamana.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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 <algorithm>
#include <memory>
#include <numeric>
#include <string>
#include <string_view>
#include <vector>

namespace {

// Logger capture helpers for the per-index logger tests.
struct CapturingLogger {
std::shared_ptr<std::vector<std::string>> messages =
std::make_shared<std::vector<std::string>>();
svs::logging::logger_ptr logger;

explicit CapturingLogger(const std::string& name) {
auto sink = std::make_shared<spdlog::sinks::callback_sink_mt>(
[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<spdlog::logger>(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 <typename DataLoaderT, typename DistanceT>
void test_build(DataLoaderT&& data_loader, DistanceT distance = DistanceT()) {
auto expected_result = test_dataset::vamana::expected_build_results(
Expand Down Expand Up @@ -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<float>::load(test_dataset::data_svs_file());
const size_t n = data.size();
const size_t half = n / 2;
auto first_data = svs::data::SimpleData<float>(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<float>(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<size_t> first_ids(half);
std::iota(first_ids.begin(), first_ids.end(), 0);
std::vector<size_t> 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<float>(
build_params,
first_data,
first_ids,
distance,
num_threads,
svs::data::Blocked<svs::HugepageAllocator<uint32_t>>(),
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<float>(
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<float>(
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<float>(
config_dir,
svs::GraphLoader(graph_dir),
svs::VectorDataLoader<float>(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<float>(
config_dir,
svs::GraphLoader(graph_dir),
svs::VectorDataLoader<float>(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());
}
}
Loading
Loading