From 90ccb43cab3d544f25fbe64860f93dcd40c9ceb3 Mon Sep 17 00:00:00 2001 From: bubbliiiing <3323290568@qq.com> Date: Thu, 24 Sep 2026 10:37:19 +0800 Subject: [PATCH] Update qwen image 21 control --- .../qwenimage21_fun/predict_i2i_inpaint.py | 51 +++++++++++-------- .../qwenimage21_fun/predict_t2i_control.py | 18 ++++--- .../pipeline/pipeline_qwenimage21_control.py | 38 ++++++++++---- 3 files changed, 68 insertions(+), 39 deletions(-) diff --git a/examples/qwenimage21_fun/predict_i2i_inpaint.py b/examples/qwenimage21_fun/predict_i2i_inpaint.py index d9eb3cf..2752d98 100644 --- a/examples/qwenimage21_fun/predict_i2i_inpaint.py +++ b/examples/qwenimage21_fun/predict_i2i_inpaint.py @@ -5,7 +5,6 @@ import torch from diffusers import FlowMatchEulerDiscreteScheduler from omegaconf import OmegaConf -from PIL import Image current_file_path = os.path.abspath(__file__) project_roots = [os.path.dirname(current_file_path), os.path.dirname(os.path.dirname(current_file_path)), os.path.dirname(os.path.dirname(os.path.dirname(current_file_path)))] @@ -23,6 +22,7 @@ from videox_fun.utils import (register_auto_device_hook, from videox_fun.utils.fp8_optimization import (convert_model_weight_to_float8, convert_weight_dtype_wrapper) from videox_fun.utils.lora_utils import merge_lora, unmerge_lora +from videox_fun.utils.utils import get_image_latent # GPU memory mode, which can be chosen in [model_full_load, model_full_load_and_qfloat8, model_cpu_offload, model_cpu_offload_and_qfloat8, model_group_offload, sequential_cpu_offload]. # model_full_load means that the entire model will be moved to the GPU. @@ -73,13 +73,9 @@ use_kv_cache = True # Use torch.float16 if GPU does not support torch.bfloat16 # Some graphics cards, such as v100, 2080ti, do not support torch.bfloat16 weight_dtype = torch.bfloat16 -# Inpaint source image, and its mask. In the mask, WHITE (>= 0.5) marks the region to REGENERATE and BLACK -# marks the region to KEEP -- matching the mask convention used during training and the *_mask assets in asset/. -image_path = "asset/pose.jpg" -mask_path = "asset/mask.png" -# Optional edge/pose control map layered on top of the inpaint branch. None -> its 64 channels are zeroed, -# which is exactly the trained "inpaint only" regime. Set a path to run control + inpaint together. -control_image_path = None +control_image = "asset/pose.jpg" +inpaint_image = "asset/8.png" +mask_image = "asset/mask.png" # Strength of the control/union branch. 1.0 is the value the adapter is trained to consume. control_context_scale = 1.0 @@ -172,10 +168,8 @@ if ulysses_degree > 1 or ring_degree > 1: pipeline.transformer = shard_fn(pipeline.transformer) print("Add FSDP DIT") if fsdp_text_encoder: - from functools import partial - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.model.language_model.layers) - text_encoder = shard_fn(text_encoder) + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=pipeline.text_encoder.model.language_model.layers) + pipeline.text_encoder = shard_fn(pipeline.text_encoder) print("Add FSDP TEXT ENCODER") if compile_dit: @@ -206,14 +200,27 @@ generator = torch.Generator(device=device).manual_seed(seed) if lora_path is not None: pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype) -# Load the source image and mask as PIL images; the pipeline's image_processor / mask_processor resize them to -# (height, width) and build control_context = [control_latents(64) | mask(1) | masked-image latents(64)] = 129 ch. -inpaint_image = Image.open(image_path).convert("RGB") -mask_image = Image.open(mask_path) -if control_image_path is not None: - control_image = Image.open(control_image_path).convert("RGB") +# Load the conditions through get_image_latent -- the same single-frame tensors that +# scripts/qwenimage21_fun/train_control.py validation builds -- so the resize / normalization the model was +# trained with is reproduced here; the pipeline only preprocesses them further and assembles +# control_context = [control_latents(64) | mask(1) | masked-image latents(64)] = 129 ch. +if inpaint_image is not None: + inpaint_image_input = get_image_latent(inpaint_image, sample_size=sample_size)[:, :, 0] else: - control_image = None + inpaint_image_input = torch.zeros([1, 3, sample_size[0], sample_size[1]]) + +# In the mask, WHITE (>= 0.5) marks the region to REGENERATE and BLACK the region to KEEP -- the convention used +# during training. get_image_latent opens it through convert("RGB"), which also resolves palette (mode "P") PNGs +# by their rendered grey value instead of the raw palette index (asset/mask.png: 54.6% vs 0.29% regenerate area). +if mask_image is not None: + mask_image_input = get_image_latent(mask_image, sample_size=sample_size)[:, :1, 0] +else: + mask_image_input = torch.ones([1, 1, sample_size[0], sample_size[1]]) * 255 + +if control_image is not None: + control_image_input = get_image_latent(control_image, sample_size=sample_size)[:, :, 0] +else: + control_image_input = None with torch.no_grad(): sample = pipeline( @@ -224,9 +231,9 @@ with torch.no_grad(): generator = generator, true_cfg_scale = guidance_scale, num_inference_steps = num_inference_steps, - image = inpaint_image, - mask_image = mask_image, - control_image = control_image, + image = inpaint_image_input, + mask_image = mask_image_input, + control_image = control_image_input, control_context_scale = control_context_scale, use_kv_cache = use_kv_cache, ).images diff --git a/examples/qwenimage21_fun/predict_t2i_control.py b/examples/qwenimage21_fun/predict_t2i_control.py index f272290..0023270 100644 --- a/examples/qwenimage21_fun/predict_t2i_control.py +++ b/examples/qwenimage21_fun/predict_t2i_control.py @@ -165,10 +165,8 @@ if ulysses_degree > 1 or ring_degree > 1: pipeline.transformer = shard_fn(pipeline.transformer) print("Add FSDP DIT") if fsdp_text_encoder: - from functools import partial - from videox_fun.dist import set_multi_gpus_devices, shard_model - shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=text_encoder.model.language_model.layers) - text_encoder = shard_fn(text_encoder) + shard_fn = partial(shard_model, device_id=device, param_dtype=weight_dtype, module_to_wrapper=pipeline.text_encoder.model.language_model.layers) + pipeline.text_encoder = shard_fn(pipeline.text_encoder) print("Add FSDP TEXT ENCODER") if compile_dit: @@ -199,9 +197,13 @@ generator = torch.Generator(device=device).manual_seed(seed) if lora_path is not None: pipeline = merge_lora(pipeline, lora_path, lora_weight, device=device, dtype=weight_dtype) -# Load the control image as a single-frame (1, 3, h, w) tensor, matching scripts/qwenimage21_fun/train_control.py -# validation (get_image_latent(... )[:, :, 0]) so inference preprocessing is identical to training. -control_image = get_image_latent(control_image, sample_size=(sample_size[0], sample_size[1]))[:, :, 0] +# Load the control image through get_image_latent -- the same single-frame (1, 3, h, w) tensor that +# scripts/qwenimage21_fun/train_control.py validation builds, so the resize / normalization the model was trained +# with is reproduced here instead of relying on the pipeline's own PIL preprocessing. +if control_image is not None: + control_image_input = get_image_latent(control_image, sample_size=sample_size)[:, :, 0] +else: + control_image_input = None with torch.no_grad(): sample = pipeline( @@ -212,7 +214,7 @@ with torch.no_grad(): generator = generator, true_cfg_scale = guidance_scale, num_inference_steps = num_inference_steps, - control_image = control_image, + control_image = control_image_input, control_context_scale = control_context_scale, use_kv_cache = use_kv_cache, ).images diff --git a/videox_fun/pipeline/pipeline_qwenimage21_control.py b/videox_fun/pipeline/pipeline_qwenimage21_control.py index a0c14ef..b9128c8 100644 --- a/videox_fun/pipeline/pipeline_qwenimage21_control.py +++ b/videox_fun/pipeline/pipeline_qwenimage21_control.py @@ -502,8 +502,19 @@ class QwenImage21ControlPipeline(DiffusionPipeline): prompt (`str` or `list[str]`, *optional*): The prompt to guide image generation. Pass `prompt_embeds` instead to supply embeddings directly. image (`PipelineImageInput`, *optional*): - One or more condition images, as a PIL image or a numpy array. A list is one set of images shared by - every prompt in the batch, not one entry per prompt. + The inpaint source, as a PIL image, a numpy array or a (1, 3, h, w) tensor in [0, 1] -- the form + `videox_fun.utils.utils.get_image_latent` returns and the training validations use. Unlike in the base + pipeline this is never shown to the text encoder; it is masked and VAE-encoded into `control_context` + (step 3b), so a preprocessed tensor is accepted here. A list is one set of images shared by every + prompt in the batch, not one entry per prompt. + mask_image (`PipelineImageInput`, *optional*): + Binary mask over `image`, in the same pixel forms. WHITE (>= 0.5) marks the region to regenerate and + black the region to keep; omitted, the whole image is regenerated. + control_image (`PipelineImageInput`, *optional*): + The control signal (pose / depth / canny / ...), in the same pixel forms. Omitted, the control + channels of `control_context` are zeros -- the pure-inpaint recipe from training. + control_context_scale (`float`, *optional*, defaults to 1.0): + Multiplier on the assembled `control_context` conditioning. negative_prompt (`str` or `list[str]`, *optional*): The prompt not to guide image generation. Ignored when `true_cfg_scale` is not greater than 1. true_cfg_scale (`float`, *optional*, defaults to 1.0): @@ -555,8 +566,10 @@ class QwenImage21ControlPipeline(DiffusionPipeline): with the generated images. """ if image is not None: - # The text encoder reads each condition image as vision context, so the pixels have to be there. Normalize - # to PIL up front, and everything downstream — the aspect ratio below, the resize, the VAE — sees one type. + # `image` is only the inpaint source here: step 1 leaves the condition images that reach the text encoder + # empty, and step 3b consumes `image` through the VAE image processor, which also takes tensors. So the + # base pipeline's PIL-only requirement (its condition images really do go into the vision slots) does not + # apply -- normalize to PIL / tensor and let everything downstream see one of those two types. image = image if isinstance(image, list) else [image] condition_images = [] for img in image: @@ -564,6 +577,8 @@ class QwenImage21ControlPipeline(DiffusionPipeline): condition_images.append(img) elif isinstance(img, np.ndarray): condition_images.append(PILImage.fromarray(img)) + elif isinstance(img, torch.Tensor): + condition_images.append(img) elif isinstance(img, (list, tuple)): raise ValueError( "`image` is one flat set of condition images that applies to every prompt in the batch, so it " @@ -572,13 +587,18 @@ class QwenImage21ControlPipeline(DiffusionPipeline): ) else: raise ValueError( - f"`image` accepts a PIL image or a numpy array, or a list of either, but got " - f"{type(img).__name__}. Latents cannot stand in for a condition image here, because the text " - f"encoder has to see the image itself." + "`image` accepts a PIL image, a numpy array or a (1, 3, h, w) tensor, or a list of either, but " + f"got {type(img).__name__}." ) image = condition_images + # Aspect ratio of the last condition image: PIL / numpy carry (w, h) in `.size`, a tensor carries + # (..., h, w) in its trailing dims. + if isinstance(image[-1], torch.Tensor): + image_width, image_height = image[-1].shape[-1], image[-1].shape[-2] + else: + image_width, image_height = image[-1].size calculated_width, calculated_height, _ = calculate_dimensions( - output_resolution * output_resolution, image[-1].size[0] / image[-1].size[1] + output_resolution * output_resolution, image_width / image_height ) height = height or calculated_height width = width or calculated_width @@ -677,7 +697,7 @@ class QwenImage21ControlPipeline(DiffusionPipeline): mask_condition = _to_ctrl_batch(mask_condition).to(device) if image is not None: - inpaint_image = self.image_processor.preprocess(image, height=height, width=width) + inpaint_image = self.image_processor.preprocess(image, height=height, width=width).to(device) inpaint_image = inpaint_image * (mask_condition < 0.5) # zero out the region to regenerate inpaint_latent = self._encode_vae_image( _rgba(inpaint_image).unsqueeze(2).to(device=device, dtype=weight_dtype), generator