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/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 8dcab3830..3e6b0f0e5 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,14 +19,19 @@ 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: + 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) self.compiled_phases = {} self.lazy_load = self.config.get("lazy_load", False) + # Event block offload does not support lazy_load. if self.lazy_load: self.offload_manager.init_lazy_load(num_workers=self.config.get("num_disk_workers", 4)) @@ -173,3 +179,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..2090fe2a5 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,99 @@ 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; 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 {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 +197,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 +253,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..c482c9f13 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,38 @@ ) +def _resolve_resident_block_indices(value, num_blocks, policy="interleaved"): + # 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}") + 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 +55,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 +76,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 +144,25 @@ 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) + 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 081f9c14b..aa915f80f 100755 --- a/lightx2v/models/runners/qwen_image/qwen_image_runner.py +++ b/lightx2v/models/runners/qwen_image/qwen_image_runner.py @@ -100,16 +100,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() @@ -203,6 +209,23 @@ 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_rank_aware_model_offload() + + def _prepare_rank_aware_model_offload(self): + if not self.config.get("qwen_image_rank_aware_model_offload", False): + return + + # 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") + 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 load_transformer(self): qwen_image_model_kwargs = { @@ -227,6 +250,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 @@ -240,7 +265,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 @@ -365,24 +390,123 @@ 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._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 _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): + # Text Encoder outputs are BF16 tensors shaped [batch, sequence, hidden]. + rank = dist.get_rank() + metadata_device = torch.device(AI_DEVICE) + + 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: + 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], + dtype=torch.long, + device=metadata_device, + ) + packed_embeds = torch.cat(tensors) + else: + metadata = torch.empty(7, dtype=torch.long, device=metadata_device) + + 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() + + 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=torch.bfloat16, device=metadata_device) + + 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 + 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_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") + 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") + 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: + 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"]) @@ -413,23 +537,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}") - - 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("step_post"): - self.model.scheduler.step_post() - - if self.progress_callback: - self.progress_callback(((step_index + 1) / total_steps) * 100, 100) + 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("🚀 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() + + 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