AI GPT架构折腾手记
最近在做一个对话系统,用了开源 GPT 模型,生成质量时好时坏,温度、top-p 只会调参不懂原理,推理速度也摸不准瓶颈——根子是对 Decoder 生成链路不熟,于是把整条链路从头捋了一遍。
为什么写这篇
上面那几个问题,本质上都是没搞清 GPT 生成时在算什么、怎么采样。下面是在那次梳理后整理的一份"用代码说话"的记录。
先搞清楚:GPT 不是"预测下一个字"那么简单
很多资料都说:GPT 就是"预测下一个 token"。这话没错,但不完整。
完整的过程是:在给定上下文的情况下,模型计算所有可能的下一个 token 的概率分布,然后从这堆概率里"采样"出一个 token,把这个 token 当作输入,再预测下一个,如此循环。
关键有两个:
- 概率分布怎么算:这是 Decoder 架构的核心
- 采样策略怎么选:这决定了生成结果的质量和多样性
用一张流程图可能更清楚:
文本生成是一步步迭代的:每轮基于已有上下文算下一 token 的概率,再采样出一个写回去,循环直到结束。
Decoder 的核心机制:自回归
GPT 用的就是典型的自回归(Autoregressive)架构。这个概念说得很玄,其实很简单:
- 自回归:每次预测只依赖"已经生成的东西",不看未来
- 自编码:每次预测可以看到完整上下文
GPT 属于前者。这也解释了为什么 GPT 生成是串行的:你不知道下一个 token 是什么,就没法预测下下个。
用一个最小化例子说明:
假设我们想让模型学 “我 爱 编 程” 这个句子:
# 训练数据准备
sentences = ["我 爱 编 程"]
# 模型需要学习的是:
# P("爱" | "我") = ?
# P("编" | "我 爱") = ?
# P("程" | "我 爱 编") = ?
# P("<EOS>" | "我 爱 编 程") = ?
# 在推理时:
# 输入 "我" → 输出 "爱"
# 输入 "我 爱" → 输出 "编"
# 输入 "我 爱 编" → 输出 "程"
# 输入 "我 爱 编 程" → 输出 "<EOS>"
这就是自回归的本质:每一步都基于之前的所有历史。
掩码注意力:确保不偷看未来
在训练时,模型能看到完整句子,但推理时看不到未来。怎么让训练和推理对齐?
答案是:掩码注意力(Masked Attention)。
import numpy as np
# 假设我们有 4 个 token 的序列
seq_len = 4
# 创建掩码矩阵:下三角矩阵为 0(可见),上三角为 -inf(不可见)
mask = np.triu(np.full((seq_len, seq_len), -np.inf), k=1)
print("Mask 矩阵:")
print(mask)
这个掩码矩阵会在注意力分数计算时发挥作用:
# 注意力分数
scores = attention_scores + mask # 不可见的位置加上 -inf
# 经过 Softmax 后,这些位置的权重会变成 0
attention_weights = softmax(scores)
这样,当模型预测第 3 个 token 时,它只能用前 2 个 token 的信息,第 4 个和后面的都被"遮住"了。
这是 GPT 能"自言自语"的基础:每一步都严格依赖历史,不偷看未来。
位置编码:让模型知道"顺序"
Attention 机制本身没有"顺序"概念,如果输入换顺序,输出是一样的。所以需要显式告诉模型"哪个 token 在哪个位置"。
GPT 用的是学习到的位置编码:
import torch
import torch.nn as nn
# 位置编码层
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=512):
super().__init__()
# 创建一个可学习的位置编码矩阵
self.pe = nn.Parameter(torch.randn(max_len, d_model))
def forward(self, x):
# x: (batch_size, seq_len, d_model)
seq_len = x.size(1)
# 取对应长度的位置编码
return x + self.pe[:seq_len]
这样,每个位置都有一个独特的编码向量,模型通过学习来"记住"位置信息。
采样策略:从概率分布到具体输出
模型输出的是所有可能 token 的概率分布,但最终只能选一个。怎么选?
这就是采样策略的问题。最简单的是"贪心采样":直接选概率最高的。
def greedy_sample(logits):
"""贪心采样:直接选概率最高的 token"""
probs = torch.softmax(logits, dim=-1)
next_token = torch.argmax(probs, dim=-1)
return next_token
但贪心采样的问题是:容易陷入重复、生成单调。于是有了更高级的策略:
温度采样
温度(Temperature)是一个很常见的参数,本质上是"控制输出的随机性":
def temperature_sample(logits, temperature=1.0):
"""温度采样"""
# 1. 缩放 logits
scaled_logits = logits / temperature
# 2. 转成概率分布
probs = torch.softmax(scaled_logits, dim=-1)
# 3. 按概率采样
next_token = torch.multinomial(probs, num_samples=1)
return next_token
温度越小,输出越"确定"(接近贪心);温度越大,输出越"随机"。
# 温度对输出的影响
high_temp_logits = torch.tensor([1.0, 2.0, 3.0])
low_temp_logits = torch.tensor([1.0, 2.0, 3.0])
print("温度=1.0:", temperature_sample(high_temp_logits, temperature=1.0))
print("温度=0.1:", temperature_sample(low_temp_logits, temperature=0.1))
print("温度=2.0:", temperature_sample(high_temp_logits, temperature=2.0))
Top-k 和 Top-p
这两个参数是为了避免采样到"非常低概率"的 token:
- Top-k:只从概率最高的 k 个 token 中采样
- Top-p:从累计概率达到 p 的 token 中采样(也叫 Nucleus Sampling)
def top_k_sample(logits, k=50):
"""Top-k 采样"""
# 1. 找到概率最高的 k 个 token
top_k_probs, top_k_indices = torch.topk(logits, k)
# 2. 其他位置设为 -inf
logits_masked = torch.full_like(logits, -np.inf)
logits_masked.scatter_(1, top_k_indices, top_k_probs)
# 3. 按概率采样
probs = torch.softmax(logits_masked, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
return next_token
def top_p_sample(logits, p=0.9):
"""Top-p 采样"""
# 1. 按概率排序
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
# 2. 计算累计概率
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
# 3. 移除累计概率超过 p 的 token
sorted_indices_to_remove = cumulative_probs > p
# 但至少保留一个
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
# 4. 遮罩
sorted_logits[sorted_indices_to_remove] = -np.inf
# 5. 恢复原始顺序
logits_masked = torch.gather(sorted_logits, 1, sorted_indices.argsort(-1))
# 6. 按概率采样
probs = torch.softmax(logits_masked, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
return next_token
这两个参数在实际用的时候很常见,合理设置能显著提升输出质量。
踩过的坑:这些参数真的会卡住你
在实践中,这几个参数调节不当会导致明显问题:
温度设置过低
现象:输出非常"机械",容易陷入重复循环。
# 错误示例
output = model.generate(
input_ids,
temperature=0.1, # 太低了
max_length=100
)
# 可能输出:
# "好的。好的。好的。好的。好的。好的。好的。好的。好的。好的。好的。好的。"
温度设置过高
现象:输出开始"胡说八道",不再连贯。
# 错误示例
output = model.generate(
input_ids,
temperature=2.0, # 太高了
max_length=100
)
# 可能输出:
# "太阳上的鱼正在用键盘给月亮写代码,但是彩虹告诉我昨天明天..."
没有设置 Top-p
现象:偶尔会采样到非常奇怪的 token。
# 错误示例
output = model.generate(
input_ids,
temperature=0.8,
# 没有设置 top_p
max_length=100
)
# 可能输出:
# "我今天吃了[UNK],感觉很[UNK]。"
合理的参数组合
根据经验,一个比较安全的起点是:
# 推荐参数组合
output = model.generate(
input_ids,
temperature=0.7, # 适中的随机性
top_p=0.9, # 避免极端低概率 token
top_k=50, # 限制候选集
max_length=512, # 合理的输出长度
repetition_penalty=1.2 # 避免重复
)
当然,具体场景需要具体调,但这是一个不错的起点。
实战:从零实现一个简化的生成循环
为了真正理解这个过程,我写了一个最小化的生成函数:
import torch
import torch.nn.functional as F
def simple_generate(model, tokenizer, prompt, max_length=100,
temperature=0.7, top_p=0.9):
"""一个简化的文本生成函数"""
# 1. 编码输入
input_ids = tokenizer.encode(prompt, return_tensors='pt')
# 2. 自回归生成
for _ in range(max_length):
# 2.1 前向传播,得到下一个 token 的 logits
with torch.no_grad():
outputs = model(input_ids)
next_token_logits = outputs.logits[:, -1, :] # 取最后一个位置
# 2.2 温度缩放
next_token_logits = next_token_logits / temperature
# 2.3 Top-p 过滤
filtered_logits = top_p_filter(next_token_logits, top_p)
# 2.4 采样
probs = F.softmax(filtered_logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
# 2.5 判断是否结束
if next_token.item() == tokenizer.eos_token_id:
break
# 2.6 拼接到输入
input_ids = torch.cat([input_ids, next_token], dim=-1)
# 3. 解码输出
generated_text = tokenizer.decode(input_ids[0], skip_special_tokens=True)
return generated_text
def top_p_filter(logits, top_p):
"""Top-p 过滤"""
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(torch.softmax(sorted_logits, dim=-1), dim=-1)
# 移除累计概率超过 top_p 的 token
sorted_indices_to_remove = cumulative_probs > top_p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
sorted_logits[sorted_indices_to_remove] = float('-inf')
return torch.gather(sorted_logits, 1, sorted_indices.argsort(-1))
这个函数虽然简单,但包含了 GPT 生成的核心逻辑。
优化推理速度:KV Cache 的作用
在实际用的时候,发现生成速度是个大问题。原因很简单:
每生成一个 token,都要重新计算之前所有 token 的注意力,效率很低。
KV Cache(键值缓存)就是为了解决这个问题:
def generate_with_kv_cache(model, tokenizer, prompt, max_length=100):
"""使用 KV Cache 的生成函数"""
input_ids = tokenizer.encode(prompt, return_tensors='pt')
# 初始化 KV cache
past_key_values = None
for _ in range(max_length):
# 如果有 past_key_values,只传入最后一个 token
if past_key_values is not None:
input_ids = input_ids[:, -1:]
with torch.no_grad():
outputs = model(
input_ids,
past_key_values=past_key_values,
use_cache=True
)
logits = outputs.logits[:, -1, :]
next_token = torch.argmax(logits, dim=-1)
# 更新 KV cache
past_key_values = outputs.past_key_values
if next_token.item() == tokenizer.eos_token_id:
break
# 只保留最新的 token
input_ids = torch.cat([input_ids[:, :-1], next_token], dim=-1)
generated_text = tokenizer.decode(input_ids[0], skip_special_tokens=True)
return generated_text
KV Cache 的本质是:把之前计算过的 Key 和 Value 缓存起来,新生成 token 时,只需要计算这个 token 和之前所有 token 的注意力,而不用重新计算之前 token 之间的注意力。
这个优化在实际项目中非常关键,能显著提升推理速度。
结果:理解带来更好的控制
梳理完这些,再回头看最初的几个问题:
生成质量时好时坏:本质上是采样策略和参数设置问题。温度、top-p、top_k 这些参数,直接影响输出的多样性和连贯性。
温度、top-p 的本质:控制"从概率分布到具体输出"的策略,不是玄学旋钮。搞懂采样,调参才有依据。
优化推理速度:少做重复计算是主线。KV Cache 是基础手段,还有 beam search、early stopping 等可选策略。
更重要的是,理解了生成过程,就能更有针对性地优化。比如:
- 想要更有创意的输出?提高温度
- 想要更稳定的输出?降低温度,提高 top-p
- 想要避免重复?设置 repetition_penalty
- 想要更快的推理?用 KV Cache,限制 beam search 的 beam size
结语
GPT 从 token 化、掩码注意力、位置编码到采样和 KV Cache,是一条完整链路,不是单个"预测下一个字"能概括的。
这篇只覆盖 Decoder 到生成实践的主线,离完整架构还远,但日常调参和排障够用了。
搞懂设计取舍,才知道该动哪、不该动哪——这次梳理对我帮助最大的是这个。
版权声明: 本文首发于 指尖魔法屋-AI GPT架构折腾手记(https://blog.thinkmoon.cn/post/343-ai-gpt-architecture-decoder-generation-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。