SDXL Prompt2Prompt 相似图对生成
更新时间:2026-09-14
简介
SDXL 交叉注意力控制生成相似图对的算子:给定一对 caption,用仓库内置的 Prompt2PromptPipeline(daft/aihc/functions/image/_prompt2prompt_pipeline.py)以 edit_type="refine" 的注意力替换,一次生成两张仅局部语义不同的图,常用于造 image-diff / 细粒度对比训练数据。默认权重 stabilityai/stable-diffusion-xl-base-1.0。
功能描述
- 输入是两个 prompt 列,没有图像输入:两个 prompt 作为一个 batch 一次采样,共享同一条噪声轨迹,第二张图只在注意力 refine 覆盖的位置上与第一张不同
cross_attention_kwargs固定edit_type="refine",n_self_replace与n_cross_replace由参数控制seed = -1时每行现取random.randint(0, 9999);seed >= 0时全列复用同一 seed,结果可复现- 生成尺寸可用
height/width指定;尺寸非法(非正数或非 8 的倍数)在构造期就报错,因为transform会把行级异常吞成空列表 - 构造期同时校验
num_inference_steps > 0、n_self_replace ∈ [0, 1]、n_cross_replace ∈ [0, 1] - 推理在
torch.no_grad()下执行;输出恒为两张图,不支持一行出多对 - 输出文件名
image_pair_row<行内下标>_<seed>_1.jpg/_2.jpg;output_bosdir非空时上传并返回 BOS 路径 - prompt 对中任一为空串或非字符串时跳过该行、返回空列表并记 warning;行级异常记日志后返回空列表
- 权重只从本地
{model_path}/{model_name}加载,目录不存在时构造期抛FileNotFoundError,不联网下载、不做运行时 pip 安装
算子参数
输入
| 输入 | 含义 |
|---|---|
| caption | 第一张图的 prompt 文本 |
| caption_second | 第二张图的 prompt 文本,通常与第一个 prompt 只差一个局部词 |
输出
| 输出 | 含义 |
|---|---|
| image_pair | list<large_string>:[第一张图路径, 第二张图路径](output_bosdir 非空时为 BOS 路径);prompt 对非法或该行失败时为空列表 |
参数
| 参数名称 | 类型 | 默认值 | 描述 |
|---|---|---|---|
| model_path | str | "/opt/aihc/model" | 权重根目录 |
| model_name | str | "stabilityai/stable-diffusion-xl-base-1.0" | 相对 model_path 的 SDXL 权重子目录 |
| dtype | str | "float16" | 权重精度:float32 / float16 / bfloat16。GPU 上建议 float16 |
| variant | str | "" | 权重文件 variant 后缀;只下了 *.fp16.safetensors 时需传 "fp16",空串表示默认权重 |
| num_inference_steps | int | 50 | 采样步数 |
| guidance_scale | float | 7.5 | CFG 强度 |
| n_self_replace | float | 0.4 | self-attention 替换比例,取值 [0, 1] |
| n_cross_replace | float | 1.0 | cross-attention 替换比例,取值 [0, 1] |
| seed | int | -1 | 随机种子;-1 表示每行随机 |
| height | int 或 None | None | 生成高度,必须是 8 的正整数倍;None 走 SDXL 默认 1024 |
| width | int 或 None | None | 生成宽度,语义同 height |
| output_dir | str | "/tmp/aihc_sdxl_prompt2prompt" | 本地输出目录 |
| output_bosdir | str | "" | 非空时把生成图上传到该 BOS 目录,并返回 BOS 路径 |
| rank | int | 0 | 多卡场景 worker 序号,设备取 cuda:(rank % 可见卡数) |
注意事项
- 显存:prompt2prompt 要显式拿到注意力矩阵,走不了 SDPA。1024×1024 峰值实测 29.6G,24G 卡放不下;降到 768×768 是 14.3G。小卡上务必显式设
height/width。 - 若权重只有 fp16 分片,必须同时传
variant="fp16",否则from_pretrained找不到权重文件。 - 输出文件名里的行内下标是「本次微批内」的下标,多个微批会从 0 重新计数:固定
seed且batch_size < 总行数时不同批之间会互相覆盖,批量场景建议按批设置不同的output_dir。 - 只有
aihc_udf(num_gpus=...)大于 0 时才用 GPU。
调用示例
Python
1from __future__ import annotations
2
3import os
4
5import daft
6from daft import col
7
8from daft.aihc.common.udf import aihc_udf
9from daft.aihc.functions.image.sdxl_prompt2prompt import SdxlPrompt2Prompt
10
11if __name__ == "__main__":
12 if os.getenv("DAFT_RUNNER", "native") == "ray":
13 import ray
14 ray.init(dashboard_host="0.0.0.0", ignore_reinit_error=True)
15 daft.set_runner_ray()
16 daft.set_execution_config(actor_udf_ready_timeout=6000, min_cpu_per_task=0)
17
18samples = {
19 "caption": ["a photo of a squirrel eating a burger"],
20 "caption_second": ["a photo of a squirrel eating a confetti burger"],
21 }
22 ds = daft.from_pydict(samples)
23 ds = ds.with_column(
24 "image_pair",
25 aihc_udf(
26 SdxlPrompt2Prompt,
27 construct_args={
28 "model_path": "/path/to/models",
29 "model_name": "stabilityai/stable-diffusion-xl-base-1.0",
30 "dtype": "float16",
31 "variant": "fp16",
32 "num_inference_steps": 20,
33 "seed": 42,
34 "height": 768,
35 "width": 768,
36 "output_dir": "/tmp/aihc_sdxl_prompt2prompt",
37 },
38 num_cpus=1,
39 num_gpus=1,
40 concurrency=1,
41 batch_size=1,
42 )(col("caption"), col("caption_second")),
43 )
44 ds.show()
评价此篇文章
