-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtokenizer.hpp
More file actions
199 lines (153 loc) · 7.44 KB
/
Copy pathtokenizer.hpp
File metadata and controls
199 lines (153 loc) · 7.44 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
#pragma once
#include <queue>
#include <format>
#include "hash_map.hpp"
#include "gguf.hpp"
namespace misha::gemma4 {
struct tokenizer {
std::vector<std::string_view> tokens;
std::vector<std::string_view> control_tokens;
hash_map<std::string_view, int> token_to_id;
hash_map<std::string_view, int> merge_ranks;
uint32_t bos_token_id = 2;
uint32_t eos_token_id = 1;
uint32_t unknown_token_id = 3;
uint32_t padding_token_id = 0;
uint32_t mask_token_id = 5;
bool add_bos_token = true;
bool add_space_prefix = false;
static constexpr std::string_view space_replacement = "\xE2\x96\x81"; // U+2581 '_'
auto encode(std::string_view text) -> std::vector<int> {
assert(text.length() <= std::numeric_limits<int>::max());
struct piece_t {
int prev = -1, next = -1, start = 0, len = 0;
};
struct merged_piece_t {
int left, right, rank, total_len;
auto operator<(const merged_piece_t& other) const noexcept -> bool {
return rank > other.rank;
}
};
std::string processed_text;
std::vector<piece_t> pieces;
std::priority_queue<merged_piece_t> merged_pieces_queue;
std::string merged_buffer;
auto push_merged_piece = [&](int left, int right) noexcept {
merged_buffer.clear();
merged_buffer.append(processed_text.data() + pieces[left].start, pieces[left].len);
merged_buffer.push_back(' ');
merged_buffer.append(processed_text.data() + pieces[right].start, pieces[right].len);
auto rank_opt = merge_ranks.find(merged_buffer);
if (rank_opt)
merged_pieces_queue.emplace(left, right, *rank_opt, pieces[left].len + pieces[right].len);
};
auto push_piece = [&](std::string_view s) noexcept {
int start = processed_text.length();
processed_text.append(s);
int piece_idx = pieces.size();
pieces.emplace_back(piece_idx - 1, piece_idx + 1, start, s.length());
if (piece_idx != 0)
push_merged_piece(piece_idx - 1, piece_idx);
};
if (add_space_prefix) {
processed_text += space_replacement;
pieces.emplace_back(-1, 1, 0, space_replacement.length());
}
for (int i = 0; i < std::ssize(text);) {
unsigned char utf8_first_byte = text[i];
int char_len = 1;
if ((utf8_first_byte & 0b1000'0000) == 0b0000'0000) char_len = 1; // 0xxxxxxx
else if ((utf8_first_byte & 0b1110'0000) == 0b1100'0000) char_len = 2; // 110xxxxx
else if ((utf8_first_byte & 0b1111'0000) == 0b1110'0000) char_len = 3; // 1110xxxx
else if ((utf8_first_byte & 0b1111'1000) == 0b1111'0000) char_len = 4; // 11110xxx
char_len = std::min<int>(char_len, text.length() - i);
auto utf8_char = std::string_view(text.data() + i, char_len);
if (utf8_char == "<")
for (auto control_token : control_tokens)
if (text.subview(i).starts_with(control_token)) {
push_piece(control_token);
i += control_token.length();
goto continue_label;
}
if (utf8_char == " ")
push_piece(space_replacement);
else if (token_to_id.contains(utf8_char))
push_piece(utf8_char);
else [[unlikely]]
for (uint8_t byte : utf8_char)
push_piece(std::format("<0x{:02X}>", byte));
i += char_len;
continue_label:;
}
if (not pieces.empty())
pieces.back().next = -1;
std::vector<int> output;
if (add_bos_token) output.push_back(bos_token_id);
while (not merged_pieces_queue.empty()) {
auto [left_idx, right_idx, rank, total_len] = merged_pieces_queue.top();
merged_pieces_queue.pop();
auto& left = pieces[left_idx];
auto& right = pieces[right_idx];
if (rank == std::numeric_limits<int>::max()) break;
if (left.len + right.len != total_len) continue;
left.len += right.len;
left.next = right.next;
right.len = left.len;
if (right.next != -1)
pieces[right.next].prev = left_idx;
if (left.prev != -1)
push_merged_piece(left.prev, left_idx);
if (right.next != -1)
push_merged_piece(left_idx, right.next);
}
for (int i = 0; i != -1; i = pieces[i].next) {
auto piece_str = std::string_view(processed_text.data() + pieces[i].start, pieces[i].len);
output.push_back(token_to_id.find(piece_str).value_or(unknown_token_id));
}
return output;
}
auto decode(int token_id) -> std::string& {
auto token = tokens[token_id < std::ssize(tokens) ? token_id : unknown_token_id];
thread_local std::string decoded_token;
decoded_token.clear();
for (size_t i = 0; i < token.size(); ) {
if (token.subview(i).starts_with(space_replacement)) {
decoded_token += ' ';
i += space_replacement.size();
}
else decoded_token += token[i++];
}
return decoded_token;
}
};
inline auto load_tokenizer_from_gguf(const gguf::gguf_model& model) -> tokenizer {
using enum gguf::gguf_type;
tokenizer res;
res.bos_token_id = model.get_metadata("tokenizer.ggml.bos_token_id", GGUF_TYPE_UINT32).value;
res.unknown_token_id = model.get_metadata("tokenizer.ggml.unknown_token_id", GGUF_TYPE_UINT32).value;
res.padding_token_id = model.get_metadata("tokenizer.ggml.padding_token_id", GGUF_TYPE_UINT32).value;
res.mask_token_id = model.get_metadata("tokenizer.ggml.mask_token_id", GGUF_TYPE_UINT32).value;
res.add_bos_token = model.get_metadata("tokenizer.ggml.add_bos_token", GGUF_TYPE_BOOL).value;
res.add_space_prefix = model.get_metadata("tokenizer.ggml.add_space_prefix", GGUF_TYPE_BOOL).value;
auto gguf_tokens = model.get_metadata("tokenizer.ggml.tokens", GGUF_TYPE_ARRAY, GGUF_TYPE_STRING);
auto gguf_token_type = model.get_metadata("tokenizer.ggml.token_type", GGUF_TYPE_ARRAY, GGUF_TYPE_INT32);
auto gguf_tokens_reader = gguf::gguf_reader { .mem = { (std::byte*)gguf_tokens.value, model.data.data() + model.data.size() } };
auto gguf_token_type_reader = gguf::gguf_reader { .mem = { (std::byte*)gguf_token_type.value, model.data.data() + model.data.size() } };
res.tokens.reserve(gguf_tokens.count);
for (size_t i = 0; i < gguf_tokens.count; i++) {
std::string_view token = gguf_tokens_reader.read_string();
int32_t token_type = gguf_token_type_reader.read<int32_t>();
res.tokens.push_back(token);
res.token_to_id[token] = i;
if (token_type == gguf::token_type::LLAMA_TOKEN_TYPE_CONTROL) [[unlikely]]
res.control_tokens.push_back(token);
}
auto gguf_merges = model.get_metadata("tokenizer.ggml.merges", GGUF_TYPE_ARRAY, GGUF_TYPE_STRING);
auto gguf_merges_reader = gguf::gguf_reader { .mem = { (std::byte*)gguf_merges.value, model.data.data() + model.data.size() } };
for (size_t i = 0; i < gguf_merges.count; i++) {
std::string_view merge = gguf_merges_reader.read_string();
res.merge_ranks[merge] = i;
}
return res;
}
}