1+ import time
2+
13import torch
24import torch .distributed as dist
5+ from loguru import logger
36from torch .nn import functional as F
47
58from lightx2v .models .networks .base_model import BaseTransformerModel
912from lightx2v .models .networks .qwen_image .infer .transformer_infer import QwenImageTransformerInfer
1013from lightx2v .models .networks .qwen_image .weights .post_weights import QwenImagePostWeights
1114from 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+ )
1319from lightx2v .utils .envs import *
1420from 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
1726class 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