两种生成场景下模型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()后再输入模型,但问题仍存在,且该操作会破坏计算图,不符合需求。
提问
- 该问题的原因是什么?
- 如何修复?
- 是否存在无需深度重训原模型的解决方案?
内容的提问来源于stack exchange,提问作者Nikio
相关产品推荐
相关产品推荐

