如何在不使用池化操作的情况下获取长文本的Embedding?
长文本无池化生成Embedding方案(基于GPT类模型)
针对你要处理2000token左右长文本、且完全不使用池化操作生成Embedding的需求,直接用GPT类模型提取特定位置token的隐藏状态即可——因为池化是对多个token向量做聚合(均值、最大、拼接等),而取单个token的向量不属于池化操作。
核心思路
用支持长序列的GPT变体模型(原生GPT2仅支持1024token,满足不了2000token需求),直接提取文本序列中最后一个有效token的隐藏状态作为整个文本的Embedding。自回归模型的最后一个token隐藏状态天然包含了前文所有信息,用作文本Embedding合理性拉满。
完整代码示例
import torch from transformers import AutoTokenizer, AutoModel # 选择支持2048token的长序列GPT类模型,这里用轻量化的gpt-neo-1.3B,也可以用更大的2.7B版本 model_name = "EleutherAI/gpt-neo-1.3B" tokenizer = AutoTokenizer.from_pretrained(model_name) # GPT系列模型默认没有pad token,用eos token替代 tokenizer.pad_token = tokenizer.eos_token model = AutoModel.from_pretrained(model_name) # 替换成你的2000token左右长文本 long_text = "你的长文本内容..." # 编码文本,设置最大长度为模型支持的2048,超长时自动截断 inputs = tokenizer( long_text, return_tensors="pt", max_length=2048, truncation=True, padding="max_length" ) # 推理获取隐藏状态,避免计算梯度节省资源 with torch.no_grad(): outputs = model(**inputs) # last_hidden_state形状: (batch_size, sequence_length, hidden_size) # 找到有效token的长度(排除pad) valid_token_len = inputs["attention_mask"].sum(dim=1) # 取最后一个有效token的索引 last_token_idx = valid_token_len - 1 # 提取该位置的向量作为文本Embedding text_embedding = outputs.last_hidden_state[0, last_token_idx, :]
额外说明
- 如果不需要单向量的文本Embedding,而是要保留整个序列所有token的Embedding,直接取
outputs.last_hidden_state即可,这也完全没有用到池化操作。 - 若你坚持要用原生GPT2,只能把长文本拆成多个1024token的片段,每个片段单独提取最后一个token的Embedding,但这种方式不如直接用长序列模型高效。
内容的提问来源于stack exchange,提问作者Penguin
相关产品推荐
相关产品推荐

