Skip to content

Commit 6bdc174

Browse files
author
Super User
committed
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.
1 parent 1e86807 commit 6bdc174

5 files changed

Lines changed: 527 additions & 65 deletions

File tree

lightx2v/models/input_encoders/hf/qwen25/qwen25_vlforconditionalgeneration.py

Lines changed: 19 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -73,6 +73,7 @@ def __init__(self, config):
7373
self.cpu_offload = config.get("qwen25vl_cpu_offload", config.get("cpu_offload", False))
7474
self.dtype = torch.bfloat16
7575
self.load()
76+
self._is_on_device = not self.cpu_offload
7677

7778
def load(self):
7879
if self.config.get("qwen25vl_quantized", False):
@@ -152,11 +153,24 @@ def get_image_caption(self, prompt_image):
152153
output_text = self.vl_processor.batch_decode(generated_ids_trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
153154
return output_text.strip()
154155

155-
@torch.no_grad()
156-
def infer(self, text, image_list=None):
156+
def load_to_device(self):
157157
if self.cpu_offload:
158-
if not hasattr(self, "device_map") or self.device_map == AI_DEVICE:
158+
if (not hasattr(self, "device_map") or self.device_map == AI_DEVICE) and not self._is_on_device:
159159
self.text_encoder.to(AI_DEVICE)
160+
self._is_on_device = True
161+
162+
def offload_to_cpu(self):
163+
if self.cpu_offload:
164+
if (not hasattr(self, "device_map") or self.device_map == AI_DEVICE) and self._is_on_device:
165+
self.text_encoder.to(torch.device("cpu"))
166+
self._is_on_device = False
167+
torch_device_module.empty_cache()
168+
gc.collect()
169+
170+
@torch.no_grad()
171+
def infer(self, text, image_list=None, manage_cpu_offload=True):
172+
if manage_cpu_offload:
173+
self.load_to_device()
160174

161175
if self.is_layered:
162176
text = [self.get_image_caption(image_list[0])]
@@ -248,10 +262,7 @@ def infer(self, text, image_list=None):
248262
prompt_embeds_mask = prompt_embeds_mask.repeat(1, 1, 1)
249263
prompt_embeds_mask = prompt_embeds_mask.view(1 * 1, seq_len)
250264

251-
if self.cpu_offload:
252-
if not hasattr(self, "device_map") or self.device_map == AI_DEVICE:
253-
self.text_encoder.to(torch.device("cpu"))
254-
torch_device_module.empty_cache()
255-
gc.collect()
265+
if manage_cpu_offload:
266+
self.offload_to_cpu()
256267

257268
return prompt_embeds, prompt_embeds_mask, image_info

lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py

Lines changed: 81 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
import torch
22

3+
from lightx2v.common.offload.event_manager import EventSlotWeightAsyncStreamManager
34
from lightx2v.common.offload.manager import WeightAsyncStreamManager
45
from lightx2v.models.networks.qwen_image.infer.transformer_infer import (
56
QwenImageTransformerInfer,
@@ -18,15 +19,23 @@ def __init__(self, config):
1819
self.offload_ratio = self.config.get("offload_ratio", 1)
1920
offload_granularity = self.config.get("offload_granularity", "block")
2021
if offload_granularity == "block":
21-
self.infer_func = self.infer_with_blocks_offload
22-
self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity)
22+
if self.config.get("use_event_offload", False):
23+
self.infer_func = self.infer_with_event_offload
24+
self.offload_manager = EventSlotWeightAsyncStreamManager(offload_granularity=offload_granularity)
25+
else:
26+
if self.config.get("offload_resident_blocks", 0) not in (None, 0):
27+
raise ValueError("Qwen-Image resident block offload requires use_event_offload=true")
28+
self.infer_func = self.infer_with_blocks_offload
29+
self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity)
2330
elif offload_granularity == "phase":
2431
self.infer_func = self.infer_with_phases_offload
2532
self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity)
2633
self.compiled_phases = {}
2734

2835
self.lazy_load = self.config.get("lazy_load", False)
2936
if self.lazy_load:
37+
if isinstance(self.offload_manager, EventSlotWeightAsyncStreamManager):
38+
raise NotImplementedError("Qwen-Image event block offload does not support lazy_load")
3039
self.offload_manager.init_lazy_load(num_workers=self.config.get("num_disk_workers", 4))
3140

