Skip to content

Commit 16980e4

Browse files
committed
feature(metric): add support for FP16 inner product calculations
1 parent f567928 commit 16980e4

10 files changed

Lines changed: 1537 additions & 20 deletions

File tree

‎cpp/deglib/include/config.h‎

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,17 +2,25 @@
22

33
#include <cstdint>
44

5+
// Architecture flags. Define DEGLIB_X86 if the code is for x86 machines. ARM often has separate coding pathes.
56
#if defined(__x86_64__) || defined(_M_X64) || defined(__i386__) || defined(_M_IX86)
67
#define DEGLIB_X86 1
78
#endif
89

9-
// Target attribute for AVX-512 functions on GCC/Clang
10+
// Compile methods with this attribute for AVX-512 functions on GCC/Clang
1011
#if defined(DEGLIB_X86) && (defined(__GNUC__) || defined(__clang__))
1112
#define DEGLIB_TARGET_AVX512 __attribute__((target("avx512f,avx512dq,avx512bw,fma")))
1213
#else
1314
#define DEGLIB_TARGET_AVX512
1415
#endif
1516

17+
// Compile methods with this attribute for F16C functions on GCC/Clang
18+
#if defined(DEGLIB_X86) && (defined(__GNUC__) || defined(__clang__))
19+
#define DEGLIB_TARGET_F16C __attribute__((target("f16c,avx")))
20+
#else
21+
#define DEGLIB_TARGET_F16C
22+
#endif
23+
1624
// Architecture intrinsic headers
1725
#if defined(DEGLIB_X86)
1826
#ifdef _MSC_VER

‎cpp/deglib/include/distance/fp16.h‎

Lines changed: 322 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,332 @@
11
#pragma once
22

3+
#include <cmath>
4+
#include <cstdint>
5+
#include <cstring>
36
#include "config.h"
47

