Skip to content

Commit 7ca2929

Browse files
committed
feat: add RVV batch distance operators
1 parent 2f0b573 commit 7ca2929

7 files changed

Lines changed: 581 additions & 0 deletions

src/ailego/math_batch/euclidean_distance_batch_dispatch.cc

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -87,9 +87,39 @@ void compute_one_to_many_squared_euclidean_avx2_fp16_12(
8787
// float *results);
8888
#endif
8989

90+
#if defined(__riscv_vector)
91+
void compute_one_to_many_squared_euclidean_rvv_fp32_1(
92+
const float *query, const float **ptrs,
93+
std::array<const float *, 1> &prefetch_ptrs, size_t dimensionality,
94+
float *results);
95+
96+
void compute_one_to_many_squared_euclidean_rvv_fp32_12(
97+
const float *query, const float **ptrs,
98+
std::array<const float *, 12> &prefetch_ptrs, size_t dimensionality,
99+
float *results);
100+
#endif
101+
102+
#if defined(__riscv_zvfh)
103+
void compute_one_to_many_squared_euclidean_rvv_fp16_1(
104+
const ailego::Float16 *query, const ailego::Float16 **ptrs,
105+
std::array<const ailego::Float16 *, 1> &prefetch_ptrs,
106+
size_t dimensionality, float *results);
107+
108+
void compute_one_to_many_squared_euclidean_rvv_fp16_12(
109+
const ailego::Float16 *query, const ailego::Float16 **ptrs,
110+
std::array<const ailego::Float16 *, 12> &prefetch_ptrs,
111+
size_t dimensionality, float *results);
112+
#endif
113+
90114
void SquaredEuclideanDistanceBatchImpl<float, 1>::compute_one_to_many(
91115
const ValueType *query, const ValueType **ptrs,
92116
std::array<const ValueType *, 1> &prefetch_ptrs, size_t dim, float *sums) {
117+
#if defined(__riscv_vector)
118+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_VECTOR) {
119+
return compute_one_to_many_squared_euclidean_rvv_fp32_1(
120+
query, ptrs, prefetch_ptrs, dim, sums);
121+
}
122+
#endif
93123
#if defined(__AVX2__)
94124
if (zvec::ailego::internal::CpuFeatures::static_flags_.AVX2) {
95125
return compute_one_to_many_squared_euclidean_avx2_fp32_1(
@@ -104,6 +134,12 @@ void SquaredEuclideanDistanceBatchImpl<ailego::Float16, 1>::compute_one_to_many(
104134
const ailego::Float16 *query, const ailego::Float16 **ptrs,
105135
std::array<const ailego::Float16 *, 1> &prefetch_ptrs, size_t dim,
106136
float *sums) {
137+
#if defined(__riscv_zvfh)
138+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_ZVFH) {
139+
return compute_one_to_many_squared_euclidean_rvv_fp16_1(
140+
query, ptrs, prefetch_ptrs, dim, sums);
141+
}
142+
#endif
107143
#if defined(__AVX512FP16__)
108144
if (zvec::ailego::internal::CpuFeatures::static_flags_.AVX512_FP16) {
109145
return compute_one_to_many_squared_euclidean_avx512fp16_fp16_1(
@@ -129,6 +165,12 @@ void SquaredEuclideanDistanceBatchImpl<ailego::Float16, 1>::compute_one_to_many(
129165
void SquaredEuclideanDistanceBatchImpl<float, 12>::compute_one_to_many(
130166
const float *query, const float **ptrs,
131167
std::array<const float *, 12> &prefetch_ptrs, size_t dim, float *sums) {
168+
#if defined(__riscv_vector)
169+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_VECTOR) {
170+
return compute_one_to_many_squared_euclidean_rvv_fp32_12(
171+
query, ptrs, prefetch_ptrs, dim, sums);
172+
}
173+
#endif
132174
#if defined(__AVX512F__)
133175
if (zvec::ailego::internal::CpuFeatures::static_flags_.AVX512F) {
134176
return compute_one_to_many_squared_euclidean_avx512f_fp32_12(
@@ -151,6 +193,12 @@ void SquaredEuclideanDistanceBatchImpl<ailego::Float16, 12>::
151193
const ailego::Float16 **ptrs,
152194
std::array<const ailego::Float16 *, 12> &prefetch_ptrs,
153195
size_t dim, float *sums) {
196+
#if defined(__riscv_zvfh)
197+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_ZVFH) {
198+
return compute_one_to_many_squared_euclidean_rvv_fp16_12(
199+
query, ptrs, prefetch_ptrs, dim, sums);
200+
}
201+
#endif
154202
#if defined(__AVX512FP16__)
155203
if (zvec::ailego::internal::CpuFeatures::static_flags_.AVX512_FP16) {
156204
return compute_one_to_many_squared_euclidean_avx512fp16_fp16_12(
Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,97 @@
1+
// Copyright 2025-present the zvec project
2+
//
3+
// Licensed under the Apache License, Version 2.0 (the "License");
4+
// you may not use this file except in compliance with the License.
5+
// You may obtain a copy of the License at
6+
//
7+
// http://www.apache.org/licenses/LICENSE-2.0
8+
//
9+
// Unless required by applicable law or agreed to in writing, software
10+
// distributed under the License is distributed on an "AS IS" BASIS,
11+
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
// See the License for the specific language governing permissions and
13+
// limitations under the License.
14+
15+
#include <array>
16+
#include <zvec/ailego/internal/platform.h>
17+
#include <zvec/ailego/utility/type_helper.h>
18+
19+
namespace zvec::ailego::DistanceBatch {
20+
21+
#if defined(__riscv_zvfh)
22+
23+
void compute_one_to_many_squared_euclidean_rvv_fp16_1(
24+
const ailego::Float16 *query, const ailego::Float16 **ptrs,
25+
std::array<const ailego::Float16 *, 1> &prefetch_ptrs,
26+
size_t dimensionality, float *results) {
27+
const _Float16 *q_ptr = reinterpret_cast<const _Float16 *>(query);
28+
const _Float16 *m_ptr = reinterpret_cast<const _Float16 *>(ptrs[0]);
29+
30+
const size_t vlmax = __riscv_vsetvlmax_e16m4();
31+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
32+
33+
size_t dim = 0;
34+
while (dim < dimensionality) {
35+
size_t vl = __riscv_vsetvl_e16m4(dimensionality - dim);
36+
37+
vfloat16m4_t q = __riscv_vle16_v_f16m4(q_ptr + dim, vl);
38+
vfloat16m4_t m = __riscv_vle16_v_f16m4(m_ptr + dim, vl);
39+
40+
// Widen the fp16 difference to fp32 (matches the scalar/AVX reference,
41+
// which converts to fp32 before subtracting), then square-accumulate.
42+
vfloat32m8_t diff = __riscv_vfwsub_vv_f32m8(q, m, vl);
43+
acc = __riscv_vfmacc_vv_f32m8_tu(acc, diff, diff, vl);
44+
45+
if (prefetch_ptrs[0] != nullptr) {
46+
ailego_prefetch(prefetch_ptrs[0] + dim);
47+
}
48+
49+
dim += vl;
50+
}
51+
52+
vfloat32m1_t zero = __riscv_vfmv_v_f_f32m1(0.0f, 1);
53+
vfloat32m1_t red = __riscv_vfredusum_vs_f32m8_f32m1(acc, zero, vlmax);
54+
results[0] = __riscv_vfmv_f_s_f32m1_f32(red);
55+
}
56+
57+
void compute_one_to_many_squared_euclidean_rvv_fp16_12(
58+
const ailego::Float16 *query, const ailego::Float16 **ptrs,
59+
std::array<const ailego::Float16 *, 12> &prefetch_ptrs,
60+
size_t dimensionality, float *results) {
61+
const _Float16 *q_ptr = reinterpret_cast<const _Float16 *>(query);
62+
63+
const size_t vlmax = __riscv_vsetvlmax_e16m4();
64+
65+
for (size_t channel = 0; channel < 12; ++channel) {
66+
const _Float16 *m_ptr = reinterpret_cast<const _Float16 *>(ptrs[channel]);
67+
68+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
69+
70+
size_t dim = 0;
71+
while (dim < dimensionality) {
72+
size_t vl = __riscv_vsetvl_e16m4(dimensionality - dim);
73+
74+
vfloat16m4_t q = __riscv_vle16_v_f16m4(q_ptr + dim, vl);
75+
vfloat16m4_t m = __riscv_vle16_v_f16m4(m_ptr + dim, vl);
76+
77+
// Widen the fp16 difference to fp32 (matches the scalar/AVX reference,
78+
// which converts to fp32 before subtracting), then square-accumulate.
79+
vfloat32m8_t diff = __riscv_vfwsub_vv_f32m8(q, m, vl);
80+
acc = __riscv_vfmacc_vv_f32m8_tu(acc, diff, diff, vl);
81+
82+
if (prefetch_ptrs[channel] != nullptr) {
83+
ailego_prefetch(prefetch_ptrs[channel] + dim);
84+
}
85+
86+
dim += vl;
87+
}
88+
89+
vfloat32m1_t zero = __riscv_vfmv_v_f_f32m1(0.0f, 1);
90+
vfloat32m1_t red = __riscv_vfredusum_vs_f32m8_f32m1(acc, zero, vlmax);
91+
results[channel] = __riscv_vfmv_f_s_f32m1_f32(red);
92+
}
93+
}
94+
95+
#endif
96+
97+
} // namespace zvec::ailego::DistanceBatch
Lines changed: 91 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,91 @@
1+
// Copyright 2025-present the zvec project
2+
//
3+
// Licensed under the Apache License, Version 2.0 (the "License");
4+
// you may not use this file except in compliance with the License.
5+
// You may obtain a copy of the License at
6+
//
7+
// http://www.apache.org/licenses/LICENSE-2.0
8+
//
9+
// Unless required by applicable law or agreed to in writing, software
10+
// distributed under the License is distributed on an "AS IS" BASIS,
11+
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
// See the License for the specific language governing permissions and
13+
// limitations under the License.
14+
15+
#include <array>
16+
#include <zvec/ailego/internal/platform.h>
17+
18+
namespace zvec::ailego::DistanceBatch {
19+
20+
#if defined(__riscv_vector)
21+
22+
void compute_one_to_many_squared_euclidean_rvv_fp32_1(
23+
const float *query, const float **ptrs,
24+
std::array<const float *, 1> &prefetch_ptrs, size_t dimensionality,
25+
float *results) {
26+
const float *q_ptr = query;
27+
const float *m_ptr = ptrs[0];
28+
29+
const size_t vlmax = __riscv_vsetvlmax_e32m8();
30+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
31+
32+
size_t dim = 0;
33+
while (dim < dimensionality) {
34+
size_t vl = __riscv_vsetvl_e32m8(dimensionality - dim);
35+
36+
vfloat32m8_t q = __riscv_vle32_v_f32m8(q_ptr + dim, vl);
37+
vfloat32m8_t m = __riscv_vle32_v_f32m8(m_ptr + dim, vl);
38+
39+
vfloat32m8_t diff = __riscv_vfsub_vv_f32m8(q, m, vl);
40+
acc = __riscv_vfmacc_vv_f32m8_tu(acc, diff, diff, vl);
41+
42+
if (prefetch_ptrs[0] != nullptr) {
43+
ailego_prefetch(prefetch_ptrs[0] + dim);
44+
}
45+
46+
dim += vl;
47+
}
48+
49+
vfloat32m1_t zero = __riscv_vfmv_v_f_f32m1(0.0f, 1);
50+
vfloat32m1_t red = __riscv_vfredusum_vs_f32m8_f32m1(acc, zero, vlmax);
51+
results[0] = __riscv_vfmv_f_s_f32m1_f32(red);
52+
}
53+
54+
void compute_one_to_many_squared_euclidean_rvv_fp32_12(
55+
const float *query, const float **ptrs,
56+
std::array<const float *, 12> &prefetch_ptrs, size_t dimensionality,
57+
float *results) {
58+
const float *q_ptr = query;
59+
60+
const size_t vlmax = __riscv_vsetvlmax_e32m8();
61+
62+
for (size_t channel = 0; channel < 12; ++channel) {
63+
const float *m_ptr = ptrs[channel];
64+
vfloat32m8_t acc = __riscv_vfmv_v_f_f32m8(0.0f, vlmax);
65+
66+
size_t dim = 0;
67+
while (dim < dimensionality) {
68+
size_t vl = __riscv_vsetvl_e32m8(dimensionality - dim);
69+
70+
vfloat32m8_t q = __riscv_vle32_v_f32m8(q_ptr + dim, vl);
71+
vfloat32m8_t m = __riscv_vle32_v_f32m8(m_ptr + dim, vl);
72+
73+
vfloat32m8_t diff = __riscv_vfsub_vv_f32m8(q, m, vl);
74+
acc = __riscv_vfmacc_vv_f32m8_tu(acc, diff, diff, vl);
75+
76+
if (prefetch_ptrs[channel] != nullptr) {
77+
ailego_prefetch(prefetch_ptrs[channel] + dim);
78+
}
79+
80+
dim += vl;
81+
}
82+
83+
vfloat32m1_t zero = __riscv_vfmv_v_f_f32m1(0.0f, 1);
84+
vfloat32m1_t red = __riscv_vfredusum_vs_f32m8_f32m1(acc, zero, vlmax);
85+
results[channel] = __riscv_vfmv_f_s_f32m1_f32(red);
86+
}
87+
}
88+
89+
#endif
90+
91+
} // namespace zvec::ailego::DistanceBatch

src/ailego/math_batch/inner_product_distance_batch_dispatch.cc

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,9 +93,51 @@ void compute_one_to_many_inner_product_avx2_int8_12(
9393
float *results);
9494
#endif
9595

96+
#if defined(__riscv_vector)
97+
void compute_one_to_many_inner_product_rvv_fp32_1(
98+
const float *query, const float **ptrs,
99+
std::array<const float *, 1> &prefetch_ptrs, size_t dimensionality,
100+
float *results);
101+
102+
void compute_one_to_many_inner_product_rvv_fp32_12(
103+
const float *query, const float **ptrs,
104+
std::array<const float *, 12> &prefetch_ptrs, size_t dimensionality,
105+
float *results);
106+
#endif
107+
108+
#if defined(__riscv_zvfh)
109+
void compute_one_to_many_inner_product_rvv_fp16_1(
110+
const ailego::Float16 *query, const ailego::Float16 **ptrs,
111+
std::array<const ailego::Float16 *, 1> &prefetch_ptrs,
112+
size_t dimensionality, float *results);
113+
114+
void compute_one_to_many_inner_product_rvv_fp16_12(
115+
const ailego::Float16 *query, const ailego::Float16 **ptrs,
116+
std::array<const ailego::Float16 *, 12> &prefetch_ptrs,
117+
size_t dimensionality, float *results);
118+
#endif
119+
120+
#if defined(__riscv_vector)
121+
void compute_one_to_many_inner_product_rvv_int8_1(
122+
const int8_t *query, const int8_t **ptrs,
123+
std::array<const int8_t *, 1> &prefetch_ptrs, size_t dimensionality,
124+
float *results);
125+
126+
void compute_one_to_many_inner_product_rvv_int8_12(
127+
const int8_t *query, const int8_t **ptrs,
128+
std::array<const int8_t *, 12> &prefetch_ptrs, size_t dimensionality,
129+
float *results);
130+
#endif
131+
96132
void InnerProductDistanceBatchImpl<float, 1>::compute_one_to_many(
97133
const ValueType *query, const ValueType **ptrs,
98134
std::array<const ValueType *, 1> &prefetch_ptrs, size_t dim, float *sums) {
135+
#if defined(__riscv_vector)
136+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_VECTOR) {
137+
return compute_one_to_many_inner_product_rvv_fp32_1(
138+
query, ptrs, prefetch_ptrs, dim, sums);
139+
}
140+
#endif
99141
#if defined(__AVX2__)
100142
if (zvec::ailego::internal::CpuFeatures::static_flags_.AVX2) {
101143
return compute_one_to_many_inner_product_avx2_fp32_1(
@@ -110,6 +152,12 @@ void InnerProductDistanceBatchImpl<ailego::Float16, 1>::compute_one_to_many(
110152
const ailego::Float16 *query, const ailego::Float16 **ptrs,
111153
std::array<const ailego::Float16 *, 1> &prefetch_ptrs, size_t dim,
112154
float *sums) {
155+
#if defined(__riscv_zvfh)
156+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_ZVFH) {
157+
return compute_one_to_many_inner_product_rvv_fp16_1(
158+
query, ptrs, prefetch_ptrs, dim, sums);
159+
}
160+
#endif
113161
#if defined(__AVX512FP16__)
114162
if (zvec::ailego::internal::CpuFeatures::static_flags_.AVX512_FP16) {
115163
return compute_one_to_many_inner_product_avx512fp16_fp16_1(
@@ -135,6 +183,12 @@ void InnerProductDistanceBatchImpl<ailego::Float16, 1>::compute_one_to_many(
135183
void InnerProductDistanceBatchImpl<int8_t, 1>::compute_one_to_many(
136184
const int8_t *query, const int8_t **ptrs,
137185
std::array<const int8_t *, 1> &prefetch_ptrs, size_t dim, float *sums) {
186+
#if defined(__riscv_vector)
187+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_VECTOR) {
188+
return compute_one_to_many_inner_product_rvv_int8_1(
189+
query, ptrs, prefetch_ptrs, dim, sums);
190+
}
191+
#endif
138192
// #if defined(__AVX512BW__) // TODO: this version is problematic
139193
// return compute_one_to_many_avx512_int8<ValueType, BatchSize>(
140194
// query, ptrs, prefetch_ptrs, dim, sums);
@@ -167,6 +221,12 @@ InnerProductDistanceBatchImpl<int8_t, 1>::GetQueryPreprocessFunc() {
167221
void InnerProductDistanceBatchImpl<float, 12>::compute_one_to_many(
168222
const ValueType *query, const ValueType **ptrs,
169223
std::array<const ValueType *, 12> &prefetch_ptrs, size_t dim, float *sums) {
224+
#if defined(__riscv_vector)
225+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_VECTOR) {
226+
return compute_one_to_many_inner_product_rvv_fp32_12(
227+
query, ptrs, prefetch_ptrs, dim, sums);
228+
}
229+
#endif
170230
#if defined(__AVX2__)
171231
if (zvec::ailego::internal::CpuFeatures::static_flags_.AVX2) {
172232
return compute_one_to_many_inner_product_avx2_fp32_12(
@@ -181,6 +241,12 @@ void InnerProductDistanceBatchImpl<ailego::Float16, 12>::compute_one_to_many(
181241
const ailego::Float16 *query, const ailego::Float16 **ptrs,
182242
std::array<const ailego::Float16 *, 12> &prefetch_ptrs, size_t dim,
183243
float *sums) {
244+
#if defined(__riscv_zvfh)
245+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_ZVFH) {
246+
return compute_one_to_many_inner_product_rvv_fp16_12(
247+
query, ptrs, prefetch_ptrs, dim, sums);
248+
}
249+
#endif
184250
#if defined(__AVX512FP16__)
185251
if (zvec::ailego::internal::CpuFeatures::static_flags_.AVX512_FP16) {
186252
return compute_one_to_many_inner_product_avx512fp16_fp16_12(
@@ -206,6 +272,12 @@ void InnerProductDistanceBatchImpl<ailego::Float16, 12>::compute_one_to_many(
206272
void InnerProductDistanceBatchImpl<int8_t, 12>::compute_one_to_many(
207273
const int8_t *query, const int8_t **ptrs,
208274
std::array<const int8_t *, 12> &prefetch_ptrs, size_t dim, float *sums) {
275+
#if defined(__riscv_vector)
276+
if (zvec::ailego::internal::CpuFeatures::static_flags_.RISCV_VECTOR) {
277+
return compute_one_to_many_inner_product_rvv_int8_12(
278+
query, ptrs, prefetch_ptrs, dim, sums);
279+
}
280+
#endif
209281
// #if defined(__AVX512BW__) // TODO: this version is problematic
210282
// return compute_one_to_many_avx512_int8<ValueType, BatchSize>(
211283
// query, ptrs, prefetch_ptrs, dim, sums);

0 commit comments

Comments
 (0)