3241
def get_compile_block_key(self, _block_idx, block):
@@ -173,3 +182,73 @@ def infer_with_blocks_offload(
173182
self.offload_manager.swap_blocks()
174183

175184
return hidden_states
185+
186+
def infer_with_event_offload(
187+
self,
188+
blocks,
189+
hidden_states,
190+
encoder_hidden_states,
191+
temb_img_silu,
192+
temb_txt_silu,
193+
image_rotary_emb,
194+
image_rotary_positions,
195+
modulate_index,
196+
):
197+
resident_indices = set(getattr(self.block_weights, "resident_block_indices", ()))
198+
offloaded_indices = [idx for idx in range(self.num_blocks) if idx not in resident_indices]
199+
200+
device_module = self.offload_manager.device_module
201+
current_stream = device_module.current_stream()
202+
compute_stream = self.offload_manager.compute_stream
203+
compute_stream.wait_stream(current_stream)
204+
205+
scheduled_slots = {}
206+
next_offloaded = 0
207+
208+
def prefetch_next(slot_idx):
209+
nonlocal next_offloaded
210+
if next_offloaded >= len(offloaded_indices):
211+
return
212+
block_idx = offloaded_indices[next_offloaded]
213+
self.offload_manager.prefetch_to_slot(slot_idx, block_idx, blocks)
214+
scheduled_slots[block_idx] = slot_idx
215+
next_offloaded += 1
216+
217+
if offloaded_indices:
218+
for slot_idx in range(min(self.offload_manager.slot_count, len(offloaded_indices))):
219+
prefetch_next(slot_idx)
220+
221+
for block_idx, resident_block in enumerate(blocks):
222+
if block_idx in resident_indices:
223+
block = resident_block
224+
slot_idx = None
225+
else:
226+
slot_idx = scheduled_slots.pop(block_idx)
227+
block = self.offload_manager.wait_ready(slot_idx)
228+
229+
with device_module.stream(compute_stream):
230+
encoder_hidden_states, hidden_states = self.run_block(
231+
block_idx,
232+
block,
233+
hidden_states,
234+
encoder_hidden_states,
235+
temb_img_silu,
236+
temb_txt_silu,
237+
image_rotary_emb,
238+
image_rotary_positions,
239+
modulate_index,
240+
)
241+
242+
if slot_idx is not None:
243+
self.offload_manager.record_free(slot_idx)
244+
prefetch_next(slot_idx)
245+
246+
with device_module.stream(compute_stream):
247+
final_done = compute_stream.record_event()
248+
current_stream.wait_event(final_done)
249+
hidden_states.record_stream(current_stream)
250+
return hidden_states
251+
252+
def infer(self, block_weights, pre_infer_out):
253+
self.block_weights = block_weights
254+
return super().infer(block_weights, pre_infer_out)

lightx2v/models/networks/qwen_image/model.py

Lines changed: 114 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,8 @@
1+
import time
2+
13
import torch
24
import torch.distributed as dist
5+
from loguru import logger
36
from torch.nn import functional as F
47

58
from lightx2v.models.networks.base_model import BaseTransformerModel
@@ -9,9 +12,15 @@
912
from lightx2v.models.networks.qwen_image.infer.transformer_infer import QwenImageTransformerInfer
1013
from lightx2v.models.networks.qwen_image.weights.post_weights import QwenImagePostWeights
1114
from lightx2v.models.networks.qwen_image.weights.pre_weights import QwenImagePreWeights
12-
from lightx2v.models.networks.qwen_image.weights.transformer_weights import QwenImageTransformerWeights
15+
from lightx2v.models.networks.qwen_image.weights.transformer_weights import (
16+
QwenImageTransformerWeights,
17+
release_weight_module_device_tensors,
18+
)
1319
from lightx2v.utils.envs import *
1420
from lightx2v.utils.utils import *
21+
from lightx2v_platform.base.global_var import AI_DEVICE
22+
23+
torch_device_module = getattr(torch, AI_DEVICE)
1524

1625

1726
class QwenImageTransformerModel(BaseTransformerModel):
@@ -21,6 +30,7 @@ class QwenImageTransformerModel(BaseTransformerModel):
2130

2231
def __init__(self, model_path, config, device, lora_path=None, lora_strength=1.0):
2332
super().__init__(model_path, config, device, None, lora_path, lora_strength)
33+
self._offload_weights_active = False
2434
self.in_channels = self.config["in_channels"]
2535
self.attention_kwargs = {}
2636
if self.lazy_load:
@@ -49,6 +59,103 @@ def _init_infer(self):
4959
if hasattr(self.transformer_infer, "offload_manager"):
5060
self._init_offload_manager()
5161

62+
def _init_offload_manager(self):
63+
if hasattr(self.transformer_weights, "offload_block_cuda_buffers"):
64+
self.transformer_infer.offload_manager.init_cuda_buffer(
65+
blocks_cuda_buffer=self.transformer_weights.offload_block_cuda_buffers,
66+
)
67+
if self.lazy_load and hasattr(self.transformer_weights, "offload_block_cpu_buffers"):
68+
self.transformer_infer.offload_manager.init_cpu_buffer(
69+
blocks_cpu_buffer=self.transformer_weights.offload_block_cpu_buffers,
70+
)
71+
72+
def prepare_offload_weights(self):
73+
"""Keep the largest configured set of weights resident for the DiT loop."""
74+
if not self.cpu_offload:
75+
return
76+
if self._offload_weights_active:
77+
if (self.offload_granularity == "block" and self.config.get("offload_persistent_resident_blocks", False)) or self._keep_model_weights_resident_on_this_rank():
78+
return
79+
raise RuntimeError("Qwen-Image offload weights are already active")
80+
81+
self._offload_weights_active = True
82+
if self.offload_granularity == "model":
83+
transfer_start = time.perf_counter()
84+
self.to_cuda()
85+
if self.config.get("qwen_image_rank_aware_model_offload", False):
86+
torch_device_module.synchronize()
87+
rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0
88+
free_bytes, total_bytes = torch_device_module.mem_get_info()
89+
logger.info(
90+
f"[QwenImage] Rank {rank}: full DiT model H2D completed in "
91+
f"{time.perf_counter() - transfer_start:.3f}s; device free={free_bytes / 2**30:.2f} GiB, "
92+
f"total={total_bytes / 2**30:.2f} GiB"
93+
)
94+
else:
95+
self.pre_weight.to_cuda()
96+
self.post_weight.to_cuda()
97+
self.transformer_weights.resident_blocks_to_cuda()
98+
rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0
99+
resident_count = len(self.transformer_weights.resident_block_indices)
100+
if hasattr(torch_device_module, "mem_get_info"):
101+
free_bytes, total_bytes = torch_device_module.mem_get_info()
102+
logger.info(
103+
f"[QwenImage] Rank {rank}: prepared {resident_count}/{self.config['num_layers']} resident DiT blocks; "
104+
f"device free={free_bytes / 2**30:.2f} GiB, total={total_bytes / 2**30:.2f} GiB"
105+
)
106+
107+
def finish_offload_weights(self):
108+
"""Finish one DiT pass while optionally retaining immutable weights."""
109+
if not self.cpu_offload or not self._offload_weights_active:
110+
return
111+
if self._keep_model_weights_resident_on_this_rank():
112+
torch_device_module.synchronize()
113+
return
114+
if self.offload_granularity == "block" and self.config.get("offload_persistent_resident_blocks", False):
115+
torch_device_module.synchronize()
116+
if hasattr(self.transformer_infer.offload_manager, "reset_slots"):
117+
self.transformer_infer.offload_manager.reset_slots()
118+
return
119+
self.force_cleanup_offload_weights()
120+
121+
def force_cleanup_offload_weights(self):
122+
"""Release DiT-resident weights without copying immutable weights back to CPU."""
123+
if not self.cpu_offload or not self._offload_weights_active:
124+
return
125+
126+
torch_device_module.synchronize()
127+
if self.offload_granularity == "model":
128+
if self.config.get("qwen_image_model_offload_release_only", False):
129+
release_start = time.perf_counter()
130+
release_weight_module_device_tensors(self.pre_weight)
131+
release_weight_module_device_tensors(self.transformer_weights)
132+
release_weight_module_device_tensors(self.post_weight)
133+
torch_device_module.empty_cache()
134+
rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0
135+
logger.info(
136+
f"[QwenImage] Rank {rank}: released full DiT device replica without D2H in "
137+
f"{time.perf_counter() - release_start:.3f}s"
138+
)
139+
else:
140+
self.to_cpu()
141+
else:
142+
if hasattr(self.transformer_infer.offload_manager, "reset_slots"):
143+
self.transformer_infer.offload_manager.reset_slots()
144+
release_weight_module_device_tensors(self.pre_weight)
145+
release_weight_module_device_tensors(self.post_weight)
146+
self.transformer_weights.release_resident_blocks()
147+
self._offload_weights_active = False
148+
149+
def _keep_model_weights_resident_on_this_rank(self):
150+
return (
151+
self.offload_granularity == "model"
152+
and self.config.get("qwen_image_rank_aware_model_offload", False)
153+
and dist.is_available()
154+
and dist.is_initialized()
155+
and dist.get_world_size() > 1
156+
and dist.get_rank() != 0
157+
)
158+
52159
@torch.no_grad()
53160
def _infer_cond_uncond(self, latents_input, prompt_embeds, infer_condition=True):
54161
self.scheduler.infer_condition = infer_condition
@@ -94,7 +201,9 @@ def _seq_parallel_post_process(self, noise_pred):
94201
@torch.no_grad()
95202
def infer(self, inputs):
96203
if self.cpu_offload:
97-
if self.offload_granularity == "model" and self.scheduler.step_index == 0:
204+
if self._offload_weights_active:
205+
pass
206+
elif self.offload_granularity == "model" and self.scheduler.step_index == 0:
98207
self.to_cuda()
99208
elif self.offload_granularity != "model":
100209
self.pre_weight.to_cuda()
@@ -148,7 +257,9 @@ def infer(self, inputs):
148257
self.scheduler.noise_pred = noise_pred
149258

150259
if self.cpu_offload:
151-
if self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1:
260+
if self._offload_weights_active:
261+
pass
262+
elif self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1:
152263
self.to_cpu()
153264
elif self.offload_granularity != "model":
154265
self.pre_weight.to_cpu()

0 commit comments

Comments
 (0)