为自定义GPT-NEO模型实现do_sampling采样功能的技术咨询
问题原因
你自定义模型中使用的是贪心解码逻辑,每次固定选取概率最高的下一个token,不存在随机性,所以每次生成结果完全相同。而官方generate方法开启do_sample=True后会基于token的概率分布做随机采样,因此每次返回结果有差异。
修改方案
我们需要把原有的argmax贪心逻辑替换为随机采样逻辑,同时可以添加温度缩放、top-k过滤等常用采样策略,和官方生成逻辑对齐。
核心修改点
- 补充缺失的
torch依赖导入 - 替换
NEO类中forward方法的token选取逻辑,新增采样参数控制 - 可根据需求调整温度、top-k等超参数控制生成效果
修改后完整代码
import torch import numpy as np from transformers import GPTNeoForCausalLM, GPT2Tokenizer import coremltools as ct tokenizer = GPT2Tokenizer.from_pretrained("gpt2") sentence_fragment = "The Oceans are" class NEO(torch.nn.Module): def __init__(self, model, temperature=1.0, top_k=50): super(NEO, self).__init__() self.next_token_predictor = model self.temperature = temperature self.top_k = top_k def forward(self, x): sentence = x predictions, _ = self.next_token_predictor(sentence) # 取出最后一个位置的logits next_token_logits = predictions[-1, :] # 温度缩放:调整随机性强弱 if self.temperature > 0: next_token_logits = next_token_logits / self.temperature # Top-K过滤:只保留概率最高的k个token,避免生成异常内容 if self.top_k > 0: topk_vals, _ = torch.topk(next_token_logits, self.top_k) k_th_val = topk_vals[-1] next_token_logits[next_token_logits < k_th_val] = -float('inf') # 转换为概率分布 probs = torch.softmax(next_token_logits, dim=-1) # 随机采样1个token替代原固定取最大值的逻辑 token = torch.multinomial(probs, num_samples=1) sentence = torch.cat((sentence, token), 0) return sentence token_predictor = GPTNeoForCausalLM.from_pretrained("EleutherAI/gpt-neo-125M", torchscript=True).eval() context = torch.tensor(tokenizer.encode(sentence_fragment)) random_tokens = torch.randint(10000, (5,)) traced_token_predictor = torch.jit.trace(token_predictor, random_tokens) # 初始化模型时可自定义采样参数 model = NEO(model=traced_token_predictor, temperature=0.7, top_k=50) scripted_model = torch.jit.script(model) # 自定义模型推理 sentence_fragment = "The Oceans are" for i in range(10): context = torch.tensor(tokenizer.encode(sentence_fragment)) torch_out = scripted_model(context) sentence_fragment = tokenizer.decode(torch_out) print("Custom model: {}".format(sentence_fragment)) # 官方模型推理对照 model = GPTNeoForCausalLM.from_pretrained("EleutherAI/gpt-neo-125M", torchscript=True).eval() sentence_fragment = "The Oceans are" input_ids = tokenizer(sentence_fragment, return_tensors="pt").input_ids gen_tokens = model.generate(input_ids, do_sample=True, max_length=20, temperature=0.7, top_k=50) gen_text = tokenizer.batch_decode(gen_tokens)[0] print("Stock model: "+gen_text)
参数说明
temperature:温度系数,取值越大生成随机性越强,取值越小越接近贪心解码,设为0等价于原贪心逻辑top_k:只保留概率最高的k个候选token,避免采样到极低概率的异常token,设为0则不做过滤- 如果需要和官方逻辑完全对齐,还可以自行补充top-p采样、重复惩罚等逻辑
修改完成后运行代码,自定义模型每次生成的结果就会具备随机性,和官方采样效果对齐。
内容的提问来源于stack exchange,提问作者Olexander Korenyuk
相关产品推荐
相关产品推荐

