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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions include/turboquant/qjl.h
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,15 @@ class QJL {
float estimate_inner_product(const Vec& y, std::span<const int8_t> 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<const int8_t> signs,
float residual_norm) const;

private:
int dim_;
Mat S_;
Expand Down
4 changes: 4 additions & 0 deletions include/turboquant/quantizer_mse.h
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,10 @@ class QuantizerMSE {
QuantizedMSE quantize(const Vec& x) const;
Vec dequantize(const QuantizedMSE& q) const;

// <y, dequantize(q)> given rotated_y = rotation() * y, in O(dim): the
// rotation is orthogonal, so <y, Pi^T y_hat> = <Pi y, y_hat>.
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_; }
Expand Down
13 changes: 13 additions & 0 deletions include/turboquant/quantizer_prod.h
Original file line number Diff line number Diff line change
Expand Up @@ -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_; }

Expand Down
13 changes: 11 additions & 2 deletions src/qjl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -33,11 +33,20 @@ Vec QJL::dequantize(std::span<const int8_t> signs, float residual_norm) const {

float QJL::estimate_inner_product(const Vec& y, std::span<const int8_t> 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<const int8_t> signs,
float residual_norm) const {
// Avoids full dequantization: <y, dequant> = scale * gamma * <S*y, z>
Vec Sy = S_ * y;
float dot = 0.0f;
for (int i = 0; i < dim_; ++i)
dot += Sy(i) * static_cast<float>(signs[i]);
dot += projected_y(i) * static_cast<float>(signs[i]);

return scale_ * residual_norm * dot;
}
Expand Down
8 changes: 8 additions & 0 deletions src/quantizer_mse.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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
17 changes: 15 additions & 2 deletions src/quantizer_prod.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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;
}

Expand Down