From 4d99b00b9f57c616c0b662ec8b401a318d36216a Mon Sep 17 00:00:00 2001 From: Andrew Mikhail Date: Sun, 13 Sep 2026 19:26:57 -0700 Subject: [PATCH] feat: prepared-query inner product estimate, O(dim) per code QuantizerProd::estimate_inner_product did two O(dim^2) products for every code: S * y in QJL, and Pi^T * y_hat in the MSE dequantize. Both depend only on the query once the MSE term is rewritten as = (Pi is orthogonal), so scoring a query against N codes paid N * O(dim^2). - QuantizerProd::prepare_query(y) computes Pi*y and S*y once. - QuantizerProd::estimate_inner_product(PreparedQuery, q) is O(dim) per code. - QJL::project / estimate_inner_product_projected and QuantizerMSE::inner_product_rotated are the building blocks. - The existing estimate_inner_product(y, q) now prepares and delegates, so the two entry points produce identical values. turboquant_example output is byte-identical to main (all estimates, 6 d.p.). Measured through recall's MediaIndex (100 entries, d=1536, -c opt): search 23.1 ms -> 0.56 ms on Apple Silicon, 204 ms -> 3.0 ms on Jetson Orin Nano. Co-Authored-By: Claude Opus 5 --- include/turboquant/qjl.h | 9 +++++++++ include/turboquant/quantizer_mse.h | 4 ++++ include/turboquant/quantizer_prod.h | 13 +++++++++++++ src/qjl.cpp | 13 +++++++++++-- src/quantizer_mse.cpp | 8 ++++++++ src/quantizer_prod.cpp | 17 +++++++++++++++-- 6 files changed, 60 insertions(+), 4 deletions(-) diff --git a/include/turboquant/qjl.h b/include/turboquant/qjl.h index ed2aa50..5d70207 100644 --- a/include/turboquant/qjl.h +++ b/include/turboquant/qjl.h @@ -16,6 +16,15 @@ class QJL { float estimate_inner_product(const Vec& y, std::span signs, float residual_norm) const; + // Query-only half of estimate_inner_product: S * y, O(dim^2). Compute once + // per query and reuse it across codes. + Vec project(const Vec& y) const; + + // estimate_inner_product given project(y): O(dim) per code. + float estimate_inner_product_projected(const Vec& projected_y, + std::span signs, + float residual_norm) const; + private: int dim_; Mat S_; diff --git a/include/turboquant/quantizer_mse.h b/include/turboquant/quantizer_mse.h index 9050447..4dd3a0f 100644 --- a/include/turboquant/quantizer_mse.h +++ b/include/turboquant/quantizer_mse.h @@ -18,6 +18,10 @@ class QuantizerMSE { QuantizedMSE quantize(const Vec& x) const; Vec dequantize(const QuantizedMSE& q) const; + // given rotated_y = rotation() * y, in O(dim): the + // rotation is orthogonal, so = . + float inner_product_rotated(const Vec& rotated_y, const QuantizedMSE& q) const; + int dim() const { return dim_; } int bitwidth() const { return codebook_.bitwidth; } const ScalarCodebook& codebook() const { return codebook_; } diff --git a/include/turboquant/quantizer_prod.h b/include/turboquant/quantizer_prod.h index 6f2ab8d..2902acf 100644 --- a/include/turboquant/quantizer_prod.h +++ b/include/turboquant/quantizer_prod.h @@ -16,10 +16,23 @@ class QuantizerProd { public: QuantizerProd(int dim, int bitwidth, std::mt19937& rng); + // The query-only work of estimate_inner_product — the rotation and the QJL + // projection of y, both O(dim^2). Prepare once per query, then score many + // codes at O(dim) each. + struct PreparedQuery { + Vec rotated; // Pi * y + Vec projected; // S * y + }; + QuantizedProd quantize(const Vec& x) const; Vec dequantize(const QuantizedProd& q) const; float estimate_inner_product(const Vec& y, const QuantizedProd& q) const; + PreparedQuery prepare_query(const Vec& y) const; + // Same arithmetic as estimate_inner_product(y, q) for p = prepare_query(y), + // so the two agree exactly. + float estimate_inner_product(const PreparedQuery& p, const QuantizedProd& q) const; + int dim() const { return dim_; } int bitwidth() const { return total_bitwidth_; } diff --git a/src/qjl.cpp b/src/qjl.cpp index d2e165f..e84364d 100644 --- a/src/qjl.cpp +++ b/src/qjl.cpp @@ -33,11 +33,20 @@ Vec QJL::dequantize(std::span signs, float residual_norm) const { float QJL::estimate_inner_product(const Vec& y, std::span signs, float residual_norm) const { + return estimate_inner_product_projected(project(y), signs, residual_norm); +} + +Vec QJL::project(const Vec& y) const { + return S_ * y; +} + +float QJL::estimate_inner_product_projected(const Vec& projected_y, + std::span signs, + float residual_norm) const { // Avoids full dequantization: = scale * gamma * - Vec Sy = S_ * y; float dot = 0.0f; for (int i = 0; i < dim_; ++i) - dot += Sy(i) * static_cast(signs[i]); + dot += projected_y(i) * static_cast(signs[i]); return scale_ * residual_norm * dot; } diff --git a/src/quantizer_mse.cpp b/src/quantizer_mse.cpp index 6d49276..54f9606 100644 --- a/src/quantizer_mse.cpp +++ b/src/quantizer_mse.cpp @@ -43,4 +43,12 @@ Vec QuantizerMSE::dequantize(const QuantizedMSE& q) const { return q.norm * (Pi_.transpose() * y_hat); } +float QuantizerMSE::inner_product_rotated(const Vec& rotated_y, + const QuantizedMSE& q) const { + float dot = 0.0f; + for (int j = 0; j < dim_; ++j) + dot += rotated_y(j) * codebook_.centroids[q.indices[j]]; + return q.norm * dot; +} + } // namespace turboquant \ No newline at end of file diff --git a/src/quantizer_prod.cpp b/src/quantizer_prod.cpp index d6949c2..506f794 100644 --- a/src/quantizer_prod.cpp +++ b/src/quantizer_prod.cpp @@ -42,8 +42,21 @@ Vec QuantizerProd::dequantize(const QuantizedProd& q) const { } float QuantizerProd::estimate_inner_product(const Vec& y, const QuantizedProd& q) const { - float ip_mse = y.dot(mse_quantizer_.dequantize(q.mse_part)); - float ip_qjl = qjl_.estimate_inner_product(y, q.qjl_signs, q.residual_norm); + return estimate_inner_product(prepare_query(y), q); +} + +QuantizerProd::PreparedQuery QuantizerProd::prepare_query(const Vec& y) const { + return PreparedQuery{ + .rotated = mse_quantizer_.rotation() * y, + .projected = qjl_.project(y), + }; +} + +float QuantizerProd::estimate_inner_product(const PreparedQuery& p, + const QuantizedProd& q) const { + float ip_mse = mse_quantizer_.inner_product_rotated(p.rotated, q.mse_part); + float ip_qjl = qjl_.estimate_inner_product_projected(p.projected, q.qjl_signs, + q.residual_norm); return ip_mse + ip_qjl; }