1+ #pragma once
2+
3+ #include < assert.h>
4+ #include < stdio.h>
5+ #include < cstring>
6+ #include < unordered_map>
7+
8+ #include < filesystem>
9+ #include < fstream>
10+ #include < iostream>
11+
12+ namespace deglib
13+ {
14+ /* *
15+ * A repository of float feature vectors.
16+ */
17+ class FeatureRepository
18+ {
19+ public:
20+ virtual size_t dims () const = 0;
21+ virtual size_t size () const = 0;
22+ virtual const std::byte* getFeature (const uint32_t vertexid) const = 0;
23+ virtual void clear () = 0;
24+ };
25+
26+ /* *
27+ * A repository of float feature vectors. Since the repository deals
28+ * with static data, a single contiguous array is preserved internally.
29+ */
30+ class StaticFeatureRepository : public FeatureRepository
31+ {
32+ public:
33+ StaticFeatureRepository (std::unique_ptr<std::byte[]> contiguous_features, const size_t dims, const size_t count, const size_t bytes_per_dim)
34+ : bytes_per_dim_{bytes_per_dim}, dims_{dims}, count_{count}, contiguous_features_{std::move (contiguous_features), }
35+ {
36+ }
37+
38+ size_t dims () const override { return dims_; }
39+ size_t size () const override { return count_; }
40+ const std::byte* getFeature (const uint32_t idx) const override { return &contiguous_features_[idx * dims_ * bytes_per_dim_]; }
41+ void clear () override { contiguous_features_.reset (); }
42+
43+ private:
44+ const size_t bytes_per_dim_;
45+ const size_t dims_;
46+ const size_t count_;
47+ std::unique_ptr<std::byte[]> contiguous_features_;
48+ };
49+
50+
51+ /* ****************************************************
52+ * I/O functions for fvecs and ivecs
53+ * Reference
54+ * https://github.com/facebookresearch/faiss/blob/e86bf8cae1a0ecdaee1503121421ed262ecee98c/demos/demo_sift1M.cpp
55+ *****************************************************/
56+
57+ auto fvecs_read (const char * fname, size_t & d_out, size_t & n_out)
58+ {
59+ std::error_code ec{};
60+ auto file_size = std::filesystem::file_size (fname, ec);
61+ if (ec != std::error_code{})
62+ {
63+ std::fprintf (stderr, " error when accessing file %s, size is: %ju message: %s \n " , fname, file_size, ec.message ().c_str ());
64+ perror (" " );
65+ abort ();
66+ }
67+
68+ // open as binary
69+ auto ifstream = std::ifstream (fname, std::ios::binary);
70+ if (!ifstream.is_open ())
71+ {
72+ std::fprintf (stderr, " could not open %s\n " , fname);
73+ perror (" " );
74+ abort ();
75+ }
76+
77+ // read dimension header
78+ uint32_t dims = 0 ;
79+ ifstream.read (reinterpret_cast <char *>(&dims), sizeof (dims));
80+ assert ((dims > 0 && dims < 1'000'000 ) && " unreasonable dimension" );
81+
82+ // compute number of rows
83+ assert (file_size % ((dims + 1 ) * sizeof (float )) == 0 || !" weird file size" );
84+ size_t n = (size_t )file_size / ((dims + 1 ) * sizeof (float ));
85+ d_out = dims;
86+ n_out = n;
87+
88+ // read data rows (each row starts with its dimension, which is 4 bytes)
89+ auto x = std::make_unique<std::byte[]>(file_size);
90+ ifstream.seekg (0 );
91+ ifstream.read (reinterpret_cast <char *>(x.get ()), file_size);
92+ if (!ifstream) assert (ifstream.gcount () == static_cast <int >(file_size) || !" could not read whole file" );
93+
94+ // shift array to remove row headers
95+ for (size_t i = 0 ; i < n; i++) std::memmove (&x[i * dims * sizeof (float )], &x[sizeof (dims) + i * (dims + 1 ) * sizeof (float )], dims * sizeof (float ));
96+
97+ ifstream.close ();
98+ return x;
99+ }
100+
101+ auto u8vecs_read (const char * fname, size_t & d_out, size_t & n_out)
102+ {
103+ // get total file size
104+ std::error_code ec{};
105+ auto file_size = std::filesystem::file_size (fname, ec);
106+ if (ec != std::error_code{})
107+ {
108+ std::fprintf (stderr, " error when accessing file %s, size is: %ju message: %s \n " , fname, file_size, ec.message ().c_str ());
109+ std::abort ();
110+ }
111+
112+ // open as binary
113+ auto ifstream = std::ifstream (fname, std::ios::binary);
114+ if (!ifstream.is_open ())
115+ {
116+ std::fprintf (stderr, " could not open %s\n " , fname);
117+ std::abort ();
118+ }
119+
120+ // read dimension header
121+ uint32_t dims = 0 ;
122+ ifstream.read (reinterpret_cast <char *>(&dims), sizeof (dims));
123+ assert ((dims > 0 && dims < 1'000'000 ) && " unreasonable dimension" );
124+
125+ // compute number of rows
126+ assert (file_size % (dims + 4 ) == 0 || !" weird file size" );
127+ size_t n = (size_t )file_size / (dims + 4 );
128+ d_out = dims;
129+ n_out = n;
130+
131+ // read data rows (each row starts with its dimension, which is 4 bytes)
132+ auto x = std::make_unique<std::byte[]>(file_size);
133+ ifstream.seekg (0 );
134+ ifstream.read (reinterpret_cast <char *>(x.get ()), file_size);
135+ if (!ifstream) assert (ifstream.gcount () == static_cast <int >(file_size) || !" could not read whole file" );
136+
137+ // shift array to remove row headers
138+ for (size_t i = 0 ; i < n; i++) std::memmove (&x[i * dims], &x[sizeof (dims) + i * (dims + sizeof (dims))], dims);
139+
140+ ifstream.close ();
141+ return x;
142+ }
143+
144+ bool string_ends_with (const char * str, const char * suffix) {
145+ size_t str_len = std::strlen (str);
146+ size_t suffix_len = std::strlen (suffix);
147+
148+ if (suffix_len > str_len) {
149+ return false ;
150+ }
151+ return std::strcmp (str + str_len - suffix_len, suffix) == 0 ;
152+ }
153+
154+ StaticFeatureRepository load_static_repository (const char * path_repository)
155+ {
156+ if (string_ends_with (path_repository, " fvecs" )) {
157+ size_t dims;
158+ size_t count;
159+ auto contiguous_features = fvecs_read (path_repository, dims, count);
160+ return StaticFeatureRepository (std::move (contiguous_features), dims, count, sizeof (float ));
161+ } else if (string_ends_with (path_repository, " u8vecs" )) {
162+ size_t dims;
163+ size_t count;
164+ auto contiguous_features = u8vecs_read (path_repository, dims, count);
165+ return StaticFeatureRepository (std::move (contiguous_features), dims, count, sizeof (uint8_t ));
166+ }
167+
168+ std::fprintf (stderr, " unsupported file extension, only fvecs and u8vecs are supported, but got %s \n " , path_repository);
169+ std::perror (" " );
170+ std::abort ();
171+ }
172+
173+ } // namespace deglib
0 commit comments