从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())
错误原因
- 输入类型错误:
GPT2Model的输入参数要求是词表范围内的整数input_ids,但你将embedding张量转成long类型后,数值远超出GPT-2词表的索引范围(GPT-2词表大小为50257),导致索引越界。 - 序列截断问题:提取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
相关产品推荐
相关产品推荐

