mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-27 21:46:57 +02:00
sampling : simplify temp sampling
This commit is contained in:
+4
-23
@@ -1597,32 +1597,13 @@ static void llama_sampler_backend_temp_sampling(
|
||||
struct ggml_tensor * max_idx = ggml_argmax(ctx, data->logits);
|
||||
ggml_set_name(max_idx, "temp_max_idx");
|
||||
|
||||
// Reshape to 2D and so we can use get_rows.
|
||||
struct ggml_tensor * logits_2d = ggml_reshape_2d(ctx, data->logits, 1, data->logits->ne[0]);
|
||||
ggml_set_name(logits_2d, "temp_logits_2d");
|
||||
struct ggml_tensor * max_logit = ggml_get_rows(ctx, logits_2d, max_idx);
|
||||
ggml_set_name(max_logit, "temp_max_logit");
|
||||
data->candidates = max_idx;
|
||||
|
||||
// Subtract the max_logit from all logits.
|
||||
struct ggml_tensor * diff = ggml_sub(ctx, data->logits, max_logit);
|
||||
ggml_set_name(diff, "temp_diff");
|
||||
struct ggml_tensor * logit = ggml_reshape_2d(ctx, data->logits, 1, data->logits->ne[0]);
|
||||
|
||||
// Add small epsilon to make max position strictly positive.
|
||||
struct ggml_tensor * diff_eps = ggml_scale_bias(ctx, diff, 1.0f, 1e-6f);
|
||||
ggml_set_name(diff_eps, "temp_diff_eps");
|
||||
|
||||
// Create the mask for the max logit.
|
||||
struct ggml_tensor * mask = ggml_step(ctx, diff_eps);
|
||||
ggml_set_name(mask, "temp_mask");
|
||||
|
||||
// Create the bias.
|
||||
const float large_val = 1e9f;
|
||||
struct ggml_tensor * bias = ggml_scale_bias(ctx, mask, large_val, -large_val);
|
||||
ggml_set_name(bias, "temp_bias");
|
||||
|
||||
// Add the bias to the logits.
|
||||
data->logits = ggml_add(ctx, data->logits, bias);
|
||||
data->logits = ggml_get_rows(ctx, logit, max_idx);
|
||||
ggml_build_forward_expand(gf, data->logits);
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user