@@ -31,8 +31,7 @@ def __init__(self, config, model_path, device):
3131 # event offload, with block-level CPU offload, BF16 Default weights, and no LoRA,
3232 # lazy loading, or tensor parallelism.
3333 self .use_block_slab_offload = self .config .get ("offload_use_block_slab" , False )
34- self ._offload_weights_loaded = False
35- self ._offload_weights_preparing = False
34+ self ._offload_weights_active = False
3635 self .in_channels = self .config .get ("transformer_in_channels" , self .config .get ("in_channels" , 64 ))
3736 self .attention_kwargs = {}
3837 self ._combined_img_ids_cache = None
@@ -298,61 +297,43 @@ def _init_offload_manager(self):
298297 if single_slabs :
299298 self .transformer_infer .offload_manager_single .init_block_slabs (single_slabs )
300299
301- def _prepare_infer_weights (self ):
302- if not self .cpu_offload or self ._offload_weights_loaded :
300+ def prepare_offload_weights (self ):
301+ """Load weights kept resident for one runner invocation."""
302+ if not self .cpu_offload :
303303 return
304- if self ._offload_weights_preparing :
305- raise RuntimeError ("Flux2 offload weights are already being prepared " )
304+ if self ._offload_weights_active :
305+ raise RuntimeError ("Flux2 offload weights are already active " )
306306
307- self ._offload_weights_preparing = True
308- try :
309- if self .offload_granularity == "model" :
310- self .to_cuda ()
311- else :
312- # These weights are used on every diffusion step, so keep them
313- # resident for the complete denoising loop.
314- preserve_weight_module_cpu_tensors (self .pre_weight )
315- preserve_weight_module_cpu_tensors (self .post_weight )
316- self .pre_weight .to_cuda ()
317- self .post_weight .to_cuda ()
318- self .transformer_weights .non_block_weights_to_cuda ()
319- self .transformer_weights .resident_blocks_to_cuda ()
320- except BaseException as prepare_error :
321- try :
322- self .force_cleanup_offload_weights ()
323- except BaseException as cleanup_error :
324- if hasattr (prepare_error , "add_note" ):
325- prepare_error .add_note (f"Flux2 offload cleanup also failed: { cleanup_error !r} " )
326- raise
307+ # Mark the model active before moving weights so runner cleanup also
308+ # handles a partially completed preparation.
309+ self ._offload_weights_active = True
310+ if self .offload_granularity == "model" :
311+ self .to_cuda ()
327312 else :
328- self ._offload_weights_loaded = True
329- self ._offload_weights_preparing = False
330-
331- def _release_infer_weights (self ):
332- if not self .cpu_offload or self .scheduler .step_index != self .scheduler .infer_steps - 1 :
333- return
334- self .force_cleanup_offload_weights ()
335-
336- def _flush_event_offload_managers (self ):
337- if not getattr (self .transformer_infer , "use_event_offload" , False ):
338- return
339- for manager_name in ("offload_manager_double" , "offload_manager_single" ):
340- manager = getattr (self .transformer_infer , manager_name , None )
341- if manager is not None and hasattr (manager , "flush" ):
342- manager .flush ()
313+ # These weights are used on every diffusion step, so keep them
314+ # resident for the complete denoising loop.
315+ preserve_weight_module_cpu_tensors (self .pre_weight )
316+ preserve_weight_module_cpu_tensors (self .post_weight )
317+ self .pre_weight .to_cuda ()
318+ self .post_weight .to_cuda ()
319+ self .transformer_weights .non_block_weights_to_cuda ()
320+ self .transformer_weights .resident_blocks_to_cuda ()
343321
344322 def force_cleanup_offload_weights (self ):
345323 """Release loaded offload weights and reset event-slot state.
346324
347325 This method is intentionally idempotent and may be called from a
348326 runner ``finally`` block after a short run, cancellation, or error.
349327 """
350- if not self .cpu_offload or not ( self ._offload_weights_loaded or self . _offload_weights_preparing ) :
351- return False
328+ if not self .cpu_offload or not self ._offload_weights_active :
329+ return
352330
353- self ._flush_event_offload_managers ()
354331 # Device execution and non-blocking H2D copies are asynchronous.
355332 self ._sync_device ()
333+ if getattr (self .transformer_infer , "use_event_offload" , False ):
334+ self .transformer_infer .offload_manager_double .reset_slots ()
335+ self .transformer_infer .offload_manager_single .reset_slots ()
336+
356337 if self .offload_granularity == "model" :
357338 self .to_cpu ()
358339 else :
@@ -361,9 +342,7 @@ def force_cleanup_offload_weights(self):
361342 self .transformer_weights .release_non_block_weights ()
362343 self .transformer_weights .release_resident_blocks ()
363344
364- self ._offload_weights_loaded = False
365- self ._offload_weights_preparing = False
366- return True
345+ self ._offload_weights_active = False
367346
368347 def _get_combined_img_ids (self , img_ids , input_image_ids ):
369348 cached = self ._combined_img_ids_cache
@@ -459,8 +438,6 @@ def _init_infer_class(self):
459438 @compiled_method ()
460439 @torch .no_grad ()
461440 def infer (self , inputs ):
462- self ._prepare_infer_weights ()
463-
464441 latents = self .scheduler .latents
465442 do_cfg = self .config .get ("enable_cfg" , True ) and self .config .get ("sample_guide_scale" , 1.0 ) > 1.0
466443
@@ -544,8 +521,6 @@ def infer(self, inputs):
544521 )
545522 self .scheduler .noise_pred = noise_pred
546523
547- self ._release_infer_weights ()
548-
549524
550525class Flux2DevTransformerModel (_Flux2TransformerModelBase ):
551526 """Flux2 Dev transformer: single forward pass with embedded guidance (no CFG)."""
@@ -571,8 +546,6 @@ def _init_infer_class(self):
571546 @compiled_method ()
572547 @torch .no_grad ()
573548 def infer (self , inputs ):
574- self ._prepare_infer_weights ()
575-
576549 latents = self .scheduler .latents
577550 txt_ids = inputs ["text_encoder_output" ].get ("text_ids" , None )
578551 img_ids = getattr (self .scheduler , "latent_image_ids" , None )
@@ -585,5 +558,3 @@ def infer(self, inputs):
585558 img_ids = img_ids ,
586559 )
587560 self .scheduler .noise_pred = noise_pred
588-
589- self ._release_infer_weights ()
0 commit comments