mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-08-11 01:16:31 +00:00
llama: move suppress_tokens handling to common/sampling (#26276)
* llama: move suppress_tokens handling to common/sampling * address security issues * rm has_logit_bias
This commit is contained in:
parent
caa596ab3f
commit
afeebe103b
5 changed files with 32 additions and 44 deletions
|
|
@ -142,33 +142,6 @@ static ggml_tensor * ggml_view_2d_slice(ggml_context * ctx0, ggml_tensor * x, in
|
|||
idx * x->ne[0] * x->ne[1] * ggml_element_size(x));
|
||||
}
|
||||
|
||||
// TODO @ngxson : maybe improve this in the future
|
||||
class llm_graph_input_logits_bias : public llm_graph_input_i {
|
||||
public:
|
||||
llm_graph_input_logits_bias(const llama_vocab & vocab) {
|
||||
arr.resize(vocab.n_tokens(), 0.0f);
|
||||
for (llama_token id : vocab.get_suppress_tokens()) {
|
||||
if (0 <= id && id < (int32_t)vocab.n_tokens()) {
|
||||
arr[id] = -INFINITY;
|
||||
}
|
||||
}
|
||||
}
|
||||
virtual ~llm_graph_input_logits_bias() = default;
|
||||
|
||||
void set_input(const llama_ubatch * /*ubatch*/) override {
|
||||
const int64_t n_vocab = arr.size();
|
||||
ggml_backend_tensor_set(logits_bias, arr.data(), 0, n_vocab*ggml_element_size(logits_bias));
|
||||
}
|
||||
|
||||
bool can_reuse(const llm_graph_params & /*params*/) override {
|
||||
return true;
|
||||
}
|
||||
|
||||
ggml_tensor * logits_bias = nullptr; // F32 [n_vocab]
|
||||
|
||||
std::vector<float> arr;
|
||||
};
|
||||
|
||||
llama_model_gemma4::graph::graph(const llama_model & model, const llm_graph_params & params) :
|
||||
llm_graph_context(params),
|
||||
model(model),
|
||||
|
|
@ -429,16 +402,6 @@ llama_model_gemma4::graph::graph(const llama_model & model, const llm_graph_para
|
|||
cur = ggml_scale(ctx0, cur, hparams.f_final_logit_softcapping);
|
||||
}
|
||||
|
||||
// apply logits bias if needed (e.g. for gemma4_unified patch)
|
||||
// this is to mirror the suppress_tokens patch on transformers, to avoid model from outputing <image|> and <audio|> tokens (which is a known issue related to the checkpoint)
|
||||
// TODO: maybe handle this inside the sampling system in the future
|
||||
if (!model.vocab.get_suppress_tokens().empty()) {
|
||||
auto inp_bias = std::make_unique<llm_graph_input_logits_bias>(model.vocab);
|
||||
inp_bias->logits_bias = ggml_new_tensor_1d(ctx0, GGML_TYPE_F32, inp_bias->arr.size());
|
||||
cur = ggml_add(ctx0, cur, inp_bias->logits_bias);
|
||||
res->add_input(std::move(inp_bias));
|
||||
}
|
||||
|
||||
cb(cur, "result_output", -1);
|
||||
res->t_logits = cur;
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue