power law sampler added

This commit is contained in:
Concedo
2025-12-27 09:46:06 +08:00
parent dfa1b72d2f
commit 91d8863f18
5 changed files with 987 additions and 110 deletions
+34 -19
View File
@@ -125,6 +125,8 @@ static std::vector<gpt_vocab::id> current_context_tokens;
static size_t mem_per_token = 0;
static std::vector<float> logits;
static std::vector<int> smartcontext;
static float power_law_weighted_sum = 0; //power law sampling state vars
static float power_law_total_weight = 0;
static std::vector<std::string> stop_sequence;
static std::vector<int> special_stop_sequence; //for stop sequences that don't have a string representation
static std::vector<std::string> banned_tokens;
@@ -1271,39 +1273,41 @@ float & weighted_sum, // persistent EMA state
float & total_weight, // persistent EMA state
llama_token_data_array * cur_p)
{
const float width = 0.3; // DISTRIBUTION_WIDTH
const float peak_logit = 5.0; // PEAK_LOGIT_VALUE
const float width = 0.3; // DISTRIBUTION_WIDTH
const float peak_logit = 5.0; // PEAK_LOGIT_VALUE
const float inv_width = 1.0f / width; // INV_WIDTH
if (target <= 0.0f || cur_p->size == 0) {
return;
}
const float inv_width = 1.0f / width;
// target is the desired average probability for selected tokens (0.0 to 1.0)
// higher values favor more probable tokens (more stable and predictable)
// lower values favor less probable tokens (more creative)
// Step 1: softmax to get original probabilities
sample_softmax(cur_p);
// Step 2: compute adaptive target (EMA feedback)
float computed_target;
if (total_weight == 0.0f) {
computed_target = target;
} else {
computed_target = 2.0f * target - (weighted_sum / total_weight);
computed_target = std::clamp(computed_target, 0.0f, 1.0f);
}
// compute the adapted target probability for the current sampling step
float computed_target = std::clamp(total_weight == 0.0f ? target : 2.0f * target - (weighted_sum / total_weight),0.0f, 1.0f);
// Step 3: apply power-law shaping in logit space
// power law transform
for (size_t i = 0; i < cur_p->size; ++i) {
float dist = (cur_p->data[i].p - computed_target) * inv_width;
float score = peak_logit / (1.0f + dist * dist);
cur_p->data[i].logit = score;
float dist = (cur_p->data[i].p - computed_target) * inv_width;
cur_p->data[i].logit = (peak_logit / (1.0f + dist * dist));
}
cur_p->sorted = false;
sample_softmax(cur_p);
// Step 4: update EMA history AFTER sampling, update_power_law_history(original_prob[idx])
//update EMA history AFTER sampling, update_power_law_history(original_prob[idx])
}
inline void power_law_update_history(float selected_token_prob, float & weighted_sum, float & total_weight) {
// decay controls how quickly history influence fades (0.0 to 0.99)
// lower values = faster adaptation, more reactive to recent tokens
// higher values = slower adaptation, more stable over time
// effective history length ≈ 1/(1-decay) tokens
// example: decay=0.5 --> ~2 tokens; decay=0.9 --> ~10 tokens; decay=0.95 --> ~20 tokens
// keep <= 0.99 to prevent unbounded accumulation
const float power_law_decay = 0.90f;
weighted_sum = selected_token_prob + power_law_decay * weighted_sum;
total_weight = 1.0f + power_law_decay * total_weight;
@@ -1735,7 +1739,7 @@ void sample_guidance(struct llama_context * ctx, struct llama_context * guidance
int SampleLogits(const float * logits, int n_ctx, int n_vocab, int rep_pen_range, float rep_pen, float rep_pen_slope, float presence_penalty, float top_k, float top_a, float top_p, float min_p, float typical_p, float tfs, float nsigma, float temp, std::mt19937 & rng,
int mirostat, float mirostat_tau, float mirostat_eta, float dry_multiplier, float dry_base, int dry_allowed_length, int dry_penalty_last_n, float xtc_threshold, float xtc_probability,
const std::vector<samplers> & sampler_order, llama_grammar * grammar, float dynatemp_range, float dynatemp_exponent, float smoothing_factor, float smoothing_curve)
const std::vector<samplers> & sampler_order, llama_grammar * grammar, float dynatemp_range, float dynatemp_exponent, float smoothing_factor, float smoothing_curve, float power_law_target)
{
// printf("SampleLogits called with: n_ctx=%d, n_vocab=%d, rep_pen_range=%d, rep_pen=%f, rep_pen_slope=%f, presence_penalty=%f, top_k=%f, top_a=%f, top_p=%f, min_p=%f, typical_p=%f, tfs=%f, nsigma=%f, temp=%f, mirostat=%d, mirostat_tau=%f, mirostat_eta=%f, dry_multiplier=%f, dry_base=%f, dry_allowed_length=%d, dry_penalty_last_n=%d, xtc_threshold=%f, xtc_probability=%f, sampler_order_size=%zu, dynatemp_range=%f, dynatemp_exponent=%f, smoothing_factor=%f\n",
// n_ctx, n_vocab, rep_pen_range, rep_pen, rep_pen_slope, presence_penalty, top_k, top_a, top_p, min_p, typical_p, tfs, nsigma, temp, mirostat, mirostat_tau, mirostat_eta, dry_multiplier, dry_base, dry_allowed_length, dry_penalty_last_n, xtc_threshold, xtc_probability, sampler_order.size(), dynatemp_range, dynatemp_exponent, smoothing_factor);
@@ -1841,6 +1845,8 @@ const std::vector<samplers> & sampler_order, llama_grammar * grammar, float dyna
}
//xtc always last
sample_xtc(&candidates_p, xtc_threshold, xtc_probability, rng);
//power law must be last, it messes up all probs
sample_power_law(power_law_target, power_law_weighted_sum, power_law_total_weight, &candidates_p);
id = sample_token(&candidates_p, rng);
}
@@ -3436,6 +3442,9 @@ generation_outputs gpttype_generate(const generation_inputs inputs)
}
}
power_law_weighted_sum = 0;
power_law_total_weight = 0;
//handle custom token bans and antislop phrase banning
banned_phrases.clear();
delayed_generated_tokens_limit = 0;
@@ -3644,6 +3653,7 @@ generation_outputs gpttype_generate(const generation_inputs inputs)
kcpp_data->n_ctx = inputs.max_context_length;
kcpp_data->smoothing_factor = inputs.smoothing_factor;
kcpp_data->smoothing_curve = inputs.smoothing_curve;
kcpp_data->power_law_target = inputs.power_law_target;
// Parse dry sequence breakers / restart sequences
kcpp_data->dry_sequence_breakers.clear();
@@ -4606,7 +4616,12 @@ generation_outputs gpttype_generate(const generation_inputs inputs)
kcpp_data->mirostat, kcpp_data->mirostat_tau, kcpp_data->mirostat_eta,
kcpp_data->dry_multiplier, kcpp_data->dry_base,
kcpp_data->dry_allowed_length, kcpp_data->dry_penalty_last_n, kcpp_data->xtc_threshold, kcpp_data->xtc_probability,
sampler_order, grammar, dynatemp_range, dynatemp_exponent, smoothing_factor, smoothing_curve);
sampler_order, grammar, dynatemp_range, dynatemp_exponent, smoothing_factor, smoothing_curve, power_law_target);
if (power_law_target > 0.0f) {
float original_prob = original_candidates[id].p;
power_law_update_history(original_prob, power_law_weighted_sum, power_law_total_weight);
}
if(draft_used)
{