Skip to content

Commit 2bf44ea

Browse files
committed
Fix half type numpy conversions
1 parent 6d9b2c6 commit 2bf44ea

10 files changed

Lines changed: 103 additions & 19 deletions

File tree

cpp/include/cuvs/util/file_io.hpp

Lines changed: 4 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@
77
#include <raft/core/error.hpp>
88
#include <raft/core/serialize.hpp>
99

10+
#include <cuvs/util/numpy_dtype.hpp>
11+
1012
#include <algorithm>
1113
#include <cstring>
1214
#include <istream>
@@ -188,15 +190,8 @@ std::pair<file_descriptor, size_t> create_numpy_file(const std::string& path,
188190
// Open file
189191
file_descriptor fd(path, O_CREAT | O_RDWR | O_TRUNC, 0644);
190192

191-
// Build header
192-
const auto dtype = raft::detail::numpy_serializer::get_numpy_dtype<T>();
193-
const bool fortran_order = false;
194-
const raft::detail::numpy_serializer::header_t header = {dtype, fortran_order, shape};
195-
196-
std::stringstream ss;
197-
raft::detail::numpy_serializer::write_header(ss, header);
198-
std::string header_str = ss.str();
199-
size_t header_size = header_str.size();
193+
const std::string header_str = make_numpy_header_string<T>(shape);
194+
size_t header_size = header_str.size();
200195

201196
// Calculate data size from shape
202197
size_t data_bytes = sizeof(T);
Lines changed: 86 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,86 @@
1+
/*
2+
* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION.
3+
* SPDX-License-Identifier: Apache-2.0
4+
*/
5+
#pragma once
6+
7+
#include <cuvs/core/cuda_fp16.hpp>
8+
#include <cuvs/core/export.hpp>
9+
10+
#include <raft/core/detail/mdspan_numpy_serializer.hpp>
11+
#include <raft/core/error.hpp>
12+
13+
#include <cstddef>
14+
#include <cstdint>
15+
#include <limits>
16+
#include <sstream>
17+
#include <string>
18+
#include <type_traits>
19+
#include <vector>
20+
21+
namespace CUVS_EXPORT cuvs {
22+
namespace util {
23+
24+
template <typename T>
25+
inline auto numpy_dtype_string() -> std::string
26+
{
27+
if constexpr (std::is_same_v<T, half>) {
28+
return "<f2";
29+
} else {
30+
return raft::detail::numpy_serializer::get_numpy_dtype<T>().to_string();
31+
}
32+
}
33+
34+
inline auto make_numpy_header_from_dtype(const std::string& dtype, const std::vector<size_t>& shape)
35+
-> std::string
36+
{
37+
std::stringstream dict;
38+
dict << "{'descr': '" << dtype << "', 'fortran_order': False, 'shape': (";
39+
for (size_t i = 0; i < shape.size(); ++i) {
40+
if (i != 0) { dict << ", "; }
41+
dict << shape[i];
42+
}
43+
if (shape.size() == 1) { dict << ","; }
44+
dict << "), }";
45+
46+
std::string header = dict.str();
47+
constexpr size_t preamble_size =
48+
6 + 2 + sizeof(uint16_t); // magic string, version, v1 header length
49+
const size_t remainder = (preamble_size + header.size() + 1) % 16;
50+
if (remainder != 0) { header.append(16 - remainder, ' '); }
51+
header.push_back('\n');
52+
53+
RAFT_EXPECTS(header.size() <= std::numeric_limits<uint16_t>::max(),
54+
"NumPy v1 header is too large: %zu bytes",
55+
header.size());
56+
57+
const auto header_len = static_cast<uint16_t>(header.size());
58+
std::string result;
59+
result.reserve(preamble_size + header.size());
60+
result.append("\x93NUMPY", 6);
61+
result.push_back(1);
62+
result.push_back(0);
63+
result.push_back(static_cast<char>(header_len & 0xff));
64+
result.push_back(static_cast<char>((header_len >> 8) & 0xff));
65+
result.append(header);
66+
return result;
67+
}
68+
69+
template <typename T>
70+
inline auto make_numpy_header_string(const std::vector<size_t>& shape) -> std::string
71+
{
72+
if constexpr (std::is_same_v<T, half>) {
73+
return make_numpy_header_from_dtype(numpy_dtype_string<T>(), shape);
74+
} else {
75+
const auto dtype = raft::detail::numpy_serializer::get_numpy_dtype<T>();
76+
const bool fortran_order = false;
77+
const raft::detail::numpy_serializer::header_t header = {dtype, fortran_order, shape};
78+
79+
std::stringstream ss;
80+
raft::detail::numpy_serializer::write_header(ss, header);
81+
return ss.str();
82+
}
83+
}
84+
85+
} // namespace util
86+
} // namespace CUVS_EXPORT cuvs

cpp/src/neighbors/brute_force_serialize.cu

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@ void serialize(raft::resources const& handle,
2626
RAFT_LOG_DEBUG(
2727
"Saving brute force index, size %zu, dim %u", static_cast<size_t>(index.size()), index.dim());
2828

29-
auto dtype_string = raft::detail::numpy_serializer::get_numpy_dtype<T>().to_string();
29+
auto dtype_string = cuvs::util::numpy_dtype_string<T>();
3030
dtype_string.resize(4);
3131
os << dtype_string;
3232

cpp/src/neighbors/detail/cagra/cagra_serialize.cuh

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -55,7 +55,7 @@ void serialize(raft::resources const& res,
5555
RAFT_LOG_DEBUG(
5656
"Saving CAGRA index, size %zu, dim %u", static_cast<size_t>(index_.size()), index_.dim());
5757

58-
std::string dtype_string = raft::detail::numpy_serializer::get_numpy_dtype<T>().to_string();
58+
std::string dtype_string = cuvs::util::numpy_dtype_string<T>();
5959
dtype_string.resize(4);
6060
os << dtype_string;
6161

cpp/src/neighbors/detail/hnsw.hpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -740,7 +740,7 @@ inline auto open_npy_file(const std::string& path) -> npy_file
740740
template <typename T>
741741
inline void validate_npy_file(const npy_file& file, const std::string& path, const char* name)
742742
{
743-
const auto expected_dtype = raft::detail::numpy_serializer::get_numpy_dtype<T>().to_string();
743+
const auto expected_dtype = cuvs::util::numpy_dtype_string<T>();
744744
RAFT_EXPECTS(file.dtype == expected_dtype,
745745
"%s dtype (%s) does not match expected dtype (%s): %s",
746746
name,
@@ -802,7 +802,7 @@ inline auto open_layered_dataset_file(const std::string& path) -> npy_file
802802
return {std::move(fd),
803803
sizeof(uint32_t) * header.size(),
804804
{header[0], header[1]},
805-
raft::detail::numpy_serializer::get_numpy_dtype<T>().to_string(),
805+
cuvs::util::numpy_dtype_string<T>(),
806806
false};
807807
}
808808

cpp/src/neighbors/ivf_flat/ivf_flat_serialize.cuh

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -45,7 +45,7 @@ void serialize(raft::resources const& handle, std::ostream& os, const index<T, I
4545
RAFT_LOG_DEBUG(
4646
"Saving IVF-Flat index, size %zu, dim %u", static_cast<size_t>(index_.size()), index_.dim());
4747

48-
std::string dtype_string = raft::detail::numpy_serializer::get_numpy_dtype<T>().to_string();
48+
std::string dtype_string = cuvs::util::numpy_dtype_string<T>();
4949
dtype_string.resize(4);
5050
os << dtype_string;
5151

cpp/src/neighbors/ivf_sq/ivf_sq_serialize.cuh

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@ void serialize(raft::resources const& handle, std::ostream& os, const index<Code
2929
RAFT_LOG_DEBUG(
3030
"Saving IVF-SQ index, size %zu, dim %u", static_cast<size_t>(index_.size()), index_.dim());
3131

32-
std::string dtype_string = raft::detail::numpy_serializer::get_numpy_dtype<CodeT>().to_string();
32+
std::string dtype_string = cuvs::util::numpy_dtype_string<CodeT>();
3333
dtype_string.resize(4);
3434
os << dtype_string;
3535

cpp/src/neighbors/mg/snmg.cuh

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
#include <cuvs/neighbors/ivf_flat.hpp>
2222
#include <cuvs/neighbors/ivf_pq.hpp>
2323
#include <cuvs/neighbors/knn_merge_parts.hpp>
24+
#include <cuvs/util/numpy_dtype.hpp>
2425

2526
#include <fstream>
2627

@@ -738,7 +739,7 @@ void serialize(const raft::resources& clique,
738739
std::ofstream of(filename, std::ios::out | std::ios::binary);
739740
if (!of) { RAFT_FAIL("Cannot open file %s", filename.c_str()); }
740741

741-
std::string dtype_string = raft::detail::numpy_serializer::get_numpy_dtype<T>().to_string();
742+
std::string dtype_string = cuvs::util::numpy_dtype_string<T>();
742743
dtype_string.resize(4);
743744
of << dtype_string;
744745

cpp/src/util/serialize_validation.hpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,8 @@
66

77
#include <cuvs/distance/distance.hpp>
88
#include <cuvs/neighbors/ivf_pq.hpp>
9+
#include <cuvs/util/numpy_dtype.hpp>
910

10-
#include <raft/core/detail/mdspan_numpy_serializer.hpp>
1111
#include <raft/core/error.hpp>
1212

1313
#include <algorithm>
@@ -44,7 +44,7 @@ inline bool validate_serialized_dtype(const char* dtype_prefix, std::size_t dtyp
4444
{
4545
if (dtype_prefix == nullptr || dtype_prefix_size != 4) { return false; }
4646

47-
auto expected_dtype = raft::detail::numpy_serializer::get_numpy_dtype<T>().to_string();
47+
auto expected_dtype = cuvs::util::numpy_dtype_string<T>();
4848
expected_dtype.resize(dtype_prefix_size, '\0');
4949

5050
return std::equal(dtype_prefix, dtype_prefix + dtype_prefix_size, expected_dtype.begin());

examples/cpp/CMakeLists.txt

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -67,7 +67,9 @@ target_link_libraries(
6767
DYNAMIC_BATCHING_EXAMPLE PRIVATE cuvs::cuvs $<TARGET_NAME_IF_EXISTS:conda_env> Threads::Threads
6868
)
6969
target_link_libraries(HNSW_ACE_EXAMPLE PRIVATE cuvs::cuvs $<TARGET_NAME_IF_EXISTS:conda_env>)
70-
target_link_libraries(HNSW_ACE_LAYERED_EXAMPLE PRIVATE cuvs::cuvs $<TARGET_NAME_IF_EXISTS:conda_env>)
70+
target_link_libraries(
71+
HNSW_ACE_LAYERED_EXAMPLE PRIVATE cuvs::cuvs $<TARGET_NAME_IF_EXISTS:conda_env>
72+
)
7173
target_link_libraries(HNSW_OPENAI_EXAMPLE PRIVATE cuvs::cuvs $<TARGET_NAME_IF_EXISTS:conda_env>)
7274
target_link_libraries(IVF_PQ_EXAMPLE PRIVATE cuvs::cuvs $<TARGET_NAME_IF_EXISTS:conda_env>)
7375
target_link_libraries(IVF_FLAT_EXAMPLE PRIVATE cuvs::cuvs $<TARGET_NAME_IF_EXISTS:conda_env>)

0 commit comments

Comments
 (0)