无token嵌入时,如何计算CLIP文本编码器的池化投影输出?
问题与解决方案
问题概述
需要将text_embeddings与CLIP文本编码器中间层输出结合,编码器输入为随机初始化的可学习提示嵌入。目标输出形状为[batch_size, embed_dim],但当前输出是[batch_size, seq_len, embed_dim];参考Transformers库的CLIP池化实现时发现其依赖input_ids,但当前场景没有该输入;尝试通过dummy输入添加CLS/SOS/EOS token,不确定方法是否正确。
核心解决方案
1. 正确构造含SOS/EOS的输入序列
CLIP的池化逻辑依赖EOS token的位置,因此需要给可学习提示嵌入ctx添加标准的SOS(起始token)和EOS(结束token)嵌入,构造符合CLIP输入格式的序列:[SOS, 提示嵌入, EOS]。可以通过dummy输入获取这两个特殊token的嵌入。
2. 中间层融合的正确实现
在指定的编码器中间层(示例为第4层)处理后,提取除原提示嵌入外的部分,与text_embeddings拼接,确保维度对齐后继续后续层的编码。
3. 无input_ids时的池化方法
由于没有真实的input_ids,可以利用我们构造序列时EOS的固定位置来提取池化输出——因为我们明确知道EOS在序列的最后一个有效位置,直接索引该位置即可。
修改后的完整代码
import torch from transformers import CLIPTextModelWithProjection, CLIPTokenizerFast # 初始化输入张量 text_embeddings = torch.randn(2, 4, 512) # 要融合的外部文本嵌入 ctx = torch.randn(2, 16, 512) # 可学习的提示嵌入 # 加载模型和分词器 hf_tokenizer = CLIPTokenizerFast.from_pretrained("wisdomik/QuiltNet-B-32") hf_text_encoder = CLIPTextModelWithProjection.from_pretrained("wisdomik/QuiltNet-B-32") device = torch.device("cuda" if torch.cuda.is_available() else "cpu") hf_text_encoder.to(device) text_embeddings = text_embeddings.to(device) ctx = ctx.to(device) # 提取CLIP文本编码器的核心组件 transformer = hf_text_encoder.text_model final_ln = transformer.final_layer_norm proj = hf_text_encoder.text_projection encoder = transformer.encoder embeddings = transformer.embeddings # 获取SOS和EOS的嵌入:通过dummy输入获取特殊token的嵌入 dummy_text = [""] # 空文本会被编码为[SOS, EOS] dummy_input = hf_tokenizer(dummy_text, return_tensors="pt").to(device) with torch.no_grad(): dummy_embeds = embeddings(input_ids=dummy_input.input_ids) sos_emb = dummy_embeds[:, 0:1, :] # SOS token嵌入,形状[1,1,512] eos_emb = dummy_embeds[:, 1:2, :] # EOS token嵌入,形状[1,1,512] # 构造完整输入序列:[SOS, 提示嵌入, EOS] x = torch.cat([ sos_emb.repeat(ctx.shape[0], 1, 1), # 扩展到batch_size ctx, eos_emb.repeat(ctx.shape[0], 1, 1), ], dim=1) # 遍历编码器层,在第4层后融合text_embeddings for idx, layer in enumerate(encoder.layers): x = layer(x, attention_mask=None, causal_attention_mask=None)[0] if idx == 4: # 去掉原序列中的SOS+提示嵌入部分,保留后续编码结果 x = x[:, 1 + ctx.shape[1]:, :] # 拼接text_embeddings和中间层输出 x = torch.cat([text_embeddings, x], dim=1) # 最终层归一化 x = final_ln(x) # 池化:取EOS位置的输出(此时EOS在序列最后一位) batch_size = x.shape[0] pooled_output = x[torch.arange(batch_size, device=device), x.shape[1]-1, :] # 投影得到最终输出 o = proj(pooled_output) print(f"Output shape: {o.shape}") # 应为[2, 512](假设embed_dim=512)
Transformers库CLIP池化实现代码(中文注释)
if self.eos_token_id == 2: # PR #24773之前的eos_token_id设置有误,此处保留原有逻辑 # 此类配置的CLIP模型无法正确处理新增的token # ------------------------------------------------------------ # text_embeds形状为[batch_size, sequence_length, transformer.width] # 取EOT(文本结束)token的特征(EOT是每个序列中ID最大的token) # 转换为torch.int是为了兼容ONNX:opset 14的argmax不支持int64输入 pooled_output = last_hidden_state[ torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device), input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1), ] else: # 配置已通过PR #24773更新了eos_token_id,支持新增token pooled_output = last_hidden_state[ torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device), # 需要找到第一个eos_token_id的位置(pad_token_id可能等于eos_token_id) # 注意:假设每个序列(batch维度)都包含一个eos_token_id(由分词器生成) (input_ids.to(dtype=torch.int, device=last_hidden_state.device) == self.eos_token_id) .int() .argmax(dim=-1), ]
内容的提问来源于stack exchange,提问作者user22617340
相关产品推荐
相关产品推荐

