基于Qwen的SR——ODTSR

虽然比较heavey,有20B参数,几乎使用了T2I的全部框架,但是效果很好。

ODTSR

使用Qwen-Image来做ISR。o指one step,D和T指diffusion和transformer。另外ODTSR还是Controllable的,支持双语bilingual prompt,而且Fidelity Weight 也是可调的参数。即便没有在特定数据集上训练,在real-world scene text image super-resolution (STISR) 上也能有很好的表现。

  • 混合噪声视觉流(NVS, Noise-hybrid Visual Stream)设计
    引入了一个全新的视觉流来接收带有可调噪声(Control Noise) 的低质量图像(LQ),而原有的视觉流则接收带有**一致噪声(Prior Noise)**的低质量图像。这种双管齐下的设计有效融合了保真度与控制力。
  • 保真度感知对抗训练(FAA, Fidelity-aware Adversarial Training)
    ODTSR 进一步采用了 FAA 机制,在增强模型可控性的同时,成功实现了单步推理(One-step inference),大幅提升了效率。

Flow matching模型,对时刻t的intermediate latent variable进行建模:

x1是Gaussian noise,x0是真实分布,vt是t时刻对应的velocity。通过最小化MSE,模型就可以预测任意t的velocity:

QwenImagePipeline

下载代码后,还需要下载模型文件Qwen-image,Qwen-Image放在path2model/Qwen-Image中,使用ODTSR-main/examples/qwen_image/test_gan.sh 推理,会根据export qwen_path去读取文件。Generator就会用这些去初始化得到model:

复制代码
pretrained_qwen_path = os.environ["qwen_path"]

    sd_safe_tensor_path_json_format = f'''[
        [
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00001-of-00009.safetensors",
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00002-of-00009.safetensors",
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00003-of-00009.safetensors",
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00004-of-00009.safetensors",
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00005-of-00009.safetensors",
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00006-of-00009.safetensors",
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00007-of-00009.safetensors",
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00008-of-00009.safetensors",
            "{pretrained_qwen_path}/transformer/diffusion_pytorch_model-00009-of-00009.safetensors"
        ],
        [
            "{pretrained_qwen_path}/text_encoder/model-00001-of-00004.safetensors",
            "{pretrained_qwen_path}/text_encoder/model-00002-of-00004.safetensors",
            "{pretrained_qwen_path}/text_encoder/model-00003-of-00004.safetensors",
            "{pretrained_qwen_path}/text_encoder/model-00004-of-00004.safetensors"
        ],
        "{pretrained_qwen_path}/vae/diffusion_pytorch_model.safetensors"
    ]'''

    model = Generator(
        torch_dtype = torch.bfloat16,
        pretrained_weights=sd_safe_tensor_path_json_format,
        tokenizer_path = f"{pretrained_qwen_path}/tokenizer",
        learning_rate=0,
        use_gradient_checkpointing=False,
        pretrained_ckpt_path_gen = args.trained_ckpt
    )

其中args.trained_ckpt指的是trained ODTSR model weight: huggingface

Generator继承自BaseModelForT2ILoRA,核心pipe依赖于QwenImagePipeline,由上面的几个模组进行初始化。

model_configs对应sd_safe_tensor_path_json_format,包括了transformer,text_encoder,vae。而把tokenizer_path作为单独的变量传递过去,这是因为tokenizer_path只是分词器,严格意义上不算网络的一部分:

python 复制代码
if tokenizer_path is not None:
     self.pipe = QwenImagePipeline.from_pretrained(torch_dtype=torch.bfloat16, device="cpu", model_configs=model_configs, tokenizer_config=ModelConfig(tokenizer_path))
else:
     self.pipe = QwenImagePipeline.from_pretrained(torch_dtype=torch.bfloat16, device="cpu", model_configs=model_configs)

这几个模组的作用如下:

