大模型手撕-核心题¶
实现贪心解码函数 greedy_decode(logits),逐步选取概率最大的 token 拼接成序列。¶
1.1 原理¶
贪心解码(Greedy Decoding)是自回归生成中最简单、最确定性的策略。在每一个生成步骤,模型输出一个形状为 (vocab_size,) 的 logits 张量,经过 softmax 后得到概率分布。贪心解码直接选择概率最高的 token 作为本步的输出,将其拼接到已生成序列的末尾,然后该新序列作为下一步的输入继续生成,直到遇到结束符(EOS)或达到预设的最大长度。
贪心解码不进行任何随机采样,因此生成结果是完全确定的(假设模型固定且无随机性 dropout 等)。其优点是速度极快,无需维护候选集;缺点是缺乏全局最优性,容易陷入局部最优,在长文本生成中常导致重复和退化(如循环输出相同短语)。
1.2 实现¶
import torch
def greedy_decode(model, input_ids, max_length=50, eos_token_id=None):
"""
贪心解码自回归生成。
Args:
model: 自回归语言模型,调用 model(input_ids) 返回 logits (batch_size, seq_len, vocab_size)。
假设输入 batch_size=1,已处理 prompt。
input_ids: (1, prompt_len) 初始提示词 token ID 张量。
max_length: 最大生成总长度(包括 prompt)。
eos_token_id: 结束符 token ID,遇到则停止。
Returns:
generated: 生成的完整 token ID 列表(包括 prompt)。
"""
model.eval()
generated = input_ids.tolist()[0] # 转为Python列表
with torch.no_grad():
for _ in range(max_length - len(generated)):
# 将当前序列转为tensor输入模型
inputs = torch.tensor([generated], device=input_ids.device)
outputs = model(inputs)
# 取最后一个位置的logits
next_logits = outputs[0, -1, :] # shape: (vocab_size,)
# 贪心选择最大概率的token
next_token_id = torch.argmax(next_logits).item()
generated.append(next_token_id)
# 如果遇到eos则停止
if eos_token_id is not None and next_token_id == eos_token_id:
break
return generated
1.3 关键细节¶
-
KV Cache 优化:上述实现每次都将整个历史序列重新输入模型,计算量随生成长度平方增长。生产环境中必须使用 KV Cache:在每步仅输入最新 token,利用缓存的历史 Key/Value 张量进行注意力计算。实现 KV Cache 需要模型支持
past_key_values参数。这里为了清晰展示解码逻辑,采用简化形式,但面试中应主动指出此优化。 -
模型调用约定:通常自回归模型的
forward接受input_ids和past_key_values,返回logits和新的past_key_values。上述函数可扩展为接受past_key_values参数并返回更新后的缓存。 -
EOS 处理:遇到 EOS 时应当立即停止,避免生成无意义的后续 token。同时应确保生成的序列中不会因为 EOS 后的 token 而被截断。
-
Tensor 操作:使用
torch.argmax直接获取索引,高效简洁。
1.4 贪心解码的局限与改进方向¶
-
无法产生多样性输出。
-
容易陷入重复循环,尤其在开放域对话中。
-
对于需要探索的任务,通常与温度缩放、Top-k/Top-p 结合使用,或在束搜索中作为最终评分基准。
实现基础束搜索 beam_search(logits_fn, k, max_len),每步扩展并保留 k 条最优序列。¶
2.1 原理¶
束搜索(Beam Search)在每一步维护一个大小为 k(束宽)的候选序列集合,而不是仅保留一个最佳序列。具体流程:
-
初始化:起始序列为 prompt,得分(对数概率和)为 0。
-
对当前束中的每一条序列,模型给出下一步的 logits,计算所有可能的下一个 token 的对数概率。
-
将所有
(当前序列, 下一个token)的组合视为新候选,共k * |V|个(实际常只取每个序列的 top-k 候选以加速),每个候选的得分为原序列得分 + log P(token|sequence)。 -
按得分降序排列,保留得分最高的
k个候选作为下一轮的束。 -
重复直到满足终止条件:所有序列都生成了 EOS 或达到最大长度。
-
最终从所有已完成和未完成的序列中选择得分最高的一条作为输出。
束搜索通过维护多条路径,能在一定程度上避免贪心的局部最优,找到全局更好的序列。
2.2 实现¶
def beam_search(model, input_ids, beam_width=3, max_length=50, eos_token_id=None):
"""
基础束搜索。
Args:
model: 自回归模型,返回 logits。
input_ids: (1, prompt_len) 初始提示词。
beam_width (k): 束宽。
max_length: 最大生成长度(包括prompt)。
eos_token_id: EOS token ID。
Returns:
best_sequence: (list) 得分最高的token序列。
"""
device = input_ids.device
prompt = input_ids[0].tolist()
# 束中的每个元素为 (tokens_list, cumulative_log_prob)
beams = [(prompt, 0.0)]
finished = [] # 已完成的序列
for step in range(max_length - len(prompt)):
if not beams: # 所有束都已完成
break
all_candidates = []
# 为每个束计算下一步 logits
for tokens, score in beams:
if eos_token_id is not None and tokens[-1] == eos_token_id:
finished.append((tokens, score))
continue
input_tensor = torch.tensor([tokens], device=device)
outputs = model(input_tensor)
next_logits = outputs[0, -1, :] # (vocab_size,)
log_probs = torch.log_softmax(next_logits, dim=-1)
# 为效率,可仅考虑 top-k 个候选,此处取全部
topk_log_probs, topk_indices = torch.topk(log_probs, beam_width, dim=-1)
for i in range(beam_width):
next_token = topk_indices[i].item()
new_score = score + topk_log_probs[i].item()
new_tokens = tokens + [next_token]
all_candidates.append((new_tokens, new_score))
if not all_candidates:
break
# 按得分降序排序,保留 beam_width 个最佳
ordered = sorted(all_candidates, key=lambda x: x[1], reverse=True)
beams = ordered[:beam_width]
# 合并已完成和剩余束
final_candidates = finished + beams
# 选取得分最高的
best_tokens, best_score = max(final_candidates, key=lambda x: x[1])
return best_tokens
2.3 重要说明¶
-
计算效率:上述实现每步对束中的每条序列独立调用模型,没有利用 batch。实际工程中应将束中所有序列 batch 化:将
beams中所有tokens堆叠为(beam_width, cur_len)的张量,一次前向得到所有序列的下一步 logits,然后扩展、排序、选择。这能大幅提升 GPU 利用率。 -
对数概率:使用
log_softmax将乘法转换为加法,避免浮点下溢。 -
EOS 处理:当某序列生成 EOS 时,将其移入
finished列表,并停止对其扩展,但不影响其他序列。 -
束宽选择:
beam_width通常取 3~10,过大计算量剧增,且可能引入更多噪音。
2.4 局限性¶
-
倾向于产生较短序列(因为对数概率为负,累加导致长序列得分低),需要结合长度惩罚。
-
多样性不足:多个束可能收敛到相似的高分序列,缺乏差异性。
实现带长度惩罚的束搜索,每条序列得分除以 len^α 后排序。¶
3.1 原理¶
长度惩罚(Length Penalty)用于缓解束搜索偏好短序列的问题。对于序列长度 L(通常只计生成部分),修正后的得分 score' = score / L^α,其中 α 是长度惩罚系数(常为 0.6~1.0)。排序时使用修正得分,使得长度适中的序列更可能胜出。
3.2 实现¶
在基础束搜索的基础上,修改排序逻辑和最终选择逻辑。
def beam_search_len_penalty(model, input_ids, beam_width=3, max_length=50,
eos_token_id=None, alpha=0.7):
device = input_ids.device
prompt_len = input_ids.shape[1]
prompt = input_ids[0].tolist()
beams = [(prompt, 0.0)] # (tokens, raw_score)
finished = []
for step in range(max_length - prompt_len):
if not beams:
break
all_candidates = []
for tokens, score in beams:
if eos_token_id is not None and tokens[-1] == eos_token_id:
finished.append((tokens, score))
continue
input_tensor = torch.tensor([tokens], device=device)
outputs = model(input_tensor)
next_logits = outputs[0, -1, :]
log_probs = torch.log_softmax(next_logits, dim=-1)
topk_log_probs, topk_indices = torch.topk(log_probs, beam_width)
for i in range(beam_width):
next_token = topk_indices[i].item()
new_score = score + topk_log_probs[i].item()
new_tokens = tokens + [next_token]
all_candidates.append((new_tokens, new_score))
if not all_candidates:
break
# 计算长度惩罚得分并排序
def len_pen_score(tokens, raw_score):
gen_len = len(tokens) - prompt_len
return raw_score / (gen_len ** alpha) if gen_len > 0 else raw_score
# 为每个候选添加惩罚得分
scored_candidates = [(tokens, raw_score, len_pen_score(tokens, raw_score))
for tokens, raw_score in all_candidates]
# 按惩罚得分降序
scored_candidates.sort(key=lambda x: x[2], reverse=True)
beams = [(tokens, raw_score) for tokens, raw_score, _ in scored_candidates[:beam_width]]
# 同样用长度惩罚选择最终最佳
final = finished + beams
best = max(final, key=lambda x: len_pen_score(x[0], x[1]))
return best[0]
3.3 参数 α 的选择¶
-
α = 0:等同于无惩罚。
-
α = 1.0:强惩罚,倾向最短序列。
-
常用值 0.6~0.8,在翻译、摘要等任务中微调。
3.4 注意事项¶
-
长度计算方式:应只计算生成部分的长度,不包括 prompt。否则 prompt 长度影响所有序列的惩罚,不公平。
-
惩罚应在最终选择时也应用,而不仅仅是每步排序。因为最终比较的是不同长度的已完成序列。
实现束搜索的提前终止逻辑,所有束生成 EOS 时停止,并正确处理已完成序列。¶
4.1 原理¶
提前终止(Early Stopping)是束搜索的标准优化。当束中所有活跃序列都已完成(即最后一个 token 都是 EOS),继续生成无意义,应立即终止。同时,已完成序列需要被安全存放,并在最终评选中参与竞争。
实现的关键是维护 beams 和 finished 两个列表:每步将遇到 EOS 的序列移入 finished,并从 beams 中移除;当 beams 为空时终止循环。
4.2 实现¶
在基础束搜索代码中,我们已实现了提前终止的核心逻辑。以下是更健壮的版本,明确处理活跃束与完成束的分离。
def beam_search_early_stop(model, input_ids, beam_width=3, max_length=50, eos_token_id=None):
device = input_ids.device
prompt = input_ids[0].tolist()
prompt_len = len(prompt)
# active_beams 存储 (tokens, score)
active_beams = [(prompt, 0.0)]
finished_beams = []
for step in range(max_length - prompt_len):
if not active_beams:
break
next_candidates = []
for tokens, score in active_beams:
# 注意:这里假设 tokens 的最后一个不会是EOS,因为我们在上一轮已经移除了
input_tensor = torch.tensor([tokens], device=device)
outputs = model(input_tensor)
next_logits = outputs[0, -1, :]
log_probs = torch.log_softmax(next_logits, dim=-1)
topk_log_probs, topk_indices = torch.topk(log_probs, beam_width)
for i in range(beam_width):
next_token = topk_indices[i].item()
new_score = score + topk_log_probs[i].item()
new_tokens = tokens + [next_token]
if eos_token_id is not None and next_token == eos_token_id:
finished_beams.append((new_tokens, new_score))
else:
next_candidates.append((new_tokens, new_score))
# 从扩展后的候选(不含EOS结束的)中选出 beam_width 个继续
next_candidates.sort(key=lambda x: x[1], reverse=True)
active_beams = next_candidates[:beam_width]
# 将所有序列合并,选取最佳(如果希望,可应用长度惩罚等)
all_sequences = finished_beams + active_beams
all_sequences.sort(key=lambda x: x[1], reverse=True)
best_tokens, _ = all_sequences[0]
return best_tokens
4.3 边界情况¶
-
如果 prompt 本身已经以 EOS 结尾?不应发生,但可预先检查。
-
如果某束刚刚生成 EOS,但在
next_candidates中已被剔除,则active_beams数量减少,这是正确的。 -
可能出现一种情况:所有活跃束都扩展后,
next_candidates数量少于beam_width,此时active_beams会缩减,甚至变为 0,循环退出。
4.4 优化建议¶
- 批量前向:将活跃束中的所有 tokens 整理为一批,一次前向,可大幅提速。
实现多样束搜索(Diverse Beam Search),在束间施加差异性惩罚。¶
5.1 原理¶
多样束搜索(Diverse Beam Search, DBS)解决普通束搜索多样性不足的问题。它将 k 个束分成 g 组,每组独立进行束搜索,同时引入组间差异惩罚:若某个组选择的 token 与之前组(按顺序)的选择相同,则对该 token 的得分施加惩罚。这使得各组被迫探索不同的生成方向。
典型做法:
-
将
k分成g组,每组有k/g个束。 -
在第
t步,依次处理每组:对于组i,计算所有候选的原始得分,然后对于候选中的每个 token,如果该 token 与之前组(0..i-1)在同样位置t生成的 token 重复,则减去惩罚值 δ(或乘以惩罚因子)。 -
惩罚强度 δ 控制多样性程度。
5.2 实现¶
def diverse_beam_search(model, input_ids, beam_width=5, groups=2, diversity_strength=0.5,
max_length=50, eos_token_id=None):
"""
diverse beam search
Args:
beam_width: 总束宽,需被groups整除。
groups: 分组数。
diversity_strength: 惩罚强度(得分减去该值乘以重复次数)。
"""
assert beam_width % groups == 0
group_size = beam_width // groups
prompt = input_ids[0].tolist()
prompt_len = len(prompt)
device = input_ids.device
# 为每组初始化束
group_beams = [[(prompt, 0.0)] for _ in range(groups)]
finished = [[] for _ in range(groups)]
for step in range(max_length - prompt_len):
# 存储本步各组选出的token,用于后续组计算惩罚
chosen_tokens_per_group = [] # list of list of token ids
new_group_beams = []
for g in range(groups):
candidates = []
# 收集本组当前束的所有扩展
for tokens, score in group_beams[g]:
if eos_token_id is not None and tokens[-1] == eos_token_id:
finished[g].append((tokens, score))
continue
input_t = torch.tensor([tokens], device=device)
outputs = model(input_t)
next_logits = outputs[0, -1, :]
log_probs = torch.log_softmax(next_logits, dim=-1)
# 为加速,取 topk*2 以便惩罚后仍有足够候选
topk_log_probs, topk_indices = torch.topk(log_probs, beam_width * 2)
for i in range(topk_indices.shape[0]):
next_token = topk_indices[i].item()
new_score = score + topk_log_probs[i].item()
# 多样性惩罚:检查是否与前面组选中的token重复
penalty = 0.0
for prev_g in range(g):
# 本步中前面组已决定的token(如果是EOS则不考虑)
if step < len(chosen_tokens_per_group[prev_g]):
if chosen_tokens_per_group[prev_g][step] == next_token:
penalty += diversity_strength
new_score -= penalty
new_tokens = tokens + [next_token]
candidates.append((new_tokens, new_score))
if not candidates:
new_group_beams.append([])
chosen_tokens_per_group.append([])
continue
# 选择本组得分最高的 group_size 个束,并记录本步选择的token(用于后续组)
candidates.sort(key=lambda x: x[1], reverse=True)
selected = candidates[:group_size]
new_group_beams.append(selected)
# 记录本组本步选出的所有token(注意可能由于有多个束,我们只对束内最高分的token?通常是对每个束记录其最后token)
# 更标准的DBS是记录本组所有束的最后token,后续组若生成相同token就惩罚。
group_step_tokens = [seq[-1] for seq, _ in selected]
chosen_tokens_per_group.append(group_step_tokens)
group_beams = new_group_beams
# 如果所有组的活跃束都为空,则结束
if all(len(b) == 0 for b in group_beams):
break
# 合并所有组的所有序列,选取最佳
all_final = []
for g in range(groups):
all_final.extend(finished[g])
for tokens, score in group_beams[g]:
if not (eos_token_id and tokens[-1] == eos_token_id): # 未完成的也加入
all_final.append((tokens, score))
best = max(all_final, key=lambda x: x[1])
return best[0]
5.3 说明¶
-
上述实现中,多样性惩罚是基于“前面组在同一时间步生成的 token”进行的。更复杂的实现可以考虑整个序列历史。
-
惩罚强度
diversity_strength需要调优,太大可能导致生成不连贯,太小则多样性不足。 -
在 NLP 任务中,多样束搜索常用于生成多个不同图像描述或对话回复。
实现 Top‑k 采样 top_k_sampling(logits, k, temperature)。¶
6.1 原理¶
Top‑k 采样在生成时截断概率分布,只保留概率最高的 k 个 token,将其余 token 的概率置零,然后重新归一化,从该截断分布中随机采样。这能有效避免选取到低概率的无关 token(噪声),同时保持一定的随机性和多样性。温度参数用于在截断前缩放 logits,控制分布的尖锐度。
6.2 实现¶
def top_k_sampling(logits, k, temperature=1.0):
"""
Args:
logits: (vocab_size,) 浮点张量。
k: 保留的候选 token 数量。
temperature: 温度参数,>1 更平滑,<1 更尖锐。
Returns:
sampled_token_id: int
"""
# 温度缩放
if temperature != 1.0:
logits = logits / temperature
# 获取 top-k 的 logits 及其索引
topk_logits, topk_indices = torch.topk(logits, k, dim=-1)
# 创建一个新的 logits 张量,填充 -inf
masked_logits = torch.full_like(logits, float('-inf'))
masked_logits.scatter_(0, topk_indices, topk_logits)
# softmax 得到概率
probs = torch.softmax(masked_logits, dim=-1)
# 采样
token_id = torch.multinomial(probs, num_samples=1).item()
return token_id
6.3 参数指南¶
-
k:通常取 10~100。值太小会导致输出重复、缺乏多样性;太大则引入噪音,可能生成不连贯文本。 -
temperature:T=1 无缩放;T<1 使高概率 token 更突出(接近贪心);T>1 使分布更平坦,增加低概率 token 被选中的机会。
6.4 与贪心的比较¶
-
贪心解码等价于 T=0 且 k=1 的极限情况。
-
Top‑k 通过保留 k 个候选并随机采样,打破了贪心解码的确定性,能产生更多样化的输出。
实现 Top‑p(Nucleus)采样 top_p_sampling(logits, p, temperature)。¶
7.1 原理¶
Top‑p 采样(核采样)不固定候选数量,而是动态选择最小集合,使得这些 token 的累积概率达到阈值 p。具体步骤:将 logits 按概率降序排序,累加概率,找到第一个使得累积概率超过 p 的位置,保留该位置及之前的所有 token,将其余 token 的概率置零,重新归一化后采样。这种方法可以适应不同置信度的分布,避免 Top‑k 有时截断过多或过少的弊端。
7.2 实现¶
def top_p_sampling(logits, p, temperature=1.0):
"""
Args:
logits: (vocab_size,) 张量。
p: 累积概率阈值,通常 0.9~0.95。
temperature: 温度缩放。
Returns:
token_id: int
"""
if temperature != 1.0:
logits = logits / temperature
# 按降序排序
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
# 计算累积概率
sorted_probs = torch.softmax(sorted_logits, dim=-1)
cumsum_probs = torch.cumsum(sorted_probs, dim=-1)
# 找到需要截断的位置:累积概率 > p 的第一个索引,但要确保至少保留一个token
# 方法:创建掩码,将超过p的部分设为True
# 我们希望保留 cumsum_probs <= p 的部分,以及刚好越过p的那个token(为了至少达到p)
# 常用技巧:将第一个超过p的位置保留,所以将其掩码置为False
mask = cumsum_probs > p
# 将第一个超过p的位置的mask设为False(即保留该token)
mask[1:] = mask[:-1].clone()
mask[0] = False
# 将mask对应的logits设为 -inf
sorted_logits[mask] = float('-inf')
# 重新映射回原始顺序
filtered_logits = torch.full_like(logits, float('-inf'))
filtered_logits.scatter_(0, sorted_indices, sorted_logits)
# 采样
probs = torch.softmax(filtered_logits, dim=-1)
token_id = torch.multinomial(probs, num_samples=1).item()
return token_id
7.3 与 Top‑k 的对比¶
-
Top‑k 固定候选数量,对分布变化不敏感:当模型非常确定时(少数 token 占绝大部分概率),Top‑k 仍保留 k 个,可能引入噪音;当模型不确定时(概率分散),Top‑k 可能遗漏重要的长尾 token。
-
Top‑p 动态调整候选集大小,更灵活,通常在多样性和质量间取得更好平衡。
实现温度缩放函数,并将其与贪心解码结合输出单个 token。¶
8.1 温度缩放原理¶
温度 T 对 logits 进行线性缩放:scaled_logits = logits / T。当 T > 1 时,概率分布趋于均匀,采样更随机;T < 1 时,分布更陡峭,接近 one-hot(极端时 T→0 退化为 argmax)。温度缩放在贪心解码中直接改变 logits,然后取 argmax 等同于先缩放再取最大值。因此可封装一个温度缩放函数,然后用 argmax 或采样。
8.2 实现¶
def temperature_scale(logits, temperature):
"""对 logits 进行温度缩放"""
if temperature == 0:
raise ValueError("Temperature must be > 0; for deterministic decoding use argmax directly.")
return logits / temperature
def temperature_greedy_decode(logits, temperature):
"""结合温度缩放和贪心解码,返回一个token"""
scaled_logits = temperature_scale(logits, temperature)
token_id = torch.argmax(scaled_logits).item()
return token_id
8.3 解析¶
-
当
temperature非常小(如 0.1),输出近乎确定性,但可能仍受浮点精度影响。 -
通常不建议在贪心中使用温度缩放,因为贪心本意是取最大值,温度缩放并不改变 argmax 的结果(单调变换,argmax 不变)。严格来说,
argmax(softmax(logits/T))等价于argmax(logits),因为 softmax 是单调递增函数,温度缩放不改变相对顺序。因此温度缩放结合贪心解码不会改变输出!这是一个常见的误解,面试中需明确指出。如果希望温度影响贪心,那是不可能的;要引入随机性必须使用采样。所以这里的结合实际是指“如果贪心解码,温度缩放无效”。我们可以讨论这一点,并说明正确的结合方式:温度缩放在采样策略中有效,而贪心不是采样。
修正:面试中我们应指出“温度缩放不会改变 argmax 的结果”,因此贪心+温度是无意义的。如果要求输出单个 token 且希望温度影响结果,应当使用温度缩放 + 随机采样,而不是贪心。所以题目可能是在考察对温度的理解。我们可以实现温度缩放的函数,然后说明其与采样结合,而不是贪心。
根据题意“实现温度缩放函数,并将其与贪心解码结合输出单个 token”,或许出题人希望我们结合的是:先温度缩放,然后贪心(取argmax)。但如上分析,argmax 不变。我们可以指出这一点,然后提供温度采样函数。
此处我实现温度缩放函数,并用采样来展示温度的影响,并讨论贪心为何不受影响。
def temperature_sampling(logits, temperature):
"""温度缩放后随机采样一个token"""
scaled = logits / temperature
probs = torch.softmax(scaled, dim=-1)
return torch.multinomial(probs, num_samples=1).item()
8.4 结论¶
-
温度缩放是采样策略的预处理步骤,不是贪心解码的组成部分。
-
贪心解码永远选择概率最高的 token,不受温度影响。
实现组合 Top‑k + Top‑p 采样,先 Top‑k 再 Top‑p 截断后采样。¶
9.1 原理¶
组合采样结合了 Top‑k 和 Top‑p 的优点:先用 Top‑k 丢弃长尾噪声,再在剩余高质量候选上应用 Top‑p 动态截断,以应对不同的分布形状。步骤如下:
-
温度缩放 logits。
-
选出 top-k 个最高概率的 logits,其余置
-inf。 -
在截断后的 logits 上应用 Top‑p 过滤。
-
softmax 后采样。
9.2 实现¶
def top_k_top_p_sampling(logits, k, p, temperature=1.0):
"""
先 Top-k 再 Top-p,最后采样。
"""
if temperature != 1.0:
logits = logits / temperature
# 第一步:Top‑k 过滤
topk_logits, topk_indices = torch.topk(logits, k, dim=-1)
# 创建 masked logits
k_masked = torch.full_like(logits, float('-inf'))
k_masked.scatter_(0, topk_indices, topk_logits)
# 第二步:在 Top‑k 结果上应用 Top‑p
sorted_logits, sorted_indices = torch.sort(k_masked, descending=True)
sorted_probs = torch.softmax(sorted_logits, dim=-1)
cumsum = sorted_probs.cumsum(dim=-1)
# 确定截断掩码
mask = cumsum > p
mask[1:] = mask[:-1].clone()
mask[0] = False
sorted_logits[mask] = float('-inf')
# 重建原始顺序的 logits
final_logits = torch.full_like(logits, float('-inf'))
final_logits.scatter_(0, sorted_indices, sorted_logits)
# 采样
probs = torch.softmax(final_logits, dim=-1)
token = torch.multinomial(probs, 1).item()
return token
9.3 参数推荐¶
-
k: 50~200,视词汇表大小而定。 -
p: 0.9~0.95。 -
temperature: 0.7~1.0。
该组合在实际应用(如 GPT 系列)中被广泛采用,能很好地平衡流畅性与多样性。
实现重复惩罚函数,支持频率惩罚和存在惩罚两种模式。¶
10.1 原理¶
重复惩罚用于减少模型生成重复文本。通过在生成每个 token 前,对已出现过的 token 施加惩罚,降低其被再次选中的概率。两种常见模式:
-
频率惩罚(Frequency Penalty):根据 token 在已生成序列中出现的次数按比例惩罚,出现越多,惩罚越重。公式:
logits[t] -= freq[t] * penalty。 -
存在惩罚(Presence Penalty):只要 token 出现过至少一次,就施加一个固定惩罚,与出现次数无关。公式:
logits[t] -= exists[t] * penalty。
通常 penalty 为正值(如 0.1~1.0),直接减少 logits。
10.2 实现¶
def apply_repetition_penalty(logits, generated_ids, penalty, mode='frequency'):
"""
Args:
logits: (vocab_size,) 当前步的 logits。
generated_ids: list of int,已生成的 token ID 序列(不包括当前)。
penalty: 惩罚系数,通常 0.1~1.0。
mode: 'frequency' 或 'presence'。
Returns:
penalized_logits: 同形状张量。
"""
if not generated_ids or penalty == 0.0:
return logits
# 统计频率或存在性
if mode == 'frequency':
# 计算每个token的出现次数
counts = torch.bincount(torch.tensor(generated_ids, device=logits.device),
minlength=logits.shape[0])
# 对出现过的token施加惩罚:logits -= counts * penalty
# 注意:我们希望出现次数多的惩罚更重,因此直接按次数乘以系数
penalized = logits - counts.float() * penalty
elif mode == 'presence':
# 只要出现过就惩罚固定值
unique_ids = set(generated_ids)
mask = torch.zeros(logits.shape[0], dtype=torch.bool, device=logits.device)
mask[list(unique_ids)] = True
penalized = logits - mask.float() * penalty
else:
raise ValueError("mode must be 'frequency' or 'presence'")
return penalized
10.3 使用方式¶
在每一步解码前,先对 logits 调用此函数施加惩罚,然后再进行 softmax 和采样或贪心。注意惩罚不应应用于 EOS token?通常 EOS 也需要适当惩罚以防止提前结束,但也可以选择对特殊 token 免除惩罚。
10.4 注意事项¶
-
惩罚强度过大可能导致输出不连贯或语法错误。
-
通常结合 temperature 和其他采样策略使用。
-
实现时需注意
generated_ids应包含 prompt 部分的 token 吗?通常只对生成部分的重复进行惩罚,但有时 prompt 中的高频词也可能导致生成时被惩罚,可根据需要决定。
实现硬约束解码,禁止生成与已有序列末尾形成重复 n-gram 的 token。¶
11.1 原理¶
硬约束解码(Hard N-gram Blocking)用于防止模型生成重复的连续片段。在每一步生成时,检查当前序列末尾是否存在一个长度为 n 的片段,它在已生成的文本中出现过(尤其是最近的输出),如果下一个 token 会导致即将形成的 n-gram 已经存在,则将该 token 的概率强制置零(logits 设为 -inf)。这能有效避免模型陷入重复循环,提升生成文本的多样性。
11.2 实现¶
def ngram_blocking_processor(logits, generated_ids, n=3):
"""
将导致形成重复 n-gram 的 token 的 logits 置为 -inf。
Args:
logits: (vocab_size,) 当前步的 logits。
generated_ids: list of int,已生成的 token ID 序列。
n: n-gram 长度,通常 3 或 4。
Returns:
processed_logits: 同形状张量。
"""
if len(generated_ids) < n - 1:
# 序列太短,无法形成 n-gram,不做约束
return logits
# 取当前序列末尾的 (n-1) 个 token 作为前缀
prefix = tuple(generated_ids[-(n-1):]) # 长度 n-1
# 扫描已生成序列中所有长度为 n 的 n-gram,记录 (prefix, next_token) 对
blocked_tokens = set()
for i in range(len(generated_ids) - n + 1):
if tuple(generated_ids[i:i+n-1]) == prefix:
# 找到了相同前缀的 n-gram,记录它的最后一个 token
next_tok = generated_ids[i+n-1]
blocked_tokens.add(next_tok)
# 如果末尾序列本身不足 n-1,但我们可以只检查最后的 (n-1) 长度历史
# 将 blocked_tokens 中的 token 的 logits 置为 -inf
processed_logits = logits.clone()
for tok in blocked_tokens:
processed_logits[tok] = float('-inf')
return processed_logits
11.3 使用方式¶
将此函数作为 LogitsProcessor 插入解码循环的 softmax 之前,通常与采样策略配合:
logits = model(...)
logits = ngram_blocking_processor(logits, generated_ids, n=3)
# 然后进行温度缩放、Top-k/Top-p、采样
11.4 注意事项¶
-
n值不宜太小(如2),否则会过度限制常见短语的正常生成;也不宜太大(如5+),则几乎不生效。 -
这种硬约束可能在某些任务中导致生成不连贯,因此常与软惩罚(如重复惩罚)结合,或仅在检测到严重重复时启用。
-
实现中
prefix是(n-1)个 token,这是为了检查“如果下一个 token 是 X,将形成一个新的 n-gram 是否已存在”。所以扫描的是整个序列中已有的所有长度为n的片段,并记录以相同前缀结尾的那些 token。 -
注意:我们只应检查生成部分(不含 prompt)的重复,但具体取决于任务,通常对整个序列(含 prompt)检查也是可行的,因为 prompt 中也可能有需要避免的重复模式。
实现 min‑p 采样,过滤概率低于 min_p * max_prob 的 token 后归一化采样。¶
12.1 原理¶
Min‑p 采样是一种动态的截断策略,由 Nguyen et al. 提出。它计算当前概率分布中的最大概率值 max_prob,然后将所有概率低于 min_p * max_prob 的 token 丢弃(概率置零),再重新归一化进行采样。与 Top‑p 不同,它不累积概率,而是基于最大概率的比例来截断。当分布非常尖锐时,只有少数 token 保留;当分布平坦时,保留较多 token。这使其适应性更强,尤其在创意性和连贯性之间取得较好平衡。
12.2 实现¶
def min_p_sampling(logits, min_p=0.1, temperature=1.0):
"""
Args:
logits: (vocab_size,) 张量。
min_p: 最小概率比例,通常 0.05~0.2。
temperature: 温度缩放。
Returns:
sampled_token_id: int
"""
if temperature != 1.0:
logits = logits / temperature
probs = torch.softmax(logits, dim=-1)
max_prob = probs.max()
# 计算阈值
threshold = min_p * max_prob
# 将低于阈值的 logits 置为 -inf
filtered_logits = logits.clone()
filtered_logits[probs < threshold] = float('-inf')
# 重新计算概率并采样
filtered_probs = torch.softmax(filtered_logits, dim=-1)
token_id = torch.multinomial(filtered_probs, num_samples=1).item()
return token_id
12.3 参数调优¶
-
min_p设为 0.1 时,若最大概率为 0.8,则保留概率 ≥ 0.08 的所有 token;若 max_prob 为 0.3,则保留 ≥ 0.03 的 token。可见其动态调整候选数量。 -
较小的
min_p(如 0.05)接近全量采样,较大的min_p(如 0.2)则更保守,类似 Top‑p 但机制不同。 -
可与温度缩放联用,先温度缩放再 min‑p 过滤。
12.4 与 Top‑p 的比较¶
- Top‑p 总是确保保留累积概率为 p 的最小集合,可能导致在平坦分布时保留过多 token(因为需要累积到 p),而 min‑p 在平坦分布时阈值也较低,因此保留更多,但在尖锐分布时,min‑p 可能仅保留 1~2 个 token,更接近贪心。两者可结合使用。
实现典型采样(Typical Sampling),基于局部熵选择 token 集合。¶
13.1 原理¶
典型采样(Typical Sampling, Meister et al., 2022)不是基于绝对概率,而是基于信息论中的“典型集”概念。它计算每个 token 的信息量 I = -log P,并与分布的局部熵比较。只保留信息量接近局部熵(即典型概率附近)的 token。具体做法:计算分布 P 的局部熵 H = -∑ P log P;对于每个 token,计算其信息量 I_i = -log P_i;保留那些满足 |I_i - H| < τ * H 的 token(τ 为阈值因子,如 0.2)。然后重新归一化采样。
此方法能筛选出具有“典型”概率值的 token,避免长尾噪声,同时在高概率 token 过多时限制候选集,生成更符合人类习惯的文本。
13.2 实现¶
def typical_sampling(logits, tau=0.2, temperature=1.0):
"""
Args:
logits: (vocab_size,) 张量。
tau: 阈值因子,通常 0.2~0.5。
temperature: 温度。
Returns:
sampled_token_id: int
"""
if temperature != 1.0:
logits = logits / temperature
probs = torch.softmax(logits, dim=-1)
# 计算局部熵 H
log_probs = torch.log(probs)
entropy = -torch.sum(probs * log_probs) # 标量
# 计算每个 token 的信息量
I = -log_probs # shape (vocab_size,)
# 保留接近熵的 token
mask = torch.abs(I - entropy) < tau * entropy
# 将被过滤的 logits 置为 -inf
filtered_logits = logits.clone()
filtered_logits[~mask] = float('-inf')
# 重新 softmax 并采样
filtered_probs = torch.softmax(filtered_logits, dim=-1)
token_id = torch.multinomial(filtered_probs, num_samples=1).item()
return token_id
13.3 参数说明¶
-
tau:较小的 tau 使得候选集更小、更保守;较大的 tau 放宽约束。实践中 0.2 是一个较好的起点。 -
局部熵
H反映了分布的不确定性。当分布平坦时 H 大,阈值放宽,保留更多 token;当分布尖锐时 H 小,阈值收紧,只保留少数高概率 token。
13.4 与其他采样策略的差异¶
- Top‑k/Top‑p 依赖于绝对概率排序;Min‑p 依赖于最大概率比例;Typical Sampling 依赖信息量与局部熵的偏离,理论基础更深厚,在某些场景下能产生更自然连贯的文本。
在束搜索的每步扩展中用采样替代固定 top‑k 选择,实现采样式束搜索。¶
14.1 原理¶
传统束搜索每步扩展时选择 top-k 个最高得分的 token,这可能导致束多样性不足。采样式束搜索(Stochastic Beam Search)在扩展时采用随机采样(如 top-k 采样、温度采样)从概率分布中抽取下一个 token,而不是确定性地选择 top-k。这样可以在保持束搜索框架的同时增加生成多样性,避免所有束收敛到相同序列。
注意:当引入采样后,束的得分不再是纯对数概率和,还需要考虑采样带来的随机性,通常仍使用累积对数概率作为得分。
14.2 实现¶
def stochastic_beam_search(model, input_ids, beam_width=3, max_length=50,
temperature=1.0, top_k=10, eos_token_id=None):
"""
每步从 top-k 中按概率采样一个 token 来扩展束。
"""
device = input_ids.device
prompt = input_ids[0].tolist()
prompt_len = len(prompt)
beams = [(prompt, 0.0)] # (tokens, score)
finished = []
for step in range(max_length - prompt_len):
if not beams:
break
next_candidates = []
for tokens, score in beams:
if eos_token_id is not None and tokens[-1] == eos_token_id:
finished.append((tokens, score))
continue
input_t = torch.tensor([tokens], device=device)
outputs = model(input_t)
next_logits = outputs[0, -1, :]
# 采样式选择下一个 token
next_token = top_k_sampling(next_logits, k=top_k, temperature=temperature)
new_tokens = tokens + [next_token]
# 计算对数概率(用于排序得分)
log_probs = torch.log_softmax(next_logits, dim=-1)
new_score = score + log_probs[next_token].item()
if eos_token_id is not None and next_token == eos_token_id:
finished.append((new_tokens, new_score))
else:
next_candidates.append((new_tokens, new_score))
# 按得分排序保留 beam_width 个最佳
next_candidates.sort(key=lambda x: x[1], reverse=True)
beams = next_candidates[:beam_width]
# 选取最佳序列
final = finished + beams
final.sort(key=lambda x: x[1], reverse=True)
return final[0][0]
14.3 讨论¶
-
因为每步是采样而非选 top-k,束中可能出现重复序列,可以通过维护已生成序列的哈希集合来去重。
-
采样参数(temperature, top_k)强烈影响多样性和质量,需要调优。
-
采样式束搜索常用于需要多样性的任务,如图像描述生成多个候选。
构建灵活的 LogitsProcessor 管道,依次施加温度、Top‑k、Top‑p、重复惩罚等。¶
15.1 设计思路¶
借鉴 transformers 库的 LogitsProcessorList,我们可以构建一个可组合的处理器管道。每个处理器是一个函数,接受 (logits, context) 返回修改后的 logits。管道按顺序调用。
15.2 实现¶
class LogitsProcessor:
def __call__(self, logits, generated_ids):
return logits
class TemperatureProcessor(LogitsProcessor):
def __init__(self, temperature):
self.temperature = temperature
def __call__(self, logits, generated_ids):
if self.temperature != 1.0:
logits = logits / self.temperature
return logits
class TopKProcessor(LogitsProcessor):
def __init__(self, k):
self.k = k
def __call__(self, logits, generated_ids):
topk_logits, topk_indices = torch.topk(logits, self.k, dim=-1)
mask = torch.full_like(logits, float('-inf'))
mask.scatter_(0, topk_indices, topk_logits)
return mask
class TopPProcessor(LogitsProcessor):
def __init__(self, p):
self.p = p
def __call__(self, logits, generated_ids):
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
sorted_probs = torch.softmax(sorted_logits, dim=-1)
cumsum = sorted_probs.cumsum(dim=-1)
mask = cumsum > self.p
mask[1:] = mask[:-1].clone()
mask[0] = False
sorted_logits[mask] = float('-inf')
result = torch.full_like(logits, float('-inf'))
result.scatter_(0, sorted_indices, sorted_logits)
return result
class RepetitionPenaltyProcessor(LogitsProcessor):
def __init__(self, penalty, mode='frequency'):
self.penalty = penalty
self.mode = mode
def __call__(self, logits, generated_ids):
if not generated_ids or self.penalty == 0.0:
return logits
if self.mode == 'frequency':
counts = torch.bincount(torch.tensor(generated_ids, device=logits.device),
minlength=logits.shape[0])
return logits - counts.float() * self.penalty
elif self.mode == 'presence':
unique_ids = set(generated_ids)
mask = torch.zeros(logits.shape[0], dtype=torch.bool, device=logits.device)
mask[list(unique_ids)] = True
return logits - mask.float() * self.penalty
return logits
class LogitsProcessorPipeline:
def __init__(self, processors):
self.processors = processors
def __call__(self, logits, generated_ids):
for proc in self.processors:
logits = proc(logits, generated_ids)
return logits
15.3 使用示例¶
pipeline = LogitsProcessorPipeline([
TemperatureProcessor(0.8),
RepetitionPenaltyProcessor(0.5, mode='frequency'),
TopKProcessor(50),
TopPProcessor(0.9),
])
# 在每步解码时
logits = pipeline(logits, generated_ids)
probs = torch.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, 1).item()
15.4 扩展性¶
-
可以方便地添加其他约束(如白名单、黑名单、长度惩罚、EOS 禁止等)。
-
注意处理顺序:温度应最先,因为缩放后影响后续概率;重复惩罚也应在截断前;Top‑k/Top‑p 通常在最后。
实现流式生成器,每次 yield 一个 token,支持实时调整采样参数。¶
16.1 设计¶
生成器是一个可迭代对象,每一步生成一个 token 并通过 yield 返回。为了支持动态调整参数,生成器可以接受一个外部队列或回调函数,在每步生成前检查是否有新参数传入。简化实现:在 yield 返回 token 的同时,外部可以通过 send 方法传入新的参数(如 temperature),利用 Python 生成器的双向通信。
16.2 实现¶
def stream_generate(model, input_ids, max_length=50, eos_token_id=None, **initial_params):
"""
流式生成器,每次 yield 一个 token。
可以通过 generator.send(params) 实时调整参数。
"""
# 初始参数
params = initial_params.copy()
generated = input_ids[0].tolist()
model.eval()
with torch.no_grad():
while len(generated) < max_length:
inputs = torch.tensor([generated], device=input_ids.device)
outputs = model(inputs)
logits = outputs[0, -1, :]
# 应用当前参数(如温度、top_k、top_p 等)
temperature = params.get('temperature', 1.0)
top_k = params.get('top_k', 0)
top_p = params.get('top_p', 1.0)
if temperature != 1.0:
logits = logits / temperature
probs = torch.softmax(logits, dim=-1)
if top_k > 0:
topk_probs, topk_indices = torch.topk(probs, top_k)
filtered_probs = torch.zeros_like(probs)
filtered_probs.scatter_(0, topk_indices, topk_probs)
probs = filtered_probs / filtered_probs.sum()
if top_p < 1.0:
sorted_probs, sorted_indices = torch.sort(probs, descending=True)
cumsum = sorted_probs.cumsum(dim=-1)
mask = cumsum > top_p
mask[1:] = mask[:-1].clone()
mask[0] = False
sorted_probs[mask] = 0.0
probs = torch.zeros_like(probs)
probs.scatter_(0, sorted_indices, sorted_probs)
probs = probs / probs.sum()
# 采样
next_token = torch.multinomial(probs, 1).item()
generated.append(next_token)
# yield 返回 token,并接收可能的新参数
new_params = yield next_token
if new_params is not None:
params.update(new_params)
if eos_token_id is not None and next_token == eos_token_id:
break
16.3 使用示例¶
gen = stream_generate(model, input_ids, max_length=50, temperature=0.9, top_p=0.95)
first_token = next(gen) # 启动生成器
# 后续可以动态调整
gen.send({'temperature': 1.2}) # 在下一个 token 生成时应用新温度
16.4 注意¶
-
生成器的双向通信需要先调用
next(gen)或gen.send(None)启动。 -
状态更新应确保参数安全,例如检查参数合法性。
-
对于需要更多复杂处理(如重复惩罚)的参数调整,可扩展参数字典。
实现带最小生成长度约束的束搜索,未达长度前禁止输出 EOS。¶
17.1 原理¶
在一些任务(如强制生成摘要)中,我们希望模型生成长度至少达到 min_length。实现方式:在束搜索的扩展阶段,如果当前序列长度(生成部分)小于 min_length,则禁止选择 EOS token——将其 logits 置为 -inf。但注意 EOS 的禁止可能导致序列无法终止,因此当所有束长度都超过 min_length 后,正常允许 EOS。
17.2 实现¶
在束搜索扩展循环中,对于每条序列,判断其生成部分的长度(len(tokens) - prompt_len),如果小于 min_length,则在获取 logits 后强制将 EOS token 的 logits 置为 -inf。
def beam_search_min_length(model, input_ids, beam_width=3, max_length=50,
min_length=10, eos_token_id=None):
device = input_ids.device
prompt = input_ids[0].tolist()
prompt_len = len(prompt)
beams = [(prompt, 0.0)]
finished = []
for step in range(max_length - prompt_len):
if not beams:
break
next_candidates = []
for tokens, score in beams:
# 若当前序列已生成部分 >= min_length 且遇到 EOS,则移入finished
gen_len = len(tokens) - prompt_len
if eos_token_id is not None and tokens[-1] == eos_token_id and gen_len >= min_length:
finished.append((tokens, score))
continue
# 否则继续扩展
input_t = torch.tensor([tokens], device=device)
outputs = model(input_t)
next_logits = outputs[0, -1, :]
# 强制禁止 EOS 如果 gen_len < min_length
if gen_len < min_length and eos_token_id is not None:
next_logits[eos_token_id] = float('-inf')
log_probs = torch.log_softmax(next_logits, dim=-1)
topk_log_probs, topk_indices = torch.topk(log_probs, beam_width)
for i in range(beam_width):
next_token = topk_indices[i].item()
new_score = score + topk_log_probs[i].item()
new_tokens = tokens + [next_token]
next_candidates.append((new_tokens, new_score))
if not next_candidates:
break
next_candidates.sort(key=lambda x: x[1], reverse=True)
beams = next_candidates[:beam_width]
final = finished + beams
final.sort(key=lambda x: x[1], reverse=True)
return final[0][0]
17.3 注意事项¶
-
若
min_length过大,可能导致束搜索无法在max_length内终止,最后可能选择未完成序列。 -
更精细的实现:对于已完成但长度不足的序列,不放入
finished,而是强制继续扩展(禁止 EOS)。当所有束都超过min_length后,允许正常结束。
在束搜索中加入去重机制,检测并排除重复序列路径。¶
18.1 原理¶
普通束搜索可能产生内容完全相同的多个束,浪费计算且降低多样性。去重机制通过维护一个哈希集合记录已经出现过的序列(或最近 n 步的关键特征),当某条新候选序列与已存在的序列(通常指已完成的束或同一束内其他分支)重复时,降低其得分或直接移除。
18.2 实现¶
可以在每步扩展后,对候选列表进行去重:如果两条候选序列的 token 列表完全相同,则只保留得分最高的那条(或直接丢弃重复)。也可以基于 n-gram 哈希进行去重。
def beam_search_with_dedup(model, input_ids, beam_width=3, max_length=50, eos_token_id=None):
device = input_ids.device
prompt = input_ids[0].tolist()
prompt_len = len(prompt)
beams = [(prompt, 0.0)]
finished = []
for step in range(max_length - prompt_len):
if not beams:
break
next_candidates = []
for tokens, score in beams:
if eos_token_id is not None and tokens[-1] == eos_token_id:
finished.append((tokens, score))
continue
input_t = torch.tensor([tokens], device=device)
outputs = model(input_t)
next_logits = outputs[0, -1, :]
log_probs = torch.log_softmax(next_logits, dim=-1)
topk_log_probs, topk_indices = torch.topk(log_probs, beam_width)
for i in range(beam_width):
next_token = topk_indices[i].item()
new_score = score + topk_log_probs[i].item()
new_tokens = tokens + [next_token]
next_candidates.append((new_tokens, new_score))
# 去重:相同 token 序列只保留得分最高的
dedup = {}
for tokens, score in next_candidates:
key = tuple(tokens) # 以完整序列作为键
if key not in dedup or score > dedup[key][1]:
dedup[key] = (tokens, score)
unique_candidates = list(dedup.values())
unique_candidates.sort(key=lambda x: x[1], reverse=True)
beams = unique_candidates[:beam_width]
final = finished + beams
final.sort(key=lambda x: x[1], reverse=True)
return final[0][0]
18.3 优化¶
-
完全基于完整序列去重在长序列下开销较大,可改用滑动窗口的 n-gram 哈希集合来近似去重。
-
去重可能减少束的有效数量,如果剩余束不足
beam_width,可允许保留较少束。 -
在多样束搜索中,去重尤为重要。
实现对比束搜索,结合参考模型惩罚项 score = log P_model - λ * log P_ref。¶
19.1 原理¶
对比解码(Contrastive Decoding)旨在通过对比一个较弱参考模型(如小模型或 amatuer model)来增强主模型的生成质量。束搜索中,可以修改得分函数:新得分 = log P_main(token) - λ * log P_ref(token)。当主模型和参考模型对某 token 的预测一致时,参考项会抵消一部分主模型的得分,抑制高频但无意义的词;当主模型明显偏好而参考模型不喜欢时,得分被增强,从而鼓励主模型生成更自信、更有信息量的内容。
19.2 实现¶
def contrastive_beam_search(main_model, ref_model, input_ids, beam_width=3, max_length=50,
lambda_val=0.5, eos_token_id=None):
"""
main_model: 主模型
ref_model: 参考模型(较弱)
"""
device = input_ids.device
prompt = input_ids[0].tolist()
prompt_len = len(prompt)
beams = [(prompt, 0.0)]
finished = []
for step in range(max_length - prompt_len):
if not beams:
break
next_candidates = []
for tokens, score in beams:
if eos_token_id is not None and tokens[-1] == eos_token_id:
finished.append((tokens, score))
continue
input_t = torch.tensor([tokens], device=device)
# 主模型 logits
main_logits = main_model(input_t)[0, -1, :]
# 参考模型 logits
ref_logits = ref_model(input_t)[0, -1, :]
# 计算对比得分
main_log_probs = torch.log_softmax(main_logits, dim=-1)
ref_log_probs = torch.log_softmax(ref_logits, dim=-1)
scores = main_log_probs - lambda_val * ref_log_probs
# 选择 topk 进行扩展
topk_scores, topk_indices = torch.topk(scores, beam_width, dim=-1)
for i in range(beam_width):
next_token = topk_indices[i].item()
new_score = score + topk_scores[i].item()
new_tokens = tokens + [next_token]
next_candidates.append((new_tokens, new_score))
next_candidates.sort(key=lambda x: x[1], reverse=True)
beams = next_candidates[:beam_width]
final = finished + beams
final.sort(key=lambda x: x[1], reverse=True)
return final[0][0]
19.3 注意事项¶
-
lambda_val控制参考模型的惩罚力度,值越大,越抑制常见词,鼓励主模型与众不同的输出。 -
参考模型通常应较弱,否则会削弱主模型的表现。
-
该方法可有效减少文本退化,提升信息量。
实现基于 token 白名单的约束解码,将不允许的 token 置为 -inf。¶
20.1 原理¶
在某些应用(如受控词汇生成、特定领域的输出格式化)中,我们只想从一组预先定义好的 token 集合(白名单)中选择下一个 token。实现很简单:在每一步解码前,将不在白名单中的 token 的 logits 强制设为 -inf,然后 softmax 时这些 token 的概率即为 0,保证采样或贪心时不会被选中。
20.2 实现¶
def whitelist_constraint(logits, allowed_token_ids):
"""
allowed_token_ids: list or tensor of token IDs that are allowed.
其他 token 置为 -inf。
"""
mask = torch.full_like(logits, float('-inf'))
allowed = torch.tensor(allowed_token_ids, device=logits.device, dtype=torch.long)
mask[allowed] = logits[allowed]
return mask
在解码循环中:
20.3 扩展¶
-
可以结合黑名单(将特定 token 排除)。
-
白名单通常用于:
- 强制输出 JSON 键或特定标点。
- 限制回答只在特定选项内。
- 受控词汇(如医疗术语)生成。
实现批量束搜索,支持 batch 内多样本的并行束搜索。¶
21.1 原理¶
在实际应用中,我们通常需要同时为 batch 中的多个样本进行束搜索。可以通过将 batch 内所有样本的束展开为一个大的 batch 维度,利用 GPU 并行计算。实现时,维护一个形状为 (batch_size * beam_width, cur_len) 的输入张量,每个样本的束占据连续的 beam_width 行。扩展时,从该大 batch 计算 logits,然后为每个样本独立执行 top‑k 选择、得分更新、束重组。注意需要使用索引偏移来区分不同样本的束。
21.2 实现(简化版)¶
假设模型接受 (batch_size, seq_len) 并返回 (batch_size, seq_len, vocab_size) 的 logits。
def batched_beam_search(model, input_ids, beam_width=3, max_length=50, eos_token_id=None):
"""
input_ids: (batch_size, prompt_len) 初始 prompt 批量。
返回: (batch_size, seq_len) 每个样本的最佳序列。
"""
batch_size, prompt_len = input_ids.shape
device = input_ids.device
# 初始束:将每个样本的 prompt 重复 beam_width 次
# 形状: (batch_size * beam_width, prompt_len)
expanded_input = input_ids.unsqueeze(1).expand(-1, beam_width, -1).reshape(batch_size * beam_width, prompt_len)
# 每个束的得分,初始 0
beam_scores = torch.zeros(batch_size * beam_width, device=device)
# 标记束是否已完成
done = torch.zeros(batch_size * beam_width, dtype=torch.bool, device=device)
# 存储最佳序列(每个样本独立)
# 此处简化:仅返回最后得分最高的束对应的序列。实际还需考虑提前终止。
for step in range(max_length - prompt_len):
if done.all():
break
# 前向
outputs = model(expanded_input)
next_logits = outputs[:, -1, :] # (B*BW, vocab)
log_probs = torch.log_softmax(next_logits, dim=-1)
# 计算所有可能扩展的得分:每个束扩展 beam_width 个候选,共 (B*BW * beam_width) 个
# 为简便,我们取每个束的 top-beams 个 token
topk_log_probs, topk_indices = torch.topk(log_probs, beam_width, dim=-1) # (B*BW, beam_width)
# 计算新得分
new_scores = beam_scores.unsqueeze(1) + topk_log_probs # (B*BW, BW)
# 重塑为 (B, BW*BW) 以便为每个样本独立选取 top BW
new_scores = new_scores.view(batch_size, beam_width * beam_width)
topk_scores, topk_pos = torch.topk(new_scores, beam_width, dim=-1) # (B, BW)
# topk_pos 指示在 (BW*BW) 中的索引,需要还原为原束索引和 token 索引
beam_indices = topk_pos // beam_width # (B, BW) 原束索引
token_indices = topk_pos % beam_width # (B, BW) token 索引(在 topk_indices 中的位置)
# 获取实际 token ID
next_tokens = torch.gather(topk_indices, 1, token_indices + beam_indices * beam_width) # 需正确索引,简化可改用循环。
# 为简洁,此处略去详细索引实现,实际可参考 HuggingFace 源码。
# 更新 expanded_input 和 beam_scores
# ...
# 最终从每个样本的束中选择最佳序列(略)
return best_sequences
由于完整实现较复杂,面试中可描述原理并给出核心索引逻辑,强调 batch 内并行带来的效率提升。
21.3 关键点¶
-
需要正确维护
beam_scores和done标志,以及处理 EOS 提前终止。 -
通过将 beam_width 维度融入 batch 维度,充分利用 GPU 并行。
-
最终选择每个样本的最高分序列作为输出。