You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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类实现完全等效的功能,步骤如下:

  1. 导入所需类:
from transformers import TopKLogitsWarper, TopPLogitsWarper
from transformers.generation.logits_process import LogitsProcessorList
import torch
  1. 创建处理器列表,组合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)
  1. 调用处理器完成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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.21 22:05:00