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

两种生成场景下模型Token差异问题:循环求解错误Token暴增的修复

模型生成动作时的Token错误问题

模型与训练背景

我有一个模型,输入状态配置(比如魔方的整数序列)后会生成动作(取值0-5),该动作可将当前状态转换为另一个状态。模型初始训练目标是通过循环生成各类Token(动作、状态中的数字或编码空格的数字),将配置转换为目标解状态identical_state。

另有一个确定性重排状态的函数reindex,相关代码如下:

def reindex(state,moves):
    # args are tensors
    for m in dict_generators[moves]:
        state = state[m]
    return(state)

def build_state(n):
    seq_moves = torch.randint(0, 5, (n,), device=device)
    state = reindex(identical_state,seq_moves)
    return state

for i in range(400):
    # encoding
    x = build_state(int(i / 20)+1)
    x = x + 8
    # add "next line" tokens, generate encoded torch.Size([26])
    x = torch.cat((x, torch.tensor([7], device=device)), dim=0)
    x = torch.cat((torch.tensor([7], device=device), x), dim=0)
    print(model.generate(x[None],temperature=temperature,top_k=top_k))

编码规则

模型接收的是编码后的输入:

  • 状态中所有数值加8
  • Token0-5:表示动作
  • Token6:表示状态已达解状态
  • Token7:分隔两个配置
  • Token8:状态中0的编码,是所有配置的起始标记

上述单次生成代码运行效果良好,几乎仅生成动作Token,仅极少数情况下会返回错误Token8。

核心问题

当运行循环生成完整求解序列的代码时(通过reindex更新状态,同时解码、重新编码),错误Token8的数量激增——尽管输入模型的仍是符合格式的配置,只是由前一状态重排而来,而非新生成的配置。相关代码如下:

def predict_solution(x):
    # arg tensor torch.Size([24])
    sol = []
    list_probs = []
    # encode:
    x = x + 8
    # add "next line" tokens, generate encoded torch.Size([26])
    x = torch.cat((x, torch.tensor([7], device=device)), dim=0)
    x = torch.cat((torch.tensor([7], device=device), x), dim=0)
    # make it torch.Size([1,26]) to pass to the model
    x = x[None]
    for it in range(max_sol_length):
        y,probs = model.generate(x, temperature=temperature, top_k=top_k)
        # y, probs are two gradients of Size [1,1], probs keeps gradient

        if torch.equal(y, torch.tensor([[8]])):
            print("wrong token")
            # force it someway to be a movement in order to apply reindex below

        list_probs.append(probs)
        sol.append(y)

        # decode x into a tensor of torch.Size([24]) to pass to reindex
        x = x.squeeze(0)
        x = x-8
        x = x[1:-1]
        x = reindex(x,y)
        x = x[None] # unsqueeze(0)

        if torch.equal(x, identical_state[None]):
            return sol, list_probs

        # re-encode to pass to the model again
        x = x + 8
        x = torch.cat((x, torch.tensor([[7]], device=device)), dim=1)
        x = torch.cat((torch.tensor([[7]], device=device), x), dim=1)
    
    if it == max_sol_length - 1:
        print("solution not found")
    return sol, list_probs


for i in range(400):
    x = build_state(int(i / 20)+1)
    predict_solution(x)

我需要保留模型采样动作时的概率梯度以用于后续强化学习,曾尝试在循环末尾对x执行detach()后再输入模型,但问题仍存在,且该操作会破坏计算图,不符合需求。

提问

  1. 该问题的原因是什么?
  2. 如何修复?
  3. 是否存在无需深度重训原模型的解决方案?

内容的提问来源于stack exchange,提问作者Nikio

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 16:13:11