关于替换Stable Diffusion v2.1文本编码器为图像编码器的技术疑问
我尝试将Stable Diffusion的文本编码器替换为对应的图像编码器,以便输入图像而非文本。根据Stable Diffusion文档,该模型使用来自OpenCLIP的ViT/H预训练文本编码器。由于CLIP的文本编码器与图像编码器共享同一潜在空间,理论上可直接替换无需额外训练即可正常运行。
但实际中,我发现两者生成的文本嵌入存在差异。
通过Stable Diffusion文本编码器获取嵌入的代码
prompt = 'dress, long sleeve' model_key = "./models--stabilityai--stable-diffusion-2-1-base/" pipe = StableDiffusionPipeline.from_pretrained(model_key, torch_dtype=self.precision_t) self.tokenizer = pipe.tokenizer self.text_encoder = pipe.text_encoder inputs = self.tokenizer(prompt, padding='max_length', max_length=self.tokenizer.model_max_length, return_tensors='pt') embeddings = self.text_encoder(inputs.input_ids.to(self.device))[0]
通过OpenCLIP文本编码器获取嵌入的代码
model, _, preprocess = open_clip.create_model_and_transforms('ViT-H-14', pretrained='laion2b_s32b_b79k') model.eval() tokenizer = open_clip.get_tokenizer('ViT-H-14') text = tokenizer([prompt]) text_features = model.encode_text(text)
主要差异在于,Stable Diffusion文本编码器生成的embeddings维度为(1, 77, 1024),而OpenCLIP文本编码器生成的text_features维度为(1, 1024)。
我有两个技术问题:
- 应使用OpenCLIP的哪个文本编码器才能得到与Stable Diffusion文本编码器相同的嵌入?
- 与Stable Diffusion文本编码器共享同一潜在空间的对应图像编码器是哪个?
问题解答
1. 获取与Stable Diffusion匹配的文本嵌入
要得到和Stable Diffusion文本编码器相同的嵌入,需要使用OpenCLIP中未做池化处理的文本编码器输出。
Stable Diffusion的文本编码器输出的是包含所有token(包括[CLS]和padding token)的序列嵌入,维度为(batch_size, seq_len, hidden_size);而OpenCLIP的encode_text()方法默认返回的是经过CLS token池化后的全局文本特征,维度为(batch_size, hidden_size)。
你需要直接调用OpenCLIP文本编码器的底层模块,获取原始的token级嵌入,示例代码如下:
import torch import open_clip prompt = 'dress, long sleeve' model, _, preprocess = open_clip.create_model_and_transforms('ViT-H-14', pretrained='laion2b_s32b_b79k') model.eval() tokenizer = open_clip.get_tokenizer('ViT-H-14') text_tokens = tokenizer([prompt]) with torch.no_grad(): # 获取未池化的原始token嵌入 text_embeddings = model.text_encoder(text_tokens)[0]
此时text_embeddings的维度会是(1, 77, 1024),和Stable Diffusion的输出完全匹配。Stable Diffusion使用的文本编码器就是该OpenCLIP模型的文本分支,权重完全对齐,因此只需获取未池化的输出即可。
2. 对应的图像编码器
与Stable Diffusion文本编码器共享同一潜在空间的图像编码器,就是OpenCLIP的ViT-H-14图像编码器(预训练权重为laion2b_s32b_b79k)。
该图像编码器是同一OpenCLIP模型的视觉分支,和文本分支共享潜在空间。默认情况下,encode_image()方法输出的是池化后的全局图像特征(维度(1, 1024)),若要适配Stable Diffusion的输入格式(需要(1, 77, 1024)维度的序列嵌入),可以通过扩展维度的方式调整,示例代码如下:
# 假设已用preprocess处理好输入图像,得到image_tensor(维度为(1, 3, 224, 224)) with torch.no_grad(): image_features = model.encode_image(image_tensor) # 将全局图像特征扩展为77长度的序列,匹配文本嵌入的序列长度 image_embeddings = image_features.unsqueeze(1).repeat(1, 77, 1)
这种方式是基础的格式适配,理论上无需额外训练即可生成对应图像风格的内容;若追求更好效果,可在此基础上进行微调。
内容的提问来源于stack exchange,提问作者Nagabhushan S N

