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

无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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 16:20:56