-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.cpp
More file actions
126 lines (110 loc) · 3.87 KB
/
Copy pathmain.cpp
File metadata and controls
126 lines (110 loc) · 3.87 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
import std;
import llamacpp;
namespace {
void log_callback(enum ggml_log_level, const char * text, void *) {
if (text) std::cerr << text;
}
int usage(const char * program) {
std::cerr << "usage: " << program
<< " <model.gguf> [prompt]\n";
return 2;
}
} // namespace
int main(int argc, char ** argv) {
if (argc != 2 && argc != 3) return usage(argv[0]);
const std::string prompt = argc == 3
? argv[2]
: "User: Hello! Who are you?\nAssistant:";
llama_log_set(log_callback, nullptr);
llama_backend_init();
llama_model_params model_params = llama_model_default_params();
model_params.n_gpu_layers = 0;
llama_model * model = llama_model_load_from_file(argv[1], model_params);
if (!model) {
std::cerr << "failed to load model: " << argv[1] << '\n';
llama_backend_free();
return 4;
}
const llama_vocab * vocabulary = llama_model_get_vocab(model);
const int prompt_size = -llama_tokenize(
vocabulary, prompt.data(), prompt.size(), nullptr, 0, true, true
);
if (prompt_size <= 0) {
std::cerr << "failed to measure prompt tokens\n";
llama_model_free(model);
llama_backend_free();
return 5;
}
std::vector<llama_token> prompt_tokens(prompt_size);
if (llama_tokenize(
vocabulary,
prompt.data(),
prompt.size(),
prompt_tokens.data(),
prompt_tokens.size(),
true,
true
) < 0) {
std::cerr << "failed to tokenize prompt\n";
llama_model_free(model);
llama_backend_free();
return 5;
}
constexpr int max_generated_tokens = 32;
llama_context_params context_params = llama_context_default_params();
context_params.n_ctx = prompt_size + max_generated_tokens;
context_params.n_batch = prompt_size;
llama_context * context = llama_init_from_model(model, context_params);
if (!context) {
std::cerr << "failed to create context\n";
llama_model_free(model);
llama_backend_free();
return 6;
}
llama_sampler * sampler = llama_sampler_chain_init(
llama_sampler_chain_default_params()
);
llama_sampler_chain_add(sampler, llama_sampler_init_top_k(40));
llama_sampler_chain_add(sampler, llama_sampler_init_top_p(0.9F, 1));
llama_sampler_chain_add(sampler, llama_sampler_init_temp(0.8F));
llama_sampler_chain_add(sampler, llama_sampler_init_dist(1234));
llama_batch batch = llama_batch_get_one(
prompt_tokens.data(), prompt_tokens.size()
);
std::cout << prompt;
std::cout.flush();
int generated = 0;
int exit_code = 0;
llama_token sampled = LLAMA_TOKEN_NULL;
for (; generated < max_generated_tokens; ++generated) {
const int decode_result = llama_decode(context, batch);
if (decode_result != 0) {
std::cerr << "\ndecode failed: " << decode_result << '\n';
exit_code = 7;
break;
}
sampled = llama_sampler_sample(sampler, context, -1);
if (llama_vocab_is_eog(vocabulary, sampled)) break;
char piece[256] = {};
const int piece_size = llama_token_to_piece(
vocabulary, sampled, piece, sizeof(piece), 0, true
);
if (piece_size < 0 || piece_size > static_cast<int>(sizeof(piece))) {
std::cerr << "\nfailed to render sampled token " << sampled << '\n';
exit_code = 8;
break;
}
std::cout << std::string_view(piece, piece_size);
std::cout.flush();
batch = llama_batch_get_one(&sampled, 1);
}
std::cout << '\n';
std::cerr << "backend=cpu"
<< " params=" << llama_model_n_params(model)
<< " generated_tokens=" << generated << '\n';
llama_sampler_free(sampler);
llama_free(context);
llama_model_free(model);
llama_backend_free();
return exit_code;
}