GPT 自回归文本生成实战:让训练好的 GPT 模型开口说话(Make GPT Talk Back)
发布时间:2026/9/18 1:57:54 锦皓数字建站
`)
GPT 自回归文本生成实战让训练好的 GPT 模型开口说话Make GPT Talk Back【免费下载链接】leetcodeLeetcode solutions项目地址: https://gitcode.com/GitHub_Trending/leetcode1/leetcode本文聚焦 GPT 系列课程中的收官环节——自回归文本生成如何把训练好的 GPT 模型接入一个生成循环让它逐 token 采样、逐字符拼出与训练数据风格相似的新文本。读完本文你将掌握生成循环的完整流程、torch.multinomial采样与贪心解码的本质区别、上下文窗口裁剪的边界处理以及温度缩放temperature与 top-k/top-p 等生产级解码策略的原理。前置知识动手生成前需要掌握的三个概念原文档指出在尝试本问题之前你需要对以下内容足够熟悉GPT 架构GPT Architecture一个完整的模型——输入一串 token ID 上下文在每个位置输出一个概率分布。它的内部由 token 嵌入、位置嵌入、N 个 Transformer 块、最终 LayerNorm 与词表投影组成对应仓库中的 GPT 架构专题文章自回归生成Autoregressive Generation一次只生成一个 token把上一次的预测结果重新作为输入喂回模型。生成本质上是一个模型前向传播 token 采样交替进行的循环从分布中采样Sampling from Distributions使用torch.multinomial从概率分布中抽取一个 token从而为输出引入受控的随机性。这三点分别对应模型是什么生成怎么循环如何挑选下一个词是理解整篇解决方案的三块基石。核心概念GPT 的文本生成就是自回归循环GPT 的文本生成在推理阶段进行是一个自回归循环模型一次生成一个 token把它追加到上下文末尾然后重复这一过程。这正是把训练好的语言模型变成文本生成器的推理期流程。生成循环的七个步骤裁剪Crop如果上下文已经长得超过模型的最大上下文长度就把最前面的部分裁掉只保留最近的部分前向传播Forward pass把上下文喂给模型得到每个位置的概率分布原始文档中模型输出 logits需要自行 softmax 得到概率提取末位Extract last position只有最后一个位置的分布才有意义因为它预测的是下一个 token采样Sample用torch.multinomial从分布中抽取一个 token追加Append把采样得到的 token 追加到上下文末尾解码Decode把 token ID 转换成字符重复Repeat按需要的字符数重复以上过程。其中第 3 步的依据来自模型本身的因果设计GPT 注意力层内部使用因果掩码causal masking保证位置 $t$ 只能看到位置 $0$ 到 $t$ 的信息参见 GPT 架构专题文章因此当前位置的输出只能预测下一个 token这是整个自回归范式成立的前提。采样 vs. 贪心解码策略做法特点贪心解码greedy / argmax每次都取概率最大的 token确定性输出但文本常常重复、单调采样sampling从完整分布中按概率随机抽取引入多样性文本更自然在生产环境中通常还会叠加两类控制手段来调节创造性 vs. 连贯性的平衡温度缩放temperature scaling$\text{probs} \text{softmax}(\text{logits} / T)$。当 $T \to 0$ 时分布趋近于 one-hot近似贪心当 $T$ 较大时分布趋于均匀更随机、更具创造性top-k / top-p 过滤只允许在概率最高的前 $k$ 个 token 中采样top-k或只保留累积概率达到 $p$ 的最小 token 集合top-p也称 nucleus sampling。两者都是把采样空间限制在合理候选内避免尾部噪声 token 拉低质量。上下文窗口是有限的上下文窗口context window是有限的。一旦上下文长度超过最大长度 $C$更早的 token 就会被裁剪掉。模型的记忆仅限于最近 $C$ 个 token——这正是 GPT 模型存在上下文长度上限的原因GPT-2 为 1024现代模型多为 8192 及以上。这也解释了为什么生成时必须显式裁剪位置嵌入表只覆盖 $0 \sim C-1$ 的位置索引超长输入会直接触发索引越界。补充一个底层细节GPT 使用可学习的位置嵌入learned position embeddings对应 GPT 架构专题文章 中position_embeddings nn.Embedding(context_length, model_dim)的实现——位置表长度就是context_length这也从源码层面印证了超出上下文长度会导致位置嵌入索引错误这一裁剪必要性的根源。解决方案完整的生成函数实现思路Intuition循环new_chars次每次迭代执行需要时裁剪上下文 → 前向传播 → 提取末位概率 → 采样一个 token → 追加到上下文 → 解码。为了可复现性需要对随机数生成器RNG状态进行管理。完整实现import torch import torch.nn as nn from torchtyping import TensorType class Solution: def generate(self, model, new_chars: int, context: TensorType[int], context_length: int, int_to_char: dict) - str: generator torch.manual_seed(0) initial_state generator.get_state() result [] for _ in range(new_chars): # Crop context to max length the model can handle if context.shape[1] context_length: context context[:, -context_length:] # Forward pass - logits for every position logits model(context) # (1, T, vocab_size) last_logits logits[:, -1, :] # (1, vocab_size) probs nn.functional.softmax(last_logits, dim-1) # Sample next token and reset RNG for reproducibility next_token torch.multinomial(probs, 1, generatorgenerator) generator.set_state(initial_state) # Append token to context and decode context torch.cat((context, next_token), dim-1) result.append(int_to_char[next_token.item()]) return .join(result)关键行逐一拆解代码作用备注context[:, -context_length:]裁剪超长上下文只保留最后context_length个 token配合位置嵌入表大小参见 code-gpt.mdmodel(context)前向传播输出形状 $(1, T, V)$即每个位置一份 logitslogits[:, -1, :]提取末位只有最后一个位置预测下一个 tokennn.functional.softmax(last_logits, dim-1)logits → 概率训练时cross_entropy内部自带 softmax生成时必须手动做参见 train-your-gpt.md 与 softmax.mdtorch.multinomial(probs, 1, generatorgenerator)按概率采样返回形状 $(1, 1)$ 的 token IDtorch.cat((context, next_token), dim-1)追加新 token生成的新 token 成为下一次前向的输入int_to_char[next_token.item()]解码把 token ID 还原为字符并累积成字符串其中torch.manual_seed(0)generator.set_state(initial_state)的组合值得单独说明每次采样前把 RNG 恢复到初始状态保证每一步采样都是可复现的——这是该问题测试用例能够精确校验输出文本的关键也是整个课程中用torch.manual_seed保证可复现性这一模式的延续同样的模式出现在 gpt-data-loader.md 和 train-your-gpt.md 中。逐步推演Walkthrough假设模型已训练完毕context_length 8初始context [[5, 12, 3]]new_chars 4步骤上下文长度动作采样到的 Token解码结果13无需裁剪对 3 个 token 前向从末位采样7e24无需裁剪对 4 个 token 前向15l35无需裁剪对 5 个 token 前向15l46无需裁剪对 6 个 token 前向3o最终输出ello。一旦上下文增长超过 8 个 token就会被裁剪为最后 8 个。注意每步上下文长度 1 的规律这正体现了生成一个、追加一个的循环本质与训练阶段输入窗口 偏移 1 的目标窗口参见 gpt-data-loader.md形成前后呼应——训练时模型学习看到前缀预测下一个生成时模型把学到的这个能力逐 token 兑现。时间与空间复杂度时间$O(n \cdot T^2 \cdot d)$其中 $n$ 是新生成字符数$T$ 是不断增长的上下文长度。二次项来自注意力机制对序列长度的依赖注意力是 $O(T^2)$这一点在 gpt-dataset.md 中也有提及空间$O(T \cdot V T^2)$用于存放概率分布和注意力矩阵。$T \cdot V$ 是 logits / 概率张量$T^2$ 是注意力分数矩阵。值得一提的实践优化既然每步只用到末位 logits实际推理时可以复用已计算的中间表示、只对最后一个位置做增量计算这正是 KV Cache 技术的动机可参考仓库中 kv-cache.md不过就本问题而言朴素实现已经足够因为它的目的是讲清楚生成循环本身。常见陷阱Common Pitfalls陷阱一忘记裁剪上下文不裁剪的话上下文会无限增长超过模型最大长度后位置嵌入索引越界直接崩溃。# Wrong: context grows indefinitely logits model(context) # crashes when context context_length # Correct: crop to max length before forward pass if context.shape[1] context_length: context context[:, -context_length:] logits model(context)这个陷阱的根源再次指向位置嵌入表nn.Embedding(context_length, model_dim)只能索引 $0 \sim C-1$见 code-gpt.md所以裁剪到context_length不是可选项而是硬性约束。陷阱二用 argmax 代替采样贪心解码argmax输出确定但常常重复的文本。本题期望的是用torch.multinomial做采样。# Wrong: deterministic, repetitive output next_token torch.argmax(probs, dim-1, keepdimTrue) # Correct: sampling introduces variety next_token torch.multinomial(probs, 1, generatorgenerator)argmax 与采样的差异本质上对应概率最大的 token与按概率分布随机抽取的 token之间的差异——前者总选同一个高峰后者允许以较小概率选中次优 token这正是文本多样性的来源。陷阱三取所有位置的 logits 而不是最后一个模型在每个位置都输出 logits但只有最后一个位置预测下一个 token。用其他位置的 logits得到的是对已经存在的 token 的预测。# Wrong: using all positions probs nn.functional.softmax(logits, dim-1) # (1, T, V) # Correct: only the last position predicts the next token last_logits logits[:, -1, :] # (1, V) probs nn.functional.softmax(last_logits, dim-1)为什么只能取末位因为训练时模型在每个位置学到的目标是该位置之后的下一个 token参见 train-your-gpt.md 中 logits 展平为 $(B \cdot T, V)$ 后与偏移 1 的目标做交叉熵的描述。位置 $t$ 的 logits 对应的是 token $t1$所以生成下一个字符时只有最后一个位置的预测才是当前上下文之外的信息。在 GPT 项目中的位置原文档说明这一步对应课程项目中的generate.py是整个课程的最后一道题。在完成以下链路之后——数据准备gpt-dataset.md文档将原始文本分词并构造输入-目标对数据加载gpt-data-loader.md文档批量生成输入窗口 偏移 1 目标窗口的训练张量模型构建code-gpt.md文档把嵌入层、Transformer 块、最终归一化与词表投影组装成完整 GPT模型训练train-your-gpt.md文档用 AdamW 交叉熵驱动模型学会预测下一个字符——最后一步就是这个生成函数训练好的 GPT 模型靠它真正开口说话。你只需喂入一段种子上下文哪怕只有一个字符模型就会逐字符生成产出模仿其训练数据模式的新文本。生成与训练在数学上互为镜像训练时把真实文本的 logits 与偏移 1 的目标做交叉熵softmax 由损失函数内部完成生成时手动 softmax 末位 logits 再采样——同一个分布、同一个模型只是从学切换到了用。关键要点Key Takeaways自回归生成一次只产出一个 token并把每次预测结果作为下一次的输入循环往复上下文窗口把模型的记忆限制在最近 $C$ 个 token 内更早的 token 被裁剪并遗忘裁剪不是优化而是位置嵌入表长度带来的硬约束采样优于贪心从分布中采样而非取 argmax能产出更多样、更自然的文本同时也是温度缩放、top-p 采样等高级解码策略的基础——掌握这一步你就掌握了所有现代 LLM 推理服务背后最核心的生成原语。【免费下载链接】leetcodeLeetcode solutions项目地址: https://gitcode.com/GitHub_Trending/leetcode1/leetcode创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
锦
锦皓数字建站
深耕本土企业品牌数字化升级,专注原创端正雅致商务官网,从视觉设计到稳定运维全程保驾护航。