Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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])]
Expand Down Expand Up @@ -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
12 changes: 2 additions & 10 deletions lightx2v/models/networks/flux2/weights/transformer_weights.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}")
Expand Down
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -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))

Expand Down Expand Up @@ -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)
113 changes: 110 additions & 3 deletions lightx2v/models/networks/qwen_image/model.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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):
Expand All @@ -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:
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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()
Expand Down
Loading
Loading