尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

基于Qwen的SR——ODTSR

基于Qwen的SR——ODTSR 虽然比较heavey有20B参数几乎使用了T2I的全部框架但是效果很好。ODTSR使用Qwen-Image来做ISR。o指one stepD和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 TrainingODTSR 进一步采用了 FAA 机制在增强模型可控性的同时成功实现了单步推理One-step inference大幅提升了效率。Flow matching模型对时刻t的intermediate latent variable进行建模x1是Gaussian noisex0是真实分布vt是t时刻对应的velocity。通过最小化MSE模型就可以预测任意t的velocityQwenImagePipeline下载代码后还需要下载模型文件Qwen-imageQwen-Image放在path2model/Qwen-Image中使用ODTSR-main/examples/qwen_image/test_gan.sh 推理会根据export qwen_path去读取文件。Generator就会用这些去初始化得到modelpretrained_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_weightssd_safe_tensor_path_json_format, tokenizer_path f{pretrained_qwen_path}/tokenizer, learning_rate0, use_gradient_checkpointingFalse, pretrained_ckpt_path_gen args.trained_ckpt )其中args.trained_ckpt指的是trained ODTSR model weight: huggingfaceGenerator继承自BaseModelForT2ILoRA核心pipe依赖于QwenImagePipeline由上面的几个模组进行初始化。model_configs对应sd_safe_tensor_path_json_format包括了transformertext_encodervae。而把tokenizer_path作为单独的变量传递过去这是因为tokenizer_path只是分词器严格意义上不算网络的一部分if tokenizer_path is not None: self.pipe QwenImagePipeline.from_pretrained(torch_dtypetorch.bfloat16, devicecpu, model_configsmodel_configs, tokenizer_configModelConfig(tokenizer_path)) else: self.pipe QwenImagePipeline.from_pretrained(torch_dtypetorch.bfloat16, devicecpu, model_configsmodel_configs)这几个模组的作用如下tokenizer把 prompt 字符串切成 token id供text_encodertext_encoder负责把 prompt 编成prompt_emb和prompt_emb_maskvae负责 RGB 图像和 latent 之间转换image - vae.encode - latentlatent - vae.decode - imagedit核心生成网络也就是Diffusion Transformer / DiT。基于这几个模块QwenImagePipeline还有其他成员scheduler扩散/Flow Matching 的噪声调度器。它负责生成 timestep / sigma以及执行加噪add_noise、去噪unit_runner预处理流水线执行器。它会按顺序跑下面的units把普通输入变成模型需要的 tensor / latent / embedding。unitsQwenImageUnit_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_maskmodel_fnQwenImagePipeline的推理函数DiT 前向的包装函数two streamvariational auto-encoder (VAE)是连接数据与隐式空间latent space z的桥梁由encoder(E) and decoder(D) 构成。关键的DIT就在latent space中生效。标准T2I是visual streamText streamODTSR使用了两个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比较当输入改成LQhigh noise时的效果又更好因为此时图像画质很差high noise相当于发挥空间更大更多地依赖文本提示词Prompt和模型自身的先验知识去“脑补”和重构细节。为了更好的利用预训练模型的这种特性ODTSR的做法是把一个控制噪声的t变成两个Prior Noise和Control Noise分别负责提升和保真t决定了噪声的强度这里可以看到两个分支的t是明显不同的并且条件分支的t和f挂钩f越大t越低相当于噪声也更少。使用(1-f), 控制了条件分支在原始LQ和生成分支间线性过渡。项目生成 visual streamLQ 条件 visual stream名字Prior Noise stream先验噪声流Control Noise stream控制噪声流作用利用模型的先验知识来“脑补”细节从而提升画面的感知质量Perceptual quality。牢牢锁定原图的特征确保生成结果不偏离原图从而保证保真度Fidelity。原始图像LQLQ编码器可训练的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层# 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# 结构修改 fp8降低显存 lora_base_model dit # hard core lora_rank 128 self.add_custom_dual_lora( getattr(self.pipe, lora_base_model), lora_ranklora_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, ranklora_rank, alpha10, alpha2lora_rank, use_fp8 True)lora是一种低秩分解Low-Rank Factorization的数学思想不替代原始权重而是提供增量原始权重则被冻结。一个lora由两个矩阵构成两个矩阵的乘积作为权重更新的增量DualLoRALinear里面有两套 LoRAlora_A1 / lora_B1 lora_A2 / lora_B2两个lora分别有scaling1和scaling2用来对增量delta进行加权。看论文的fig 3control noise分支有lora旁边画了一把火。LoRA 1 (alpha10)LoRA 2 (alpha2lora_rank)特征缩放因子 alpha 为 0。这意味着这个 LoRA 的权重更新被完全屏蔽了它实际上不起任何作用或者作为一个占位符/直通通道。缩放因子 alpha 等于 rank即全量激活。这意味着这个 LoRA 会全力工作极大地改变原始模型的权重分布。对应 Stream这对应 Prior Noise stream先验噪声流。Prior stream 冻结是为了守住 T2I 去噪先验这通常对应 Control Noise stream控制噪声流。Control stream 加 LoRA 是为了让模型学会读取可变噪声的 LQ 条件。DITDIT的输入有三路Prior visual Control visual Text。两路visual虽然因为有lora的差异但是两路 visual 的特征尺寸完全相同所以还是可以合并并行计算。比如共用一套QKV投影和位置编码。这样可以最大程度复用T2I的结构。然后把img和text的QKV拼接再计算QKVjoint_q torch.cat([txt_q, img_q], dim2) joint_k torch.cat([txt_k, img_k], dim2) joint_v torch.cat([txt_v, img_v], dim2)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输入是被拼接在一起处理的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_checkpointingTrue ) return out如果cfg_scale!1.0还会根据负提示词送入self.pipe.model_fn按照CFGClassifier-Free Guidance把正负提示词得到的结果的diff进行加权noise_pred noise_pred_nega cfg_scale * (noise_pred_posi - noise_pred_nega)最终得到的noise_pred不是最终图像 latent而是从当前 noisy latent 往干净 latent 走的“方向/速度”。所以还需要Flow Matching 的一步更新# one step prediction training_pred noisy_latents (0 - one_step_sigma) * noise_pred然后decoder就得到最终的图像# Decode image self.pipe.vae.decode(training_pred, deviceself.device, tiledtiled, tile_sizetile_size, tile_stridetile_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下判别器会严厉惩罚那些偏离原图结构的生成结果。Metricsfull-reference (FR)no-reference (NR)PSNRMUSIQSSIMMANIQALPIPSDISTSNED text similarity
返回列表