You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

为自定义GPT-NEO模型实现do_sampling采样功能的技术咨询

问题原因

你自定义模型中使用的是贪心解码逻辑,每次固定选取概率最高的下一个token,不存在随机性,所以每次生成结果完全相同。而官方generate方法开启do_sample=True后会基于token的概率分布做随机采样,因此每次返回结果有差异。

修改方案

我们需要把原有的argmax贪心逻辑替换为随机采样逻辑,同时可以添加温度缩放、top-k过滤等常用采样策略,和官方生成逻辑对齐。

核心修改点

  1. 补充缺失的torch依赖导入
  2. 替换NEO类中forward方法的token选取逻辑,新增采样参数控制
  3. 可根据需求调整温度、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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 11:45:03