23 if (!tok || !
token)
return -1;
25 return info ? info->id : -1;
32 char *tmp = stack_buf;
33 if (
text_len >= (
int)
sizeof(stack_buf)) {
34 tmp = (
char *)malloc((
size_t)
text_len + 1);
40 if (tmp != stack_buf) free(tmp);
50 if (match_len) *match_len = 0;
51 if (!tok || !
text || text_len <= 0 || pos >= (
size_t)
text_len)
return -1;
59 while (cur < (
size_t)
text_len && node) {
60 unsigned char c = (
unsigned char)
text[cur];
71 if (match_len) *match_len = best_len;
81 for (
int len =
max_len; len >= 1; len--) {
82 memcpy(tmp,
text + pos, (
size_t)len);
85 if (info && info->id >= 0 && info->is_special) {
86 if (match_len) *match_len = (size_t)len;
108#define GGUF_TOKEN_NORMAL 1
109#define GGUF_TOKEN_UNKNOWN 2
110#define GGUF_TOKEN_CONTROL 3
111#define GGUF_TOKEN_USER_DEFINED 4
112#define GGUF_TOKEN_UNUSED 5
113#define GGUF_TOKEN_BYTE 6
118 if (!tok->
types || token_id < 0 || token_id >= (int32_t)tok->
vocab_size) {
121 uint8_t t = tok->
types[token_id];
128 if (!tok->
types || token_id < 0 || token_id >= (int32_t)tok->
vocab_size) {
142 int len = snprintf(byte_token,
sizeof(byte_token),
"<0x%02X>", byte_val);
143 if (len <= 0)
return tok->
unk_id;
158 tok->
byte_token_id = (int32_t *)malloc(256 *
sizeof(int32_t));
163 for (
int i = 0; i < 256; i++) {
172 size_t len = strlen(
token);
176 unsigned char byte_val = (
unsigned char)
token[0];
180 unsigned int byte_val;
181 if (sscanf(
token,
"<0x%02X>", &byte_val) == 1 && byte_val < 256) {
190 if ((c & 0x80) == 0x00)
return 1;
191 if ((c & 0xE0) == 0xC0)
return 2;
192 if ((c & 0xF0) == 0xE0)
return 3;
193 if ((c & 0xF8) == 0xF0)
return 4;
207 if (
out_len + 3 > out_max)
return -1;
214 if (
text[i] ==
' ') {
224 if (
out_len + 3 > out_max)
return -1;
229 if (
out_len + run > out_max)
return -1;
230 for (
int k = 0; k < run; k++) {
236 if (
out_len + 1 > out_max)
return -1;
260 const SpmLlamaNode *nodes,
265 if (!tok || !nodes || node_id < 0 || !ids || out_idx >=
max_ids) {
269 const SpmLlamaNode *node = &nodes[node_id];
272 ids[out_idx++] = token_id;
276 if (node->left >= 0 && node->right >= 0) {
282 for (
int i = 0; i < node->n && out_idx <
max_ids; i++) {
284 ids[out_idx++] = (byte_token >= 0) ? byte_token : tok->
unk_id;
299 char preprocessed[8192];
302 if (pp_len < 0)
return 0;
303 preprocessed[pp_len] =
'\0';
306 for (
int offs = 0; offs < pp_len;) {
307 int char_len =
utf8_len((
unsigned char)preprocessed[offs]);
308 if (char_len <= 0) char_len = 1;
309 if (offs + char_len > pp_len) char_len = pp_len - offs;
313 if (num_symbols <= 0)
return 0;
315 SpmLlamaSymbol *symbols = (SpmLlamaSymbol *)calloc((
size_t)num_symbols,
sizeof(SpmLlamaSymbol));
316 int node_cap = 2 * num_symbols + 1;
317 SpmLlamaNode *nodes = (SpmLlamaNode *)calloc((
size_t)node_cap,
sizeof(SpmLlamaNode));
318 if (!symbols || !nodes) {
319 if (symbols) free(symbols);
320 if (nodes) free(nodes);
325 for (
int offs = 0; offs < pp_len && index < num_symbols;) {
326 int char_len =
utf8_len((
unsigned char)preprocessed[offs]);
327 if (char_len <= 0) char_len = 1;
328 if (offs + char_len > pp_len) char_len = pp_len - offs;
330 symbols[index].text = preprocessed + offs;
331 symbols[index].n = char_len;
332 symbols[index].prev = index - 1;
333 symbols[index].next = (index + 1 < num_symbols) ? (index + 1) : -1;
334 symbols[index].node_id = index;
336 nodes[index].text = preprocessed + offs;
337 nodes[index].n = char_len;
338 nodes[index].left = -1;
339 nodes[index].right = -1;
345 int node_count = num_symbols;
349 float best_score = -1e30f;
353 if (
right < 0)
continue;
355 int pair_len = symbols[
left].n + symbols[
right].n;
357 if (token_id < 0 || token_id >= (int32_t)tok->
vocab_size)
continue;
364 if (best_left < 0 || score > best_score || (
score == best_score &&
left < best_left)) {
371 if (best_left < 0 || best_right < 0)
break;
372 if (node_count >= node_cap)
break;
374 SpmLlamaSymbol *
left = &symbols[best_left];
375 SpmLlamaSymbol *
right = &symbols[best_right];
377 int new_node_id = node_count++;
378 nodes[new_node_id].text =
left->text;
379 nodes[new_node_id].n =
left->n +
right->n;
380 nodes[new_node_id].left =
left->node_id;
381 nodes[new_node_id].right =
right->node_id;
384 left->node_id = new_node_id;
386 if (
right->next >= 0) {
387 symbols[
right->next].prev = best_left;
396 for (
int i = 0; i != -1 && out_idx <
max_ids; i = symbols[i].next) {
416 while (lead_spaces <
text_len &&
text[lead_spaces] ==
' ') {
421 int trail_spaces = 0;
422 while (trail_spaces <
text_len - lead_spaces &&
428 int content_len =
text_len - lead_spaces - trail_spaces;
429 int starts_with_prefix = (
text_len >= 3 &&
430 (
unsigned char)
text[0] == 0xE2 &&
431 (
unsigned char)
text[1] == 0x96 &&
432 (
unsigned char)
text[2] == 0x81);
433 int inserted_prefix = 0;
435 if (
out_len + 3 > out_max)
return -1;
444 int last_was_space = (starts_with_prefix || inserted_prefix) ? 1 : 0;
445 while (i <
text_len - trail_spaces) {
446 if (
text[i] ==
' ') {
447 if (!last_was_space) {
449 if (
out_len + 3 > out_max)
return -1;
457 if (
out_len + 1 > out_max)
return -1;
469 size_t pos, int32_t *candidates,
int max_candidates);
488 const int dbg = getenv(
"CK_DEBUG_SPM_ENCODE") ? 1 : 0;
490 fprintf(stderr,
"[SPM] encode start: text_len=%d max_ids=%d\n",
text_len,
max_ids);
494 char preprocessed[8192];
497 if (pp_len < 0)
return 0;
498 preprocessed[pp_len] =
'\0';
500 fprintf(stderr,
"[SPM] preprocessed len=%d: \"%.*s\"\n", pp_len, pp_len, preprocessed);
504 size_t n = (size_t)pp_len + 1;
505 float *best_score = (
float *)malloc(n *
sizeof(
float));
506 int32_t *best_prev = (int32_t *)malloc(n *
sizeof(int32_t));
507 int32_t *best_token = (int32_t *)malloc(n *
sizeof(int32_t));
509 fprintf(stderr,
"[SPM] DP alloc n=%zu\n", n);
512 if (!best_score || !best_prev || !best_token) {
513 if (best_score) free(best_score);
514 if (best_prev) free(best_prev);
515 if (best_token) free(best_token);
520 const float neg_inf = -1e30f;
521 const float unknown_penalty = -10.0f;
522 for (
size_t i = 0; i < n; i++) {
523 best_score[i] = neg_inf;
527 best_score[0] = 0.0f;
530 for (
size_t pos = 0; pos < n; pos++) {
531 if (best_score[pos] == neg_inf)
continue;
534 int32_t candidates[64];
536 if (dbg && pos < 8) {
537 fprintf(stderr,
"[SPM] pos=%zu cand=%d\n", pos, num_cand);
540 for (
int c = 0; c < num_cand; c++) {
541 int32_t token_id = candidates[c];
550 if (!
token)
continue;
553 int token_len = (int)strlen(
token);
556 if (token_id == tok->
unk_id) {
558 if (token_len == 0) token_len = 1;
561 size_t next_pos = pos + token_len;
563 if (next_pos >= n)
continue;
566 float token_score = 0.0f;
568 token_score = tok->
scores[token_id];
572 if (tok->
types && token_id >= 0 && token_id < (int32_t)tok->
types_size) {
578 if (token_id == tok->
unk_id) {
579 token_score += unknown_penalty;
583 float new_score = best_score[pos] + token_score;
585 if (new_score > best_score[next_pos]) {
586 best_score[next_pos] = new_score;
587 best_prev[next_pos] = (int32_t)pos;
588 best_token[next_pos] = token_id;
594 int32_t *reverse_ids = (int32_t *)malloc(
max_ids *
sizeof(int32_t));
603 int32_t curr = (int32_t)(n - 1);
606 while (curr > 0 && best_token[curr] < 0) {
607 curr = best_prev[curr];
613 while (curr > 0 && num_tokens <
max_ids) {
614 int32_t token_id = best_token[curr];
617 int token_start = best_prev[curr];
620 if (token_start != last_start) {
621 reverse_ids[num_tokens++] = token_id;
622 last_start = token_start;
625 curr = best_prev[curr];
628 fprintf(stderr,
"[SPM] backtrack tokens=%d curr=%d\n", num_tokens, curr);
639 for (
int i = 0; i < num_tokens / 2; i++) {
640 int32_t tmp = reverse_ids[i];
641 reverse_ids[i] = reverse_ids[num_tokens - 1 - i];
642 reverse_ids[num_tokens - 1 - i] = tmp;
647 for (
int i = 0; i < num_tokens && out_idx <
max_ids; i++) {
648 int32_t token_id = reverse_ids[i];
651 if (token_id == tok->
unk_id && out_idx > 0 &&
ids[out_idx - 1] == tok->
unk_id) {
654 ids[out_idx++] = token_id;
657 fprintf(stderr,
"[SPM] encode done: out=%d\n", out_idx);
663 if (num_tokens == 0) {
693 unsigned char byte_val = (
unsigned char)
text[i];
697 if (byte_token >= 0 && byte_token != tok->
unk_id) {
698 ids[count++] = byte_token;
708 size_t pos, int32_t *candidates,
int max_candidates) {
717 for (
int len =
max_len; len >= 1 && num_found < max_candidates; len--) {
718 memcpy(tmp,
text + pos, len);
722 if (info && info->id >= 0 && info->id != tok->
unk_id) {
730 for (
int j = 0; j < num_found; j++) {
731 if (candidates[j] == info->id) {
737 candidates[num_found++] = info->id;
746 if (num_found == 0 && tok->
unk_id >= 0 && max_candidates > 0) {
749 candidates[num_found++] = tok->
unk_id;
760 while (pos + run < (
size_t)
text_len) {
762 if (pos + run + 3 <= (
size_t)
text_len &&
763 (
unsigned char)
text[pos + run] == 0xE2 &&
764 (
unsigned char)
text[pos + run + 1] == 0x96 &&
765 (
unsigned char)
text[pos + run + 2] == 0x81) {
775 for (
int len =
max_len; len >= 1; len--) {
777 memcpy(tmp,
text + pos + run, len);
811 int segment_start = 0;
813 size_t special_len = 0;
815 if (special_id < 0 || special_len == 0) {
820 if (segment_start < pos) {
823 text + segment_start,
828 if (n <= 0)
return n;
833 ids[out_idx++] = special_id;
834 pos += (int)special_len;
841 text + segment_start,
846 if (n <= 0)
return n;
863 const uint8_t *types,
884 if (!tok->
scores)
return -1;
903 float score = scores ? scores[i] : 0.0f;
912 int count_normal = 0, count_unknown = 0, count_control = 0, count_byte = 0, count_other = 0;
915 uint8_t t = tok->
types[i];
916 if (t > max_type) max_type = t;
922 default: count_other++;
break;
925 fprintf(stderr,
"[TOKENIZER] Loaded %d tokens: normal=%d, unknown=%d, control=%d, byte=%d, other=%d\n",
926 vocab_size, count_normal, count_unknown, count_control, count_byte, count_other);
928 fprintf(stderr,
"[TOKENIZER] Warning: Unexpected token type %d\n", max_type);
int32_t ck_tokenizer_lookup(const CKTokenizer *tok, const char *token, int len)
int32_t ck_tokenizer_add_token(CKTokenizer *tok, const char *token, int len)
const char * ck_tokenizer_id_to_token(const CKTokenizer *tok, int32_t id)
void * ck_tokenizer_hash_table_lookup(CKTokenizerHashTable *table, const char *key)
CKTokenizerHashTable * vocab
struct CKTrieNode * children[256]
void ck_tokenizer_reset(CKTokenizer *tok)
#define CK_TOKENIZER_MAX_TOKEN_LEN
const int32_t int int * out_len
static int spm_find_candidates_at_pos(const CKTokenizer *tok, const char *text, int text_len, size_t pos, int32_t *candidates, int max_candidates)
static int32_t ck_tokenizer_lookup_exact(const CKTokenizer *tok, const char *token)
static int preprocess_spm_llama_text(const char *text, int text_len, char *out, int out_max, bool add_space_prefix)
static bool spm_token_is_byte_format(const char *token)
int ck_tokenizer_load_binary_with_scores(CKTokenizer *tok, int vocab_size, const int32_t *offsets, const char *strings, const float *scores, const uint8_t *types, int num_merges, const int32_t *merges)
static int ck_tokenizer_encode_spm_llama_impl(const CKTokenizer *tok, const char *text, int text_len, int32_t *ids, int max_ids)
static bool spm_token_allowed_in_dp(const CKTokenizer *tok, int32_t token_id)
#define GGUF_TOKEN_CONTROL
static void spm_build_byte_lookup(CKTokenizer *tok, const char *strings, const int32_t *offsets, int vocab_size)
static int preprocess_spm_text(const char *text, int text_len, char *out, int out_max, bool add_space_prefix)
static int spm_encode_byte_fallback(const CKTokenizer *tok, const char *text, int text_len, int32_t *ids, int max_ids)
int ck_tokenizer_encode_spm_dispatch(const CKTokenizer *tok, const char *text, int text_len, int32_t *ids, int max_ids)
static int spm_llama_resegment_node(const CKTokenizer *tok, const SpmLlamaNode *nodes, int node_id, int32_t *ids, int max_ids, int out_idx)
static int32_t ck_tokenizer_lookup_exact_n(const CKTokenizer *tok, const char *text, int text_len)
#define GGUF_TOKEN_USER_DEFINED
static int utf8_len(unsigned char c)
static int32_t spm_find_special_token_at_pos(const CKTokenizer *tok, const char *text, int text_len, size_t pos, size_t *match_len)
#define GGUF_TOKEN_UNUSED
static bool spm_is_byte_token(const CKTokenizer *tok, int32_t token_id)
static int32_t spm_get_byte_token(const CKTokenizer *tok, unsigned char byte_val)
static int spm_count_unknown_run(const CKTokenizer *tok, const char *text, int text_len, size_t pos)
#define GGUF_TOKEN_UNKNOWN
static int ck_tokenizer_encode_spm_impl(const CKTokenizer *tok, const char *text, int text_len, int32_t *ids, int max_ids)
#define GGUF_TOKEN_NORMAL
static int ck_tokenizer_encode_spm_plain_segment(const CKTokenizer *tok, const char *text, int text_len, int32_t *ids, int max_ids)
int const int32_t const char int num_merges
int const int32_t const char * strings
int const int32_t const char int const int32_t * merges
const int32_t int char int max_len
int const int32_t * offsets
const char int int32_t int max_ids
const char const char * right