transformers 4.41.x中top_k_top_p_filtering函数替代方案问询
问题描述
在transformers 4.41.x版本中,top_k_top_p_filtering函数已被移除。此前版本中使用如下代码调用该函数:
next_token_logscores = top_k_top_p_filtering(logits, top_k=k, top_p=p)
其中k为保留元素数量,p为累积概率,该函数的具体实现如下:
def top_k_top_p_filtering( logits: Tensor, top_k: int = 0, top_p: float = 1.0, filter_value: float = -float("Inf"), min_tokens_to_keep: int = 1, ) -> Tensor: """Filter a distribution of logits using top-k and/or nucleus (top-p) filtering Args: logits: logits distribution shape (batch size, vocabulary size) if top_k > 0: keep only top k tokens with highest probability (top-k filtering). if top_p < 1.0: keep the top tokens with cumulative probability >= top_p (nucleus filtering). Nucleus filtering is described in Holtzman et al. (http://arxiv.org/abs/1904.09751) Make sure we keep at least min_tokens_to_keep per batch example in the output From: https://gist.github.com/thomwolf/1a5a29f6962089e871b94cbd09daf317 """ if top_k > 0: top_k = min(max(top_k, min_tokens_to_keep), logits.size(-1)) # Safety check # Remove all tokens with a probability less than the last token of the top-k indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None] logits[indices_to_remove] = filter_value if top_p < 1.0: sorted_logits, sorted_indices = torch.sort(logits, descending=True) cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1) # Remove tokens with cumulative probability above the threshold (token with 0 are kept) sorted_indices_to_remove = cumulative_probs > top_p if min_tokens_to_keep > 1: # Keep at least min_tokens_to_keep (set to min_tokens_to_keep-1 because we add the first one below) sorted_indices_to_remove[..., :min_tokens_to_keep] = 0 # Shift the indices to the right to keep also the first token above the threshold sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone() sorted_indices_to_remove[..., 0] = 0 # scatter sorted tensors to original indexing indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove) logits[indices_to_remove] = filter_value return logits
如何仅使用transformers 4.41.x提供的函数改写上述调用?
解决方案
transformers 4.41.x将原top_k_top_p_filtering的逻辑拆分为标准化的LogitsWarper实现,你可以通过组合TopKLogitsWarper和TopPLogitsWarper类实现完全等效的功能,步骤如下:
- 导入所需类:
from transformers import TopKLogitsWarper, TopPLogitsWarper from transformers.generation.logits_process import LogitsProcessorList import torch
- 创建处理器列表,组合top-k和top-p过滤逻辑:
logits_processors = LogitsProcessorList() # 添加top-k过滤(仅当k>0时) if k > 0: top_k_warper = TopKLogitsWarper( top_k=k, filter_value=-float("Inf"), min_tokens_to_keep=1 ) logits_processors.append(top_k_warper) # 添加top-p过滤(仅当p<1.0时) if p < 1.0: top_p_warper = TopPLogitsWarper( top_p=p, filter_value=-float("Inf"), min_tokens_to_keep=1 ) logits_processors.append(top_p_warper)
- 调用处理器完成logits过滤:
# 需传入与logits batch维度匹配的token_ids(无历史token时传空张量) batch_size = logits.size(0) dummy_token_ids = torch.zeros((batch_size, 0), dtype=torch.long, device=logits.device) next_token_logscores = logits_processors(logits, dummy_token_ids)
关键说明
TopKLogitsWarper和TopPLogitsWarper的参数与原函数一一对应,内部逻辑完全一致,包括安全检查、最小保留token数等细节。- 处理器按添加顺序执行过滤,先top-k后top-p,和原函数执行顺序保持一致。
- 若不需要某一种过滤(如
top_k=0或top_p=1.0),跳过对应处理器的添加即可。
内容的提问来源于stack exchange,提问作者Vladimir Canic
相关产品推荐
相关产品推荐

