Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
76 changes: 54 additions & 22 deletions common/common.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1503,7 +1503,8 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
}

if (llama_model_has_encoder(model)) {
llama_encode(lctx, llama_batch_get_one(tmp.data(), tmp.size()));
common_batch batch = common_batch_get_one(lctx, tmp);
llama_process(lctx, LLAMA_PROCESS_TYPE_ENCODE, batch.get());
llama_token decoder_start_token_id = llama_model_decoder_start_token(model);
if (decoder_start_token_id == LLAMA_TOKEN_NULL) {
decoder_start_token_id = bos;
Expand All @@ -1512,7 +1513,9 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
tmp.push_back(decoder_start_token_id);
}
if (llama_model_has_decoder(model)) {
llama_decode(lctx, llama_batch_get_one(tmp.data(), std::min(tmp.size(), (size_t) params.n_batch)));
tmp.resize(std::min(tmp.size(), (size_t) params.n_batch));
common_batch batch = common_batch_get_one(lctx, tmp);
llama_process(lctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
llama_memory_clear(llama_get_memory(lctx), true);
llama_synchronize(lctx);
Expand Down Expand Up @@ -1571,9 +1574,13 @@ common_context_seq_rm_type common_context_can_seq_rm(llama_context * ctx) {
tmp.push_back(0);
tmp.push_back(0);

int ret = llama_decode(ctx, llama_batch_get_one(tmp.data(), tmp.size()));
int ret;
{
common_batch batch = common_batch_get_one(ctx, tmp);
ret = llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get());
}
if (ret != 0) {
COM_ERR("llama_decode() failed: %d\n", ret);
COM_ERR("llama_process() failed: %d\n", ret);
res = COMMON_CONTEXT_SEQ_RM_TYPE_NO;
goto done;
}
Expand Down Expand Up @@ -2163,31 +2170,58 @@ float lr_opt::get_lr(float epoch) const {
}

bool common_replay_last_token(struct llama_context * ctx, llama_token last_token, int32_t pos) {
llama_batch batch = llama_batch_get_one(&last_token, 1);
batch.pos = &pos;
if (llama_decode(ctx, batch)) {
common_batch batch(ctx);
batch.add(last_token, pos, 0, true);

if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
LOG_ERR("%s: failed to replay last token\n", __func__);
return false;
}
return true;
}

llama_batch_ext_ptr common_batch_ext_get_one(llama_context * ctx, const llama_tokens & tokens) {
llama_batch_ext_ptr batch(llama_batch_ext_init(ctx));
void common_batch::clear() {
tokens.clear();
llama_batch_ext_clear(batch.get());
}

int32_t common_batch::add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output) {
const int32_t idx = llama_batch_ext_add_token(batch.get(), seq_id, id);
if (idx < 0) {
return idx;
}
llama_batch_ext_set_pos(batch.get(), idx, &pos);
if (output) {
llama_batch_ext_set_output_logits(batch.get(), idx, true);
}
tokens.push_back({ id, pos, seq_id, output });
return idx;
}

bool common_batch::set_output(int32_t idx, bool value) {
if (idx < 0 || idx >= (int32_t) tokens.size()) {
return false;
}
tokens[idx].output = value;
return llama_batch_ext_set_output_logits(batch.get(), idx, value);
}

bool common_batch::set_embd(int32_t idx, llama_embd embd) {
return llama_batch_ext_set_embd_token(batch.get(), idx, embd);
}

common_batch common_batch_get_one(llama_context * ctx, const llama_tokens & tokens) {
common_batch batch(ctx);

auto mem = llama_get_memory(ctx);
llama_pos pos = mem ? llama_memory_seq_pos_max(mem, 0) + 1 : 0;
llama_pos pos = llama_memory_seq_pos_max(mem, 0) + 1; // -1 + 1 == 0 when the memory is empty

for (size_t i = 0; i < tokens.size(); ++i) {
const int32_t idx = llama_batch_ext_add_token(batch.get(), 0, tokens[i]);
llama_batch_ext_set_pos(batch.get(), idx, &pos);
const bool output = i == tokens.size() - 1;
batch.add(tokens[i], pos, 0, output);
pos++;
}

if (!tokens.empty()) {
llama_batch_ext_set_output_logits(batch.get(), (int32_t) tokens.size() - 1, true);
}

return batch;
}

Expand Down Expand Up @@ -2215,7 +2249,7 @@ bool common_prompt_batch_decode(
// memory, so we can't just remove the last token from the memory and replay the last token which
// is the reason for this logic.
llama_tokens prefix_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_tokens_before_last);
llama_batch_ext_ptr batch_prefix = common_batch_ext_get_one(ctx, prefix_tokens);
common_batch batch_prefix = common_batch_get_one(ctx, prefix_tokens);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_prefix.get())) {
COM_ERR("%s", "failed to eval\n");
return false;
Expand All @@ -2225,10 +2259,8 @@ bool common_prompt_batch_decode(
llama_state_save_file(ctx, state_path.data(), all_tokens.data(), all_tokens.size());
COM_INF("saved session before last token to %s, n_new = %zu\n", state_path.data(), all_tokens.size());

llama_token last_token = all_tokens.back();
llama_batch_ext_ptr batch_last = common_batch_ext_get_one(ctx, { last_token });
llama_pos pos = n_past;
llama_batch_ext_set_pos(batch_last.get(), 0, &pos);
common_batch batch_last(ctx);
batch_last.add(all_tokens.back(), n_past, 0, true);

if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch_last.get())) {
COM_ERR("%s", "failed to eval last token\n");
Expand All @@ -2237,7 +2269,7 @@ bool common_prompt_batch_decode(
n_past++;
} else {
llama_tokens new_tokens(all_tokens.begin() + offset, all_tokens.begin() + offset + n_new);
llama_batch_ext_ptr batch = common_batch_ext_get_one(ctx, new_tokens);
common_batch batch = common_batch_get_one(ctx, new_tokens);
if (llama_process(ctx, LLAMA_PROCESS_TYPE_DECODE, batch.get())) {
COM_ERR("%s", "failed to eval\n");
return false;
Expand Down
32 changes: 31 additions & 1 deletion common/common.h
Original file line number Diff line number Diff line change
Expand Up @@ -1003,9 +1003,39 @@ void common_batch_add(
const std::vector<llama_seq_id> & seq_ids,
bool logits);

// wrapper around llama_batch_ext that provide getter functions for downstream code
struct common_batch {
struct token {
llama_token id;
llama_pos pos;
llama_seq_id seq_id;
bool output;
};

std::vector<token> tokens; // mirror of the entries, tokens[i] describes batch index i
llama_batch_ext_ptr batch;

common_batch() = default;
common_batch(struct llama_context * ctx) : batch(llama_batch_ext_init(ctx)) {}

llama_batch_ext * get() const { return batch.get(); }

void clear();

// returns the batch index (>= 0), or a negative error from llama_batch_ext_add_token()
int32_t add(llama_token id, llama_pos pos, llama_seq_id seq_id, bool output);

bool set_output(int32_t idx, bool value);

// attach a token embedding to the entry at idx, can only be set once per entry
bool set_embd(int32_t idx, llama_embd embd);

int32_t size() const { return (int32_t) tokens.size(); }
};

// create a single-sequence batch from a list of tokens
// last token always have output_logits set to true
llama_batch_ext_ptr common_batch_ext_get_one(struct llama_context * ctx, const llama_tokens & tokens);
common_batch common_batch_get_one(struct llama_context * ctx, const llama_tokens & tokens);

// decodes a single batch of tokens for a prompt and manages session tokens
//
Expand Down
Loading
Loading