revised power law sampling

This commit is contained in:
Concedo
2025-12-27 10:59:46 +08:00
parent 91d8863f18
commit 9bb362cce9
+3 -1
View File
@@ -1291,9 +1291,11 @@ llama_token_data_array * cur_p)
float computed_target = std::clamp(total_weight == 0.0f ? target : 2.0f * target - (weighted_sum / total_weight),0.0f, 1.0f);
// power law transform
const float k = 4.0f; // controls sharpness
for (size_t i = 0; i < cur_p->size; ++i) {
float dist = (cur_p->data[i].p - computed_target) * inv_width;
cur_p->data[i].logit = (peak_logit / (1.0f + dist * dist));
float abs_dist = fabs(dist);
cur_p->data[i].logit = peak_logit - k * abs_dist * (abs_dist / (1.0f + abs_dist));
}
cur_p->sorted = false;