|--------------|-----------------------------------------------------------------------------------------|
| | |
| tokenizer | 把 prompt 字符串切成 token id,供 text_encoder |
| text_encoder | 负责把 prompt 编成 prompt_embprompt_emb_mask |
| vae | 负责 RGB 图像和 latent 之间转换, image -> vae.encode -> latent latent -> vae.decode -> image |
| dit | 核心生成网络,也就是 Diffusion Transformer / DiT。 |

基于这几个模块,QwenImagePipeline还有其他成员:

|-------------|-----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| | |
| scheduler | 扩散/Flow Matching 的噪声调度器 。它负责生成 timestep / sigma,以及执行加噪add_noise、去噪 |
| unit_runner | 预处理流水线执行器。它会按顺序跑下面的 units,把普通输入变成模型需要的 tensor / latent / embedding。 |
| units | QwenImageUnit_ShapeChecker(),检查/修正 height, width,让尺寸符合模型要求 QwenImageUnit_NoiseInitializer(),生成 noise,形状通常是 1, 16, H/8, W/8 QwenImageUnit_InputImageEmbedder(),-> preprocess_image -> vae.encode -> 得到 condition_latents, condition_rgb, input_latents 等 QwenImageUnit_PromptEmbedder(),得到prompt_emb, prompt_emb_mask |
| model_fn | QwenImagePipeline的推理函数,DiT 前向的包装函数 |

two stream

variational auto-encoder (VAE)是连接数据与隐式空间latent space z的桥梁,由encoder(E) and decoder(D) 构成。关键的DIT就在latent space中生效。

标准T2I是visual stream+Text stream,ODTSR使用了两个visual stream。一个是control noise,通过loRA进行微调,另外一个prior noise被冻结。

为什么要使用两个visual stream?简单回答就是为了平衡Generative和Fidelity, 两路结构让模型一边"生成",一边持续"看原图"**。**模型在t时刻的预测受限于noised latent x_t和text prompt c。下图中通过使用不同的t,可以看出噪声强度对结果的影响。从a和b图可以得到结论,low-noise时和原图的一致性consistency比high noise的结果更好,虽然文字这种结构性强的会有损失,但是可以通过补充text prompt弥补。从b和c比较,当输入改成LQ,high noise时的效果又更好,因为此时图像画质很差,high noise相当于发挥空间更大,更多地依赖文本提示词(Prompt)和模型自身的先验知识去"脑补"和重构细节。

为了更好的利用预训练模型的这种特性,ODTSR的做法是把一个控制噪声的t变成两个,Prior Noise和Control Noise,分别负责提升和保真:

t决定了噪声的强度,这里可以看到两个分支的t是明显不同的,并且条件分支的t和f挂钩,f越大,t越低,相当于噪声也更少。使用(1-f), 控制了条件分支在原始LQ和生成分支间线性过渡。

| 项目 | 生成 visual stream | LQ 条件 visual stream |
| 名字 | Prior Noise stream(先验噪声流) | Control Noise stream(控制噪声流) |

作用 利用模型的先验知识来"脑补"细节,从而提升画面的感知质量(Perceptual quality) 牢牢锁定原图的特征,确保生成结果不偏离原图,从而保证保真度(Fidelity)
原始图像 LQ LQ
编码器 可训练的 new_vae.encoder 冻结的原始 vae.encoder
初始 latent
加噪索引 固定为 750,非线性 exponential shift之后对应0.43 训练时随机取 \([750,1000)\)
作用 被模型恢复、产生最终输出 向生成流提供结构和内容条件
最终是否输出 否,最后被裁掉

条件分支因为只需要针对原始的LQ,所以也使用原始的VAE进行encoder得到lq_latents,对应 ODTSR 里的 Control Noise 那一路

而生成分支因为需要更大的生成能力,所以对vae的encoder进行了微调。new_vae从pipe.vae中deepcopy得到,并只解冻了它的 encoder.conv_in 层:

