以下代码实现了温度采样、Top-K 采样和 Top-P 采样,[1] 和 [2] 处应分别填入:
import torch import torch.nn.functional as F
def sample_with_strategies(logits, temperature=1.0, top_k=0, top_p=0.0): """ 对模型输出的 logits 应用温度、Top-K、Top-P 采样策略 参数: logits: 模型输出的原始分数,shape (vocab_size,) temperature: 温度系数,控制分布的平滑程度 top_k: 只保留概率最高的 k 个 token (0 表示不启用) top_p: 只保留累积概率达到 p 的最少 token 集合 (0.0 表示不启用)
返回: 采样得到的 token 索引 """ # [1] 温度缩放:用 temperature 对 logits 进行缩放 ____________________Top-K 采样:只保留概率最高的 k 个 token
if top_k > 0:
top_k_values, _ = torch.topk(scaled_logits, top_k)
min_top_k = top_k_values[-1]
scaled_logits = scaled_logits.masked_fill(scaled_logits < min_top_k, float('-inf'))Top-P (Nucleus) 采样
if top_p > 0.0:
sorted_logits, sorted_indices = torch.sort(scaled_logits, descending=True)
prob_sorted = F.softmax(sorted_logits, dim=-1)
cumulative_probs = torch.cumsum(prob_sorted, dim=-1)# [2] 创建掩码:移除累积概率超过 top_p 的 token,但保留第一个超过阈值的 token ____________________
[1] scaled_logits = logits / temperature
[2] sorted_indices_to_remove = cumulative_probs > top_p
[1] scaled_logits = logits * temperature
[2] sorted_indices_to_remove = cumulative_probs > top_p
[1] scaled_logits = logits / temperature
[2] sorted_indices_to_remove = cumulative_probs < top_p
[1] scaled_logits = logits - temperature
[2] sorted_indices_to_remove = probs_sorted > top_p