58
namespace deglib::distances {
69

710
// Shared FP16 distance utilities and declarations.
811
// Base header for future FP16 metric modules (fp16_l2.h, fp16_ip.h).
912

13+
// ---------------------------------------------------------------------------
14+
// FP16 <-> Float conversion utilities
15+
// ---------------------------------------------------------------------------
16+
// IEEE 754 half-precision (binary16) conversion functions.
17+
// FP16 vectors are stored as uint16_t bit patterns (IEEE 754 half-precision).
18+
// These functions provide bit-exact conversion between float (binary32)
19+
// and half-precision, using hardware F16C when available and a
20+
// bit-manipulation scalar fallback otherwise.
21+
// ---------------------------------------------------------------------------
22+
23+
namespace fp16 {
24+
25+
// ---------------------------------------------------------------------------
26+
// GCC/Clang F16C intrinsics (target-attributed)
27+
// ---------------------------------------------------------------------------
28+
// On GCC/Clang, __attribute__((target("f16c,avx"))) forces the compiler to
29+
// generate F16C instructions without requiring global -mf16c flags.
30+
// On MSVC, these intrinsics are not available in the same form, so we use
31+
// _mm_cvtps_ph / _mm_cvtph_ps (SSE intrinsics) instead.
32+
// ---------------------------------------------------------------------------
33+
34+
#if defined(DEGLIB_X86) && (defined(__GNUC__) || defined(__clang__))
35+
36+
DEGLIB_TARGET_F16C inline uint16_t float_to_fp16_gcc(float f) {
37+
return _cvtss_sh(f, 0);
38+
}
39+
40+
DEGLIB_TARGET_F16C inline float fp16_to_float_gcc(uint16_t h) {
41+
return _cvtsh_ss(h);
42+
}
43+
44+
DEGLIB_TARGET_F16C inline void floats_to_fp16_gcc(const float* floats, uint16_t* fp16_vals, size_t count) {
45+
size_t i = 0;
46+
// Process 8 floats per step with _mm256_cvtps_ph
47+
for (; i + 8 <= count; i += 8) {
48+
__m256 va = _mm256_loadu_ps(floats + i);
49+
__m128i vhp = _mm256_cvtps_ph(va, 0);
50+
_mm_storeu_si128(reinterpret_cast<__m128i*>(fp16_vals + i), vhp);
51+
}
52+
// Process 4 floats per step with _mm_cvtps_ph
53+
for (; i + 4 <= count; i += 4) {
54+
__m128 va = _mm_loadu_ps(floats + i);
55+
__m128i vhp = _mm_cvtps_ph(va, 0);
56+
_mm_storel_epi64(reinterpret_cast<__m128i*>(fp16_vals + i), vhp);
57+
}
58+
// Scalar fallback for remaining 0-3 elements
59+
for (; i < count; ++i) {
60+
fp16_vals[i] = _cvtss_sh(floats[i], 0);
61+
}
62+
}
63+
64+
DEGLIB_TARGET_F16C inline void fp16_to_floats_gcc(const uint16_t* fp16_vals, float* floats, size_t count) {
65+
size_t i = 0;
66+
// Process 8 uint16_t per step with _mm256_cvtph_ps
67+
for (; i + 8 <= count; i += 8) {
68+
__m128i vhp = _mm_loadu_si128(reinterpret_cast<const __m128i*>(fp16_vals + i));
69+
__m256 va = _mm256_cvtph_ps(vhp);
70+
_mm256_storeu_ps(floats + i, va);
71+
}
72+
// Process 4 uint16_t per step with _mm_cvtph_ps
73+
for (; i + 4 <= count; i += 4) {
74+
__m128i vhp = _mm_loadl_epi64(reinterpret_cast<const __m128i*>(fp16_vals + i));
75+
__m128 va = _mm_cvtph_ps(vhp);
76+
_mm_storeu_ps(floats + i, va);
77+
}
78+
// Scalar fallback for remaining 0-3 elements
79+
for (; i < count; ++i) {
80+
floats[i] = _cvtsh_ss(fp16_vals[i]);
81+
}
82+
}
83+
84+
#endif // defined(DEGLIB_X86) && (defined(__GNUC__) || defined(__clang__))
85+
86+
// ---------------------------------------------------------------------------
87+
// MSVC F16C intrinsics (using SSE _mm_cvtps_ph / _mm_cvtph_ps)
88+
// ---------------------------------------------------------------------------
89+
// MSVC's <immintrin.h> doesn't define GCC-style F16C intrinsics
90+
// (_cvtss_sh, _cvtsh_ss). Instead, we use _mm_cvtps_ph / _mm_cvtph_ps
91+
// which operate on __m128/__m128i and are available with /arch:AVX2.
92+
// ---------------------------------------------------------------------------
93+
94+
#if defined(DEGLIB_X86) && defined(_MSC_VER)
95+
96+
inline uint16_t float_to_fp16_msvc(float f) {
97+
__m128 f_val = _mm_set_ss(f);
98+
__m128i h_val = _mm_cvtps_ph(f_val, 0);
99+
return static_cast<uint16_t>(_mm_cvtsi128_si32(h_val));
100+
}
101+
102+
inline float fp16_to_float_msvc(uint16_t h) {
103+
__m128i h_val = _mm_cvtsi32_si128(static_cast<int>(h));
104+
__m128 f_val = _mm_cvtph_ps(h_val);
105+
return _mm_cvtss_f32(f_val);
106+
}
107+
108+
inline void floats_to_fp16_msvc(const float* floats, uint16_t* fp16_vals, size_t count) {
109+
size_t i = 0;
110+
// Process 4 floats per step with _mm_cvtps_ph
111+
for (; i + 4 <= count; i += 4) {
112+
__m128 va = _mm_loadu_ps(floats + i);
113+
__m128i vhp = _mm_cvtps_ph(va, 0);
114+
// Store 4 uint16_t values from the __m128i
115+
alignas(16) uint16_t temp[4];
116+
_mm_store_si128(reinterpret_cast<__m128i*>(temp), vhp);
117+
fp16_vals[i] = temp[0];
118+
fp16_vals[i + 1] = temp[1];
119+
fp16_vals[i + 2] = temp[2];
120+
fp16_vals[i + 3] = temp[3];
121+
}
122+
// Scalar fallback for remaining 0-3 elements
123+
for (; i < count; ++i) {
124+
fp16_vals[i] = float_to_fp16_msvc(floats[i]);
125+
}
126+
}
127+
128+
inline void fp16_to_floats_msvc(const uint16_t* fp16_vals, float* floats, size_t count) {
129+
size_t i = 0;
130+
// Process 4 uint16_t per step with _mm_cvtph_ps
131+
for (; i + 4 <= count; i += 4) {
132+
// Load 4 uint16_t values into __m128i
133+
alignas(16) uint16_t temp[4] = {fp16_vals[i], fp16_vals[i + 1], fp16_vals[i + 2], fp16_vals[i + 3]};
134+
__m128i vhp = _mm_load_si128(reinterpret_cast<const __m128i*>(temp));
135+
__m128 va = _mm_cvtph_ps(vhp);
136+
_mm_storeu_ps(floats + i, va);
137+
}
138+
// Scalar fallback for remaining 0-3 elements
139+
for (; i < count; ++i) {
140+
floats[i] = fp16_to_float_msvc(fp16_vals[i]);
141+
}
142+
}
143+
144+
#endif // defined(DEGLIB_X86) && defined(_MSC_VER)
145+
146+
// ---------------------------------------------------------------------------
147+
// Scalar IEEE 754 Round-to-Nearest-Even fallback
148+
// ---------------------------------------------------------------------------
149+
// Used when F16C is not available at runtime (or not compiled in).
150+
// Implements proper round-to-nearest-even rounding per IEEE 754.
151+
// ---------------------------------------------------------------------------
152+
153+
inline uint16_t float_to_fp16_scalar(float f) {
154+
uint32_t x;
155+
std::memcpy(&x, &f, sizeof(f));
156+
uint32_t sign = (x >> 31) & 0x1;
157+
uint32_t mantissa = x & 0x7FFFFF;
158+
int32_t exponent = ((x >> 23) & 0xFF) - 127;
159+
uint32_t fp16_exp;
160+
161+
if (exponent < -24) {
162+
// Too small, underflow to zero
163+
return static_cast<uint16_t>(sign << 15);
164+
} else if (exponent < -14) {
165+
// Subnormal FP16
166+
int32_t shift = -14 - exponent;
167+
uint32_t mantissa_with_hidden = mantissa | 0x800000;
168+
uint32_t fp16_mantissa = mantissa_with_hidden >> (23 + shift);
169+
// Round to nearest even
170+
uint32_t remainder = mantissa_with_hidden & ((1u << (23 + shift)) - 1);
171+
if (remainder > (1u << (22 + shift)) ||
172+
(remainder == (1u << (22 + shift)) && (fp16_mantissa & 1))) {
173+
fp16_mantissa++;
174+
}
175+
fp16_exp = 0;
176+
return static_cast<uint16_t>((sign << 15) | (fp16_exp << 10) | (fp16_mantissa & 0x3FF));
177+
} else if (exponent <= 15) {
178+
// Normal FP16
179+
fp16_exp = static_cast<uint32_t>(exponent + 15);
180+
uint32_t fp16_mantissa = mantissa >> 13;
181+
// Round to nearest even
182+
uint32_t remainder = mantissa & 0x1FFF;
183+
if (remainder > 0x1000 ||
184+
(remainder == 0x1000 && (fp16_mantissa & 1))) {
185+
fp16_mantissa++;
186+
if (fp16_mantissa > 0x3FF) {
187+
fp16_mantissa = 0;
188+
fp16_exp++;
189+
}
190+
}
191+
return static_cast<uint16_t>((sign << 15) | (fp16_exp << 10) | (fp16_mantissa & 0x3FF));
192+
} else if (exponent >= 128) {
193+
// NaN or Inf (exponent field is 0xFF in the float)
194+
if (mantissa == 0) {
195+
return static_cast<uint16_t>((sign << 15) | (0x1F << 10));
196+
} else {
197+
return static_cast<uint16_t>((sign << 15) | (0x1F << 10) | 0x200 | (mantissa >> 13));
198+
}
199+
} else {
200+
// Too large, overflow to infinity
201+
return static_cast<uint16_t>((sign << 15) | (0x1F << 10));
202+
}
203+
}
204+
205+
inline float fp16_to_float_scalar(uint16_t h) {
206+
uint32_t sign = (h >> 15) & 0x1;
207+
uint32_t fp16_exp = (h >> 10) & 0x1F;
208+
uint32_t fp16_mantissa = h & 0x3FF;
209+
uint32_t float_bits;
210+
211+
if (fp16_exp == 0) {
212+
// Zero or subnormal
213+
if (fp16_mantissa == 0) {
214+
float_bits = sign << 31;
215+
} else {
216+
// Subnormal: normalize
217+
uint32_t mantissa = fp16_mantissa;
218+
int32_t shift = 0;
219+
while ((mantissa & 0x400) == 0) {
220+
mantissa <<= 1;
221+
shift--;
222+
}
223+
mantissa &= 0x3FF; // Remove the implicit leading 1
224+
float_bits = (sign << 31) | ((127 + (-14) + shift) << 23) | (mantissa << 13);
225+
}
226+
} else if (fp16_exp == 0x1F) {
227+
// Inf or NaN
228+
if (fp16_mantissa == 0) {
229+
float_bits = (sign << 31) | (0xFF << 23);
230+
} else {
231+
float_bits = (sign << 31) | (0xFF << 23) | 0x7FFFFF;
232+
}
233+
} else {
234+
// Normal number
235+
float_bits = (sign << 31) | ((fp16_exp + 127 - 15) << 23) | (fp16_mantissa << 13);
236+
}
237+
238+
float result;
239+
std::memcpy(&result, &float_bits, sizeof(float));
240+
return result;
241+
}
242+
243+
// ---------------------------------------------------------------------------
244+
// Public API: float_to_fp16, fp16_to_float, floats_to_fp16, fp16_to_floats
245+
// ---------------------------------------------------------------------------
246+
// Runtime dispatch via deglib::cpu::has_f16c().
247+
// On GCC/Clang: uses DEGLIB_TARGET_F16C-attributed intrinsics.
248+
// On MSVC: uses _mm_cvtps_ph / _mm_cvtph_ps (SSE intrinsics).
249+
// Fallback: scalar IEEE 754 Round-to-Nearest-Even.
250+
// ---------------------------------------------------------------------------
251+
252+
inline uint16_t float_to_fp16(float f) {
253+
#if defined(DEGLIB_X86)
254+
if (deglib::cpu::has_f16c()) {
255+
#if defined(__GNUC__) || defined(__clang__)
256+
return float_to_fp16_gcc(f);
257+
#elif defined(_MSC_VER)
258+
return float_to_fp16_msvc(f);
259+
#endif
260+
}
261+
#endif
262+
return float_to_fp16_scalar(f);
263+
}
264+
265+
inline float fp16_to_float(uint16_t h) {
266+
#if defined(DEGLIB_X86)
267+
if (deglib::cpu::has_f16c()) {
268+
#if defined(__GNUC__) || defined(__clang__)
269+
return fp16_to_float_gcc(h);
270+
#elif defined(_MSC_VER)
271+
return fp16_to_float_msvc(h);
272+
#endif
273+
}
274+
#endif
275+
return fp16_to_float_scalar(h);
276+
}
277+
278+
inline void floats_to_fp16(const float* floats, uint16_t* fp16_vals, size_t count) {
279+
#if defined(DEGLIB_X86)
280+
if (deglib::cpu::has_f16c()) {
281+
#if defined(__GNUC__) || defined(__clang__)
282+
floats_to_fp16_gcc(floats, fp16_vals, count);
283+
return;
284+
#elif defined(_MSC_VER)
285+
floats_to_fp16_msvc(floats, fp16_vals, count);
286+
return;
287+
#endif
288+
}
289+
#endif
290+
// Scalar fallback
291+
for (size_t i = 0; i < count; ++i) {
292+
fp16_vals[i] = float_to_fp16_scalar(floats[i]);
293+
}
294+
}
295+
296+
inline void fp16_to_floats(const uint16_t* fp16_vals, float* floats, size_t count) {
297+
#if defined(DEGLIB_X86)
298+
if (deglib::cpu::has_f16c()) {
299+
#if defined(__GNUC__) || defined(__clang__)
300+
fp16_to_floats_gcc(fp16_vals, floats, count);
301+
return;
302+
#elif defined(_MSC_VER)
303+
fp16_to_floats_msvc(fp16_vals, floats, count);
304+
return;
305+
#endif
306+
}
307+
#endif
308+
// Scalar fallback
309+
for (size_t i = 0; i < count; ++i) {
310+
floats[i] = fp16_to_float_scalar(fp16_vals[i]);
311+
}
312+
}
313+
314+
// Naive scalar inner product for FP16 vectors (used for testing and fallback).
315+
// Computes the raw dot product (without 1.f -) using std::fma for precision.
316+
inline float fp16_ip_naive(const void* pVect1v, const void* pVect2v, const void* qty_ptr) {
317+
const uint16_t* a = static_cast<const uint16_t*>(pVect1v);
318+
const uint16_t* b = static_cast<const uint16_t*>(pVect2v);
319+
size_t size = *((size_t*)qty_ptr);
320+
321+
float result = 0.0f;
322+
for (size_t i = 0; i < size; ++i) {
323+
float fa = fp16_to_float(a[i]);
324+
float fb = fp16_to_float(b[i]);
325+
result = std::fma(fa, fb, result);
326+
}
327+
return result;
328+
}
329+
330+
} // namespace fp16
331+
10332
} // end namespace deglib::distances

0 commit comments

Comments
 (0)