Skip to content

Commit 578859a

Browse files
committed
feat: dequantize (#431)
closes #427 Signed-off-by: usamoi <usamoi@outlook.com>
1 parent c833c98 commit 578859a

6 files changed

Lines changed: 93 additions & 8 deletions

File tree

src/datatype/functions_rabitq4.rs

Lines changed: 31 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -12,12 +12,13 @@
1212
//
1313
// Copyright (c) 2025 TensorChord Inc.
1414

15-
use crate::datatype::memory_halfvec::HalfvecInput;
16-
use crate::datatype::memory_rabitq4::Rabitq4Output;
17-
use crate::datatype::memory_vector::VectorInput;
15+
use crate::datatype::memory_halfvec::{HalfvecInput, HalfvecOutput};
16+
use crate::datatype::memory_rabitq4::{Rabitq4Input, Rabitq4Output};
17+
use crate::datatype::memory_vector::{VectorInput, VectorOutput};
1818
use simd::{Floating, f16};
1919
use vector::VectorBorrowed;
2020
use vector::rabitq4::Rabitq4Borrowed;
21+
use vector::vect::VectBorrowed;
2122

2223
#[pgrx::pg_extern(sql = "")]
2324
fn _vchord_vector_quantize_to_rabitq4(vector: VectorInput) -> Rabitq4Output {
@@ -54,3 +55,30 @@ fn _vchord_halfvec_quantize_to_rabitq4(vector: HalfvecInput) -> Rabitq4Output {
5455
&elements,
5556
))
5657
}
58+
59+
#[pgrx::pg_extern(sql = "")]
60+
fn _vchord_rabitq4_dequantize_to_vector(vector: Rabitq4Input) -> VectorOutput {
61+
let vector = vector.as_borrowed();
62+
let scale = vector.sum_of_x2().sqrt() / vector.norm_of_lattice();
63+
let mut result = Vec::with_capacity(vector.dim() as _);
64+
for c in vector.unpacked_code() {
65+
let base = -0.5 * ((1 << 4) - 1) as f32;
66+
result.push((base + c as f32) * scale);
67+
}
68+
rabitq::rotate::rotate_reversed_inplace(&mut result);
69+
VectorOutput::new(VectBorrowed::new(&result))
70+
}
71+
72+
#[pgrx::pg_extern(sql = "")]
73+
fn _vchord_rabitq4_dequantize_to_halfvec(vector: Rabitq4Input) -> HalfvecOutput {
74+
let vector = vector.as_borrowed();
75+
let scale = vector.sum_of_x2().sqrt() / vector.norm_of_lattice();
76+
let mut result = Vec::with_capacity(vector.dim() as _);
77+
for c in vector.unpacked_code() {
78+
let base = -0.5 * ((1 << 4) - 1) as f32;
79+
result.push((base + c as f32) * scale);
80+
}
81+
rabitq::rotate::rotate_reversed_inplace(&mut result);
82+
let result = f16::vector_from_f32(&result);
83+
HalfvecOutput::new(VectBorrowed::new(&result))
84+
}

src/datatype/functions_rabitq8.rs

Lines changed: 31 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -12,12 +12,13 @@
1212
//
1313
// Copyright (c) 2025 TensorChord Inc.
1414

15-
use crate::datatype::memory_halfvec::HalfvecInput;
16-
use crate::datatype::memory_rabitq8::Rabitq8Output;
17-
use crate::datatype::memory_vector::VectorInput;
15+
use crate::datatype::memory_halfvec::{HalfvecInput, HalfvecOutput};
16+
use crate::datatype::memory_rabitq8::{Rabitq8Input, Rabitq8Output};
17+
use crate::datatype::memory_vector::{VectorInput, VectorOutput};
1818
use simd::{Floating, f16};
1919
use vector::VectorBorrowed;
2020
use vector::rabitq8::Rabitq8Borrowed;
21+
use vector::vect::VectBorrowed;
2122

2223
#[pgrx::pg_extern(sql = "")]
2324
fn _vchord_vector_quantize_to_rabitq8(vector: VectorInput) -> Rabitq8Output {
@@ -54,3 +55,30 @@ fn _vchord_halfvec_quantize_to_rabitq8(vector: HalfvecInput) -> Rabitq8Output {
5455
&elements,
5556
))
5657
}
58+
59+
#[pgrx::pg_extern(sql = "")]
60+
fn _vchord_rabitq8_dequantize_to_vector(vector: Rabitq8Input) -> VectorOutput {
61+
let vector = vector.as_borrowed();
62+
let scale = vector.sum_of_x2().sqrt() / vector.norm_of_lattice();
63+
let mut result = Vec::with_capacity(vector.dim() as _);
64+
for c in vector.unpacked_code() {
65+
let base = -0.5 * ((1 << 8) - 1) as f32;
66+
result.push((base + c as f32) * scale);
67+
}
68+
rabitq::rotate::rotate_reversed_inplace(&mut result);
69+
VectorOutput::new(VectBorrowed::new(&result))
70+
}
71+
72+
#[pgrx::pg_extern(sql = "")]
73+
fn _vchord_rabitq8_dequantize_to_halfvec(vector: Rabitq8Input) -> HalfvecOutput {
74+
let vector = vector.as_borrowed();
75+
let scale = vector.sum_of_x2().sqrt() / vector.norm_of_lattice();
76+
let mut result = Vec::with_capacity(vector.dim() as _);
77+
for c in vector.unpacked_code() {
78+
let base = -0.5 * ((1 << 8) - 1) as f32;
79+
result.push((base + c as f32) * scale);
80+
}
81+
rabitq::rotate::rotate_reversed_inplace(&mut result);
82+
let result = f16::vector_from_f32(&result);
83+
HalfvecOutput::new(VectBorrowed::new(&result))
84+
}

