本文深入剖析LLM文本生成的核心解码算法,对比Greedy、Beam Search、Top-k、Top-p及典型惩罚机制的数学原理与工程实现,并结合KV缓存、动态批处理等优化手段,为高吞吐、低延迟的AI写作服务提供实战指南。
无论是ChatGPT的对话续写,还是Notion AI的文案生成,其底层都依赖自回归语言模型(Autoregressive LM)。给定前缀提示(Prompt),模型逐token预测下一个词的概率分布,并通过解码策略决定最终输出。
解码策略直接决定生成文本的多样性、连贯性、创造性。工业级AI写作系统往往需要在质量与速度之间做精细权衡。本文将系统性地梳理主流解码方法,并给出高性能服务端的工程化方案。
设模型为 PθPθ,给定输入 token 序列 x1:tx1:t,生成下一个 token 的概率为:
P(xt+1∣x1:t)=softmax(W⋅ht+b)P(xt+1∣x1:t)=softmax(W⋅ht+b)
其中 htht 为Transformer最后层的隐藏状态。完整生成序列 Y={y1,...,yT}Y={y1,...,yT} 的联合概率为:
P(Y∣X)=∏i=1TP(yi∣X,y<i)P(Y∣X)=i=1∏TP(yi∣X,y<i)
解码的目标是找到使联合概率最大的序列(或近似最优序列),同时控制生成质量。
每一步直接选取概率最大的 token:
yt=argmaxw∈VP(w∣y<t)yt=argw∈VmaxP(w∣y<t)
def greedy_decode(logits):
return torch.argmax(logits, dim=-1)维护 KK 个候选序列,每一步扩展所有候选,保留总对数概率最高的 KK 个序列。
score(y1:t)=∑i=1tlogP(yi∣y<i)score(y1:t)=i=1∑tlogP(yi∣y<i)
最终选择得分最高的序列。
# 伪代码:Beam Search核心步骤
def beam_search(model, input_ids, beam_size=4, max_len=50):
sequences = [[input_ids, 0.0]]
for _ in range(max_len):
candidates = []
for seq, score in sequences:
logits = model(seq)
topk = torch.topk(logits[-1], beam_size)
for token, log_prob in zip(topk.indices, topk.values):
candidates.append([seq + token, score + log_prob])
sequences = sorted(candidates, key=lambda x: x[1], reverse=True)[:beam_size]
return sequences[0][0]为了提升创造性和多样性,工业级AI写作服务几乎都采用采样方法。
对 logits 除以温度 TT 后再 softmax,控制概率分布的“锐利”程度:
P(w)=exp(zw/T)∑jexp(zj/T)P(w)=∑jexp(zj/T)exp(zw/T)
工程注意:温度不能过高,否则会采样到大量无意义 token。
def temperature_sampling(logits, temperature=1.0):
scaled = logits / temperature
probs = torch.softmax(scaled, dim=-1)
return torch.multinomial(probs, num_samples=1)仅保留概率最高的 kk 个 token,重新归一化后采样。避免采样到低概率的“尾部分布”。
动态选择概率累积达到阈值 pp 的最小 token 集合,再采样。这是目前 ChatGPT 等产品使用的默认方案。
Vp=argminV′{∑w∈V′P(w)≥p}Vp=argV′min{w∈V′∑P(w)≥p}
def top_p_sampling(logits, p=0.9, temperature=1.0):
probs = torch.softmax(logits / temperature, dim=-1)
sorted_probs, sorted_indices = torch.sort(probs, descending=True)
cumsum = torch.cumsum(sorted_probs, dim=-1)
mask = cumsum < p
# 确保至少包含一个token
mask[..., 0] = True
sorted_probs[~mask] = 0.0
sorted_probs = sorted_probs / sorted_probs.sum()
idx = torch.multinomial(sorted_probs, 1)
return sorted_indices[..., idx]AI写作常见问题——重复短语、重复句子。可通过频次惩罚(Frequency Penalty)和存在惩罚(Presence Penalty)缓解:
OpenAI API 中即支持这两个参数。工程实现需维护当前序列的 token 计数。
在生产环境,AI写作服务需要支撑高并发,解码延迟是关键瓶颈。以下是通用优化手段:
Transformer 自注意力中,每个 token 的 Key/Value 可被缓存,避免重复计算。实现时需使用增量式推理:
# 伪代码:带KV缓存的生成
def generate_with_cache(model, input_ids, past_kv=None):
for step in range(max_tokens):
outputs = model(input_ids[:, -1:], past_key_values=past_kv)
past_kv = outputs.past_key_values
next_token = sample(outputs.logits)
input_ids = torch.cat([input_ids, next_token], dim=1)
return input_ids现代 LLM 库(HuggingFace Transformers)已内置此功能,只需设置 use_cache=True。
传统静态批处理(Static Batching)需等最慢请求完成才能返回。连续批处理允许新请求动态加入,已完成的请求先退出,大幅提升 GPU 利用率。框架如 vLLM、TensorRT-LLM 均支持。
使用一个较小的“草稿模型”快速生成多个候选 token,再由大模型并行验证,可在不降低质量的前提下提升 2~3 倍吞吐。适用于对延迟要求极苛刻的场景。
基于多个生产环境经验,推荐默认配置:
参数 | 推荐值 | 说明 |
|---|---|---|
温度 | 0.7 ~ 0.85 | 创意类写作偏高,事实类偏低 |
Top-p | 0.9 | 广泛适用 |
Top-k | 40 | 可选,作为前过滤 |
频次惩罚 | 0.3 ~ 0.5 | 防止重复 |
存在惩罚 | 0.0 ~ 0.2 | 鼓励新话题 |
最大长度 | 动态 | 根据任务调整 |
此外,针对不同风格的写作(诗歌、新闻、代码),可以维护不同的“策略模板”,通过 A/B 测试优化。
下面提供一个融合 Top-k、Top-p、温度、惩罚的通用采样函数,可直接集成到推理服务中:
def sample_with_all(logits, tokens_seen, temperature=0.8, top_k=40, top_p=0.9,
frequency_penalty=0.3, presence_penalty=0.1):
# 应用惩罚
for token_id, count in tokens_seen.items():
logits[token_id] -= frequency_penalty * count
logits[token_id] -= presence_penalty * (count > 0)
# 温度缩放
logits = logits / temperature
# Top-k 过滤
if top_k > 0:
indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]
logits[indices_to_remove] = -float('Inf')
# Top-p 过滤
if top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cum_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
sorted_indices_to_remove = cum_probs > top_p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = False
indices_to_remove = sorted_indices[sorted_indices_to_remove]
logits[indices_to_remove] = -float('Inf')
probs = torch.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
return next_token策略 | 生成质量 | 吞吐量 (tokens/s) | 多样性 | 适用场景 |
|---|---|---|---|---|
Greedy | 中 | 最高 | 最低 | 代码补全、翻译 |
Beam Search (K=4) | 高 | 中 | 低 | 摘要、机器翻译 |
Top-p (p=0.9) | 高 | 中高 | 高 | 对话、创意写作 |
Top-p + 惩罚 | 很高 | 中 | 极高 | 长文本生成 |
在线服务通常采用Top-p + 温度 + 惩罚组合,并配合 KV 缓存与连续批处理。
AI写作的解码策略并非一成不变,需要根据任务类型、用户反馈、硬件资源动态调整。未来趋势包括:
原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。
如有侵权,请联系 cloudcommunity@tencent.com 删除。