From b497012b870ce057bd393ae622c8f26e92996273 Mon Sep 17 00:00:00 2001 From: Alexandr Guzhva Date: Mon, 28 Sep 2026 18:06:51 -0400 Subject: [PATCH 1/3] feat: add SINDI query-mass pruning and FP16 candidate refinement Signed-off-by: Alexandr Guzhva --- benchmark/CMakeLists.txt | 13 + benchmark/benchmark_sparse.cpp | 581 +++++++++++++++++++++++ src/index/sparse/inverted_index.h | 7 + src/index/sparse/inverted_index_format.h | 1 + src/index/sparse/sindi_inverted_index.h | 260 +++++++++- src/index/sparse/sindi_refinement.h | 100 ++++ src/index/sparse/sparse_index_config.h | 47 +- src/index/sparse/sparse_index_node.cc | 36 +- tests/ut/test_sparse.cc | 289 +++++++++++ 9 files changed, 1319 insertions(+), 15 deletions(-) create mode 100644 benchmark/benchmark_sparse.cpp create mode 100644 src/index/sparse/sindi_refinement.h diff --git a/benchmark/CMakeLists.txt b/benchmark/CMakeLists.txt index 2c70642d6..6f9372ccc 100644 --- a/benchmark/CMakeLists.txt +++ b/benchmark/CMakeLists.txt @@ -84,3 +84,16 @@ if(CMAKE_SYSTEM_PROCESSOR MATCHES "^(x86_64|AMD64|amd64|X86_64)$") else() message(STATUS "Skipping sparse SIMD benchmark on ${CMAKE_SYSTEM_PROCESSOR} (x86_64 only)") endif() + +if(CMAKE_SYSTEM_NAME STREQUAL "Linux" AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$") +# Sparse CSR baseline/refinement experiment; all scoring uses the Knowhere API. +add_executable(benchmark_sparse benchmark_sparse.cpp) +execute_process(COMMAND git rev-parse HEAD WORKING_DIRECTORY ${CMAKE_SOURCE_DIR} + OUTPUT_VARIABLE sparse_benchmark_revision OUTPUT_STRIP_TRAILING_WHITESPACE) +target_compile_definitions(benchmark_sparse PRIVATE KNOWHERE_SPARSE_BENCHMARK_REVISION="${sparse_benchmark_revision}") +target_link_libraries(benchmark_sparse knowhere milvus-common::milvus-common ${CMAKE_DL_LIBS}) +if(NOT APPLE) + target_link_options(benchmark_sparse PRIVATE "LINKER:--no-as-needed") + target_link_libraries(benchmark_sparse atomic) +endif() +endif() diff --git a/benchmark/benchmark_sparse.cpp b/benchmark/benchmark_sparse.cpp new file mode 100644 index 000000000..64fdef8e0 --- /dev/null +++ b/benchmark/benchmark_sparse.cpp @@ -0,0 +1,581 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy at http://www.apache.org/licenses/LICENSE-2.0 + +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "knowhere/comp/brute_force.h" +#include "knowhere/comp/knowhere_config.h" +#include "knowhere/index/index_factory.h" +#include "knowhere/sparse_utils.h" +#include "knowhere/version.h" +#include "src/index/sparse/sindi_refinement.h" +#include "src/index/sparse/sindi_simd.h" + +namespace { +using Clock = std::chrono::steady_clock; +using Row = knowhere::sparse::SparseRow; +using knowhere::Json; + +void +Check(bool ok, std::string_view message) { + if (!ok) { + throw std::runtime_error(std::string(message)); + } +} + +double +Seconds(Clock::time_point start) { + return std::chrono::duration(Clock::now() - start).count(); +} + +// The challenge files are little-endian, structure-of-arrays CSR / k-NN files. +class MappedFile { + public: + explicit MappedFile(const std::string& path) { + Check(std::endian::native == std::endian::little, "Only little-endian hosts are supported"); + int fd = open(path.c_str(), O_RDONLY); + Check(fd >= 0, "Cannot open " + path); + struct stat st {}; + if (fstat(fd, &st) != 0 || st.st_size <= 0) { + close(fd); + throw std::runtime_error("Cannot stat or empty file: " + path); + } + size = st.st_size; + void* ptr = mmap(nullptr, size, PROT_READ, MAP_PRIVATE, fd, 0); + close(fd); + Check(ptr != MAP_FAILED, "Cannot mmap " + path); + data = static_cast(ptr); + } + ~MappedFile() { + munmap(const_cast(data), size); + } + MappedFile(const MappedFile&) = delete; + MappedFile& + operator=(const MappedFile&) = delete; + + template + T + Read(uint64_t offset) const { + Check(offset <= size && sizeof(T) <= size - offset, "Truncated binary file"); + T value; + std::memcpy(&value, data + offset, sizeof(T)); + return value; + } + const char* data = nullptr; + uint64_t size = 0; +}; + +struct SparseData { + uint64_t file_rows = 0, dim = 0, nnz = 0; + std::vector rows; + + knowhere::DataSetPtr + Dataset() const { + auto ds = knowhere::GenDataSet(rows.size(), dim, rows.data()); + ds->SetIsSparse(true); + return ds; + } +}; + +SparseData +LoadCsr(const std::string& path, const std::vector* selection, uint64_t limit = 0) { + MappedFile f(path); + SparseData out; + out.file_rows = f.Read(0); + out.dim = f.Read(8); + const auto nnz = f.Read(16); + Check(f.size >= 32 && out.file_rows > 0 && out.file_rows <= (f.size - 32) / 8 && out.dim > 0 && + out.dim <= std::numeric_limits::max(), + "Invalid CSR dimensions: " + path); + const uint64_t indices_offset = 24 + 8 * (out.file_rows + 1); + Check(nnz <= (f.size - indices_offset) / 8 && indices_offset + 8 * nnz == f.size, + "CSR file size does not match header: " + path); + const uint64_t values_offset = indices_offset + 4 * nnz; + Check(f.Read(24) == 0 && f.Read(24 + 8 * out.file_rows) == nnz, "Invalid CSR endpoint offsets"); + uint64_t previous = 0; + for (uint64_t i = 1; i <= out.file_rows; ++i) { + const auto next = f.Read(24 + 8 * i); + Check(next >= previous && next <= nnz, "Invalid CSR row offsets"); + previous = next; + } + const auto count = selection ? selection->size() : (limit ? std::min(limit, out.file_rows) : out.file_rows); + out.rows.reserve(count); + for (uint64_t i = 0; i < count; ++i) { + const auto id = selection ? selection->at(i) : i; + Check(id < out.file_rows, "Query ID outside CSR file"); + const auto start = f.Read(24 + 8 * id); + const auto end = f.Read(24 + 8 * (id + 1)); + out.rows.emplace_back(end - start); + uint32_t last = 0; + for (auto j = start; j < end; ++j) { + const auto col = f.Read(indices_offset + 4 * j); + const auto value = f.Read(values_offset + 4 * j); + Check(col < out.dim && (j == start || col > last), "CSR columns must be sorted and unique"); + Check(std::isfinite(value) && value >= 0, "Expected finite nonnegative SPLADE weights"); + out.rows.back().set_at(j - start, col, value); + last = col; + } + out.nnz += end - start; + } + return out; +} + +struct Truth { + std::vector ids; + std::vector scores; +}; + +Truth +LoadTruth(const std::string& path, const std::vector& selection, uint64_t nq, uint64_t nb, int k) { + MappedFile f(path); + const uint64_t rows = f.Read(0), width = f.Read(4); + Check(rows == nq && width >= static_cast(k) && f.size >= 8 && rows > 0 && + width <= (f.size - 8) / 8 / rows && 8 + rows * width * 8 == f.size, + "Invalid ground-truth header or unsupported k"); + Truth out; + for (auto q : selection) { + std::set unique; + float previous = std::numeric_limits::infinity(); + for (int j = 0; j < k; ++j) { + auto id = f.Read(8 + 4 * (q * width + j)); + auto score = f.Read(8 + 4 * rows * width + 4 * (q * width + j)); + Check( + id < nb && unique.insert(id).second && std::isfinite(score) && + score <= previous + 4 * std::numeric_limits::epsilon() * std::max(1.0f, std::abs(previous)), + "Invalid ground-truth IDs or score order"); + out.ids.push_back(id); + out.scores.push_back(score); + previous = score; + } + } + return out; +} + +struct Options { + std::string data, output = "sparse_results.csv", method = "all"; + uint64_t queries = 100, offset = 0, base_limit = 0; + bool refine = false; + float mass = 1, factor = 1, drop = 0; + std::string sweep, split = "dev"; + int threads = std::max(2u, std::thread::hardware_concurrency()), repeats = 3, k = 10; + std::optional seed; +}; + +Options +Parse(int argc, char** argv) { + Options o; + for (int i = 1; i < argc; ++i) { + const std::string key = argv[i]; + if (key == "--help") { + std::cout + << "benchmark_sparse --data-dir DIR [--queries 100] [--query-offset 0] [--seed N]\n" + " [--threads N] [--repeats 3] [--k 10] [--output sparse_results.csv]\n" + " [--refine 0|1] [--mass 1] [--factor 1] [--drop 0]\n" + " [--sweep mass|count|drop] [--split dev|hidden]\n" + " [--method " + "all|TAAT_NAIVE|DAAT_WAND|DAAT_MAXSCORE|BLOCK_MAX_MAXSCORE|BLOCK_MAX_WAND|SINDI|SPARSE_WAND]\n" + " [--base-limit N] Reduced-base validation against brute force, not supplied ground truth.\n"; + std::exit(0); + } + Check(i + 1 < argc, "Missing value for " + key); + const std::string value = argv[++i]; + if (key == "--data-dir") + o.data = value; + else if (key == "--output") + o.output = value; + else if (key == "--refine") + o.refine = std::stoi(value) != 0; + else if (key == "--mass") + o.mass = std::stof(value); + else if (key == "--factor") + o.factor = std::stof(value); + else if (key == "--drop") + o.drop = std::stof(value); + else if (key == "--sweep") + o.sweep = value; + else if (key == "--split") + o.split = value; + else if (key == "--method") + o.method = value; + else { + Check(!value.empty() && value.find_first_not_of("0123456789") == std::string::npos, + "Expected unsigned integer for " + key); + const uint64_t n = std::stoull(value); + if (key == "--queries") + o.queries = n; + else if (key == "--query-offset") + o.offset = n; + else if (key == "--base-limit") + o.base_limit = n; + else if (key == "--seed") { + Check(n <= UINT32_MAX, "Seed out of range"); + o.seed = n; + } else { + Check(n > 0 && n <= INT32_MAX, "Option out of range: " + key); + if (key == "--threads") + o.threads = n; + else if (key == "--repeats") + o.repeats = n; + else if (key == "--k") + o.k = n; + else + throw std::runtime_error("Unknown option: " + key); + } + } + } + Check(!o.data.empty() && o.queries > 0 && o.threads >= 2, "Provide --data-dir, positive queries, and >=2 threads"); + return o; +} + +// ID recall deliberately exposes ties and FP16 rounding relative to FP32 ground truth. +double +Recall(const knowhere::DataSetPtr& result, const Truth& truth, size_t nq, int k, size_t nb) { + size_t hits = 0, total = 0; + const auto ids = result->GetIds(); + const auto scores = result->GetDistance(); + for (size_t q = 0; q < nq; ++q) { + std::set seen; + float previous = std::numeric_limits::infinity(); + bool padding = false; + for (int j = 0; j < k; ++j) { + const auto pos = q * k + j; + if (truth.ids[pos] >= 0) + ++total; + if (ids[pos] == -1) { + padding = true; + continue; + } + Check(!padding && ids[pos] >= 0 && static_cast(ids[pos]) < nb && seen.insert(ids[pos]).second && + std::isfinite(scores[pos]) && scores[pos] <= previous, + "Invalid search result IDs or score order"); + previous = scores[pos]; + auto begin = truth.ids.begin() + q * k; + hits += std::find(begin, begin + k, ids[pos]) != begin + k; + } + } + return total ? double(hits) / total : 1.0; +} + +Json +RecallMismatches(const knowhere::DataSetPtr& result, const Truth& truth, const std::vector& selected, int k) { + Json mismatches = Json::array(); + for (size_t q = 0; q < selected.size(); ++q) { + const auto offset = q * k; + std::set expected(truth.ids.begin() + offset, truth.ids.begin() + offset + k); + std::set actual(result->GetIds() + offset, result->GetIds() + offset + k); + if (expected != actual) { + mismatches.push_back( + {{"query_id", selected[q]}, + {"expected_ids", std::vector(truth.ids.begin() + offset, truth.ids.begin() + offset + k)}, + {"expected_scores", + std::vector(truth.scores.begin() + offset, truth.scores.begin() + offset + k)}, + {"actual_ids", std::vector(result->GetIds() + offset, result->GetIds() + offset + k)}, + {"actual_scores", + std::vector(result->GetDistance() + offset, result->GetDistance() + offset + k)}}); + } + } + return mismatches; +} +} // namespace + +int +main(int argc, char** argv) { + try { + const auto o = Parse(argc, argv); + const auto& ip_kernels = knowhere::sparse::inverted::sindi::get_ip_kernels(); + Dl_info accumulate_info{}, insert_info{}; + Check(dladdr(reinterpret_cast(ip_kernels.accumulate), &accumulate_info) != 0 && + accumulate_info.dli_sname != nullptr, + "Cannot identify IP kernel"); + Check(dladdr(reinterpret_cast(ip_kernels.batch_insert), &insert_info) != 0 && + insert_info.dli_sname != nullptr, + "Cannot identify selection kernel"); + Check(std::string(accumulate_info.dli_sname).find("sve") != std::string::npos && + std::string(insert_info.dli_sname).find("sve") != std::string::npos, + "Baseline requires SVE kernels"); + const int sve_vl = prctl(PR_SVE_GET_VL); + Check(sve_vl > 0, "Cannot read SVE vector length"); + cpu_set_t affinity; + CPU_ZERO(&affinity); + Check(sched_getaffinity(0, sizeof(affinity), &affinity) == 0, "Cannot read CPU affinity"); + std::vector cpus; + for (int i = 0; i < CPU_SETSIZE; ++i) + if (CPU_ISSET(i, &affinity)) + cpus.push_back(i); + std::cerr << "IP kernel=" << accumulate_info.dli_sname << " SVE bytes=" << (sve_vl & PR_SVE_VL_LEN_MASK) + << std::endl; + const std::vector methods = {"TAAT_NAIVE", "DAAT_WAND", "DAAT_MAXSCORE", "BLOCK_MAX_MAXSCORE", + "BLOCK_MAX_WAND", "SINDI", "SPARSE_WAND"}; + Check(o.method == "all" || std::find(methods.begin(), methods.end(), o.method) != methods.end(), + "Unknown method: " + o.method); + knowhere::KnowhereConfig::SetBuildThreadPoolSize(o.threads); + knowhere::KnowhereConfig::SetSearchThreadPoolSize(o.threads); + const auto version = knowhere::Version::GetMaximumVersion().VersionNumber(); + Check(o.split == "dev" || o.split == "hidden", "Invalid split"); + const auto query_path = o.data + "/queries." + o.split + ".csr"; + uint64_t file_queries; + { + MappedFile query_file(query_path); + file_queries = query_file.Read(0); + Check(file_queries <= query_file.size / 8, "Invalid query count"); + } + Check(o.offset < file_queries && o.queries <= file_queries - o.offset, "Query selection outside dataset"); + std::vector selected(file_queries - o.offset); + std::iota(selected.begin(), selected.end(), o.offset); + if (o.seed) { + std::mt19937 rng(*o.seed); + std::shuffle(selected.begin(), selected.end(), rng); + } + selected.resize(o.queries); + auto queries = LoadCsr(query_path, &selected); + auto load_start = Clock::now(); + std::cerr << "Loading base vectors..." << std::endl; + auto base = LoadCsr(o.data + "/base_full.csr", nullptr, o.base_limit); + Check(base.rows.size() >= static_cast(o.k), "Base smaller than k"); + auto load_seconds = Seconds(load_start); + // Allow a query file with a smaller declared vocabulary; use a shared dimension. + base.dim = queries.dim = std::max(base.dim, queries.dim); + auto base_ds = base.Dataset(), query_ds = queries.Dataset(); + Json search = {{"metric_type", "IP"}, {"k", o.k}, {"drop_ratio_search", 0.0}, + {"dim_max_score_ratio", 1.05}, {"refine_factor", 1}, {"search_algo", "INHERIT"}}; + Truth truth; + if (o.base_limit) { + std::cerr << "Computing reduced-base FP32 brute-force ground truth..." << std::endl; + auto gt = knowhere::BruteForce::SearchSparse(base_ds, query_ds, search, nullptr); + Check(gt.has_value(), "Brute force failed: " + gt.what()); + truth.ids.assign(gt.value()->GetIds(), gt.value()->GetIds() + o.queries * o.k); + truth.scores.assign(gt.value()->GetDistance(), gt.value()->GetDistance() + o.queries * o.k); + } else { + truth = LoadTruth(o.data + "/base_full." + o.split + ".gt", selected, file_queries, base.file_rows, o.k); + } + struct utsname host {}; + uname(&host); + Json metadata = { + {"data_dir", std::filesystem::absolute(o.data).string()}, + {"base_rows", base.rows.size()}, + {"base_nnz", base.nnz}, + {"dimension", base.dim}, + {"query_ids", selected}, + {"query_nnz", queries.nnz}, + {"threads", o.threads}, + {"repeats", o.repeats}, + {"index_version", version}, + {"quant_type", "fp16"}, + {"search_config", search}, + {"base_load_seconds", load_seconds}, + {"machine", host.machine}, + {"hostname", host.nodename}, + {"kernel", host.release}, + {"compiler", __VERSION__}, + {"source_revision", KNOWHERE_SPARSE_BENCHMARK_REVISION}, + {"ground_truth", o.base_limit ? "reduced-base FP32 brute force" : "base_full." + o.split + ".gt"}, + {"warmup_batches", 1}, + {"runs", Json::array()}}; + metadata["effective_ip_kernel"] = accumulate_info.dli_sname; + metadata["effective_selection_kernel"] = insert_info.dli_sname; + metadata["sve_vector_bytes"] = sve_vl & PR_SVE_VL_LEN_MASK; + metadata["cpu_affinity"] = cpus; + std::ifstream cpu("/proc/cpuinfo"), mem("/proc/meminfo"); + metadata["cpuinfo"] = std::string(std::istreambuf_iterator(cpu), {}); + metadata["meminfo"] = std::string(std::istreambuf_iterator(mem), {}); + std::ofstream csv(o.output); + Check(csv.good(), "Cannot write " + o.output); + csv << "method,index_type,index_version,quant_type,codec,base_rows,queries,k,threads,build_seconds,index_bytes," + "repeat,search_seconds,qps,recall,mass,refine_k,drop_ratio,refine\n" + << std::setprecision(10); + for (const auto& method : methods) { + if (o.method != "all" && o.method != method) + continue; + const bool alias = method == "SPARSE_WAND"; + const auto type = + alias ? knowhere::IndexEnum::INDEX_SPARSE_WAND : knowhere::IndexEnum::INDEX_SPARSE_INVERTED_INDEX; + const std::string algo = alias ? "DAAT_WAND" : method; + Json build = { + {"dim", base.dim}, {"metric_type", "IP"}, {"inverted_index_algo", algo}, {"quant_type", "fp16"}}; + const std::string codec = algo == "SINDI" ? "fixed_docid_windows" : "block_streamvbyte"; + if (algo == "SINDI") { + build["sindi_window_size"] = 4096; + build["refine"] = o.refine; + } + if (algo != "SINDI") + build["inverted_index_codec"] = codec; + std::cerr << "Building " << method << "..." << std::endl; + auto created = knowhere::IndexFactory::Instance().Create(type, version); + Check(created.has_value(), "Create failed: " + created.what()); + auto index = std::move(created.value()); + auto start = Clock::now(); + auto status = index.Build(base_ds, build); + Check(status == knowhere::Status::success, method + " Build failed, status=" + std::to_string(int(status))); + const double build_seconds = Seconds(start); + const auto bytes = index.Size(); + std::vector settings; + auto add_setting = [&](float mass, float factor, float drop) { + auto cfg = search; + cfg["sindi_query_mass"] = mass; + cfg["refine_k"] = factor; + cfg["drop_ratio_search"] = drop; + settings.push_back(cfg); + }; + if (o.sweep == "mass") { + for (float mass : {1.0f, .9f, .8f, .7f, .6f, .5f}) + for (float factor : {1.0f, 5.0f, 10.0f}) add_setting(mass, factor, 0); + } else if (o.sweep == "count") { + for (float keep : {1.0f, .7f, .5f, .3f, .1f}) + for (float factor : {1.0f, 5.0f, 10.0f}) add_setting(keep, factor, 0); + } else if (o.sweep == "drop") { + for (float drop : {0.0f, .3f, .5f, .7f, .9f}) add_setting(1, 1, drop); + } else + add_setting(o.mass, o.factor, o.drop); + size_t setting_id = 0; + for (const auto& setting : settings) { + search = setting; + const std::string hit_path = o.output + ".case" + std::to_string(setting_id++) + ".hits.csv"; + { + auto warmup = index.Search(query_ds, search, nullptr); + Check(warmup.has_value(), "Warmup failed: " + warmup.what()); + } + Json run = {{"method", method}, {"index_type", type}, + {"build_config", build}, {"effective_codec", codec}, + {"search_config", search}, {"build_seconds", build_seconds}, + {"index_bytes", bytes}, {"measurements", Json::array()}}; + std::vector first_ids; + std::vector first_scores; + for (int repeat = 0; repeat < o.repeats; ++repeat) { + start = Clock::now(); + auto result = index.Search(query_ds, search, nullptr); + const double seconds = Seconds(start); + Check(result.has_value(), "Search failed: " + result.what()); + const size_t result_count = o.queries * o.k; + const auto* ids = result.value()->GetIds(); + const auto* scores = result.value()->GetDistance(); + uint64_t result_hash = 14695981039346656037ULL; + for (size_t i = 0; i < result_count; ++i) { + result_hash = (result_hash ^ static_cast(ids[i])) * 1099511628211ULL; + result_hash = (result_hash ^ std::bit_cast(scores[i])) * 1099511628211ULL; + } + if (repeat == 0) { + first_ids.assign(ids, ids + result_count); + first_scores.assign(scores, scores + result_count); + std::ofstream hits(hit_path); + hits << "query_id,rank,document_id,score\n" << std::setprecision(9); + for (size_t i = 0; i < result_count; ++i) + hits << selected[i / o.k] << ',' << i % o.k + 1 << ',' << ids[i] << ',' << scores[i] + << '\n'; + hits.flush(); + Check(hits.good(), "Cannot write hit file"); + } else { + Check(std::memcmp(ids, first_ids.data(), result_count * sizeof(*ids)) == 0 && + std::memcmp(scores, first_scores.data(), result_count * sizeof(*scores)) == 0, + "Results changed across measured repetitions"); + } + const auto recall = Recall(result.value(), truth, o.queries, o.k, base.rows.size()); + if (repeat == 0) { + run["recall_mismatches"] = RecallMismatches(result.value(), truth, selected, o.k); + } + const double qps = o.queries / seconds; + csv << method << ',' << type << ',' << version << ",fp16," << codec << ',' << base.rows.size() + << ',' << o.queries << ',' << o.k << ',' << o.threads << ',' << build_seconds << ',' << bytes + << ',' << repeat + 1 << ',' << seconds << ',' << qps << ',' << recall << ',' + << search["sindi_query_mass"] << ',' << search["refine_k"] << ',' << search["drop_ratio_search"] + << ',' << o.refine << '\n'; + csv.flush(); + Check(csv.good(), "Failed writing results"); + run["measurements"].push_back({{"repeat", repeat + 1}, + {"search_seconds", seconds}, + {"qps", qps}, + {"recall", recall}, + {"result_hash", result_hash}}); + std::cout << method << " repeat=" << repeat + 1 << " seconds=" << seconds << " QPS=" << qps + << " recall@" << o.k << '=' << recall << std::endl; + } + // Separate, untimed coverage diagnostics. Materialize exactly the coarse query. + SparseData selected_data; + selected_data.dim = queries.dim; + double mass_sum = 0; + size_t retained_nnz = 0; + for (const auto& q : queries.rows) { + Row selected; + if (o.refine) + selected = knowhere::sparse::inverted::sindi::retain_query_mass( + q, search["sindi_query_mass"].get()); + else { + std::vector weights; + for (size_t j = 0; j < q.size(); ++j) weights.push_back(q[j].val); + std::sort(weights.begin(), weights.end()); + const size_t count = size_t(search["drop_ratio_search"].get() * weights.size()); + float threshold = weights.empty() ? 0 : weights[std::min(count, weights.size() - 1)]; + std::vector> terms; + for (size_t j = 0; j < q.size(); ++j) + if (q[j].val >= threshold) + terms.emplace_back(q[j].id, q[j].val); + selected = Row(terms); + } + double total = 0, retained = 0; + for (size_t j = 0; j < q.size(); ++j) total += q[j].val; + for (size_t j = 0; j < selected.size(); ++j) retained += selected[j].val; + mass_sum += total ? retained / total : 1; + retained_nnz += selected.size(); + selected_data.rows.push_back(std::move(selected)); + } + const size_t pool = knowhere::sparse::inverted::sindi::refinement_pool_size( + o.k, search["refine_k"].get(), base.rows.size()); + auto diagnostic = search; + diagnostic["k"] = pool; + diagnostic["sindi_query_mass"] = 1; + diagnostic["refine_k"] = 1; + diagnostic["drop_ratio_search"] = 0; + auto coarse = index.Search(selected_data.Dataset(), diagnostic, nullptr); + Check(coarse.has_value(), "Coarse coverage diagnostic failed"); + size_t covered = 0; + for (size_t q = 0; q < o.queries; ++q) + for (int j = 0; j < o.k; ++j) + covered += + std::find(coarse.value()->GetIds() + q * pool, coarse.value()->GetIds() + (q + 1) * pool, + truth.ids[q * o.k + j]) != coarse.value()->GetIds() + (q + 1) * pool; + run["candidate_pool"] = pool; + run["candidate_coverage"] = double(covered) / (o.queries * o.k); + run["retained_query_nnz_mean"] = double(retained_nnz) / o.queries; + run["retained_mass_mean"] = mass_sum / o.queries; + metadata["runs"].push_back(std::move(run)); + std::ofstream meta(o.output + ".json"); + meta << metadata.dump(2) << '\n'; + meta.flush(); + Check(meta.good(), "Cannot write metadata"); + } // search settings + } + return 0; + } catch (const std::exception& e) { + std::cerr << "benchmark_sparse: " << e.what() << '\n'; + return 1; + } +} diff --git a/src/index/sparse/inverted_index.h b/src/index/sparse/inverted_index.h index fee72ba2c..cba83985a 100644 --- a/src/index/sparse/inverted_index.h +++ b/src/index/sparse/inverted_index.h @@ -142,6 +142,8 @@ struct InvertedIndexSearchParams { InvertedIndexAlgo algo; IndexScorerConfig scorer_config; size_t bulk_query_nnz_threshold; + float sindi_query_mass = 1.0f; + float refine_k = 1.0f; struct { float drop_ratio_search; @@ -200,6 +202,11 @@ class InvertedIndex { virtual ~InvertedIndex() = default; + virtual bool + refinement_enabled() const noexcept { + return false; + } + /** * @brief Get total size of the index in bytes */ diff --git a/src/index/sparse/inverted_index_format.h b/src/index/sparse/inverted_index_format.h index 750432ba4..fc74addbd 100644 --- a/src/index/sparse/inverted_index_format.h +++ b/src/index/sparse/inverted_index_format.h @@ -81,6 +81,7 @@ enum class InvertedIndexSectionType : uint32_t { PROMETHEUS_BUILD_STATS = 6, DIM_MAP_MPHF = 7, BM25_U8_OVERFLOWS = 8, + SINDI_REFINEMENT = 9, }; struct InvertedIndexSectionHeader { diff --git a/src/index/sparse/sindi_inverted_index.h b/src/index/sparse/sindi_inverted_index.h index f90e9b05a..791f1a36e 100644 --- a/src/index/sparse/sindi_inverted_index.h +++ b/src/index/sparse/sindi_inverted_index.h @@ -2,7 +2,6 @@ // Reference: https://arxiv.org/abs/2509.08395 #pragma once - #include #include #include @@ -29,6 +28,7 @@ #include "knowhere/bitsetview.h" #include "knowhere/operands.h" #include "simd/hook.h" +#include "sindi_refinement.h" namespace knowhere::sparse::inverted { @@ -58,7 +58,8 @@ class SindiInvertedIndex : public DimMapInvertedIndex::value_type); @@ -707,6 +718,24 @@ class SindiInvertedIndex : public DimMapInvertedIndex(knowhere::fp16(data[i][j].val)))) { + return Status::invalid_args; + } + } + } + } + const size_t old_nr_rows = this->nr_rows_; this->max_dim_ = std::max(this->max_dim_, static_cast(dim)); LOG_KNOWHERE_INFO_ << "SindiInvertedIndex build started: rows=" << rows << ", existing_rows=" << old_nr_rows @@ -753,6 +782,11 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_rows_ = rows; + + if (refine_) { + rebuild_refinement_seek(); + } + LOG_KNOWHERE_INFO_ << "SindiInvertedIndex build completed: rows=" << this->nr_rows_ << ", inner_dims=" << this->nr_inner_dims_ << ", windows=" << nr_windows_ << ", index_bytes=" << size(); @@ -839,6 +873,9 @@ class SindiInvertedIndex : public DimMapInvertedIndex(max_scores_per_dim_.data(), max_scores_per_dim_.size()); + if (refine_) + rebuild_refinement_seek(); + LOG_KNOWHERE_INFO_ << "SindiInvertedIndex incremental build completed: rows=" << this->nr_rows_ << ", inner_dims=" << this->nr_inner_dims_ << ", windows=" << nr_windows_ << ", index_bytes=" << size(); @@ -884,6 +921,10 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_inner_dims_; @@ -948,6 +989,14 @@ class SindiInvertedIndex : public DimMapInvertedIndex(InvertedIndexQuantType::IP_FP16)}; + writer.write(metadata, sizeof(metadata)); + } + return Status::success; } } @@ -1025,6 +1081,9 @@ class SindiInvertedIndex : public DimMapInvertedIndex(reader.data() + reader.tellg()); bm25_u8_overflow_offsets_span_ = std::span(offsets, overflow_count); reader.advance(static_cast(overflow_count) * sizeof(uint32_t)); const auto* values = reinterpret_cast(reader.data() + reader.tellg()); bm25_u8_overflow_values_span_ = std::span(values, overflow_count); reader.advance(static_cast(overflow_count) * sizeof(uint16_t)); + if (!std::is_sorted(bm25_u8_overflow_offsets_span_.begin(), bm25_u8_overflow_offsets_span_.end()) || (!bm25_u8_overflow_offsets_span_.empty() && @@ -1261,6 +1322,24 @@ class SindiInvertedIndex : public DimMapInvertedIndex> bits; + count = packed & mask; + if (wid >= nr_windows_ || (i && seek.windows.back() >= wid)) { + throw std::runtime_error("Invalid refinement sparse windows"); + } + + seek.windows.push_back(wid); + } else { + uint16_t value; + std::memcpy(&value, counts.data() + i * 2, 2); + count = value; + } + + if (count > window_size_ || sum + count > ids.size()) { + throw std::runtime_error("Invalid refinement posting count"); + } + + seek.offsets.push_back(static_cast(sum)); + // Builders emit document order; validate it also on deserialization. + for (size_t j = sum; j < sum + count; ++j) { + const float represented = static_cast(posting_vals(dim)[j]); + if (!std::isfinite(represented) || represented < 0 || + uint64_t(wid) * window_size_ + ids[j] >= this->nr_rows_ || ids[j] >= window_size_ || + (j > sum && ids[j - 1] >= ids[j])) { + // no good + throw std::runtime_error("Invalid refinement posting IDs"); + } + } + + sum += count; + } + + if (sum != ids.size() || sum > std::numeric_limits::max()) { + throw std::runtime_error("Invalid refinement posting offsets"); + } + + seek.offsets.push_back(static_cast(sum)); + } + + refinement_seek_.swap(seeks); + } + + std::pair + window_posting_range(uint32_t dim, uint32_t window) const { + return refinement_seek_.at(dim).range(window); + } + + // Physical U16/FP16 lookup backend. Packed ID/value formats can replace this + // operation without changing mass selection, pool sizing, or final selection. + void + score_candidate_ids(uint32_t dim, uint32_t begin, uint32_t end, float weight, std::span candidates, + std::span scores) const { + const auto ids = posting_ids(dim); + const auto vals = posting_vals(dim); + if (begin == end) + return; + auto cursor = ids.begin() + begin; + const auto stop = ids.begin() + end; + for (size_t i = 0; i < candidates.size(); ++i) { + const uint32_t local = candidates[i] % window_size_; + cursor = std::lower_bound(cursor, stop, local); + if (cursor != stop && *cursor == local) + scores[i] = std::fma(weight, static_cast(vals[cursor - ids.begin()]), scores[i]); + } + } + + void + search_refined(const SparseRow& query, size_t k, float* distances, label_t* labels, + const BitsetView& bitset, const InvertedIndexSearchParams& params) const { + const auto count = sindi::refinement_pool_size(k, params.refine_k, this->nr_rows_); + if (count == 0) { + return; + } + + const auto selected = sindi::retain_query_mass(query, params.sindi_query_mass); + auto coarse_query = parse_query_with_dim_map(selected, this->dim_map_, 0.0f); + + std::vector coarse_scores(count, std::numeric_limits::quiet_NaN()); + std::vector coarse_ids(count, -1); + search_coarse(std::move(coarse_query), count, coarse_scores.data(), coarse_ids.data(), bitset, params); + + std::vector candidates; + candidates.reserve(count); + for (auto id : coarse_ids) { + if (id >= 0 && static_cast(id) < this->nr_rows_ && (bitset.empty() || !bitset.test(id))) { + candidates.push_back(static_cast(id)); + } + } + std::sort(candidates.begin(), candidates.end()); + candidates.erase(std::unique(candidates.begin(), candidates.end()), candidates.end()); + + std::vector scores(candidates.size(), 0.0f); + auto full = parse_query_with_dim_map(query, this->dim_map_, 0.0f); + std::sort(full.begin(), full.end(), [&](const auto& a, const auto& b) { + const float x = a.second * max_scores_per_dim_span_[a.first]; + const float y = b.second * max_scores_per_dim_span_[b.first]; + return x != y ? x > y : a.first < b.first; + }); + + for (const auto& [dim, weight] : full) { + for (size_t start = 0; start < candidates.size();) { + const auto window = candidates[start] / window_size_; + size_t stop = start + 1; + while (stop < candidates.size() && candidates[stop] / window_size_ == window) { + ++stop; + } + + const auto [begin, end] = window_posting_range(dim, window); + score_candidate_ids(dim, begin, end, weight, + std::span(candidates).subspan(start, stop - start), + std::span(scores).subspan(start, stop - start)); + start = stop; + } + } + + std::vector order(candidates.size()); + std::iota(order.begin(), order.end(), 0); + std::sort(order.begin(), order.end(), [&](size_t a, size_t b) { + return scores[a] != scores[b] ? (scores[a] > scores[b]) : (candidates[a] < candidates[b]); + }); + + for (size_t i = 0; i < std::min(k, order.size()); ++i) { + labels[i] = candidates[order[i]]; + distances[i] = scores[order[i]]; + } + } + uint32_t window_size_{max_window_size}; uint32_t nr_windows_{0}; diff --git a/src/index/sparse/sindi_refinement.h b/src/index/sparse/sindi_refinement.h new file mode 100644 index 000000000..1594b9239 --- /dev/null +++ b/src/index/sparse/sindi_refinement.h @@ -0,0 +1,100 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// Licensed under the Apache License, Version 2.0. +#pragma once + +#include +#include +#include +#include +#include +#include + +#include "knowhere/sparse_utils.h" + +namespace knowhere::sparse::inverted::sindi { + +inline size_t +refinement_pool_size(size_t k, float factor, size_t count) { + if (!std::isfinite(factor) || factor < 1) { + throw std::invalid_argument("Invalid refine_k"); + } + + const long double requested = std::ceil(static_cast(k) * factor); + return requested >= count ? count : static_cast(requested); +} + +template +bool +valid_refinement_row(const SparseRow& row) { + for (size_t i = 0; i < row.size(); ++i) { + const auto [dim, val] = row[i]; + if (!std::isfinite(val) || val < 0 || (i && row[i - 1].id >= dim)) { + return false; + } + } + + return true; +} + +// Select using original weights and external IDs, before dimension-map lookup. +// Query term selection does not depend on the physical posting value/ID codecs. +template +SparseRow +retain_query_mass(const SparseRow& query, float mass) { + if (!std::isfinite(mass) || mass <= 0 || mass > 1 || !valid_refinement_row(query)) { + throw std::invalid_argument("Refinement requires finite nonnegative, sorted unique query coordinates"); + } + + std::vector order(query.size()); + std::iota(order.begin(), order.end(), 0); + + double total = 0; + for (size_t i = 0; i < query.size(); ++i) { + total += query[i].val; + } + + std::sort(order.begin(), order.end(), [&](size_t a, size_t b) { + return (query[a].val != query[b].val) ? (query[a].val > query[b].val) : (query[a].id < query[b].id); + }); + + size_t keep = 0; + double sum = 0; + while (keep < order.size() && query[order[keep]].val > 0 && (mass == 1 || sum < mass * total)) { + sum += query[order[keep++]].val; + } + + order.resize(keep); + std::sort(order.begin(), order.end()); + + SparseRow out(keep); + for (size_t i = 0; i < keep; ++i) { + out.set_at(i, query[order[i]].id, query[order[i]].val); + } + + return out; +} + +// Logical posting offsets, not byte offsets. Future packed backends may locate/decode +// blocks behind window_posting_range and score_candidate_ids without changing search. +struct RefinementSeek { + bool sparse = false; + std::vector windows; + std::vector offsets; + + std::pair inline range(uint32_t window) const { + size_t pos = window; + + if (sparse) { + auto it = std::lower_bound(windows.begin(), windows.end(), window); + if (it == windows.end() || *it != window) { + return {0, 0}; + } + + pos = it - windows.begin(); + } + + return {offsets.at(pos), offsets.at(pos + 1)}; + } +}; + +} // namespace knowhere::sparse::inverted::sindi diff --git a/src/index/sparse/sparse_index_config.h b/src/index/sparse/sparse_index_config.h index c3bd90c33..ab173e5fa 100644 --- a/src/index/sparse/sparse_index_config.h +++ b/src/index/sparse/sparse_index_config.h @@ -14,6 +14,7 @@ #include #include +#include #include #include @@ -59,6 +60,11 @@ class SparseInvertedIndexConfig : public BaseConfig { CFG_FLOAT drop_ratio_build; CFG_FLOAT drop_ratio_search; CFG_INT refine_factor; + // Keep the legacy integer field: changing its type is not assumed ABI-compatible. + // Match HNSW's floating-point multiplier, including fractional pool sizes. + CFG_BOOL refine; + CFG_FLOAT refine_k; + CFG_FLOAT sindi_query_mass; CFG_FLOAT dim_max_score_ratio; CFG_INT bulk_query_nnz_threshold; CFG_INT block_max_block_size; @@ -70,6 +76,25 @@ class SparseInvertedIndexConfig : public BaseConfig { CFG_INT sindi_window_size; KNOWHERE_DECLARE_CONFIG(SparseInvertedIndexConfig) { + KNOWHERE_CONFIG_DECLARE_FIELD(refine) + .description("build SINDI IP full-query refinement support") + .set_default(false) + .for_train() + .for_static(); + KNOWHERE_CONFIG_DECLARE_FIELD(refine_k) + .description("SINDI coarse candidate multiplier; ceil(k * refine_k)") + .set_default(1.0f) + .set_range(1.0f, std::numeric_limits::max()) + .for_search() + .for_range_search() + .for_iterator(); + KNOWHERE_CONFIG_DECLARE_FIELD(sindi_query_mass) + .description("SINDI coarse query retained weight mass, not term-count fraction") + .set_default(1.0f) + .set_range(0.0f, 1.0f, false, true) + .for_search() + .for_range_search() + .for_iterator(); // NOTE: drop_ratio_build has been deprecated, it won't change anything KNOWHERE_CONFIG_DECLARE_FIELD(drop_ratio_build) .description("drop ratio for build") @@ -83,17 +108,10 @@ class SparseInvertedIndexConfig : public BaseConfig { .for_search() .for_range_search() .for_iterator(); - /** - * refine_factor is used for approximate search. - * refine_factor == 1 means no refinement, and is the default value. - * refine_factor > 1 means refinement. The larger the value, the more - * accurate the approximate result will be, but the slower the - * performance. - * Be aware that if you opt to use a large drop_ratio_search, it is - * necessary for you to manually modify this value. - */ + // Legacy compatibility field. Current sparse search does not consume it; + // SINDI refinement uses the separate floating-point refine_k parameter. KNOWHERE_CONFIG_DECLARE_FIELD(refine_factor) - .description("refine factor for approximate search") + .description("legacy unused integer multiplier; SINDI refinement uses refine_k") .set_default(1) .for_search(); /** @@ -183,6 +201,15 @@ class SparseInvertedIndexConfig : public BaseConfig { Status CheckAndAdjust(PARAM_TYPE param_type, std::string* err_msg) override { + const float mass = sindi_query_mass.value_or(1.0f); + const float factor = refine_k.value_or(1.0f); + if (!std::isfinite(mass) || mass <= 0 || mass > 1 || !std::isfinite(factor) || factor < 1) { + return HandleError(err_msg, "Invalid SINDI query mass or refine_k", Status::invalid_args); + } + if ((param_type & (RANGE_SEARCH | ITERATOR)) && (mass != 1 || factor != 1)) { + return HandleError(err_msg, "SINDI refinement controls support top-k search only", Status::invalid_args); + } + if (inverted_index_algo.has_value() && !IsSupportedSparseInvertedIndexAlgo(inverted_index_algo.value())) { return HandleError( err_msg, diff --git a/src/index/sparse/sparse_index_node.cc b/src/index/sparse/sparse_index_node.cc index 904c38a07..102c646d1 100644 --- a/src/index/sparse/sparse_index_node.cc +++ b/src/index/sparse/sparse_index_node.cc @@ -301,8 +301,28 @@ class SparseInvertedIndexNode : public IndexNode { } auto search_params = search_params_or.value(); + const bool refine = index_->refinement_enabled(); + if ((!refine && (cfg.sindi_query_mass.value_or(1) != 1 || cfg.refine_k.value_or(1) != 1)) || + (refine && (cfg.drop_ratio_search.value_or(0) != 0 || cfg.refine_factor.value_or(1) != 1))) { + return expected::Err( + Status::invalid_args, + "SINDI refinement requires a refined index, drop_ratio_search=0 and refine_factor=1; use refine_k"); + } + + search_params.sindi_query_mass = cfg.sindi_query_mass.value_or(1); + search_params.refine_k = cfg.refine_k.value_or(1); + auto queries = static_cast*>(dataset->GetTensor()); auto nq = dataset->GetRows(); + + if (refine) { + for (int64_t i = 0; i < nq; ++i) { + if (!sparse::inverted::sindi::valid_refinement_row(queries[i])) { + return expected::Err(Status::invalid_args, "Invalid SINDI refinement query"); + } + } + } + auto k = cfg.k.value(); auto p_id = std::make_unique(nq * k); auto p_dist = std::make_unique(nq * k); @@ -701,9 +721,11 @@ class SparseInvertedIndexNode : public IndexNode { cfg.sindi_window_size.value_or(sparse::inverted::SindiInvertedIndexIP::max_window_size); IndexPtr index; if (is_growable) { - index = std::make_unique(window_size); + index = std::make_unique( + window_size, cfg.refine.value_or(false)); } else { - index = std::make_unique(window_size); + index = std::make_unique(window_size, + cfg.refine.value_or(false)); } ConfigureSindiSerialization(index.get()); index->set_build_algo(algo); @@ -753,6 +775,16 @@ class SparseInvertedIndexNode : public IndexNode { expected>> CreateIndex(const SparseInvertedIndexConfig& cfg, bool is_growable = false, std::optional encoding = std::nullopt) const { + if (cfg.refine.value_or(false)) { + const auto algo = NormalizeInvertedIndexAlgo(cfg.inverted_index_algo.value_or("")); + if (index_version_ < 11 || !IsMetricType(cfg.metric_type.value(), metric::IP) || + (!algo.empty() && algo != "SINDI") || + (!cfg.quant_type.value_or("").empty() && cfg.quant_type.value() != "fp16")) { + return expected>>::Err( + Status::invalid_args, "refine requires SINDI FP16 IP version >= 11"); + } + } + const auto explicit_algo = NormalizeInvertedIndexAlgo(cfg.inverted_index_algo.value_or("")); const auto status = ValidateInvertedIndexAlgo(explicit_algo); if (status != Status::success) { diff --git a/tests/ut/test_sparse.cc b/tests/ut/test_sparse.cc index 9ddbc4c8f..86f71f210 100644 --- a/tests/ut/test_sparse.cc +++ b/tests/ut/test_sparse.cc @@ -2693,3 +2693,292 @@ TEST_CASE("Test SINDI Index Default Algo for Version 10", "[sparse][sindi]") { results = idx.Search(query_ds, search_json, nullptr); REQUIRE(!results.has_value()); } + +#include "index/sparse/inverted_index.h" +#include "index/sparse/sindi_refinement.h" + +namespace { +using RefineRow = knowhere::sparse::SparseRow; +RefineRow +RefineTestRow(std::initializer_list> values) { + RefineRow row(values.size()); + size_t i = 0; + for (auto [id, value] : values) row.set_at(i++, id, value); + return row; +} +knowhere::DataSetPtr +RefineDataset(const std::vector& rows) { + auto data = knowhere::GenDataSet(rows.size(), 100001, rows.data()); + data->SetIsSparse(true); + return data; +} +knowhere::Json +RefineBuild(bool enabled = true) { + return {{"metric_type", "IP"}, + {"inverted_index_algo", "SINDI"}, + {"refine", enabled}, + {"sindi_window_size", 4096}, + {"quant_type", "fp16"}}; +} +knowhere::Json +RefineSearch() { + return {{"metric_type", "IP"}, {"k", 1}, {"refine_k", 2.0}, {"sindi_query_mass", 0.6}}; +} +} // namespace + +TEST_CASE("SINDI mass and legacy count selection contracts", "[sparse][sindi][refinement]") { + using namespace knowhere::sparse::inverted; + auto q = RefineTestRow({{0, 50}, {1, 20}, {2, 10}, {3, 8}, {4, 5}, {5, 3}, {6, 2}, {7, 1}, {8, .6f}, {9, .4f}}); + auto selected = sindi::retain_query_mass(q, .7f); + REQUIRE(selected.size() == 2); + REQUIRE(selected[1].id == 1); + std::vector weights; + for (size_t i = 0; i < q.size(); ++i) weights.push_back(q[i].val); + REQUIRE(get_query_drop_threshold(weights, .3f) == 2); + auto tied = RefineTestRow({{0, 1}, {1, 1}, {2, 1}, {3, 1}}); + REQUIRE(sindi::retain_query_mass(tied, .5f).size() == 2); + REQUIRE(sindi::retain_query_mass(tied, 1).size() == 4); + REQUIRE(sindi::retain_query_mass(RefineTestRow({{0, 0}}), .5f).size() == 0); + REQUIRE_FALSE(sindi::valid_refinement_row(RefineTestRow({{1, 1}, {0, 1}}))); + REQUIRE_FALSE(sindi::valid_refinement_row(RefineTestRow({{0, -1}}))); + REQUIRE(sindi::refinement_pool_size(3, 1.5f, 100) == 5); + REQUIRE(sindi::refinement_pool_size(10, std::numeric_limits::max(), 100) == 100); + REQUIRE_THROWS(sindi::refinement_pool_size(10, std::numeric_limits::infinity(), 100)); +} + +TEST_CASE("SINDI full-query refinement changes rank and persists", "[sparse][sindi][refinement]") { + using namespace knowhere; + std::vector base; + base.push_back(RefineTestRow({{0, 10}})); + base.push_back(RefineTestRow({{0, 9}, {100000, 9}})); + std::vector queries; + queries.push_back(RefineTestRow({{0, 2}, {100000, 1}})); + auto index = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(index.Build(RefineDataset(base), RefineBuild()) == Status::success); + const auto search = RefineSearch(); + auto check = [&](auto& idx) { + auto result = idx.Search(RefineDataset(queries), search, nullptr); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == 1); + REQUIRE(result.value()->GetDistance()[0] == 27); + }; + check(index); + auto one = search; + one["refine_k"] = 1; + auto result = index.Search(RefineDataset(queries), one, nullptr); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == 0); // reranking cannot recover an absent candidate + REQUIRE(result.value()->GetDistance()[0] == 20); + BinarySet binary; + REQUIRE(index.Serialize(binary) == Status::success); + auto restored = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(restored.Deserialize(binary, Json{{"metric_type", "IP"}}) == Status::success); + check(restored); + SparseQuantIndexFile file(binary.GetByName(IndexEnum::INDEX_SPARSE_INVERTED_INDEX)); + auto mapped = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(mapped.DeserializeFromFile(file.path, Json{{"metric_type", "IP"}, {"enable_mmap", true}}) == + Status::success); + check(mapped); + for (auto field : {"drop_ratio_search", "refine_factor"}) { + auto invalid = search; + invalid[field] = field == std::string("refine_factor") ? 2.0 : .2; + REQUIRE_FALSE(index.Search(RefineDataset(queries), invalid, nullptr).has_value()); + } + for (auto field : {"sindi_query_mass", "refine_k"}) { + auto invalid = search; + invalid[field] = 0; + REQUIRE_FALSE(index.Search(RefineDataset(queries), invalid, nullptr).has_value()); + } + auto legacy = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(legacy.Build(RefineDataset(base), RefineBuild(false)) == Status::success); + REQUIRE_FALSE(legacy.Search(RefineDataset(queries), search, nullptr).has_value()); + auto legacy_search = Json{{"metric_type", "IP"}, {"k", 1}, {"drop_ratio_search", .5}, {"refine_factor", 1}}; + auto a = legacy.Search(RefineDataset(queries), legacy_search, nullptr); + legacy_search["refine_factor"] = 10; + auto b = legacy.Search(RefineDataset(queries), legacy_search, nullptr); + REQUIRE(a.has_value()); + REQUIRE(b.has_value()); + REQUIRE(a.value()->GetIds()[0] == 0); + REQUIRE(b.value()->GetIds()[0] == 0); + REQUIRE(a.value()->GetDistance()[0] == b.value()->GetDistance()[0]); +} + +TEST_CASE("SINDI refinement windows filters and growable Add", "[sparse][sindi][refinement]") { + using namespace knowhere; + const auto window = GENERATE(1024, 4096, 65535); + const bool growable = GENERATE(false, true); + std::vector base(70001); + std::vector ids = {0, 4095, 4096, 65534, 65535, 70000}; + for (size_t i = 0; i < ids.size(); ++i) base[ids[i]] = RefineTestRow({{0, float(10 - i)}, {100000, float(i * 4)}}); + std::vector queries; + queries.push_back(RefineTestRow({{0, 2}, {100000, 1}})); + auto build = RefineBuild(); + build["sindi_window_size"] = window; + auto type = growable ? IndexEnum::INDEX_SPARSE_INVERTED_INDEX_CC : IndexEnum::INDEX_SPARSE_INVERTED_INDEX; + auto idx = IndexFactory::Instance().Create(type, 11).value(); + if (growable) { + std::vector first(base.begin(), base.begin() + 4096); + std::vector rest(base.begin() + 4096, base.end()); + REQUIRE(idx.Build(RefineDataset(first), build) == Status::success); + REQUIRE(idx.Add(RefineDataset(rest), build) == Status::success); + } else + REQUIRE(idx.Build(RefineDataset(base), build) == Status::success); + auto search = RefineSearch(); + search["refine_k"] = 10; + auto result = idx.Search(RefineDataset(queries), search, nullptr); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == 70000); + REQUIRE(result.value()->GetDistance()[0] == 30); + std::vector mask((base.size() + 7) / 8, 0); + mask[70000 / 8] |= 1u << (70000 % 8); + BitsetView bitset(mask.data(), base.size()); + result = idx.Search(RefineDataset(queries), search, bitset); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == 65535); + REQUIRE(result.value()->GetDistance()[0] == 28); +} + +TEST_CASE("SINDI refinement represented-value oracle and validation", "[sparse][sindi][refinement]") { + using namespace knowhere; + std::vector base; + for (size_t i = 0; i < 1500; ++i) + base.push_back(RefineTestRow({{0, 1.003f + float(i % 7) * .0007f}, {7, .5f}, {100000, float(i) / 1500}})); + std::vector queries; + queries.push_back(RefineTestRow({{0, 2}, {7, .5f}, {100000, 1}})); + auto idx = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(idx.Build(RefineDataset(base), RefineBuild()) == Status::success); + auto search = RefineSearch(); + search["k"] = 2000; + search["refine_k"] = 1.5; + search["sindi_query_mass"] = .7; + auto result = idx.Search(RefineDataset(queries), search, nullptr); + REQUIRE(result.has_value()); + for (size_t i = 0; i < 1500; ++i) { + const auto id = result.value()->GetIds()[i]; + REQUIRE(id >= 0); + REQUIRE(id < 1500); + float reference = 0; + for (auto term : {0, 2, 1}) + reference = std::fma(queries[0][term].val, float(fp16(base[id][term].val)), reference); + REQUIRE(std::abs(result.value()->GetDistance()[i] - reference) < 1e-6f); + } + for (size_t i = 1500; i < 2000; ++i) REQUIRE(result.value()->GetIds()[i] == -1); + auto invalid = RefineBuild(); + invalid["inverted_index_algo"] = "DAAT_WAND"; + auto bad = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(bad.Build(RefineDataset(base), invalid) != Status::success); + auto v10 = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 10).value(); + REQUIRE(v10.Build(RefineDataset(base), RefineBuild()) != Status::success); + std::vector negative; + negative.push_back(RefineTestRow({{0, -1}})); + REQUIRE_FALSE(idx.Search(RefineDataset(negative), RefineSearch(), nullptr).has_value()); + auto range = RefineSearch(); + range["radius"] = 1; + REQUIRE_FALSE(idx.RangeSearch(RefineDataset(queries), range, nullptr).has_value()); + BinarySet bytes; + REQUIRE(idx.Serialize(bytes) == Status::success); + auto binary = bytes.GetByName(IndexEnum::INDEX_SPARSE_INVERTED_INDEX); + using namespace sparse::inverted; + uint32_t sections = 0; + std::memcpy(§ions, binary->data.get() + kInvertedIndexFileHeaderSize, 4); + bool corrupted = false; + for (uint32_t i = 0; i < sections; ++i) { + InvertedIndexSectionHeader header; + std::memcpy(&header, binary->data.get() + kInvertedIndexFileHeaderSize + 4 + i * sizeof(header), + sizeof(header)); + if (header.type == InvertedIndexSectionType::SINDI_REFINEMENT) { + uint32_t unsupported_version = 99; + std::memcpy(binary->data.get() + header.offset, &unsupported_version, 4); + corrupted = true; + } + } + REQUIRE(corrupted); + REQUIRE(bad.Deserialize(bytes, Json{{"metric_type", "IP"}}) != Status::success); +} + +TEST_CASE("SINDI count threshold agrees with sorted reference", "[sparse][sindi][refinement]") { + using namespace knowhere::sparse::inverted; + for (size_t n : {0, 1, 2, 3, 10, 101}) + for (float ratio : {0.f, .1f, .3f, .7f, .99f}) { + std::vector values(n); + for (size_t i = 0; i < n; ++i) values[i] = float((i * 17) % 11); // zeros and ties + auto sorted = values; + std::sort(sorted.begin(), sorted.end()); + const auto count = static_cast(ratio * n); + const float threshold = count ? sorted[count] : 0; + REQUIRE(get_query_drop_threshold(values, ratio) == threshold); + REQUIRE(std::count_if(values.begin(), values.end(), [&](float v) { return v < threshold; }) <= count); + } +} + +TEST_CASE("Historical candidate-filter refinement agrees with direct lookup", "[sparse][sindi][refinement]") { + // Port of the pre-eb69fcc1 algorithm's control flow, using current FP16 SINDI + // for both passes. This is not a build of the historical FP32/WAND backend. + using namespace knowhere; + std::vector base; + base.push_back(RefineTestRow({{0, 10}})); + base.push_back(RefineTestRow({{0, 9}, {100000, 9}})); + base.push_back(RefineTestRow({{0, .1f}, {100000, 100}})); + std::vector queries; + queries.push_back(RefineTestRow({{0, 2}, {100000, 1}})); + auto make = [] { + return IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + }; + auto legacy = make(), refined = make(); + REQUIRE(legacy.Build(RefineDataset(base), RefineBuild(false)) == Status::success); + REQUIRE(refined.Build(RefineDataset(base), RefineBuild(true)) == Status::success); + Json coarse_cfg = {{"metric_type", "IP"}, {"k", 2}, {"drop_ratio_search", .5}, {"dim_max_score_ratio", 1.05}}; + auto coarse = legacy.Search(RefineDataset(queries), coarse_cfg, nullptr); + REQUIRE(coarse.has_value()); + uint8_t mask = 0xff; + for (size_t i = 0; i < 2; ++i) mask &= ~(1u << coarse.value()->GetIds()[i]); + auto full_cfg = coarse_cfg; + full_cfg["k"] = 1; + full_cfg["drop_ratio_search"] = 0; + BitsetView allowed(&mask, base.size()); + auto filtered = legacy.Search(RefineDataset(queries), full_cfg, allowed); + auto direct = refined.Search(RefineDataset(queries), RefineSearch(), nullptr); + REQUIRE(filtered.has_value()); + REQUIRE(direct.has_value()); + REQUIRE(filtered.value()->GetIds()[0] == 1); + REQUIRE(filtered.value()->GetIds()[0] == direct.value()->GetIds()[0]); + REQUIRE(filtered.value()->GetDistance()[0] == direct.value()->GetDistance()[0]); + auto unfiltered = legacy.Search(RefineDataset(queries), full_cfg, nullptr); + REQUIRE(unfiltered.has_value()); + REQUIRE(unfiltered.value()->GetIds()[0] == 2); +} + +TEST_CASE("SINDI refinement empty candidates and fractional API pool", "[sparse][sindi][refinement]") { + using namespace knowhere; + std::vector base; + base.push_back(RefineTestRow({{0, 10}})); + base.push_back(RefineTestRow({{0, 9}, {100000, 9}})); + auto idx = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(idx.Build(RefineDataset(base), RefineBuild()) == Status::success); + std::vector query; + query.push_back(RefineTestRow({{0, 2}, {100000, 1}})); + auto search = RefineSearch(); + search["refine_k"] = 1.01; // ceil(1 * 1.01) must retrieve two candidates. + auto result = idx.Search(RefineDataset(query), search, nullptr); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == 1); + REQUIRE(result.value()->GetDistance()[0] == 27); + uint8_t mask = 0xff; + result = idx.Search(RefineDataset(query), search, BitsetView(&mask, base.size())); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == -1); + for (auto row : {RefineTestRow({}), RefineTestRow({{0, 0}}), RefineTestRow({{1, 10}, {100000, 1}})}) { + std::vector empty_coarse; + empty_coarse.push_back(std::move(row)); + result = idx.Search(RefineDataset(empty_coarse), search, nullptr); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == -1); + } + auto unsupported = RefineBuild(); + unsupported["metric_type"] = "BM25"; + auto other = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(other.Build(RefineDataset(base), unsupported) != Status::success); + unsupported = RefineBuild(); + unsupported["quant_type"] = "fp32"; + REQUIRE(other.Build(RefineDataset(base), unsupported) != Status::success); +} From 81bb8520fdfa3286196022c5e4b4dd605ce405dc Mon Sep 17 00:00:00 2001 From: Alexandr Guzhva Date: Wed, 30 Sep 2026 11:08:17 -0400 Subject: [PATCH 2/3] add U12/E5M7 for SINDI IP for 4096 window size Signed-off-by: Alexandr Guzhva --- src/index/sparse/inverted_index.h | 1 + src/index/sparse/inverted_index_format.h | 1 + src/index/sparse/sindi_inverted_index.h | 455 ++++++++++++++++++++--- src/index/sparse/sindi_packed12.h | 71 ++++ src/index/sparse/sindi_packed12_x86.h | 22 ++ src/index/sparse/sindi_simd.cc | 37 ++ src/index/sparse/sindi_simd.h | 17 + src/index/sparse/sindi_simd_avx2.cc | 47 +++ src/index/sparse/sindi_simd_avx512.cc | 52 +++ src/index/sparse/sindi_simd_sve.cc | 150 ++++++++ src/index/sparse/sparse_index_config.h | 7 +- src/index/sparse/sparse_index_node.cc | 88 ++++- tests/ut/test_sparse.cc | 239 ++++++++++++ 13 files changed, 1114 insertions(+), 73 deletions(-) create mode 100644 src/index/sparse/sindi_packed12.h create mode 100644 src/index/sparse/sindi_packed12_x86.h diff --git a/src/index/sparse/inverted_index.h b/src/index/sparse/inverted_index.h index cba83985a..dfd78d94b 100644 --- a/src/index/sparse/inverted_index.h +++ b/src/index/sparse/inverted_index.h @@ -50,6 +50,7 @@ enum class InvertedIndexEncoding : uint32_t { BLOCK_MASKEDVBYTE = 2, FIXED_DOCID_WINDOWS = 3, BLOCK_ADAPTIVE = 4, + FIXED_DOCID_WINDOWS_U12_E5M7 = 5, }; enum class InvertedIndexPrometheusBuildStats : uint32_t { DATASET_NNZ_STATS = 0, POSTING_LIST_LENGTH_STATS = 1 }; diff --git a/src/index/sparse/inverted_index_format.h b/src/index/sparse/inverted_index_format.h index fc74addbd..ba8a289bc 100644 --- a/src/index/sparse/inverted_index_format.h +++ b/src/index/sparse/inverted_index_format.h @@ -34,6 +34,7 @@ enum class InvertedIndexQuantType : uint32_t { BM25_U8 = 3, BM25_U16 = 4, BM25_U32 = 5, + IP_E5M7 = 6, }; static_assert(sizeof(InvertedIndexQuantType) == sizeof(uint32_t)); diff --git a/src/index/sparse/sindi_inverted_index.h b/src/index/sparse/sindi_inverted_index.h index 791f1a36e..9ab3f517a 100644 --- a/src/index/sparse/sindi_inverted_index.h +++ b/src/index/sparse/sindi_inverted_index.h @@ -28,6 +28,7 @@ #include "knowhere/bitsetview.h" #include "knowhere/operands.h" #include "simd/hook.h" +#include "sindi_packed12.h" #include "sindi_refinement.h" namespace knowhere::sparse::inverted { @@ -58,8 +59,10 @@ class SindiInvertedIndex : public DimMapInvertedIndex::value_type); - if (!total_plists_ids_flat_span_.empty()) { + if (packed_ready_) { + res += packed_ids_.empty() ? packed_ids_span_.size() : packed_ids_.capacity(); + res += packed_vals_.empty() ? packed_vals_span_.size() : packed_vals_.capacity(); + } else if (!total_plists_ids_flat_span_.empty()) { res += total_plists_ids_flat_span_.size() * sizeof(uint16_t); res += total_plists_vals_flat_span_.size() * sizeof(QuantType); } else { @@ -117,6 +128,8 @@ class SindiInvertedIndex : public DimMapInvertedIndex uint64_t(std::numeric_limits::max()) - this->nr_rows_) { + return Status::invalid_args; + } + + uint64_t postings = plists_dim_offsets_span_.empty() ? 0 : plists_dim_offsets_span_.back(); + for (size_t i = 0; i < rows; ++i) { + for (size_t j = 0; j < data[i].size(); ++j) { + postings += std::abs(data[i][j].val) >= std::numeric_limits::epsilon(); + } + + if (postings > std::numeric_limits::max()) { + return Status::invalid_args; + } + } + } + + if (refine_ || packed_) { if constexpr (!is_ip) { return Status::invalid_args; } @@ -729,13 +759,20 @@ class SindiInvertedIndex : public DimMapInvertedIndex(knowhere::fp16(data[i][j].val)))) { + if (!std::isfinite(static_cast(knowhere::fp16(data[i][j].val))) || + (packed_ && std::signbit(data[i][j].val))) { return Status::invalid_args; } } } } + if constexpr (AllowIncremental) { + if (packed_ready_) { + expand_for_add(); + } + } + const size_t old_nr_rows = this->nr_rows_; this->max_dim_ = std::max(this->max_dim_, static_cast(dim)); LOG_KNOWHERE_INFO_ << "SindiInvertedIndex build started: rows=" << rows << ", existing_rows=" << old_nr_rows @@ -782,6 +819,8 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_rows_ = rows; + if (packed_) + pack_postings(); if (refine_) { rebuild_refinement_seek(); @@ -873,8 +912,13 @@ class SindiInvertedIndex : public DimMapInvertedIndex(max_scores_per_dim_.data(), max_scores_per_dim_.size()); - if (refine_) + if (packed_) { + pack_postings(); + } + + if (refine_) { rebuild_refinement_seek(); + } LOG_KNOWHERE_INFO_ << "SindiInvertedIndex incremental build completed: rows=" << this->nr_rows_ << ", inner_dims=" << this->nr_inner_dims_ << ", windows=" << nr_windows_ @@ -901,7 +945,7 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_rows_, sizeof(uint32_t)); writer.write(&this->max_dim_, sizeof(uint32_t)); writer.write(&this->nr_inner_dims_, sizeof(uint32_t)); - const auto quant_type = posting_quant_type(); + const auto quant_type = packed_ ? InvertedIndexQuantType::IP_E5M7 : posting_quant_type(); writer.write(&quant_type, sizeof(quant_type)); const std::array reserved{}; writer.write(reserved.data(), reserved.size()); @@ -933,7 +977,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex uint64_t { - size_t res = sizeof(uint32_t) * 3; + size_t res = sizeof(uint32_t) * (packed_ ? 6 : 3); const size_t mask_sz = (nr_dims + 7) / 8; res += mask_sz * sizeof(uint8_t); @@ -951,8 +995,8 @@ class SindiInvertedIndex : public DimMapInvertedIndex(InvertedIndexEncoding::FIXED_DOCID_WINDOWS); + uint32_t index_encoding_type = + static_cast(packed_ ? InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7 + : InvertedIndexEncoding::FIXED_DOCID_WINDOWS); write_padding_until(writer, section_headers[0].offset); writer.write(&index_encoding_type, sizeof(uint32_t)); writer.write(&this->window_size_, sizeof(uint32_t)); writer.write(&this->nr_windows_, sizeof(uint32_t)); + if (packed_) { + // Version 1, independent U12 layout tag 2 and E5M7 value-codec tag 6. + const uint32_t descriptor[] = {1, 2, static_cast(InvertedIndexQuantType::IP_E5M7)}; + writer.write(descriptor, sizeof(descriptor)); + } // write plists_woffsets_formats_mask and plists_window_nnzs writer.write(plists_wnnzs_fmts_msk_span_.data(), sizeof(uint8_t), plists_wnnzs_fmts_msk_span_.size()); @@ -1020,7 +1071,10 @@ class SindiInvertedIndex : public DimMapInvertedIndex 0 && nr_dims > 0) { + if (packed_) { + writer.write(packed_ids_span_.data(), packed_ids_span_.size()); + writer.write(packed_vals_span_.data(), packed_vals_span_.size()); + } else if (nr_windows_ > 0 && nr_dims > 0) { // ids for (size_t dim_id = 0; dim_id < nr_dims; ++dim_id) { const auto ids = posting_ids(dim_id); @@ -1065,7 +1119,9 @@ class SindiInvertedIndex : public DimMapInvertedIndex(InvertedIndexQuantType::IP_FP16)}; + const uint32_t metadata[] = { + 1, packed_ ? 2u : 1u, + static_cast(packed_ ? InvertedIndexQuantType::IP_E5M7 : InvertedIndexQuantType::IP_FP16)}; writer.write(metadata, sizeof(metadata)); } @@ -1083,6 +1139,9 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_inner_dims_, sizeof(uint32_t)); InvertedIndexQuantType quant_type{}; reader.read(&quant_type, sizeof(quant_type)); - if (!validate_posting_quant_type(quant_type)) { + if (packed_ ? quant_type != InvertedIndexQuantType::IP_E5M7 + : !validate_posting_quant_type(quant_type)) { return Status::invalid_serialized_index_type; } reader.advance(kInvertedIndexHeaderReservedBytes); @@ -1144,13 +1204,17 @@ class SindiInvertedIndex : public DimMapInvertedIndex(InvertedIndexEncoding::FIXED_DOCID_WINDOWS)) { + static_cast(packed_ ? InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7 + : InvertedIndexEncoding::FIXED_DOCID_WINDOWS)) { return Status::invalid_serialized_index_type; } // check window params reader.read(&this->window_size_, sizeof(uint32_t)); reader.read(&this->nr_windows_, sizeof(uint32_t)); + if (packed_) { + reader.advance(3 * sizeof(uint32_t)); // descriptor checked in preflight + } if (this->window_size_ == 0 || this->window_size_ >= 65536) { LOG_KNOWHERE_INFO_ << "SindiInvertedIndex::deserialize invalid window_size_=" @@ -1170,7 +1234,7 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_inner_dims_; - const uint64_t bytes_header = static_cast(sizeof(uint32_t) * 3); + const uint64_t bytes_header = static_cast(sizeof(uint32_t) * (packed_ ? 6 : 3)); window_index_plists_sz_.clear(); window_index_plists_sz_spans_.clear(); @@ -1217,15 +1281,23 @@ class SindiInvertedIndex : public DimMapInvertedIndex( - reinterpret_cast(reader.data() + reader.tellg()), nr_dims + 1); - reader.advance((nr_dims + 1) * sizeof(uint32_t)); + if (packed_) { + // Packed payload has byte alignment. Metadata uses memcpy, never typed byte casts. + plists_dim_offsets_.resize(nr_dims + 1); + reader.read(plists_dim_offsets_.data(), (nr_dims + 1) * sizeof(uint32_t)); + plists_dim_offsets_span_ = plists_dim_offsets_; + } else { + plists_dim_offsets_span_ = std::span( + reinterpret_cast(reader.data() + reader.tellg()), nr_dims + 1); + reader.advance((nr_dims + 1) * sizeof(uint32_t)); + } total_postings = plists_dim_offsets_span_[nr_dims]; // Validate total_postings against section size const uint64_t bytes_dim_offsets = static_cast(nr_dims + 1) * sizeof(uint32_t); - const uint64_t bytes_postings_data = - static_cast(total_postings) * (sizeof(uint16_t) + sizeof(QuantType)); + const uint64_t bytes_postings_data = packed_ ? 2 * sindi::packed12_bytes(total_postings) + : static_cast(total_postings) * + (sizeof(uint16_t) + sizeof(QuantType)); const uint64_t expected_section_bytes = bytes_header + bytes_mask + bytes_win_nnzs + bytes_dim_offsets + bytes_postings_data; if (expected_section_bytes != section_header.size) { @@ -1235,19 +1307,25 @@ class SindiInvertedIndex : public DimMapInvertedIndex(reader.data() + reader.tellg()); - const uint64_t bytes_ids = static_cast(total_postings) * sizeof(uint16_t); - total_plists_ids_flat_span_ = std::span(ids_region, total_postings); - reader.advance(total_postings * sizeof(uint16_t)); - - // vals region (per-dim contiguous, concatenated) - const QuantType* vals_region = - reinterpret_cast(reader.data() + reader.tellg()); - const uint64_t bytes_vals = static_cast(total_postings) * sizeof(QuantType); - total_plists_vals_flat_span_ = std::span(vals_region, total_postings); - reader.advance(total_postings * sizeof(QuantType)); + const uint64_t bytes_ids = + packed_ ? sindi::packed12_bytes(total_postings) : total_postings * sizeof(uint16_t); + const uint64_t bytes_vals = packed_ ? bytes_ids : total_postings * sizeof(QuantType); + if (packed_) { + packed_ids_span_ = {reader.data() + reader.tellg(), static_cast(bytes_ids)}; + reader.advance(bytes_ids); + packed_vals_span_ = {reader.data() + reader.tellg(), static_cast(bytes_vals)}; + reader.advance(bytes_vals); + packed_ready_ = true; + } else { + total_plists_ids_flat_span_ = { + reinterpret_cast(reader.data() + reader.tellg()), + static_cast(total_postings)}; + reader.advance(bytes_ids); + total_plists_vals_flat_span_ = { + reinterpret_cast(reader.data() + reader.tellg()), + static_cast(total_postings)}; + reader.advance(bytes_vals); + } // Log breakdown for POSTING_LISTS section LOG_KNOWHERE_DEBUG_ << "SindiInvertedIndex::deserialize POSTING_LISTS breakdown: " @@ -1266,10 +1344,16 @@ class SindiInvertedIndex : public DimMapInvertedIndex( - reinterpret_cast(reader.data() + section_header.offset), - this->nr_inner_dims_); - reader.advance(sizeof(float) * this->nr_inner_dims_); + if (packed_) { + max_scores_per_dim_.resize(this->nr_inner_dims_); + reader.read(max_scores_per_dim_.data(), sizeof(float) * this->nr_inner_dims_); + max_scores_per_dim_span_ = max_scores_per_dim_; + } else { + max_scores_per_dim_span_ = std::span( + reinterpret_cast(reader.data() + section_header.offset), + this->nr_inner_dims_); + reader.advance(sizeof(float) * this->nr_inner_dims_); + } break; } case InvertedIndexSectionType::BM25_U8_OVERFLOWS: { @@ -1332,8 +1416,9 @@ class SindiInvertedIndex : public DimMapInvertedIndex(InvertedIndexQuantType::IP_FP16)) { + if (metadata[0] != 1 || metadata[1] != (packed_ ? 2u : 1u) || + metadata[2] != static_cast(packed_ ? InvertedIndexQuantType::IP_E5M7 + : InvertedIndexQuantType::IP_FP16)) { return Status::invalid_serialized_index_type; } @@ -1357,9 +1442,25 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_inner_dims_; ++dim) { + float maximum = 0; + for (size_t j = 0; j < posting_count(dim); ++j) { + maximum = std::max(maximum, posting_value_at(dim, j)); + } + + if (max_scores_per_dim_span_[dim] != maximum) { + return Status::invalid_serialized_index_type; + } + } + + if (!refine_) { + std::vector{}.swap(refinement_seek_); + } + } } catch (const std::exception&) { return Status::invalid_serialized_index_type; } @@ -1465,8 +1566,8 @@ class SindiInvertedIndex : public DimMapInvertedIndex{} : posting_ids(qid); + const auto vals = packed_ready_ ? std::span{} : posting_vals(qid); bool is_sparse = !plists_wnnzs_fmts_msk_span_.empty() && ((plists_wnnzs_fmts_msk_span_[qid >> 3] & static_cast(0x1u << (qid & 0x7))) != 0); @@ -1500,6 +1601,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex curr_max_score) { curr_max_score = dispatch_max; } @@ -1697,8 +1799,8 @@ class SindiInvertedIndex : public DimMapInvertedIndex{} : posting_ids(qid); + const auto vals = packed_ready_ ? std::span{} : posting_vals(qid); bool is_sparse = !plists_wnnzs_fmts_msk_span_.empty() && ((plists_wnnzs_fmts_msk_span_[qid >> 3] & static_cast(0x1u << (qid & 0x7))) != 0); @@ -1736,6 +1838,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex(std::min(bm25_u16_value(value), 255)); } else if constexpr (is_bm25) { return static_cast(bm25_u16_value(value)); } else { - return static_cast(static_cast(value)); + const auto half = static_cast(static_cast(value)); + return packed_ ? sindi::decode_e5m7_half(sindi::encode_e5m7(half)) : half; } } @@ -2022,6 +2128,210 @@ class SindiInvertedIndex : public DimMapInvertedIndex row_sums_; std::span row_sums_span_; + static bool + validate_packed_file(const MemoryIOReader& reader) { + try { + const auto* data = reader.data(); + const size_t size = reader.total_; + auto u32 = [&](size_t offset) { + if (offset > size || size - offset < 4) + throw std::runtime_error("Truncated packed header"); + uint32_t value; + std::memcpy(&value, data + offset, 4); + return value; + }; + + if (size < 36 || u32(0) != kInvertedIndexFileFormatVersion || + u32(16) != static_cast(InvertedIndexQuantType::IP_E5M7)) { + return false; + } + + const size_t dims = u32(12), sections = u32(32); + if (sections < 3 || sections > (size - 36) / sizeof(InvertedIndexSectionHeader)) { + return false; + } + + const size_t directory_end = 36 + sections * sizeof(InvertedIndexSectionHeader); + std::vector headers(sections); + std::memcpy(headers.data(), data + 36, sections * sizeof(InvertedIndexSectionHeader)); + std::unordered_set types; + for (const auto& h : headers) { + switch (h.type) { + case InvertedIndexSectionType::POSTING_LISTS: + case InvertedIndexSectionType::DIM_MAP_REVERSE: + case InvertedIndexSectionType::DIM_MAP_MPHF: + case InvertedIndexSectionType::MAX_SCORES_PER_DIM: + case InvertedIndexSectionType::SINDI_REFINEMENT: + break; + default: + return false; + } + + if (!types.insert(static_cast(h.type)).second || h.offset < directory_end || + h.offset > size || h.size > size - h.offset) { + return false; + } + } + + auto by_offset = headers; + std::sort(by_offset.begin(), by_offset.end(), + [](const auto& a, const auto& b) { return a.offset < b.offset; }); + for (size_t i = 1; i < by_offset.size(); ++i) { + if (by_offset[i - 1].offset + by_offset[i - 1].size > by_offset[i].offset) { + return false; + } + } + + auto* h = find_section_header(headers, InvertedIndexSectionType::POSTING_LISTS); + auto* reverse = find_section_header(headers, InvertedIndexSectionType::DIM_MAP_REVERSE); + if (!reverse || ((reinterpret_cast(data) + reverse->offset) % alignof(uint32_t))) { + return false; + } + + auto* maxima = find_section_header(headers, InvertedIndexSectionType::MAX_SCORES_PER_DIM); + if (!h || !maxima || !find_section_header(headers, InvertedIndexSectionType::DIM_MAP_REVERSE) || + maxima->size != dims * 4 || h->size < 24 || + u32(h->offset) != static_cast(InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7) || + u32(h->offset + 4) != 4096 || u32(h->offset + 12) != 1 || u32(h->offset + 16) != 2 || + u32(h->offset + 20) != static_cast(InvertedIndexQuantType::IP_E5M7)) { + return false; + } + + if (u32(h->offset + 8) != (uint64_t(u32(4)) + 4095) / 4096) { + return false; + } + + size_t pos = h->offset + 24, end = h->offset + h->size; + auto advance = [&](size_t bytes) { + if (pos > end || bytes > end - pos) { + throw std::runtime_error("Truncated packed postings"); + } + pos += bytes; + }; + + advance((dims + 7) / 8); + if (dims > (end - pos) / 4) { + return false; + } + + for (size_t i = 0; i < dims; ++i) { + advance(4); + const auto bytes = u32(pos - 4); + advance(bytes); + } + + const size_t offsets = pos; + advance((dims + 1) * 4); + if (u32(offsets) != 0) { + return false; + } + + for (size_t i = 0; i < dims; ++i) { + if (u32(offsets + 4 * i) > u32(offsets + 4 * (i + 1))) { + return false; + } + } + + const auto count = u32(offsets + dims * 4); + const auto bytes = sindi::packed12_bytes(count); + advance(bytes); + advance(bytes); + + if (pos != end) { + return false; + } + if ((count & 1) && ((data[end - 1] & 0xf0) || (data[end - bytes - 1] & 0xf0))) { + return false; + } + + return true; + } catch (const std::exception&) { + return false; + } + } + + bool packed_ = false; + bool packed_ready_ = false; + std::vector packed_ids_, packed_vals_; + std::span packed_ids_span_, packed_vals_span_; + + size_t + posting_count(size_t dim) const { + return plists_dim_offsets_span_[dim + 1] - plists_dim_offsets_span_[dim]; + } + + uint16_t + posting_id_at(size_t dim, size_t pos) const { + return packed_ready_ ? sindi::unpack12(packed_ids_span_.data(), plists_dim_offsets_span_[dim] + pos) + : posting_ids(dim)[pos]; + } + + float + posting_value_at(size_t dim, size_t pos) const { + return packed_ready_ + ? sindi::decode_e5m7(sindi::unpack12(packed_vals_span_.data(), plists_dim_offsets_span_[dim] + pos)) + : static_cast(posting_vals(dim)[pos]); + } + + void + pack_postings() { + if constexpr (is_ip) { + const size_t count = plists_dim_offsets_span_.back(); + std::vector ids(sindi::packed12_bytes(count), 0), vals(ids.size(), 0); + max_scores_per_dim_.assign(this->nr_inner_dims_, 0); + // Sequential pair ownership avoids shared-nibble races at term boundaries. + for (size_t dim = 0; dim < this->nr_inner_dims_; ++dim) { + const auto source_ids = posting_ids(dim); + const auto source_vals = posting_vals(dim); + const size_t start = plists_dim_offsets_span_[dim]; + for (size_t j = 0; j < source_ids.size(); ++j) { + const auto code = sindi::encode_e5m7(source_vals[j]); + sindi::pack12(ids.data(), start + j, source_ids[j]); + sindi::pack12(vals.data(), start + j, code); + max_scores_per_dim_[dim] = std::max(max_scores_per_dim_[dim], sindi::decode_e5m7(code)); + } + } + + packed_ids_.swap(ids); + packed_vals_.swap(vals); + + packed_ids_span_ = packed_ids_; + packed_vals_span_ = packed_vals_; + max_scores_per_dim_span_ = max_scores_per_dim_; + aligned_u16_vec{}.swap(total_plists_ids_flat_); + aligned_quant_vec{}.swap(total_plists_vals_flat_); + std::vector{}.swap(total_plists_ids_); + std::vector{}.swap(total_plists_vals_); + std::vector>{}.swap(total_plists_ids_spans_); + std::vector>{}.swap(total_plists_vals_spans_); + total_plists_ids_flat_span_ = {}; + total_plists_vals_flat_span_ = {}; + packed_ready_ = true; + } + } + + void + expand_for_add() { + if constexpr (is_ip && AllowIncremental) { + total_plists_ids_.resize(this->nr_inner_dims_); + total_plists_vals_.resize(this->nr_inner_dims_); + for (size_t dim = 0; dim < this->nr_inner_dims_; ++dim) { + const size_t count = posting_count(dim); + auto& ids = total_plists_ids_[dim]; + auto& vals = total_plists_vals_[dim]; + ids.resize(count); + vals.resize(count); + for (size_t j = 0; j < count; ++j) { + ids[j] = posting_id_at(dim, j); + vals[j] = static_cast(posting_value_at(dim, j)); + } + } + + packed_ready_ = false; + // append_window_indexes publishes fresh typed spans before packing again. + } + } + bool legacy_dim_map_mphf_trailer_workaround_{true}; bool refine_ = false; std::vector refinement_seek_; @@ -2035,7 +2345,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex> 3] & (1u << (dim & 7))) != 0; const size_t entries = seek.sparse ? counts.size() / 4 : nr_windows_; @@ -2067,17 +2377,18 @@ class SindiInvertedIndex : public DimMapInvertedIndex window_size_ || sum + count > ids.size()) { + if (count > window_size_ || sum + count > count_ids) { throw std::runtime_error("Invalid refinement posting count"); } seek.offsets.push_back(static_cast(sum)); // Builders emit document order; validate it also on deserialization. for (size_t j = sum; j < sum + count; ++j) { - const float represented = static_cast(posting_vals(dim)[j]); + const float represented = posting_value_at(dim, j); if (!std::isfinite(represented) || represented < 0 || - uint64_t(wid) * window_size_ + ids[j] >= this->nr_rows_ || ids[j] >= window_size_ || - (j > sum && ids[j - 1] >= ids[j])) { + uint64_t(wid) * window_size_ + posting_id_at(dim, j) >= this->nr_rows_ || + posting_id_at(dim, j) >= window_size_ || + (j > sum && posting_id_at(dim, j - 1) >= posting_id_at(dim, j))) { // no good throw std::runtime_error("Invalid refinement posting IDs"); } @@ -2086,7 +2397,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex std::numeric_limits::max()) { + if (sum != count_ids || sum > std::numeric_limits::max()) { throw std::runtime_error("Invalid refinement posting offsets"); } @@ -2106,6 +2417,28 @@ class SindiInvertedIndex : public DimMapInvertedIndex candidates, std::span scores) const { + if (packed_) { + for (size_t i = 0; i < candidates.size(); ++i) { + const uint32_t local = candidates[i] % window_size_; + uint32_t low = begin, high = end; + while (low < high) { + const auto mid = low + (high - low) / 2; + if (posting_id_at(dim, mid) < local) { + low = mid + 1; + } else { + high = mid; + } + } + + begin = low; + if (low < end && posting_id_at(dim, low) == local) { + scores[i] = std::fma(weight, posting_value_at(dim, low), scores[i]); + } + } + + return; + } + const auto ids = posting_ids(dim); const auto vals = posting_vals(dim); if (begin == end) diff --git a/src/index/sparse/sindi_packed12.h b/src/index/sparse/sindi_packed12.h new file mode 100644 index 000000000..9141c04f8 --- /dev/null +++ b/src/index/sparse/sindi_packed12.h @@ -0,0 +1,71 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// Licensed under the Apache License, Version 2.0. +#pragma once +#include +#include +#include +#include +#include + +#include "knowhere/operands.h" + +namespace knowhere::sparse::inverted::sindi { + +// Two little-endian 12-bit words in three bytes; odd final words use two bytes. +inline size_t +packed12_bytes(size_t count) { + if (count > (std::numeric_limits::max() / 3) * 2) { + throw std::overflow_error("Packed posting length overflow"); + } + + return (count / 2) * 3 + (count % 2) * 2; +} + +inline uint16_t +unpack12(const uint8_t* bytes, size_t position) { + const size_t offset = (position / 2) * 3 + (position & 1); + return ((uint16_t(bytes[offset]) | (uint16_t(bytes[offset + 1]) << 8)) >> ((position & 1) * 4)) & 4095; +} + +// Sequential writes or independently owned complete pairs only: adjacent words share a byte. +inline void +pack12(uint8_t* bytes, size_t position, uint16_t code) { + if (code > 4095) { + throw std::invalid_argument("U12 overflow"); + } + + const size_t offset = (position / 2) * 3; + if (position & 1) { + bytes[offset + 1] = (bytes[offset + 1] & 15) | ((code & 15) << 4); + bytes[offset + 2] = code >> 4; + } else { + bytes[offset] = code & 255; + bytes[offset + 1] = code >> 8; + } +} + +inline uint16_t +encode_e5m7(knowhere::fp16 value) { + const auto bits = std::bit_cast(value); + if ((bits & 0x8000) || (bits & 0x7c00) == 0x7c00) { + throw std::invalid_argument("E5M7 requires finite nonnegative FP16, including positive zero"); + } + + return bits >> 3; +} + +inline knowhere::fp16 +decode_e5m7_half(uint16_t code) { + if (code >= 0xf80) { + throw std::invalid_argument("Reserved E5M7 code"); + } + + uint16_t bits = code << 3; + return std::bit_cast(bits); +} + +inline float +decode_e5m7(uint16_t code) { + return static_cast(decode_e5m7_half(code)); +} +} // namespace knowhere::sparse::inverted::sindi diff --git a/src/index/sparse/sindi_packed12_x86.h b/src/index/sparse/sindi_packed12_x86.h new file mode 100644 index 000000000..f2c395a74 --- /dev/null +++ b/src/index/sparse/sindi_packed12_x86.h @@ -0,0 +1,22 @@ +#pragma once + +#include + +#include + +namespace knowhere::sparse::inverted::sindi { + +// Eight contiguous U12 codes occupy exactly 12 bytes. Two bounded loads avoid +// reading beyond an exact-sized stream. Shuffle and extraction stay in registers. +static inline __m128i +unpack12_eight_x86(const uint8_t* p) { + uint32_t last; + std::memcpy(&last, p + 8, sizeof(last)); + + const auto raw = _mm_insert_epi32(_mm_loadl_epi64(reinterpret_cast(p)), last, 2); + const auto pairs = _mm_shuffle_epi8(raw, _mm_setr_epi8(0, 1, 1, 2, 3, 4, 4, 5, 6, 7, 7, 8, 9, 10, 10, 11)); + const auto aligned = _mm_blend_epi16(pairs, _mm_srli_epi16(pairs, 4), 0xaa); + return _mm_and_si128(aligned, _mm_set1_epi16(4095)); +} + +} // namespace knowhere::sparse::inverted::sindi diff --git a/src/index/sparse/sindi_simd.cc b/src/index/sparse/sindi_simd.cc index 33fd477a9..8e9a8522c 100644 --- a/src/index/sparse/sindi_simd.cc +++ b/src/index/sparse/sindi_simd.cc @@ -1,9 +1,46 @@ #include "index/sparse/sindi_simd.h" +#include "index/sparse/sindi_packed12.h" #include "simd/hook.h" namespace knowhere::sparse::inverted::sindi { +float +ip_accumulate_scalar_u12_e5m7(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out) { + float maximum = 0; + for (int32_t i = 0; i < n; ++i) { + auto id = unpack12(ids, start + i); + out[id] = std::fma(q, decode_e5m7(unpack12(vals, start + i)), out[id]); + maximum = std::max(maximum, out[id]); + } + + return maximum; +} + +packed_ip_accumulate_fn_t +get_packed_ip_kernel() { +#if defined(__x86_64__) + namespace cpu = faiss::cppcontrib::knowhere; + if (cpu::cpu_support_f16c() && __builtin_cpu_supports("fma")) { + if (cpu::use_avx512 && cpu::cpu_support_avx512() && __builtin_cpu_supports("avx512vl") && + __builtin_cpu_supports("avx512cd") && __builtin_cpu_supports("avx512f")) { + return ip_accumulate_avx512_u12_e5m7; + } + if (cpu::use_avx2 && cpu::cpu_support_avx2() && __builtin_cpu_supports("avx2")) { + return ip_accumulate_avx2_u12_e5m7; + } + } +#endif + +#if defined(__aarch64__) && defined(KNOWHERE_USE_SVE) + if (faiss::cppcontrib::knowhere::supports_sve()) { + return ip_accumulate_sve_u12_e5m7; + } +#endif + + return ip_accumulate_scalar_u12_e5m7; +} + float ip_accumulate_scalar_fp16(float qval, const knowhere::fp16* vals, const uint16_t* ids, int32_t num, float* out) { float max_val = 0.0f; diff --git a/src/index/sparse/sindi_simd.h b/src/index/sparse/sindi_simd.h index ac1c5b3fa..45678dded 100644 --- a/src/index/sparse/sindi_simd.h +++ b/src/index/sparse/sindi_simd.h @@ -8,6 +8,19 @@ namespace knowhere::sparse::inverted::sindi { +using packed_ip_accumulate_fn_t = float (*)(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*); + +packed_ip_accumulate_fn_t +get_packed_ip_kernel(); + +float +ip_accumulate_scalar_u12_e5m7(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*); + +#if defined(__aarch64__) && defined(KNOWHERE_USE_SVE) +float +ip_accumulate_sve_u12_e5m7(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*); +#endif + using ip_accumulate_fn_t = float (*)(float qval, const knowhere::fp16* vals, const uint16_t* ids, int32_t num, float* out); @@ -57,6 +70,10 @@ batch_insert_scalar(const float* scores, size_t docid_start, size_t count, knowhere::ResultMinHeap& topk_q, float& threshold, const BitsetView& bitset); #if defined(__x86_64__) +float +ip_accumulate_avx2_u12_e5m7(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*); +float +ip_accumulate_avx512_u12_e5m7(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*); // AVX2 implementations (compiled separately with -mavx2) float ip_accumulate_avx2_fp16(float qval, const knowhere::fp16* vals, const uint16_t* ids, int32_t num, float* out); diff --git a/src/index/sparse/sindi_simd_avx2.cc b/src/index/sparse/sindi_simd_avx2.cc index ee101b922..19c59871f 100644 --- a/src/index/sparse/sindi_simd_avx2.cc +++ b/src/index/sparse/sindi_simd_avx2.cc @@ -3,8 +3,55 @@ #if defined(__x86_64__) #include +#include "index/sparse/sindi_packed12_x86.h" + namespace knowhere::sparse::inverted::sindi { +float +ip_accumulate_avx2_u12_e5m7(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out) { + if (n <= 0) { + return 0; + } + + float maximum = 0; + if (start & 1) { + maximum = ip_accumulate_scalar_u12_e5m7(q, vals, ids, start, 1, out); + ++start; + --n; + } + + const auto vq = _mm256_set1_ps(q); + auto vmax = _mm256_setzero_ps(); + + int32_t i = 0; + for (; i + 8 <= n; i += 8) { + const auto* ip = ids + ((start + i) / 2) * 3; + const auto* vp = vals + ((start + i) / 2) * 3; + const auto id = _mm256_cvtepu16_epi32(unpack12_eight_x86(ip)); + const auto value = _mm256_cvtph_ps(_mm_slli_epi16(unpack12_eight_x86(vp), 3)); + const auto sum = _mm256_fmadd_ps(value, vq, _mm256_i32gather_ps(out, id, 4)); + + // AVX2 has gather but no scatter. + alignas(32) uint32_t indices[8]; + alignas(32) float scores[8]; + _mm256_store_si256(reinterpret_cast<__m256i*>(indices), id); + _mm256_store_ps(scores, sum); + for (int lane = 0; lane < 8; ++lane) { + out[indices[lane]] = scores[lane]; + } + + vmax = _mm256_max_ps(vmax, sum); + } + + alignas(32) float maxima[8]; + _mm256_store_ps(maxima, vmax); + for (float value : maxima) { + maximum = std::max(maximum, value); + } + + return std::max(maximum, ip_accumulate_scalar_u12_e5m7(q, vals, ids, start + i, n - i, out)); +} + float ip_accumulate_avx2_fp16(float qval, const knowhere::fp16* vals, const uint16_t* ids, int32_t num, float* out) { int32_t i = 0; diff --git a/src/index/sparse/sindi_simd_avx512.cc b/src/index/sparse/sindi_simd_avx512.cc index d9d45f548..7b0a87d6e 100644 --- a/src/index/sparse/sindi_simd_avx512.cc +++ b/src/index/sparse/sindi_simd_avx512.cc @@ -3,8 +3,60 @@ #if defined(__x86_64__) #include +#include "index/sparse/sindi_packed12_x86.h" + namespace knowhere::sparse::inverted::sindi { +float +ip_accumulate_avx512_u12_e5m7(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out) { + if (n <= 0) { + return 0; + } + + float maximum = 0; + if (start & 1) { + maximum = ip_accumulate_scalar_u12_e5m7(q, vals, ids, start, 1, out); + ++start; + --n; + } + + const auto vq = _mm512_set1_ps(q); + auto vmax = _mm512_setzero_ps(); + + int32_t i = 0; + for (; i + 16 <= n; i += 16) { + const auto* ip = ids + ((start + i) / 2) * 3; + const auto* vp = vals + ((start + i) / 2) * 3; + const auto id16 = _mm256_set_m128i(unpack12_eight_x86(ip + 12), unpack12_eight_x86(ip)); + const auto code = _mm256_set_m128i(unpack12_eight_x86(vp + 12), unpack12_eight_x86(vp)); + const auto id = _mm512_cvtepu16_epi32(id16); + const auto value = _mm512_cvtph_ps(_mm256_slli_epi16(code, 3)); + const auto sum = _mm512_fmadd_ps(value, vq, _mm512_i32gather_ps(id, out, 4)); + _mm512_i32scatter_ps(out, id, sum, 4); + vmax = _mm512_max_ps(vmax, sum); + } + + maximum = std::max(maximum, _mm512_reduce_max_ps(vmax)); + + // Match the FP16 kernel's eight-posting tail before the scalar remainder. + // The 256-bit scatter is available through AVX-512VL in this translation unit. + if (i + 8 <= n) { + const auto* ip = ids + ((start + i) / 2) * 3; + const auto* vp = vals + ((start + i) / 2) * 3; + const auto id = _mm256_cvtepu16_epi32(unpack12_eight_x86(ip)); + const auto value = _mm256_cvtph_ps(_mm_slli_epi16(unpack12_eight_x86(vp), 3)); + const auto sum = _mm256_fmadd_ps(value, _mm256_set1_ps(q), _mm256_i32gather_ps(out, id, 4)); + _mm256_i32scatter_ps(out, id, sum, 4); + auto tail_max = _mm_max_ps(_mm256_castps256_ps128(sum), _mm256_extractf128_ps(sum, 1)); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(2, 3, 0, 1))); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(1, 0, 3, 2))); + maximum = std::max(maximum, _mm_cvtss_f32(tail_max)); + i += 8; + } + + return std::max(maximum, ip_accumulate_scalar_u12_e5m7(q, vals, ids, start + i, n - i, out)); +} + float ip_accumulate_avx512_fp16(float qval, const knowhere::fp16* vals, const uint16_t* ids, int32_t num, float* out) { int32_t i = 0; diff --git a/src/index/sparse/sindi_simd_sve.cc b/src/index/sparse/sindi_simd_sve.cc index f7a5986be..27f7b1ce7 100644 --- a/src/index/sparse/sindi_simd_sve.cc +++ b/src/index/sparse/sindi_simd_sve.cc @@ -6,6 +6,156 @@ namespace knowhere::sparse::inverted::sindi { +// Direct U12-to-U32 expansion, adapted from the standalone packed12 SVE kernel. +// Tables are generated for the current VL and nibble phase. Byte predicates bound +// every load, including odd starts and final partial vectors. +static inline svuint32_t +load12(const std::uint8_t* base, std::uint32_t n, std::uint32_t phase, svuint32_t table, svuint32_t shift) { + const auto bytes = (n * 12 + phase + 7) / 8; + const auto raw = svld1_u8(svwhilelt_b8(0u, bytes), base); + const auto pairs = svreinterpret_u32_u8(svtbl_u8(raw, svreinterpret_u8_u32(table))); + return svand_n_u32_x(svptrue_b32(), svlsr_u32_x(svptrue_b32(), svand_n_u32_x(svptrue_b32(), pairs, 65535), shift), + 4095); +} + +// Decode two full U32 vectors from one packed-byte load. The second +// table advances by 3*vl/2 bytes; vl is always a multiple of four. +// Preserve the starting nibble phase and predicate the exact input byte count. +static inline svuint32x2_t +load_pair(const std::uint8_t* packed, svbool_t bytes, svuint8_t table0, svuint8_t table1, svuint32_t shift) { + const auto pg = svptrue_b32(); + const auto raw = svld1_u8(bytes, packed); + // Upper shuffled bytes are immaterial after the shift and U12 mask. + const auto a = svand_n_u32_x(pg, svlsr_u32_x(pg, svreinterpret_u32_u8(svtbl_u8(raw, table0)), shift), 4095); + const auto b = svand_n_u32_x(pg, svlsr_u32_x(pg, svreinterpret_u32_u8(svtbl_u8(raw, table1)), shift), 4095); + return svcreate2_u32(a, b); +} + +static inline svfloat32_t +expand_e5(svbool_t pg, svuint32_t code) { + // FCVT consumes the low binary16 half of each U32 lane, including subnormals. + return svcvt_f32_f16_x(pg, svreinterpret_f16_u32(svlsl_n_u32_x(pg, code, 3))); +} + +float +ip_accumulate_sve_u12_e5m7(float q, const uint8_t* __restrict vals, const uint8_t* __restrict ids, size_t start, + int32_t count, float* __restrict out) { + if (count <= 0) { + return 0; + } + + const auto vl = static_cast(svcntw()); + const auto pg = svptrue_b32(); + const auto lane = svindex_u32(0, 1); + const auto phase = static_cast((start & 1) * 4); + const auto byte = svlsr_n_u32_x(pg, svadd_n_u32_x(pg, svmul_n_u32_x(pg, lane, 12), phase), 3); + const auto table = svorr_u32_x(pg, byte, svlsl_n_u32_x(pg, svadd_n_u32_x(pg, byte, 1), 8)); + const auto shift = + svlsl_n_u32_x(pg, svand_n_u32_x(pg, sveor_n_u32_x(pg, lane, static_cast(start & 1)), 1), 2); + const auto table0 = svreinterpret_u8_u32(table); + const auto table1 = svadd_n_u8_x(svptrue_b8(), table0, static_cast(vl * 3 / 2)); + + // All loop increments are even posting counts, so the phase and full-pair + // byte predicate remain constant. Hoist these instead of recomputing them + // for each of the four pairs in the unrolled loop. + const auto pair_bytes = svwhilelt_b8(0u, 3 * vl + static_cast(start & 1)); + const size_t byte_start = (start / 2) * 3 + (start & 1); + vals += byte_start; + ids += byte_start; + + const auto vq = svdup_f32(q); + auto vmax = svdup_f32(0); + + uint32_t i = 0; + uint32_t n = static_cast(count); + + for (; i + 8 * vl <= n; i += 8 * vl) { + const auto ids0 = load_pair(ids + (i / 2) * 3 + 0 * vl * 3, pair_bytes, table0, table1, shift); + const auto words0 = load_pair(vals + (i / 2) * 3 + 0 * vl * 3, pair_bytes, table0, table1, shift); + const auto id0 = svget2_u32(ids0, 0); + const auto id1 = svget2_u32(ids0, 1); + const auto v0 = expand_e5(pg, svget2_u32(words0, 0)); + const auto v1 = expand_e5(pg, svget2_u32(words0, 1)); + const auto ids2 = load_pair(ids + (i / 2) * 3 + 1 * vl * 3, pair_bytes, table0, table1, shift); + const auto words2 = load_pair(vals + (i / 2) * 3 + 1 * vl * 3, pair_bytes, table0, table1, shift); + const auto id2 = svget2_u32(ids2, 0); + const auto id3 = svget2_u32(ids2, 1); + const auto v2 = expand_e5(pg, svget2_u32(words2, 0)); + const auto v3 = expand_e5(pg, svget2_u32(words2, 1)); + const auto ids4 = load_pair(ids + (i / 2) * 3 + 2 * vl * 3, pair_bytes, table0, table1, shift); + const auto words4 = load_pair(vals + (i / 2) * 3 + 2 * vl * 3, pair_bytes, table0, table1, shift); + const auto id4 = svget2_u32(ids4, 0); + const auto id5 = svget2_u32(ids4, 1); + const auto v4 = expand_e5(pg, svget2_u32(words4, 0)); + const auto v5 = expand_e5(pg, svget2_u32(words4, 1)); + const auto ids6 = load_pair(ids + (i / 2) * 3 + 3 * vl * 3, pair_bytes, table0, table1, shift); + const auto words6 = load_pair(vals + (i / 2) * 3 + 3 * vl * 3, pair_bytes, table0, table1, shift); + const auto id6 = svget2_u32(ids6, 0); + const auto id7 = svget2_u32(ids6, 1); + const auto v6 = expand_e5(pg, svget2_u32(words6, 0)); + const auto v7 = expand_e5(pg, svget2_u32(words6, 1)); + const auto old0 = svld1_gather_u32index_f32(pg, out, id0); + const auto old1 = svld1_gather_u32index_f32(pg, out, id1); + const auto old2 = svld1_gather_u32index_f32(pg, out, id2); + const auto old3 = svld1_gather_u32index_f32(pg, out, id3); + const auto old4 = svld1_gather_u32index_f32(pg, out, id4); + const auto old5 = svld1_gather_u32index_f32(pg, out, id5); + const auto old6 = svld1_gather_u32index_f32(pg, out, id6); + const auto old7 = svld1_gather_u32index_f32(pg, out, id7); + const auto sum0 = svmad_f32_x(pg, v0, vq, old0); + svst1_scatter_u32index_f32(pg, out, id0, sum0); + vmax = svmax_f32_x(pg, vmax, sum0); + const auto sum1 = svmad_f32_x(pg, v1, vq, old1); + svst1_scatter_u32index_f32(pg, out, id1, sum1); + vmax = svmax_f32_x(pg, vmax, sum1); + const auto sum2 = svmad_f32_x(pg, v2, vq, old2); + svst1_scatter_u32index_f32(pg, out, id2, sum2); + vmax = svmax_f32_x(pg, vmax, sum2); + const auto sum3 = svmad_f32_x(pg, v3, vq, old3); + svst1_scatter_u32index_f32(pg, out, id3, sum3); + vmax = svmax_f32_x(pg, vmax, sum3); + const auto sum4 = svmad_f32_x(pg, v4, vq, old4); + svst1_scatter_u32index_f32(pg, out, id4, sum4); + vmax = svmax_f32_x(pg, vmax, sum4); + const auto sum5 = svmad_f32_x(pg, v5, vq, old5); + svst1_scatter_u32index_f32(pg, out, id5, sum5); + vmax = svmax_f32_x(pg, vmax, sum5); + const auto sum6 = svmad_f32_x(pg, v6, vq, old6); + svst1_scatter_u32index_f32(pg, out, id6, sum6); + vmax = svmax_f32_x(pg, vmax, sum6); + const auto sum7 = svmad_f32_x(pg, v7, vq, old7); + svst1_scatter_u32index_f32(pg, out, id7, sum7); + vmax = svmax_f32_x(pg, vmax, sum7); + } + + for (; i + 2 * vl <= n; i += 2 * vl) { + const auto pair_ids = load_pair(ids + (i / 2) * 3, pair_bytes, table0, table1, shift); + const auto words = load_pair(vals + (i / 2) * 3, pair_bytes, table0, table1, shift); + const auto id0 = svget2_u32(pair_ids, 0), id1 = svget2_u32(pair_ids, 1); + const auto v0 = expand_e5(pg, svget2_u32(words, 0)); + const auto v1 = expand_e5(pg, svget2_u32(words, 1)); + const auto sum0 = svmad_f32_x(pg, v0, vq, svld1_gather_u32index_f32(pg, out, id0)); + svst1_scatter_u32index_f32(pg, out, id0, sum0); + vmax = svmax_f32_x(pg, vmax, sum0); + const auto sum1 = svmad_f32_x(pg, v1, vq, svld1_gather_u32index_f32(pg, out, id1)); + svst1_scatter_u32index_f32(pg, out, id1, sum1); + vmax = svmax_f32_x(pg, vmax, sum1); + } + + for (; i < n; i += vl) { + const auto active = std::min(vl, n - i); + const auto tail = svwhilelt_b32(0u, active); + const auto tail_ids = load12(ids + (i / 2) * 3, active, phase, table, shift); + const auto v = expand_e5(tail, load12(vals + (i / 2) * 3, active, phase, table, shift)); + const auto old = svld1_gather_u32index_f32(tail, out, tail_ids); + const auto sum = svmad_f32_x(tail, v, vq, old); + svst1_scatter_u32index_f32(tail, out, tail_ids, sum); + vmax = svmax_f32_m(tail, vmax, sum); + } + + return svmaxv_f32(pg, vmax); +} + float ip_accumulate_sve_fp16(float qval, const knowhere::fp16* __restrict vals, const uint16_t* __restrict ids, int32_t num, float* __restrict out) { diff --git a/src/index/sparse/sparse_index_config.h b/src/index/sparse/sparse_index_config.h index ab173e5fa..78eb5190c 100644 --- a/src/index/sparse/sparse_index_config.h +++ b/src/index/sparse/sparse_index_config.h @@ -176,7 +176,8 @@ class SparseInvertedIndexConfig : public BaseConfig { .for_deserialize_from_file(); KNOWHERE_CONFIG_DECLARE_FIELD(quant_type) .description( - "quantization type for posting list values: fp16/fp32 for IP, u8/u16/u32/auto for BM25; u8 is " + "quantization type for posting list values: fp16/fp32/e5m7 for IP (e5m7 requires SINDI, window=4096), " + "u8/u16/u32/auto for BM25; u8 is " "supported only by sealed SINDI with index version >= 11; BM25 auto requires index version >= 11 " "and resolves to u8/u16 for sealed SINDI or u16 for other indexes; the concrete type is persisted " "in the index and restored automatically on load; the load parameter is used only for legacy " @@ -239,9 +240,9 @@ class SparseInvertedIndexConfig : public BaseConfig { auto qt = quant_type.value(); auto mt = metric_type.value(); if (mt == metric::IP) { - if (qt != "fp16" && qt != "fp32") { + if (qt != "fp16" && qt != "fp32" && qt != "e5m7") { if (err_msg) { - *err_msg = "quant_type for IP metric must be 'fp16' or 'fp32', got '" + qt + "'"; + *err_msg = "quant_type for IP metric must be 'fp16', 'fp32', or 'e5m7', got '" + qt + "'"; } return Status::invalid_args; } diff --git a/src/index/sparse/sparse_index_node.cc b/src/index/sparse/sparse_index_node.cc index 102c646d1..8d3998498 100644 --- a/src/index/sparse/sparse_index_node.cc +++ b/src/index/sparse/sparse_index_node.cc @@ -222,6 +222,10 @@ class SparseInvertedIndexNode : public IndexNode { std::string resolved_quant_type; bool is_ip = false; switch (posting_type.value()) { + case InvertedIndexQuantType::IP_E5M7: + resolved_quant_type = "e5m7"; + is_ip = true; + break; case InvertedIndexQuantType::IP_FP16: resolved_quant_type = "fp16"; is_ip = true; @@ -247,6 +251,10 @@ class SparseInvertedIndexNode : public IndexNode { if (!IsMetricType(cfg.metric_type.value(), is_ip ? metric::IP : metric::BM25)) { return Status::invalid_serialized_index_type; } + if (resolved_quant_type == "e5m7") { + cfg.inverted_index_algo = "SINDI"; + cfg.sindi_window_size = 4096; + } cfg.quant_type = resolved_quant_type; return Status::success; } @@ -315,7 +323,7 @@ class SparseInvertedIndexNode : public IndexNode { auto queries = static_cast*>(dataset->GetTensor()); auto nq = dataset->GetRows(); - if (refine) { + if (refine || PackedStorageEnabled()) { for (int64_t i = 0; i < nq; ++i) { if (!sparse::inverted::sindi::valid_refinement_row(queries[i])) { return expected::Err(Status::invalid_args, "Invalid SINDI refinement query"); @@ -362,6 +370,15 @@ class SparseInvertedIndexNode : public IndexNode { } auto search_params = search_params_or.value(); + if (PackedStorageEnabled()) { + const auto* queries = static_cast*>(dataset->GetTensor()); + for (int64_t i = 0; i < nq; ++i) { + if (!sparse::inverted::sindi::valid_refinement_row(queries[i])) { + return expected>::Err(Status::invalid_args, + "Invalid E5M7 query"); + } + } + } auto vec = std::vector>(nq, nullptr); const auto& id_map = this->GetIdMap(); const auto* result_id_map = this->SearchResultIdMap(id_map); @@ -470,6 +487,19 @@ class SparseInvertedIndexNode : public IndexNode { LOG_KNOWHERE_ERROR_ << "Failed to create index from BinarySet with name " << Type(); return index_or.error(); } + + if (cfg.quant_type.value_or("") == "e5m7") { + auto candidate = std::move(index_or.value()); + MemoryIOReader packed_reader(binary->data.get(), binary->size); + const auto status = candidate->deserialize(packed_reader); + if (status != Status::success) { + return status; + } + index_ = std::move(candidate); + binary_ = binary; + return Status::success; + } + index_ = std::move(index_or.value()); // deserialize index from binary @@ -525,6 +555,19 @@ class SparseInvertedIndexNode : public IndexNode { if (!index_or.has_value()) { return index_or.error(); } + + if (cfg.quant_type.value_or("") == "e5m7") { + auto candidate = std::move(index_or.value()); + MemoryIOReader packed_reader(reinterpret_cast(mapped_memory), map_size); + const auto status = candidate->deserialize(packed_reader); + if (status != Status::success) { + return status; + } + index_ = std::move(candidate); + this->mmap_guard_ = std::move(mmap_guard); + return Status::success; + } + index_ = std::move(index_or.value()); // deserialize index from mapped memory @@ -714,18 +757,20 @@ class SparseInvertedIndexNode : public IndexNode { // When encoding is available and not FIXED_DOCID_WINDOWS, the file was not // built with SINDI, so use create_index_before_v10 to match the actual file encoding. bool use_sindi = - algo == "SINDI" && (!encoding.has_value() || - encoding.value() == sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS); + algo == "SINDI" && + (!encoding.has_value() || + (encoding.value() == sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS || + encoding.value() == sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7)); if (use_sindi) { auto window_size = cfg.sindi_window_size.value_or(sparse::inverted::SindiInvertedIndexIP::max_window_size); IndexPtr index; if (is_growable) { index = std::make_unique( - window_size, cfg.refine.value_or(false)); + window_size, cfg.refine.value_or(false), cfg.quant_type.value_or("") == "e5m7"); } else { - index = std::make_unique(window_size, - cfg.refine.value_or(false)); + index = std::make_unique( + window_size, cfg.refine.value_or(false), cfg.quant_type.value_or("") == "e5m7"); } ConfigureSindiSerialization(index.get()); index->set_build_algo(algo); @@ -775,13 +820,25 @@ class SparseInvertedIndexNode : public IndexNode { expected>> CreateIndex(const SparseInvertedIndexConfig& cfg, bool is_growable = false, std::optional encoding = std::nullopt) const { + if (cfg.quant_type.value_or("") == "e5m7") { + const bool loading = encoding.has_value(); + if (index_version_ < 11 || !IsMetricType(cfg.metric_type.value(), metric::IP) || + (!loading && NormalizeInvertedIndexAlgo(cfg.inverted_index_algo.value_or("")) != "SINDI") || + (!loading && cfg.sindi_window_size.value_or(4096) != 4096) || + !cfg.inverted_index_codec.value_or("").empty() || + (loading && *encoding != sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7)) { + return expected>>::Err( + Status::invalid_args, "e5m7 requires SINDI IP, window=4096, version>=11, and no block codec"); + } + } if (cfg.refine.value_or(false)) { const auto algo = NormalizeInvertedIndexAlgo(cfg.inverted_index_algo.value_or("")); if (index_version_ < 11 || !IsMetricType(cfg.metric_type.value(), metric::IP) || (!algo.empty() && algo != "SINDI") || - (!cfg.quant_type.value_or("").empty() && cfg.quant_type.value() != "fp16")) { + (!cfg.quant_type.value_or("").empty() && cfg.quant_type.value() != "fp16" && + cfg.quant_type.value() != "e5m7")) { return expected>>::Err( - Status::invalid_args, "refine requires SINDI FP16 IP version >= 11"); + Status::invalid_args, "refine requires SINDI FP16/E5M7 IP version >= 11"); } } @@ -806,7 +863,7 @@ class SparseInvertedIndexNode : public IndexNode { using sparse::inverted::IndexScorerType; if (IsMetricType(cfg.metric_type.value(), metric::IP)) { // version < threshold forces fp32; version >= threshold defaults to fp16, user can override to fp32 - bool use_fp16 = version_support_fp16_quant_for_ip() && (qt == "fp16" || qt.empty()); + bool use_fp16 = version_support_fp16_quant_for_ip() && (qt == "fp16" || qt == "e5m7" || qt.empty()); if (use_fp16) { return CreateIndexImpl(cfg, is_growable, encoding); } else { @@ -829,6 +886,19 @@ class SparseInvertedIndexNode : public IndexNode { } } + bool + PackedStorageEnabled() const { + if (auto* index = dynamic_cast(index_.get())) { + return index->packed_storage_enabled(); + } + + if (auto* index = dynamic_cast(index_.get())) { + return index->packed_storage_enabled(); + } + + return false; + } + bool version_sindi_uses_mphf_section() const { return index_version_ >= kSindiMphfSectionMinVersion; diff --git a/tests/ut/test_sparse.cc b/tests/ut/test_sparse.cc index 86f71f210..338b7847d 100644 --- a/tests/ut/test_sparse.cc +++ b/tests/ut/test_sparse.cc @@ -2982,3 +2982,242 @@ TEST_CASE("SINDI refinement empty candidates and fractional API pool", "[sparse] unsupported["quant_type"] = "fp32"; REQUIRE(other.Build(RefineDataset(base), unsupported) != Status::success); } + +#include "index/sparse/sindi_packed12.h" +#include "index/sparse/sindi_simd.h" +namespace { + +float +E5Oracle(uint16_t code) { + const int exponent = code >> 7, fraction = code & 127; + return exponent == 0 ? std::ldexp(float(fraction), -21) : std::ldexp(1.0f + float(fraction) / 128, exponent - 15); +} + +float +E5Represented(float value) { + knowhere::fp16 half(value); + uint16_t bits; + std::memcpy(&bits, &half, 2); + return E5Oracle(bits >> 3); +} + +} // namespace + +TEST_CASE("SINDI U12 E5M7 exhaustive codecs", "[sparse][sindi][u12]") { + using namespace knowhere::sparse::inverted::sindi; + std::vector bytes(packed12_bytes(4097), 0); + for (size_t i = 0; i < 4097; ++i) pack12(bytes.data(), i, i & 4095); + for (size_t i = 0; i < 4097; ++i) REQUIRE(unpack12(bytes.data(), i) == (i & 4095)); + REQUIRE((bytes.back() & 0xf0) == 0); + for (uint16_t c = 0; c < 4096; ++c) { + if (c < 3968) + REQUIRE(decode_e5m7(c) == E5Oracle(c)); + else + REQUIRE_THROWS(decode_e5m7(c)); + } + for (uint32_t b = 0; b < 65536; ++b) { + uint16_t bits = b; + knowhere::fp16 half; + std::memcpy(&half, &bits, 2); + if ((b & 0x8000) || (b & 0x7c00) == 0x7c00) + REQUIRE_THROWS(encode_e5m7(half)); + else + REQUIRE(encode_e5m7(half) == (b >> 3)); + } + REQUIRE_THROWS(packed12_bytes(std::numeric_limits::max())); +} + +TEST_CASE("SINDI U12 E5M7 dispatched kernel tails and unroll", "[sparse][sindi][u12]") { + using namespace knowhere::sparse::inverted::sindi; + std::vector kernels = {ip_accumulate_scalar_u12_e5m7, get_packed_ip_kernel()}; +#if defined(__x86_64__) + if (__builtin_cpu_supports("avx2") && __builtin_cpu_supports("f16c") && __builtin_cpu_supports("fma")) + kernels.push_back(ip_accumulate_avx2_u12_e5m7); + if (__builtin_cpu_supports("avx512f") && __builtin_cpu_supports("avx512bw") && __builtin_cpu_supports("avx512vl") && + __builtin_cpu_supports("avx512dq") && __builtin_cpu_supports("avx512cd") && __builtin_cpu_supports("f16c") && + __builtin_cpu_supports("fma")) + kernels.push_back(ip_accumulate_avx512_u12_e5m7); +#endif + std::vector lengths = {0}; + for (size_t i = 1; i <= 129; ++i) lengths.push_back(i); + for (auto i : {255, 256, 257, 511, 512, 513, 4095, 4096}) lengths.push_back(i); + for (size_t start = 0; start < 4; ++start) + for (size_t n : lengths) { + // Deliberately unaligned byte bases; allocation ends exactly at stream end. + std::vector ids(packed12_bytes(start + n) + 1), vals(ids.size()); + for (size_t j = 0; j < start + n; ++j) { + pack12(ids.data() + 1, j, (j * 17) % 4096); + pack12(vals.data() + 1, j, (j * 37) % 3968); + } + std::vector scores(4096, .25f), expected = scores; + float maximum = 0; + for (size_t j = start; j < start + n; ++j) { + const auto id = (j * 17) % 4096; + expected[id] = std::fma(.75f, E5Oracle((j * 37) % 3968), expected[id]); + maximum = std::max(maximum, expected[id]); + } + for (auto kernel : kernels) { + std::fill(scores.begin(), scores.end(), .25f); + REQUIRE(kernel(.75f, vals.data() + 1, ids.data() + 1, start, n, scores.data()) == maximum); + REQUIRE(scores == expected); + } + } +} + +TEST_CASE("SINDI U12 E5M7 API oracle Add and persistence", "[sparse][sindi][u12]") { + using namespace knowhere; + const bool refine = GENERATE(false, true); + const bool growable = GENERATE(false, true); + std::vector base; + for (size_t i = 0; i < 8203; ++i) { + if (i % 4095 == 0) + base.push_back(RefineTestRow({{0, 1.003f + float(i % 97) / 100}, {7, .1257f}, {100000, 2.019f}})); + else + base.push_back(RefineTestRow({{0, 1.003f + float(i % 97) / 100}})); + } + auto build = RefineBuild(refine); + build["quant_type"] = "e5m7"; + const auto type = growable ? IndexEnum::INDEX_SPARSE_INVERTED_INDEX_CC : IndexEnum::INDEX_SPARSE_INVERTED_INDEX; + auto idx = IndexFactory::Instance().Create(type, 11).value(); + if (growable) { + std::vector first(base.begin(), base.begin() + 4095), rest(base.begin() + 4095, base.end()); + REQUIRE(idx.Build(RefineDataset(first), build) == Status::success); + REQUIRE(idx.Add(RefineDataset(rest), build) == Status::success); + std::vector invalid{RefineTestRow({{0, -0.0f}})}; + REQUIRE(idx.Add(RefineDataset(invalid), build) != Status::success); + REQUIRE(idx.Count() == base.size()); + } else + REQUIRE(idx.Build(RefineDataset(base), build) == Status::success); + std::vector query{RefineTestRow({{0, 2}, {7, .5f}, {100000, 1}})}; + Json search{{"metric_type", "IP"}, {"k", base.size() + 3}}; + if (refine) { + search["sindi_query_mass"] = .5f; + search["refine_k"] = 1.5f; + } + auto check = [&](auto& index) { + auto result = index.Search(RefineDataset(query), search, nullptr); + REQUIRE(result.has_value()); + std::set seen; + for (size_t j = 0; j < base.size(); ++j) { + auto id = result.value()->GetIds()[j]; + REQUIRE(id >= 0); + REQUIRE(id < int64_t(base.size())); + REQUIRE(seen.insert(id).second); + float score = 2 * E5Represented(base[id][0].val); + if (base[id].size() > 1) { + score = std::fma(1.0f, E5Represented(base[id][2].val), score); + score = std::fma(.5f, E5Represented(base[id][1].val), score); + } + REQUIRE(std::abs(result.value()->GetDistance()[j] - score) < 1e-6f); + } + REQUIRE(result.value()->GetIds()[base.size()] == -1); + std::vector mask((base.size() + 7) / 8, 255); + auto empty = index.Search(RefineDataset(query), search, BitsetView(mask.data(), base.size())); + REQUIRE(empty.has_value()); + REQUIRE(empty.value()->GetIds()[0] == -1); + }; + check(idx); + if (!growable) { + BinarySet bytes; + REQUIRE(idx.Serialize(bytes) == Status::success); + auto restored = IndexFactory::Instance().Create(type, 11).value(); + REQUIRE(restored.Deserialize(bytes, Json{{"metric_type", "IP"}}) == Status::success); + check(restored); + SparseQuantIndexFile file(bytes.GetByName(type)); + auto mapped = IndexFactory::Instance().Create(type, 11).value(); + REQUIRE(mapped.DeserializeFromFile(file.path, Json{{"metric_type", "IP"}, {"enable_mmap", true}}) == + Status::success); + check(mapped); + auto blob = bytes.GetByName(type); + using namespace sparse::inverted; + uint32_t count; + std::memcpy(&count, blob->data.get() + 32, 4); + for (size_t i = 0; i < count; ++i) { + InvertedIndexSectionHeader h; + std::memcpy(&h, blob->data.get() + 36 + i * sizeof(h), sizeof(h)); + if (h.type == InvertedIndexSectionType::POSTING_LISTS) { + uint32_t unsupported = 99; + std::memcpy(blob->data.get() + h.offset + 12, &unsupported, 4); + } + } + REQUIRE(restored.Deserialize(bytes, Json{{"metric_type", "IP"}}) != Status::success); + } +} + +TEST_CASE("SINDI U12 E5M7 invalid build contracts", "[sparse][sindi][u12]") { + using namespace knowhere; + std::vector rows{RefineTestRow({{0, 1}})}; + auto build = RefineBuild(false); + build["quant_type"] = "e5m7"; + for (auto change : {Json{{"metric_type", "BM25"}}, Json{{"inverted_index_algo", "DAAT_WAND"}}, + Json{{"sindi_window_size", 1024}}, Json{{"inverted_index_codec", "block_streamvbyte"}}}) { + auto cfg = build; + cfg.update(change); + auto idx = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(idx.Build(RefineDataset(rows), cfg) != Status::success); + } + for (float value : + {-1.f, -0.f, std::numeric_limits::infinity(), std::numeric_limits::quiet_NaN(), 1e10f}) { + auto idx = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + std::vector invalid{RefineTestRow({{0, value}})}; + REQUIRE(idx.Build(RefineDataset(invalid), build) != Status::success); + } +} + +TEST_CASE("SINDI U12 E5M7 sparse windows new dimensions and failure isolation", "[sparse][sindi][u12]") { + using namespace knowhere; + for (bool growable : {false, true}) { + const auto type = growable ? IndexEnum::INDEX_SPARSE_INVERTED_INDEX_CC : IndexEnum::INDEX_SPARSE_INVERTED_INDEX; + std::vector base(70001); + base[0] = RefineTestRow({{0, 10}}); + base[4095] = RefineTestRow({{0, 9}}); + base[4096] = RefineTestRow({{0, 8}, {100000, 20}}); + base[70000] = RefineTestRow({{0, 7}, {100000, 30}}); + auto cfg = RefineBuild(); + cfg["quant_type"] = "e5m7"; + auto idx = IndexFactory::Instance().Create(type, 11).value(); + if (growable) { + std::vector first(base.begin(), base.begin() + 4096), last(base.begin() + 4096, base.end()); + REQUIRE(idx.Build(RefineDataset(first), cfg) == Status::success); + REQUIRE(idx.Add(RefineDataset(last), cfg) == Status::success); + } else + REQUIRE(idx.Build(RefineDataset(base), cfg) == Status::success); + std::vector query{RefineTestRow({{0, 2}, {100000, 1}})}; + auto search = RefineSearch(); + search["refine_k"] = 4; + auto result = idx.Search(RefineDataset(query), search, nullptr); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == 70000); + REQUIRE(result.value()->GetDistance()[0] == 44); + std::vector mask((base.size() + 7) / 8, 0); + std::fill(mask.begin(), mask.begin() + 4096 / 8, 255); + mask[70000 / 8] |= 1u << (70000 % 8); + auto filtered = idx.Search(RefineDataset(query), search, BitsetView(mask.data(), base.size())); + REQUIRE(filtered.has_value()); + REQUIRE(filtered.value()->GetIds()[0] == 4096); + REQUIRE(filtered.value()->GetDistance()[0] == 36); + Json range{{"metric_type", "IP"}, {"radius", 30}}; + auto ranged = idx.RangeSearch(RefineDataset(query), range, nullptr); + REQUIRE(ranged.has_value()); + REQUIRE(ranged.value()->GetLims()[1] == 2); + if (!growable) { + BinarySet bytes; + REQUIRE(idx.Serialize(bytes) == Status::success); + auto blob = bytes.GetByName(type); + auto original = std::vector(blob->data.get(), blob->data.get() + blob->size); + using namespace sparse::inverted; + InvertedIndexSectionHeader h; + std::memcpy(&h, blob->data.get() + 36, sizeof(h)); + const std::vector corrupt_positions = {size_t(h.offset + 12), size_t(h.offset + 16), + size_t(h.offset + h.size - 1)}; + for (auto position : corrupt_positions) { + blob->data[position] = 255; + REQUIRE(idx.Deserialize(bytes, Json{{"metric_type", "IP"}}) != Status::success); + auto after = idx.Search(RefineDataset(query), search, nullptr); + REQUIRE(after.has_value()); + REQUIRE(after.value()->GetIds()[0] == 70000); + std::memcpy(blob->data.get(), original.data(), original.size()); + } + } + } +} From 7944e20e6bc35831ccd63c2703da067ed34bc80c Mon Sep 17 00:00:00 2001 From: Alexandr Guzhva Date: Mon, 5 Oct 2026 19:09:25 -0400 Subject: [PATCH 3/3] SINDI BM25 U4 LUT Signed-off-by: Alexandr Guzhva --- benchmark/benchmark_sparse.cpp | 387 +++++++++++++--- benchmark/benchmark_sparse_bm25_compact.cpp | 466 ++++++++++++++++++++ src/index/sparse/inverted_index.h | 2 + src/index/sparse/inverted_index_format.h | 2 + src/index/sparse/sindi_bm25_u4.h | 125 ++++++ src/index/sparse/sindi_inverted_index.h | 269 ++++++++--- src/index/sparse/sindi_simd.cc | 64 ++- src/index/sparse/sindi_simd.h | 31 ++ src/index/sparse/sindi_simd_avx2.cc | 103 +++++ src/index/sparse/sindi_simd_avx512.cc | 127 ++++++ src/index/sparse/sindi_simd_sve.cc | 72 +++ src/index/sparse/sparse_index_config.h | 20 +- src/index/sparse/sparse_index_node.cc | 83 +++- tests/ut/test_sparse.cc | 294 ++++++++++++ 14 files changed, 1917 insertions(+), 128 deletions(-) create mode 100644 benchmark/benchmark_sparse_bm25_compact.cpp create mode 100644 src/index/sparse/sindi_bm25_u4.h diff --git a/benchmark/benchmark_sparse.cpp b/benchmark/benchmark_sparse.cpp index 64fdef8e0..8a3920ba8 100644 --- a/benchmark/benchmark_sparse.cpp +++ b/benchmark/benchmark_sparse.cpp @@ -38,6 +38,8 @@ #include "knowhere/index/index_factory.h" #include "knowhere/sparse_utils.h" #include "knowhere/version.h" +#include "src/index/sparse/inverted_index_format.h" +#include "src/index/sparse/sindi_packed12.h" #include "src/index/sparse/sindi_refinement.h" #include "src/index/sparse/sindi_simd.h" @@ -184,9 +186,13 @@ LoadTruth(const std::string& path, const std::vector& selection, uint6 struct Options { std::string data, output = "sparse_results.csv", method = "all"; uint64_t queries = 100, offset = 0, base_limit = 0; - bool refine = false; + bool refine = false, size_only = false; + std::string quant_type = "fp16"; float mass = 1, factor = 1, drop = 0; - std::string sweep, split = "dev"; + std::string sweep, split = "dev", scheme = "mass"; + int samples = 100; + uint32_t pareto_seed = 42; + float factor_min = 1, factor_max = 20, mass_min = .4f, mass_max = 1; int threads = std::max(2u, std::thread::hardware_concurrency()), repeats = 3, k = 10; std::optional seed; }; @@ -194,14 +200,19 @@ struct Options { Options Parse(int argc, char** argv) { Options o; + std::set supplied; for (int i = 1; i < argc; ++i) { const std::string key = argv[i]; if (key == "--help") { std::cout << "benchmark_sparse --data-dir DIR [--queries 100] [--query-offset 0] [--seed N]\n" " [--threads N] [--repeats 3] [--k 10] [--output sparse_results.csv]\n" + " [--quant-type fp16|e5m7] [--size-only 0|1] (SINDI only)\n" " [--refine 0|1] [--mass 1] [--factor 1] [--drop 0]\n" - " [--sweep mass|count|drop] [--split dev|hidden]\n" + " [--sweep mass|count|drop|pareto] [--split dev|hidden]\n" + " Pareto: [--scheme mass] [--samples 100] [--pareto-seed 42]\n" + " [--factor-min 1] [--factor-max 20] [--mass-min 0.4] [--mass-max 1]\n" + " Writes .points.csv, .frontier.csv and .pareto.svg; builds once.\n" " [--method " "all|TAAT_NAIVE|DAAT_WAND|DAAT_MAXSCORE|BLOCK_MAX_MAXSCORE|BLOCK_MAX_WAND|SINDI|SPARSE_WAND]\n" " [--base-limit N] Reduced-base validation against brute force, not supplied ground truth.\n"; @@ -209,10 +220,15 @@ Parse(int argc, char** argv) { } Check(i + 1 < argc, "Missing value for " + key); const std::string value = argv[++i]; + supplied.insert(key); if (key == "--data-dir") o.data = value; else if (key == "--output") o.output = value; + else if (key == "--quant-type") + o.quant_type = value; + else if (key == "--size-only") + o.size_only = std::stoi(value) != 0; else if (key == "--refine") o.refine = std::stoi(value) != 0; else if (key == "--mass") @@ -221,6 +237,16 @@ Parse(int argc, char** argv) { o.factor = std::stof(value); else if (key == "--drop") o.drop = std::stof(value); + else if (key == "--scheme") + o.scheme = value; + else if (key == "--factor-min") + o.factor_min = std::stof(value); + else if (key == "--factor-max") + o.factor_max = std::stof(value); + else if (key == "--mass-min") + o.mass_min = std::stof(value); + else if (key == "--mass-max") + o.mass_max = std::stof(value); else if (key == "--sweep") o.sweep = value; else if (key == "--split") @@ -237,7 +263,10 @@ Parse(int argc, char** argv) { o.offset = n; else if (key == "--base-limit") o.base_limit = n; - else if (key == "--seed") { + else if (key == "--pareto-seed") { + Check(n <= UINT32_MAX, "Pareto seed out of range"); + o.pareto_seed = n; + } else if (key == "--seed") { Check(n <= UINT32_MAX, "Seed out of range"); o.seed = n; } else { @@ -246,6 +275,8 @@ Parse(int argc, char** argv) { o.threads = n; else if (key == "--repeats") o.repeats = n; + else if (key == "--samples") + o.samples = n; else if (key == "--k") o.k = n; else @@ -253,10 +284,111 @@ Parse(int argc, char** argv) { } } } + Check(o.quant_type == "fp16" || o.quant_type == "e5m7", "Unknown quantization"); + Check((!o.size_only && o.quant_type == "fp16") || o.method == "SINDI", "Packed/size mode requires SINDI"); + Check(!o.size_only || o.sweep.empty(), "Size measurement does not run a sweep"); Check(!o.data.empty() && o.queries > 0 && o.threads >= 2, "Provide --data-dir, positive queries, and >=2 threads"); + Check(o.sweep.empty() || o.sweep == "mass" || o.sweep == "count" || o.sweep == "drop" || o.sweep == "pareto", + "Unknown sweep: " + o.sweep); + Check(o.scheme == "mass", "Only the mass Pareto scheme is supported"); + for (const auto* key : + {"--scheme", "--samples", "--pareto-seed", "--factor-min", "--factor-max", "--mass-min", "--mass-max"}) + Check(o.sweep == "pareto" || !supplied.count(key), std::string(key) + " requires --sweep pareto"); + if (o.sweep == "pareto") { + Check(o.method == "SINDI" && o.refine && o.drop == 0, + "Pareto sweep requires --method SINDI --refine 1 --drop 0"); + Check(!supplied.count("--mass") && !supplied.count("--factor"), + "Use range options for a Pareto sweep, not --mass or --factor"); + Check(std::isfinite(o.factor_min) && std::isfinite(o.factor_max) && o.factor_min >= 1 && + o.factor_min <= o.factor_max, + "Require finite 1 <= factor-min <= factor-max"); + Check(std::isfinite(o.mass_min) && std::isfinite(o.mass_max) && o.mass_min > 0 && o.mass_min <= o.mass_max && + o.mass_max <= 1, + "Require finite 0 < mass-min <= mass-max <= 1"); + } return o; } +struct ParetoPoint { + size_t id; + float factor, mass, drop; + double qps, recall; +}; + +// Maximize both axes. Exact objective ties remain on the frontier. +void +WritePareto(const std::string& output, const std::vector& points, const std::string& scheme) { + std::vector order(points.size()); + std::iota(order.begin(), order.end(), 0); + std::sort(order.begin(), order.end(), [&](size_t a, size_t b) { + if (points[a].recall != points[b].recall) + return points[a].recall > points[b].recall; + if (points[a].qps != points[b].qps) + return points[a].qps > points[b].qps; + return points[a].id < points[b].id; + }); + std::vector frontier(points.size(), false); + double best_qps = -1, best_recall = -1; + for (auto i : order) { + const auto& p = points[i]; + frontier[i] = p.qps > best_qps || (p.qps == best_qps && p.recall == best_recall); + if (p.qps > best_qps) { + best_qps = p.qps; + best_recall = p.recall; + } + } + std::ofstream all(output + ".points.csv"), edge(output + ".frontier.csv"), svg(output + ".pareto.svg"); + for (auto* stream : {&all, &edge}) + *stream << "sample_id,scheme,refine_k,refine_factor,sindi_query_mass,drop_ratio_search,mean_qps,mean_recall," + "pareto\n" + << std::setprecision(17); + for (auto i : order) { + const auto& p = points[i]; + for (auto* stream : {&all, &edge}) { + if (stream == &edge && !frontier[i]) + continue; + *stream << p.id << ',' << scheme << ',' << (scheme == "mass" ? p.factor : 1) << ',' + << (scheme == "count" ? p.factor : 1) << ',' << p.mass << ',' << p.drop << ',' << p.qps << ',' + << p.recall << ',' << frontier[i] << '\n'; + } + } + const double max_qps = std::max(1.0, best_qps * 1.05); + auto x = [](double recall) { return 80 + 800 * recall; }; + auto y = [&](double qps) { return 530 - 460 * qps / max_qps; }; + svg << "" + "" + "" + "SINDI " + << scheme << " QPS / recall Pareto frontier" + << "Blue: all samples; red: nondominated samples (higher is better)"; + for (int tick = 0; tick <= 5; ++tick) { + const double recall = tick / 5.0, qps = max_qps * tick / 5; + svg << "" + << "" << recall << "" + << "" << std::fixed << std::setprecision(0) + << qps << std::defaultfloat << std::setprecision(6) << ""; + } + svg << "Recall@k" + "QPS" + ""; + for (auto i : order) { + const auto& p = points[i]; + svg << "sample=" << p.id + << (scheme == "count" ? " refine_factor=" : " refine_k=") << p.factor << " mass=" << p.mass + << " drop_ratio_search=" << p.drop << " QPS=" << p.qps << " recall=" << p.recall << ""; + } + svg << "\n"; + for (auto* stream : {&all, &edge, &svg}) { + stream->flush(); + Check(stream->good(), "Cannot write Pareto artifacts for " + output); + } +} + // ID recall deliberately exposes ties and FP16 rounding relative to FP32 ground truth. double Recall(const knowhere::DataSetPtr& result, const Truth& truth, size_t nq, int k, size_t nb) { @@ -314,17 +446,24 @@ main(int argc, char** argv) { const auto o = Parse(argc, argv); const auto& ip_kernels = knowhere::sparse::inverted::sindi::get_ip_kernels(); Dl_info accumulate_info{}, insert_info{}; - Check(dladdr(reinterpret_cast(ip_kernels.accumulate), &accumulate_info) != 0 && + Check(dladdr(o.quant_type == "e5m7" + ? reinterpret_cast(knowhere::sparse::inverted::sindi::get_packed_ip_kernel()) + : reinterpret_cast(ip_kernels.accumulate), + &accumulate_info) != 0 && accumulate_info.dli_sname != nullptr, "Cannot identify IP kernel"); Check(dladdr(reinterpret_cast(ip_kernels.batch_insert), &insert_info) != 0 && insert_info.dli_sname != nullptr, "Cannot identify selection kernel"); + int sve_bytes = 0; +#if defined(__aarch64__) Check(std::string(accumulate_info.dli_sname).find("sve") != std::string::npos && std::string(insert_info.dli_sname).find("sve") != std::string::npos, "Baseline requires SVE kernels"); const int sve_vl = prctl(PR_SVE_GET_VL); Check(sve_vl > 0, "Cannot read SVE vector length"); + sve_bytes = sve_vl & PR_SVE_VL_LEN_MASK; +#endif cpu_set_t affinity; CPU_ZERO(&affinity); Check(sched_getaffinity(0, sizeof(affinity), &affinity) == 0, "Cannot read CPU affinity"); @@ -332,8 +471,7 @@ main(int argc, char** argv) { for (int i = 0; i < CPU_SETSIZE; ++i) if (CPU_ISSET(i, &affinity)) cpus.push_back(i); - std::cerr << "IP kernel=" << accumulate_info.dli_sname << " SVE bytes=" << (sve_vl & PR_SVE_VL_LEN_MASK) - << std::endl; + std::cerr << "IP kernel=" << accumulate_info.dli_sname << " SVE bytes=" << sve_bytes << std::endl; const std::vector methods = {"TAAT_NAIVE", "DAAT_WAND", "DAAT_MAXSCORE", "BLOCK_MAX_MAXSCORE", "BLOCK_MAX_WAND", "SINDI", "SPARSE_WAND"}; Check(o.method == "all" || std::find(methods.begin(), methods.end(), o.method) != methods.end(), @@ -390,7 +528,7 @@ main(int argc, char** argv) { {"threads", o.threads}, {"repeats", o.repeats}, {"index_version", version}, - {"quant_type", "fp16"}, + {"quant_type", o.quant_type}, {"search_config", search}, {"base_load_seconds", load_seconds}, {"machine", host.machine}, @@ -401,9 +539,13 @@ main(int argc, char** argv) { {"ground_truth", o.base_limit ? "reduced-base FP32 brute force" : "base_full." + o.split + ".gt"}, {"warmup_batches", 1}, {"runs", Json::array()}}; + if (o.sweep == "pareto") + metadata["pareto"] = {{"scheme", o.scheme}, {"samples", o.samples}, {"seed", o.pareto_seed}, + {"factor_min", o.factor_min}, {"factor_max", o.factor_max}, {"mass_min", o.mass_min}, + {"mass_max", o.mass_max}}; metadata["effective_ip_kernel"] = accumulate_info.dli_sname; metadata["effective_selection_kernel"] = insert_info.dli_sname; - metadata["sve_vector_bytes"] = sve_vl & PR_SVE_VL_LEN_MASK; + metadata["sve_vector_bytes"] = sve_bytes; metadata["cpu_affinity"] = cpus; std::ifstream cpu("/proc/cpuinfo"), mem("/proc/meminfo"); metadata["cpuinfo"] = std::string(std::istreambuf_iterator(cpu), {}); @@ -421,7 +563,7 @@ main(int argc, char** argv) { alias ? knowhere::IndexEnum::INDEX_SPARSE_WAND : knowhere::IndexEnum::INDEX_SPARSE_INVERTED_INDEX; const std::string algo = alias ? "DAAT_WAND" : method; Json build = { - {"dim", base.dim}, {"metric_type", "IP"}, {"inverted_index_algo", algo}, {"quant_type", "fp16"}}; + {"dim", base.dim}, {"metric_type", "IP"}, {"inverted_index_algo", algo}, {"quant_type", o.quant_type}}; const std::string codec = algo == "SINDI" ? "fixed_docid_windows" : "block_streamvbyte"; if (algo == "SINDI") { build["sindi_window_size"] = 4096; @@ -438,6 +580,92 @@ main(int argc, char** argv) { Check(status == knowhere::Status::success, method + " Build failed, status=" + std::to_string(int(status))); const double build_seconds = Seconds(start); const auto bytes = index.Size(); + if (o.size_only) { + auto smoke = index.Search(query_ds, search, nullptr); + Check(smoke.has_value(), "Pre-serialize smoke search failed: " + smoke.what()); + Json measured = {{"build_config", build}, + {"build_seconds", build_seconds}, + {"index_bytes", bytes}, + {"refine", o.refine}, + {"quant_type", o.quant_type}}; + const std::string path = o.output + ".index"; + { + knowhere::BinarySet binary; + start = Clock::now(); + Check(index.Serialize(binary) == knowhere::Status::success, "Serialize failed"); + measured["serialize_seconds"] = Seconds(start); + auto blob = binary.GetByName(type); + Check(blob != nullptr, "Missing serialized index"); + measured["file_bytes"] = blob->size; + size_t total_bytes = 0; + for (const auto& item : binary.binary_map_) total_bytes += item.second->size; + measured["binaryset_bytes"] = total_bytes; + auto u32 = [&](size_t pos) { + uint32_t v; + std::memcpy(&v, blob->data.get() + pos, 4); + return v; + }; + const auto dims = u32(12), section_count = u32(32); + measured["dimensions"] = dims; + measured["sections"] = Json::array(); + for (size_t i = 0; i < section_count; ++i) { + knowhere::sparse::inverted::InvertedIndexSectionHeader h; + std::memcpy(&h, blob->data.get() + 36 + i * sizeof(h), sizeof(h)); + measured["sections"].push_back({{"type", uint32_t(h.type)}, {"bytes", h.size}}); + if (h.type != knowhere::sparse::inverted::InvertedIndexSectionType::POSTING_LISTS) + continue; + size_t pos = h.offset + (o.quant_type == "e5m7" ? 24 : 12) + (dims + 7) / 8; + for (size_t d = 0; d < dims; ++d) pos += 4 + u32(pos); + const uint64_t count = u32(pos + dims * 4); + const uint64_t payload = o.quant_type == "e5m7" + ? 2 * knowhere::sparse::inverted::sindi::packed12_bytes(count) + : 4 * count; + measured["retained_postings"] = count; + measured["posting_payload_bytes"] = payload; + measured["reported_nonpayload_bytes"] = uint64_t(bytes) - payload; + measured["serialized_nonpayload_bytes"] = blob->size - payload; + measured["payload_bytes_per_posting"] = count ? double(payload) / count : 0; + if (o.quant_type == "e5m7") { + const auto stream_bytes = knowhere::sparse::inverted::sindi::packed12_bytes(count); + const auto* vals = blob->data.get() + pos + (dims + 1) * 4 + stream_bytes; + uint64_t zeros = 0; + for (uint64_t j = 0; j < count; ++j) + zeros += knowhere::sparse::inverted::sindi::unpack12(vals, j) == 0; + measured["represented_zero_postings"] = zeros; + } + } + std::ofstream file(path, std::ios::binary); + file.write(reinterpret_cast(blob->data.get()), blob->size); + file.close(); + Check(file.good(), "Cannot write index file"); + } + // Release the built index and serialization buffer before mapping the file. + auto replacement = knowhere::IndexFactory::Instance().Create(type, version); + Check(replacement.has_value(), "Create load target failed"); + index = std::move(replacement.value()); + start = Clock::now(); + Check(index.DeserializeFromFile(path, Json{{"metric_type", "IP"}, {"enable_mmap", true}}) == + knowhere::Status::success, + "Mmap load failed"); + measured["mmap_load_seconds"] = Seconds(start); + measured["loaded_index_bytes"] = index.Size(); + auto loaded = index.Search(query_ds, search, nullptr); + Check(loaded.has_value(), "Post-load smoke failed"); + const size_t n = o.queries * o.k; + Check(std::memcmp(smoke.value()->GetIds(), loaded.value()->GetIds(), n * sizeof(int64_t)) == 0 && + std::memcmp(smoke.value()->GetDistance(), loaded.value()->GetDistance(), n * sizeof(float)) == + 0, + "Persistence changed results"); + measured["smoke_recall"] = Recall(loaded.value(), truth, o.queries, o.k, base.rows.size()); + measured["persistence_result_identity"] = true; + metadata["runs"].push_back(measured); + std::ofstream meta(o.output + ".json"); + meta << metadata.dump(2) << '\n'; + meta.close(); + Check(meta.good(), "Cannot write size metadata"); + std::cout << measured.dump() << std::endl; + continue; + } std::vector settings; auto add_setting = [&](float mass, float factor, float drop) { auto cfg = search; @@ -446,7 +674,20 @@ main(int argc, char** argv) { cfg["drop_ratio_search"] = drop; settings.push_back(cfg); }; - if (o.sweep == "mass") { + if (o.sweep == "pareto") { + std::mt19937 rng(o.pareto_seed); + // Match the donor's reproducible engine-to-float mapping. + auto sample = [&](float low, float high) { + const double unit = double(rng()) / 4294967296.0; + return static_cast(double(low) + (double(high) - low) * unit); + }; + for (int i = 0; i < o.samples; ++i) { + const float factor = sample(o.factor_min, o.factor_max); + const float mass = sample(o.mass_min, o.mass_max); + add_setting(mass, factor, 0); + } + metadata["pareto"]["settings"] = settings; + } else if (o.sweep == "mass") { for (float mass : {1.0f, .9f, .8f, .7f, .6f, .5f}) for (float factor : {1.0f, 5.0f, 10.0f}) add_setting(mass, factor, 0); } else if (o.sweep == "count") { @@ -456,6 +697,7 @@ main(int argc, char** argv) { for (float drop : {0.0f, .3f, .5f, .7f, .9f}) add_setting(1, 1, drop); } else add_setting(o.mass, o.factor, o.drop); + std::vector points; size_t setting_id = 0; for (const auto& setting : settings) { search = setting; @@ -468,6 +710,11 @@ main(int argc, char** argv) { {"build_config", build}, {"effective_codec", codec}, {"search_config", search}, {"build_seconds", build_seconds}, {"index_bytes", bytes}, {"measurements", Json::array()}}; + run["sample_id"] = setting_id - 1; + run["scheme"] = o.sweep == "pareto" ? "mass" : "existing"; + double qps_sum = 0, recall_sum = 0; + std::cout << "Setting " << setting_id << '/' << settings.size() << " refine_k=" << search["refine_k"] + << " mass=" << search["sindi_query_mass"] << std::endl; std::vector first_ids; std::vector first_scores; for (int repeat = 0; repeat < o.repeats; ++repeat) { @@ -499,15 +746,18 @@ main(int argc, char** argv) { "Results changed across measured repetitions"); } const auto recall = Recall(result.value(), truth, o.queries, o.k, base.rows.size()); - if (repeat == 0) { + if (repeat == 0 && o.sweep != "pareto") { run["recall_mismatches"] = RecallMismatches(result.value(), truth, selected, o.k); } + Check(seconds > 0 && std::isfinite(seconds), "Invalid search duration"); const double qps = o.queries / seconds; - csv << method << ',' << type << ',' << version << ",fp16," << codec << ',' << base.rows.size() - << ',' << o.queries << ',' << o.k << ',' << o.threads << ',' << build_seconds << ',' << bytes - << ',' << repeat + 1 << ',' << seconds << ',' << qps << ',' << recall << ',' - << search["sindi_query_mass"] << ',' << search["refine_k"] << ',' << search["drop_ratio_search"] - << ',' << o.refine << '\n'; + qps_sum += qps; + recall_sum += recall; + csv << method << ',' << type << ',' << version << ',' << o.quant_type << ',' << codec << ',' + << base.rows.size() << ',' << o.queries << ',' << o.k << ',' << o.threads << ',' + << build_seconds << ',' << bytes << ',' << repeat + 1 << ',' << seconds << ',' << qps << ',' + << recall << ',' << search["sindi_query_mass"] << ',' << search["refine_k"] << ',' + << search["drop_ratio_search"] << ',' << o.refine << '\n'; csv.flush(); Check(csv.good(), "Failed writing results"); run["measurements"].push_back({{"repeat", repeat + 1}, @@ -518,54 +768,65 @@ main(int argc, char** argv) { std::cout << method << " repeat=" << repeat + 1 << " seconds=" << seconds << " QPS=" << qps << " recall@" << o.k << '=' << recall << std::endl; } - // Separate, untimed coverage diagnostics. Materialize exactly the coarse query. - SparseData selected_data; - selected_data.dim = queries.dim; - double mass_sum = 0; - size_t retained_nnz = 0; - for (const auto& q : queries.rows) { - Row selected; - if (o.refine) - selected = knowhere::sparse::inverted::sindi::retain_query_mass( - q, search["sindi_query_mass"].get()); - else { - std::vector weights; - for (size_t j = 0; j < q.size(); ++j) weights.push_back(q[j].val); - std::sort(weights.begin(), weights.end()); - const size_t count = size_t(search["drop_ratio_search"].get() * weights.size()); - float threshold = weights.empty() ? 0 : weights[std::min(count, weights.size() - 1)]; - std::vector> terms; - for (size_t j = 0; j < q.size(); ++j) - if (q[j].val >= threshold) - terms.emplace_back(q[j].id, q[j].val); - selected = Row(terms); + run["mean_qps"] = qps_sum / o.repeats; + run["mean_recall"] = recall_sum / o.repeats; + if (o.sweep == "pareto") { + const float factor = search["refine_k"].get(); + run["candidate_pool"] = + knowhere::sparse::inverted::sindi::refinement_pool_size(o.k, factor, base.rows.size()); + points.push_back({setting_id - 1, factor, search["sindi_query_mass"].get(), 0, + qps_sum / o.repeats, recall_sum / o.repeats}); + WritePareto(o.output, points, "mass"); + } else { + // Separate, untimed coverage diagnostics. Materialize exactly the coarse query. + SparseData selected_data; + selected_data.dim = queries.dim; + double mass_sum = 0; + size_t retained_nnz = 0; + for (const auto& q : queries.rows) { + Row selected; + if (o.refine) + selected = knowhere::sparse::inverted::sindi::retain_query_mass( + q, search["sindi_query_mass"].get()); + else { + std::vector weights; + for (size_t j = 0; j < q.size(); ++j) weights.push_back(q[j].val); + std::sort(weights.begin(), weights.end()); + const size_t count = size_t(search["drop_ratio_search"].get() * weights.size()); + float threshold = weights.empty() ? 0 : weights[std::min(count, weights.size() - 1)]; + std::vector> terms; + for (size_t j = 0; j < q.size(); ++j) + if (q[j].val >= threshold) + terms.emplace_back(q[j].id, q[j].val); + selected = Row(terms); + } + double total = 0, retained = 0; + for (size_t j = 0; j < q.size(); ++j) total += q[j].val; + for (size_t j = 0; j < selected.size(); ++j) retained += selected[j].val; + mass_sum += total ? retained / total : 1; + retained_nnz += selected.size(); + selected_data.rows.push_back(std::move(selected)); } - double total = 0, retained = 0; - for (size_t j = 0; j < q.size(); ++j) total += q[j].val; - for (size_t j = 0; j < selected.size(); ++j) retained += selected[j].val; - mass_sum += total ? retained / total : 1; - retained_nnz += selected.size(); - selected_data.rows.push_back(std::move(selected)); + const size_t pool = knowhere::sparse::inverted::sindi::refinement_pool_size( + o.k, search["refine_k"].get(), base.rows.size()); + auto diagnostic = search; + diagnostic["k"] = pool; + diagnostic["sindi_query_mass"] = 1; + diagnostic["refine_k"] = 1; + diagnostic["drop_ratio_search"] = 0; + auto coarse = index.Search(selected_data.Dataset(), diagnostic, nullptr); + Check(coarse.has_value(), "Coarse coverage diagnostic failed"); + size_t covered = 0; + for (size_t q = 0; q < o.queries; ++q) + for (int j = 0; j < o.k; ++j) + covered += std::find(coarse.value()->GetIds() + q * pool, + coarse.value()->GetIds() + (q + 1) * pool, + truth.ids[q * o.k + j]) != coarse.value()->GetIds() + (q + 1) * pool; + run["candidate_pool"] = pool; + run["candidate_coverage"] = double(covered) / (o.queries * o.k); + run["retained_query_nnz_mean"] = double(retained_nnz) / o.queries; + run["retained_mass_mean"] = mass_sum / o.queries; } - const size_t pool = knowhere::sparse::inverted::sindi::refinement_pool_size( - o.k, search["refine_k"].get(), base.rows.size()); - auto diagnostic = search; - diagnostic["k"] = pool; - diagnostic["sindi_query_mass"] = 1; - diagnostic["refine_k"] = 1; - diagnostic["drop_ratio_search"] = 0; - auto coarse = index.Search(selected_data.Dataset(), diagnostic, nullptr); - Check(coarse.has_value(), "Coarse coverage diagnostic failed"); - size_t covered = 0; - for (size_t q = 0; q < o.queries; ++q) - for (int j = 0; j < o.k; ++j) - covered += - std::find(coarse.value()->GetIds() + q * pool, coarse.value()->GetIds() + (q + 1) * pool, - truth.ids[q * o.k + j]) != coarse.value()->GetIds() + (q + 1) * pool; - run["candidate_pool"] = pool; - run["candidate_coverage"] = double(covered) / (o.queries * o.k); - run["retained_query_nnz_mean"] = double(retained_nnz) / o.queries; - run["retained_mass_mean"] = mass_sum / o.queries; metadata["runs"].push_back(std::move(run)); std::ofstream meta(o.output + ".json"); meta << metadata.dump(2) << '\n'; diff --git a/benchmark/benchmark_sparse_bm25_compact.cpp b/benchmark/benchmark_sparse_bm25_compact.cpp new file mode 100644 index 000000000..4b901b740 --- /dev/null +++ b/benchmark/benchmark_sparse_bm25_compact.cpp @@ -0,0 +1,466 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy at http://www.apache.org/licenses/LICENSE-2.0 + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include "knowhere/comp/brute_force.h" +#include "knowhere/comp/knowhere_config.h" +#include "knowhere/index/index_factory.h" +#include "knowhere/sparse_utils.h" +#include "knowhere/version.h" +#include "src/index/sparse/inverted_index_format.h" +#include "src/index/sparse/sindi_bm25_u4.h" +#include "src/index/sparse/sindi_refinement.h" +#include "src/index/sparse/sindi_simd.h" + +namespace { +using Clock = std::chrono::steady_clock; +using Row = knowhere::sparse::SparseRow; +using knowhere::Json; + +void +Check(bool ok, std::string_view message) { + if (!ok) { + throw std::runtime_error(std::string(message)); + } +} + +double +Seconds(Clock::time_point start) { + return std::chrono::duration(Clock::now() - start).count(); +} + +// The challenge files are little-endian, structure-of-arrays CSR / k-NN files. +class MappedFile { + public: + explicit MappedFile(const std::string& path) { + Check(std::endian::native == std::endian::little, "Only little-endian hosts are supported"); + int fd = open(path.c_str(), O_RDONLY); + Check(fd >= 0, "Cannot open " + path); + struct stat st {}; + if (fstat(fd, &st) != 0 || st.st_size <= 0) { + close(fd); + throw std::runtime_error("Cannot stat or empty file: " + path); + } + size = st.st_size; + void* ptr = mmap(nullptr, size, PROT_READ, MAP_PRIVATE, fd, 0); + close(fd); + Check(ptr != MAP_FAILED, "Cannot mmap " + path); + data = static_cast(ptr); + } + ~MappedFile() { + munmap(const_cast(data), size); + } + MappedFile(const MappedFile&) = delete; + MappedFile& + operator=(const MappedFile&) = delete; + + template + T + Read(uint64_t offset) const { + Check(offset <= size && sizeof(T) <= size - offset, "Truncated binary file"); + T value; + std::memcpy(&value, data + offset, sizeof(T)); + return value; + } + const char* data = nullptr; + uint64_t size = 0; +}; + +struct SparseData { + uint64_t file_rows = 0, dim = 0, nnz = 0; + std::vector rows; + + knowhere::DataSetPtr + Dataset() const { + auto ds = knowhere::GenDataSet(rows.size(), dim, rows.data()); + ds->SetIsSparse(true); + return ds; + } +}; + +SparseData +LoadCsr(const std::string& path, const std::vector* selection, uint64_t limit = 0) { + MappedFile f(path); + SparseData out; + out.file_rows = f.Read(0); + out.dim = f.Read(8); + const auto nnz = f.Read(16); + Check(f.size >= 32 && out.file_rows > 0 && out.file_rows <= (f.size - 32) / 8 && out.dim > 0 && + out.dim <= std::numeric_limits::max(), + "Invalid CSR dimensions: " + path); + const uint64_t indices_offset = 24 + 8 * (out.file_rows + 1); + Check(nnz <= (f.size - indices_offset) / 8 && indices_offset + 8 * nnz == f.size, + "CSR file size does not match header: " + path); + const uint64_t values_offset = indices_offset + 4 * nnz; + Check(f.Read(24) == 0 && f.Read(24 + 8 * out.file_rows) == nnz, "Invalid CSR endpoint offsets"); + uint64_t previous = 0; + for (uint64_t i = 1; i <= out.file_rows; ++i) { + const auto next = f.Read(24 + 8 * i); + Check(next >= previous && next <= nnz, "Invalid CSR row offsets"); + previous = next; + } + const auto count = selection ? selection->size() : (limit ? std::min(limit, out.file_rows) : out.file_rows); + out.rows.reserve(count); + for (uint64_t i = 0; i < count; ++i) { + const auto id = selection ? selection->at(i) : i; + Check(id < out.file_rows, "Query ID outside CSR file"); + const auto start = f.Read(24 + 8 * id); + const auto end = f.Read(24 + 8 * (id + 1)); + out.rows.emplace_back(end - start); + uint32_t last = 0; + for (auto j = start; j < end; ++j) { + const auto col = f.Read(indices_offset + 4 * j); + const auto value = f.Read(values_offset + 4 * j); + Check(col < out.dim && (j == start || col > last), "CSR columns must be sorted and unique"); + Check(std::isfinite(value) && value >= 0, "Expected finite nonnegative SPLADE weights"); + out.rows.back().set_at(j - start, col, value); + last = col; + } + out.nnz += end - start; + } + return out; +} + +Json +ReadJson(const std::string& path) { + std::ifstream f(path); + Check(f.good(), "Cannot read " + path); + Json j; + f >> j; + return j; +} + +double +Score(const Row& query, const Row& document, double dl, double k1, double b, double avgdl, + const knowhere::sparse::inverted::sindi::Bm25U4Lut* lut) { + double score = 0; + size_t i = 0, j = 0; + while (i < query.size() && j < document.size()) { + if (query[i].id < document[j].id) + ++i; + else if (query[i].id > document[j].id) + ++j; + else { + double tf = document[j].val; + if (lut) + tf = lut->decode[lut->encode[std::min(tf, 255)]]; + score += query[i].val * (k1 + 1) * tf / (tf + k1 * (1 - b + b * dl / avgdl)); + ++i; + ++j; + } + } + return score; +} +} // namespace + +int +main(int argc, char** argv) { + try { + std::map args; + for (int i = 1; i + 1 < argc; i += 2) args[argv[i]] = argv[i + 1]; + auto get = [&](std::string name, std::string fallback) { return args.count(name) ? args.at(name) : fallback; }; + const std::string data = get("--data-dir", "/volume/data_sparse/nq"), split = get("--split", "test"), + quant = get("--quant-type", "u8"), output = get("--output", "bm25_compact.json"); + const bool u4 = quant == "u4_lut" || quant == "u4_lut_u12" || quant == "u4_lut_u16"; + const bool u16_ids = quant == "u4_lut_u16"; + const size_t k = std::stoul(get("--k", "10")), limit = std::stoul(get("--queries", "0")); + const int threads = std::stoi(get("--threads", "8")), requested = std::stoi(get("--repeats", "3")); + Check(k > 0 && (k == 10 || k == 100) && threads > 0 && requested >= 3, "Invalid benchmark arguments"); + const auto simd = get("--simd", "auto"); + auto simd_type = knowhere::KnowhereConfig::SimdType::AUTO; + if (simd == "avx2") + simd_type = knowhere::KnowhereConfig::SimdType::AVX2; + else if (simd == "scalar") + simd_type = knowhere::KnowhereConfig::SimdType::GENERIC; + else + Check(simd == "auto", "Unsupported --simd mode (auto/avx2/scalar)"); + const auto selected_simd = knowhere::KnowhereConfig::SetSimdType(simd_type); + knowhere::KnowhereConfig::SetBuildThreadPoolSize(threads); + knowhere::KnowhereConfig::SetSearchThreadPoolSize(threads); + const auto manifest = ReadJson(data + "/manifest.json"); + auto base = LoadCsr(data + "/sparse/base_tf.csr", nullptr); + auto queries = LoadCsr(data + "/sparse/queries." + split + ".idf.csr", nullptr, limit); + const float k1 = manifest.at("bm25_k1"), b = manifest.at("bm25_b"), avgdl = manifest.at("avgdl"); + const std::string truth_path = + get("--truth", data + "/ground_truth/" + split + ".k" + std::to_string(k) + ".bm25"); + auto signature = ReadJson(truth_path + ".json"); + Check(signature.at("base_rows") == base.rows.size(), "Truth corpus mismatch"); + Check(signature.at("query_ids").size() == queries.file_rows, "Truth query count mismatch"); + for (size_t q = 0; q < queries.file_rows; ++q) + Check(signature["query_ids"][q] == q, "Truth query order mismatch"); + for (auto key : {"bm25_k1", "bm25_b", "bm25_avgdl"}) { + const float value = signature.at("search").at(key); + const float expected = std::string(key) == "bm25_k1" ? k1 : (std::string(key) == "bm25_b" ? b : avgdl); + Check(std::abs(value - expected) < 1e-5 * std::max(1.f, std::abs(expected)), "Truth scorer mismatch"); + } + Check(signature.at("search").at("k") == k, "Truth k mismatch"); + if (signature.contains("query_csr_sha256")) { + Check(signature["query_csr_sha256"] == manifest.at("sha256").at("queries." + split + ".idf.csr"), + "Truth query hash mismatch"); + } + Check(signature.at("sha256").at("base_tf.csr") == manifest.at("sha256").at("base_tf.csr"), + "Truth base hash mismatch"); + MappedFile truth(truth_path); + const auto truth_n = truth.Read(0); + Check(truth_n == queries.file_rows * k && truth.size == 8 + 12 * truth_n, "Truth payload mismatch"); + size_t valid_truth = 0; + for (size_t q = 0; q < queries.rows.size(); ++q) { + std::set seen; + bool padding = false; + for (size_t j = 0; j < k; ++j) { + const auto id = truth.Read(8 + 8 * (q * k + j)); + if (id < 0) { + padding = true; + continue; + } + Check(!padding && uint64_t(id) < base.rows.size() && seen.insert(id).second, + "Invalid truth ordering/ID"); + ++valid_truth; + } + } + Check(valid_truth > 0, "No valid ground-truth neighbors"); + const auto type = knowhere::IndexEnum::INDEX_SPARSE_INVERTED_INDEX; + auto index = knowhere::IndexFactory::Instance().Create(type, 11).value(); + Json build{{"metric_type", "BM25"}, {"inverted_index_algo", "SINDI"}, + {"quant_type", quant}, {"sindi_window_size", 4096}, + {"refine", false}, {"bm25_k1", k1}, + {"bm25_b", b}, {"bm25_avgdl", avgdl}}; + auto t = Clock::now(); + Check(index.Build(base.Dataset(), build) == knowhere::Status::success, "Build failed"); + const auto build_seconds = Seconds(t); + const auto index_bytes = index.Size(); + knowhere::BinarySet bytes; + Check(index.Serialize(bytes) == knowhere::Status::success, "Serialize failed"); + auto blob = bytes.GetByName(type); + knowhere::sparse::inverted::sindi::Bm25U4Lut lut; + Json meta{{"dataset", data}, + {"split", split}, + {"quant_type", quant}, + {"threads", threads}, + {"k", k}, + {"queries", queries.rows.size()}, + {"base_rows", base.rows.size()}, + {"postings", base.nnz}, + {"build_seconds", build_seconds}, + {"index_bytes", index_bytes}, + {"serialized_bytes", blob->size}, + {"compiler", __VERSION__}, + {"requested_simd", simd}, + {"selected_simd", selected_simd}, + {"build", build}, + {"runs", Json::array()}}; + using namespace knowhere::sparse::inverted::sindi; + void* kernel = u4 ? reinterpret_cast(get_packed_bm25_kernel(u16_ids)) + : quant == "u8" ? reinterpret_cast(get_bm25_u8_kernels().accumulate) + : reinterpret_cast(get_bm25_kernels().accumulate); + Dl_info dispatch{}; + if (dladdr(kernel, &dispatch) && dispatch.dli_sname) + meta["kernel_symbol"] = dispatch.dli_sname; + if (u4) { + using namespace knowhere::sparse::inverted; + uint32_t sections; + std::memcpy(§ions, blob->data.get() + 32, 4); + for (size_t i = 0; i < sections; ++i) { + InvertedIndexSectionHeader h; + std::memcpy(&h, blob->data.get() + 36 + i * sizeof(h), sizeof(h)); + if (h.type == InvertedIndexSectionType::POSTING_LISTS) { + std::memcpy(lut.decode.data(), blob->data.get() + h.offset + 24, 16); + std::memcpy(lut.ends.data(), blob->data.get() + h.offset + 40, 16); + } + } + lut.k1 = k1; + lut.validate_and_encode(); + meta["lut"] = { + {"decode", lut.decode}, {"ends", lut.ends}, {"encode", lut.encode}, {"lut_id", lut.fingerprint()}}; + meta["posting_payload_bytes"] = + u16_ids ? 2 * base.nnz + base.nnz / 2 + base.nnz % 2 : 2 * base.nnz + base.nnz % 2; + } + std::vector lengths(base.rows.size()); + for (size_t d = 0; d < base.rows.size(); ++d) + for (size_t j = 0; j < base.rows[d].size(); ++j) lengths[d] += base.rows[d][j].val; + Json search{{"metric_type", "BM25"}, + {"k", k}, + {"bm25_k1", k1}, + {"bm25_b", b}, + {"bm25_avgdl", avgdl}, + {"drop_ratio_search", 0}, + {"dim_max_score_ratio", 1.05}}; + auto warm = index.Search(queries.Dataset(), search, nullptr); + Check(warm.has_value(), "Warmup failed"); + std::vector reference_ids(warm.value()->GetIds(), warm.value()->GetIds() + queries.rows.size() * k); + std::vector reference_scores(warm.value()->GetDistance(), + warm.value()->GetDistance() + queries.rows.size() * k); + auto validate = [&](const auto& result) { + Check(result.has_value(), "Search failed"); + size_t correct = 0; + for (size_t q = 0; q < queries.rows.size(); ++q) { + std::set seen; + for (size_t j = 0; j < k; ++j) { + const auto id = result.value()->GetIds()[q * k + j]; + Check(id == reference_ids[q * k + j], "Repeated IDs differ"); + if (id < 0) + continue; + Check(uint64_t(id) < base.rows.size() && seen.insert(id).second, "Invalid/duplicate result ID"); + const auto score = result.value()->GetDistance()[q * k + j]; + Check(std::abs(score - reference_scores[q * k + j]) < 1e-5f, "Repeated scores differ"); + const auto expected = + Score(queries.rows[q], base.rows[id], lengths[id], k1, b, avgdl, u4 ? &lut : nullptr); + Check(std::abs(score - expected) < 3e-5 * std::max(1., std::abs(expected)), + "Returned score disagrees with oracle"); + for (size_t g = 0; g < k; ++g) + if (id == truth.Read(8 + 8 * (q * k + g))) { + ++correct; + break; + } + } + } + return double(correct) / valid_truth; + }; + meta["valid_truth_positions"] = valid_truth; + meta["recall_denominator"] = "nonnegative ground-truth IDs; exclude sentinel padding"; + size_t outside_truth = 0, near_tie_hits = 0, below_original_threshold = 0; + for (size_t q = 0; q < queries.rows.size(); ++q) { + double threshold = 0; + std::set truth_ids; + for (size_t j = 0; j < k; ++j) { + const auto id = truth.Read(8 + 8 * (q * k + j)); + if (id >= 0) { + truth_ids.insert(id); + threshold = truth.Read(8 + 8 * truth_n + 4 * (q * k + j)); + } + } + for (size_t j = 0; j < k; ++j) { + const auto id = reference_ids[q * k + j]; + if (id < 0 || truth_ids.count(id)) + continue; + ++outside_truth; + const auto score = Score(queries.rows[q], base.rows[id], lengths[id], k1, b, avgdl, nullptr); + if (score + 3e-5 * std::max(1., std::abs(threshold)) >= threshold) + ++near_tie_hits; + else + ++below_original_threshold; + } + } + meta["original_tf_diagnostics"] = {{"returned_hits_outside_exact_id_truth", outside_truth}, + {"within_or_above_threshold_tolerance", near_tie_hits}, + {"below_threshold_tolerance", below_original_threshold}, + {"relative_tolerance", 3e-5}}; + const auto recall = validate(warm); + double timed = 0; + for (int repetition = 0; repetition < requested || timed < 5; ++repetition) { + t = Clock::now(); + auto result = index.Search(queries.Dataset(), search, nullptr); + const double seconds = Seconds(t); + timed += seconds; + const auto repeated_recall = validate(result); + Check(repeated_recall == recall, "Repeated recall differs"); + meta["runs"].push_back({{"seconds", seconds}, {"qps", queries.rows.size() / seconds}, {"recall", recall}}); + } + auto load = build; + load.erase("quant_type"); + load.erase("sindi_window_size"); + auto restored = knowhere::IndexFactory::Instance().Create(type, 11).value(); + t = Clock::now(); + Check(restored.Deserialize(bytes, load) == knowhere::Status::success, "Reload failed"); + meta["load_seconds"] = Seconds(t); + validate(restored.Search(queries.Dataset(), search, nullptr)); + const std::string save = get("--index-save", ""); + if (!save.empty()) { + Check(!std::filesystem::exists(save), "Refuse to overwrite index"); + std::ofstream f(save, std::ios::binary); + f.write(reinterpret_cast(blob->data.get()), blob->size); + Check(f.good(), "Index save failed"); + f.close(); + auto mapped = knowhere::IndexFactory::Instance().Create(type, 11).value(); + t = Clock::now(); + Check(mapped.DeserializeFromFile(save, load) == knowhere::Status::success, "File/mmap reload failed"); + meta["mmap_load_seconds"] = Seconds(t); + validate(mapped.Search(queries.Dataset(), search, nullptr)); + } + // Independent exhaustive scans of deterministic samples verify the + // decoded score domain and quantify original-TF ties/losses separately. + meta["sample_oracles"] = Json::array(); + for (size_t q = 0; q < std::min(8, queries.rows.size()); ++q) { + std::vector> top; + for (size_t d = 0; d < base.rows.size(); ++d) { + const auto score = Score(queries.rows[q], base.rows[d], lengths[d], k1, b, avgdl, u4 ? &lut : nullptr); + if (score > 0) { + top.emplace_back(score, d); + std::push_heap(top.begin(), top.end(), std::greater<>()); + if (top.size() > k) { + std::pop_heap(top.begin(), top.end(), std::greater<>()); + top.pop_back(); + } + } + } + const double threshold = top.empty() ? 0 : top.front().first; + size_t missing = 0; + for (size_t j = 0; j < k; ++j) { + auto id = reference_ids[q * k + j]; + if (id < 0) + continue; + const auto score = + Score(queries.rows[q], base.rows[id], lengths[id], k1, b, avgdl, u4 ? &lut : nullptr); + missing += score + 3e-5 * std::max(1., std::abs(threshold)) < threshold; + } + Check(missing == 0, "Full-corpus decoded-score oracle ranking mismatch"); + meta["sample_oracles"].push_back({{"query", q}, {"threshold", threshold}, {"below_threshold", missing}}); + } + struct rusage resources {}; + getrusage(RUSAGE_SELF, &resources); + meta["peak_rss_kib"] = resources.ru_maxrss; + meta["recall"] = recall; + meta["mean_qps"] = queries.rows.size() * meta["runs"].size() / timed; + cpu_set_t affinity; + CPU_ZERO(&affinity); + sched_getaffinity(0, sizeof(affinity), &affinity); + std::vector cores; + for (int i = 0; i < CPU_SETSIZE; ++i) + if (CPU_ISSET(i, &affinity)) + cores.push_back(i); + meta["affinity"] = cores; +#ifdef __aarch64__ + meta["sve_vl_bytes"] = prctl(PR_SVE_GET_VL) & PR_SVE_VL_LEN_MASK; +#endif + std::ofstream f(output); + f << meta.dump(2) << '\n'; + Check(f.good(), "Output write failed"); + std::cout << quant << " QPS=" << meta["mean_qps"] << " recall=" << recall << " bytes=" << index_bytes << "\n"; + return 0; + } catch (const std::exception& e) { + std::cerr << e.what() << '\n'; + return 1; + } +} diff --git a/src/index/sparse/inverted_index.h b/src/index/sparse/inverted_index.h index dfd78d94b..cef13179c 100644 --- a/src/index/sparse/inverted_index.h +++ b/src/index/sparse/inverted_index.h @@ -51,6 +51,8 @@ enum class InvertedIndexEncoding : uint32_t { FIXED_DOCID_WINDOWS = 3, BLOCK_ADAPTIVE = 4, FIXED_DOCID_WINDOWS_U12_E5M7 = 5, + FIXED_DOCID_WINDOWS_U12_U4_LUT = 6, + FIXED_DOCID_WINDOWS_U16_U4_LUT = 7, }; enum class InvertedIndexPrometheusBuildStats : uint32_t { DATASET_NNZ_STATS = 0, POSTING_LIST_LENGTH_STATS = 1 }; diff --git a/src/index/sparse/inverted_index_format.h b/src/index/sparse/inverted_index_format.h index ba8a289bc..390102eea 100644 --- a/src/index/sparse/inverted_index_format.h +++ b/src/index/sparse/inverted_index_format.h @@ -35,6 +35,8 @@ enum class InvertedIndexQuantType : uint32_t { BM25_U16 = 4, BM25_U32 = 5, IP_E5M7 = 6, + BM25_U4_LUT_U12 = 7, // Renamed symbol; preserve existing U12 on-disk identity. + BM25_U4_LUT_U16 = 8, }; static_assert(sizeof(InvertedIndexQuantType) == sizeof(uint32_t)); diff --git a/src/index/sparse/sindi_bm25_u4.h b/src/index/sparse/sindi_bm25_u4.h new file mode 100644 index 000000000..9421f03d2 --- /dev/null +++ b/src/index/sparse/sindi_bm25_u4.h @@ -0,0 +1,125 @@ +// Copyright (C) 2026 Zilliz. All rights reserved. +// Licensed under the Apache License, Version 2.0. +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace knowhere::sparse::inverted::sindi { + +inline uint8_t +unpack_bm25_u4(const uint8_t* values, size_t posting) noexcept { + return (values[posting / 2] >> (4 * (posting & 1))) & 15; +} + +// Compact sections need only byte alignment, including U16 ID streams. +inline uint16_t +unpack_bm25_u16_id(const uint8_t* ids, size_t posting) noexcept { + uint16_t id; + std::memcpy(&id, ids + posting * sizeof(id), sizeof(id)); + return id; +} + +// A code is an index, not a TF. Endpoints define all inputs, including unseen TFs. +struct Bm25U4Lut { + std::array decode{}; + std::array ends{}; + std::array encode{}; + float k1 = 1.2f; + + void + validate_and_encode() { + if (!std::isfinite(k1) || k1 < 0 || decode[0] != 0 || ends[0] != 0 || decode[1] != 1 || ends[1] != 1 || + ends[15] != 255) + throw std::invalid_argument("Invalid BM25 u4 LUT contract"); + for (size_t c = 1; c < 16; ++c) { + if (ends[c] <= ends[c - 1] || decode[c] <= ends[c - 1] || decode[c] > ends[c]) + throw std::invalid_argument("Invalid BM25 u4 LUT interval/representative"); + for (unsigned t = unsigned(ends[c - 1]) + 1; t <= ends[c]; ++t) encode[t] = c; + } + encode[0] = 0; + } + + // Stable descriptor identity, not a cryptographic integrity checksum. + uint64_t + fingerprint() const { + uint64_t h = 14695981039346656037ULL; + auto add = [&](uint8_t byte) { h = (h ^ byte) * 1099511628211ULL; }; + add(1); // codec/fitting version + for (auto t : decode) add(t); + for (auto t : ends) add(t); + uint32_t bits; + std::memcpy(&bits, &k1, sizeof(bits)); + for (unsigned i = 0; i < 4; ++i) add((bits >> (8 * i)) & 255); + return h; + } +}; + +// Globally optimal contiguous-bin fit for the declared reference-length BM25 +// squared-error objective. TF 1 is exact. Stable ties choose lower representatives +// and earlier split points. Empty bins have a deterministic lowest representative. +inline Bm25U4Lut +fit_bm25_u4_lut(const std::array& histogram, float k1) { + if (!std::isfinite(k1) || k1 < 0) + throw std::invalid_argument("Invalid LUT fitting k1"); + std::array g{}, w{}, a{}, z{}; + for (unsigned t = 1; t < 256; ++t) { + g[t] = (static_cast(k1) + 1) * t / (t + static_cast(k1)); + w[t] = w[t - 1] + histogram[t]; + a[t] = a[t - 1] + histogram[t] * g[t]; + z[t] = z[t - 1] + histogram[t] * g[t] * g[t]; + } + constexpr size_t stride = 256; + std::vector cost(stride * stride); + std::vector rep(stride * stride); + for (unsigned lo = 2; lo < 256; ++lo) + for (unsigned hi = lo; hi < 256; ++hi) { + const auto weight = w[hi] - w[lo - 1], sum = a[hi] - a[lo - 1], square = z[hi] - z[lo - 1]; + auto loss = [&](unsigned r) { return std::max(0.L, square - 2 * g[r] * sum + g[r] * g[r] * weight); }; + unsigned r = lo; + if (weight > 0) { + const auto mean = sum / weight; + const auto it = std::lower_bound(g.begin() + lo, g.begin() + hi + 1, mean); + r = std::min(hi, it - g.begin()); + if (r > lo && loss(r - 1) <= loss(r)) + --r; + } + cost[lo * stride + hi] = loss(r); + rep[lo * stride + hi] = r; + } + const auto inf = std::numeric_limits::infinity(); + std::array, 15> dp{}; + std::array, 15> split{}; + for (auto& row : dp) row.fill(inf); + dp[0][1] = 0; + for (unsigned bins = 1; bins <= 14; ++bins) + for (unsigned hi = bins + 1; hi < 256; ++hi) + for (unsigned prev = bins; prev < hi; ++prev) { + const auto candidate = dp[bins - 1][prev] + cost[(prev + 1) * stride + hi]; + if (candidate < dp[bins][hi]) { + dp[bins][hi] = candidate; + split[bins][hi] = prev; + } + } + Bm25U4Lut lut; + lut.k1 = k1; + lut.ends[1] = lut.decode[1] = 1; + unsigned hi = 255; + for (unsigned bins = 14; bins > 0; --bins) { + const unsigned prev = split[bins][hi]; + lut.ends[bins + 1] = hi; + lut.decode[bins + 1] = rep[(prev + 1) * stride + hi]; + hi = prev; + } + lut.validate_and_encode(); + return lut; +} + +} // namespace knowhere::sparse::inverted::sindi diff --git a/src/index/sparse/sindi_inverted_index.h b/src/index/sparse/sindi_inverted_index.h index 9ab3f517a..5789b904c 100644 --- a/src/index/sparse/sindi_inverted_index.h +++ b/src/index/sparse/sindi_inverted_index.h @@ -28,6 +28,7 @@ #include "knowhere/bitsetview.h" #include "knowhere/operands.h" #include "simd/hook.h" +#include "sindi_bm25_u4.h" #include "sindi_packed12.h" #include "sindi_refinement.h" @@ -59,10 +60,17 @@ class SindiInvertedIndex : public DimMapInvertedIndex (packed_u16_ids_ ? max_window_size : 4096))))) + throw std::invalid_argument("Invalid packed SINDI representation/window/refinement"); } SindiInvertedIndex(const SindiInvertedIndex& rhs) = delete; @@ -128,8 +136,9 @@ class SindiInvertedIndex : public DimMapInvertedIndex(knowhere::fp16(data[i][j].val))) || - (packed_ && std::signbit(data[i][j].val))) { + const float value = data[i][j].val; + if ((is_ip && !std::isfinite(static_cast(knowhere::fp16(value)))) || + (packed_ && std::signbit(value)) || + (is_bm25 && (value > 65535 || std::floor(value) != value))) { return Status::invalid_args; } } @@ -945,7 +957,7 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_rows_, sizeof(uint32_t)); writer.write(&this->max_dim_, sizeof(uint32_t)); writer.write(&this->nr_inner_dims_, sizeof(uint32_t)); - const auto quant_type = packed_ ? InvertedIndexQuantType::IP_E5M7 : posting_quant_type(); + const auto quant_type = packed_ ? packed_quant_type() : posting_quant_type(); writer.write(&quant_type, sizeof(quant_type)); const std::array reserved{}; writer.write(reserved.data(), reserved.size()); @@ -977,7 +989,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex uint64_t { - size_t res = sizeof(uint32_t) * (packed_ ? 6 : 3); + size_t res = posting_header_bytes(); const size_t mask_sz = (nr_dims + 7) / 8; res += mask_sz * sizeof(uint8_t); @@ -995,7 +1007,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex(packed_ ? InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7 - : InvertedIndexEncoding::FIXED_DOCID_WINDOWS); + static_cast(packed_ ? packed_encoding() : InvertedIndexEncoding::FIXED_DOCID_WINDOWS); write_padding_until(writer, section_headers[0].offset); writer.write(&index_encoding_type, sizeof(uint32_t)); writer.write(&this->window_size_, sizeof(uint32_t)); writer.write(&this->nr_windows_, sizeof(uint32_t)); if (packed_) { - // Version 1, independent U12 layout tag 2 and E5M7 value-codec tag 6. - const uint32_t descriptor[] = {1, 2, static_cast(InvertedIndexQuantType::IP_E5M7)}; + // Version 1: explicit ID layout and independent quantization identity. + const uint32_t descriptor[] = {1, packed_id_layout(), static_cast(packed_quant_type())}; writer.write(descriptor, sizeof(descriptor)); + if constexpr (is_bm25) { + const auto& cfg = this->build_scorer_->config().scorer_params.bm25; + writer.write(lut_.decode.data(), lut_.decode.size()); + writer.write(lut_.ends.data(), lut_.ends.size()); + const float params[] = {cfg.k1, cfg.b, cfg.avgdl}; + writer.write(params, sizeof(params)); + } } // write plists_woffsets_formats_mask and plists_window_nnzs @@ -1155,8 +1173,7 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_inner_dims_, sizeof(uint32_t)); InvertedIndexQuantType quant_type{}; reader.read(&quant_type, sizeof(quant_type)); - if (packed_ ? quant_type != InvertedIndexQuantType::IP_E5M7 - : !validate_posting_quant_type(quant_type)) { + if (packed_ ? quant_type != packed_quant_type() : !validate_posting_quant_type(quant_type)) { return Status::invalid_serialized_index_type; } reader.advance(kInvertedIndexHeaderReservedBytes); @@ -1204,7 +1221,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex(packed_ ? InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7 + static_cast(packed_ ? packed_encoding() : InvertedIndexEncoding::FIXED_DOCID_WINDOWS)) { return Status::invalid_serialized_index_type; } @@ -1214,6 +1231,18 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_windows_, sizeof(uint32_t)); if (packed_) { reader.advance(3 * sizeof(uint32_t)); // descriptor checked in preflight + if constexpr (is_bm25) { + reader.read(lut_.decode.data(), lut_.decode.size()); + reader.read(lut_.ends.data(), lut_.ends.size()); + float params[3]; + reader.read(params, sizeof(params)); + lut_.k1 = params[0]; + lut_.validate_and_encode(); + this->set_build_scorer(IndexScorerConfig{ + .scorer_type = IndexScorerType::BM25, + .scorer_params = { + .bm25 = {.k1 = params[0], .b = params[1], .avgdl = params[2]}}}); + } } if (this->window_size_ == 0 || this->window_size_ >= 65536) { @@ -1234,7 +1263,7 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_inner_dims_; - const uint64_t bytes_header = static_cast(sizeof(uint32_t) * (packed_ ? 6 : 3)); + const uint64_t bytes_header = static_cast(posting_header_bytes()); window_index_plists_sz_.clear(); window_index_plists_sz_spans_.clear(); @@ -1295,7 +1324,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex(nr_dims + 1) * sizeof(uint32_t); - const uint64_t bytes_postings_data = packed_ ? 2 * sindi::packed12_bytes(total_postings) + const uint64_t bytes_postings_data = packed_ ? packed_payload_bytes(total_postings) : static_cast(total_postings) * (sizeof(uint16_t) + sizeof(QuantType)); const uint64_t expected_section_bytes = @@ -1308,8 +1337,9 @@ class SindiInvertedIndex : public DimMapInvertedIndex(bytes_ids)}; reader.advance(bytes_ids); @@ -1442,15 +1472,17 @@ class SindiInvertedIndex : public DimMapInvertedIndexnr_inner_dims_; ++dim) { - float maximum = 0; - for (size_t j = 0; j < posting_count(dim); ++j) { - maximum = std::max(maximum, posting_value_at(dim, j)); - } + float maximum = packed_dimension_maximum(dim, refinement_seek_[dim]); if (max_scores_per_dim_span_[dim] != maximum) { return Status::invalid_serialized_index_type; @@ -1670,6 +1702,7 @@ class SindiInvertedIndex : public DimMapInvertedIndexwindow_begin != overflow_cur->window_end) { dispatch_max = apply_bm25_u8_overflow_corrections( @@ -1873,6 +1911,7 @@ class SindiInvertedIndex : public DimMapInvertedIndexwindow_begin != overflow_cur->window_end) { (void)apply_bm25_u8_overflow_corrections( @@ -2128,7 +2172,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex row_sums_; std::span row_sums_span_; - static bool + bool validate_packed_file(const MemoryIOReader& reader) { try { const auto* data = reader.data(); @@ -2142,7 +2186,7 @@ class SindiInvertedIndex : public DimMapInvertedIndex(InvertedIndexQuantType::IP_E5M7)) { + u32(16) != static_cast(packed_quant_type())) { return false; } @@ -2163,6 +2207,10 @@ class SindiInvertedIndex : public DimMapInvertedIndexsize != dims * 4 || h->size < 24 || - u32(h->offset) != static_cast(InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7) || - u32(h->offset + 4) != 4096 || u32(h->offset + 12) != 1 || u32(h->offset + 16) != 2 || - u32(h->offset + 20) != static_cast(InvertedIndexQuantType::IP_E5M7)) { + maxima->size != dims * 4 || h->size < posting_header_bytes() || + u32(h->offset) != static_cast(packed_encoding()) || + (u32(h->offset + 4) < min_window_size || + u32(h->offset + 4) > (packed_u16_ids_ ? max_window_size : 4096) || + (is_ip && u32(h->offset + 4) != 4096)) || + u32(h->offset + 12) != 1 || u32(h->offset + 16) != packed_id_layout() || + u32(h->offset + 20) != static_cast(packed_quant_type())) { return false; } - if (u32(h->offset + 8) != (uint64_t(u32(4)) + 4095) / 4096) { + if (u32(h->offset + 8) != (uint64_t(u32(4)) + u32(h->offset + 4) - 1) / u32(h->offset + 4)) { return false; } - size_t pos = h->offset + 24, end = h->offset + h->size; + if constexpr (is_bm25) { + if (types.count(static_cast(InvertedIndexSectionType::SINDI_REFINEMENT))) + return false; + const auto* lengths = find_section_header(headers, InvertedIndexSectionType::ROW_SUMS); + if ((u32(4) && !lengths) || (lengths && (lengths->size != uint64_t(u32(4)) * 4 || lengths->offset % 4))) + return false; + sindi::Bm25U4Lut lut; + std::memcpy(lut.decode.data(), data + h->offset + 24, 16); + std::memcpy(lut.ends.data(), data + h->offset + 40, 16); + float params[3]; + std::memcpy(params, data + h->offset + 56, 12); + if (!std::isfinite(params[1]) || params[1] < 0 || params[1] > 1 || !std::isfinite(params[2]) || + params[2] < 1) + return false; + lut.k1 = params[0]; + lut.validate_and_encode(); + } + size_t pos = h->offset + posting_header_bytes(), end = h->offset + h->size; auto advance = [&](size_t bytes) { if (pos > end || bytes > end - pos) { throw std::runtime_error("Truncated packed postings"); @@ -2233,14 +2301,15 @@ class SindiInvertedIndex : public DimMapInvertedIndex packed_ids_, packed_vals_; + sindi::Bm25U4Lut lut_; + + constexpr InvertedIndexQuantType + packed_quant_type() const { + return is_ip ? InvertedIndexQuantType::IP_E5M7 + : (packed_u16_ids_ ? InvertedIndexQuantType::BM25_U4_LUT_U16 + : InvertedIndexQuantType::BM25_U4_LUT_U12); + } + constexpr uint32_t + packed_id_layout() const { + return packed_u16_ids_ ? 1u : 2u; + } + constexpr InvertedIndexEncoding + packed_encoding() const { + return is_ip ? InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_E5M7 + : (packed_u16_ids_ ? InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U16_U4_LUT + : InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_U4_LUT); + } + size_t + packed_id_bytes(size_t n) const { + return packed_u16_ids_ ? n * sizeof(uint16_t) : sindi::packed12_bytes(n); + } + size_t + posting_header_bytes() const { + return (packed_ ? 24 : 12) + (packed_ && is_bm25 ? 44 : 0); + } + static size_t + packed_value_bytes(size_t n) { + return is_ip ? sindi::packed12_bytes(n) : n / 2 + n % 2; + } + size_t + packed_payload_bytes(size_t n) const { + return packed_id_bytes(n) + packed_value_bytes(n); + } std::span packed_ids_span_, packed_vals_span_; size_t @@ -2262,19 +2366,82 @@ class SindiInvertedIndex : public DimMapInvertedIndex(posting_vals(dim)[pos]); + if (!packed_ready_) + return static_cast(posting_vals(dim)[pos]); + const size_t posting = plists_dim_offsets_span_[dim] + pos; + if constexpr (is_bm25) + return lut_.decode[sindi::unpack_bm25_u4(packed_vals_span_.data(), posting)]; + else + return sindi::decode_e5m7(sindi::unpack12(packed_vals_span_.data(), posting)); + } + + float + packed_dimension_maximum(size_t dim, const sindi::RefinementSeek& seek) const { + float maximum = 0; + for (size_t w = 0; w + 1 < seek.offsets.size(); ++w) { + const uint32_t wid = seek.sparse ? seek.windows[w] : w; + for (size_t j = seek.offsets[w]; j < seek.offsets[w + 1]; ++j) { + float score = posting_value_at(dim, j); + if constexpr (is_bm25) { + const auto& p = this->build_scorer_->config().scorer_params.bm25; + const float dl = row_sums_span_[size_t(wid) * window_size_ + posting_id_at(dim, j)]; + score = (p.k1 + 1) * score / (score + p.k1 * (1 - p.b) + p.k1 * p.b / p.avgdl * dl); + } + maximum = std::max(maximum, score); + } + } + return maximum; } void pack_postings() { + if constexpr (is_bm25_u8) { + std::array histogram{}; + for (const auto tf : total_plists_vals_flat_span_) ++histogram[tf]; + lut_ = sindi::fit_bm25_u4_lut(histogram, this->build_scorer_->config().scorer_params.bm25.k1); + const size_t count = plists_dim_offsets_span_.back(); + packed_ids_.assign(packed_id_bytes(count), 0); + packed_vals_.assign(packed_value_bytes(count), 0); + for (size_t i = 0; i < count; ++i) { + if (packed_u16_ids_) + std::memcpy(packed_ids_.data() + i * sizeof(uint16_t), &total_plists_ids_flat_span_[i], + sizeof(uint16_t)); + else + sindi::pack12(packed_ids_.data(), i, total_plists_ids_flat_span_[i]); + packed_vals_[i / 2] |= lut_.encode[total_plists_vals_flat_span_[i]] << (4 * (i & 1)); + } + packed_ids_span_ = packed_ids_; + packed_vals_span_ = packed_vals_; + packed_ready_ = true; + rebuild_refinement_seek(); // temporary window validation; never retained by unrefined storage + max_scores_per_dim_.assign(this->nr_inner_dims_, 0); + for (size_t d = 0; d < this->nr_inner_dims_; ++d) + max_scores_per_dim_[d] = packed_dimension_maximum(d, refinement_seek_[d]); + max_scores_per_dim_span_ = max_scores_per_dim_; + std::vector{}.swap(refinement_seek_); + aligned_u16_vec{}.swap(total_plists_ids_flat_); + aligned_quant_vec{}.swap(total_plists_vals_flat_); + std::vector{}.swap(total_plists_ids_); + std::vector{}.swap(total_plists_vals_); + std::vector>{}.swap(total_plists_ids_spans_); + std::vector>{}.swap(total_plists_vals_spans_); + total_plists_ids_flat_span_ = {}; + total_plists_vals_flat_span_ = {}; + std::vector{}.swap(bm25_u8_overflow_offsets_); + std::vector{}.swap(bm25_u8_overflow_values_); + bm25_u8_overflow_offsets_span_ = {}; + bm25_u8_overflow_values_span_ = {}; + } if constexpr (is_ip) { const size_t count = plists_dim_offsets_span_.back(); std::vector ids(sindi::packed12_bytes(count), 0), vals(ids.size(), 0); diff --git a/src/index/sparse/sindi_simd.cc b/src/index/sparse/sindi_simd.cc index 8e9a8522c..4ebb1ef10 100644 --- a/src/index/sparse/sindi_simd.cc +++ b/src/index/sparse/sindi_simd.cc @@ -1,5 +1,6 @@ #include "index/sparse/sindi_simd.h" +#include "index/sparse/sindi_bm25_u4.h" #include "index/sparse/sindi_packed12.h" #include "simd/hook.h" @@ -150,12 +151,13 @@ get_bm25_kernels() { BM25Kernels k{}; #if defined(__x86_64__) const bool support_f16c = faiss::cppcontrib::knowhere::cpu_support_f16c(); - if (support_f16c && faiss::cppcontrib::knowhere::cpu_support_avx512()) { + if (faiss::cppcontrib::knowhere::use_avx512 && support_f16c && + faiss::cppcontrib::knowhere::cpu_support_avx512()) { k.accumulate = bm25_accumulate_avx512_u16; k.batch_insert = batch_insert_avx512; return k; } - if (support_f16c && faiss::cppcontrib::knowhere::cpu_support_avx2()) { + if (faiss::cppcontrib::knowhere::use_avx2 && support_f16c && faiss::cppcontrib::knowhere::cpu_support_avx2()) { k.accumulate = bm25_accumulate_avx2_u16; k.batch_insert = batch_insert_avx2; return k; @@ -179,12 +181,12 @@ get_bm25_u8_kernels() { static const BM25U8Kernels kernels = []() { BM25U8Kernels k{}; #if defined(__x86_64__) - if (faiss::cppcontrib::knowhere::cpu_support_avx512()) { + if (faiss::cppcontrib::knowhere::use_avx512 && faiss::cppcontrib::knowhere::cpu_support_avx512()) { k.accumulate = bm25_accumulate_avx512_u8; k.batch_insert = batch_insert_avx512; return k; } - if (faiss::cppcontrib::knowhere::cpu_support_avx2()) { + if (faiss::cppcontrib::knowhere::use_avx2 && faiss::cppcontrib::knowhere::cpu_support_avx2()) { k.accumulate = bm25_accumulate_avx2_u8; k.batch_insert = batch_insert_avx2; return k; @@ -203,4 +205,58 @@ get_bm25_u8_kernels() { return kernels; } +float +bm25_accumulate_scalar_u12_u4_lut(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out, + float k1, float b, float avgdl, const float* lengths, const uint8_t* lut) { + const float p1 = q * (k1 + 1), p2 = k1 * (1 - b), p3 = k1 * b / avgdl; + float maximum = 0; + for (int32_t i = 0; i < n; ++i) { + const size_t position = start + i; + const auto id = unpack12(ids, position); + const float tf = lut[(vals[position / 2] >> (4 * (position & 1))) & 15]; + const float contribution = p1 * tf / (tf + p2 + p3 * lengths[id]); + out[id] += contribution; + maximum = std::max(maximum, out[id]); + } + return maximum; +} + +float +bm25_accumulate_scalar_u16_u4_lut(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out, + float k1, float b, float avgdl, const float* lengths, const uint8_t* lut) { + const float p1 = q * (k1 + 1), p2 = k1 * (1 - b), p3 = k1 * b / avgdl; + float maximum = 0; + for (int32_t i = 0; i < n; ++i) { + const size_t position = start + i; + const auto id = unpack_bm25_u16_id(ids, position); + const float tf = lut[(vals[position / 2] >> (4 * (position & 1))) & 15]; + const float contribution = p1 * tf / (tf + p2 + p3 * lengths[id]); + out[id] += contribution; + maximum = std::max(maximum, out[id]); + } + return maximum; +} + +packed_bm25_accumulate_fn_t +get_packed_bm25_kernel(bool u16_ids) { +#if defined(__x86_64__) + namespace cpu = faiss::cppcontrib::knowhere; + if (cpu::use_avx512 && cpu::cpu_support_avx512() && __builtin_cpu_supports("avx512vl") && + __builtin_cpu_supports("avx512cd") && __builtin_cpu_supports("avx2") && __builtin_cpu_supports("fma") && + cpu::cpu_support_f16c()) { + return u16_ids ? bm25_accumulate_avx512_u16_u4_lut : bm25_accumulate_avx512_u12_u4_lut; + } + if (cpu::use_avx2 && cpu::cpu_support_avx2() && __builtin_cpu_supports("avx2") && __builtin_cpu_supports("fma") && + cpu::cpu_support_f16c()) { + return u16_ids ? bm25_accumulate_avx2_u16_u4_lut : bm25_accumulate_avx2_u12_u4_lut; + } +#endif +#if defined(__aarch64__) && defined(KNOWHERE_USE_SVE) + if (faiss::cppcontrib::knowhere::supports_sve()) { + return u16_ids ? bm25_accumulate_sve_u16_u4_lut : bm25_accumulate_sve_u12_u4_lut; + } +#endif + return u16_ids ? bm25_accumulate_scalar_u16_u4_lut : bm25_accumulate_scalar_u12_u4_lut; +} + } // namespace knowhere::sparse::inverted::sindi diff --git a/src/index/sparse/sindi_simd.h b/src/index/sparse/sindi_simd.h index 45678dded..bdb9b958c 100644 --- a/src/index/sparse/sindi_simd.h +++ b/src/index/sparse/sindi_simd.h @@ -13,6 +13,25 @@ using packed_ip_accumulate_fn_t = float (*)(float, const uint8_t*, const uint8_t packed_ip_accumulate_fn_t get_packed_ip_kernel(); +using packed_bm25_accumulate_fn_t = float (*)(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, + float, float, const float*, const uint8_t*); +packed_bm25_accumulate_fn_t +get_packed_bm25_kernel(bool u16_ids = false); +float +bm25_accumulate_scalar_u12_u4_lut(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, float, float, + const float*, const uint8_t*); +float +bm25_accumulate_scalar_u16_u4_lut(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, float, float, + const float*, const uint8_t*); +#if defined(__aarch64__) && defined(KNOWHERE_USE_SVE) +float +bm25_accumulate_sve_u12_u4_lut(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, float, float, + const float*, const uint8_t*); +float +bm25_accumulate_sve_u16_u4_lut(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, float, float, + const float*, const uint8_t*); +#endif + float ip_accumulate_scalar_u12_e5m7(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*); @@ -71,6 +90,18 @@ batch_insert_scalar(const float* scores, size_t docid_start, size_t count, #if defined(__x86_64__) float +bm25_accumulate_avx2_u12_u4_lut(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, float, float, + const float*, const uint8_t*); +float +bm25_accumulate_avx2_u16_u4_lut(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, float, float, + const float*, const uint8_t*); +float +bm25_accumulate_avx512_u12_u4_lut(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, float, float, + const float*, const uint8_t*); +float +bm25_accumulate_avx512_u16_u4_lut(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*, float, float, float, + const float*, const uint8_t*); +float ip_accumulate_avx2_u12_e5m7(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*); float ip_accumulate_avx512_u12_e5m7(float, const uint8_t*, const uint8_t*, size_t, int32_t, float*); diff --git a/src/index/sparse/sindi_simd_avx2.cc b/src/index/sparse/sindi_simd_avx2.cc index 19c59871f..7c1cce88c 100644 --- a/src/index/sparse/sindi_simd_avx2.cc +++ b/src/index/sparse/sindi_simd_avx2.cc @@ -52,6 +52,109 @@ ip_accumulate_avx2_u12_e5m7(float q, const uint8_t* vals, const uint8_t* ids, si return std::max(maximum, ip_accumulate_scalar_u12_e5m7(q, vals, ids, start + i, n - i, out)); } +float +bm25_accumulate_avx2_u12_u4_lut(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out, + float k1, float b, float avgdl, const float* lengths, const uint8_t* table) { + if (n <= 0) { + return 0; + } + + float maximum = 0; + if (start & 1) { + maximum = bm25_accumulate_scalar_u12_u4_lut(q, vals, ids, start, 1, out, k1, b, avgdl, lengths, table); + ++start; + --n; + } + const auto lut = _mm_loadu_si128(reinterpret_cast(table)); + const auto nibble_mask = _mm_set1_epi8(15); + const auto vqp1 = _mm256_set1_ps(q * (k1 + 1.0f)); + const auto vp2 = _mm256_set1_ps(k1 * (1.0f - b)), vp3 = _mm256_set1_ps(k1 * b / avgdl); + auto vmax = _mm256_setzero_ps(); + + int32_t i = 0; + for (; i + 8 <= n; i += 8) { + const auto id = _mm256_cvtepu16_epi32(unpack12_eight_x86(ids + ((start + i) / 2) * 3)); + // Eight codes occupy exactly four bytes; the LUT load is exactly 16. + uint32_t packed; + std::memcpy(&packed, vals + (start + i) / 2, sizeof(packed)); + const auto bytes = _mm_cvtsi32_si128(packed); + const auto codes = _mm_and_si128(_mm_unpacklo_epi8(bytes, _mm_srli_epi16(bytes, 4)), nibble_mask); + const auto tf = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(_mm_shuffle_epi8(lut, codes))); + const auto dl = _mm256_i32gather_ps(lengths, id, 4); + const auto denominator = _mm256_add_ps(tf, _mm256_fmadd_ps(dl, vp3, vp2)); + const auto contribution = _mm256_div_ps(_mm256_mul_ps(tf, vqp1), denominator); + const auto sum = _mm256_add_ps(_mm256_i32gather_ps(out, id, 4), contribution); + + // Match existing AVX2 BM25 writeback: gather is available, scatter is not. + alignas(32) uint32_t indices[8]; + alignas(32) float scores[8]; + _mm256_store_si256(reinterpret_cast<__m256i*>(indices), id); + _mm256_store_ps(scores, sum); + for (int lane = 0; lane < 8; ++lane) { + out[indices[lane]] = scores[lane]; + } + vmax = _mm256_max_ps(vmax, sum); + } + auto tail_max = _mm_max_ps(_mm256_castps256_ps128(vmax), _mm256_extractf128_ps(vmax, 1)); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(2, 3, 0, 1))); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(1, 0, 3, 2))); + maximum = std::max(maximum, _mm_cvtss_f32(tail_max)); + return std::max( + maximum, bm25_accumulate_scalar_u12_u4_lut(q, vals, ids, start + i, n - i, out, k1, b, avgdl, lengths, table)); +} + +float +bm25_accumulate_avx2_u16_u4_lut(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out, + float k1, float b, float avgdl, const float* lengths, const uint8_t* table) { + if (n <= 0) { + return 0; + } + + float maximum = 0; + if (start & 1) { + maximum = bm25_accumulate_scalar_u16_u4_lut(q, vals, ids, start, 1, out, k1, b, avgdl, lengths, table); + ++start; + --n; + } + const auto lut = _mm_loadu_si128(reinterpret_cast(table)); + const auto nibble_mask = _mm_set1_epi8(15); + const auto vqp1 = _mm256_set1_ps(q * (k1 + 1.0f)); + const auto vp2 = _mm256_set1_ps(k1 * (1.0f - b)), vp3 = _mm256_set1_ps(k1 * b / avgdl); + auto vmax = _mm256_setzero_ps(); + + int32_t i = 0; + for (; i + 8 <= n; i += 8) { + const auto id = _mm256_cvtepu16_epi32( + _mm_loadu_si128(reinterpret_cast(ids + (start + i) * sizeof(uint16_t)))); + // Eight codes occupy exactly four bytes; the LUT load is exactly 16. + uint32_t packed; + std::memcpy(&packed, vals + (start + i) / 2, sizeof(packed)); + const auto bytes = _mm_cvtsi32_si128(packed); + const auto codes = _mm_and_si128(_mm_unpacklo_epi8(bytes, _mm_srli_epi16(bytes, 4)), nibble_mask); + const auto tf = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(_mm_shuffle_epi8(lut, codes))); + const auto dl = _mm256_i32gather_ps(lengths, id, 4); + const auto denominator = _mm256_add_ps(tf, _mm256_fmadd_ps(dl, vp3, vp2)); + const auto contribution = _mm256_div_ps(_mm256_mul_ps(tf, vqp1), denominator); + const auto sum = _mm256_add_ps(_mm256_i32gather_ps(out, id, 4), contribution); + + // Match existing AVX2 BM25 writeback: gather is available, scatter is not. + alignas(32) uint32_t indices[8]; + alignas(32) float scores[8]; + _mm256_store_si256(reinterpret_cast<__m256i*>(indices), id); + _mm256_store_ps(scores, sum); + for (int lane = 0; lane < 8; ++lane) { + out[indices[lane]] = scores[lane]; + } + vmax = _mm256_max_ps(vmax, sum); + } + auto tail_max = _mm_max_ps(_mm256_castps256_ps128(vmax), _mm256_extractf128_ps(vmax, 1)); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(2, 3, 0, 1))); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(1, 0, 3, 2))); + maximum = std::max(maximum, _mm_cvtss_f32(tail_max)); + return std::max( + maximum, bm25_accumulate_scalar_u16_u4_lut(q, vals, ids, start + i, n - i, out, k1, b, avgdl, lengths, table)); +} + float ip_accumulate_avx2_fp16(float qval, const knowhere::fp16* vals, const uint16_t* ids, int32_t num, float* out) { int32_t i = 0; diff --git a/src/index/sparse/sindi_simd_avx512.cc b/src/index/sparse/sindi_simd_avx512.cc index 7b0a87d6e..7151652ac 100644 --- a/src/index/sparse/sindi_simd_avx512.cc +++ b/src/index/sparse/sindi_simd_avx512.cc @@ -57,6 +57,133 @@ ip_accumulate_avx512_u12_e5m7(float q, const uint8_t* vals, const uint8_t* ids, return std::max(maximum, ip_accumulate_scalar_u12_e5m7(q, vals, ids, start + i, n - i, out)); } +float +bm25_accumulate_avx512_u12_u4_lut(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out, + float k1, float b, float avgdl, const float* lengths, const uint8_t* table) { + if (n <= 0) { + return 0; + } + + float maximum = 0; + // Align both canonical packed streams to a pair. No expanded posting cache. + if (start & 1) { + maximum = bm25_accumulate_scalar_u12_u4_lut(q, vals, ids, start, 1, out, k1, b, avgdl, lengths, table); + ++start; + --n; + } + + const auto lut = _mm_loadu_si128(reinterpret_cast(table)); + const auto nibble_mask = _mm_set1_epi8(15); + const float qp1 = q * (k1 + 1.0f), p2 = k1 * (1.0f - b), p3 = k1 * b / avgdl; + const auto vqp1 = _mm512_set1_ps(qp1), vp2 = _mm512_set1_ps(p2), vp3 = _mm512_set1_ps(p3); + auto vmax = _mm512_setzero_ps(); + + int32_t i = 0; + for (; i + 16 <= n; i += 16) { + const auto* ip = ids + ((start + i) / 2) * 3; + const auto id16 = _mm256_set_m128i(unpack12_eight_x86(ip + 12), unpack12_eight_x86(ip)); + const auto id = _mm512_cvtepu16_epi32(id16); + // Sixteen codes occupy exactly eight bytes. Interleave low/high nibbles + // and look up their TF representatives in a register, before widening. + const auto bytes = _mm_loadl_epi64(reinterpret_cast(vals + (start + i) / 2)); + const auto codes = _mm_and_si128(_mm_unpacklo_epi8(bytes, _mm_srli_epi16(bytes, 4)), nibble_mask); + const auto tf = _mm512_cvtepi32_ps(_mm512_cvtepu8_epi32(_mm_shuffle_epi8(lut, codes))); + const auto dl = _mm512_i32gather_ps(id, lengths, 4); + const auto denominator = _mm512_add_ps(tf, _mm512_fmadd_ps(dl, vp3, vp2)); + const auto contribution = _mm512_div_ps(_mm512_mul_ps(tf, vqp1), denominator); + const auto sum = _mm512_add_ps(_mm512_i32gather_ps(id, out, 4), contribution); + _mm512_i32scatter_ps(out, id, sum, 4); + vmax = _mm512_max_ps(vmax, sum); + } + maximum = std::max(maximum, _mm512_reduce_max_ps(vmax)); + + // AVX512VL handles eight postings without reading sixteen codes or IDs. + if (i + 8 <= n) { + const auto id = _mm256_cvtepu16_epi32(unpack12_eight_x86(ids + ((start + i) / 2) * 3)); + uint32_t packed; + std::memcpy(&packed, vals + (start + i) / 2, sizeof(packed)); + const auto bytes = _mm_cvtsi32_si128(packed); + const auto codes = _mm_and_si128(_mm_unpacklo_epi8(bytes, _mm_srli_epi16(bytes, 4)), nibble_mask); + const auto tf = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(_mm_shuffle_epi8(lut, codes))); + const auto dl = _mm256_i32gather_ps(lengths, id, 4); + const auto denominator = _mm256_add_ps(tf, _mm256_fmadd_ps(dl, _mm256_set1_ps(p3), _mm256_set1_ps(p2))); + const auto contribution = _mm256_div_ps(_mm256_mul_ps(tf, _mm256_set1_ps(qp1)), denominator); + const auto sum = _mm256_add_ps(_mm256_i32gather_ps(out, id, 4), contribution); + _mm256_i32scatter_ps(out, id, sum, 4); + auto tail_max = _mm_max_ps(_mm256_castps256_ps128(sum), _mm256_extractf128_ps(sum, 1)); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(2, 3, 0, 1))); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(1, 0, 3, 2))); + maximum = std::max(maximum, _mm_cvtss_f32(tail_max)); + i += 8; + } + return std::max( + maximum, bm25_accumulate_scalar_u12_u4_lut(q, vals, ids, start + i, n - i, out, k1, b, avgdl, lengths, table)); +} + +float +bm25_accumulate_avx512_u16_u4_lut(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t n, float* out, + float k1, float b, float avgdl, const float* lengths, const uint8_t* table) { + if (n <= 0) { + return 0; + } + + float maximum = 0; + // Align the packed TF stream to a pair. No expanded posting cache. + if (start & 1) { + maximum = bm25_accumulate_scalar_u16_u4_lut(q, vals, ids, start, 1, out, k1, b, avgdl, lengths, table); + ++start; + --n; + } + + const auto lut = _mm_loadu_si128(reinterpret_cast(table)); + const auto nibble_mask = _mm_set1_epi8(15); + const float qp1 = q * (k1 + 1.0f), p2 = k1 * (1.0f - b), p3 = k1 * b / avgdl; + const auto vqp1 = _mm512_set1_ps(qp1), vp2 = _mm512_set1_ps(p2), vp3 = _mm512_set1_ps(p3); + auto vmax = _mm512_setzero_ps(); + + int32_t i = 0; + for (; i + 16 <= n; i += 16) { + const auto* ip = ids + (start + i) * sizeof(uint16_t); + const auto id16 = _mm256_loadu_si256(reinterpret_cast(ip)); + const auto id = _mm512_cvtepu16_epi32(id16); + // Sixteen codes occupy exactly eight bytes. Interleave low/high nibbles + // and look up their TF representatives in a register, before widening. + const auto bytes = _mm_loadl_epi64(reinterpret_cast(vals + (start + i) / 2)); + const auto codes = _mm_and_si128(_mm_unpacklo_epi8(bytes, _mm_srli_epi16(bytes, 4)), nibble_mask); + const auto tf = _mm512_cvtepi32_ps(_mm512_cvtepu8_epi32(_mm_shuffle_epi8(lut, codes))); + const auto dl = _mm512_i32gather_ps(id, lengths, 4); + const auto denominator = _mm512_add_ps(tf, _mm512_fmadd_ps(dl, vp3, vp2)); + const auto contribution = _mm512_div_ps(_mm512_mul_ps(tf, vqp1), denominator); + const auto sum = _mm512_add_ps(_mm512_i32gather_ps(id, out, 4), contribution); + _mm512_i32scatter_ps(out, id, sum, 4); + vmax = _mm512_max_ps(vmax, sum); + } + maximum = std::max(maximum, _mm512_reduce_max_ps(vmax)); + + // AVX512VL handles eight postings without reading sixteen codes or IDs. + if (i + 8 <= n) { + const auto id = _mm256_cvtepu16_epi32( + _mm_loadu_si128(reinterpret_cast(ids + (start + i) * sizeof(uint16_t)))); + uint32_t packed; + std::memcpy(&packed, vals + (start + i) / 2, sizeof(packed)); + const auto bytes = _mm_cvtsi32_si128(packed); + const auto codes = _mm_and_si128(_mm_unpacklo_epi8(bytes, _mm_srli_epi16(bytes, 4)), nibble_mask); + const auto tf = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(_mm_shuffle_epi8(lut, codes))); + const auto dl = _mm256_i32gather_ps(lengths, id, 4); + const auto denominator = _mm256_add_ps(tf, _mm256_fmadd_ps(dl, _mm256_set1_ps(p3), _mm256_set1_ps(p2))); + const auto contribution = _mm256_div_ps(_mm256_mul_ps(tf, _mm256_set1_ps(qp1)), denominator); + const auto sum = _mm256_add_ps(_mm256_i32gather_ps(out, id, 4), contribution); + _mm256_i32scatter_ps(out, id, sum, 4); + auto tail_max = _mm_max_ps(_mm256_castps256_ps128(sum), _mm256_extractf128_ps(sum, 1)); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(2, 3, 0, 1))); + tail_max = _mm_max_ps(tail_max, _mm_shuffle_ps(tail_max, tail_max, _MM_SHUFFLE(1, 0, 3, 2))); + maximum = std::max(maximum, _mm_cvtss_f32(tail_max)); + i += 8; + } + return std::max( + maximum, bm25_accumulate_scalar_u16_u4_lut(q, vals, ids, start + i, n - i, out, k1, b, avgdl, lengths, table)); +} + float ip_accumulate_avx512_fp16(float qval, const knowhere::fp16* vals, const uint16_t* ids, int32_t num, float* out) { int32_t i = 0; diff --git a/src/index/sparse/sindi_simd_sve.cc b/src/index/sparse/sindi_simd_sve.cc index 27f7b1ce7..51a10218c 100644 --- a/src/index/sparse/sindi_simd_sve.cc +++ b/src/index/sparse/sindi_simd_sve.cc @@ -498,6 +498,78 @@ batch_insert_sve(const float* scores, size_t docid_start, size_t count, } } +// Both streams are decoded directly into matching U32 lanes. All byte loads +// are predicated to their exact extent, including odd starts and partial tails. +float +bm25_accumulate_sve_u12_u4_lut(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t count, + float* out, float k1, float b, float avgdl, const float* lengths, const uint8_t* lut) { + if (count <= 0) + return 0; + const uint32_t vl = svcntw(), phase = start & 1; + const auto full = svptrue_b32(); + const auto lane = svindex_u32(0, 1); + const auto bit12 = svadd_n_u32_x(full, svmul_n_u32_x(full, lane, 12), phase * 4); + const auto byte12 = svlsr_n_u32_x(full, bit12, 3); + const auto table12 = svorr_u32_x(full, byte12, svlsl_n_u32_x(full, svadd_n_u32_x(full, byte12, 1), 8)); + const auto shift12 = svand_n_u32_x(full, bit12, 7); + const auto byte4 = svlsr_n_u32_x(full, svadd_n_u32_x(full, lane, phase), 1); + const auto shift4 = svlsl_n_u32_x(full, svand_n_u32_x(full, svadd_n_u32_x(full, lane, phase), 1), 2); + const auto decode = svld1_u8(svwhilelt_b8(0u, 16u), lut); + const auto p1 = svdup_f32(q * (k1 + 1)); + const float p2 = k1 * (1 - b), p3 = k1 * b / avgdl; + auto maximum = svdup_f32(0); + for (uint32_t i = 0; i < uint32_t(count); i += vl) { + const uint32_t n = std::min(vl, uint32_t(count) - i); + const auto pg = svwhilelt_b32(0u, n); + const auto id = load12(ids + 3 * ((start + i) / 2) + phase, n, phase * 4, table12, shift12); + const auto raw = svld1_u8(svwhilelt_b8(0u, (n + phase + 1) / 2), vals + (start + i) / 2); + const auto bytes = svreinterpret_u32_u8(svtbl_u8(raw, svreinterpret_u8_u32(byte4))); + const auto code = svand_n_u32_x(full, svlsr_u32_x(full, bytes, shift4), 15); + const auto tfword = svreinterpret_u32_u8(svtbl_u8(decode, svreinterpret_u8_u32(code))); + const auto tf = svcvt_f32_u32_x(pg, tfword); + const auto dl = svld1_gather_u32index_f32(pg, lengths, id); + const auto denominator = svadd_f32_x(pg, svadd_n_f32_x(pg, tf, p2), svmul_n_f32_x(pg, dl, p3)); + const auto contribution = svdiv_f32_x(pg, svmul_f32_x(pg, p1, tf), denominator); + const auto sum = svadd_f32_x(pg, svld1_gather_u32index_f32(pg, out, id), contribution); + svst1_scatter_u32index_f32(pg, out, id, sum); + maximum = svmax_f32_m(pg, maximum, sum); + } + return svmaxv_f32(full, maximum); +} + +float +bm25_accumulate_sve_u16_u4_lut(float q, const uint8_t* vals, const uint8_t* ids, size_t start, int32_t count, + float* out, float k1, float b, float avgdl, const float* lengths, const uint8_t* lut) { + if (count <= 0) + return 0; + const uint32_t vl = svcntw(), phase = start & 1; + const auto full = svptrue_b32(); + const auto lane = svindex_u32(0, 1); + const auto byte4 = svlsr_n_u32_x(full, svadd_n_u32_x(full, lane, phase), 1); + const auto shift4 = svlsl_n_u32_x(full, svand_n_u32_x(full, svadd_n_u32_x(full, lane, phase), 1), 2); + const auto decode = svld1_u8(svwhilelt_b8(0u, 16u), lut); + const auto p1 = svdup_f32(q * (k1 + 1)); + const float p2 = k1 * (1 - b), p3 = k1 * b / avgdl; + auto maximum = svdup_f32(0); + for (uint32_t i = 0; i < uint32_t(count); i += vl) { + const uint32_t n = std::min(vl, uint32_t(count) - i); + const auto pg = svwhilelt_b32(0u, n); + const auto id = svld1uh_u32(pg, reinterpret_cast(ids + (start + i) * sizeof(uint16_t))); + const auto raw = svld1_u8(svwhilelt_b8(0u, (n + phase + 1) / 2), vals + (start + i) / 2); + const auto bytes = svreinterpret_u32_u8(svtbl_u8(raw, svreinterpret_u8_u32(byte4))); + const auto code = svand_n_u32_x(full, svlsr_u32_x(full, bytes, shift4), 15); + const auto tfword = svreinterpret_u32_u8(svtbl_u8(decode, svreinterpret_u8_u32(code))); + const auto tf = svcvt_f32_u32_x(pg, tfword); + const auto dl = svld1_gather_u32index_f32(pg, lengths, id); + const auto denominator = svadd_f32_x(pg, svadd_n_f32_x(pg, tf, p2), svmul_n_f32_x(pg, dl, p3)); + const auto contribution = svdiv_f32_x(pg, svmul_f32_x(pg, p1, tf), denominator); + const auto sum = svadd_f32_x(pg, svld1_gather_u32index_f32(pg, out, id), contribution); + svst1_scatter_u32index_f32(pg, out, id, sum); + maximum = svmax_f32_m(pg, maximum, sum); + } + return svmaxv_f32(full, maximum); +} + } // namespace knowhere::sparse::inverted::sindi #endif diff --git a/src/index/sparse/sparse_index_config.h b/src/index/sparse/sparse_index_config.h index 78eb5190c..136dd9871 100644 --- a/src/index/sparse/sparse_index_config.h +++ b/src/index/sparse/sparse_index_config.h @@ -35,6 +35,16 @@ namespace knowhere { +inline bool +IsBm25U4QuantType(const std::string& quant) { + return quant == "u4_lut" || quant == "u4_lut_u12" || quant == "u4_lut_u16"; +} + +inline std::string +NormalizeBm25U4QuantType(const std::string& quant) { + return quant == "u4_lut" ? "u4_lut_u12" : quant; +} + inline std::string NormalizeSparseInvertedIndexAlgo(std::string algo) { std::transform(algo.begin(), algo.end(), algo.begin(), @@ -177,7 +187,8 @@ class SparseInvertedIndexConfig : public BaseConfig { KNOWHERE_CONFIG_DECLARE_FIELD(quant_type) .description( "quantization type for posting list values: fp16/fp32/e5m7 for IP (e5m7 requires SINDI, window=4096), " - "u8/u16/u32/auto for BM25; u8 is " + "u4_lut_u12/u4_lut_u16/u8/u16/u32/auto for BM25; u4_lut aliases u4_lut_u12; U4 requires sealed SINDI, " + "no refinement; U12 window<=4096; u8 is " "supported only by sealed SINDI with index version >= 11; BM25 auto requires index version >= 11 " "and resolves to u8/u16 for sealed SINDI or u16 for other indexes; the concrete type is persisted " "in the index and restored automatically on load; the load parameter is used only for legacy " @@ -247,9 +258,12 @@ class SparseInvertedIndexConfig : public BaseConfig { return Status::invalid_args; } } else if (mt == metric::BM25) { - if (qt != "u8" && qt != "u16" && qt != "u32" && qt != "auto") { + if (!IsBm25U4QuantType(qt) && qt != "u8" && qt != "u16" && qt != "u32" && qt != "auto") { if (err_msg) { - *err_msg = "quant_type for BM25 metric must be 'u8', 'u16', 'u32', or 'auto', got '" + qt + "'"; + *err_msg = + "quant_type for BM25 metric must be 'u4_lut'/'u4_lut_u12'/'u4_lut_u16', 'u8', 'u16', " + "'u32', or 'auto', got '" + + qt + "'"; } return Status::invalid_args; } diff --git a/src/index/sparse/sparse_index_node.cc b/src/index/sparse/sparse_index_node.cc index 8d3998498..21940758d 100644 --- a/src/index/sparse/sparse_index_node.cc +++ b/src/index/sparse/sparse_index_node.cc @@ -108,6 +108,34 @@ class SparseInvertedIndexNode : public IndexNode { index_version_(version) { } + Status + Build(const DataSetPtr dataset, std::shared_ptr config, bool use_knowhere_build_pool = true) override { + const auto& cfg = static_cast(*config); + if (!IsBm25U4QuantType(cfg.quant_type.value_or(""))) + return IndexNode::Build(dataset, std::move(config), use_knowhere_build_pool); + // Sealed compact builds publish only after validation, fitting and packing. + if (Type() != IndexEnum::INDEX_SPARSE_INVERTED_INDEX || !dataset || !cfg.bm25_k1.has_value() || + !cfg.bm25_b.has_value() || !cfg.bm25_avgdl.has_value()) + return Status::invalid_args; + try { + auto candidate = CreateIndex(cfg); + if (!candidate.has_value()) + return candidate.error(); + const auto status = + candidate.value()->add(static_cast*>(dataset->GetTensor()), + dataset->GetRows(), dataset->GetDim()); + if (status != Status::success) + return status; + index_ = std::move(candidate.value()); + binary_.reset(); + mmap_guard_.reset(); + return Status::success; + } catch (const std::exception& e) { + LOG_KNOWHERE_WARNING_ << "Failed compact SINDI build: " << e.what(); + return Status::sparse_inner_error; + } + } + Status Train(const DataSetPtr dataset, std::shared_ptr config, bool use_knowhere_build_pool) override { auto cfg = static_cast(*config); @@ -222,6 +250,12 @@ class SparseInvertedIndexNode : public IndexNode { std::string resolved_quant_type; bool is_ip = false; switch (posting_type.value()) { + case InvertedIndexQuantType::BM25_U4_LUT_U12: + resolved_quant_type = "u4_lut_u12"; + break; + case InvertedIndexQuantType::BM25_U4_LUT_U16: + resolved_quant_type = "u4_lut_u16"; + break; case InvertedIndexQuantType::IP_E5M7: resolved_quant_type = "e5m7"; is_ip = true; @@ -251,7 +285,10 @@ class SparseInvertedIndexNode : public IndexNode { if (!IsMetricType(cfg.metric_type.value(), is_ip ? metric::IP : metric::BM25)) { return Status::invalid_serialized_index_type; } - if (resolved_quant_type == "e5m7") { + if (IsBm25U4QuantType(resolved_quant_type) && !cfg.quant_type.value_or("").empty() && + NormalizeBm25U4QuantType(cfg.quant_type.value()) != resolved_quant_type) + return Status::invalid_serialized_index_type; + if (resolved_quant_type == "e5m7" || IsBm25U4QuantType(resolved_quant_type)) { cfg.inverted_index_algo = "SINDI"; cfg.sindi_window_size = 4096; } @@ -488,7 +525,7 @@ class SparseInvertedIndexNode : public IndexNode { return index_or.error(); } - if (cfg.quant_type.value_or("") == "e5m7") { + if (cfg.quant_type.value_or("") == "e5m7" || IsBm25U4QuantType(cfg.quant_type.value_or(""))) { auto candidate = std::move(index_or.value()); MemoryIOReader packed_reader(binary->data.get(), binary->size); const auto status = candidate->deserialize(packed_reader); @@ -556,7 +593,7 @@ class SparseInvertedIndexNode : public IndexNode { return index_or.error(); } - if (cfg.quant_type.value_or("") == "e5m7") { + if (cfg.quant_type.value_or("") == "e5m7" || IsBm25U4QuantType(cfg.quant_type.value_or(""))) { auto candidate = std::move(index_or.value()); MemoryIOReader packed_reader(reinterpret_cast(mapped_memory), map_size); const auto status = candidate->deserialize(packed_reader); @@ -785,14 +822,19 @@ class SparseInvertedIndexNode : public IndexNode { const std::string algo = get_inverted_index_algo("DAAT_MAXSCORE"); const std::string codec = cfg.inverted_index_codec.value_or("block_streamvbyte"); bool use_sindi = - algo == "SINDI" && (!encoding.has_value() || - encoding.value() == sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS); + algo == "SINDI" && + (!encoding.has_value() || + (encoding.value() == sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS || + encoding.value() == sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_U4_LUT || + encoding.value() == sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U16_U4_LUT)); if (use_sindi) { auto window_size = cfg.sindi_window_size.value_or(sparse::inverted::SindiInvertedIndexBM25::max_window_size); IndexPtr index; if constexpr (std::is_same_v) { - index = std::make_unique(window_size); + index = std::make_unique( + window_size, false, IsBm25U4QuantType(cfg.quant_type.value_or("")), + cfg.quant_type.value_or("") == "u4_lut_u16"); } else { if (is_growable) { index = std::make_unique(window_size); @@ -820,6 +862,23 @@ class SparseInvertedIndexNode : public IndexNode { expected>> CreateIndex(const SparseInvertedIndexConfig& cfg, bool is_growable = false, std::optional encoding = std::nullopt) const { + if (IsBm25U4QuantType(cfg.quant_type.value_or(""))) { + const bool loading = encoding.has_value(); + const auto window = cfg.sindi_window_size.value_or(4096); + const bool u16_ids = cfg.quant_type.value_or("") == "u4_lut_u16"; + const auto expected_encoding = + u16_ids ? sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U16_U4_LUT + : sparse::inverted::InvertedIndexEncoding::FIXED_DOCID_WINDOWS_U12_U4_LUT; + if (index_version_ < 11 || is_growable || cfg.refine.value_or(false) || + !IsMetricType(cfg.metric_type.value(), metric::BM25) || + (!loading && (NormalizeInvertedIndexAlgo(cfg.inverted_index_algo.value_or("")) != "SINDI" || + window < 1024 || window > (u16_ids ? 65535 : 4096))) || + !cfg.inverted_index_codec.value_or("").empty() || (loading && *encoding != expected_encoding)) + return expected>>::Err( + Status::invalid_args, + "U4 LUT requires sealed SINDI BM25, U12 window 1024..4096/U16 1024..65535, version>=11, no " + "refinement/block codec"); + } if (cfg.quant_type.value_or("") == "e5m7") { const bool loading = encoding.has_value(); if (index_version_ < 11 || !IsMetricType(cfg.metric_type.value(), metric::IP) || @@ -870,7 +929,7 @@ class SparseInvertedIndexNode : public IndexNode { return CreateIndexImpl(cfg, is_growable, encoding); } } else { - if (qt == "u8") { + if (qt == "u8" || IsBm25U4QuantType(qt)) { if (index_version_ < kBm25AutoU8MinVersion || requested_algo != "SINDI" || is_growable) { return expected>>::Err( Status::invalid_args, "u8 quantization requires sealed SINDI with index version >= 11"); @@ -888,6 +947,9 @@ class SparseInvertedIndexNode : public IndexNode { bool PackedStorageEnabled() const { + if (auto* index = dynamic_cast(index_.get())) { + return index->packed_storage_enabled(); + } if (auto* index = dynamic_cast(index_.get())) { return index->packed_storage_enabled(); } @@ -1013,6 +1075,13 @@ class SparseInvertedIndexNode : public IndexNode { config.bm25_k1.value_or(index_->get_scorer_config().scorer_params.bm25.k1); search_params.scorer_config.scorer_params.bm25.b = config.bm25_b.value_or(index_->get_scorer_config().scorer_params.bm25.b); + if (PackedStorageEnabled()) { + const auto& stored = index_->get_scorer_config().scorer_params.bm25; + const auto& requested = search_params.scorer_config.scorer_params.bm25; + if (requested.k1 != stored.k1 || requested.b != stored.b || requested.avgdl != stored.avgdl) + return expected::Err( + Status::invalid_args, "Packed BM25 search parameters must match the built index"); + } if (search_params.algo == sparse::inverted::InvertedIndexAlgo::DAAT_WAND || search_params.algo == sparse::inverted::InvertedIndexAlgo::DAAT_MAXSCORE || search_params.algo == sparse::inverted::InvertedIndexAlgo::BLOCK_MAX_WAND || diff --git a/tests/ut/test_sparse.cc b/tests/ut/test_sparse.cc index 338b7847d..6da4926e2 100644 --- a/tests/ut/test_sparse.cc +++ b/tests/ut/test_sparse.cc @@ -25,6 +25,7 @@ #include "catch2/catch_test_macros.hpp" #include "catch2/generators/catch_generators.hpp" #include "index/sparse/inverted_index_format.h" +#include "index/sparse/sindi_bm25_u4.h" #include "io/memory_io.h" #include "knowhere/bitsetview.h" #include "knowhere/comp/brute_force.h" @@ -32,6 +33,7 @@ #include "knowhere/comp/knowhere_check.h" #include "knowhere/comp/knowhere_config.h" #include "knowhere/index/index_factory.h" +#include "simd/hook.h" #include "utils.h" void @@ -3221,3 +3223,295 @@ TEST_CASE("SINDI U12 E5M7 sparse windows new dimensions and failure isolation", } } } + +TEST_CASE("SINDI BM25 U12 LUT fitting and exact byte kernels", "[sparse][sindi][u4_u12]") { + using namespace knowhere::sparse::inverted::sindi; + std::array histogram{}; + for (size_t t = 1; t < 256; ++t) histogram[t] = t < 16 ? 1000000 : 1; + const auto lut = fit_bm25_u4_lut(histogram, 1.2f); + REQUIRE(lut.fingerprint() == fit_bm25_u4_lut(histogram, 1.2f).fingerprint()); + REQUIRE(lut.encode[0] == 0); + REQUIRE(lut.decode[lut.encode[1]] == 1); + for (size_t t = 1; t < 256; ++t) { + REQUIRE(lut.encode[t] > 0); + REQUIRE(lut.encode[t] < 16); + REQUIRE(lut.ends[lut.encode[t]] >= t); + REQUIRE(lut.encode[t] >= lut.encode[t - 1]); + } + REQUIRE_NOTHROW(fit_bm25_u4_lut({}, 0)); + REQUIRE_THROWS(fit_bm25_u4_lut({}, -1)); + std::vector kernels{bm25_accumulate_scalar_u12_u4_lut, get_packed_bm25_kernel()}; +#if defined(__x86_64__) + if (__builtin_cpu_supports("avx2") && __builtin_cpu_supports("fma") && __builtin_cpu_supports("f16c")) { + kernels.push_back(bm25_accumulate_avx2_u12_u4_lut); + } +#endif + std::vector lengths(4096), oracle(4096), actual(4096); + for (size_t i = 0; i < 4096; ++i) lengths[i] = 1 + i % 211; + for (size_t start : {0u, 1u, 2u, 3u}) { + for (size_t n : + {0u, 1u, 2u, 3u, 4u, 7u, 8u, 9u, 15u, 16u, 17u, 23u, 24u, 25u, 31u, 32u, 33u, 63u, 64u, 129u, 4096u}) { + const size_t count = start + n; + std::vector ids(packed12_bytes(count), 0), values(count / 2 + count % 2, 0); + for (size_t i = 0; i < count; ++i) { + const auto id = (i * 17) % 4096; + pack12(ids.data(), i, id); + values[i / 2] |= uint8_t(i % 16) << (4 * (i & 1)); + } + std::fill(oracle.begin(), oracle.end(), .125f); + for (size_t i = start; i < count; ++i) { + const auto id = (i * 17) % 4096; + const double tf = lut.decode[i % 16]; + oracle[id] += float(.7 * 2.2 * tf / (tf + 1.2 * (1 - .75) + 1.2 * .75 / 57 * lengths[id])); + } + for (auto fn : kernels) { + std::fill(actual.begin(), actual.end(), .125f); + const auto maximum = fn(.7f, values.data(), ids.data(), start, n, actual.data(), 1.2f, .75f, 57.f, + lengths.data(), lut.decode.data()); + REQUIRE(std::abs(maximum - (n ? *std::max_element(oracle.begin(), oracle.end()) : 0.f)) < 2e-6f); + for (size_t i = 0; i < 4096; ++i) REQUIRE(std::abs(actual[i] - oracle[i]) < 2e-6f); + } + } + } +} + +TEST_CASE("SINDI BM25 U12 LUT compact public lifecycle and decoded oracle", "[sparse][sindi][u4_u12]") { + using namespace knowhere; + using namespace knowhere::sparse::inverted; + for (const std::string quant : {"u4_lut", "u4_lut_u12", "u4_lut_u16"}) { + const bool u16_ids = quant == "u4_lut_u16"; + for (uint32_t window : + (u16_ids ? std::vector{1024, 4096, 8192, 65535} : std::vector{1024, 4096})) { + std::vector rows(9001); + std::array histogram{}; + for (size_t i = 0; i < rows.size(); ++i) { + const float tf = i == 9000 ? 65535 : 1 + i % 255; + rows[i] = RefineTestRow({{0, tf}, {7, float(1 + i % 13)}, {100000, float(1 + i % 17)}}); + for (size_t j = 0; j < rows[i].size(); ++j) ++histogram[std::min(rows[i][j].val, 255)]; + } + const auto lut = sindi::fit_bm25_u4_lut(histogram, 1.2f); + Json build{{"metric_type", "BM25"}, {"inverted_index_algo", "SINDI"}, + {"quant_type", quant}, {"sindi_window_size", window}, + {"refine", false}, {"bm25_k1", 1.2f}, + {"bm25_b", .75f}, {"bm25_avgdl", 200.f}}; + auto idx = + IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(idx.Build(RefineDataset(rows), build) == Status::success); + auto original = build; + original["quant_type"] = "u8"; + auto baseline = + IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(baseline.Build(RefineDataset(rows), original) == Status::success); + REQUIRE(idx.Size() < baseline.Size()); + std::vector queries{RefineTestRow({{0, .7f}, {7, 1.3f}, {100000, .2f}})}; + Json search{{"metric_type", "BM25"}, {"k", 100}, + {"bm25_k1", 1.2f}, {"bm25_b", .75f}, + {"bm25_avgdl", 200.f}, {"dim_max_score_ratio", 1.05f}}; + std::vector oracle(rows.size()); + for (size_t i = 0; i < rows.size(); ++i) { + double dl = 0, score = 0; + for (size_t j = 0; j < rows[i].size(); ++j) dl += rows[i][j].val; + for (size_t j = 0; j < rows[i].size(); ++j) { + const double tf = lut.decode[lut.encode[std::min(rows[i][j].val, 255)]]; + score += queries[0][j].val * 2.2 * tf / (tf + 1.2 * (1 - .75 + .75 * dl / 200)); + } + oracle[i] = score; + } + auto sorted = oracle; + std::sort(sorted.begin(), sorted.end(), std::greater()); + auto check = [&](auto& index) { + auto result = index.Search(RefineDataset(queries), search, nullptr); + REQUIRE(result.has_value()); + std::set seen; + for (size_t j = 0; j < 100; ++j) { + const auto id = result.value()->GetIds()[j]; + REQUIRE(id >= 0); + REQUIRE(id < int64_t(rows.size())); + REQUIRE(seen.insert(id).second); + REQUIRE(std::abs(result.value()->GetDistance()[j] - oracle[id]) < 3e-6f); + REQUIRE(oracle[id] >= sorted[99] - 3e-6f); + } + std::vector mask((rows.size() + 7) / 8, 255); + auto none = index.Search(RefineDataset(queries), search, BitsetView(mask.data(), rows.size())); + REQUIRE(none.has_value()); + REQUIRE(none.value()->GetIds()[0] == -1); + std::fill(mask.begin(), mask.end(), 0); + std::fill(mask.begin(), mask.begin() + std::min(window / 8, mask.size()), 255); + auto filtered = index.Search(RefineDataset(queries), search, BitsetView(mask.data(), rows.size())); + REQUIRE(filtered.has_value()); + for (size_t j = 0; j < 100; ++j) { + if (window < rows.size()) + REQUIRE(filtered.value()->GetIds()[j] >= window); + else + REQUIRE(filtered.value()->GetIds()[j] == -1); + } + }; + check(idx); + std::vector invalid_rebuild{RefineTestRow({{0, -1}})}; + REQUIRE(idx.Build(RefineDataset(invalid_rebuild), build) != Status::success); + check(idx); + auto mismatched_search = search; + mismatched_search["bm25_avgdl"] = 201.f; + REQUIRE_FALSE(idx.Search(RefineDataset(queries), mismatched_search, nullptr).has_value()); + BinarySet bytes; + REQUIRE(idx.Serialize(bytes) == Status::success); + auto blob = bytes.GetByName(IndexEnum::INDEX_SPARSE_INVERTED_INDEX); + const auto sections = ReadSparseIndexSections(blob); + REQUIRE(FindSection(sections, InvertedIndexSectionType::BM25_U8_OVERFLOWS) == nullptr); + REQUIRE(FindSection(sections, InvertedIndexSectionType::SINDI_REFINEMENT) == nullptr); + const auto* h = FindSection(sections, InvertedIndexSectionType::POSTING_LISTS); + REQUIRE(h != nullptr); + // Three dense terms, each with an odd number of postings. Canonical + // concatenation pays a single odd tail, not one per term/window. + const size_t windows = (rows.size() + window - 1) / window, n = rows.size() * 3; + REQUIRE(h->size == + 68 + 1 + 3 * (4 + windows * 2) + 4 * 4 + (u16_ids ? 2 * n + n / 2 + n % 2 : 2 * n + n % 2)); + uint32_t serialized_quant; + std::memcpy(&serialized_quant, blob->data.get() + 16, 4); + REQUIRE(serialized_quant == static_cast(u16_ids ? InvertedIndexQuantType::BM25_U4_LUT_U16 + : InvertedIndexQuantType::BM25_U4_LUT_U12)); + REQUIRE(std::memcmp(blob->data.get() + h->offset + 24, lut.decode.data(), 16) == 0); + REQUIRE(std::memcmp(blob->data.get() + h->offset + 40, lut.ends.data(), 16) == 0); + auto load = build; + load.erase("quant_type"); + load.erase("sindi_window_size"); + auto restored = + IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(restored.Deserialize(bytes, load) == Status::success); + check(restored); + auto conflict = load; + conflict["quant_type"] = u16_ids ? "u4_lut_u12" : "u4_lut_u16"; + REQUIRE(restored.Deserialize(bytes, conflict) != Status::success); + check(restored); + SparseQuantIndexFile file(blob); + auto mapped = + IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(mapped.DeserializeFromFile(file.path, load) == Status::success); + check(mapped); + auto original_bytes = std::vector(blob->data.get(), blob->data.get() + blob->size); + for (size_t pos : {size_t(h->offset + 12), size_t(h->offset + 16), size_t(h->offset + 20), + size_t(h->offset + 24), size_t(h->offset + 40), size_t(h->offset + h->size - 1)}) { + blob->data[pos] = 255; + REQUIRE(restored.Deserialize(bytes, load) != Status::success); + check(restored); + std::memcpy(blob->data.get(), original_bytes.data(), original_bytes.size()); + } + auto invalid_search = search; + invalid_search["refine_k"] = 2; + REQUIRE_FALSE(idx.Search(RefineDataset(queries), invalid_search, nullptr).has_value()); + } + } +} + +TEST_CASE("SINDI BM25 U12 LUT unsupported requests and invalid input", "[sparse][sindi][u4_u12]") { + using namespace knowhere; + Json build{{"metric_type", "BM25"}, {"inverted_index_algo", "SINDI"}, + {"quant_type", "u4_lut"}, {"sindi_window_size", 4096}, + {"refine", false}, {"bm25_k1", 1.2}, + {"bm25_b", .75}, {"bm25_avgdl", 20}}; + std::vector rows{RefineTestRow({{0, 1}})}; + for (auto change : + {Json{{"refine", true}}, Json{{"metric_type", "IP"}}, Json{{"sindi_window_size", 8192}}, + Json{{"inverted_index_algo", "DAAT_WAND"}}, Json{{"inverted_index_codec", "block_streamvbyte"}}}) { + auto cfg = build; + cfg.update(change); + auto idx = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(idx.Build(RefineDataset(rows), cfg) != Status::success); + } + for (float tf : + {-1.f, -0.f, .5f, 65536.f, std::numeric_limits::infinity(), std::numeric_limits::quiet_NaN()}) { + auto idx = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + std::vector invalid{RefineTestRow({{0, tf}})}; + REQUIRE(idx.Build(RefineDataset(invalid), build) != Status::success); + } + auto growable = + IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX_CC, 11).value(); + REQUIRE(growable.Build(RefineDataset(rows), build) != Status::success); +} + +TEST_CASE("SINDI BM25 U16 LUT kernels and full U16 window", "[sparse][sindi][u4_u16]") { + using namespace knowhere; + using namespace knowhere::sparse::inverted; + using namespace knowhere::sparse::inverted::sindi; + std::array hist{}; + for (size_t t = 1; t < 256; ++t) hist[t] = 256 - t; + const auto lut = fit_bm25_u4_lut(hist, 1.2f); + std::vector kernels{bm25_accumulate_scalar_u16_u4_lut, get_packed_bm25_kernel(true)}; +#if defined(__x86_64__) + if (__builtin_cpu_supports("avx2") && __builtin_cpu_supports("fma") && __builtin_cpu_supports("f16c")) + kernels.push_back(bm25_accumulate_avx2_u16_u4_lut); +#endif + std::vector lengths(65535), oracle(65535), actual(65535); + for (size_t i = 0; i < lengths.size(); ++i) lengths[i] = 1 + i % 211; + for (size_t start : {0u, 1u, 2u, 3u}) { + for (size_t n : {0u, 1u, 7u, 8u, 9u, 15u, 16u, 17u, 23u, 24u, 25u, 31u, 32u, 33u, 129u, 4096u, 65535u}) { + const size_t count = start + n; + // Deliberately unaligned U16 streams model canonical mapped sections. + std::vector storage(2 * count + 1), values(count / 2 + count % 2); + auto* ids = storage.data() + 1; + for (size_t i = 0; i < count; ++i) { + const uint16_t id = (i * 31) % 65535; + std::memcpy(ids + 2 * i, &id, 2); + values[i / 2] |= uint8_t(i % 16) << (4 * (i & 1)); + } + std::fill(oracle.begin(), oracle.end(), .125f); + for (size_t i = start; i < count; ++i) { + const size_t id = (i * 31) % 65535; + const double tf = lut.decode[i % 16]; + oracle[id] += float(.7 * 2.2 * tf / (tf + 1.2 * (1 - .75) + 1.2 * .75 / 57 * lengths[id])); + } + for (auto fn : kernels) { + std::fill(actual.begin(), actual.end(), .125f); + const auto maximum = fn(.7f, values.data(), ids, start, n, actual.data(), 1.2f, .75f, 57, + lengths.data(), lut.decode.data()); + REQUIRE(std::equal(actual.begin(), actual.end(), oracle.begin(), + [](float a, float b) { return std::abs(a - b) < 2e-6f; })); + REQUIRE(std::abs(maximum - (n ? *std::max_element(oracle.begin(), oracle.end()) : 0.f)) < 2e-6f); + } + } + } + // A dense 65535-posting window exercises the U16 count and highest valid ID. + std::vector rows(65535); + for (auto& r : rows) r = RefineTestRow({{0, 1}}); + Json build{{"metric_type", "BM25"}, + {"inverted_index_algo", "SINDI"}, + {"quant_type", "u4_lut_u16"}, + {"sindi_window_size", 65535}, + {"refine", false}, + {"bm25_k1", 1.2f}, + {"bm25_b", .75f}, + {"bm25_avgdl", 1.f}}; + auto idx = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(idx.Build(RefineDataset(rows), build) == Status::success); + BinarySet bytes; + REQUIRE(idx.Serialize(bytes) == Status::success); + auto blob = bytes.GetByName(IndexEnum::INDEX_SPARSE_INVERTED_INDEX); + const auto sections = ReadSparseIndexSections(blob); + const auto* h = FindSection(sections, InvertedIndexSectionType::POSTING_LISTS); + REQUIRE(h != nullptr); + REQUIRE(h->size == 68 + 1 + 4 + 2 + 8 + 2 * rows.size() + (rows.size() + 1) / 2); + auto cfg = build; + cfg.erase("quant_type"); + cfg.erase("sindi_window_size"); + auto loaded = IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX, 11).value(); + REQUIRE(loaded.Deserialize(bytes, cfg) == Status::success); + std::vector queries{RefineTestRow({{0, 1}})}; + Json search{{"metric_type", "BM25"}, {"k", 10}, {"bm25_k1", 1.2f}, {"bm25_b", .75f}, {"bm25_avgdl", 1.f}}; + std::vector mask((rows.size() + 7) / 8, 255); + mask[65534 / 8] &= ~(1u << (65534 % 8)); + auto result = loaded.Search(RefineDataset(queries), search, BitsetView(mask.data(), rows.size())); + REQUIRE(result.has_value()); + REQUIRE(result.value()->GetIds()[0] == 65534); + REQUIRE(std::abs(result.value()->GetDistance()[0] - 1.f) < 2e-6f); + for (size_t j = 1; j < 10; ++j) REQUIRE(result.value()->GetIds()[j] == -1); + for (auto change : + {Json{{"sindi_window_size", 65536}}, Json{{"refine", true}}, Json{{"metric_type", "IP"}}, + Json{{"inverted_index_algo", "DAAT_WAND"}}, Json{{"inverted_index_codec", "block_streamvbyte"}}}) { + auto bad = build; + bad.update(change); + REQUIRE(idx.Build(RefineDataset(rows), bad) != Status::success); + } + auto growable = + IndexFactory::Instance().Create(IndexEnum::INDEX_SPARSE_INVERTED_INDEX_CC, 11).value(); + REQUIRE(growable.Build(RefineDataset(rows), build) != Status::success); +}