python 复制代码
# copy a new vae
self.pipe.new_vae = deepcopy(self.pipe.vae)
self.unfrozen(self.pipe.new_vae.encoder, type(self.pipe.new_vae.encoder.conv_in))

经过了训练。训练时候通过loss约束:

复制代码
new_lq_latents_rgb = generator.module.pipe.vae.decode(new_lq_latents)
loss_new_vae_lq = mse(new_lq_latents_rgb, gt_rgb)

Generator是QwenImagePipeline的上一级。noisy_latents和lq_latents是输入,为了兼容这样的输入,需要把 Qwen DiT 里指定的一批 Linear 层替换成"双 LoRA"版本,支持 ODTSR 的双 visual stream:

python 复制代码
# 结构修改 & fp8降低显存
lora_base_model = 'dit' # hard core
lora_rank = 128
self.add_custom_dual_lora(
     getattr(self.pipe, lora_base_model),
     lora_rank=lora_rank)

def add_custom_dual_lora(self, model, lora_rank):
     patterns = [
             "img_in",
            "img_mod.1",
            "attn.to_q",
            "attn.to_k",
            "attn.to_v",
            "to_out.0",
            "img_mlp.net.0.proj",
            "img_mlp.net.2",
        ]
      replace_linear_with_duallora(model, patterns, rank=lora_rank, alpha1=0, alpha2=lora_rank, use_fp8 = True)

lora是一种低秩分解(Low-Rank Factorization)的数学思想,不替代原始权重,而是提供增量,原始权重则被冻结。一个lora由两个矩阵构成,两个矩阵的乘积作为权重更新的增量:

DualLoRALinear 里面有两套 LoRA:

python 复制代码
lora_A1 / lora_B1
lora_A2 / lora_B2

两个lora分别有scaling1和scaling2,用来对增量delta进行加权。看论文的fig 3,control noise分支有lora,旁边画了一把火。

|-----------|--------------------------------------------------------------------|--------------------------------------------------------------------------------|
| | LoRA 1 (alpha1=0) | LoRA 2 (alpha2=lora_rank) |
| 特征 | 缩放因子 alpha 为 0。这意味着这个 LoRA 的权重更新被完全屏蔽了,它实际上不起任何作用(或者作为一个占位符/直通通道)。 | 缩放因子 alpha 等于 rank(即全量激活)。这意味着这个 LoRA 会全力工作,极大地改变原始模型的权重分布。 |
| 对应 Stream | 这对应 Prior Noise stream(先验噪声流)。Prior stream 冻结是为了守住 T2I 去噪先验 | 这通常对应 Control Noise stream(控制噪声流)。Control stream 加 LoRA 是为了让模型学会读取可变噪声的 LQ 条件。 |

DIT

DIT的输入有三路:Prior visual + Control visual + Text。

两路visual虽然因为有lora的差异,但是两路 visual 的特征尺寸完全相同,所以还是可以合并并行计算。比如,共用一套QKV投影和位置编码。这样可以最大程度复用T2I的结构。

然后把img和text的QKV拼接再计算QKV:

复制代码
joint_q = torch.cat([txt_q, img_q], dim=2)
joint_k = torch.cat([txt_k, img_k], dim=2)
joint_v = torch.cat([txt_v, img_v], dim=2)

img_q 内部已经是 [Prior, Control]。因此注意力矩阵实际上可以看作 3x3的交互:

注意力之后,再分别经过 visual MLP 和 text MLP。

最终,Control token 被丢弃,只把 Prior token送入输出层,最后经过 VAE decoder 得到 SR 图像。

predict

速度场(Velocity)由 self.pipe.model_fn预测得到。self.pipe.model_fn是扩散模型在去噪(Denoising)过程中的核心前向传播函数(Forward Function) 。本质上是一个封装好的函数引用 ,它指向底层的 Diffusion Transformer (DiT) 模型。这里的DIT还是MMDIT,输入是被拼接在一起处理的:

