Skip to content

Commit 35455b3

Browse files
committed
refactor(flux2): simplify offload lifecycle
1 parent 7264e5b commit 35455b3

3 files changed

Lines changed: 41 additions & 74 deletions

File tree

lightx2v/common/offload/event_manager.py

Lines changed: 3 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -28,13 +28,14 @@ def __init__(self, offload_granularity, load_stream=None, compute_stream=None):
2828
self.device_module = torch_device_module
2929
self._ready_events = [torch_device_module.Event() for _ in range(self._EVENT_SLOT_COUNT)]
3030
self._free_events = [torch_device_module.Event() for _ in range(self._EVENT_SLOT_COUNT)]
31-
self._reset_slot_state()
31+
self.reset_slots()
3232

3333
@property
3434
def slot_count(self):
3535
return self._EVENT_SLOT_COUNT
3636

37-
def _reset_slot_state(self):
37+
def reset_slots(self):
38+
"""Reset slot bookkeeping; synchronize pending device work first."""
3839
self._slot_pending = [False] * self.slot_count
3940
self._slot_ready_waited = [False] * self.slot_count
4041
self._slot_free_recorded = [False] * self.slot_count
@@ -117,13 +118,6 @@ def record_free(self, slot_idx, stream=None):
117118
self._slot_ready_waited[slot_idx] = False
118119
self._slot_free_recorded[slot_idx] = True
119120

120-
def flush(self):
121-
"""Drain the offload streams and reset the slot state."""
122-
self.cuda_load_stream.synchronize()
123-
if self.compute_stream is not self.cuda_load_stream:
124-
self.compute_stream.synchronize()
125-
self._reset_slot_state()
126-
127121
def init_block_slabs(self, block_slabs, staging_raw=None):
128122
"""Prepare the shared device buffer used for block-slab copies."""
129123
block_slabs = dict(block_slabs or {})

lightx2v/models/networks/flux2/model.py

Lines changed: 26 additions & 55 deletions
Original file line numberDiff line numberDiff line change
@@ -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

550525
class 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()

lightx2v/models/runners/flux2/flux2_runner.py

Lines changed: 12 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -247,7 +247,11 @@ def run(self, total_steps=None):
247247
total_steps = self.model.scheduler.infer_steps
248248
run_error = None
249249
try:
250-
for step_index in range(total_steps):
250+
step_indices = range(total_steps)
251+
if step_indices:
252+
self.model.prepare_offload_weights()
253+
254+
for step_index in step_indices:
251255
logger.info(f"==> step_index: {step_index + 1} / {total_steps}")
252256

253257
with ProfilingContext4DebugL1("step_pre"):
@@ -267,15 +271,13 @@ def run(self, total_steps=None):
267271
run_error = caught_error
268272
raise
269273
finally:
270-
cleanup = getattr(self.model, "force_cleanup_offload_weights", None)
271-
if cleanup is not None:
272-
try:
273-
cleanup()
274-
except BaseException as cleanup_error:
275-
if run_error is None:
276-
raise
277-
if hasattr(run_error, "add_note"):
278-
run_error.add_note(f"Flux2 offload cleanup also failed: {cleanup_error!r}")
274+
try:
275+
self.model.force_cleanup_offload_weights()
276+
except BaseException as cleanup_error:
277+
if run_error is None:
278+
raise
279+
if hasattr(run_error, "add_note"):
280+
run_error.add_note(f"Flux2 offload cleanup also failed: {cleanup_error!r}")
279281

280282
def get_custom_shape(self):
281283
default_aspect_ratios = {

0 commit comments

Comments
 (0)