Skip to content
Open
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
144 changes: 143 additions & 1 deletion common/sampling.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,76 @@ struct ring_buffer {
std::vector<T> data;
};

// Top k of the logits without building the full token array. It repeats the steps of libstdc++ std::partial_sort
// (make_heap on the first k, replace the heap top by every larger logit in index order, sort_heap) so that tokens with
// equal logits end up in the same order as with llama_sampler_top_k.
static void topk_from_logits(const float * logits, int32_t n_vocab, int32_t k, std::vector<llama_token_data> & out) {
static const auto comp = [](const llama_token_data & a, const llama_token_data & b) {
return a.logit > b.logit;
};

out.resize(k);
for (int32_t i = 0; i < k; i++) {
out[i] = llama_token_data{ i, logits[i], 0.0f };
}
std::make_heap(out.begin(), out.end(), comp);

llama_token_data * heap = out.data();

// std::__pop_heap(first, middle, result) with value = the new token: sift down to a leaf, then sift up
const auto replace_top = [&](const llama_token_data & value) {
ptrdiff_t hole = 0;
ptrdiff_t child = 0;
const ptrdiff_t len = k;
while (child < (len - 1) / 2) {
child = 2 * (child + 1);
if (comp(heap[child], heap[child - 1])) {
child--;
}
heap[hole] = heap[child];
hole = child;
}
if ((len & 1) == 0 && child == (len - 2) / 2) {
child = 2 * (child + 1);
heap[hole] = heap[child - 1];
hole = child - 1;
}
ptrdiff_t parent = (hole - 1) / 2;
while (hole > 0 && comp(heap[parent], value)) {
heap[hole] = heap[parent];
hole = parent;
parent = (hole - 1) / 2;
}
heap[hole] = value;
};

constexpr int32_t BLOCK = 16;
float thr = heap[0].logit;
int32_t i = k;
for (; i + BLOCK <= n_vocab; i += BLOCK) {
bool any = false;
for (int32_t j = 0; j < BLOCK; j++) {
any |= logits[i + j] > thr;
}
if (!any) {
continue;
}
for (int32_t j = 0; j < BLOCK; j++) {
if (logits[i + j] > heap[0].logit) {
replace_top(llama_token_data{ i + j, logits[i + j], 0.0f });
}
}
thr = heap[0].logit;
}
for (; i < n_vocab; i++) {
if (logits[i] > heap[0].logit) {
replace_top(llama_token_data{ i, logits[i], 0.0f });
}
}

std::sort_heap(out.begin(), out.end(), comp);
}

struct common_sampler {
common_params_sampling params;

Expand All @@ -122,6 +192,9 @@ struct common_sampler {

llama_token_data_array cur_p;

// > 0: the chain starts with top-k (all earlier samplers are no-ops), so cur_p can be built from the top k logits only
int32_t fast_topk = 0;

void reset() {
prev.clear();

Expand Down Expand Up @@ -162,6 +235,27 @@ struct common_sampler {
cur_p = { cur.data(), cur.size(), -1, false };
}

// cur_p = the k largest logits in descending order, same entries and order as top-k on the full array (set_logits + std::partial_sort)
// returns false if the logits are not plain CPU logits (backend sampling)
bool set_logits_topk(struct llama_context * ctx, int idx, int32_t k) {
if (llama_get_sampled_probs_ith(ctx, idx) || llama_get_sampled_logits_ith(ctx, idx)) {
return false;
}

const float * logits = llama_get_logits_ith(ctx, idx);
GGML_ASSERT(logits != nullptr);

const int32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(llama_get_model(ctx)));
if (n_vocab <= k) {
return false;
}

topk_from_logits(logits, n_vocab, k, cur);

cur_p = { cur.data(), cur.size(), -1, true };
return true;
}

common_time_meas tm() {
return common_time_meas(t_total_us, params.no_perf);
}
Expand All @@ -185,6 +279,39 @@ std::string common_params_sampling::print() const {
return std::string(result);
}

// k > 0 if the sampler chain starts with top-k (every sampler before it does nothing with these settings)
static int32_t common_sampler_fast_topk(const common_params_sampling & params, bool has_logit_bias) {
if (has_logit_bias || params.mirostat != 0 || params.backend_sampling || params.top_k <= 0 || params.top_k > 128) {
return 0;
}

for (const auto cnstr : params.samplers) {
switch (cnstr) {
case COMMON_SAMPLER_TYPE_TOP_K:
return params.top_k;
case COMMON_SAMPLER_TYPE_PENALTIES:
if (params.penalty_last_n != 0 && (params.penalty_repeat != 1.0f || params.penalty_freq != 0.0f || params.penalty_present != 0.0f)) {
return 0;
}
break;
case COMMON_SAMPLER_TYPE_DRY:
if (params.dry_multiplier != 0.0f && params.dry_base >= 1.0f && params.dry_penalty_last_n != 0) {
return 0;
}
break;
case COMMON_SAMPLER_TYPE_TOP_N_SIGMA:
if (params.top_n_sigma > 0.0f) {
return 0;
}
break;
default:
return 0;
}
}

return 0;
}

struct common_sampler * common_sampler_init(
const struct llama_model * model,
struct common_params_sampling & params) {
Expand Down Expand Up @@ -323,6 +450,8 @@ struct common_sampler * common_sampler_init(
}
}

bool has_logit_bias = false;

// logit bias: user biases + model suppress tokens (-INFINITY)
{
std::vector<llama_logit_bias> merged = params.logit_bias;
Expand All @@ -334,6 +463,7 @@ struct common_sampler * common_sampler_init(
}

if (!merged.empty()) {
has_logit_bias = true;
samplers.push_back(llama_sampler_init_logit_bias(llama_vocab_n_tokens(vocab), merged.size(), merged.data()));
}
}
Expand Down Expand Up @@ -435,6 +565,8 @@ struct common_sampler * common_sampler_init(
/* .cur_p = */ {},
};

result->fast_topk = common_sampler_fast_topk(result->params, has_logit_bias);

return result;
}

Expand Down Expand Up @@ -516,6 +648,7 @@ struct common_sampler * common_sampler_clone(common_sampler * gsmpl) {
/* .prev = */ gsmpl->prev,
/* .cur = */ gsmpl->cur,
/* .cur_p = */ gsmpl->cur_p,
/* .fast_topk = */ gsmpl->fast_topk,
};
}

Expand All @@ -536,6 +669,7 @@ void common_sampler_copy(const common_sampler * src, common_sampler * dst) {
dst->cur = src->cur;
dst->cur_p = src->cur_p;
dst->cur_p.data = src->cur_p.data ? dst->cur.data() : nullptr; // re-point to dst's buffer
dst->fast_topk = src->fast_topk;
dst->t_total_us = src->t_total_us;
}

Expand Down Expand Up @@ -605,7 +739,15 @@ llama_token common_sampler_sample(struct common_sampler * gsmpl, struct llama_co
auto & chain = gsmpl->chain;
auto & cur_p = gsmpl->cur_p; // initialized by set_logits

gsmpl->set_logits(ctx, idx);
// top-k directly from the logits when nothing before the top-k sampler reads the full array
const bool fast = gsmpl->fast_topk > 0 &&
!(grammar_first && grammar_should_apply(gsmpl)) &&
!(rbudget && common_reasoning_budget_get_state(rbudget) == REASONING_BUDGET_FORCING) &&
gsmpl->set_logits_topk(ctx, idx, gsmpl->fast_topk);

if (!fast) {
gsmpl->set_logits(ctx, idx);
}

// Check if a backend sampler has already sampled a token in which case we
// return that token id directly.
Expand Down