python 复制代码
def forward(self, noisy_latents, condition_latent, timestep, prompt_emb, prompt_emb_mask):
        b,c,h,w = noisy_latents.shape
        out = self.pipe.model_fn(self.pipe.dit, 
                                 noisy_latents,
                                 condition_latent,
                                 timestep,
                                 prompt_emb,
                                 prompt_emb_mask,
                                 h*8,
                                 w*8,
                                 use_gradient_checkpointing=True
                                 )
        return out

如果cfg_scale!=1.0,还会根据负提示词送入self.pipe.model_fn,按照CFG(Classifier-Free Guidance)把正负提示词得到的结果的diff进行加权:

python 复制代码
noise_pred = noise_pred_nega + cfg_scale * (noise_pred_posi - noise_pred_nega)

最终得到的noise_pred不是最终图像 latent,而是从当前 noisy latent 往干净 latent 走的"方向/速度"。所以还需要Flow Matching 的一步更新:

python 复制代码
# one step prediction
training_pred = noisy_latents + (0 - one_step_sigma) * noise_pred

然后decoder就得到最终的图像:

python 复制代码
# Decode
image = self.pipe.vae.decode(training_pred, device=self.device, tiled=tiled, tile_size=tile_size, tile_stride=tile_stride)
image = self.pipe.vae_output_to_image(image)

需要 40GB GPU memory,不过可以利用QwenImagePipeline的enable_vram_management,灵活地把需要的VAE或者DIT搬到GPU上,而不是一下子全部搬到GPU上。

loss

重建损失肯定是必要的,通过计算预测图和GT的MSE和LPIPS加权得到:

使用了相对GAN损失,优化生成器:

最终的loss大小还会根据fidelity的值调整,这就是**(FAA, Fidelity-aware Adversarial Training)。**输入画质高时,fidelity可以更高,f更高,adv loss可以更低,避免引入artifacts。

  • 在低 fidelity 下,判别器允许生成结果与原图有较大差异,只要细节逼真即可。
  • 在高 fidelity 下,判别器会严厉惩罚那些偏离原图结构的生成结果。

Metrics

|---------------------|-------------------|
| full-reference (FR) | no-reference (NR) |
| PSNR | MUSIQ |
| SSIM | MANIQA |
| LPIPS | |
| DISTS | |
| NED text similarity | |

相关推荐
小马9261 小时前
从单模型到多模型编排:GitHub HydraFusion 如何让编程 Agent 降本 36%-67%
人工智能·github
“AI国潮设计-小江”1 小时前
【Python实战】SDXL精准控制“普宁英歌舞×星空蛋糕”IP落地,附核心Prompt与商用授权思路
开发语言·人工智能·python·prompt·aigc
HySpark1 小时前
从“能识别”到“稳定识别”:离线ASR在真实会议场景中的问题与工程优化实践
人工智能·语音识别
weixin_446260851 小时前
CABAL:用于追踪同行评审中合谋投标影响的多智能体仿真框架
人工智能·算法·机器学习
万象新讯1 小时前
数据中心运维管理软件平台,有哪些合适的产品可以选择?
人工智能
广州灵眸科技有限公司1 小时前
瑞芯微(EASY EAI)RV1126B 星闪使用
运维·人工智能·科技·docker·容器
昇腾知识体系1 小时前
msprobe/msdebug 全家桶:昇腾精度比对、溢出检测、msSanitizer 内存检测与 msOpProf 算子调优实战
人工智能·华为·知识图谱
醍醐实验室1 小时前
分布式通信算子剖析:All-Reduce、All-Gather 与 Reduce-Scatter 底层算法
人工智能·all-reduce
愚公搬代码1 小时前
【愚公系列】《造浪者:AI创业实战地图》001-回望来路:四次技术浪潮的创业逻辑
人工智能