如何解码PyTorch中OpenAIGPTModel的输出结果?
解决OpenAIGPTModel输出解码问题
你当前使用的OpenAIGPTModel是基础编码器模型,仅输出隐藏层特征(last_hidden_state),没有附带语言建模(LM)头部,所以无法直接将其解码为文本。正确的做法是使用带LM头部的OpenAIGPTLMHeadModel,它会输出对应词表的预测得分(logits),可直接用于解码。
完整实现代码
from transformers import OpenAIGPTTokenizer, OpenAIGPTLMHeadModel import torch # 加载tokenizer和带LM头部的模型 tokenizer = OpenAIGPTTokenizer.from_pretrained('openai-gpt') model = OpenAIGPTLMHeadModel.from_pretrained('openai-gpt') # 输入处理 inputs = tokenizer("How does a kite fly?", return_tensors="pt") outputs = model(**inputs) # 获取logits:形状为(batch_size, sequence_length, vocab_size) logits = outputs.logits # 对每个位置取概率最高的token索引 predicted_token_ids = torch.argmax(logits, dim=-1)[0].tolist() # 解码为文本 decoded_text = tokenizer.decode(predicted_token_ids) print(decoded_text)
错误原因说明
- 直接解码
last_hidden_state:基础模型的隐藏层向量是高维特征(维度为hidden_size),并非对应词表的概率分布,tokenizer.decode无法将这类向量转换为有效文本,因此得到<unk>或乱码。 - 误将
last_hidden_state当logits:last_hidden_state的最后一维是hidden_size(而非词表大小vocab_size),强行取argmax得到的索引不在词表有效范围内,解码自然无意义。
补充说明
如果需要生成后续文本(而非仅预测输入序列的对应token),可以结合模型的生成方法:
# 生成后续文本,max_length控制总长度 generated_output = model.generate(**inputs, max_length=50) generated_text = tokenizer.decode(generated_output[0], skip_special_tokens=True) print(generated_text)
内容的提问来源于stack exchange,提问作者Jellicle
相关产品推荐
相关产品推荐

