-
Notifications
You must be signed in to change notification settings - Fork 72
Expand file tree
/
Copy pathreka.cpp
More file actions
111 lines (96 loc) · 3.64 KB
/
Copy pathreka.cpp
File metadata and controls
111 lines (96 loc) · 3.64 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
#include "llama.h"
namespace chatllm::reka::flash
{
typedef llama::v3::Config Config;
class ChatHistoryEncoder : public BaseHistoryEncoder
{
public:
void append_sys_prompt(std::vector<int> &ids) const override
{
}
void append_ai(int round_idx, const std::string &ai, std::vector<int> &ids) const override
{
std::ostringstream oss;
oss << "assistant: " << ai << " <sep> ";
tokenizer->encode(oss.str(), ids);
}
void append_user(int round_idx, const std::string &user, std::vector<int> &ids) const override
{
std::ostringstream oss;
oss << "human: ";
if ((0 == round_idx) && (tokenizer->get_system_prompt().size() > 0))
{
oss << tokenizer->get_system_prompt() << " ";
}
oss << user << " <sep> ";
tokenizer->encode(oss.str(), ids);
}
void append_ai_opening(int round_idx, std::vector<int> &ids) const override
{
std::ostringstream oss;
oss << "assistant: ";
tokenizer->encode(oss.str(), ids);
}
};
static ChatHistoryEncoder _chat_encoder;
class Tokenizer : public BaseTokenizer
{
public:
Tokenizer(const Config &config)
: Tokenizer(config, &_chat_encoder)
{}
Tokenizer(const BaseConfig &config, BaseHistoryEncoder *encoder,
BaseHistoryEncoder *qa_encoder = nullptr,
BaseHistoryEncoder *completion_encoder = nullptr)
: BaseTokenizer::BaseTokenizer(config, encoder, qa_encoder, completion_encoder)
{
sys_prompt = "";
}
size_t load(tokenizer::DataReader *buffer, int n_vocab) override
{
tp = new tokenizer::BPEProcessor2(
{
"(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
}
);
size_t size = tp->Load(buffer, n_vocab);
tp->Encode(" <sep>", &chat_terminate_seq);
return size;
}
public:
std::vector<int> chat_terminate_seq;
};
class ConditionalGeneration : public llama::v3::ConditionalGeneration
{
public:
ConditionalGeneration() = default;
ConditionalGeneration(const Config &config, const RuntimeConfig &runtime_config)
: llama::v3::ConditionalGeneration(config, runtime_config, ModelType::MODEL_TYPE_REKA_FLASH3)
{}
protected:
bool is_output_terminated(const std::vector<int> &output_ids, int &keep_idx, int &pop_output) override
{
if (output_ids.size() < 1) return false;
Tokenizer *tokenizer = dynamic_cast<Tokenizer *>(this->tokenizer);
int len = 0;
switch (tokenizer->get_chat_format())
{
case ChatFormat::CHAT:
if (match_output_sequence(output_ids, tokenizer->chat_terminate_seq))
{
pop_output = (int)tokenizer->chat_terminate_seq.size();
return true;
}
len = (int)tokenizer->chat_terminate_seq.size();
break;
default:
;
}
if (BaseModelForConditionalGeneration::is_output_terminated(output_ids, keep_idx, pop_output))
return true;
keep_idx = (int)(output_ids.size()) - len + 1;
return false;
}
};
REGISTER_MODEL_LOADER(REKA_FLASH3, reka::flash, 1);
}