Repository navigation
feat(generation): min-p sampling and OpenAI-style frequency/presence penalties - #3731
RajaBabu15 wants to merge 1 commit into
Conversation
Add Sampling::MinP { p, temperature } to LogitsProcessor: keep tokens
whose probability is at least p times that of the most likely token,
then sample from the renormalized survivors (arXiv:2407.01082).
Add apply_frequency_presence_penalty with OpenAI semantics:
logit[t] -= frequency_penalty * count(t) + presence_penalty, distinct
from utils::apply_repeat_penalty's multiplicative rescaling.
Both are expected by OpenAI-compatible servers built on candle.
Fixes huggingface#3654
…nd OpenAI-style frequency/presence penalties
|
Thanks for this — min-p sampling and OpenAI-compatible frequency/presence penalties are both useful additions for servers built on candle, and the implementation is solid: sample_minp correctly reuses the existing "zero out below threshold, let WeightedIndex renormalize" pattern already established by sample_topp, and the penalty helper correctly deduplicates context tokens via a HashMap so presence penalty isn't accidentally scaled by occurrence count. One precision issue found while integrating this into a fork: apply_frequency_presence_penalty computes the adjustment in f32 but then casts the result back to the caller's original dtype before returning: Tensor::from_vec(logits_v, logits_len, device)?.to_dtype(dtype) Suggested fix: drop the trailing .to_dtype(dtype) (and the now-unused dtype variable), including in the zero-penalty no-op branch, so the function's return dtype is uniformly f32 regardless of whether a penalty was applied — matching apply_repeat_penalty's convention. Pushed this fix on top of the PR in my fork, with a regression test reproducing the f16 rounding scenario above: astorise/candle@fcd9d03. Happy to open it as a PR against this branch or main if useful. |
Upstream huggingface#3731: min-p sampling and OpenAI-style frequency/presence penalties
Summary
Implements #3654:
min_psampling and OpenAI-style frequency/presence penalties incandle_transformers::generation— both expected by anyone exposing an OpenAI-compatible API on top of candle.Sampling::MinP { p, temperature }Min-p sampling (arXiv:2407.01082) keeps tokens whose probability is at least
p * max_prob, then samples from the renormalized survivors. Wired throughLogitsProcessor::samplelike the existing TopK/TopP variants;p <= 0degrades gracefully to plain multinomial sampling.apply_frequency_presence_penaltyOpenAI semantics, applied to logits before sampling:
Distinct from
utils::apply_repeat_penalty(multiplicative, count-insensitive). Out-of-vocab context tokens are ignored; zero penalties are a no-op returning the input tensor. Input dtype is preserved.Tests
sample_with_min_p_dominant_token— deterministic across seeds when one token dominatessample_with_min_p_filters_and_renormalizes— sub-threshold token never sampled; survivor odds match the renormalized distribution over 10k drawssample_with_min_p_zero_keeps_all_tokens—p = 0keeps the full distributionfrequency_presence_penalty_matches_openai_semantics— exact penalty arithmeticfrequency_presence_penalty_noop_and_out_of_range— zero-penalty no-op, out-of-vocab ids ignoredAll existing tests pass;
cargo clippy -p candle-transformers --tests -- -D warningsandcargo fmt --checkare clean.Fixes #3654