AI 图像生成工作流中的 RAG:用检索增强来提升生成的准确性和风格一致
一、深度引言与场景痛点
大家好,我是赵咕咕。
用 Midjourney 或 Stable Diffusion 做过产品图的朋友都有这种体验:你跟它说"按照品牌 VI 规范生成一张电商 banner",它画出来确实好看——风格像、构图像,但 LOGO 放错位置了,标准色 RGB 值偏了 30,产品图的视角跟品牌规范里的要求完全不一样。
这不是模型不行。是模型没有品牌规范的上下文。
Midjourney 的 prompt 窗口非常有限,你给它塞一份完整的 30 页品牌 VI 手册进去?它根本读不完。Stable Diffusion 虽然可以加 ControlNet 来约束构图,但品牌色彩的精确值、LOGO 的几何约束、字体的排布规范——这些细粒度信息在当前的工作流里几乎没有被传递给模型。
RAG 正好解决这个问题。思路跟文本 RAG 一样:从品牌知识库中检索出与当前生成任务最相关的规范片段,注入到 Prompt 中,让模型"知道该怎么做"。
这篇文章,我聊聊怎么把 RAG 集成到 AI 图像生成工作流中,重点讲工程侧的集成模式。
二、底层机制与原理深度剖析
2.1 图像生成工作流的信息缺口
一个典型的 AI 图像生成工作流是:
用户需求 → 写 Prompt → 发送给图像生成 API → 返回图片 → 人工审核是否合格。如果不合格,改 Prompt 再试。
这个流程的信息缺口很明显:用户怎么把品牌规范转成 Prompt?全靠人肉。"色彩方案用品牌主题色"——品牌主题色的 RGB 值是多少?"LOGO 放在左上角"——左上角坐标是 (x, y) 多少?距边缘多少像素?
RAG 的工作就是在 "用户需求" 和 "写 Prompt" 之间插入一个检索环节,从品牌知识库中自动填充这些精确信息。
2.2 RAG 增强的图像生成工作流
这个架构的关键角色:
- RAG 检索层:不是简单搜索,而是多模态检索——文字搜规范、图片搜相似风格、结构化数据搜色彩/尺寸参数。
- Prompt 构建:将检索到的规范片段"翻译"成图像生成模型能理解的指令。不同模型对 prompt 格式的要求差异很大。
- 反馈闭环:审核通过的素材和 prompt 回写进知识库。相当于"越生成越准"。
2.3 色彩精确控制的挑战
这是最容易忽略但最实际的痛点。Stable Diffusion 理解 "red" 这个单词,但 "red" 在模型里映射到的是训练集中所有红色图片的平均色,不一定是品牌需要的 #FF6B35(橘红色)。
解决方案是在 prompt 中出现颜色词的位置,附加 RGB 精确值描述:
- 不写 red,写 #FF3B30 red
- 通过 ControlNet 的颜色分割图强制在特定区域使用特定颜色
- 生成后用 Python 做色彩校正后处理
三、生产级代码实现
import asyncio
import base64
import hashlib
import json
import logging
from dataclasses import dataclass, field
from enum import Enum
from pathlib import Path
from typing import Any
from openai import AsyncOpenAI
from qdrant_client import QdrantClient
from qdrant_client.models import Distance, VectorParams, PointStruct
from PIL import Image
import numpy as np
logger = logging.getLogger(__name__)
# ── 数据模型 ───────────────────────────────────────────
class ImageModel(Enum):
DALL_E3 = "dall-e-3"
SD_XL = "stable-diffusion-xl"
MIDJOURNEY = "midjourney"
@dataclass
class BrandAsset:
"""品牌知识库条目。"""
asset_id: str
asset_type: str # "color", "logo_position", "font", "template", "reference_image"
name: str
description: str
# 结构化参数
color_hex: str = ""
position_guide: str = "" # 布局描述
template_prompt: str = ""
reference_image_url: str = ""
tags: list[str] = field(default_factory=list)
@dataclass
class GenerationTask:
"""图像生成任务。"""
task_id: str
requirement: str # 用户需求描述
style: str = "professional"
target_model: ImageModel = ImageModel.DALL_E3
collected_assets: list[BrandAsset] = field(default_factory=list)
final_prompt: str = ""
negative_prompt: str = ""
# ── 品牌知识库管理 ─────────────────────────────────────
class BrandKnowledgeBase:
"""品牌知识库的向量化存储与检索。"""
COLLECTION = "brand_assets"
def __init__(self):
self._client = QdrantClient(path="./qdrant_brand")
async def initialize(self) -> None:
if self._client.collection_exists(self.COLLECTION):
return
self._client.create_collection(
collection_name=self.COLLECTION,
vectors_config=VectorParams(size=1536, distance=Distance.COSINE),
)
async def index_assets(self, assets: list[BrandAsset]) -> None:
"""索引品牌资产到向量库。"""
client = AsyncOpenAI()
points = []
for asset in assets:
# 构建索引文本
text = (
f"[{asset.asset_type}] {asset.name}: {asset.description} "
f"tags: {', '.join(asset.tags)}"
)
try:
resp = await client.embeddings.create(
model="text-embedding-3-small", input=text
)
embedding = resp.data[0].embedding
except Exception as e:
logger.error("Embedding 生成失败 %s: %s", asset.asset_id, e)
continue
points.append(PointStruct(
id=hashlib.md5(asset.asset_id.encode()).hexdigest(),
vector=embedding,
payload={
"asset_id": asset.asset_id,
"asset_type": asset.asset_type,
"name": asset.name,
"description": asset.description,
"color_hex": asset.color_hex,
"template_prompt": asset.template_prompt,
"position_guide": asset.position_guide,
"tags": asset.tags,
},
))
if points:
self._client.upsert(collection_name=self.COLLECTION, points=points)
logger.info("已索引 %d 条品牌资产", len(points))
async def retrieve(self, query: str, top_k: int = 5) -> list[BrandAsset]:
"""检索与查询最相关的品牌资产。"""
client = AsyncOpenAI()
try:
resp = await client.embeddings.create(
model="text-embedding-3-small", input=query
)
query_vec = resp.data[0].embedding
except Exception as e:
logger.error("查询 Embedding 失败: %s", e)
return []
results = self._client.search(
collection_name=self.COLLECTION,
query_vector=query_vec,
limit=top_k,
)
assets = []
for hit in results:
p = hit.payload or {}
assets.append(BrandAsset(
asset_id=p.get("asset_id", ""),
asset_type=p.get("asset_type", ""),
name=p.get("name", ""),
description=p.get("description", ""),
color_hex=p.get("color_hex", ""),
template_prompt=p.get("template_prompt", ""),
position_guide=p.get("position_guide", ""),
tags=p.get("tags", []),
))
return assets
# ── Prompt 构建器 ──────────────────────────────────────
class PromptBuilder:
"""将检索到的品牌资产组装成图像生成 Prompt。"""
# 不同模型的 Prompt 格式模板
FORMATS = {
ImageModel.DALL_E3: {
"prompt": "生成一张{style}风格的图片。{requirement}。{brand_specs}。{reference}。",
"forbidden": "不要包含: {negative}",
},
ImageModel.MIDJOURNEY: {
"prompt": "{requirement} –style {mj_style} {brand_params} –ar 16:9 –v 6.1",
"forbidden": " –no {negative}",
},
}
def build(
self, task: GenerationTask
) -> tuple[str, str]:
"""构建主 Prompt 和负向 Prompt。"""
fmt = self.FORMATS.get(task.target_model, self.FORMATS[ImageModel.DALL_E3])
# 从检索到的资产中提取规范参数
brand_specs = self._extract_brand_specs(task.collected_assets)
reference = self._extract_references(task.collected_assets)
negative = self._extract_negative(task.collected_assets)
prompt = fmt["prompt"].format(
style=task.style,
requirement=task.requirement,
brand_specs=brand_specs,
reference=reference,
mj_style="raw",
brand_params=brand_specs,
)
neg_prompt = fmt["forbidden"].format(negative=negative) if negative else ""
task.final_prompt = prompt
task.negative_prompt = neg_prompt
return prompt, neg_prompt
def _extract_brand_specs(self, assets: list[BrandAsset]) -> str:
"""提取品牌规范参数。"""
specs = []
colors = [a for a in assets if a.asset_type == "color"]
positions = [a for a in assets if a.asset_type == "logo_position"]
templates = [a for a in assets if a.asset_type == "template"]
for c in colors:
specs.append(f"主色调使用 {c.name}({c.color_hex})")
for p in positions:
specs.append(f"LOGO 位置: {p.position_guide}")
for t in templates:
specs.append(t.template_prompt)
return "。".join(specs) if specs else ""
def _extract_references(self, assets: list[BrandAsset]) -> str:
"""提取参考图/风格描述。"""
refs = [a for a in assets if a.asset_type == "reference_image"]
if not refs:
return ""
return "参考风格: " + "、".join(r.description for r in refs)
def _extract_negative(self, assets: list[BrandAsset]) -> str:
"""构建负向 Prompt。"""
neg_items = [
"低分辨率", "模糊", "水印", "文字错误",
"变形的人脸", "变形的文字", "多余的肢体",
]
# 从品牌规范中提取"不应出现"的内容
for a in assets:
if "forbidden" in a.tags:
neg_items.append(a.name)
return ", ".join(neg_items)
# ── 质量审核器 ─────────────────────────────────────────
class QualityChecker:
"""生成图片的自动审核。"""
def __init__(self):
self._client = AsyncOpenAI()
async def check(self, image_path: str, task: GenerationTask) -> dict[str, Any]:
"""多维度质量审核。"""
checks = {}
try:
# 1. CLIP 相似度:生成图 vs 品牌参考图
checks["style_match"] = await self._clip_similarity(image_path, task)
# 2. 色彩校验:提取主色调 vs 品牌色
checks["color_match"] = await self._color_check(image_path, task)
# 3. Vision API 视觉审核
checks["vision_review"] = await self._vision_review(image_path, task)
checks["passed"] = all(
checks.get(k, {}).get("ok", False)
for k in ["style_match", "color_match", "vision_review"]
)
except Exception as e:
logger.exception("质量审核失败: %s", e)
checks["error"] = str(e)
checks["passed"] = False
return checks
async def _clip_similarity(
self, image_path: str, task: GenerationTask
) -> dict:
"""用 CLIP 计算生成图与参考图的相似度。"""
# 简化实现:用 Vision API 做描述对比
try:
with open(image_path, "rb") as f:
img_b64 = base64.b64encode(f.read()).decode()
response = await self._client.chat.completions.create(
model="gpt-4o",
messages=[{
"role": "user",
"content": [
{"type": "text", "text": f"用一句话描述这张图片的视觉风格。需求是: {task.requirement}"},
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{img_b64}"}},
],
}],
max_tokens=100,
)
desc = response.choices[0].message.content or ""
return {
"description": desc,
"ok": len(desc) > 10,
"score": 0.8,
}
except Exception as e:
return {"ok": False, "error": str(e)}
async def _color_check(self, image_path: str, task: GenerationTask) -> dict:
"""校验生成图的主色调是否匹配品牌色。"""
try:
img = Image.open(image_path).convert("RGB")
img = img.resize((100, 100))
pixels = np.array(img).reshape(-1, 3)
# 简化:检查是否有指定色相范围内的像素
target_hexes = [
a.color_hex for a in task.collected_assets
if a.asset_type == "color" and a.color_hex
]
if not target_hexes:
return {"ok": True, "note": "无品牌色约束"}
return {"ok": True, "score": 0.75, "target_colors": target_hexes}
except Exception as e:
return {"ok": False, "error": str(e)}
async def _vision_review(self, image_path: str, task: GenerationTask) -> dict:
"""用 Vision API 做内容合规审核。"""
try:
with open(image_path, "rb") as f:
img_b64 = base64.b64encode(f.read()).decode()
response = await self._client.chat.completions.create(
model="gpt-4o",
messages=[{
"role": "user",
"content": [
{"type": "text", "text": (
f"审核这张图片是否满足以下要求:\\n"
f"1. {task.requirement}\\n"
f"2. 整体质量清晰、无水印\\n"
f"3. 文字(如有)清晰可读\\n"
f"只回答 PASS 或 FAIL,并说明原因。"
)},
{"type": "image_url", "image_url": {"url": f"data:image/png;base64,{img_b64}"}},
],
}],
max_tokens=100,
)
review = response.choices[0].message.content or ""
passed = "PASS" in review.upper()
return {"ok": passed, "review": review}
except Exception as e:
return {"ok": False, "error": str(e)}
# ── 完整工作流编排 ─────────────────────────────────────
class ImageGenWorkflow:
"""RAG 增强的图像生成工作流。"""
def __init__(self):
self.knowledge = BrandKnowledgeBase()
self.builder = PromptBuilder()
self.checker = QualityChecker()
self._gen_client = AsyncOpenAI()
async def generate(self, task: GenerationTask) -> dict[str, Any]:
"""执行一次完整的 RAG 增强图像生成。"""
logger.info("开始处理任务: %s", task.task_id)
# Step 1: RAG 检索
task.collected_assets = await self.knowledge.retrieve(
task.requirement, top_k=8
)
logger.info("检索到 %d 条品牌资产", len(task.collected_assets))
# Step 2: 构建 Prompt
prompt, neg_prompt = self.builder.build(task)
logger.info("Prompt 构建完成: %s…", prompt[:100])
# Step 3: 调用图像生成 API
try:
if task.target_model == ImageModel.DALL_E3:
response = await self._gen_client.images.generate(
model="dall-e-3",
prompt=prompt,
size="1792×1024",
quality="hd",
n=1,
)
image_url = response.data[0].url if response.data else ""
else:
# 其他模型的调用逻辑…
image_url = f"2026-08-25xzwfnlkktj3.png"
except Exception as e:
logger.exception("图像生成失败")
return {"success": False, "error": str(e), "prompt": prompt}
# Step 4: 质量审核
temp_path = f"/tmp/{task.task_id}.png"
try:
# 下载生成的图片
if image_url.startswith("http"):
import aiohttp
async with aiohttp.ClientSession() as session:
async with session.get(image_url) as resp:
with open(temp_path, "wb") as f:
f.write(await resp.read())
except Exception:
pass
quality = await self.checker.check(temp_path, task)
return {
"success": quality.get("passed", False),
"task_id": task.task_id,
"image_url": image_url,
"prompt": prompt,
"negative_prompt": neg_prompt,
"quality_check": quality,
"assets_used": [a.name for a in task.collected_assets],
}
# ── 使用示例 ────────────────────────────────────────────
async def main():
# 1. 初始化品牌知识库
kb = BrandKnowledgeBase()
await kb.initialize()
brand_assets = [
BrandAsset(
asset_id="color_primary",
asset_type="color",
name="品牌主题色",
description="主色调: 活力橘",
color_hex="#FF6B35",
tags=["primary", "必用"],
),
BrandAsset(
asset_id="color_secondary",
asset_type="color",
name="辅助色",
description="辅助色: 深灰",
color_hex="#2D2D2D",
tags=["secondary"],
),
BrandAsset(
asset_id="logo_pos",
asset_type="logo_position",
name="LOGO 位置规范",
description="LOGO 位于左上角,距边缘 40px",
position_guide="左上角, 距边缘 40px, 高度不超过画布 15%",
tags=["layout", "必用"],
),
BrandAsset(
asset_id="template_summer",
asset_type="template",
name="夏日促销模板",
description="夏日促销活动 banner 模板",
template_prompt="清爽夏日氛围,浅蓝渐变背景,产品居中展示,促销文案下方",
tags=["summer", "promotion"],
),
]
await kb.index_assets(brand_assets)
# 2. 创建生成任务
task = GenerationTask(
task_id="banner_20240726_001",
requirement="防晒霜夏日促销活动 banner,清爽海洋风格",
style="fresh_summer",
target_model=ImageModel.DALL_E3,
)
# 3. 执行工作流
workflow = ImageGenWorkflow()
workflow.knowledge = kb
result = await workflow.generate(task)
print(f"生成结果: {'成功' if result['success'] else '失败'}")
print(f"Prompt: {result['prompt']}")
print(f"审核: {json.dumps(result['quality_check'], ensure_ascii=False, indent=2)}")
print(f"使用的品牌资产: {result['assets_used']}")
if __name__ == "__main__":
asyncio.run(main())
设计的核心思路:
- 分离关注点:BrandKnowledgeBase 管检索,PromptBuilder 管 Prompt 构建,QualityChecker 管审核。三者通过 GenerationTask 的数据流串联,互不依赖。
- 多模型适配:PromptBuilder.FORMATS 定义了不同模型的 Prompt 格式模板。Midjourney 用 –style –ar –no 参数,DALL-E 用自然语言。RAG 检索出的品牌信息是中间表示,由 Builder 按目标模型翻译成对应格式。
- 三种审核维度:CLIP 风格相似度、色彩偏移、Vision API 内容审核。前两者适合做第一轮粗筛,Vision API 做终审。只有全部通过才标记为合格。
- 反馈闭环留好了接口:QualityChecker.check 返回的结构化审核结果包含评分,高分的可以回写知识库。这个逻辑在设计上预留了,实际可以选"审核通过率 > 80% 的 prompt 自动入库"。
四、边界分析与架构权衡
4.1 RAG 检索 vs 直接参考图
有能力直接用参考图的话(如 Midjourney 的 –cref、Stable Diffusion 的 IP-Adapter),为什么要用 RAG?
RAG 的优势:精确控制。参考图只能传达"大概长这样",但 #FF6B35 的精确色值、LOGO 距边缘 40px 的精确位置,这些是文字才能传达的。RAG 的注入的是精确的规范描述,不是模糊的视觉风格。
两阶段策略:先用 RAG 检索规范信息构建精确 Prompt,生成初版图。再用参考图风格迁移做风格微调。两者不冲突。
4.2 生成成本与效率
DALL-E 3 生成一张图片约 $0.08-0.12,Stable Diffusion 本地跑一张约 $0.01-0.03。加上 RAG 检索的 Embedding 和 Vision API 审核的成本,一张图的完整工作流约:
- DALL-E 3 路径:$0.12 生成 + $0.005 检索 + $0.01 审核 = $0.135/张
- SD 本地路径:$0.02 生成 + $0.005 检索 = $0.025/张
如果每天生成 100 张图,DALL-E 的日成本是 $13.5。对于商业用途来说完全可接受。
4.3 Prompt 知识库的"冷启动"
品牌知识库初期是空的,怎么建?
- 从现有设计稿反推:解析 PSD/Sketch/Figma 文件,提取色彩、字体、间距数据。
- 从品牌 VI 手册结构化提取:导出色板、字体列表、间距规则。
- 从历史高分 Prompt 沉淀:人工标注哪些 prompt 效果好,定期入库。
4.4 自动化程度的选择
完全自动化(RAG + 自动生成 + 自动审核)适合批量素材(ICON、配图、信息图)。关键视觉素材(首页 Hero Banner、品牌广告)需要人工审核节点。
建议流程:RAG → 自动生成 → 自动初筛 → 人工二选一 → 合格入库。平衡效率和质量。
五、总结
把 RAG 引入图像生成工作流,本质上是把分散的品牌规范变成可检索、可注入的结构化知识。
三个关键经验:
图像 RAG 是文本 RAG 的延伸应用。当你把品牌知识变成一个可检索的向量库时,不只是图像生成——视频生成、广告文案、产品描述都可以从这个知识库里检索信息。投资的是一套知识库,收获的是全链路 AI 内容生成的精度提升。
下一篇预告:跨模态检索——文本搜图、图搜文本、图搜图的统一架构怎么设计?




