Skip to content

Commit c66998a

Browse files
committed
refactor(quantizers): stateful EvpQuantizer, Quantizer concept and span overloads
- Add stateful EvpQuantizer class holding non_zeros, mirroring the ScalarQuantizer API (float/fp16, in-place and returning quantize) - Remove quantize_evp_* free-function wrappers, pybind quantize_batch and the int-based searcher quantizer path; DB and queries now share one quantizer instance with matching settings - Unify searcher.h on a single output_type-based quantize call and drop the local EVPQuantizer struct plus the requires-branch - Add deglib::quantization::Quantizer concept covering all 8 quantize flavors and constrain SearcherImpl/make_searcher on it - Add validated std::span overloads to all 5 quantizer classes, keeping the pointer API for bindings and the hot loop
1 parent 13f10a4 commit c66998a

21 files changed

Lines changed: 602 additions & 173 deletions

File tree

cpp/API.md

Lines changed: 11 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -379,13 +379,17 @@ vector<uint32_t> presort(
379379

380380
// --- Extreme Vector Quantization (EVP) ---
381381

382-
/// Quantize a single vector to packed EVP representation
383-
vector<byte> quantize_evp_single(float* embedding, uint32_t dim, uint32_t non_zeros);
384-
vector<byte> quantize_evp_single(uint16_t* embedding, uint32_t dim, uint32_t non_zeros);
385-
386-
/// Quantize a batch of vectors to packed EVP representation
387-
vector<byte> quantize_evp_batch(float* data, size_t count, uint32_t dim, uint32_t non_zeros, size_t num_threads = 0);
388-
vector<byte> quantize_evp_batch(uint16_t* data, size_t count, uint32_t dim, uint32_t non_zeros, size_t num_threads = 0);
382+
/// Stateful EVP quantizer holding the shared non_zeros setting.
383+
/// Construct once and reuse for database and query quantization.
384+
class EvpQuantizer {
385+
explicit EvpQuantizer(uint32_t non_zeros);
386+
void quantize(float* src, byte* dst, size_t count, uint32_t dim, size_t num_threads = 0);
387+
void quantize(uint16_t* src, byte* dst, size_t count, uint32_t dim, size_t num_threads = 0);
388+
vector<byte> quantize(float* src, size_t count, uint32_t dim, size_t num_threads = 0);
389+
vector<byte> quantize(uint16_t* src, size_t count, uint32_t dim, size_t num_threads = 0);
390+
// Same four flavors with std::span (count derived from span size, validated).
391+
};
392+
EvpQuantizer make_evp_quantizer(uint32_t non_zeros);
389393

390394
// --- MIPS to L2 Space Transformation ---
391395

cpp/deglib/include/deglib/optimization.h

Lines changed: 6 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -63,32 +63,15 @@ inline std::vector<uint32_t> presort(
6363
return result;
6464
}
6565

66-
/**
67-
* Quantize a single FP32 vector using EVP quantization.
68-
*/
69-
inline std::vector<std::byte> quantize_evp_single(const float* embedding, uint32_t dim, uint32_t non_zeros) {
70-
return deglib::quantization::evp::quantize_single(embedding, dim, non_zeros);
71-
}
72-
73-
/**
74-
* Quantize a single FP16 (uint16_t) vector using EVP quantization.
75-
*/
76-
inline std::vector<std::byte> quantize_evp_single(const uint16_t* embedding, uint32_t dim, uint32_t non_zeros) {
77-
return deglib::quantization::evp::quantize_single(embedding, dim, non_zeros);
78-
}
79-
80-
/**
81-
* Quantize a batch of FP32 vectors using EVP quantization.
82-
*/
83-
inline std::vector<std::byte> quantize_evp_batch(const float* data, size_t count, uint32_t dim, uint32_t non_zeros, size_t numThreads = 0) {
84-
return deglib::quantization::evp::quantize_batch(data, count, dim, non_zeros, numThreads);
85-
}
66+
// ========================================================================
67+
// EVP Quantizer Factory Method
68+
// ========================================================================
8669

8770
/**
88-
* Quantize a batch of FP16 (uint16_t) vectors using EVP quantization.
71+
* Make an EvpQuantizer holding the shared non_zeros setting.
8972
*/
90-
inline std::vector<std::byte> quantize_evp_batch(const uint16_t* data, size_t count, uint32_t dim, uint32_t non_zeros, size_t numThreads = 0) {
91-
return deglib::quantization::evp::quantize_batch(data, count, dim, non_zeros, numThreads);
73+
inline deglib::quantization::evp::EvpQuantizer make_evp_quantizer(uint32_t non_zeros) {
74+
return deglib::quantization::evp::EvpQuantizer(non_zeros);
9275
}
9376

9477
// ========================================================================

cpp/deglib/include/deglib/optimization/quantization/evp_quantize.h

Lines changed: 96 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,13 +5,34 @@
55
#include <algorithm>
66
#include <bit>
77
#include <cmath>
8+
#include <cstddef>
89
#include <cstring>
10+
#include <span>
911
#include <stdexcept>
1012
#include <string>
1113
#include <vector>
1214

1315
namespace deglib::quantization::evp {
1416

17+
// Validates span-based quantize inputs, returns the vector count.
18+
inline size_t checked_span_count(size_t src_size, size_t dst_size, uint32_t dim) {
19+
if (dim == 0) {
20+
if (src_size == 0 && dst_size == 0) return 0;
21+
throw std::invalid_argument("quantize: dim must be > 0");
22+
}
23+
if (dim % 8 != 0) {
24+
throw std::invalid_argument("quantize: dim must be divisible by 8");
25+
}
26+
if (src_size % dim != 0) {
27+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
28+
}
29+
const size_t count = src_size / dim;
30+
if (dst_size < count * 2 * (dim / 8)) {
31+
throw std::invalid_argument("quantize: dst span too small for src span and dim");
32+
}
33+
return count;
34+
}
35+
1536
// ============================================================================
1637
// Conversion: fp32 → EVP bytes
1738
// ============================================================================
@@ -457,4 +478,79 @@ inline std::vector<std::byte> quantize_single(const uint16_t* embedding, uint32_
457478
return result;
458479
}
459480

481+
/**
482+
* Stateful EVP quantizer holding the non_zeros parameter.
483+
*
484+
* Construct once and reuse for database and query quantization so both sides
485+
* always share the same non_zeros setting. Mirrors the ScalarQuantizer API.
486+
*/
487+
class EvpQuantizer {
488+
public:
489+
using output_type = std::byte;
490+
uint32_t non_zeros = 0;
491+
492+
EvpQuantizer() = default;
493+
explicit EvpQuantizer(uint32_t nz) : non_zeros(nz) {}
494+
495+
void quantize(const float* src, std::byte* dst, size_t count, uint32_t dim, size_t numThreads = 0) const {
496+
if (count == 0 || dim == 0) return;
497+
if (count == 1) {
498+
quantize_single_into(src, dim, non_zeros, dst);
499+
return;
500+
}
501+
auto tmp = quantize_batch(src, count, dim, non_zeros, numThreads);
502+
std::memcpy(dst, tmp.data(), tmp.size());
503+
}
504+
505+
void quantize(const uint16_t* src_fp16, std::byte* dst, size_t count, uint32_t dim, size_t numThreads = 0) const {
506+
if (count == 0 || dim == 0) return;
507+
if (count == 1) {
508+
quantize_single_into(src_fp16, dim, non_zeros, dst);
509+
return;
510+
}
511+
auto tmp = quantize_batch(src_fp16, count, dim, non_zeros, numThreads);
512+
std::memcpy(dst, tmp.data(), tmp.size());
513+
}
514+
515+
std::vector<std::byte> quantize(const float* src, size_t count, uint32_t dim, size_t numThreads = 0) const {
516+
if (count == 0 || dim == 0) return {};
517+
return quantize_batch(src, count, dim, non_zeros, numThreads);
518+
}
519+
520+
std::vector<std::byte> quantize(const uint16_t* src_fp16, size_t count, uint32_t dim, size_t numThreads = 0) const {
521+
if (count == 0 || dim == 0) return {};
522+
return quantize_batch(src_fp16, count, dim, non_zeros, numThreads);
523+
}
524+
525+
void quantize(std::span<const float> src, std::span<std::byte> dst, uint32_t dim, size_t numThreads = 0) const {
526+
quantize(src.data(), dst.data(), checked_span_count(src.size(), dst.size(), dim), dim, numThreads);
527+
}
528+
529+
void quantize(std::span<const uint16_t> src_fp16, std::span<std::byte> dst, uint32_t dim, size_t numThreads = 0) const {
530+
quantize(src_fp16.data(), dst.data(), checked_span_count(src_fp16.size(), dst.size(), dim), dim, numThreads);
531+
}
532+
533+
std::vector<std::byte> quantize(std::span<const float> src, uint32_t dim, size_t numThreads = 0) const {
534+
if (dim == 0) {
535+
if (src.empty()) return {};
536+
throw std::invalid_argument("quantize: dim must be > 0");
537+
}
538+
if (src.size() % dim != 0) {
539+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
540+
}
541+
return quantize(src.data(), src.size() / dim, dim, numThreads);
542+
}
543+
544+
std::vector<std::byte> quantize(std::span<const uint16_t> src_fp16, uint32_t dim, size_t numThreads = 0) const {
545+
if (dim == 0) {
546+
if (src_fp16.empty()) return {};
547+
throw std::invalid_argument("quantize: dim must be > 0");
548+
}
549+
if (src_fp16.size() % dim != 0) {
550+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
551+
}
552+
return quantize(src_fp16.data(), src_fp16.size() / dim, dim, numThreads);
553+
}
554+
};
555+
460556
} // namespace deglib::quantization::evp
Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,41 @@
1+
#pragma once
2+
3+
// C++20 concept describing the quantizer interface shared by all quantizer
4+
// classes (ScalarQuantizerInt8/Int8PerDim/Uint8/Uint8PerDim, EvpQuantizer).
5+
//
6+
// Every quantizer exposes its packed element type via output_type and four
7+
// quantize flavors per input precision (float / uint16_t fp16):
8+
// pointer in-place, pointer returning, span in-place, span returning.
9+
10+
#include <concepts>
11+
#include <cstddef>
12+
#include <cstdint>
13+
#include <span>
14+
#include <vector>
15+
16+
namespace deglib::quantization {
17+
18+
template <typename Q>
19+
concept Quantizer = requires(
20+
const Q& q,
21+
const float* f32,
22+
const uint16_t* f16,
23+
typename Q::output_type* dst,
24+
std::span<const float> f32_span,
25+
std::span<const uint16_t> f16_span,
26+
std::span<typename Q::output_type> dst_span,
27+
size_t count,
28+
uint32_t dim
29+
) {
30+
typename Q::output_type;
31+
{ q.quantize(f32, dst, count, dim) } -> std::same_as<void>;
32+
{ q.quantize(f16, dst, count, dim) } -> std::same_as<void>;
33+
{ q.quantize(f32, count, dim) } -> std::same_as<std::vector<typename Q::output_type>>;
34+
{ q.quantize(f16, count, dim) } -> std::same_as<std::vector<typename Q::output_type>>;
35+
{ q.quantize(f32_span, dst_span, dim) } -> std::same_as<void>;
36+
{ q.quantize(f16_span, dst_span, dim) } -> std::same_as<void>;
37+
{ q.quantize(f32_span, dim) } -> std::same_as<std::vector<typename Q::output_type>>;
38+
{ q.quantize(f16_span, dim) } -> std::same_as<std::vector<typename Q::output_type>>;
39+
};
40+
41+
} // namespace deglib::quantization

cpp/deglib/include/deglib/optimization/quantization/scalar_quantize.h

Lines changed: 137 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,13 +10,30 @@
1010
#include <format>
1111
#include <limits>
1212
#include <queue>
13+
#include <span>
1314
#include <stdexcept>
1415
#include <string>
1516
#include <utility>
1617
#include <vector>
1718

1819
namespace deglib::quantization::scalar {
1920

21+
// Validates span-based quantize inputs, returns the vector count.
22+
inline size_t checked_span_count(size_t src_size, size_t dst_size, uint32_t dim) {
23+
if (dim == 0) {
24+
if (src_size == 0 && dst_size == 0) return 0;
25+
throw std::invalid_argument("quantize: dim must be > 0");
26+
}
27+
if (src_size % dim != 0) {
28+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
29+
}
30+
const size_t count = src_size / dim;
31+
if (dst_size < count * dim) {
32+
throw std::invalid_argument("quantize: dst span too small for src span and dim");
33+
}
34+
return count;
35+
}
36+
2037
// ============================================================================
2138
// Calibration Helpers (internal percentile / min / max scanning)
2239
// ============================================================================
@@ -281,6 +298,36 @@ class ScalarQuantizerInt8 {
281298
return result;
282299
}
283300

301+
void quantize(std::span<const float> src, std::span<output_type> dst, uint32_t dim, size_t numThreads = 0) const {
302+
quantize(src.data(), dst.data(), checked_span_count(src.size(), dst.size(), dim), dim, numThreads);
303+
}
304+
305+
void quantize(std::span<const uint16_t> src_fp16, std::span<output_type> dst, uint32_t dim, size_t numThreads = 0) const {
306+
quantize(src_fp16.data(), dst.data(), checked_span_count(src_fp16.size(), dst.size(), dim), dim, numThreads);
307+
}
308+
309+
std::vector<output_type> quantize(std::span<const float> src, uint32_t dim, size_t numThreads = 0) const {
310+
if (dim == 0) {
311+
if (src.empty()) return {};
312+
throw std::invalid_argument("quantize: dim must be > 0");
313+
}
314+
if (src.size() % dim != 0) {
315+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
316+
}
317+
return quantize(src.data(), src.size() / dim, dim, numThreads);
318+
}
319+
320+
std::vector<output_type> quantize(std::span<const uint16_t> src_fp16, uint32_t dim, size_t numThreads = 0) const {
321+
if (dim == 0) {
322+
if (src_fp16.empty()) return {};
323+
throw std::invalid_argument("quantize: dim must be > 0");
324+
}
325+
if (src_fp16.size() % dim != 0) {
326+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
327+
}
328+
return quantize(src_fp16.data(), src_fp16.size() / dim, dim, numThreads);
329+
}
330+
284331
private:
285332
inline int8_t transform(float x) const {
286333
float scaled = std::round(x * scale);
@@ -414,6 +461,36 @@ class ScalarQuantizerInt8PerDim {
414461
return result;
415462
}
416463

464+
void quantize(std::span<const float> src, std::span<output_type> dst, uint32_t dim, size_t numThreads = 0) const {
465+
quantize(src.data(), dst.data(), checked_span_count(src.size(), dst.size(), dim), dim, numThreads);
466+
}
467+
468+
void quantize(std::span<const uint16_t> src_fp16, std::span<output_type> dst, uint32_t dim, size_t numThreads = 0) const {
469+
quantize(src_fp16.data(), dst.data(), checked_span_count(src_fp16.size(), dst.size(), dim), dim, numThreads);
470+
}
471+
472+
std::vector<output_type> quantize(std::span<const float> src, uint32_t dim, size_t numThreads = 0) const {
473+
if (dim == 0) {
474+
if (src.empty()) return {};
475+
throw std::invalid_argument("quantize: dim must be > 0");
476+
}
477+
if (src.size() % dim != 0) {
478+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
479+
}
480+
return quantize(src.data(), src.size() / dim, dim, numThreads);
481+
}
482+
483+
std::vector<output_type> quantize(std::span<const uint16_t> src_fp16, uint32_t dim, size_t numThreads = 0) const {
484+
if (dim == 0) {
485+
if (src_fp16.empty()) return {};
486+
throw std::invalid_argument("quantize: dim must be > 0");
487+
}
488+
if (src_fp16.size() % dim != 0) {
489+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
490+
}
491+
return quantize(src_fp16.data(), src_fp16.size() / dim, dim, numThreads);
492+
}
493+
417494
private:
418495
inline int8_t transform(float x, uint32_t d) const {
419496
float scaled = std::round(x * scales[d]);
@@ -536,6 +613,36 @@ class ScalarQuantizerUint8 {
536613
return result;
537614
}
538615

616+
void quantize(std::span<const float> src, std::span<output_type> dst, uint32_t dim, size_t numThreads = 0) const {
617+
quantize(src.data(), dst.data(), checked_span_count(src.size(), dst.size(), dim), dim, numThreads);
618+
}
619+
620+
void quantize(std::span<const uint16_t> src_fp16, std::span<output_type> dst, uint32_t dim, size_t numThreads = 0) const {
621+
quantize(src_fp16.data(), dst.data(), checked_span_count(src_fp16.size(), dst.size(), dim), dim, numThreads);
622+
}
623+
624+
std::vector<output_type> quantize(std::span<const float> src, uint32_t dim, size_t numThreads = 0) const {
625+
if (dim == 0) {
626+
if (src.empty()) return {};
627+
throw std::invalid_argument("quantize: dim must be > 0");
628+
}
629+
if (src.size() % dim != 0) {
630+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
631+
}
632+
return quantize(src.data(), src.size() / dim, dim, numThreads);
633+
}
634+
635+
std::vector<output_type> quantize(std::span<const uint16_t> src_fp16, uint32_t dim, size_t numThreads = 0) const {
636+
if (dim == 0) {
637+
if (src_fp16.empty()) return {};
638+
throw std::invalid_argument("quantize: dim must be > 0");
639+
}
640+
if (src_fp16.size() % dim != 0) {
641+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
642+
}
643+
return quantize(src_fp16.data(), src_fp16.size() / dim, dim, numThreads);
644+
}
645+
539646
private:
540647
inline uint8_t transform(float x) const {
541648
float scaled = std::round((x - min_val) * scale);
@@ -677,6 +784,36 @@ class ScalarQuantizerUint8PerDim {
677784
return result;
678785
}
679786

787+
void quantize(std::span<const float> src, std::span<output_type> dst, uint32_t dim, size_t numThreads = 0) const {
788+
quantize(src.data(), dst.data(), checked_span_count(src.size(), dst.size(), dim), dim, numThreads);
789+
}
790+
791+
void quantize(std::span<const uint16_t> src_fp16, std::span<output_type> dst, uint32_t dim, size_t numThreads = 0) const {
792+
quantize(src_fp16.data(), dst.data(), checked_span_count(src_fp16.size(), dst.size(), dim), dim, numThreads);
793+
}
794+
795+
std::vector<output_type> quantize(std::span<const float> src, uint32_t dim, size_t numThreads = 0) const {
796+
if (dim == 0) {
797+
if (src.empty()) return {};
798+
throw std::invalid_argument("quantize: dim must be > 0");
799+
}
800+
if (src.size() % dim != 0) {
801+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
802+
}
803+
return quantize(src.data(), src.size() / dim, dim, numThreads);
804+
}
805+
806+
std::vector<output_type> quantize(std::span<const uint16_t> src_fp16, uint32_t dim, size_t numThreads = 0) const {
807+
if (dim == 0) {
808+
if (src_fp16.empty()) return {};
809+
throw std::invalid_argument("quantize: dim must be > 0");
810+
}
811+
if (src_fp16.size() % dim != 0) {
812+
throw std::invalid_argument("quantize: src span size must be a multiple of dim");
813+
}
814+
return quantize(src_fp16.data(), src_fp16.size() / dim, dim, numThreads);
815+
}
816+
680817
private:
681818
inline uint8_t transform(float x, uint32_t d) const {
682819
float scaled = std::round((x - mins[d]) * scales[d]);

0 commit comments

Comments
 (0)