CoMat
通过图像描述模型反向评估生成质量,解决 AI 画图文字漏字难题
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
通过图像描述模型反向评估生成质量,解决 AI 画图文字漏字难题
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象你是一名广告设计师,需要让 AI 根据文案生成一张产品图——你输入"一位穿着红色连衣裙的女性在海边散步,背景有椰子树"。结果 AI 画出了一位女性,但没有裙子,或者椰子树"隐身"了。这种"文字漏字"(prompt misalignment)现象,是当前主流 Text-to-Image(T2I)扩散模型的最大痛点之一。
2024 年 NeurIPS 接收的论文 CoMat: Aligning Text-to-Image Diffusion Model with Image-to-Text Concept Matching 提出了一个优雅的解决方案:让 AI 生成图像后,再用图像描述模型"回译"一遍,看看哪些文字概念被遗漏,然后用这些反馈来微调图像生成模型。这形成了一个图像→文字→图像的闭环对齐机制。

图1:CoMat 通过概念匹配机制,显著提升文字描述与生成图像的对齐程度。
CoMat 的技术方案包含两个核心模块,共同解决文字漏字问题:
模块一:图像到文本概念匹配(Image-to-Text Concept Matching)
研究者使用 BLIP 图像描述模型作为"裁判"——它能准确识别图像中包含的具体概念。当扩散模型生成一张图像后,CoMat 让 BLIP 对图像进行描述,然后比对原始 prompt 中有哪些词汇在描述中缺席。缺席的词汇意味着模型"忽略"了这些概念,这些 token 会被标记为需要"重新关注"的目标。
这个机制的本质是:用视觉感知模型(BLIP)来评估文本生成器(Diffusion),形成了一种自监督的质量反馈循环。
模块二:属性浓度模块(Attribute Concentration Module)
仅有概念匹配还不够——有些属性是弥散的,比如"颜色的分布"或"光照的质感",难以通过离散的 token 匹配来捕捉。CoMat 引入了基于 Grounded-SAM 的注意力集中机制,利用分割一切模型(SAM)准确定位图像中的实体区域,再将这些区域与 prompt 中的名词短语对齐。
具体来说,CoMat 会计算 prompt 中每个名词短语的注意力热图,将其与 SAM 分割出的实体掩码进行对比,对于注意力分散到错误区域的词,会施加额外的梯度信号,引导模型将注意力集中在正确区域。
CoMat 的工程实现完全基于 HuggingFace Diffusers 框架(版本要求 ≥ 0.16.0.dev0),这是一个目前最流行的开源扩散模型库。作者通过继承 StableDiffusionPipeline 和 StableDiffusionXLPipeline,重写了前向传播(forward)方法,加入了训练专用的参数:
training_timesteps:控制训练采样的时间步early_exit:允许在扩散过程的早期退出以加速训练detach_gradient:控制梯度是否解耦bp_on_trained:对已训练部分进行反向传播# TrainableSDPipeline 核心前向方法
def forward(
self,
prompt: Union[str, List[str]] = None,
training_timesteps: Optional[List[int]] = [],
early_exit: bool = False,
train_text_encoder: bool = False,
...
):
# 重写了 StableDiffusionPipeline 的 __call__ 方法
# 支持 LoRA 微调、文本编码器联合训练等
同时,training_script.py 采用了 Accelerate 分布式训练框架,支持多 GPU 并行训练。训练脚本使用 LoRA(Low-Rank Adaptation)方法对模型进行高效微调,默认 LoRA rank 设置为 128,这意味着可以在消费级 GPU 上进行部分实验。
CoMat 的训练流程并非单纯的扩散模型微调,而是引入了对抗训练(GAN)的混合优化策略。核心组件包括:
Fidelity Preservation 模块(保真度保持):训练一个判别器(Discriminator)来区分原始模型生成的图像和微调后模型生成的图像,防止模型在拟合新概念时丧失原有的生成质量。这个判别器基于 GAN-SD 架构,支持 SD1.5 和 SDXL 两个版本。
Mask Token Loss(掩码 Token 损失):对被判定为"忽略"的 token,施加额外的交叉熵损失,权重默认为 1e-3,促使模型在生成过程中重新重视这些概念。
Mask Pixel Loss(掩码像素损失):针对注意力集中的实体区域,计算像素级重建损失,权重为 5e-5,确保属性浓度模块精确定位目标实体。
CoMat 不是一个开箱即用的推理工具,而是一套完整的训练框架。根据示例训练脚本,单卡训练 SD1.5 至少需要:
| 硬件配置 | 最低要求 | 推荐配置 |
|---|---|---|
| GPU 显存 | 16GB | 24GB+(RTX 3090/A100) |
| 内存 | 16GB | 32GB+ |
| 存储 | 20GB | 50GB+ |
| 训练时间 | ~4小时(2000步) | — |
安装过程需要手动配置多个依赖:标准 PyTorch 环境、Grounded-SAM 分割模型、以及 BLIP 描述模型,整体复杂度较高,更适合有扩散模型研究经验的开发者。
CoMat 仍有几个值得关注的局限:目前仅支持英文 prompt,对中文等多语言场景尚未优化;属性浓度模块依赖 Grounded-SAM 的分割质量,在复杂场景下可能出现分割错误;尚未公开预训练模型检查点(Checkpoint),研究者需要自行训练。
作者在 README 中列出了明确的 TODO 清单,包括发布检查点、扩展到更多基础模型等,值得持续关注。
CoMat 代表着 AI 图像生成领域从"模型越大越好"向"对齐精度驱动"的趋势转变。2024 年下半年,随着 GPT-4o、Gemini 2.0 等多模态模型的发布,图像生成的对齐问题已成为行业焦点。CoMat 的概念匹配框架为这一问题提供了一个轻量且有效的解决思路,其核心思想——用图像理解模型反过来监督图像生成模型——已被多项后续工作引用和扩展。