From 82c355098e33c9dec4717f2f27ccdff5108d4af6 Mon Sep 17 00:00:00 2001 From: Super User Date: Sat, 8 Aug 2026 15:41:31 +0800 Subject: [PATCH 1/4] feat(qwen-image): optimize distributed offload Run text encoding on rank 0 and broadcast packed embeddings. Keep the text encoder and selected DiT blocks resident, stream remaining blocks with events, and support rank-aware model offload. --- .../qwen25_vlforconditionalgeneration.py | 27 +- .../infer/offload/transformer_infer.py | 83 +++++- lightx2v/models/networks/qwen_image/model.py | 117 +++++++- .../qwen_image/weights/transformer_weights.py | 103 +++++-- .../runners/qwen_image/qwen_image_runner.py | 254 +++++++++++++++--- 5 files changed, 523 insertions(+), 61 deletions(-) diff --git a/lightx2v/models/input_encoders/hf/qwen25/qwen25_vlforconditionalgeneration.py b/lightx2v/models/input_encoders/hf/qwen25/qwen25_vlforconditionalgeneration.py index eb95f7fab..d2f6726a2 100755 --- a/lightx2v/models/input_encoders/hf/qwen25/qwen25_vlforconditionalgeneration.py +++ b/lightx2v/models/input_encoders/hf/qwen25/qwen25_vlforconditionalgeneration.py @@ -73,6 +73,7 @@ def __init__(self, config): self.cpu_offload = config.get("qwen25vl_cpu_offload", config.get("cpu_offload", False)) self.dtype = torch.bfloat16 self.load() + self._is_on_device = not self.cpu_offload def load(self): if self.config.get("qwen25vl_quantized", False): @@ -152,11 +153,24 @@ def get_image_caption(self, prompt_image): output_text = self.vl_processor.batch_decode(generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0] return output_text.strip() - @torch.no_grad() - def infer(self, text, image_list=None): + def load_to_device(self): if self.cpu_offload: - if not hasattr(self, "device_map") or self.device_map == AI_DEVICE: + if (not hasattr(self, "device_map") or self.device_map == AI_DEVICE) and not self._is_on_device: self.text_encoder.to(AI_DEVICE) + self._is_on_device = True + + def offload_to_cpu(self): + if self.cpu_offload: + if (not hasattr(self, "device_map") or self.device_map == AI_DEVICE) and self._is_on_device: + self.text_encoder.to(torch.device("cpu")) + self._is_on_device = False + torch_device_module.empty_cache() + gc.collect() + + @torch.no_grad() + def infer(self, text, image_list=None, manage_cpu_offload=True): + if manage_cpu_offload: + self.load_to_device() if self.is_layered: text = [self.get_image_caption(image_list[0])] @@ -248,10 +262,7 @@ def infer(self, text, image_list=None): prompt_embeds_mask = prompt_embeds_mask.repeat(1, 1, 1) prompt_embeds_mask = prompt_embeds_mask.view(1 * 1, seq_len) - if self.cpu_offload: - if not hasattr(self, "device_map") or self.device_map == AI_DEVICE: - self.text_encoder.to(torch.device("cpu")) - torch_device_module.empty_cache() - gc.collect() + if manage_cpu_offload: + self.offload_to_cpu() return prompt_embeds, prompt_embeds_mask, image_info diff --git a/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py b/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py index 8dcab3830..d06cec07f 100755 --- a/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py @@ -1,5 +1,6 @@ import torch +from lightx2v.common.offload.event_manager import EventSlotWeightAsyncStreamManager from lightx2v.common.offload.manager import WeightAsyncStreamManager from lightx2v.models.networks.qwen_image.infer.transformer_infer import ( QwenImageTransformerInfer, @@ -18,8 +19,14 @@ def __init__(self, config): self.offload_ratio = self.config.get("offload_ratio", 1) offload_granularity = self.config.get("offload_granularity", "block") if offload_granularity == "block": - self.infer_func = self.infer_with_blocks_offload - self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) + if self.config.get("use_event_offload", False): + self.infer_func = self.infer_with_event_offload + self.offload_manager = EventSlotWeightAsyncStreamManager(offload_granularity=offload_granularity) + else: + if self.config.get("offload_resident_blocks", 0) not in (None, 0): + raise ValueError("Qwen-Image resident block offload requires use_event_offload=true") + self.infer_func = self.infer_with_blocks_offload + self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) elif offload_granularity == "phase": self.infer_func = self.infer_with_phases_offload self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) @@ -27,6 +34,8 @@ def __init__(self, config): self.lazy_load = self.config.get("lazy_load", False) if self.lazy_load: + if isinstance(self.offload_manager, EventSlotWeightAsyncStreamManager): + raise NotImplementedError("Qwen-Image event block offload does not support lazy_load") self.offload_manager.init_lazy_load(num_workers=self.config.get("num_disk_workers", 4)) def get_compile_block_key(self, _block_idx, block): @@ -173,3 +182,73 @@ def infer_with_blocks_offload( self.offload_manager.swap_blocks() return hidden_states + + def infer_with_event_offload( + self, + blocks, + hidden_states, + encoder_hidden_states, + temb_img_silu, + temb_txt_silu, + image_rotary_emb, + image_rotary_positions, + modulate_index, + ): + resident_indices = set(getattr(self.block_weights, "resident_block_indices", ())) + offloaded_indices = [idx for idx in range(self.num_blocks) if idx not in resident_indices] + + device_module = self.offload_manager.device_module + current_stream = device_module.current_stream() + compute_stream = self.offload_manager.compute_stream + compute_stream.wait_stream(current_stream) + + scheduled_slots = {} + next_offloaded = 0 + + def prefetch_next(slot_idx): + nonlocal next_offloaded + if next_offloaded >= len(offloaded_indices): + return + block_idx = offloaded_indices[next_offloaded] + self.offload_manager.prefetch_to_slot(slot_idx, block_idx, blocks) + scheduled_slots[block_idx] = slot_idx + next_offloaded += 1 + + if offloaded_indices: + for slot_idx in range(min(self.offload_manager.slot_count, len(offloaded_indices))): + prefetch_next(slot_idx) + + for block_idx, resident_block in enumerate(blocks): + if block_idx in resident_indices: + block = resident_block + slot_idx = None + else: + slot_idx = scheduled_slots.pop(block_idx) + block = self.offload_manager.wait_ready(slot_idx) + + with device_module.stream(compute_stream): + encoder_hidden_states, hidden_states = self.run_block( + block_idx, + block, + hidden_states, + encoder_hidden_states, + temb_img_silu, + temb_txt_silu, + image_rotary_emb, + image_rotary_positions, + modulate_index, + ) + + if slot_idx is not None: + self.offload_manager.record_free(slot_idx) + prefetch_next(slot_idx) + + with device_module.stream(compute_stream): + final_done = compute_stream.record_event() + current_stream.wait_event(final_done) + hidden_states.record_stream(current_stream) + return hidden_states + + def infer(self, block_weights, pre_infer_out): + self.block_weights = block_weights + return super().infer(block_weights, pre_infer_out) diff --git a/lightx2v/models/networks/qwen_image/model.py b/lightx2v/models/networks/qwen_image/model.py index c8619b1fd..2e62e5f0d 100755 --- a/lightx2v/models/networks/qwen_image/model.py +++ b/lightx2v/models/networks/qwen_image/model.py @@ -1,5 +1,8 @@ +import time + import torch import torch.distributed as dist +from loguru import logger from torch.nn import functional as F from lightx2v.models.networks.base_model import BaseTransformerModel @@ -9,9 +12,15 @@ from lightx2v.models.networks.qwen_image.infer.transformer_infer import QwenImageTransformerInfer from lightx2v.models.networks.qwen_image.weights.post_weights import QwenImagePostWeights from lightx2v.models.networks.qwen_image.weights.pre_weights import QwenImagePreWeights -from lightx2v.models.networks.qwen_image.weights.transformer_weights import QwenImageTransformerWeights +from lightx2v.models.networks.qwen_image.weights.transformer_weights import ( + QwenImageTransformerWeights, + release_weight_module_device_tensors, +) from lightx2v.utils.envs import * from lightx2v.utils.utils import * +from lightx2v_platform.base.global_var import AI_DEVICE + +torch_device_module = getattr(torch, AI_DEVICE) class QwenImageTransformerModel(BaseTransformerModel): @@ -21,6 +30,7 @@ class QwenImageTransformerModel(BaseTransformerModel): def __init__(self, model_path, config, device, lora_path=None, lora_strength=1.0): super().__init__(model_path, config, device, None, lora_path, lora_strength) + self._offload_weights_active = False self.in_channels = self.config["in_channels"] self.attention_kwargs = {} if self.lazy_load: @@ -49,6 +59,103 @@ def _init_infer(self): if hasattr(self.transformer_infer, "offload_manager"): self._init_offload_manager() + def _init_offload_manager(self): + if hasattr(self.transformer_weights, "offload_block_cuda_buffers"): + self.transformer_infer.offload_manager.init_cuda_buffer( + blocks_cuda_buffer=self.transformer_weights.offload_block_cuda_buffers, + ) + if self.lazy_load and hasattr(self.transformer_weights, "offload_block_cpu_buffers"): + self.transformer_infer.offload_manager.init_cpu_buffer( + blocks_cpu_buffer=self.transformer_weights.offload_block_cpu_buffers, + ) + + def prepare_offload_weights(self): + """Keep the largest configured set of weights resident for the DiT loop.""" + if not self.cpu_offload: + return + if self._offload_weights_active: + if (self.offload_granularity == "block" and self.config.get("offload_persistent_resident_blocks", False)) or self._keep_model_weights_resident_on_this_rank(): + return + raise RuntimeError("Qwen-Image offload weights are already active") + + self._offload_weights_active = True + if self.offload_granularity == "model": + transfer_start = time.perf_counter() + self.to_cuda() + if self.config.get("qwen_image_rank_aware_model_offload", False): + torch_device_module.synchronize() + rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0 + free_bytes, total_bytes = torch_device_module.mem_get_info() + logger.info( + f"[QwenImage] Rank {rank}: full DiT model H2D completed in " + f"{time.perf_counter() - transfer_start:.3f}s; device free={free_bytes / 2**30:.2f} GiB, " + f"total={total_bytes / 2**30:.2f} GiB" + ) + else: + self.pre_weight.to_cuda() + self.post_weight.to_cuda() + self.transformer_weights.resident_blocks_to_cuda() + rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0 + resident_count = len(self.transformer_weights.resident_block_indices) + if hasattr(torch_device_module, "mem_get_info"): + free_bytes, total_bytes = torch_device_module.mem_get_info() + logger.info( + f"[QwenImage] Rank {rank}: prepared {resident_count}/{self.config['num_layers']} resident DiT blocks; " + f"device free={free_bytes / 2**30:.2f} GiB, total={total_bytes / 2**30:.2f} GiB" + ) + + def finish_offload_weights(self): + """Finish one DiT pass while optionally retaining immutable weights.""" + if not self.cpu_offload or not self._offload_weights_active: + return + if self._keep_model_weights_resident_on_this_rank(): + torch_device_module.synchronize() + return + if self.offload_granularity == "block" and self.config.get("offload_persistent_resident_blocks", False): + torch_device_module.synchronize() + if hasattr(self.transformer_infer.offload_manager, "reset_slots"): + self.transformer_infer.offload_manager.reset_slots() + return + self.force_cleanup_offload_weights() + + def force_cleanup_offload_weights(self): + """Release DiT-resident weights without copying immutable weights back to CPU.""" + if not self.cpu_offload or not self._offload_weights_active: + return + + torch_device_module.synchronize() + if self.offload_granularity == "model": + if self.config.get("qwen_image_model_offload_release_only", False): + release_start = time.perf_counter() + release_weight_module_device_tensors(self.pre_weight) + release_weight_module_device_tensors(self.transformer_weights) + release_weight_module_device_tensors(self.post_weight) + torch_device_module.empty_cache() + rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0 + logger.info( + f"[QwenImage] Rank {rank}: released full DiT device replica without D2H in " + f"{time.perf_counter() - release_start:.3f}s" + ) + else: + self.to_cpu() + else: + if hasattr(self.transformer_infer.offload_manager, "reset_slots"): + self.transformer_infer.offload_manager.reset_slots() + release_weight_module_device_tensors(self.pre_weight) + release_weight_module_device_tensors(self.post_weight) + self.transformer_weights.release_resident_blocks() + self._offload_weights_active = False + + def _keep_model_weights_resident_on_this_rank(self): + return ( + self.offload_granularity == "model" + and self.config.get("qwen_image_rank_aware_model_offload", False) + and dist.is_available() + and dist.is_initialized() + and dist.get_world_size() > 1 + and dist.get_rank() != 0 + ) + @torch.no_grad() def _infer_cond_uncond(self, latents_input, prompt_embeds, infer_condition=True): self.scheduler.infer_condition = infer_condition @@ -94,7 +201,9 @@ def _seq_parallel_post_process(self, noise_pred): @torch.no_grad() def infer(self, inputs): if self.cpu_offload: - if self.offload_granularity == "model" and self.scheduler.step_index == 0: + if self._offload_weights_active: + pass + elif self.offload_granularity == "model" and self.scheduler.step_index == 0: self.to_cuda() elif self.offload_granularity != "model": self.pre_weight.to_cuda() @@ -148,7 +257,9 @@ def infer(self, inputs): self.scheduler.noise_pred = noise_pred if self.cpu_offload: - if self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1: + if self._offload_weights_active: + pass + elif self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1: self.to_cpu() elif self.offload_granularity != "model": self.pre_weight.to_cpu() diff --git a/lightx2v/models/networks/qwen_image/weights/transformer_weights.py b/lightx2v/models/networks/qwen_image/weights/transformer_weights.py index 229f2b839..26f33b9c3 100755 --- a/lightx2v/models/networks/qwen_image/weights/transformer_weights.py +++ b/lightx2v/models/networks/qwen_image/weights/transformer_weights.py @@ -1,4 +1,5 @@ import torch +import torch.distributed as dist from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList from lightx2v.utils.registry_factory import ( @@ -10,6 +11,46 @@ ) +def _resolve_resident_block_indices(value, num_blocks, policy="interleaved"): + if value is None: + value = 0 + if isinstance(value, str): + if value.lower() != "all": + raise ValueError(f"offload_resident_blocks must be an integer or 'all', got {value!r}") + count = num_blocks + elif isinstance(value, bool) or not isinstance(value, int): + raise ValueError(f"offload_resident_blocks must be an integer or 'all', got {value!r}") + else: + count = value + + if not 0 <= count <= num_blocks: + raise ValueError(f"offload_resident_blocks must be between 0 and {num_blocks}, got {count}") + if count == 0: + return frozenset() + if count == num_blocks: + return frozenset(range(num_blocks)) + if policy == "prefix": + return frozenset(range(count)) + if policy == "interleaved": + return frozenset((idx * num_blocks) // count for idx in range(count)) + raise ValueError(f"offload_resident_policy must be 'prefix' or 'interleaved', got {policy!r}") + + +def release_weight_module_device_tensors(module): + """Drop immutable device weights and retain their pinned CPU masters.""" + for child in getattr(module, "_modules", {}).values(): + if child is not None: + release_weight_module_device_tensors(child) + + for _, attr_name, _ in getattr(module, "base_attrs", ()): + value = getattr(module, attr_name, None) + pin_value = getattr(module, f"pin_{attr_name}", None) + if pin_value is not None: + setattr(module, attr_name, None) + elif isinstance(value, torch.Tensor) and value.device.type != "cpu": + setattr(module, attr_name, value.to("cpu")) + + class QwenImageTransformerWeights(WeightModule): def __init__(self, config, lazy_load_path=None, lora_path=None): super().__init__() @@ -22,6 +63,7 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): if self.mm_type != "Default": assert config.get("dit_quantized") is True self.lazy_load = self.config.get("lazy_load", False) + self._configure_resident_blocks(config) blocks = WeightModuleList( QwenImageTransformerAttentionBlock( i, @@ -42,24 +84,25 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): def register_offload_buffers(self, config, lazy_load_path, lora_path): if config["cpu_offload"]: if config["offload_granularity"] == "block": - self.offload_blocks_num = 2 - self.offload_block_cuda_buffers = WeightModuleList( - [ - QwenImageTransformerAttentionBlock( - i, - self.task, - self.mm_type, - self.config, - True, - False, - "transformer_blocks", - lazy_load=self.lazy_load, - lazy_load_path=lazy_load_path, - ) - for i in range(self.offload_blocks_num) - ] - ) - self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) + if len(self.resident_block_indices) < self.blocks_num: + self.offload_blocks_num = 2 + self.offload_block_cuda_buffers = WeightModuleList( + [ + QwenImageTransformerAttentionBlock( + i, + self.task, + self.mm_type, + self.config, + True, + False, + "transformer_blocks", + lazy_load=self.lazy_load, + lazy_load_path=lazy_load_path, + ) + for i in range(self.offload_blocks_num) + ] + ) + self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) self.offload_phase_cuda_buffers = None if self.lazy_load: self.offload_blocks_num = 2 @@ -109,6 +152,30 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): self.add_module("offload_phase_cpu_buffers", self.offload_phase_cpu_buffers) self.offload_block_cpu_buffers = None + def _configure_resident_blocks(self, config): + block_offload_enabled = config.get("cpu_offload", False) and config.get("offload_granularity", "block") == "block" + resident_setting = config.get("offload_resident_blocks", 0) if block_offload_enabled else 0 + if block_offload_enabled and dist.is_available() and dist.is_initialized() and dist.get_rank() == 0: + resident_setting = config.get("offload_resident_blocks_rank0", resident_setting) + if resident_setting not in (None, 0): + if config.get("dit_quantized", False): + raise NotImplementedError("Qwen-Image resident block offload currently supports unquantized weights only") + if config.get("lora_configs"): + raise NotImplementedError("Qwen-Image resident block offload currently does not support LoRA weights") + self.resident_block_indices = _resolve_resident_block_indices( + resident_setting, + self.blocks_num, + config.get("offload_resident_policy", "interleaved"), + ) + + def resident_blocks_to_cuda(self, non_blocking=True): + for block_idx in sorted(self.resident_block_indices): + self.blocks[block_idx].to_cuda(non_blocking=non_blocking) + + def release_resident_blocks(self): + for block_idx in sorted(self.resident_block_indices): + release_weight_module_device_tensors(self.blocks[block_idx]) + class QwenImageTransformerAttentionBlock(WeightModule): def __init__( diff --git a/lightx2v/models/runners/qwen_image/qwen_image_runner.py b/lightx2v/models/runners/qwen_image/qwen_image_runner.py index a659b1564..5c370e94f 100755 --- a/lightx2v/models/runners/qwen_image/qwen_image_runner.py +++ b/lightx2v/models/runners/qwen_image/qwen_image_runner.py @@ -24,6 +24,13 @@ torch_device_module = getattr(torch, AI_DEVICE) +_TEXT_EMBED_DTYPE_TO_CODE = { + torch.float16: 0, + torch.bfloat16: 1, + torch.float32: 2, +} +_TEXT_EMBED_CODE_TO_DTYPE = {code: dtype for dtype, code in _TEXT_EMBED_DTYPE_TO_CODE.items()} + def calculate_dimensions(target_area, ratio): width = math.sqrt(target_area * ratio) @@ -103,16 +110,22 @@ def _run_warmup(self): for height, width in self._WARMUP_RESOLUTIONS: logger.info(f"Warmup: {height}x{width}") + warmup_succeeded = False try: t2i_text_cache = self._prepare_warmup_inputs(height, width, t2i_text_cache) scheduler.generator = None scheduler.prepare(self.input_info) scheduler.step_pre(step_index=0) + self.model.prepare_offload_weights() self.model.infer(self.inputs) scheduler.step_post() + self.model.finish_offload_weights() self.run_vae_decoder(scheduler.latents) torch_device_module.synchronize() + warmup_succeeded = True finally: + if not warmup_succeeded: + self.model.force_cleanup_offload_weights() if self.config.get("cpu_offload", False) and self.config.get("offload_granularity") == "model": self.model.to_cpu() self.clear_warmup_state() @@ -206,6 +219,59 @@ def load_model(self): self.image_encoder = self.load_image_encoder() self.vae = self.load_vae() self.vfi_model = self.load_vfi_model() if "video_frame_interpolation" in self.config else None + self._prepare_resident_text_encoder() + self._prepare_rank_aware_model_offload() + + def _resident_text_encoder_enabled(self): + return self.config.get("qwen_image_resident_text_encoder", False) + + def _prepare_rank_aware_model_offload(self): + if not self.config.get("qwen_image_rank_aware_model_offload", False): + return + if not self.config.get("cpu_offload", False) or self.config.get("offload_granularity") != "model": + raise ValueError("qwen_image_rank_aware_model_offload requires cpu_offload=true and offload_granularity='model'") + if self._resident_text_encoder_enabled(): + raise ValueError("qwen_image_rank_aware_model_offload requires qwen_image_resident_text_encoder=false") + if not self.config.get("qwen_image_single_rank_text_encoder", False): + raise ValueError("qwen_image_rank_aware_model_offload requires qwen_image_single_rank_text_encoder=true") + if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): + raise ValueError("qwen_image_rank_aware_model_offload does not support lazy_load or unload_modules") + if not dist.is_available() or not dist.is_initialized() or dist.get_world_size() <= 1: + raise ValueError("qwen_image_rank_aware_model_offload requires distributed execution") + + rank = dist.get_rank() + if rank != 0: + logger.info(f"[QwenImage] Rank {rank}: preloading full DiT model during initialization") + self.model.prepare_offload_weights() + torch_device_module.synchronize() + if AI_DEVICE == "cuda" and torch.cuda.is_available(): + dist.barrier(device_ids=[torch.cuda.current_device()]) + else: + dist.barrier() + logger.info(f"[QwenImage] Rank {rank}: rank-aware model-offload initialization synchronized") + + def _prepare_resident_text_encoder(self): + if not self._resident_text_encoder_enabled(): + return + if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): + raise ValueError("qwen_image_resident_text_encoder does not support lazy_load or unload_modules") + if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1: + if not self.config.get("qwen_image_single_rank_text_encoder", False): + raise ValueError("Distributed resident Qwen-Image Text Encoder requires qwen_image_single_rank_text_encoder=true") + if dist.get_rank() != 0: + return + + text_encoder = self.text_encoders[0] + if not hasattr(text_encoder, "load_to_device"): + raise ValueError("qwen_image_resident_text_encoder requires a local Text Encoder with load_to_device()") + logger.info("[QwenImage] Keeping Text Encoder resident on global rank 0") + text_encoder.load_to_device() + if hasattr(torch_device_module, "mem_get_info"): + free_bytes, total_bytes = torch_device_module.mem_get_info() + logger.info( + f"[QwenImage] Resident Text Encoder loaded; device free={free_bytes / 2**30:.2f} GiB, " + f"total={total_bytes / 2**30:.2f} GiB" + ) def load_transformer(self): qwen_image_model_kwargs = { @@ -368,24 +434,145 @@ def _run_input_encoder_local_i2i(self): def run_text_encoder(self, text, image_list=None, neg_prompt=None): if GET_RECORDER_MODE(): monitor_cli.lightx2v_input_prompt_len.observe(len(text)) + + if self._single_rank_text_encoder_enabled(): + return self._run_text_encoder_single_rank(text, image_list=image_list, neg_prompt=neg_prompt) + return self._run_text_encoder_local(text, image_list=image_list, neg_prompt=neg_prompt) + + def _single_rank_text_encoder_enabled(self): + return ( + self.config.get("qwen_image_single_rank_text_encoder", False) + and dist.is_available() + and dist.is_initialized() + and dist.get_world_size() > 1 + ) + + def _broadcast_text_encoder_tensors(self, prompt_embeds, negative_prompt_embeds, src=0): + rank = dist.get_rank() + metadata_device = torch.device(AI_DEVICE) + + if rank == src: + if prompt_embeds is None: + raise ValueError("Qwen-Image prompt embedding cannot be None") + if prompt_embeds.ndim != 3: + raise ValueError(f"Qwen-Image prompt embedding must be 3D, got shape {tuple(prompt_embeds.shape)}") + if prompt_embeds.dtype not in _TEXT_EMBED_DTYPE_TO_CODE: + raise ValueError(f"Unsupported Qwen-Image text embedding dtype: {prompt_embeds.dtype}") + + has_negative = negative_prompt_embeds is not None + negative_shape = (0, 0, 0) + tensors = [prompt_embeds.contiguous().view(-1)] + if has_negative: + if negative_prompt_embeds.ndim != 3: + raise ValueError(f"Qwen-Image negative prompt embedding must be 3D, got shape {tuple(negative_prompt_embeds.shape)}") + if negative_prompt_embeds.dtype != prompt_embeds.dtype: + raise ValueError( + f"Qwen-Image prompt embedding dtypes must match, got {prompt_embeds.dtype} and {negative_prompt_embeds.dtype}" + ) + negative_shape = tuple(negative_prompt_embeds.shape) + tensors.append(negative_prompt_embeds.contiguous().view(-1)) + + metadata = torch.tensor( + [*prompt_embeds.shape, int(has_negative), *negative_shape, _TEXT_EMBED_DTYPE_TO_CODE[prompt_embeds.dtype]], + dtype=torch.long, + device=metadata_device, + ) + packed_embeds = torch.cat(tensors) + else: + metadata = torch.empty(8, dtype=torch.long, device=metadata_device) + + dist.broadcast(metadata, src=src) + prompt_batch, prompt_seq_len, prompt_hidden_size, has_negative, negative_batch, negative_seq_len, negative_hidden_size, dtype_code = metadata.tolist() + + dtype = _TEXT_EMBED_CODE_TO_DTYPE.get(dtype_code) + if dtype is None: + raise ValueError(f"Unsupported broadcast Qwen-Image text embedding dtype code: {dtype_code}") + if rank != src: + prompt_numel = prompt_batch * prompt_seq_len * prompt_hidden_size + negative_numel = negative_batch * negative_seq_len * negative_hidden_size if has_negative else 0 + packed_embeds = torch.empty(prompt_numel + negative_numel, dtype=dtype, device=metadata_device) + + dist.broadcast(packed_embeds, src=src) + if rank == src: + return prompt_embeds, negative_prompt_embeds + + prompt_numel = prompt_batch * prompt_seq_len * prompt_hidden_size + prompt_embeds = packed_embeds[:prompt_numel].view(prompt_batch, prompt_seq_len, prompt_hidden_size) + if has_negative: + negative_prompt_embeds = packed_embeds[prompt_numel:].view(negative_batch, negative_seq_len, negative_hidden_size) + else: + negative_prompt_embeds = None + return prompt_embeds, negative_prompt_embeds + + def _run_text_encoder_single_rank(self, text, image_list=None, neg_prompt=None): + rank = dist.get_rank() + if rank == 0: + logger.info("[QwenImage] Running Text Encoder on global rank 0 and broadcasting embeddings") + text_encoder_output = self._run_text_encoder_local(text, image_list=image_list, neg_prompt=neg_prompt) + else: + text_encoder_output = {} + if image_list is not None: + text_encoder_output["image_info"] = self._prepare_local_image_info(image_list) + + prompt_embeds, negative_prompt_embeds = self._broadcast_text_encoder_tensors( + text_encoder_output.get("prompt_embeds"), + text_encoder_output.get("negative_prompt_embeds"), + ) + + self.input_info.txt_seq_lens = [prompt_embeds.shape[1]] + output = {"prompt_embeds": prompt_embeds} + if negative_prompt_embeds is not None: + self.input_info.txt_seq_lens.append(negative_prompt_embeds.shape[1]) + output["negative_prompt_embeds"] = negative_prompt_embeds + if "image_info" in text_encoder_output: + output["image_info"] = text_encoder_output["image_info"] + return output + + def _prepare_local_image_info(self, image_list): + text_encoder = self.text_encoders[0] + vae_image_list = [] + vae_image_info_list = [] + for image in image_list: + _, vae_image, _, vae_image_info = text_encoder.preprocess_image(image) + vae_image_list.append(vae_image) + vae_image_info_list.append(vae_image_info) + return { + "vae_image_list": vae_image_list, + "vae_image_info_list": vae_image_info_list, + } + + def _run_text_encoder_local(self, text, image_list=None, neg_prompt=None): text_encoder_output = {} - if self.config["task"] == "t2i": - prompt_embeds, _, _ = self.text_encoders[0].infer([text]) - self.input_info.txt_seq_lens = [prompt_embeds.shape[1]] - text_encoder_output["prompt_embeds"] = prompt_embeds - if self.config["enable_cfg"] and neg_prompt is not None: - neg_prompt_embeds, _, _ = self.text_encoders[0].infer([neg_prompt]) - self.input_info.txt_seq_lens.append(neg_prompt_embeds.shape[1]) - text_encoder_output["negative_prompt_embeds"] = neg_prompt_embeds - elif self.config["task"] == "i2i": - prompt_embeds, _, image_info = self.text_encoders[0].infer([text], image_list) - self.input_info.txt_seq_lens = [prompt_embeds.shape[1]] - text_encoder_output["prompt_embeds"] = prompt_embeds - text_encoder_output["image_info"] = image_info - if self.config["enable_cfg"] and neg_prompt is not None: - neg_prompt_embeds, _, _ = self.text_encoders[0].infer([neg_prompt], image_list) - self.input_info.txt_seq_lens.append(neg_prompt_embeds.shape[1]) - text_encoder_output["negative_prompt_embeds"] = neg_prompt_embeds + text_encoder = self.text_encoders[0] + manage_offload_externally = hasattr(text_encoder, "load_to_device") and hasattr( + text_encoder, "offload_to_cpu" + ) + keep_resident = self._resident_text_encoder_enabled() + infer_kwargs = {"manage_cpu_offload": False} if manage_offload_externally else {} + + if manage_offload_externally: + text_encoder.load_to_device() + try: + if self.config["task"] == "t2i": + prompt_embeds, _, _ = text_encoder.infer([text], **infer_kwargs) + self.input_info.txt_seq_lens = [prompt_embeds.shape[1]] + text_encoder_output["prompt_embeds"] = prompt_embeds + if self.config["enable_cfg"] and neg_prompt is not None: + neg_prompt_embeds, _, _ = text_encoder.infer([neg_prompt], **infer_kwargs) + self.input_info.txt_seq_lens.append(neg_prompt_embeds.shape[1]) + text_encoder_output["negative_prompt_embeds"] = neg_prompt_embeds + elif self.config["task"] == "i2i": + prompt_embeds, _, image_info = text_encoder.infer([text], image_list, **infer_kwargs) + self.input_info.txt_seq_lens = [prompt_embeds.shape[1]] + text_encoder_output["prompt_embeds"] = prompt_embeds + text_encoder_output["image_info"] = image_info + if self.config["enable_cfg"] and neg_prompt is not None: + neg_prompt_embeds, _, _ = text_encoder.infer([neg_prompt], image_list, **infer_kwargs) + self.input_info.txt_seq_lens.append(neg_prompt_embeds.shape[1]) + text_encoder_output["negative_prompt_embeds"] = neg_prompt_embeds + finally: + if manage_offload_externally and not keep_resident: + text_encoder.offload_to_cpu() return text_encoder_output @ProfilingContext4DebugL1("Run VAE Encoder", recorder_mode=GET_RECORDER_MODE(), metrics_func=monitor_cli.lightx2v_run_vae_encoder_image_duration, metrics_labels=["QwenImageRunner"]) @@ -416,23 +603,30 @@ def run_vae_decoder(self, latents): def run(self, total_steps=None): if total_steps is None: total_steps = self.model.scheduler.infer_steps - for step_index in range(total_steps): - logger.info(f"==> step_index: {step_index + 1} / {total_steps}") + try: + self.model.prepare_offload_weights() + for step_index in range(total_steps): + logger.info(f"==> step_index: {step_index + 1} / {total_steps}") - with ProfilingContext4DebugL1("step_pre"): - self.model.scheduler.step_pre(step_index=step_index) + with ProfilingContext4DebugL1("step_pre"): + self.model.scheduler.step_pre(step_index=step_index) - with ProfilingContext4DebugL1("🚀 infer_main"): - # Example of torch trace profile: - # with TorchTraceProfileContext() as profile: - # profile.run(self.model.infer, self.inputs) - self.model.infer(self.inputs) + with ProfilingContext4DebugL1("🚀 infer_main"): + # Example of torch trace profile: + # with TorchTraceProfileContext() as profile: + # profile.run(self.model.infer, self.inputs) + self.model.infer(self.inputs) - with ProfilingContext4DebugL1("step_post"): - self.model.scheduler.step_post() + with ProfilingContext4DebugL1("step_post"): + self.model.scheduler.step_post() - if self.progress_callback: - self.progress_callback(((step_index + 1) / total_steps) * 100, 100) + if self.progress_callback: + self.progress_callback(((step_index + 1) / total_steps) * 100, 100) + except Exception: + self.model.force_cleanup_offload_weights() + raise + else: + self.model.finish_offload_weights() return self.model.scheduler.latents, self.model.scheduler.generator From 5c9a3c8e26a623820b9a0dddc008949353c8e473 Mon Sep 17 00:00:00 2001 From: Super User Date: Sat, 8 Aug 2026 17:28:37 +0800 Subject: [PATCH 2/4] refactor(qwen-image): simplify distributed offload paths --- .../flux2/weights/transformer_weights.py | 12 +-- .../infer/offload/transformer_infer.py | 5 +- .../qwen_image/weights/transformer_weights.py | 17 +--- .../runners/qwen_image/qwen_image_runner.py | 91 +++++-------------- 4 files changed, 28 insertions(+), 97 deletions(-) diff --git a/lightx2v/models/networks/flux2/weights/transformer_weights.py b/lightx2v/models/networks/flux2/weights/transformer_weights.py index db8d96aff..9abf5b280 100644 --- a/lightx2v/models/networks/flux2/weights/transformer_weights.py +++ b/lightx2v/models/networks/flux2/weights/transformer_weights.py @@ -16,16 +16,8 @@ def _resolve_resident_block_indices(value, num_blocks, policy, config_key): at the front, which gives the offload stream regular compute windows in which to prefetch the next non-resident block. """ - if value is None: - value = 0 - if isinstance(value, str): - if value.lower() != "all": - raise ValueError(f"{config_key} must be an integer or 'all', got {value!r}") - count = num_blocks - elif isinstance(value, bool) or not isinstance(value, int): - raise ValueError(f"{config_key} must be an integer or 'all', got {value!r}") - else: - count = value + # Resident block counts are integers or "all". + count = num_blocks if value == "all" else value if not 0 <= count <= num_blocks: raise ValueError(f"{config_key} must be between 0 and {num_blocks}, got {count}") diff --git a/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py b/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py index d06cec07f..3e6b0f0e5 100755 --- a/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py @@ -23,8 +23,6 @@ def __init__(self, config): self.infer_func = self.infer_with_event_offload self.offload_manager = EventSlotWeightAsyncStreamManager(offload_granularity=offload_granularity) else: - if self.config.get("offload_resident_blocks", 0) not in (None, 0): - raise ValueError("Qwen-Image resident block offload requires use_event_offload=true") self.infer_func = self.infer_with_blocks_offload self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) elif offload_granularity == "phase": @@ -33,9 +31,8 @@ def __init__(self, config): self.compiled_phases = {} self.lazy_load = self.config.get("lazy_load", False) + # Event block offload does not support lazy_load. if self.lazy_load: - if isinstance(self.offload_manager, EventSlotWeightAsyncStreamManager): - raise NotImplementedError("Qwen-Image event block offload does not support lazy_load") self.offload_manager.init_lazy_load(num_workers=self.config.get("num_disk_workers", 4)) def get_compile_block_key(self, _block_idx, block): diff --git a/lightx2v/models/networks/qwen_image/weights/transformer_weights.py b/lightx2v/models/networks/qwen_image/weights/transformer_weights.py index 26f33b9c3..c482c9f13 100755 --- a/lightx2v/models/networks/qwen_image/weights/transformer_weights.py +++ b/lightx2v/models/networks/qwen_image/weights/transformer_weights.py @@ -12,16 +12,8 @@ def _resolve_resident_block_indices(value, num_blocks, policy="interleaved"): - if value is None: - value = 0 - if isinstance(value, str): - if value.lower() != "all": - raise ValueError(f"offload_resident_blocks must be an integer or 'all', got {value!r}") - count = num_blocks - elif isinstance(value, bool) or not isinstance(value, int): - raise ValueError(f"offload_resident_blocks must be an integer or 'all', got {value!r}") - else: - count = value + # offload_resident_blocks is an integer or "all". + count = num_blocks if value == "all" else value if not 0 <= count <= num_blocks: raise ValueError(f"offload_resident_blocks must be between 0 and {num_blocks}, got {count}") @@ -157,11 +149,6 @@ def _configure_resident_blocks(self, config): resident_setting = config.get("offload_resident_blocks", 0) if block_offload_enabled else 0 if block_offload_enabled and dist.is_available() and dist.is_initialized() and dist.get_rank() == 0: resident_setting = config.get("offload_resident_blocks_rank0", resident_setting) - if resident_setting not in (None, 0): - if config.get("dit_quantized", False): - raise NotImplementedError("Qwen-Image resident block offload currently supports unquantized weights only") - if config.get("lora_configs"): - raise NotImplementedError("Qwen-Image resident block offload currently does not support LoRA weights") self.resident_block_indices = _resolve_resident_block_indices( resident_setting, self.blocks_num, diff --git a/lightx2v/models/runners/qwen_image/qwen_image_runner.py b/lightx2v/models/runners/qwen_image/qwen_image_runner.py index 5c370e94f..8eabb1a51 100755 --- a/lightx2v/models/runners/qwen_image/qwen_image_runner.py +++ b/lightx2v/models/runners/qwen_image/qwen_image_runner.py @@ -24,13 +24,6 @@ torch_device_module = getattr(torch, AI_DEVICE) -_TEXT_EMBED_DTYPE_TO_CODE = { - torch.float16: 0, - torch.bfloat16: 1, - torch.float32: 2, -} -_TEXT_EMBED_CODE_TO_DTYPE = {code: dtype for dtype, code in _TEXT_EMBED_DTYPE_TO_CODE.items()} - def calculate_dimensions(target_area, ratio): width = math.sqrt(target_area * ratio) @@ -228,17 +221,8 @@ def _resident_text_encoder_enabled(self): def _prepare_rank_aware_model_offload(self): if not self.config.get("qwen_image_rank_aware_model_offload", False): return - if not self.config.get("cpu_offload", False) or self.config.get("offload_granularity") != "model": - raise ValueError("qwen_image_rank_aware_model_offload requires cpu_offload=true and offload_granularity='model'") - if self._resident_text_encoder_enabled(): - raise ValueError("qwen_image_rank_aware_model_offload requires qwen_image_resident_text_encoder=false") - if not self.config.get("qwen_image_single_rank_text_encoder", False): - raise ValueError("qwen_image_rank_aware_model_offload requires qwen_image_single_rank_text_encoder=true") - if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): - raise ValueError("qwen_image_rank_aware_model_offload does not support lazy_load or unload_modules") - if not dist.is_available() or not dist.is_initialized() or dist.get_world_size() <= 1: - raise ValueError("qwen_image_rank_aware_model_offload requires distributed execution") + # Rank-aware mode uses distributed model offload, rank-0 TE broadcast, and eager modules. rank = dist.get_rank() if rank != 0: logger.info(f"[QwenImage] Rank {rank}: preloading full DiT model during initialization") @@ -253,25 +237,17 @@ def _prepare_rank_aware_model_offload(self): def _prepare_resident_text_encoder(self): if not self._resident_text_encoder_enabled(): return - if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): - raise ValueError("qwen_image_resident_text_encoder does not support lazy_load or unload_modules") + + # Resident mode uses the local Text Encoder and rank-0 broadcast in distributed runs. if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1: - if not self.config.get("qwen_image_single_rank_text_encoder", False): - raise ValueError("Distributed resident Qwen-Image Text Encoder requires qwen_image_single_rank_text_encoder=true") if dist.get_rank() != 0: return text_encoder = self.text_encoders[0] - if not hasattr(text_encoder, "load_to_device"): - raise ValueError("qwen_image_resident_text_encoder requires a local Text Encoder with load_to_device()") logger.info("[QwenImage] Keeping Text Encoder resident on global rank 0") text_encoder.load_to_device() - if hasattr(torch_device_module, "mem_get_info"): - free_bytes, total_bytes = torch_device_module.mem_get_info() - logger.info( - f"[QwenImage] Resident Text Encoder loaded; device free={free_bytes / 2**30:.2f} GiB, " - f"total={total_bytes / 2**30:.2f} GiB" - ) + free_bytes, total_bytes = torch_device_module.mem_get_info() + logger.info(f"[QwenImage] Resident Text Encoder loaded; device free={free_bytes / 2**30:.2f} GiB, total={total_bytes / 2**30:.2f} GiB") def load_transformer(self): qwen_image_model_kwargs = { @@ -435,65 +411,46 @@ def run_text_encoder(self, text, image_list=None, neg_prompt=None): if GET_RECORDER_MODE(): monitor_cli.lightx2v_input_prompt_len.observe(len(text)) - if self._single_rank_text_encoder_enabled(): - return self._run_text_encoder_single_rank(text, image_list=image_list, neg_prompt=neg_prompt) + if self._rank0_text_encoder_broadcast_enabled(): + return self._run_text_encoder_rank0_broadcast(text, image_list=image_list, neg_prompt=neg_prompt) return self._run_text_encoder_local(text, image_list=image_list, neg_prompt=neg_prompt) - def _single_rank_text_encoder_enabled(self): - return ( - self.config.get("qwen_image_single_rank_text_encoder", False) - and dist.is_available() - and dist.is_initialized() - and dist.get_world_size() > 1 - ) + def _rank0_text_encoder_broadcast_enabled(self): + # rank0_broadcast is configured only for initialized multi-rank runs. + return self.config.get("text_encoder_mode") == "rank0_broadcast" - def _broadcast_text_encoder_tensors(self, prompt_embeds, negative_prompt_embeds, src=0): + def _broadcast_text_encoder_tensors(self, prompt_embeds, negative_prompt_embeds): + # Text Encoder outputs are BF16 tensors shaped [batch, sequence, hidden]. rank = dist.get_rank() metadata_device = torch.device(AI_DEVICE) - if rank == src: - if prompt_embeds is None: - raise ValueError("Qwen-Image prompt embedding cannot be None") - if prompt_embeds.ndim != 3: - raise ValueError(f"Qwen-Image prompt embedding must be 3D, got shape {tuple(prompt_embeds.shape)}") - if prompt_embeds.dtype not in _TEXT_EMBED_DTYPE_TO_CODE: - raise ValueError(f"Unsupported Qwen-Image text embedding dtype: {prompt_embeds.dtype}") - + if rank == 0: has_negative = negative_prompt_embeds is not None negative_shape = (0, 0, 0) tensors = [prompt_embeds.contiguous().view(-1)] if has_negative: - if negative_prompt_embeds.ndim != 3: - raise ValueError(f"Qwen-Image negative prompt embedding must be 3D, got shape {tuple(negative_prompt_embeds.shape)}") - if negative_prompt_embeds.dtype != prompt_embeds.dtype: - raise ValueError( - f"Qwen-Image prompt embedding dtypes must match, got {prompt_embeds.dtype} and {negative_prompt_embeds.dtype}" - ) negative_shape = tuple(negative_prompt_embeds.shape) tensors.append(negative_prompt_embeds.contiguous().view(-1)) metadata = torch.tensor( - [*prompt_embeds.shape, int(has_negative), *negative_shape, _TEXT_EMBED_DTYPE_TO_CODE[prompt_embeds.dtype]], + [*prompt_embeds.shape, int(has_negative), *negative_shape], dtype=torch.long, device=metadata_device, ) packed_embeds = torch.cat(tensors) else: - metadata = torch.empty(8, dtype=torch.long, device=metadata_device) + metadata = torch.empty(7, dtype=torch.long, device=metadata_device) - dist.broadcast(metadata, src=src) - prompt_batch, prompt_seq_len, prompt_hidden_size, has_negative, negative_batch, negative_seq_len, negative_hidden_size, dtype_code = metadata.tolist() + dist.broadcast(metadata, src=0) + prompt_batch, prompt_seq_len, prompt_hidden_size, has_negative, negative_batch, negative_seq_len, negative_hidden_size = metadata.tolist() - dtype = _TEXT_EMBED_CODE_TO_DTYPE.get(dtype_code) - if dtype is None: - raise ValueError(f"Unsupported broadcast Qwen-Image text embedding dtype code: {dtype_code}") - if rank != src: + if rank != 0: prompt_numel = prompt_batch * prompt_seq_len * prompt_hidden_size negative_numel = negative_batch * negative_seq_len * negative_hidden_size if has_negative else 0 - packed_embeds = torch.empty(prompt_numel + negative_numel, dtype=dtype, device=metadata_device) + packed_embeds = torch.empty(prompt_numel + negative_numel, dtype=torch.bfloat16, device=metadata_device) - dist.broadcast(packed_embeds, src=src) - if rank == src: + dist.broadcast(packed_embeds, src=0) + if rank == 0: return prompt_embeds, negative_prompt_embeds prompt_numel = prompt_batch * prompt_seq_len * prompt_hidden_size @@ -504,7 +461,7 @@ def _broadcast_text_encoder_tensors(self, prompt_embeds, negative_prompt_embeds, negative_prompt_embeds = None return prompt_embeds, negative_prompt_embeds - def _run_text_encoder_single_rank(self, text, image_list=None, neg_prompt=None): + def _run_text_encoder_rank0_broadcast(self, text, image_list=None, neg_prompt=None): rank = dist.get_rank() if rank == 0: logger.info("[QwenImage] Running Text Encoder on global rank 0 and broadcasting embeddings") @@ -544,9 +501,7 @@ def _prepare_local_image_info(self, image_list): def _run_text_encoder_local(self, text, image_list=None, neg_prompt=None): text_encoder_output = {} text_encoder = self.text_encoders[0] - manage_offload_externally = hasattr(text_encoder, "load_to_device") and hasattr( - text_encoder, "offload_to_cpu" - ) + manage_offload_externally = hasattr(text_encoder, "load_to_device") and hasattr(text_encoder, "offload_to_cpu") keep_resident = self._resident_text_encoder_enabled() infer_kwargs = {"manage_cpu_offload": False} if manage_offload_externally else {} From 266a1c247be944e6aa19aac1f10bdae6953ab3eb Mon Sep 17 00:00:00 2001 From: Super User Date: Sat, 8 Aug 2026 17:32:31 +0800 Subject: [PATCH 3/4] style(qwen-image): format offload logging --- lightx2v/models/networks/qwen_image/model.py | 8 ++------ 1 file changed, 2 insertions(+), 6 deletions(-) diff --git a/lightx2v/models/networks/qwen_image/model.py b/lightx2v/models/networks/qwen_image/model.py index 2e62e5f0d..2090fe2a5 100755 --- a/lightx2v/models/networks/qwen_image/model.py +++ b/lightx2v/models/networks/qwen_image/model.py @@ -100,8 +100,7 @@ def prepare_offload_weights(self): if hasattr(torch_device_module, "mem_get_info"): free_bytes, total_bytes = torch_device_module.mem_get_info() logger.info( - f"[QwenImage] Rank {rank}: prepared {resident_count}/{self.config['num_layers']} resident DiT blocks; " - f"device free={free_bytes / 2**30:.2f} GiB, total={total_bytes / 2**30:.2f} GiB" + f"[QwenImage] Rank {rank}: prepared {resident_count}/{self.config['num_layers']} resident DiT blocks; device free={free_bytes / 2**30:.2f} GiB, total={total_bytes / 2**30:.2f} GiB" ) def finish_offload_weights(self): @@ -132,10 +131,7 @@ def force_cleanup_offload_weights(self): release_weight_module_device_tensors(self.post_weight) torch_device_module.empty_cache() rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0 - logger.info( - f"[QwenImage] Rank {rank}: released full DiT device replica without D2H in " - f"{time.perf_counter() - release_start:.3f}s" - ) + logger.info(f"[QwenImage] Rank {rank}: released full DiT device replica without D2H in {time.perf_counter() - release_start:.3f}s") else: self.to_cpu() else: From 0a7f06e66086994a6ea570fff90f9e45b66fdf7f Mon Sep 17 00:00:00 2001 From: Super User Date: Sun, 9 Aug 2026 04:49:01 +0800 Subject: [PATCH 4/4] refactor(qwen-image): simplify text encoder residency --- .../runners/qwen_image/qwen_image_runner.py | 26 +++---------------- 1 file changed, 4 insertions(+), 22 deletions(-) diff --git a/lightx2v/models/runners/qwen_image/qwen_image_runner.py b/lightx2v/models/runners/qwen_image/qwen_image_runner.py index 8eabb1a51..a9b5d0590 100755 --- a/lightx2v/models/runners/qwen_image/qwen_image_runner.py +++ b/lightx2v/models/runners/qwen_image/qwen_image_runner.py @@ -212,12 +212,8 @@ def load_model(self): self.image_encoder = self.load_image_encoder() self.vae = self.load_vae() self.vfi_model = self.load_vfi_model() if "video_frame_interpolation" in self.config else None - self._prepare_resident_text_encoder() self._prepare_rank_aware_model_offload() - def _resident_text_encoder_enabled(self): - return self.config.get("qwen_image_resident_text_encoder", False) - def _prepare_rank_aware_model_offload(self): if not self.config.get("qwen_image_rank_aware_model_offload", False): return @@ -234,21 +230,6 @@ def _prepare_rank_aware_model_offload(self): dist.barrier() logger.info(f"[QwenImage] Rank {rank}: rank-aware model-offload initialization synchronized") - def _prepare_resident_text_encoder(self): - if not self._resident_text_encoder_enabled(): - return - - # Resident mode uses the local Text Encoder and rank-0 broadcast in distributed runs. - if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1: - if dist.get_rank() != 0: - return - - text_encoder = self.text_encoders[0] - logger.info("[QwenImage] Keeping Text Encoder resident on global rank 0") - text_encoder.load_to_device() - free_bytes, total_bytes = torch_device_module.mem_get_info() - logger.info(f"[QwenImage] Resident Text Encoder loaded; device free={free_bytes / 2**30:.2f} GiB, total={total_bytes / 2**30:.2f} GiB") - def load_transformer(self): qwen_image_model_kwargs = { "model_path": os.path.join(self.config["model_path"], "transformer"), @@ -272,6 +253,8 @@ def load_text_encoder(self): """ encoder_config = dict(self.config) encoder_config.update(self.config.get("lightllm_config", {})) + if self._rank0_text_encoder_broadcast_enabled() and dist.get_rank() != 0: + encoder_config["qwen25vl_cpu_offload"] = True if self.text_encoder_type == "lightllm_service": from lightx2v.models.input_encoders.lightllm import LightLLMServiceTextEncoder @@ -285,7 +268,7 @@ def load_text_encoder(self): text_encoder = LightLLMKernelTextEncoder(encoder_config) else: # baseline or default logger.info("Loading HuggingFace baseline text encoder") - text_encoder = Qwen25_VLForConditionalGeneration_TextEncoder(self.config) + text_encoder = Qwen25_VLForConditionalGeneration_TextEncoder(encoder_config) text_encoders = [text_encoder] return text_encoders @@ -502,7 +485,6 @@ def _run_text_encoder_local(self, text, image_list=None, neg_prompt=None): text_encoder_output = {} text_encoder = self.text_encoders[0] manage_offload_externally = hasattr(text_encoder, "load_to_device") and hasattr(text_encoder, "offload_to_cpu") - keep_resident = self._resident_text_encoder_enabled() infer_kwargs = {"manage_cpu_offload": False} if manage_offload_externally else {} if manage_offload_externally: @@ -526,7 +508,7 @@ def _run_text_encoder_local(self, text, image_list=None, neg_prompt=None): self.input_info.txt_seq_lens.append(neg_prompt_embeds.shape[1]) text_encoder_output["negative_prompt_embeds"] = neg_prompt_embeds finally: - if manage_offload_externally and not keep_resident: + if manage_offload_externally: text_encoder.offload_to_cpu() return text_encoder_output