Update qwen image 21 control
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user