src/datatype/memory_halfvec.rs

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -102,7 +102,6 @@ impl HalfvecOutput {
102102
}
103103
Self(q)
104104
}
105-
#[expect(dead_code)]
106105
pub fn new(vector: VectBorrowed<'_, f16>) -> Self {
107106
unsafe {
108107
let slice = vector.slice();

src/datatype/memory_vector.rs

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -101,7 +101,6 @@ impl VectorOutput {
101101
}
102102
Self(q)
103103
}
104-
#[expect(dead_code)]
105104
pub fn new(vector: VectBorrowed<'_, f32>) -> Self {
106105
unsafe {
107106
let slice = vector.slice();

src/sql/finalize.sql

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -206,12 +206,24 @@ IMMUTABLE STRICT PARALLEL SAFE LANGUAGE c AS 'MODULE_PATHNAME', '_vchord_vector_
206206
CREATE FUNCTION quantize_to_rabitq8(halfvec) RETURNS rabitq8
207207
IMMUTABLE STRICT PARALLEL SAFE LANGUAGE c AS 'MODULE_PATHNAME', '_vchord_halfvec_quantize_to_rabitq8_wrapper';
208208

209+
CREATE FUNCTION dequantize_to_vector(rabitq8) RETURNS vector
210+
IMMUTABLE STRICT PARALLEL SAFE LANGUAGE c AS 'MODULE_PATHNAME', '_vchord_rabitq8_dequantize_to_vector_wrapper';
211+
212+
CREATE FUNCTION dequantize_to_halfvec(rabitq8) RETURNS halfvec
213+
IMMUTABLE STRICT PARALLEL SAFE LANGUAGE c AS 'MODULE_PATHNAME', '_vchord_rabitq8_dequantize_to_halfvec_wrapper';
214+
209215
CREATE FUNCTION quantize_to_rabitq4(vector) RETURNS rabitq4
210216
IMMUTABLE STRICT PARALLEL SAFE LANGUAGE c AS 'MODULE_PATHNAME', '_vchord_vector_quantize_to_rabitq4_wrapper';
211217

212218
CREATE FUNCTION quantize_to_rabitq4(halfvec) RETURNS rabitq4
213219
IMMUTABLE STRICT PARALLEL SAFE LANGUAGE c AS 'MODULE_PATHNAME', '_vchord_halfvec_quantize_to_rabitq4_wrapper';
214220

221+
CREATE FUNCTION dequantize_to_vector(rabitq4) RETURNS vector
222+
IMMUTABLE STRICT PARALLEL SAFE LANGUAGE c AS 'MODULE_PATHNAME', '_vchord_rabitq4_dequantize_to_vector_wrapper';
223+
224+
CREATE FUNCTION dequantize_to_halfvec(rabitq4) RETURNS halfvec
225+
IMMUTABLE STRICT PARALLEL SAFE LANGUAGE c AS 'MODULE_PATHNAME', '_vchord_rabitq4_dequantize_to_halfvec_wrapper';
226+
215227
CREATE FUNCTION vchordrq_sampled_values(regclass) RETURNS SETOF TEXT
216228
STRICT LANGUAGE c AS 'MODULE_PATHNAME', '_vchordrq_sampled_values_wrapper';
217229

tests/general/dequantize.slt

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,19 @@
1+
query I
2+
SELECT dequantize_to_vector(quantize_to_rabitq8('[1,2,3,4,5,6,7,8]'::vector)) <-> '[1,2,3,4,5,6,7,8]'::vector < 0.07;
3+
----
4+
t
5+
6+
query I
7+
SELECT dequantize_to_halfvec(quantize_to_rabitq8('[1,2,3,4,5,6,7,8]'::halfvec)) <-> '[1,2,3,4,5,6,7,8]'::halfvec < 0.07;
8+
----
9+
t
10+
11+
query I
12+
SELECT dequantize_to_vector(quantize_to_rabitq4('[1,2,3,4,5,6,7,8]'::vector)) <-> '[1,2,3,4,5,6,7,8]'::vector < 1.00;
13+
----
14+
t
15+
16+
query I
17+
SELECT dequantize_to_halfvec(quantize_to_rabitq4('[1,2,3,4,5,6,7,8]'::halfvec)) <-> '[1,2,3,4,5,6,7,8]'::halfvec < 1.00;
18+
----
19+
t

0 commit comments

Comments
 (0)