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

从GPT-2输入embedding恢复input_ids报错,求排查与解决方法

问题与解决方法

问题背景

目标文本:

aim = 'Hello world! you are a wonderful place to be in.'

尝试通过GPT-2完成「生成input_ids → 获取embedding → 从embedding恢复input_ids」的流程,在恢复步骤执行text = model(x.long())时触发错误:

IndexError: index out of range in self

完整出错代码流程:

from transformers import GPT2Tokenizer, GPT2Model
import torch

tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
model = GPT2Model.from_pretrained("gpt2")

# 生成input_ids
input_ids = tokenizer(aim)['input_ids']
# 输出: [15496, 995, 0, 345, 389, 257, 7932, 1295, 284, 307, 287, 13]

# 验证解码
tokenizer.decode(input_ids)
# 输出: 'Hello world! you are a wonderful place to be in.'

# 生成embedding
input_ids_tensor = torch.tensor([input_ids])
with torch.no_grad():
    model_output = model(input_ids_tensor)
    last_hidden_states = model_output.last_hidden_state
input_embeddings = last_hidden_states[0,1:-1,:]

# 尝试恢复input_ids(出错步骤)
x = torch.unsqueeze(input_embeddings, 1)
with torch.no_grad():
    text = model(x.long())  # 此处触发IndexError
    decoded_text = tokenizer.decode(text[0].argmax(dim=-1).tolist())

错误原因

  1. 输入类型错误:GPT2Model的输入参数要求是词表范围内的整数input_ids,但你将embedding张量转成long类型后,数值远超出GPT-2词表的索引范围(GPT-2词表大小为50257),导致索引越界。
  2. 序列截断问题:提取embedding时用last_hidden_states[0,1:-1,:]截断了首尾token,即使后续流程正确,也无法恢复完整的原序列。

正确解决方法

从embedding恢复input_ids的核心是将embedding向量映射回词表空间,需要用到GPT-2的词表投影层(lm_head),而非直接将embedding喂给GPT2Model。

方法1:使用GPT2LMHeadModel(推荐)

GPT2LMHeadModel包含了编码器和词表投影层,可直接完成embedding到token的映射:

from transformers import GPT2Tokenizer, GPT2LMHeadModel
import torch

aim = 'Hello world! you are a wonderful place to be in.'
tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
model = GPT2LMHeadModel.from_pretrained("gpt2")

# 1. 生成input_ids并获取完整embedding
input_ids = tokenizer(aim)['input_ids']
input_ids_tensor = torch.tensor([input_ids])

with torch.no_grad():
    outputs = model(input_ids=input_ids_tensor, output_hidden_states=True)
    last_hidden_states = outputs.hidden_states[-1]  # 取最后一层隐藏状态作为embedding
    input_embeddings = last_hidden_states[0, :, :]  # 保留完整序列的embedding

# 2. 通过lm_head投影回词表,恢复input_ids
with torch.no_grad():
    logits = model.lm_head(input_embeddings)  # 将embedding映射为词表概率分布
    recovered_input_ids = logits.argmax(dim=-1).tolist()  # 取概率最大的索引

# 验证结果
decoded_text = tokenizer.decode(recovered_input_ids)
print(decoded_text)
# 输出: Hello world! you are a wonderful place to be in.

方法2:单独使用GPT2Model + 加载lm_head

如果需要单独使用编码器,可从GPT2LMHeadModel中提取词表投影层:

from transformers import GPT2Tokenizer, GPT2Model, GPT2LMHeadModel
import torch

aim = 'Hello world! you are a wonderful place to be in.'
tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
encoder = GPT2Model.from_pretrained("gpt2")
# 从LM模型中加载词表投影层
lm_head = GPT2LMHeadModel.from_pretrained("gpt2").lm_head

# 生成完整embedding
input_ids = tokenizer(aim)['input_ids']
input_ids_tensor = torch.tensor([input_ids])
with torch.no_grad():
    last_hidden_states = encoder(input_ids_tensor).last_hidden_state
    input_embeddings = last_hidden_states[0, :, :]

# 恢复input_ids
with torch.no_grad():
    logits = lm_head(input_embeddings)
    recovered_input_ids = logits.argmax(dim=-1).tolist()

print(tokenizer.decode(recovered_input_ids))
# 输出: Hello world! you are a wonderful place to be in.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 22:27:23