From 6676c54dbdadae337c5908643f670f3f48c7e4c5 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:37:13 +0800 Subject: [PATCH 01/31] refactor(offload): centralize resident block scheduling --- lightx2v/common/modules/weight_module.py | 100 +++ lightx2v/common/offload/config.py | 29 + lightx2v/common/offload/event_manager.py | 31 +- lightx2v/common/offload/manager.py | 40 +- .../transformer_infer/transformer_infer.py | 251 +++++++ lightx2v/disagg/utils.py | 8 + lightx2v/models/networks/base_model.py | 58 +- lightx2v/models/runners/base_runner.py | 24 + lightx2v/models/runners/default_runner.py | 8 +- lightx2v/pipeline.py | 8 +- lightx2v/utils/set_config.py | 2 + tests/common/offload/test_block_residency.py | 129 ++++ tests/common/offload/test_config.py | 47 ++ tests/common/offload/test_offload_schedule.py | 667 ++++++++++++++++++ 14 files changed, 1377 insertions(+), 25 deletions(-) create mode 100644 lightx2v/common/offload/config.py create mode 100644 tests/common/offload/test_block_residency.py create mode 100644 tests/common/offload/test_config.py create mode 100644 tests/common/offload/test_offload_schedule.py diff --git a/lightx2v/common/modules/weight_module.py b/lightx2v/common/modules/weight_module.py index 4a00ff539..1c6296b13 100755 --- a/lightx2v/common/modules/weight_module.py +++ b/lightx2v/common/modules/weight_module.py @@ -1,6 +1,18 @@ +from loguru import logger + +from lightx2v.common.offload.config import get_offload_granularity, get_offload_plan from lightx2v_platform.base.global_var import AI_DEVICE +def resolve_resident_block_indices(resident_count, blocks_num): + count = blocks_num if resident_count == "all" else resident_count + if not 0 <= count <= blocks_num: + raise ValueError(f"resident block count must be between 0 and {blocks_num}, got {count}") + if count == 0: + return frozenset() + return frozenset((index * blocks_num) // count for index in range(count)) + + class WeightModule: def __init__(self): self._modules = {} @@ -187,10 +199,98 @@ def to_cuda_async(self, non_blocking=True): if module is not None and hasattr(module, "to_cuda"): module.to_cuda(non_blocking=True) + def release_device_weights(self): + self.to_cpu() + + def release_non_block_weights(self): + block_groups = tuple(self.get_offload_block_groups().values()) + excluded_modules = {id(blocks) for blocks in block_groups} + for blocks in block_groups: + for name in ("offload_cuda_buffers", "offload_cpu_buffers"): + buffers = getattr(blocks, name, None) + if buffers is not None: + excluded_modules.add(id(buffers)) + + for module in self._modules.values(): + if module is not None and id(module) not in excluded_modules and hasattr(module, "to_cpu"): + module.to_cpu() + for name, parameter in self._parameters.items(): + if parameter is None or id(parameter) in excluded_modules: + continue + if hasattr(parameter, "cpu"): + self._parameters[name] = parameter.to("cpu") + setattr(self, name, self._parameters[name]) + elif hasattr(parameter, "to_cpu"): + parameter.to_cpu() + + def register_offload_block_group(self, config, name, blocks): + granularity = get_offload_granularity(config) + if not config.get("cpu_offload", False): + resident_count = "all" + elif granularity == "block": + resident_count = get_offload_plan(config).get("resident_blocks", {}).get(name, 0) + else: + resident_count = 0 + + resident_indices = resolve_resident_block_indices(resident_count, len(blocks)) + blocks.offload_group_name = name + blocks.resident_block_indices = resident_indices + blocks.offload_block_indices = tuple(index for index in range(len(blocks)) if index not in resident_indices) + + if not hasattr(self, "_offload_block_groups"): + self._offload_block_groups = {} + self._offload_block_groups[name] = blocks + slot_count = min(2, len(blocks.offload_block_indices)) if granularity == "block" and config.get("cpu_offload", False) else 0 + if config.get("cpu_offload", False) and granularity == "block": + logger.info( + "Block offload group '{}': resident={}/{}, indices={}, staging_slots={}", + name, + len(resident_indices), + len(blocks), + tuple(sorted(resident_indices)), + slot_count, + ) + return slot_count + + def register_offload_block_buffers(self, name, cuda_buffers, cpu_buffers=None): + blocks = self._offload_block_groups[name] + blocks.offload_cuda_buffers = cuda_buffers + blocks.offload_cpu_buffers = cpu_buffers + + def validate_offload_block_groups(self, config): + if not config.get("cpu_offload", False) or get_offload_granularity(config) != "block": + return + + plan = get_offload_plan(config) + configured_groups = set(plan.get("resident_blocks", {})) + registered_groups = set(self.get_offload_block_groups()) + unknown_groups = configured_groups - registered_groups + if unknown_groups: + raise ValueError(f"Unknown resident block groups {sorted(unknown_groups)}; available groups are {sorted(registered_groups)}") + + def get_offload_block_groups(self): + return getattr(self, "_offload_block_groups", {}) + + def resident_blocks_to_cuda(self, non_blocking=True): + for blocks in self.get_offload_block_groups().values(): + for block_index in blocks.resident_block_indices: + blocks[block_index].to_cuda(non_blocking=non_blocking) + + def release_resident_blocks(self): + for blocks in self.get_offload_block_groups().values(): + for block_index in blocks.resident_block_indices: + blocks[block_index].release_device_weights() + class WeightModuleList(WeightModule): def __init__(self, modules=None): super().__init__() + # Transformer block lists receive these values when registered for offload. + self.offload_group_name = None + self.resident_block_indices = frozenset() + self.offload_block_indices = () + self.offload_cuda_buffers = None + self.offload_cpu_buffers = None self._list = [] if modules is not None: for idx, module in enumerate(modules): diff --git a/lightx2v/common/offload/config.py b/lightx2v/common/offload/config.py new file mode 100644 index 000000000..d94ba6832 --- /dev/null +++ b/lightx2v/common/offload/config.py @@ -0,0 +1,29 @@ +def get_offload_plan(config): + plan = config.get("offload_plan") + if plan is not None: + return plan + + return { + "offload_granularity": config.get("offload_granularity", "block"), + "use_event_offload": config.get("use_event_offload", False), + "resident_blocks": {}, + } + + +def get_offload_granularity(config): + return get_offload_plan(config).get("offload_granularity", "block") + + +def use_event_offload(config): + return get_offload_plan(config).get("use_event_offload", False) + + +def normalize_offload_plan(config): + if "offload_plan" in config: + plan = dict(config["offload_plan"]) + else: + plan = get_offload_plan(config) + plan.setdefault("offload_granularity", "block") + plan.setdefault("use_event_offload", False) + plan.setdefault("resident_blocks", {}) + config["offload_plan"] = plan diff --git a/lightx2v/common/offload/event_manager.py b/lightx2v/common/offload/event_manager.py index 9e0a9a352..4420f9f97 100644 --- a/lightx2v/common/offload/event_manager.py +++ b/lightx2v/common/offload/event_manager.py @@ -10,9 +10,9 @@ class EventSlotWeightAsyncStreamManager(WeightAsyncStreamManager): """Weight offload with reusable buffers protected by device events.""" - _EVENT_SLOT_COUNT = 2 + uses_events = True - def __init__(self, offload_granularity, load_stream=None, compute_stream=None): + def __init__(self, offload_granularity, slot_count=2, load_stream=None, compute_stream=None): if offload_granularity != "block": raise ValueError("Event-slot weight offload only supports block granularity") @@ -22,17 +22,24 @@ def __init__(self, offload_granularity, load_stream=None, compute_stream=None): if compute_stream is not None: self.compute_stream = compute_stream - self.device_module = torch_device_module - self._ready_events = [torch_device_module.Event() for _ in range(self._EVENT_SLOT_COUNT)] - self._free_events = [torch_device_module.Event() for _ in range(self._EVENT_SLOT_COUNT)] + self._slot_count = slot_count + self._ready_events = [] + self._free_events = [] + self.reset_slots() + + def init_cuda_buffer(self, blocks_cuda_buffer=None, phases_cuda_buffer=None): + super().init_cuda_buffer(blocks_cuda_buffer, phases_cuda_buffer) + self._slot_count = len(self.cuda_buffers) + self._ready_events = [torch_device_module.Event() for _ in range(self.slot_count)] + self._free_events = [torch_device_module.Event() for _ in range(self.slot_count)] self.reset_slots() @property def slot_count(self): - return self._EVENT_SLOT_COUNT + return self._slot_count def reset_slots(self): - """Reset slot bookkeeping; synchronize pending device work first.""" + """Reset slot bookkeeping after pending device work has completed.""" self._slot_pending = [False] * self.slot_count self._slot_ready_waited = [False] * self.slot_count self._slot_free_recorded = [False] * self.slot_count @@ -45,7 +52,7 @@ def _validate_slot(self, slot_idx): if len(self.cuda_buffers) < self.slot_count: raise RuntimeError(f"Event-slot weight offload requires {self.slot_count} device buffers") - def _load_block_to_buffer(self, target_buffer, block_idx, blocks, adapter_block_idx): + def _load_block_to_buffer(self, target_buffer, block_idx, blocks, adapter_block_idx, state_dict_transform): block_slab = getattr(self, "block_slabs", {}).get(block_idx) if block_slab is not None: copy_block_slab_( @@ -67,9 +74,12 @@ def _load_block_to_buffer(self, target_buffer, block_idx, blocks, adapter_block_ if blocks is None: raise ValueError("blocks must be provided when CPU buffers have not been initialized") source = blocks[block_idx] - target_buffer.load_state_dict(source.state_dict(), block_idx, adapter_block_idx) + state_dict = source.state_dict() + if state_dict_transform is not None: + state_dict = state_dict_transform(block_idx, state_dict) + target_buffer.load_state_dict(state_dict, block_idx, adapter_block_idx) - def prefetch_to_slot(self, slot_idx, block_idx, blocks=None, adapter_block_idx=None): + def prefetch_to_slot(self, slot_idx, block_idx, blocks=None, adapter_block_idx=None, state_dict_transform=None): """Enqueue one block copy into a fixed staging slot.""" self._validate_slot(slot_idx) if self._slot_pending[slot_idx]: @@ -83,6 +93,7 @@ def prefetch_to_slot(self, slot_idx, block_idx, blocks=None, adapter_block_idx=N block_idx, blocks, adapter_block_idx, + state_dict_transform, ) self._ready_events[slot_idx].record(self.cuda_load_stream) diff --git a/lightx2v/common/offload/manager.py b/lightx2v/common/offload/manager.py index 3fb476ad6..0f97021be 100755 --- a/lightx2v/common/offload/manager.py +++ b/lightx2v/common/offload/manager.py @@ -12,10 +12,14 @@ class WeightAsyncStreamManager(object): + uses_events = False + def __init__(self, offload_granularity): self.offload_granularity = offload_granularity self.init_stream = torch_device_module.Stream(priority=0) self.need_init_first_buffer = True + self.loaded_state_dict_transform = None + self.loaded_first_adapter_block_index = None self.lazy_load = False torch_version = parse(torch.__version__.split("+")[0]) # Legacy name: this is the active device backend's weight-loading stream, not a CUDA-only stream. @@ -60,27 +64,34 @@ def _sync(self): else: self.init_stream.synchronize() - def init_first_buffer(self, blocks, adapter_block_idx=None): + def init_first_buffer(self, blocks, adapter_block_idx=None, block_idx=0, state_dict_transform=None): with torch_device_module.stream(self.init_stream): if hasattr(self, "cpu_buffers"): if self.offload_granularity == "block": - self.cuda_buffers[0].load_state_dict(self.cpu_buffers[0].state_dict(), 0, adapter_block_idx) + state_dict = self.cpu_buffers[0].state_dict() else: - self.cuda_buffers[0].load_state_dict(self.cpu_buffers[0][0].state_dict(), 0, adapter_block_idx) + state_dict = self.cpu_buffers[0][0].state_dict() else: if self.offload_granularity == "block": - self.cuda_buffers[0].load_state_dict(blocks[0].state_dict(), 0, adapter_block_idx) + state_dict = blocks[block_idx].state_dict() else: - self.cuda_buffers[0].load_state_dict(blocks[0].compute_phases[0].state_dict(), 0, adapter_block_idx) + state_dict = blocks[block_idx].compute_phases[0].state_dict() + if state_dict_transform is not None: + state_dict = state_dict_transform(block_idx, state_dict) + self.cuda_buffers[0].load_state_dict(state_dict, block_idx, adapter_block_idx) self._sync() self.need_init_first_buffer = False def prefetch_weights(self, block_idx, blocks, adapter_block_idx=None): + self.prefetch_weights_to_buffer(1, block_idx, blocks, adapter_block_idx) + + def prefetch_weights_to_buffer(self, buffer_idx, block_idx, blocks, adapter_block_idx=None, state_dict_transform=None): with torch_device_module.stream(self.cuda_load_stream): - if hasattr(self, "cpu_buffers"): - self.cuda_buffers[1].load_state_dict(self.cpu_buffers[0].state_dict(), block_idx, adapter_block_idx) - else: - self.cuda_buffers[1].load_state_dict(blocks[block_idx].state_dict(), block_idx, adapter_block_idx) + source = self.cpu_buffers[0] if hasattr(self, "cpu_buffers") else blocks[block_idx] + state_dict = source.state_dict() + if state_dict_transform is not None: + state_dict = state_dict_transform(block_idx, state_dict) + self.cuda_buffers[buffer_idx].load_state_dict(state_dict, block_idx, adapter_block_idx) def prefetch_phase(self, block_idx, phase_idx, blocks, adapter_block_idx=None): with torch_device_module.stream(self.cuda_load_stream): @@ -89,17 +100,26 @@ def prefetch_phase(self, block_idx, phase_idx, blocks, adapter_block_idx=None): else: self.cuda_buffers[phase_idx].load_state_dict(blocks[block_idx].compute_phases[phase_idx].state_dict(), block_idx, adapter_block_idx) - def swap_blocks(self): + def synchronize_block_streams(self): if AI_DEVICE == "xpu": torch_device_module.synchronize() else: self.cuda_load_stream.synchronize() self.compute_stream.synchronize() + + def swap_blocks(self): + self.synchronize_block_streams() self.cuda_buffers[0], self.cuda_buffers[1] = ( self.cuda_buffers[1], self.cuda_buffers[0], ) + def wait_for_block_compute(self): + if AI_DEVICE == "xpu": + torch_device_module.synchronize() + else: + self.compute_stream.synchronize() + def swap_phases(self): if AI_DEVICE == "xpu": torch_device_module.synchronize() diff --git a/lightx2v/common/transformer_infer/transformer_infer.py b/lightx2v/common/transformer_infer/transformer_infer.py index cff6b7d44..76a9b11fa 100644 --- a/lightx2v/common/transformer_infer/transformer_infer.py +++ b/lightx2v/common/transformer_infer/transformer_infer.py @@ -4,8 +4,56 @@ import torch from loguru import logger +from lightx2v.common.offload.config import get_offload_granularity, use_event_offload +from lightx2v.common.offload.event_manager import EventSlotWeightAsyncStreamManager +from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v_platform.base.global_var import AI_DEVICE + +torch_device_module = getattr(torch, AI_DEVICE) + class BaseTransformerInfer(ABC): + @staticmethod + def create_block_offload_manager(config, **kwargs): + manager_class = EventSlotWeightAsyncStreamManager if use_event_offload(config) else WeightAsyncStreamManager + return manager_class(offload_granularity="block", **kwargs) + + def init_block_offload(self, config, transformer_weights): + self._block_offload_managers = {} + shared_streams = {} + block_offload = config.get("cpu_offload", False) and get_offload_granularity(config) == "block" + + for name, blocks in transformer_weights.get_offload_block_groups().items(): + if not block_offload or not blocks.offload_block_indices: + self._block_offload_managers[name] = None + continue + + cuda_buffers = getattr(blocks, "offload_cuda_buffers", None) + if cuda_buffers is None: + raise RuntimeError(f"Block offload group {name!r} has no staging buffers") + expected_buffers = min(2, len(blocks.offload_block_indices)) + if len(cuda_buffers) != expected_buffers: + raise RuntimeError(f"Block offload group {name!r} requires {expected_buffers} staging buffers, got {len(cuda_buffers)}") + + manager = self.create_block_offload_manager(config, **shared_streams) + manager.init_cuda_buffer(cuda_buffers) + if config.get("lazy_load", False): + manager.init_cpu_buffer(blocks.offload_cpu_buffers) + manager.init_lazy_load(config.get("num_disk_workers", 4)) + self._block_offload_managers[name] = manager + + if manager.uses_events and not shared_streams: + shared_streams = { + "load_stream": manager.cuda_load_stream, + "compute_stream": manager.compute_stream, + } + + def get_block_offload_manager(self, blocks): + return self._block_offload_managers[blocks.offload_group_name] + + def has_block_offload_manager(self): + return any(manager is not None for manager in getattr(self, "_block_offload_managers", {}).values()) + def init_compile(self, config): self.use_compile = config.get("use_compile", False) self.compiled_blocks = {} @@ -33,6 +81,209 @@ def run_block(self, block_idx, block, *args): return self.get_compiled_block(block_idx, block)(*args) return self.infer_block(block, *args) + @staticmethod + def _adapter_block_index(adapter_block_index, block_index): + return adapter_block_index(block_index) if adapter_block_index is not None else None + + @staticmethod + def _run_offload_block(manager, run_block, block_index, block): + if AI_DEVICE == "xpu": + return run_block(block_index, block) + with torch_device_module.stream(manager.compute_stream): + return run_block(block_index, block) + + @staticmethod + def _record_block_output_stream(block_output, stream): + if isinstance(block_output, torch.Tensor): + block_output.record_stream(stream) + elif isinstance(block_output, (tuple, list)): + for value in block_output: + BaseTransformerInfer._record_block_output_stream(value, stream) + + def _finish_offload_group(self, manager, caller_stream, block_output): + if AI_DEVICE == "xpu": + return + with torch_device_module.stream(manager.compute_stream): + done = manager.compute_stream.record_event() + caller_stream.wait_event(done) + self._record_block_output_stream(block_output, caller_stream) + + def _run_blocks_with_stream_offload( + self, + manager, + blocks, + run_block, + adapter_block_index=None, + state_dict_transform=None, + ): + offloaded_indices = blocks.offload_block_indices + if not offloaded_indices: + block_output = None + for block_index, block in enumerate(blocks): + block_output = run_block(block_index, block) + return block_output + + caller_stream = torch_device_module.current_stream() + if AI_DEVICE != "xpu": + manager.compute_stream.wait_stream(caller_stream) + + original_buffers = tuple(manager.cuda_buffers) + first_index = offloaded_indices[0] + first_adapter_index = self._adapter_block_index(adapter_block_index, first_index) + block_output = None + try: + if manager.loaded_state_dict_transform is not state_dict_transform or manager.loaded_first_adapter_block_index != first_adapter_index: + manager.need_init_first_buffer = True + if manager.need_init_first_buffer: + manager.init_first_buffer( + blocks, + first_adapter_index, + block_idx=first_index, + state_dict_transform=state_dict_transform, + ) + + offloaded_position = 0 + for block_index, block in enumerate(blocks): + if block_index in blocks.resident_block_indices: + block_output = self._run_offload_block(manager, run_block, block_index, block) + continue + + if len(offloaded_indices) > 1: + next_position = (offloaded_position + 1) % len(offloaded_indices) + next_index = offloaded_indices[next_position] + manager.prefetch_weights_to_buffer( + 1, + next_index, + blocks, + self._adapter_block_index(adapter_block_index, next_index), + state_dict_transform, + ) + + block_output = self._run_offload_block(manager, run_block, block_index, manager.cuda_buffers[0]) + if len(offloaded_indices) > 1: + manager.swap_blocks() + else: + manager.wait_for_block_compute() + offloaded_position += 1 + manager.loaded_state_dict_transform = state_dict_transform + manager.loaded_first_adapter_block_index = first_adapter_index + except BaseException: + torch_device_module.synchronize() + manager.need_init_first_buffer = True + manager.loaded_state_dict_transform = None + manager.loaded_first_adapter_block_index = None + manager.cuda_buffers[:] = original_buffers + raise + + self._finish_offload_group(manager, caller_stream, block_output) + return block_output + + def _run_blocks_with_event_offload( + self, + manager, + blocks, + run_block, + adapter_block_index=None, + state_dict_transform=None, + ): + offloaded_indices = blocks.offload_block_indices + if not offloaded_indices: + block_output = None + for block_index, block in enumerate(blocks): + block_output = run_block(block_index, block) + return block_output + + caller_stream = torch_device_module.current_stream() + compute_stream = caller_stream if AI_DEVICE == "xpu" else manager.compute_stream + if AI_DEVICE != "xpu": + compute_stream.wait_stream(caller_stream) + scheduled_slots = {} + next_position = 0 + block_output = None + + def prefetch_next(slot_index): + nonlocal next_position + if next_position == len(offloaded_indices): + return + block_index = offloaded_indices[next_position] + manager.prefetch_to_slot( + slot_index, + block_index, + blocks, + self._adapter_block_index(adapter_block_index, block_index), + state_dict_transform, + ) + scheduled_slots[block_index] = slot_index + next_position += 1 + + try: + for slot_index in range(manager.slot_count): + prefetch_next(slot_index) + + for block_index, block in enumerate(blocks): + if block_index in blocks.resident_block_indices: + block_output = self._run_offload_block(manager, run_block, block_index, block) + continue + + slot_index = scheduled_slots.pop(block_index) + staged_block = manager.wait_ready(slot_index, compute_stream) + block_output = self._run_offload_block(manager, run_block, block_index, staged_block) + manager.record_free(slot_index, compute_stream) + prefetch_next(slot_index) + except BaseException: + torch_device_module.synchronize() + manager.reset_slots() + raise + + self._finish_offload_group(manager, caller_stream, block_output) + return block_output + + def run_blocks_with_offload( + self, + blocks, + run_block, + adapter_block_index=None, + state_dict_transform=None, + ): + """Run blocks in model order, staging only the non-resident weights. + + ``run_block`` returns the tensor state that leaves the block group so + its lifetime can be transferred back to the caller stream. + """ + manager = self.get_block_offload_manager(blocks) + if manager is None: + block_output = None + for block_index, block in enumerate(blocks): + block_output = run_block(block_index, block) + return block_output + if manager.uses_events: + return self._run_blocks_with_event_offload( + manager, + blocks, + run_block, + adapter_block_index, + state_dict_transform, + ) + return self._run_blocks_with_stream_offload( + manager, + blocks, + run_block, + adapter_block_index, + state_dict_transform, + ) + + def get_offload_managers(self): + managers = [manager for manager in getattr(self, "_block_offload_managers", {}).values() if manager is not None] + manager = getattr(self, "offload_manager", None) + if manager is not None and manager not in managers: + managers.append(manager) + return managers + + def clear_offload_managers(self): + self._block_offload_managers = {} + if hasattr(self, "offload_manager"): + del self.offload_manager + @abstractmethod def infer(self): pass diff --git a/lightx2v/disagg/utils.py b/lightx2v/disagg/utils.py index d596fa5db..1a6929318 100644 --- a/lightx2v/disagg/utils.py +++ b/lightx2v/disagg/utils.py @@ -57,6 +57,8 @@ def set_config( text_encoder_offload=False, image_encoder_offload=False, vae_offload=False, + resident_blocks=None, + use_event_offload=False, **kwargs, ): """ @@ -125,6 +127,12 @@ def set_config( args_dict["norm_modulate_backend"] = norm_modulate_backend args_dict.update(kwargs) + if "offload_plan" not in args_dict: + args_dict["offload_plan"] = { + "offload_granularity": args_dict.pop("offload_granularity", "block"), + "resident_blocks": {} if resident_blocks is None else resident_blocks, + "use_event_offload": args_dict.pop("use_event_offload", use_event_offload), + } # Convert to object for set_config compatibility args = ConfigObj(**args_dict) diff --git a/lightx2v/models/networks/base_model.py b/lightx2v/models/networks/base_model.py index 8bab34a75..8a737ad2c 100644 --- a/lightx2v/models/networks/base_model.py +++ b/lightx2v/models/networks/base_model.py @@ -21,6 +21,7 @@ from loguru import logger from safetensors import safe_open +from lightx2v.common.offload.config import get_offload_granularity, get_offload_plan from lightx2v.utils.envs import * from lightx2v.utils.ggml_tensor import load_gguf_sd_ckpt from lightx2v.utils.utils import * @@ -80,12 +81,17 @@ def __init__(self, model_path, config, device, model_type=None, lora_path=None, self.config = config self.cpu_offload = self.config.get("cpu_offload", False) - self.offload_granularity = self.config.get("offload_granularity", "block") + self.offload_granularity = get_offload_granularity(self.config) + self.lazy_load = self.config.get("lazy_load", False) + if self.lazy_load and self.cpu_offload and self.offload_granularity == "block": + offload_plan = get_offload_plan(self.config) + if offload_plan.get("use_event_offload", False) or any(offload_plan.get("resident_blocks", {}).values()): + raise NotImplementedError("lazy_load does not support resident blocks or event offload") + self._offload_weights_active = False if self.config["seq_parallel"]: self.seq_p_group = self.config.get("device_mesh").get_group(mesh_dim="seq_p") else: self.seq_p_group = None - self.lazy_load = self.config.get("lazy_load", False) self.clean_cuda_cache = self.config.get("clean_cuda_cache", False) self.dit_quantized = self.config.get("dit_quantized", False) if self.dit_quantized: @@ -267,6 +273,7 @@ def _init_weights(self, weight_dict=None): self.transformer_weights = self.transformer_weight_class(self.config, self.lazy_load_path, self.lora_path) else: self.transformer_weights = self.transformer_weight_class(self.config) + self.transformer_weights.validate_offload_block_groups(self.config) if hasattr(self, "post_weight_class") and self.post_weight_class is not None: self.post_weight = self.post_weight_class(self.config) @@ -282,6 +289,17 @@ def _init_infer(self): pass def _init_offload_manager(self): + block_groups = self.transformer_weights.get_offload_block_groups() + if block_groups: + self.transformer_infer.init_block_offload(self.config, self.transformer_weights) + + if ( + not self.cpu_offload + or self.offload_granularity == "model" + or (block_groups and self.offload_granularity == "block") + ): + return + self.transformer_infer.offload_manager.init_cuda_buffer(self.transformer_weights.offload_block_cuda_buffers, self.transformer_weights.offload_phase_cuda_buffers) if self.lazy_load: self.transformer_infer.offload_manager.init_cpu_buffer(self.transformer_weights.offload_block_cpu_buffers, self.transformer_weights.offload_phase_cpu_buffers) @@ -696,6 +714,42 @@ def to_cuda(self): if hasattr(self, "post_weight"): self.post_weight.to_cuda() + def prepare_offload_weights(self): + if not self.cpu_offload or self.offload_granularity != "block" or not self.transformer_weights.get_offload_block_groups() or self._offload_weights_active: + return False + + self._offload_weights_active = True + try: + self.pre_weight.to_cuda() + if hasattr(self, "post_weight"): + self.post_weight.to_cuda() + if hasattr(self.transformer_weights, "non_block_weights_to_cuda"): + self.transformer_weights.non_block_weights_to_cuda() + self.transformer_weights.resident_blocks_to_cuda() + except BaseException: + self.cleanup_offload_weights() + raise + return True + + def cleanup_offload_weights(self): + if not self._offload_weights_active: + return + + device_module = getattr(torch, AI_DEVICE) + device_module.synchronize() + for manager in self.transformer_infer.get_offload_managers(): + manager.need_init_first_buffer = True + if manager.uses_events: + manager.reset_slots() + + self.pre_weight.release_device_weights() + if hasattr(self, "post_weight"): + self.post_weight.release_device_weights() + self.transformer_weights.release_non_block_weights() + self.transformer_weights.release_resident_blocks() + self._offload_weights_active = False + device_module.empty_cache() + @abstractmethod @torch.no_grad() def infer(self, inputs): diff --git a/lightx2v/models/runners/base_runner.py b/lightx2v/models/runners/base_runner.py index 057090c9a..73ab0929f 100755 --- a/lightx2v/models/runners/base_runner.py +++ b/lightx2v/models/runners/base_runner.py @@ -2,6 +2,7 @@ import gc import os from abc import ABC +from contextlib import contextmanager import torch import torch.distributed as dist @@ -10,6 +11,15 @@ from lightx2v_platform.base.global_var import AI_DEVICE +def keep_transformer_weights_loaded(run): + @functools.wraps(run) + def wrapped(self, *args, **kwargs): + with self.transformer_offload_session(): + return run(self, *args, **kwargs) + + return wrapped + + class BaseRunner(ABC): """Abstract base class for all Runners @@ -68,6 +78,20 @@ def warmup(self): if self.config.get("warmup", False): raise NotImplementedError(f"Warmup is not supported for {type(self).__name__}") + @contextmanager + def transformer_offload_session(self, model=None): + model = getattr(self, "model", None) if model is None else model + if model is None: + yield + return + prepare = getattr(model, "prepare_offload_weights", None) + owner = prepare() if prepare is not None else False + try: + yield + finally: + if owner: + model.cleanup_offload_weights() + def set_reuse(self, reuse, reuse_prefix_segments=0): if reuse and not self.enable_reuse: raise ValueError(f"This {type(self).__name__} service does not enable reuse") diff --git a/lightx2v/models/runners/default_runner.py b/lightx2v/models/runners/default_runner.py index 2d5ff4c0f..872d62b2a 100755 --- a/lightx2v/models/runners/default_runner.py +++ b/lightx2v/models/runners/default_runner.py @@ -10,7 +10,7 @@ from PIL import Image from loguru import logger -from lightx2v.models.runners.base_runner import BaseRunner +from lightx2v.models.runners.base_runner import BaseRunner, keep_transformer_weights_loaded from lightx2v.server.metrics import monitor_cli from lightx2v.utils.envs import * from lightx2v.utils.global_paras import CALIB @@ -315,6 +315,7 @@ def set_config(self, config_modify): def set_progress_callback(self, callback): self.progress_callback = callback + @keep_transformer_weights_loaded def run_segment(self, segment_idx=0): infer_steps = self.model.scheduler.infer_steps @@ -390,7 +391,10 @@ def end_run(self): self.scheduler.transformer_infer = None models = self.model.model if hasattr(self.model, "model") and len(self.model.model) == 2 else (self.model,) for model in filter(None, models): - if hasattr(model.transformer_infer, "offload_manager"): + clear_offload_managers = getattr(model.transformer_infer, "clear_offload_managers", None) + if clear_offload_managers is not None: + clear_offload_managers() + elif hasattr(model.transformer_infer, "offload_manager"): del model.transformer_infer.offload_manager self.model = None models = model = None diff --git a/lightx2v/pipeline.py b/lightx2v/pipeline.py index ddb0eb46e..c68d9b674 100755 --- a/lightx2v/pipeline.py +++ b/lightx2v/pipeline.py @@ -356,9 +356,15 @@ def enable_offload( text_encoder_offload=False, image_encoder_offload=False, vae_offload=False, + resident_blocks=None, + use_event_offload=False, ): self.cpu_offload = cpu_offload - self.offload_granularity = offload_granularity + self.offload_plan = { + "offload_granularity": offload_granularity, + "resident_blocks": {} if resident_blocks is None else resident_blocks, + "use_event_offload": use_event_offload, + } self.vae_cpu_offload = vae_offload if self.model_cls in [ "wan2.1", diff --git a/lightx2v/utils/set_config.py b/lightx2v/utils/set_config.py index b1c9482ad..4bf6ff34d 100755 --- a/lightx2v/utils/set_config.py +++ b/lightx2v/utils/set_config.py @@ -6,6 +6,7 @@ from loguru import logger from torch.distributed.tensor.device_mesh import init_device_mesh +from lightx2v.common.offload.config import normalize_offload_plan from lightx2v.utils.input_info import ALL_INPUT_INFO_KEYS from lightx2v.utils.lockable_dict import LockableDict from lightx2v.utils.utils import find_torch_model_path, is_main_process @@ -404,6 +405,7 @@ def set_config(args): validate_model_task_args(args) config = set_args2config(args) config = auto_calc_config(config) + normalize_offload_plan(config) return config diff --git a/tests/common/offload/test_block_residency.py b/tests/common/offload/test_block_residency.py new file mode 100644 index 000000000..24411b217 --- /dev/null +++ b/tests/common/offload/test_block_residency.py @@ -0,0 +1,129 @@ +import pytest +import torch + +from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList, resolve_resident_block_indices + + +def _blocks(count): + return WeightModuleList([WeightModule() for _ in range(count)]) + + +@pytest.mark.parametrize( + ("resident_count", "blocks_num", "expected"), + [ + (0, 10, frozenset()), + (3, 10, frozenset({0, 3, 6})), + ("all", 10, frozenset(range(10))), + ], +) +def test_resident_block_indices(resident_count, blocks_num, expected): + assert resolve_resident_block_indices(resident_count, blocks_num) == expected + + +@pytest.mark.parametrize("resident_count", [-1, 11]) +def test_resident_block_count_out_of_range(resident_count): + with pytest.raises(ValueError): + resolve_resident_block_indices(resident_count, 10) + + +def test_no_offload_means_all_blocks_are_resident(): + weights = WeightModule() + blocks = _blocks(6) + + slot_count = weights.register_offload_block_group({"cpu_offload": False}, "blocks", blocks) + + assert blocks.resident_block_indices == frozenset(range(6)) + assert blocks.offload_block_indices == () + assert slot_count == 0 + + +def test_model_offload_has_no_resident_blocks(): + weights = WeightModule() + blocks = _blocks(6) + config = { + "cpu_offload": True, + "offload_plan": { + "offload_granularity": "model", + "resident_blocks": {"blocks": 4}, + }, + } + + slot_count = weights.register_offload_block_group(config, "blocks", blocks) + + assert blocks.resident_block_indices == frozenset() + assert blocks.offload_block_indices == tuple(range(6)) + assert slot_count == 0 + + +def test_unknown_resident_block_group_is_rejected(): + weights = WeightModule() + weights.register_offload_block_group( + {"cpu_offload": True, "offload_plan": {"offload_granularity": "block"}}, + "blocks", + _blocks(6), + ) + + with pytest.raises(ValueError, match="missing_blocks"): + weights.validate_offload_block_groups( + { + "cpu_offload": True, + "offload_plan": { + "offload_granularity": "block", + "resident_blocks": {"missing_blocks": 2}, + }, + } + ) + + +def test_legacy_block_offload_without_residency_does_not_require_a_group(): + weights = WeightModule() + + weights.validate_offload_block_groups( + { + "cpu_offload": True, + "offload_plan": {"offload_granularity": "block"}, + } + ) + + +@pytest.mark.parametrize( + ("resident_count", "expected_resident", "expected_offloaded", "expected_slots"), + [ + (2, frozenset({0, 3}), (1, 2, 4, 5), 2), + (5, frozenset({0, 1, 2, 3, 4}), (5,), 1), + ("all", frozenset(range(6)), (), 0), + ], +) +def test_block_offload_group_metadata(resident_count, expected_resident, expected_offloaded, expected_slots): + weights = WeightModule() + blocks = _blocks(6) + config = { + "cpu_offload": True, + "offload_plan": { + "offload_granularity": "block", + "resident_blocks": {"blocks": resident_count}, + }, + } + + slot_count = weights.register_offload_block_group(config, "blocks", blocks) + + assert blocks.resident_block_indices == expected_resident + assert blocks.offload_block_indices == expected_offloaded + assert slot_count == expected_slots + + +def test_release_uses_the_weight_objects_cpu_contract(): + class AttentionBackend: + def __init__(self): + self.runtime_cache = torch.empty(1, device="meta") + + def to_cpu(self): + pass + + weights = WeightModule() + attention = AttentionBackend() + weights.add_module("attention", attention) + + weights.release_device_weights() + + assert attention.runtime_cache.device.type == "meta" diff --git a/tests/common/offload/test_config.py b/tests/common/offload/test_config.py new file mode 100644 index 000000000..3abe7c529 --- /dev/null +++ b/tests/common/offload/test_config.py @@ -0,0 +1,47 @@ +from lightx2v.common.offload.config import get_offload_granularity, normalize_offload_plan, use_event_offload + + +def test_legacy_offload_keys_are_normalized_into_one_plan(): + config = { + "offload_granularity": "model", + "use_event_offload": True, + } + + normalize_offload_plan(config) + + assert config["offload_plan"] == { + "offload_granularity": "model", + "resident_blocks": {}, + "use_event_offload": True, + } + + +def test_explicit_offload_plan_takes_precedence(): + config = { + "offload_granularity": "model", + "use_event_offload": False, + "offload_plan": { + "offload_granularity": "block", + "resident_blocks": {"blocks": 8}, + "use_event_offload": True, + }, + } + + normalize_offload_plan(config) + + assert get_offload_granularity(config) == "block" + assert use_event_offload(config) is True + assert config["offload_plan"]["resident_blocks"] == {"blocks": 8} + + +def test_offload_plan_defaults_are_added_without_changing_model_specific_keys(): + config = {"offload_plan": {"use_block_slab": True}} + + normalize_offload_plan(config) + + assert config["offload_plan"] == { + "offload_granularity": "block", + "resident_blocks": {}, + "use_event_offload": False, + "use_block_slab": True, + } diff --git a/tests/common/offload/test_offload_schedule.py b/tests/common/offload/test_offload_schedule.py new file mode 100644 index 000000000..1dcacee19 --- /dev/null +++ b/tests/common/offload/test_offload_schedule.py @@ -0,0 +1,667 @@ +from contextlib import nullcontext +from types import SimpleNamespace + +import pytest +import torch + +from lightx2v.common.transformer_infer import transformer_infer as transformer_infer_module +from lightx2v.common.transformer_infer.transformer_infer import BaseTransformerInfer +from lightx2v.models.networks.base_model import BaseTransformerModel +from lightx2v.models.runners.base_runner import BaseRunner, keep_transformer_weights_loaded + + +class _Stream: + def __init__(self, name, log): + self.name = name + self.log = log + + def wait_stream(self, stream): + self.log.append(("wait_stream", self.name, stream.name)) + + def wait_event(self, event): + self.log.append(("wait_event", self.name, event)) + + def record_event(self): + event = f"{self.name}_done" + self.log.append(("record_event", self.name, event)) + return event + + +class _Device: + def __init__(self, caller_stream): + self.caller_stream = caller_stream + + def current_stream(self): + return self.caller_stream + + @staticmethod + def stream(stream): + return nullcontext() + + def synchronize(self): + self.caller_stream.log.append(("device_sync",)) + + +class _Blocks(list): + pass + + +class _Infer(BaseTransformerInfer): + def infer(self): + raise NotImplementedError + + +class _BindManager: + def __init__(self, uses_events=False, load_stream=None, compute_stream=None): + self.uses_events = uses_events + self.cuda_load_stream = load_stream if load_stream is not None else object() + self.compute_stream = compute_stream if compute_stream is not None else object() + self.cuda_buffers = None + + def init_cuda_buffer(self, buffers): + self.cuda_buffers = list(buffers) + + +class _BindInfer(_Infer): + def __init__(self, uses_events=False): + self.uses_events = uses_events + self.created_managers = [] + + def create_block_offload_manager(self, config, **kwargs): + manager = _BindManager(self.uses_events, **kwargs) + self.created_managers.append(manager) + return manager + + +class _TransformerWeights: + def __init__(self, groups): + self.groups = groups + + def get_offload_block_groups(self): + return self.groups + + +class _EventManager: + uses_events = True + + def __init__(self, log, slot_count=2): + self.log = log + self.slot_count = slot_count + self.compute_stream = _Stream("compute", log) + self.pending = [False] * slot_count + self.loaded_blocks = [None] * slot_count + self.load_calls = [] + + def prefetch_to_slot(self, slot_idx, block_idx, blocks, adapter_block_idx, state_dict_transform): + assert not self.pending[slot_idx] + self.pending[slot_idx] = True + self.loaded_blocks[slot_idx] = block_idx + self.load_calls.append((block_idx, adapter_block_idx, state_dict_transform)) + self.log.append(("prefetch", slot_idx, block_idx)) + + def wait_ready(self, slot_idx, stream): + assert self.pending[slot_idx] + self.log.append(("ready", slot_idx, stream.name)) + return ("slot", slot_idx, self.loaded_blocks[slot_idx]) + + def record_free(self, slot_idx, stream): + assert self.pending[slot_idx] + self.log.append(("free", slot_idx, stream.name)) + self.pending[slot_idx] = False + + def reset_slots(self): + self.log.append(("reset_slots",)) + self.pending = [False] * self.slot_count + + +class _Buffer: + def __init__(self, slot_idx): + self.slot_idx = slot_idx + self.block_idx = None + + +class _OutputTensor(torch.Tensor): + @staticmethod + def __new__(cls, log): + return torch.Tensor._make_subclass(cls, torch.empty(0), False) + + def __init__(self, log): + self.log = log + + def record_stream(self, stream): + self.log.append(("record_stream", stream.name)) + + +class _StreamManager: + uses_events = False + + def __init__(self, log, slot_count=2): + self.log = log + self.compute_stream = _Stream("compute", log) + self.cuda_buffers = [_Buffer(slot_index) for slot_index in range(slot_count)] + self.need_init_first_buffer = True + self.loaded_state_dict_transform = None + self.loaded_first_adapter_block_index = None + self.load_calls = [] + + def init_first_buffer(self, blocks, adapter_block_idx, block_idx, state_dict_transform): + self.cuda_buffers[0].block_idx = block_idx + self.need_init_first_buffer = False + self.load_calls.append((block_idx, adapter_block_idx, state_dict_transform)) + self.log.append(("init", self.cuda_buffers[0].slot_idx, block_idx)) + + def prefetch_weights_to_buffer(self, buffer_idx, block_idx, blocks, adapter_block_idx, state_dict_transform): + self.cuda_buffers[buffer_idx].block_idx = block_idx + self.load_calls.append((block_idx, adapter_block_idx, state_dict_transform)) + self.log.append(("prefetch", self.cuda_buffers[buffer_idx].slot_idx, block_idx)) + + def swap_blocks(self): + self.cuda_buffers[0], self.cuda_buffers[1] = self.cuda_buffers[1], self.cuda_buffers[0] + self.log.append(("swap",)) + + def wait_for_block_compute(self): + self.log.append(("wait_compute",)) + + +def _make_blocks(group_name="blocks"): + blocks = _Blocks([("resident", index) for index in range(6)]) + blocks.offload_group_name = group_name + blocks.resident_block_indices = frozenset({0, 3}) + blocks.offload_block_indices = (1, 2, 4, 5) + return blocks + + +def _make_one_offloaded_block(): + blocks = _Blocks([("resident", index) for index in range(4)]) + blocks.offload_group_name = "blocks" + blocks.resident_block_indices = frozenset({0, 1, 3}) + blocks.offload_block_indices = (2,) + return blocks + + +def _make_three_offloaded_blocks(): + blocks = _Blocks([("resident", index) for index in range(5)]) + blocks.offload_group_name = "blocks" + blocks.resident_block_indices = frozenset({0, 3}) + blocks.offload_block_indices = (1, 2, 4) + return blocks + + +def _make_infer(manager, group_name="blocks"): + infer = _Infer() + infer._block_offload_managers = {group_name: manager} + return infer + + +def _block_offload_config(use_events=False): + return { + "cpu_offload": True, + "offload_plan": { + "offload_granularity": "block", + "use_event_offload": use_events, + }, + } + + +def test_single_block_group_is_bound_to_its_buffers(): + blocks = _make_blocks() + buffers = [_Buffer(0), _Buffer(1)] + blocks.offload_cuda_buffers = buffers + infer = _BindInfer() + + infer.init_block_offload(_block_offload_config(), _TransformerWeights({"blocks": blocks})) + + manager = infer.get_block_offload_manager(blocks) + assert manager is infer.created_managers[0] + assert manager.cuda_buffers == buffers + + +def test_multiple_block_groups_are_bound_independently(): + double_blocks = _make_blocks("double_blocks") + single_blocks = _make_blocks("single_blocks") + double_buffers = [_Buffer(0), _Buffer(1)] + single_buffers = [_Buffer(0), _Buffer(1)] + double_blocks.offload_cuda_buffers = double_buffers + single_blocks.offload_cuda_buffers = single_buffers + infer = _BindInfer() + + infer.init_block_offload( + _block_offload_config(), + _TransformerWeights( + { + "double_blocks": double_blocks, + "single_blocks": single_blocks, + } + ), + ) + + double_manager = infer.get_block_offload_manager(double_blocks) + single_manager = infer.get_block_offload_manager(single_blocks) + assert double_manager is not single_manager + assert double_manager.cuda_buffers == double_buffers + assert single_manager.cuda_buffers == single_buffers + + +def test_fully_resident_group_is_bound_without_a_manager(): + blocks = _make_blocks() + blocks.resident_block_indices = frozenset(range(len(blocks))) + blocks.offload_block_indices = () + infer = _BindInfer() + + infer.init_block_offload(_block_offload_config(), _TransformerWeights({"blocks": blocks})) + + assert infer.get_block_offload_manager(blocks) is None + assert infer.created_managers == [] + computed = [] + infer.run_blocks_with_offload(blocks, lambda block_idx, block: computed.append((block_idx, block))) + assert computed == list(enumerate(blocks)) + + +def test_block_manager_lookup_rejects_an_unbound_group(): + infer = _Infer() + infer._block_offload_managers = {} + + with pytest.raises(KeyError): + infer.get_block_offload_manager(_make_blocks("missing")) + + +def test_get_offload_managers_includes_bound_groups_and_legacy_manager(): + first = object() + second = object() + legacy = object() + infer = _Infer() + infer._block_offload_managers = { + "first": first, + "resident": None, + "second": second, + } + infer.offload_manager = legacy + + assert infer.get_offload_managers() == [first, second, legacy] + + +def test_clear_offload_managers_clears_block_and_phase_managers(): + infer = _Infer() + infer._block_offload_managers = {"blocks": object()} + infer.offload_manager = object() + + infer.clear_offload_managers() + + assert infer._block_offload_managers == {} + assert not hasattr(infer, "offload_manager") + + +def test_unregistered_block_groups_keep_the_legacy_manager_path(): + class LegacyManager: + def init_cuda_buffer(self, block_buffers, phase_buffers): + self.buffers = block_buffers, phase_buffers + + block_buffers = object() + phase_buffers = object() + model = SimpleNamespace( + cpu_offload=True, + offload_granularity="block", + lazy_load=False, + transformer_infer=SimpleNamespace(offload_manager=LegacyManager()), + transformer_weights=SimpleNamespace( + get_offload_block_groups=lambda: {}, + offload_block_cuda_buffers=block_buffers, + offload_phase_cuda_buffers=phase_buffers, + ), + ) + + BaseTransformerModel._init_offload_manager(model) + + assert model.transformer_infer.offload_manager.buffers == (block_buffers, phase_buffers) + + +def test_event_block_groups_share_load_and_compute_streams(): + double_blocks = _make_blocks("double_blocks") + single_blocks = _make_blocks("single_blocks") + double_blocks.offload_cuda_buffers = [_Buffer(0), _Buffer(1)] + single_blocks.offload_cuda_buffers = [_Buffer(0), _Buffer(1)] + infer = _BindInfer(uses_events=True) + + infer.init_block_offload( + _block_offload_config(use_events=True), + _TransformerWeights( + { + "double_blocks": double_blocks, + "single_blocks": single_blocks, + } + ), + ) + + double_manager = infer.get_block_offload_manager(double_blocks) + single_manager = infer.get_block_offload_manager(single_blocks) + assert single_manager.cuda_load_stream is double_manager.cuda_load_stream + assert single_manager.compute_stream is double_manager.compute_stream + + +def test_nonresident_blocks_require_staging_buffers(): + blocks = _make_blocks() + infer = _BindInfer() + + with pytest.raises(RuntimeError, match="has no staging buffers"): + infer.init_block_offload( + _block_offload_config(), + _TransformerWeights({"blocks": blocks}), + ) + + +def test_staging_buffer_count_must_match_the_offloaded_blocks(): + blocks = _make_blocks() + blocks.offload_cuda_buffers = [_Buffer(0)] + infer = _BindInfer() + + with pytest.raises(RuntimeError, match="requires 2 staging buffers, got 1"): + infer.init_block_offload( + _block_offload_config(), + _TransformerWeights({"blocks": blocks}), + ) + + +def test_event_slots_overlap_nonresident_prefetch_and_preserve_block_order(monkeypatch): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _EventManager(log) + computations = [] + + _make_infer(manager).run_blocks_with_offload( + _make_blocks(), + lambda block_idx, block: computations.append((block_idx, block)), + ) + + assert [entry for entry in log if entry[0] == "prefetch"] == [ + ("prefetch", 0, 1), + ("prefetch", 1, 2), + ("prefetch", 0, 4), + ("prefetch", 1, 5), + ] + assert computations == [ + (0, ("resident", 0)), + (1, ("slot", 0, 1)), + (2, ("slot", 1, 2)), + (3, ("resident", 3)), + (4, ("slot", 0, 4)), + (5, ("slot", 1, 5)), + ] + assert ("wait_stream", "compute", "caller") in log + assert [entry for entry in log if entry[0] in {"ready", "free"}] == [ + ("ready", 0, "compute"), + ("free", 0, "compute"), + ("ready", 1, "compute"), + ("free", 1, "compute"), + ("ready", 0, "compute"), + ("free", 0, "compute"), + ("ready", 1, "compute"), + ("free", 1, "compute"), + ] + assert ("wait_event", "caller", "compute_done") in log + assert manager.pending == [False, False] + + +def test_offload_output_is_owned_by_the_caller_stream(monkeypatch): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + output = _OutputTensor(log) + + result = _make_infer(_EventManager(log)).run_blocks_with_offload( + _make_blocks(), + lambda _block_idx, _block: output, + ) + + assert result is output + assert ("wait_event", "caller", "compute_done") in log + assert ("record_stream", "caller") in log + + +@pytest.mark.parametrize("error_type", [RuntimeError, KeyboardInterrupt]) +def test_event_slots_are_reset_after_block_failure(monkeypatch, error_type): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _EventManager(log) + + def fail_on_first_offloaded_block(block_idx, _block): + if block_idx == 1: + raise error_type("block failed") + + with pytest.raises(error_type, match="block failed"): + _make_infer(manager).run_blocks_with_offload(_make_blocks(), fail_on_first_offloaded_block) + + assert ("device_sync",) in log + assert ("reset_slots",) in log + assert manager.pending == [False, False] + + +def test_residency_is_independent_of_event_scheduling(monkeypatch): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _StreamManager(log) + original_buffer_order = tuple(manager.cuda_buffers) + computations = [] + + _make_infer(manager).run_blocks_with_offload( + _make_blocks(), + lambda block_idx, block: computations.append((block_idx, block if isinstance(block, tuple) else (block.slot_idx, block.block_idx))), + ) + + assert [block_idx for block_idx, _ in computations] == list(range(6)) + assert computations[0][1] == ("resident", 0) + assert computations[3][1] == ("resident", 3) + assert [block_index for index, (_, block_index) in computations if index not in {0, 3}] == [1, 2, 4, 5] + assert tuple(manager.cuda_buffers) == original_buffer_order + + +@pytest.mark.parametrize("manager_class", [_EventManager, _StreamManager]) +def test_model_specific_weight_mapping_is_forwarded(monkeypatch, manager_class): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = manager_class(log) + + def adapter_block_index(block_index): + return block_index // 2 + + def state_dict_transform(block_index, state_dict): + return block_index, state_dict + + _make_infer(manager).run_blocks_with_offload( + _make_blocks(), + lambda _block_index, _block: None, + adapter_block_index=adapter_block_index, + state_dict_transform=state_dict_transform, + ) + + expected_loads = [ + (1, 0, state_dict_transform), + (2, 1, state_dict_transform), + (4, 2, state_dict_transform), + (5, 2, state_dict_transform), + ] + if not manager.uses_events: + expected_loads.append((1, 0, state_dict_transform)) + assert manager.load_calls == expected_loads + + +def test_stream_offload_prefetches_the_next_step_first_block(monkeypatch): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _StreamManager(log) + infer = _make_infer(manager) + blocks = _make_blocks() + + infer.run_blocks_with_offload(blocks, lambda _block_index, _block: None) + infer.run_blocks_with_offload(blocks, lambda _block_index, _block: None) + + assert [entry for entry in log if entry[0] == "init"] == [("init", 0, 1)] + assert [block_index for block_index, _, _ in manager.load_calls] == [1, 2, 4, 5, 1, 2, 4, 5, 1] + + +def test_stream_offload_preserves_staging_order_with_an_odd_block_count(monkeypatch): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _StreamManager(log) + infer = _make_infer(manager) + blocks = _make_three_offloaded_blocks() + staged_blocks = [] + + def record_staged_block(block_index, block): + if block_index in blocks.offload_block_indices: + staged_blocks.append(block.block_idx) + + infer.run_blocks_with_offload(blocks, record_staged_block) + infer.run_blocks_with_offload(blocks, record_staged_block) + + assert staged_blocks == [1, 2, 4, 1, 2, 4] + + +def test_stream_offload_reuses_a_single_staging_block(monkeypatch): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _StreamManager(log, slot_count=1) + infer = _make_infer(manager) + blocks = _make_one_offloaded_block() + computed = [] + + infer.run_blocks_with_offload(blocks, lambda block_index, _block: computed.append(block_index)) + infer.run_blocks_with_offload(blocks, lambda block_index, _block: computed.append(block_index)) + + assert computed == [0, 1, 2, 3, 0, 1, 2, 3] + assert manager.load_calls == [(2, None, None)] + + +def test_event_offload_reuses_a_single_slot_across_steps(monkeypatch): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _EventManager(log, slot_count=1) + infer = _make_infer(manager) + blocks = _make_one_offloaded_block() + staged_blocks = [] + + infer.run_blocks_with_offload(blocks, lambda block_index, block: staged_blocks.append(block[2]) if block_index == 2 else None) + infer.run_blocks_with_offload(blocks, lambda block_index, block: staged_blocks.append(block[2]) if block_index == 2 else None) + + assert staged_blocks == [2, 2] + assert manager.pending == [False] + + +def test_stream_offload_reloads_first_block_when_adapter_mapping_changes(monkeypatch): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _StreamManager(log) + infer = _make_infer(manager) + blocks = _make_blocks() + + infer.run_blocks_with_offload(blocks, lambda _block_index, _block: None, adapter_block_index=lambda _index: 0) + infer.run_blocks_with_offload(blocks, lambda _block_index, _block: None, adapter_block_index=lambda _index: 1) + + assert [entry for entry in log if entry[0] == "init"] == [ + ("init", 0, 1), + ("init", 0, 1), + ] + assert manager.load_calls[0][1] == 0 + assert manager.load_calls[5][1] == 1 + + +@pytest.mark.parametrize("error_type", [RuntimeError, KeyboardInterrupt]) +def test_stream_slots_are_restored_after_block_failure(monkeypatch, error_type): + log = [] + caller_stream = _Stream("caller", log) + monkeypatch.setattr(transformer_infer_module, "torch_device_module", _Device(caller_stream)) + monkeypatch.setattr(transformer_infer_module, "AI_DEVICE", "cuda") + manager = _StreamManager(log) + original_buffer_order = tuple(manager.cuda_buffers) + + def fail_on_second_offloaded_block(block_idx, _block): + if block_idx == 2: + raise error_type("block failed") + + with pytest.raises(error_type, match="block failed"): + _make_infer(manager).run_blocks_with_offload(_make_blocks(), fail_on_second_offloaded_block) + + assert ("device_sync",) in log + assert manager.need_init_first_buffer is True + assert tuple(manager.cuda_buffers) == original_buffer_order + + +def test_runner_entry_keeps_offload_weights_active_for_the_whole_loop(): + class Model: + def __init__(self): + self.active = False + self.prepare_count = 0 + self.cleanup_count = 0 + + def prepare_offload_weights(self): + if self.active: + return False + self.active = True + self.prepare_count += 1 + return True + + def cleanup_offload_weights(self): + self.active = False + self.cleanup_count += 1 + + class Runner(BaseRunner): + @keep_transformer_weights_loaded + def run(self): + assert self.model.active + return "done" + + runner = Runner({}) + runner.model = Model() + + assert runner.run() == "done" + assert runner.model.prepare_count == 1 + assert runner.model.cleanup_count == 1 + + +def test_runner_does_not_extend_offload_session_across_warmup(): + class Model: + def __init__(self): + self.active = False + self.prepare_count = 0 + self.cleanup_count = 0 + + def prepare_offload_weights(self): + self.active = True + self.prepare_count += 1 + return True + + def cleanup_offload_weights(self): + self.active = False + self.cleanup_count += 1 + + class Runner(BaseRunner): + def init_modules(self): + return None + + def warmup(self): + assert not self.model.active + + runner = Runner({"warmup": True}) + runner.model = Model() + runner.init_modules() + + assert runner.model.prepare_count == 0 + assert runner.model.cleanup_count == 0 From bfa99c5bcb41bcb13605ae9197facf88937420ec Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:40:17 +0800 Subject: [PATCH 02/31] feat(offload): migrate Cosmos3 block scheduling --- .../infer/offload/transformer_infer.py | 40 ++++++++++++------- lightx2v/models/networks/cosmos3/model.py | 23 ++++++----- .../cosmos3/weights/transformer_weights.py | 29 +++++++++----- .../models/runners/cosmos3/cosmos3_runner.py | 2 + 4 files changed, 59 insertions(+), 35 deletions(-) diff --git a/lightx2v/models/networks/cosmos3/infer/offload/transformer_infer.py b/lightx2v/models/networks/cosmos3/infer/offload/transformer_infer.py index 5ae0503e6..cbf6e5284 100644 --- a/lightx2v/models/networks/cosmos3/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/cosmos3/infer/offload/transformer_infer.py @@ -1,6 +1,6 @@ import torch -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.cosmos3.infer.transformer_infer import Cosmos3TransformerInfer from lightx2v_platform.base.global_var import AI_DEVICE @@ -12,43 +12,53 @@ def __init__(self, config): super().__init__(config) if not self.config.get("cpu_offload", False): return - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity != "block": raise NotImplementedError("Cosmos3 transformer supports only block-level cpu_offload.") - self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) self.lazy_load = self.config.get("lazy_load", False) - if self.lazy_load: - self.offload_manager.init_lazy_load(num_workers=self.config.get("num_disk_workers", 4)) def infer_layers(self, layers, und_seq, gen_seq, rotary_emb): + if self.lazy_load: + return self.infer_layers_with_lazy_offload(layers, und_seq, gen_seq, rotary_emb) + + def run_cosmos_block(_block_idx, block): + nonlocal und_seq, gen_seq + und_seq, gen_seq = self._infer_block(block, und_seq, gen_seq, rotary_emb) + return und_seq, gen_seq + + self.run_blocks_with_offload(layers, run_cosmos_block) + return und_seq, gen_seq + + def infer_layers_with_lazy_offload(self, layers, und_seq, gen_seq, rotary_emb): + manager = self.get_block_offload_manager(layers) current_stream = torch_device_module.current_stream() - self.offload_manager.compute_stream.wait_stream(current_stream) + manager.compute_stream.wait_stream(current_stream) for block_idx in range(len(layers)): if self.lazy_load: next_prefetch = (block_idx + 1) % len(layers) - self.offload_manager.start_prefetch_block(next_prefetch) + manager.start_prefetch_block(next_prefetch) - if self.offload_manager.need_init_first_buffer: - self.offload_manager.init_first_buffer(layers) + if manager.need_init_first_buffer: + manager.init_first_buffer(layers) if self.lazy_load: - self.offload_manager.swap_cpu_buffers() + manager.swap_cpu_buffers() - self.offload_manager.prefetch_weights((block_idx + 1) % len(layers), layers) + manager.prefetch_weights((block_idx + 1) % len(layers), layers) if AI_DEVICE == "xpu": und_seq, gen_seq = self._infer_block( - self.offload_manager.cuda_buffers[0], + manager.cuda_buffers[0], und_seq, gen_seq, rotary_emb, ) else: - with torch_device_module.stream(self.offload_manager.compute_stream): + with torch_device_module.stream(manager.compute_stream): und_seq, gen_seq = self._infer_block( - self.offload_manager.cuda_buffers[0], + manager.cuda_buffers[0], und_seq, gen_seq, rotary_emb, ) - self.offload_manager.swap_blocks() + manager.swap_blocks() return und_seq, gen_seq diff --git a/lightx2v/models/networks/cosmos3/model.py b/lightx2v/models/networks/cosmos3/model.py index e2c827e40..c8a4bb89d 100644 --- a/lightx2v/models/networks/cosmos3/model.py +++ b/lightx2v/models/networks/cosmos3/model.py @@ -2,6 +2,7 @@ import torch.distributed as dist from torch.nn import functional as F +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.base_model import BaseTransformerModel from lightx2v.models.networks.cosmos3.infer.module_io import Cosmos3PostInferModuleOutput, Cosmos3TransformerInferModuleOutput from lightx2v.models.networks.cosmos3.infer.offload.transformer_infer import Cosmos3OffloadTransformerInfer @@ -21,8 +22,8 @@ class Cosmos3TransformerModel(BaseTransformerModel): def __init__(self, model_path, config, device): if config.get("lazy_load", False): raise NotImplementedError("Cosmos3 LightX2V native transformer does not support lazy_load yet.") - if config.get("cpu_offload", False) and config.get("offload_granularity", "block") != "block": - raise NotImplementedError("Cosmos3 LightX2V native transformer supports only block-level cpu_offload.") + if config.get("cpu_offload", False) and get_offload_granularity(config) not in ("block", "model"): + raise NotImplementedError("Cosmos3 LightX2V native transformer supports block- and model-level cpu_offload.") super().__init__(model_path, config, device) self._init_infer_class() self._init_weights() @@ -30,15 +31,15 @@ def __init__(self, model_path, config, device): def _init_infer_class(self): self.pre_infer_class = Cosmos3PreInfer - self.transformer_infer_class = Cosmos3OffloadTransformerInfer if self.cpu_offload else Cosmos3TransformerInfer + use_block_offload = self.cpu_offload and self.offload_granularity == "block" + self.transformer_infer_class = Cosmos3OffloadTransformerInfer if use_block_offload else Cosmos3TransformerInfer self.post_infer_class = Cosmos3PostInfer def _init_infer(self): self.pre_infer = self.pre_infer_class(self.config) self.transformer_infer = self.transformer_infer_class(self.config) self.post_infer = self.post_infer_class(self.config) - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() + self._init_offload_manager() def set_scheduler(self, scheduler): self.scheduler = scheduler @@ -149,9 +150,9 @@ def _infer_cond_uncond(self, input_ids): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: - self.pre_weight.to_cuda() - self.post_weight.to_cuda() + if self.cpu_offload and self.offload_granularity != "block": + if self.offload_granularity == "model" and self.scheduler.step_index == 0: + self.to_cuda() text_encoder_output = inputs["text_encoder_output"] do_cfg = self.config.get("enable_cfg", True) and self.scheduler.sample_guide_scale != 1.0 @@ -171,6 +172,6 @@ def infer(self, inputs): cond = self._infer_cond_uncond(text_encoder_output["cond_input_ids"]) self._set_scheduler_noise_pred(self._detach_cfg_output(cond)) - if self.cpu_offload: - self.pre_weight.to_cpu() - self.post_weight.to_cpu() + if self.cpu_offload and self.offload_granularity != "block": + if self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1: + self.to_cpu() diff --git a/lightx2v/models/networks/cosmos3/weights/transformer_weights.py b/lightx2v/models/networks/cosmos3/weights/transformer_weights.py index e4b5694c6..32c291d9d 100644 --- a/lightx2v/models/networks/cosmos3/weights/transformer_weights.py +++ b/lightx2v/models/networks/cosmos3/weights/transformer_weights.py @@ -1,6 +1,7 @@ import torch from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.cosmos3.infer.utils import Cosmos3Rope # noqa: F401 from lightx2v.utils.registry_factory import ATTN_WEIGHT_REGISTER, MM_WEIGHT_REGISTER, RMS_WEIGHT_REGISTER, ROPE_REGISTER @@ -18,7 +19,7 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): self.lazy_load_file = lazy_load_path else: self.lazy_load_file = None - layers = WeightModuleList( + self.layers = WeightModuleList( Cosmos3TransformerLayerWeights( layer_idx, self.mm_type, @@ -31,16 +32,23 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): ) for layer_idx in range(self.layers_num) ) - self.register_offload_buffers(config, lazy_load_path, lora_path) - self.add_module("layers", layers) + slot_count = self.register_offload_block_group(config, "layers", self.layers) + self.register_offload_buffers(config, lazy_load_path, lora_path, slot_count) + self.add_module("layers", self.layers) - def register_offload_buffers(self, config, lazy_load_path, lora_path): + def register_offload_buffers(self, config, lazy_load_path, lora_path, slot_count): if not config.get("cpu_offload", False): return - if config.get("offload_granularity", "block") != "block": - raise NotImplementedError("Cosmos3 transformer supports only block-level cpu_offload.") + if get_offload_granularity(config) == "model": + return - self.offload_blocks_num = 2 + self.offload_blocks_num = slot_count + self.offload_block_cuda_buffers = None + self.offload_block_cpu_buffers = None + self.offload_phase_cuda_buffers = None + self.offload_phase_cpu_buffers = None + if slot_count == 0: + return self.offload_block_cuda_buffers = WeightModuleList( [ Cosmos3TransformerLayerWeights( @@ -59,7 +67,6 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) - self.offload_phase_cuda_buffers = None if self.lazy_load: self.offload_block_cpu_buffers = WeightModuleList( @@ -80,7 +87,11 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cpu_buffers", self.offload_block_cpu_buffers) - self.offload_phase_cpu_buffers = None + self.register_offload_block_buffers( + "layers", + self.offload_block_cuda_buffers, + self.offload_block_cpu_buffers, + ) def non_block_weights_to_cuda(self): return None diff --git a/lightx2v/models/runners/cosmos3/cosmos3_runner.py b/lightx2v/models/runners/cosmos3/cosmos3_runner.py index 0c1b753eb..34624d57e 100644 --- a/lightx2v/models/runners/cosmos3/cosmos3_runner.py +++ b/lightx2v/models/runners/cosmos3/cosmos3_runner.py @@ -17,6 +17,7 @@ from lightx2v.models.audio_encoders.hf.cosmos3.sound_tokenizer import Cosmos3SoundTokenizer from lightx2v.models.networks.cosmos3.model import Cosmos3TransformerModel +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.cosmos3.policy_runtime import ( PolicySeedSequence, build_json_policy_prompt, @@ -849,6 +850,7 @@ def _save_action_output(self, action): json.dump(action.tolist(), f) logger.info(f"Action saved: {save_action_path}") + @keep_transformer_weights_loaded def run(self, total_steps=None): if total_steps is None: total_steps = self.model.scheduler.infer_steps From 481de2c09720b5005d115ec654f39f769d3cd4f5 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:40:41 +0800 Subject: [PATCH 03/31] feat(offload): migrate Flux2 block scheduling --- .../ascend_npu/flux2_dev_t2i_1344x768.json | 15 +- .../flux2/infer/offload/transformer_infer.py | 222 +++--------------- lightx2v/models/networks/flux2/model.py | 89 +++---- .../flux2/weights/transformer_weights.py | 141 +++-------- lightx2v/models/runners/flux2/flux2_runner.py | 6 +- 5 files changed, 109 insertions(+), 364 deletions(-) diff --git a/configs/platforms/ascend_npu/flux2_dev_t2i_1344x768.json b/configs/platforms/ascend_npu/flux2_dev_t2i_1344x768.json index 9e8b51dd8..b555b15fc 100644 --- a/configs/platforms/ascend_npu/flux2_dev_t2i_1344x768.json +++ b/configs/platforms/ascend_npu/flux2_dev_t2i_1344x768.json @@ -13,12 +13,15 @@ "text_encoder_out_layers": [10, 20, 30], "attn_type": "npu_flash_attn", "cpu_offload": true, - "offload_granularity": "block", - "use_event_offload": true, - "offload_use_block_slab": true, - "offload_resident_double_blocks": "all", - "offload_resident_single_blocks": 36, - "offload_resident_policy": "interleaved", + "offload_plan": { + "offload_granularity": "block", + "resident_blocks": { + "double_blocks": "all", + "single_blocks": 36 + }, + "use_event_offload": true, + "use_block_slab": true + }, "modulate_type": "torch", "rope_type": "npu_rope", "layer_norm_type": "npu_layer_norm", diff --git a/lightx2v/models/networks/flux2/infer/offload/transformer_infer.py b/lightx2v/models/networks/flux2/infer/offload/transformer_infer.py index 827919f4f..3551ce3f0 100644 --- a/lightx2v/models/networks/flux2/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/flux2/infer/offload/transformer_infer.py @@ -1,12 +1,8 @@ import torch import torch.nn.functional as F -from lightx2v.common.offload.event_manager import EventSlotWeightAsyncStreamManager -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.flux2.infer.transformer_infer import Flux2TransformerInfer -from lightx2v_platform.base.global_var import AI_DEVICE - -torch_device_module = getattr(torch, AI_DEVICE) class Flux2OffloadTransformerInfer(Flux2TransformerInfer): @@ -14,85 +10,20 @@ class Flux2OffloadTransformerInfer(Flux2TransformerInfer): def __init__(self, config): super().__init__(config) - self.use_event_offload = self.config.get("use_event_offload", False) - resident_blocks_requested = any( - self.config.get(key, 0) not in (None, 0) - for key in ( - "offload_resident_double_blocks", - "offload_resident_single_blocks", - ) - ) - if resident_blocks_requested and not self.use_event_offload: - raise ValueError("Flux2 resident block offload requires use_event_offload=true") - if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") - if offload_granularity == "block": - if self.use_event_offload: - self.offload_manager_double = EventSlotWeightAsyncStreamManager( - offload_granularity=offload_granularity, - ) - self.offload_manager_single = EventSlotWeightAsyncStreamManager( - offload_granularity=offload_granularity, - load_stream=self.offload_manager_double.cuda_load_stream, - compute_stream=self.offload_manager_double.compute_stream, - ) - self.infer_func = self.infer_with_event_offload - else: - self.offload_manager_double = WeightAsyncStreamManager(offload_granularity=offload_granularity) - self.offload_manager_single = WeightAsyncStreamManager(offload_granularity=offload_granularity) - self.infer_func = self.infer_with_blocks_offload - elif offload_granularity == "model": - self.infer_func = super().infer - else: - raise ValueError(f"Unsupported offload_granularity: {offload_granularity}") - else: + if not config.get("cpu_offload", False): self.infer_func = super().infer - - @staticmethod - def _resident_indices(block_weights, block_kind): - attr_name = f"resident_{block_kind}_block_indices" - return set(getattr(block_weights, attr_name, ())) - - def _run_event_block_stage(self, manager, blocks, resident_indices, run_block): - """Stream non-resident blocks through event-protected slots.""" - offloaded_indices = [idx for idx in range(len(blocks)) if idx not in resident_indices] - if not offloaded_indices: - for block_idx, block in enumerate(blocks): - run_block(block_idx, block) return - 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] - manager.prefetch_to_slot(slot_idx, block_idx, blocks) - scheduled_slots[block_idx] = slot_idx - next_offloaded += 1 - - for slot_idx in range(min(manager.slot_count, len(offloaded_indices))): - prefetch_next(slot_idx) - - for block_idx, block in enumerate(blocks): - if block_idx in resident_indices: - run_block(block_idx, block) - continue - - slot_idx = scheduled_slots.pop(block_idx) - staged_block = manager.wait_ready(slot_idx) - run_block(block_idx, staged_block) - manager.record_free(slot_idx) - prefetch_next(slot_idx) + offload_granularity = get_offload_granularity(config) + if offload_granularity == "model": + self.infer_func = super().infer + return + if offload_granularity != "block": + raise ValueError(f"Unsupported offload_granularity: {offload_granularity}") - def infer_with_event_offload(self, block_weights, pre_infer_out): - """Pipeline block loading and compute with device events. + self.infer_func = self.infer_with_blocks_offload - Unlike infer_with_blocks_offload, this path reuses fixed slots without - synchronizing the host after every block. - """ + def infer_with_blocks_offload(self, block_weights, pre_infer_out): hidden_states = pre_infer_out.hidden_states encoder_hidden_states = pre_infer_out.encoder_hidden_states timestep = pre_infer_out.timestep @@ -105,130 +36,45 @@ def infer_with_event_offload(self, block_weights, pre_infer_out): double_stream_mod_txt = block_weights.double_stream_modulation_txt_linear.apply(timestep_act) single_stream_mod = block_weights.single_stream_modulation_linear.apply(timestep_act) - device_module = self.offload_manager_double.device_module - current_stream = device_module.current_stream() - compute_stream = self.offload_manager_double.compute_stream - compute_stream.wait_stream(current_stream) - - resident_double = self._resident_indices(block_weights, "double") - resident_single = self._resident_indices(block_weights, "single") - def run_double(block_idx, block): nonlocal encoder_hidden_states, hidden_states self.block_idx = block_idx - with device_module.stream(compute_stream): - encoder_hidden_states, hidden_states = self.infer_double_stream_block( - block, - hidden_states, - encoder_hidden_states, - double_stream_mod_img, - double_stream_mod_txt, - image_rotary_emb, - image_rotary_positions, - ) + encoder_hidden_states, hidden_states = self.infer_double_stream_block( + block, + hidden_states, + encoder_hidden_states, + double_stream_mod_img, + double_stream_mod_txt, + image_rotary_emb, + image_rotary_positions, + ) + return encoder_hidden_states, hidden_states - self._run_event_block_stage( - self.offload_manager_double, + self.run_blocks_with_offload( block_weights.double_blocks, - resident_double, run_double, ) - - with device_module.stream(compute_stream): - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=0) + hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=0) def run_single(block_idx, block): nonlocal hidden_states self.block_idx = block_idx - with device_module.stream(compute_stream): - hidden_states = self.infer_single_stream_block( - block, - hidden_states, - None, - single_stream_mod, - image_rotary_emb, - image_rotary_positions, - num_txt_tokens=num_txt_tokens, - ) + hidden_states = self.infer_single_stream_block( + block, + hidden_states, + None, + single_stream_mod, + image_rotary_emb, + image_rotary_positions, + num_txt_tokens=num_txt_tokens, + ) + return hidden_states - self._run_event_block_stage( - self.offload_manager_single, + self.run_blocks_with_offload( block_weights.single_blocks, - resident_single, run_single, ) - - with device_module.stream(compute_stream): - hidden_states = hidden_states[num_txt_tokens:, ...] - final_done = compute_stream.record_event() - - current_stream.wait_event(final_done) - hidden_states.record_stream(current_stream) - return hidden_states - - def infer_with_blocks_offload(self, block_weights, pre_infer_out): - """Use ping-pong buffers with host synchronization after every block.""" - hidden_states = pre_infer_out.hidden_states - encoder_hidden_states = pre_infer_out.encoder_hidden_states - timestep = pre_infer_out.timestep - image_rotary_emb = pre_infer_out.image_rotary_emb - image_rotary_positions = pre_infer_out.image_rotary_positions - - num_txt_tokens = encoder_hidden_states.shape[0] - timestep_act = F.silu(timestep) - double_stream_mod_img = block_weights.double_stream_modulation_img_linear.apply(timestep_act) - double_stream_mod_txt = block_weights.double_stream_modulation_txt_linear.apply(timestep_act) - single_stream_mod = block_weights.single_stream_modulation_linear.apply(timestep_act) - - current_stream = torch_device_module.current_stream() - self.offload_manager_double.compute_stream.wait_stream(current_stream) - for block_idx in range(len(block_weights.double_blocks)): - self.block_idx = block_idx - - if self.offload_manager_double.need_init_first_buffer: - self.offload_manager_double.init_first_buffer(block_weights.double_blocks) - - self.offload_manager_double.prefetch_weights((block_idx + 1) % len(block_weights.double_blocks), block_weights.double_blocks) - - with torch_device_module.stream(self.offload_manager_double.compute_stream): - encoder_hidden_states, hidden_states = self.infer_double_stream_block( - self.offload_manager_double.cuda_buffers[0], - hidden_states, - encoder_hidden_states, - double_stream_mod_img, - double_stream_mod_txt, - image_rotary_emb, - image_rotary_positions, - ) - - self.offload_manager_double.swap_blocks() - - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=0) - - self.offload_manager_single.compute_stream.wait_stream(self.offload_manager_double.compute_stream) - for block_idx in range(len(block_weights.single_blocks)): - self.block_idx = block_idx - - if self.offload_manager_single.need_init_first_buffer: - self.offload_manager_single.init_first_buffer(block_weights.single_blocks) - - self.offload_manager_single.prefetch_weights((block_idx + 1) % len(block_weights.single_blocks), block_weights.single_blocks) - - with torch_device_module.stream(self.offload_manager_single.compute_stream): - hidden_states = self.infer_single_stream_block( - self.offload_manager_single.cuda_buffers[0], - hidden_states, - None, - single_stream_mod, - image_rotary_emb, - image_rotary_positions, - num_txt_tokens=num_txt_tokens, - ) - - self.offload_manager_single.swap_blocks() - - hidden_states = hidden_states[num_txt_tokens:, ...] - return hidden_states + return hidden_states[num_txt_tokens:, ...] def infer(self, block_weights, pre_infer_out): return self.infer_func(block_weights, pre_infer_out) diff --git a/lightx2v/models/networks/flux2/model.py b/lightx2v/models/networks/flux2/model.py index e94acbac6..8411e7060 100644 --- a/lightx2v/models/networks/flux2/model.py +++ b/lightx2v/models/networks/flux2/model.py @@ -2,6 +2,7 @@ import torch.distributed as dist from torch.nn import functional as F +from lightx2v.common.offload.config import get_offload_plan from lightx2v.models.networks.base_model import BaseTransformerModel from lightx2v.models.networks.flux2.infer.feature_caching.transformer_infer import Flux2TransformerInferAdaCaching from lightx2v.models.networks.flux2.infer.offload.transformer_infer import Flux2OffloadTransformerInfer @@ -10,11 +11,7 @@ from lightx2v.models.networks.flux2.infer.transformer_infer import Flux2TransformerInfer from lightx2v.models.networks.flux2.weights.post_weights import Flux2PostWeights from lightx2v.models.networks.flux2.weights.pre_weights import Flux2DevPreWeights, Flux2PreWeights -from lightx2v.models.networks.flux2.weights.transformer_weights import ( - Flux2TransformerWeights, - preserve_weight_module_cpu_tensors, - release_weight_module_device_tensors, -) +from lightx2v.models.networks.flux2.weights.transformer_weights import Flux2TransformerWeights from lightx2v_platform.base import global_var @@ -29,8 +26,10 @@ def __init__(self, config, model_path, device): # Block-slab packing is platform-agnostic. Enable it only on backends that support # event offload, with block-level CPU offload, BF16 Default weights, and no LoRA, # lazy loading, or tensor parallelism. - self.use_block_slab_offload = self.config.get("offload_use_block_slab", False) - self._offload_weights_active = False + offload_plan = get_offload_plan(self.config) + self.use_block_slab_offload = offload_plan.get("use_block_slab", False) + if self.use_block_slab_offload and not offload_plan.get("use_event_offload", False): + raise ValueError("Flux2 block slab offload requires offload_plan.use_event_offload=true") self.in_channels = self.config.get("transformer_in_channels", self.config.get("in_channels", 64)) self.attention_kwargs = {} self._combined_img_ids_cache = None @@ -281,66 +280,44 @@ def _init_infer(self): self.pre_infer = self.pre_infer_class(self.config) self.post_infer = self.post_infer_class(self.config) self.pre_infer.set_rope(self.transformer_weights.double_blocks[0].rope) - if hasattr(self.transformer_infer, "offload_manager_double") and hasattr(self.transformer_infer, "offload_manager_single"): - self._init_offload_manager() + self._init_offload_manager() def _init_offload_manager(self): - if hasattr(self.transformer_weights, "offload_double_block_cuda_buffers"): - self.transformer_infer.offload_manager_double.init_cuda_buffer(blocks_cuda_buffer=self.transformer_weights.offload_double_block_cuda_buffers) - if hasattr(self.transformer_weights, "offload_single_block_cuda_buffers"): - self.transformer_infer.offload_manager_single.init_cuda_buffer(blocks_cuda_buffer=self.transformer_weights.offload_single_block_cuda_buffers) - if self.use_block_slab_offload: - double_slabs, single_slabs = self.transformer_weights.prepare_offload_block_slabs() - if double_slabs: - self.transformer_infer.offload_manager_double.init_block_slabs(double_slabs) - if single_slabs: - self.transformer_infer.offload_manager_single.init_block_slabs(single_slabs) + super()._init_offload_manager() + if not self.cpu_offload or self.offload_granularity != "block" or not self.use_block_slab_offload: + return + + double_slabs, single_slabs = self.transformer_weights.prepare_offload_block_slabs() + double_manager = self.transformer_infer.get_block_offload_manager(self.transformer_weights.double_blocks) + single_manager = self.transformer_infer.get_block_offload_manager(self.transformer_weights.single_blocks) + if double_manager is not None: + double_manager.init_block_slabs(double_slabs) + if single_manager is not None: + single_manager.init_block_slabs(single_slabs) def prepare_offload_weights(self): - """Load weights kept resident for one runner invocation.""" - if not self.cpu_offload: - return + if not self.cpu_offload or self.offload_granularity != "model": + return super().prepare_offload_weights() if self._offload_weights_active: - raise RuntimeError("Flux2 offload weights are already active") + return False - # Mark the model active before moving weights so runner cleanup also - # handles a partially completed preparation. self._offload_weights_active = True - if self.offload_granularity == "model": + try: self.to_cuda() - else: - # These weights are used on every diffusion step, so keep them - # resident for the complete denoising loop. - preserve_weight_module_cpu_tensors(self.pre_weight) - preserve_weight_module_cpu_tensors(self.post_weight) - self.pre_weight.to_cuda() - self.post_weight.to_cuda() - self.transformer_weights.non_block_weights_to_cuda() - self.transformer_weights.resident_blocks_to_cuda() - - def force_cleanup_offload_weights(self): - """Release loaded offload weights and reset event-slot state. - - This method is intentionally idempotent and may be called from a - runner ``finally`` block after a short run, cancellation, or error. - """ - if not self.cpu_offload or not self._offload_weights_active: + except BaseException: + self.cleanup_offload_weights() + raise + return True + + def cleanup_offload_weights(self): + if self.offload_granularity != "model": + super().cleanup_offload_weights() + return + if not self._offload_weights_active: return - # Device execution and non-blocking H2D copies are asynchronous. self._sync_device() - if getattr(self.transformer_infer, "use_event_offload", False): - self.transformer_infer.offload_manager_double.reset_slots() - self.transformer_infer.offload_manager_single.reset_slots() - - if self.offload_granularity == "model": - self.to_cpu() - else: - release_weight_module_device_tensors(self.pre_weight) - release_weight_module_device_tensors(self.post_weight) - self.transformer_weights.release_non_block_weights() - self.transformer_weights.release_resident_blocks() - + self.to_cpu() self._offload_weights_active = False def _get_combined_img_ids(self, img_ids, input_image_ids): diff --git a/lightx2v/models/networks/flux2/weights/transformer_weights.py b/lightx2v/models/networks/flux2/weights/transformer_weights.py index db8d96aff..ff4a52910 100644 --- a/lightx2v/models/networks/flux2/weights/transformer_weights.py +++ b/lightx2v/models/networks/flux2/weights/transformer_weights.py @@ -3,46 +3,11 @@ from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList from lightx2v.common.offload.block_slab import pack_cpu_block_slab +from lightx2v.common.offload.config import get_offload_plan from lightx2v.common.ops.utils import move_transposed_weight_module_to_device from lightx2v.utils.registry_factory import ATTN_WEIGHT_REGISTER, LN_WEIGHT_REGISTER, MM_WEIGHT_REGISTER, RMS_WEIGHT_REGISTER, ROPE_REGISTER -def _resolve_resident_block_indices(value, num_blocks, policy, config_key): - """Resolve a resident-block count into deterministic block indices. - - Resident blocks are opt-in. A count of zero therefore preserves the - original full block-streaming behaviour. ``interleaved`` spreads the - resident blocks over the whole transformer instead of concentrating them - 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 - - if not 0 <= count <= num_blocks: - raise ValueError(f"{config_key} must be between 0 and {num_blocks}, got {count}") - if count == 0: - return frozenset() - if count == num_blocks: - return frozenset(range(num_blocks)) - - if policy == "prefix": - return frozenset(range(count)) - if policy == "interleaved": - # floor(k * N / K) is duplicate-free for K <= N. For example, - # K=36 and N=48 leaves blocks 3, 7, ..., 47 for streaming. - return frozenset((idx * num_blocks) // count for idx in range(count)) - raise ValueError(f"offload_resident_policy must be 'prefix' or 'interleaved', got {policy!r}") - - def preserve_weight_module_cpu_tensors(module): """Keep CPU masters for attributes that do not have a pinned counterpart.""" for child in getattr(module, "_modules", {}).values(): @@ -256,12 +221,32 @@ def __init__(self, config): self.num_layers = config.get("num_layers", 10) self.num_single_layers = config.get("num_single_layers", 20) self.mm_type = config.get("dit_quant_scheme", "Default") - self._configure_resident_blocks(config) self.double_blocks = WeightModuleList([Flux2DoubleBlockWeights(config, i) for i in range(self.num_layers)]) self.single_blocks = WeightModuleList([Flux2SingleBlockWeights(config, i) for i in range(self.num_single_layers)]) - self.register_offload_buffers(config) + double_buffer_count = self.register_offload_block_group(config, "double_blocks", self.double_blocks) + single_buffer_count = self.register_offload_block_group(config, "single_blocks", self.single_blocks) + has_resident_blocks = ( + config.get("cpu_offload", False) + and get_offload_plan(config).get("offload_granularity", "block") == "block" + and bool(self.double_blocks.resident_block_indices or self.single_blocks.resident_block_indices) + ) + if has_resident_blocks and config.get("dit_quantized", False): + raise NotImplementedError("Flux2 resident block offload currently supports unquantized weights only") + if has_resident_blocks and config.get("lora_configs"): + raise NotImplementedError("Flux2 resident block offload currently does not support LoRA weights") + if double_buffer_count: + self.offload_double_block_cuda_buffers = WeightModuleList([Flux2DoubleBlockWeights(config, i, create_cuda_buffer=True) for i in range(double_buffer_count)]) + self.add_module("offload_double_block_cuda_buffers", self.offload_double_block_cuda_buffers) + self.register_offload_block_buffers("double_blocks", self.offload_double_block_cuda_buffers) + if single_buffer_count: + self.offload_single_block_cuda_buffers = WeightModuleList([Flux2SingleBlockWeights(config, i, create_cuda_buffer=True) for i in range(single_buffer_count)]) + self.add_module("offload_single_block_cuda_buffers", self.offload_single_block_cuda_buffers) + self.register_offload_block_buffers("single_blocks", self.offload_single_block_cuda_buffers) + + # Staging buffers must be loaded before the source blocks consume their + # entries from the shared checkpoint state dict. self.add_module("double_blocks", self.double_blocks) self.add_module("single_blocks", self.single_blocks) @@ -269,75 +254,6 @@ def __init__(self, config): self.add_module("double_stream_modulation_txt_linear", _mm_weight(config, "double_stream_modulation_txt.linear.weight")) self.add_module("single_stream_modulation_linear", _mm_weight(config, "single_stream_modulation.linear.weight")) - def _configure_resident_blocks(self, config): - block_offload_enabled = config.get("cpu_offload", False) and config.get("offload_granularity", "block") == "block" - if not block_offload_enabled: - double_setting = 0 - single_setting = 0 - else: - double_setting = config.get("offload_resident_double_blocks", 0) - single_setting = config.get("offload_resident_single_blocks", 0) - - resident_blocks_requested = double_setting not in (None, 0) or single_setting not in (None, 0) - if resident_blocks_requested and config.get("dit_quantized", False): - raise NotImplementedError("Flux2 resident block offload currently supports unquantized weights only") - if resident_blocks_requested and config.get("lora_configs"): - raise NotImplementedError("Flux2 resident block offload currently does not support LoRA weights") - - policy = config.get("offload_resident_policy", "prefix") - self.resident_double_block_indices = _resolve_resident_block_indices( - double_setting, - self.num_layers, - policy, - "offload_resident_double_blocks", - ) - self.resident_single_block_indices = _resolve_resident_block_indices( - single_setting, - self.num_single_layers, - policy, - "offload_resident_single_blocks", - ) - - def register_offload_buffers(self, config): - if config.get("cpu_offload", False) and config.get("offload_granularity", "block") == "block": - if len(self.resident_double_block_indices) < self.num_layers: - self.offload_double_block_cuda_buffers = WeightModuleList([Flux2DoubleBlockWeights(config, i, create_cuda_buffer=True) for i in range(2)]) - self.add_module("offload_double_block_cuda_buffers", self.offload_double_block_cuda_buffers) - - if len(self.resident_single_block_indices) < self.num_single_layers: - self.offload_single_block_cuda_buffers = WeightModuleList([Flux2SingleBlockWeights(config, i, create_cuda_buffer=True) for i in range(2)]) - self.add_module("offload_single_block_cuda_buffers", self.offload_single_block_cuda_buffers) - - def is_double_block_resident(self, block_idx): - return block_idx in self.resident_double_block_indices - - def is_single_block_resident(self, block_idx): - return block_idx in self.resident_single_block_indices - - def get_resident_double_block(self, block_idx): - if not self.is_double_block_resident(block_idx): - return None - return self.double_blocks[block_idx] - - def get_resident_single_block(self, block_idx): - if not self.is_single_block_resident(block_idx): - return None - return self.single_blocks[block_idx] - - def resident_blocks_to_cuda(self, non_blocking=True): - for block_idx in sorted(self.resident_double_block_indices): - preserve_weight_module_cpu_tensors(self.double_blocks[block_idx]) - self.double_blocks[block_idx].to_cuda(non_blocking=non_blocking) - for block_idx in sorted(self.resident_single_block_indices): - preserve_weight_module_cpu_tensors(self.single_blocks[block_idx]) - self.single_blocks[block_idx].to_cuda(non_blocking=non_blocking) - - def release_resident_blocks(self): - for block_idx in sorted(self.resident_double_block_indices): - release_weight_module_device_tensors(self.double_blocks[block_idx]) - for block_idx in sorted(self.resident_single_block_indices): - release_weight_module_device_tensors(self.single_blocks[block_idx]) - @staticmethod def _pack_offload_block_slabs(blocks, resident_indices): slabs = {} @@ -399,18 +315,18 @@ def collect_cpu_base_tensors(module): def prepare_offload_block_slabs(self): """Pack non-resident block weights after checkpoint loading.""" - if not self.config.get("offload_use_block_slab", False): + if not get_offload_plan(self.config).get("use_block_slab", False): return {}, {} if hasattr(self, "offload_double_block_slabs"): return self.offload_double_block_slabs, self.offload_single_block_slabs double_slabs = self._pack_offload_block_slabs( self.double_blocks, - self.resident_double_block_indices, + self.double_blocks.resident_block_indices, ) single_slabs = self._pack_offload_block_slabs( self.single_blocks, - self.resident_single_block_indices, + self.single_blocks.resident_block_indices, ) self.offload_double_block_slabs = double_slabs self.offload_single_block_slabs = single_slabs @@ -439,6 +355,11 @@ def release_non_block_weights(self): release_weight_module_device_tensors(self.double_stream_modulation_txt_linear) release_weight_module_device_tensors(self.single_stream_modulation_linear) + def release_resident_blocks(self): + for blocks in (self.double_blocks, self.single_blocks): + for block_index in blocks.resident_block_indices: + release_weight_module_device_tensors(blocks[block_index]) + def to_cuda(self, non_blocking=True): for block in self.double_blocks: block.to_cuda(non_blocking=non_blocking) diff --git a/lightx2v/models/runners/flux2/flux2_runner.py b/lightx2v/models/runners/flux2/flux2_runner.py index c0c92dc89..2c843d586 100644 --- a/lightx2v/models/runners/flux2/flux2_runner.py +++ b/lightx2v/models/runners/flux2/flux2_runner.py @@ -7,6 +7,7 @@ from loguru import logger from lightx2v.models.networks.flux2.model import Flux2DevTransformerModel, Flux2KleinTransformerModel +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.default_runner import DefaultRunner from lightx2v.models.schedulers.flux2.feature_caching.scheduler import Flux2DevSchedulerCaching, Flux2SchedulerCaching from lightx2v.models.schedulers.flux2.scheduler import Flux2DevScheduler, Flux2Scheduler @@ -295,12 +296,11 @@ def _run_dit_local_i2i(self, total_steps=None): latents, generator = self.run(total_steps) return latents, generator + @keep_transformer_weights_loaded def run(self, total_steps=None): if total_steps is None: total_steps = self.model.scheduler.infer_steps - self.model.prepare_offload_weights() - for step_index in range(total_steps): logger.info(f"==> step_index: {step_index + 1} / {total_steps}") @@ -316,8 +316,6 @@ def run(self, total_steps=None): if self.progress_callback: self.progress_callback(((step_index + 1) / total_steps) * 100, 100) - self.model.force_cleanup_offload_weights() - return self.model.scheduler.latents, self.model.scheduler.generator def get_custom_shape(self): From f3f2713a86129cc51d173a518d0ed2f45cbd3510 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:41:04 +0800 Subject: [PATCH 04/31] feat(offload): migrate HunyuanVideo block scheduling --- .../feature_caching/transformer_infer.py | 22 ++++++++++++++++--- .../infer/offload/transformer_infer.py | 22 ++++++------------- .../models/networks/hunyuan_video/model.py | 9 ++++---- .../weights/transformer_weights.py | 19 +++++++++++----- 4 files changed, 44 insertions(+), 28 deletions(-) diff --git a/lightx2v/models/networks/hunyuan_video/infer/feature_caching/transformer_infer.py b/lightx2v/models/networks/hunyuan_video/infer/feature_caching/transformer_infer.py index d35362f57..82b65faa2 100755 --- a/lightx2v/models/networks/hunyuan_video/infer/feature_caching/transformer_infer.py +++ b/lightx2v/models/networks/hunyuan_video/infer/feature_caching/transformer_infer.py @@ -5,6 +5,7 @@ import torch import torch.nn.functional as F +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.hunyuan_video.infer.offload.transformer_infer import HunyuanVideo15OffloadTransformerInfer from lightx2v_platform.base.global_var import AI_DEVICE @@ -142,14 +143,18 @@ def __init__(self, config): self.previous_modulated_input_even = None self.previous_residual_even = None - def calculate_should_calc(self, img, vec, block): + def calculate_should_calc(self, img, vec, block, release_img_mod=False): inp = img.clone() vec_ = vec.clone() img_mod_layer = block.img_branch.img_mod if self.config["cpu_offload"]: img_mod_layer.to_cuda() - img_mod1_shift, img_mod1_scale, _, _, _, _ = img_mod_layer.apply(vec_).chunk(6, dim=-1) + try: + img_mod1_shift, img_mod1_scale, _, _, _, _ = img_mod_layer.apply(vec_).chunk(6, dim=-1) + finally: + if release_img_mod: + img_mod_layer.to_cpu() inp = inp.squeeze(0) normed_inp = block.img_branch.img_norm1.apply(inp) modulated_inp = normed_inp * (1 + img_mod1_scale) + img_mod1_shift @@ -190,7 +195,18 @@ def calculate_should_calc(self, img, vec, block): return should_calc def infer(self, weights, infer_module_out): - should_calc = self.calculate_should_calc(infer_module_out.img, infer_module_out.vec, weights.double_blocks[0]) + double_blocks = weights.double_blocks + release_img_mod = ( + self.config["cpu_offload"] + and get_offload_granularity(self.config) == "block" + and 0 not in double_blocks.resident_block_indices + ) + should_calc = self.calculate_should_calc( + infer_module_out.img, + infer_module_out.vec, + double_blocks[0], + release_img_mod, + ) if not should_calc: if self.scheduler.infer_condition: infer_module_out.img += self.previous_residual_odd diff --git a/lightx2v/models/networks/hunyuan_video/infer/offload/transformer_infer.py b/lightx2v/models/networks/hunyuan_video/infer/offload/transformer_infer.py index 4700e30b0..8d10377c1 100755 --- a/lightx2v/models/networks/hunyuan_video/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/hunyuan_video/infer/offload/transformer_infer.py @@ -1,34 +1,26 @@ import torch -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.hunyuan_video.infer.transformer_infer import HunyuanVideo15TransformerInfer -from lightx2v_platform.base.global_var import AI_DEVICE - -torch_device_module = getattr(torch, AI_DEVICE) class HunyuanVideo15OffloadTransformerInfer(HunyuanVideo15TransformerInfer): def __init__(self, config): super().__init__(config) if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": self.infer_func = self.infer_with_blocks_offload elif offload_granularity == "model": self.infer_func = self.infer_without_offload else: raise NotImplementedError - if offload_granularity != "model": - self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) @torch.no_grad() def infer_with_blocks_offload(self, weights, infer_module_out): - for block_idx in range(self.double_blocks_num): + def run_hunyuan_block(block_idx, block): self.block_idx = block_idx - if block_idx == 0: - self.offload_manager.init_first_buffer(weights.double_blocks) - if block_idx < self.double_blocks_num - 1: - self.offload_manager.prefetch_weights(block_idx + 1, weights.double_blocks) - with torch_device_module.stream(self.offload_manager.compute_stream): - infer_module_out.img, infer_module_out.txt = self.infer_double_block(self.offload_manager.cuda_buffers[0], infer_module_out) - self.offload_manager.swap_blocks() + infer_module_out.img, infer_module_out.txt = self.infer_double_block(block, infer_module_out) + return infer_module_out.img, infer_module_out.txt + + self.run_blocks_with_offload(weights.double_blocks, run_hunyuan_block) diff --git a/lightx2v/models/networks/hunyuan_video/model.py b/lightx2v/models/networks/hunyuan_video/model.py index 1840c6368..ba475f10e 100755 --- a/lightx2v/models/networks/hunyuan_video/model.py +++ b/lightx2v/models/networks/hunyuan_video/model.py @@ -19,6 +19,8 @@ class HunyuanVideo15Model(BaseTransformerModel): post_weight_class = HunyuanVideo15PostWeights def __init__(self, model_path, config, device): + if config.get("lazy_load", False): + raise NotImplementedError("HunyuanVideo 1.5 transformer does not support lazy_load.") super().__init__(model_path, config, device) self.remove_keys.extend(["byt5_in", "vision_in"]) self._init_infer_class() @@ -42,8 +44,7 @@ def _init_infer(self): self.transformer_infer = self.transformer_infer_class(self.config) self.post_infer = self.post_infer_class(self.config) self.pre_infer.set_rope(self.transformer_weights.double_blocks[0].rope) - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() + self._init_offload_manager() @torch.no_grad() def _infer_cond_uncond(self, inputs, infer_condition=True): @@ -86,7 +87,7 @@ def _seq_parallel_post_process(self, x): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == 0 and "wan2.2_moe" not in self.config["model_cls"]: self.to_cuda() elif self.offload_granularity != "model": @@ -119,7 +120,7 @@ def infer(self, inputs): # ==================== No CFG ==================== self.scheduler.noise_pred = self._infer_cond_uncond(inputs, infer_condition=True) - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1 and "wan2.2_moe" not in self.config["model_cls"]: self.to_cpu() elif self.offload_granularity != "model": diff --git a/lightx2v/models/networks/hunyuan_video/weights/transformer_weights.py b/lightx2v/models/networks/hunyuan_video/weights/transformer_weights.py index bcf9a95c6..59bd9ae71 100755 --- a/lightx2v/models/networks/hunyuan_video/weights/transformer_weights.py +++ b/lightx2v/models/networks/hunyuan_video/weights/transformer_weights.py @@ -1,6 +1,7 @@ import torch from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.utils.registry_factory import ( ATTN_WEIGHT_REGISTER, LN_WEIGHT_REGISTER, @@ -19,14 +20,20 @@ def __init__(self, config): self.layer_norm_type = config.get("layer_norm_type", "Triton") self.rms_norm_type = config.get("rms_norm_type", "sgl-kernel") self.double_blocks_num = config["mm_double_blocks_depth"] - self.register_offload_buffers(config) - self.add_module("double_blocks", WeightModuleList([MMDoubleStreamBlock(i, self.task, self.config, block_prefix="double_blocks") for i in range(self.double_blocks_num)])) + self.double_blocks = WeightModuleList([MMDoubleStreamBlock(i, self.task, self.config, block_prefix="double_blocks") for i in range(self.double_blocks_num)]) + slot_count = self.register_offload_block_group(config, "double_blocks", self.double_blocks) + self.register_offload_buffers(config, slot_count) + self.add_module("double_blocks", self.double_blocks) self.add_module("final_layer", FinalLayerWeights(self.config)) - def register_offload_buffers(self, config): + def register_offload_buffers(self, config, slot_count): if config["cpu_offload"]: - if config.get("offload_granularity", "block") == "block": - self.offload_blocks_num = 2 + if get_offload_granularity(config) == "block": + self.offload_blocks_num = slot_count + self.offload_block_cuda_buffers = None + self.offload_phase_cuda_buffers = None + if slot_count == 0: + return self.offload_block_cuda_buffers = WeightModuleList( [ MMDoubleStreamBlock( @@ -40,7 +47,7 @@ def register_offload_buffers(self, config): ] ) self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) - self.offload_phase_cuda_buffers = None + self.register_offload_block_buffers("double_blocks", self.offload_block_cuda_buffers) def non_block_weights_to_cuda(self): self.final_layer.to_cuda() From df89e8d0c416d4c8d3c433544a4b846c26eeb392 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:41:31 +0800 Subject: [PATCH 05/31] feat(offload): migrate LongCat Image block scheduling --- .../infer/offload/transformer_infer.py | 87 ++++++++----------- .../models/networks/longcat_image/model.py | 24 ++--- .../weights/transformer_weights.py | 24 ++--- .../longcat_image/longcat_image_runner.py | 2 + 4 files changed, 56 insertions(+), 81 deletions(-) diff --git a/lightx2v/models/networks/longcat_image/infer/offload/transformer_infer.py b/lightx2v/models/networks/longcat_image/infer/offload/transformer_infer.py index 69bed88d9..23c20fa64 100644 --- a/lightx2v/models/networks/longcat_image/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/longcat_image/infer/offload/transformer_infer.py @@ -1,10 +1,7 @@ import torch -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.longcat_image.infer.transformer_infer import LongCatImageTransformerInfer -from lightx2v_platform.base.global_var import AI_DEVICE - -torch_device_module = getattr(torch, AI_DEVICE) class LongCatImageOffloadTransformerInfer(LongCatImageTransformerInfer): @@ -17,12 +14,9 @@ class LongCatImageOffloadTransformerInfer(LongCatImageTransformerInfer): def __init__(self, config): super().__init__(config) if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": self.infer_func = self.infer_with_blocks_offload - if offload_granularity != "model": - self.offload_manager_double = WeightAsyncStreamManager(offload_granularity=offload_granularity) - self.offload_manager_single = WeightAsyncStreamManager(offload_granularity=offload_granularity) def infer_with_blocks_offload(self, blocks, pre_infer_out): """Run transformer inference with block-level offload. @@ -42,52 +36,41 @@ def infer_with_blocks_offload(self, blocks, pre_infer_out): output_seq_len = pre_infer_out.output_seq_len hidden_states = torch.cat([hidden_states, pre_infer_out.input_image_latents], dim=0) - # Stage 1: double blocks offload - # wait for default stream - current_stream = torch_device_module.current_stream() - self.offload_manager_double.compute_stream.wait_stream(current_stream) - for block_idx in range(len(blocks.double_blocks)): + def run_double_block(block_idx, block): + nonlocal encoder_hidden_states, hidden_states self.block_idx = block_idx - - if self.offload_manager_double.need_init_first_buffer: - self.offload_manager_double.init_first_buffer(blocks.double_blocks) - - self.offload_manager_double.prefetch_weights((block_idx + 1) % len(blocks.double_blocks), blocks.double_blocks) - - with torch_device_module.stream(self.offload_manager_double.compute_stream): - encoder_hidden_states, hidden_states = self.infer_double_stream_block( - self.offload_manager_double.cuda_buffers[0], - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - image_rotary_positions, - ) - - self.offload_manager_double.swap_blocks() - - # Stage 2: single blocks offload - # wait for double stream - self.offload_manager_single.compute_stream.wait_stream(self.offload_manager_double.compute_stream) - for block_idx in range(len(blocks.single_blocks)): + encoder_hidden_states, hidden_states = self.infer_double_stream_block( + block, + hidden_states, + encoder_hidden_states, + temb, + image_rotary_emb, + image_rotary_positions, + ) + return encoder_hidden_states, hidden_states + + self.run_blocks_with_offload( + blocks.double_blocks, + run_double_block, + ) + + def run_single_block(block_idx, block): + nonlocal encoder_hidden_states, hidden_states self.block_idx = block_idx - - if self.offload_manager_single.need_init_first_buffer: - self.offload_manager_single.init_first_buffer(blocks.single_blocks) - - self.offload_manager_single.prefetch_weights((block_idx + 1) % len(blocks.single_blocks), blocks.single_blocks) - - with torch_device_module.stream(self.offload_manager_single.compute_stream): - encoder_hidden_states, hidden_states = self.infer_single_stream_block( - self.offload_manager_single.cuda_buffers[0], - hidden_states, - encoder_hidden_states, - temb, - image_rotary_emb, - image_rotary_positions, - ) - - self.offload_manager_single.swap_blocks() + encoder_hidden_states, hidden_states = self.infer_single_stream_block( + block, + hidden_states, + encoder_hidden_states, + temb, + image_rotary_emb, + image_rotary_positions, + ) + return encoder_hidden_states, hidden_states + + self.run_blocks_with_offload( + blocks.single_blocks, + run_single_block, + ) # For I2I task: only return output image latents if output_seq_len is not None: diff --git a/lightx2v/models/networks/longcat_image/model.py b/lightx2v/models/networks/longcat_image/model.py index b1073a8d4..85235d700 100755 --- a/lightx2v/models/networks/longcat_image/model.py +++ b/lightx2v/models/networks/longcat_image/model.py @@ -46,13 +46,7 @@ def _init_infer(self): self.pre_infer = self.pre_infer_class(self.config) self.post_infer = self.post_infer_class(self.config) self.pre_infer.set_rope(self.transformer_weights.double_blocks[0].rope) - if hasattr(self.transformer_infer, "offload_manager_double") and hasattr(self.transformer_infer, "offload_manager_single"): - self._init_offload_manager() - - def _init_offload_manager(self): - """Initialize offload managers for double and single block buffers.""" - self.transformer_infer.offload_manager_double.init_cuda_buffer(blocks_cuda_buffer=self.transformer_weights.offload_double_block_cuda_buffers) - self.transformer_infer.offload_manager_single.init_cuda_buffer(blocks_cuda_buffer=self.transformer_weights.offload_single_block_cuda_buffers) + self._init_offload_manager() @torch.no_grad() def _infer_cond_uncond(self, latents_input, prompt_embeds, infer_condition=True): @@ -108,12 +102,8 @@ def _seq_parallel_post_process(self, noise_pred): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: - if self.offload_granularity == "model": - self.to_cuda() - elif self.offload_granularity == "block": - self.pre_weight.to_cuda() - self.post_weight.to_cuda() + if self.cpu_offload and self.offload_granularity == "model": + self.to_cuda() latents = self.scheduler.latents @@ -166,9 +156,5 @@ def infer(self, inputs): noise_pred = self._infer_cond_uncond(latents, inputs["text_encoder_output"]["prompt_embeds"], infer_condition=True) self.scheduler.noise_pred = noise_pred - if self.cpu_offload: - if self.offload_granularity == "model": - self.to_cpu() - elif self.offload_granularity == "block": - self.pre_weight.to_cpu() - self.post_weight.to_cpu() + if self.cpu_offload and self.offload_granularity == "model": + self.to_cpu() diff --git a/lightx2v/models/networks/longcat_image/weights/transformer_weights.py b/lightx2v/models/networks/longcat_image/weights/transformer_weights.py index 082f30955..e40103bff 100755 --- a/lightx2v/models/networks/longcat_image/weights/transformer_weights.py +++ b/lightx2v/models/networks/longcat_image/weights/transformer_weights.py @@ -1,6 +1,7 @@ import torch from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.utils.registry_factory import ATTN_WEIGHT_REGISTER, LN_WEIGHT_REGISTER, MM_WEIGHT_REGISTER, RMS_WEIGHT_REGISTER, ROPE_REGISTER @@ -349,19 +350,22 @@ def __init__(self, config): # Create weight containers for each block self.double_blocks = WeightModuleList([LongCatImageDoubleBlockWeights(config, i) for i in range(self.num_layers)]) self.single_blocks = WeightModuleList([LongCatImageSingleBlockWeights(config, i) for i in range(self.num_single_layers)]) - self.register_offload_buffers(config) + double_slot_count = self.register_offload_block_group(config, "double_blocks", self.double_blocks) + single_slot_count = self.register_offload_block_group(config, "single_blocks", self.single_blocks) + self.register_offload_buffers(config, double_slot_count, single_slot_count) self.add_module("double_blocks", self.double_blocks) self.add_module("single_blocks", self.single_blocks) - def register_offload_buffers(self, config): - if config.get("cpu_offload", False) and config.get("offload_granularity", "block") == "block": - # Create 2 cuda buffer blocks for double_blocks - self.offload_double_block_cuda_buffers = WeightModuleList([LongCatImageDoubleBlockWeights(config, i, create_cuda_buffer=True) for i in range(2)]) - self.add_module("offload_double_block_cuda_buffers", self.offload_double_block_cuda_buffers) - - # Create 2 cuda buffer blocks for single_blocks - self.offload_single_block_cuda_buffers = WeightModuleList([LongCatImageSingleBlockWeights(config, i, create_cuda_buffer=True) for i in range(2)]) - self.add_module("offload_single_block_cuda_buffers", self.offload_single_block_cuda_buffers) + def register_offload_buffers(self, config, double_slot_count, single_slot_count): + if config.get("cpu_offload", False) and get_offload_granularity(config) == "block": + if double_slot_count: + self.offload_double_block_cuda_buffers = WeightModuleList([LongCatImageDoubleBlockWeights(config, i, create_cuda_buffer=True) for i in range(double_slot_count)]) + self.add_module("offload_double_block_cuda_buffers", self.offload_double_block_cuda_buffers) + self.register_offload_block_buffers("double_blocks", self.offload_double_block_cuda_buffers) + if single_slot_count: + self.offload_single_block_cuda_buffers = WeightModuleList([LongCatImageSingleBlockWeights(config, i, create_cuda_buffer=True) for i in range(single_slot_count)]) + self.add_module("offload_single_block_cuda_buffers", self.offload_single_block_cuda_buffers) + self.register_offload_block_buffers("single_blocks", self.offload_single_block_cuda_buffers) def to_cuda(self, non_blocking=True): for block in self.double_blocks: diff --git a/lightx2v/models/runners/longcat_image/longcat_image_runner.py b/lightx2v/models/runners/longcat_image/longcat_image_runner.py index fe218d279..7572edb47 100755 --- a/lightx2v/models/runners/longcat_image/longcat_image_runner.py +++ b/lightx2v/models/runners/longcat_image/longcat_image_runner.py @@ -7,6 +7,7 @@ from lightx2v.models.input_encoders.hf.longcat.longcat_text_encoder import LongCatImageTextEncoder from lightx2v.models.networks.longcat_image.model import LongCatImageTransformerModel +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.default_runner import DefaultRunner from lightx2v.models.schedulers.longcat_image.scheduler import LongCatImageScheduler from lightx2v.models.video_encoders.hf.longcat_image.vae import LongCatImageVAE @@ -279,6 +280,7 @@ def run_vae_decoder(self, latents): gc.collect() return images + @keep_transformer_weights_loaded def run(self, total_steps=None): if total_steps is None: total_steps = self.model.scheduler.infer_steps From 727968ac1f79822536c7b461d8435f46790ad0cf Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:41:34 +0800 Subject: [PATCH 06/31] feat(offload): migrate LTX2 block scheduling --- .../ltx2/infer/offload/transformer_infer.py | 58 +++++++++++-------- lightx2v/models/networks/ltx2/model.py | 7 +-- .../ltx2/weights/transformer_weights.py | 50 +++++++++++++--- lightx2v/models/runners/ltx2/ltx2_runner.py | 15 +++-- 4 files changed, 90 insertions(+), 40 deletions(-) diff --git a/lightx2v/models/networks/ltx2/infer/offload/transformer_infer.py b/lightx2v/models/networks/ltx2/infer/offload/transformer_infer.py index 64a30a759..490d7fc0b 100755 --- a/lightx2v/models/networks/ltx2/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/ltx2/infer/offload/transformer_infer.py @@ -6,7 +6,7 @@ import torch -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.ltx2.infer.module_io import LTX2PreInferModuleOutput from lightx2v.models.networks.ltx2.infer.transformer_infer import LTX2TransformerInfer from lightx2v_platform.base.global_var import AI_DEVICE @@ -37,23 +37,18 @@ def __init__(self, config): super().__init__(config) if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": # Use block-level offloading self.infer_func = self.infer_with_blocks_offload - self.offload_manager = WeightAsyncStreamManager(offload_granularity="block") elif offload_granularity == "model": # Model movement is handled by LTX2Model. self.infer_func = self.infer_without_offload else: raise ValueError(f"Unsupported offload_granularity: {offload_granularity}") - # Initialize lazy loading if enabled self.lazy_load = self.config.get("lazy_load", False) - if self.lazy_load and offload_granularity == "block": - num_workers = self.config.get("num_disk_workers", 4) - self.offload_manager.init_lazy_load(num_workers=num_workers) else: # No offloading self.infer_func = self.infer_without_offload @@ -100,24 +95,15 @@ def infer_with_blocks_offload(self, weights, pre_infer_out: LTX2PreInferModuleOu blocks = weights.blocks - # Process all transformer blocks with offloading - for block_idx in range(len(blocks)): - # Initialize first buffer on first iteration - if self.offload_manager.need_init_first_buffer: - self.offload_manager.init_first_buffer(blocks) - - # Prefetch next block to GPU (async, in background stream) - # Use modulo to handle wrap-around for warmup - next_block_idx = (block_idx + 1) % len(blocks) - self.offload_manager.prefetch_weights(next_block_idx, blocks) + if self.lazy_load: + vx, ax = self.infer_with_lazy_blocks_offload(blocks, vx, ax, pre_infer_out) + else: - # Compute current block (using compute stream) - with torch_device_module.stream(self.offload_manager.compute_stream): - # Use the block currently in cuda_buffers[0] - current_block = self.offload_manager.cuda_buffers[0] + def run_ltx_block(block_idx, block): + nonlocal vx, ax vx, ax = self.run_block( block_idx, - current_block, + block, vx, ax, pre_infer_out, @@ -126,9 +112,9 @@ def infer_with_blocks_offload(self, weights, pre_infer_out: LTX2PreInferModuleOu self._mm_skip_a2v, self._mm_skip_v2a, ) + return vx, ax - # Swap buffers: cuda_buffers[1] (prefetched) -> cuda_buffers[0] (current) - self.offload_manager.swap_blocks() + self.run_blocks_with_offload(blocks, run_ltx_block) # Clean up if needed if self.clean_cuda_cache: @@ -140,6 +126,30 @@ def infer_with_blocks_offload(self, weights, pre_infer_out: LTX2PreInferModuleOu return vx, ax, pre_infer_out.video_args.embedded_timestep, pre_infer_out.audio_args.embedded_timestep + def infer_with_lazy_blocks_offload(self, blocks, vx, ax, pre_infer_out): + manager = self.get_block_offload_manager(blocks) + for block_idx in range(len(blocks)): + next_block_idx = (block_idx + 1) % len(blocks) + manager.start_prefetch_block(next_block_idx) + if manager.need_init_first_buffer: + manager.init_first_buffer(blocks) + manager.swap_cpu_buffers() + manager.prefetch_weights(next_block_idx, blocks) + with torch_device_module.stream(manager.compute_stream): + vx, ax = self.run_block( + block_idx, + manager.cuda_buffers[0], + vx, + ax, + pre_infer_out, + block_idx in self._mm_skip_video_self_blocks, + block_idx in self._mm_skip_audio_self_blocks, + self._mm_skip_a2v, + self._mm_skip_v2a, + ) + manager.swap_blocks() + return vx, ax + def infer(self, weights, pre_infer_out: LTX2PreInferModuleOutput): """ Main inference entry point. diff --git a/lightx2v/models/networks/ltx2/model.py b/lightx2v/models/networks/ltx2/model.py index 7c6c6f37a..4b1ecc7ab 100755 --- a/lightx2v/models/networks/ltx2/model.py +++ b/lightx2v/models/networks/ltx2/model.py @@ -343,8 +343,7 @@ def _init_infer(self): self.pre_infer = self.pre_infer_class(self.config) self.post_infer = self.post_infer_class(self.config) self.transformer_infer = self.transformer_infer_class(self.config) - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() + self._init_offload_manager() @torch.no_grad() def _infer_cond_uncond(self, inputs, infer_condition=True, mm_perturb=None): @@ -667,7 +666,7 @@ def _seq_parallel_post_process(self, x, original_length=None): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == 0 and "wan2.2_moe" not in self.config["model_cls"]: self.to_cuda() elif self.offload_granularity != "model": @@ -713,7 +712,7 @@ def infer(self, inputs): self.scheduler.v_noise_pred = v_noise_pred self.scheduler.a_noise_pred = a_noise_pred - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1 and "wan2.2_moe" not in self.config["model_cls"]: self.to_cpu() elif self.offload_granularity != "model": diff --git a/lightx2v/models/networks/ltx2/weights/transformer_weights.py b/lightx2v/models/networks/ltx2/weights/transformer_weights.py index 96d6754d8..68d13225e 100755 --- a/lightx2v/models/networks/ltx2/weights/transformer_weights.py +++ b/lightx2v/models/networks/ltx2/weights/transformer_weights.py @@ -2,6 +2,7 @@ import torch.distributed as dist from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.utils.registry_factory import ( ATTN_WEIGHT_REGISTER, MM_WEIGHT_REGISTER, @@ -30,7 +31,6 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): assert not config["cpu_offload"] self.lazy_load = self.config.get("lazy_load", False) self.skip_fp8_block_index = self.config.get("skip_fp8_block_index", []) - self.register_offload_buffers(config, lazy_load_path, lora_path) self.blocks = WeightModuleList( [ LTX2TransformerBlock( @@ -47,18 +47,31 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): for i in range(self.blocks_num) ] ) + slot_count = self.register_offload_block_group(config, "blocks", self.blocks) + self.register_offload_buffers(config, lazy_load_path, lora_path, slot_count) self.add_module("blocks", self.blocks) - def register_offload_buffers(self, config, lazy_load_path, lora_path): + def register_offload_buffers(self, config, lazy_load_path, lora_path, slot_count): if config["cpu_offload"]: - if config["offload_granularity"] == "block": - self.offload_blocks_num = 2 + if get_offload_granularity(config) == "block": + self.offload_blocks_num = slot_count + self.offload_block_cuda_buffers = None + self.offload_block_cpu_buffers = None + self.offload_phase_cuda_buffers = None + self.offload_phase_cpu_buffers = None + if slot_count == 0: + return + offloaded_indices = self.blocks.offload_block_indices + offloaded_mm_types = {self.mm_type if block_index not in self.skip_fp8_block_index else "Default" for block_index in offloaded_indices} + if len(offloaded_mm_types) > 1: + raise NotImplementedError("LTX2 block offload requires all non-resident blocks to use the same weight type") + offloaded_mm_type = next(iter(offloaded_mm_types)) self.offload_block_cuda_buffers = WeightModuleList( [ LTX2TransformerBlock( - block_index=i, + block_index=offloaded_indices[i], task=self.task, - mm_type=self.mm_type if i not in self.skip_fp8_block_index else "Default", + mm_type=offloaded_mm_type, config=self.config, create_cuda_buffer=True, create_cpu_buffer=False, @@ -70,7 +83,30 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) - self.offload_phase_cuda_buffers = None + if self.lazy_load: + self.offload_block_cpu_buffers = WeightModuleList( + [ + LTX2TransformerBlock( + block_index=offloaded_indices[i], + task=self.task, + mm_type=offloaded_mm_type, + config=self.config, + create_cuda_buffer=False, + create_cpu_buffer=True, + block_prefix="transformer_blocks", + lazy_load=self.lazy_load, + lazy_load_path=lazy_load_path, + lora_path=lora_path, + ) + for i in range(self.offload_blocks_num) + ] + ) + self.add_module("offload_block_cpu_buffers", self.offload_block_cpu_buffers) + self.register_offload_block_buffers( + "blocks", + self.offload_block_cuda_buffers, + self.offload_block_cpu_buffers, + ) class LTX2TransformerBlock(WeightModule): diff --git a/lightx2v/models/runners/ltx2/ltx2_runner.py b/lightx2v/models/runners/ltx2/ltx2_runner.py index 75af1fec0..35df49090 100755 --- a/lightx2v/models/runners/ltx2/ltx2_runner.py +++ b/lightx2v/models/runners/ltx2/ltx2_runner.py @@ -6,9 +6,11 @@ import torch.distributed as dist from lightx2v.common.kvcache.utils import causal_chunk_token_range +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.input_encoders.hf.ltx2.model import LTX2TextEncoder from lightx2v.models.networks.lora_adapter import LoraAdapter from lightx2v.models.networks.ltx2.model import LTX2ARModel, LTX2Model +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.default_runner import DefaultRunner from lightx2v.models.schedulers.ltx2.scheduler import LTX2ARScheduler, LTX2Scheduler, LatentState from lightx2v.models.video_encoders.hf.ltx2.audio_vae.audio_vae import encode_audio @@ -120,7 +122,7 @@ def _run_warmup(self): scheduler = self.model.scheduler stage1_infer_steps = scheduler.infer_steps use_upsampler = bool(self.config.get("use_upsampler")) - model_offload = self.config.get("cpu_offload", False) and self.config.get("offload_granularity") == "model" + model_offload = self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "model" _, spatial_scale_h, spatial_scale_w = self.config["vae_scale_factors"] upsample_scale = 2 if use_upsampler else 1 stage_count = 2 if use_upsampler else 1 @@ -155,10 +157,11 @@ def _run_warmup(self): # unpatchifies latents so they can continue to Stage 2/VAE. last_step = scheduler.infer_steps - 1 step_indices = (0,) if last_step == 0 else (0, last_step) - for step_index in step_indices: - scheduler.step_pre(step_index=step_index) - self.model.infer(self.inputs) - scheduler.step_post() + with self.transformer_offload_session(): + for step_index in step_indices: + scheduler.step_pre(step_index=step_index) + self.model.infer(self.inputs) + scheduler.step_post() v_latent = scheduler.video_latent_state.latent a_latent = scheduler.audio_latent_state.latent @@ -1121,6 +1124,7 @@ def process_images_after_vae_decoder(self): logger.info(f"✅ Video saved successfully to: {out_path} ✅") return {"video": None} + @keep_transformer_weights_loaded def run_segment(self, segment_idx=0, stage_name=None, cleanup_inputs=None): """ Run denoising loop for a segment. @@ -1306,6 +1310,7 @@ def _load_ar_chunk(self, video_start, video_end, audio_start, audio_end): self.model.scheduler.mm_last_a_pred = None self.model.set_ar_chunk(video_start=video_start, audio_start=audio_start) + @keep_transformer_weights_loaded def run_segment(self, segment_idx=0, stage_name=None, cleanup_inputs=None): infer_steps = self.model.scheduler.infer_steps video_chunks = [] From f72faa413dd48a47e1ddf1b30db115601f789040 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:42:03 +0800 Subject: [PATCH 07/31] feat(offload): migrate MiniMax H3 block scheduling --- ...ax_h3_fp8_4step_5090_with_fp8_vae_sla.json | 1 - .../infer/offload/transformer_infer.py | 68 +++++-------------- lightx2v/models/networks/minimax_h3/model.py | 37 ++++------ .../minimax_h3/weights/transformer_weights.py | 13 ++-- .../runners/minimax_h3/minimax_h3_runner.py | 25 +++---- 5 files changed, 49 insertions(+), 95 deletions(-) diff --git a/configs/minimax_h3/dmd/minimax_h3_fp8_4step_5090_with_fp8_vae_sla.json b/configs/minimax_h3/dmd/minimax_h3_fp8_4step_5090_with_fp8_vae_sla.json index 3676c0136..10ace47db 100755 --- a/configs/minimax_h3/dmd/minimax_h3_fp8_4step_5090_with_fp8_vae_sla.json +++ b/configs/minimax_h3/dmd/minimax_h3_fp8_4step_5090_with_fp8_vae_sla.json @@ -8,7 +8,6 @@ "enable_cfg": false, "cpu_offload": true, "offload_granularity": "block", - "dit_prepost_resident": true, "text_encoder_cpu_offload": true, "text_encoder_offload_granularity": "block", "vae_cpu_offload": false, diff --git a/lightx2v/models/networks/minimax_h3/infer/offload/transformer_infer.py b/lightx2v/models/networks/minimax_h3/infer/offload/transformer_infer.py index 2093c6a6d..6167a21c8 100644 --- a/lightx2v/models/networks/minimax_h3/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/minimax_h3/infer/offload/transformer_infer.py @@ -1,10 +1,5 @@ -import torch - -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.minimax_h3.infer.transformer_infer import MiniMaxH3TransformerInfer -from lightx2v_platform.base.global_var import AI_DEVICE - -torch_device_module = getattr(torch, AI_DEVICE) class MiniMaxH3OffloadTransformerInfer(MiniMaxH3TransformerInfer): @@ -12,61 +7,32 @@ class MiniMaxH3OffloadTransformerInfer(MiniMaxH3TransformerInfer): def __init__(self, config): super().__init__(config) - offload_granularity = config.get("offload_granularity", "model") + offload_granularity = get_offload_granularity(config) if offload_granularity == "block": - self.offload_manager = WeightAsyncStreamManager(offload_granularity="block") self.infer_func = self.infer_with_blocks_offload elif offload_granularity != "model": raise NotImplementedError(f"MiniMax-H3 does not support offload_granularity={offload_granularity!r}") def get_compile_block_key(self, block_idx, block): - # block offload - if hasattr(self, "offload_manager"): + if self.has_block_offload_manager(): return id(block) - # model offload return super().get_compile_block_key(block_idx, block) - def _prefetch_weights_without_adaln(self, block_index, blocks): - with torch_device_module.stream(self.offload_manager.cuda_load_stream): - if hasattr(self.offload_manager, "cpu_buffers"): - source_block = self.offload_manager.cpu_buffers[0] - else: - source_block = blocks[block_index] - block_state_dict = source_block.state_dict() - weights_without_adaln = {} - for name, tensor in block_state_dict.items(): - if ".adaln_proj." not in name: - weights_without_adaln[name] = tensor - self.offload_manager.cuda_buffers[1].load_state_dict(weights_without_adaln, block_index) + @staticmethod + def _weights_without_adaln(_block_index, state_dict): + return {name: tensor for name, tensor in state_dict.items() if ".adaln_proj." not in name} def infer_with_blocks_offload(self, blocks, hidden_states, pre_infer_out): - num_blocks = len(blocks) - if self.use_adaln_cache and not self._adaln_cache_hit: - # The previous forward may have prefetched block 0 without AdaLN. - # Reload the full block when the current timestep misses. - self.offload_manager.need_init_first_buffer = True - current_stream = torch_device_module.current_stream() - self.offload_manager.compute_stream.wait_stream(current_stream) - - for block_index in range(num_blocks): - if self.offload_manager.need_init_first_buffer: - self.offload_manager.init_first_buffer(blocks) - - next_block_index = (block_index + 1) % num_blocks - if self.use_adaln_cache and self._adaln_cache_hit: - self._prefetch_weights_without_adaln(next_block_index, blocks) - else: - self.offload_manager.prefetch_weights(next_block_index, blocks) - block = self.offload_manager.cuda_buffers[0] + def run_h3_block(block_index, block): + nonlocal hidden_states self.block_idx = block_index - if AI_DEVICE == "xpu": - # Match Wan's XPU offload path: overlap the next weight copy on - # the load stream with current-block compute on the default - # stream, then let swap_blocks() perform the device-wide sync. - hidden_states = self.run_block(block_index, block, hidden_states, pre_infer_out) - else: - with torch_device_module.stream(self.offload_manager.compute_stream): - hidden_states = self.run_block(block_index, block, hidden_states, pre_infer_out) - self.offload_manager.swap_blocks() - + hidden_states = self.run_block(block_index, block, hidden_states, pre_infer_out) + return hidden_states + + state_dict_transform = self._weights_without_adaln if self.use_adaln_cache and self._adaln_cache_hit else None + self.run_blocks_with_offload( + blocks, + run_h3_block, + state_dict_transform=state_dict_transform, + ) return hidden_states diff --git a/lightx2v/models/networks/minimax_h3/model.py b/lightx2v/models/networks/minimax_h3/model.py index 19ca82476..8068ded65 100644 --- a/lightx2v/models/networks/minimax_h3/model.py +++ b/lightx2v/models/networks/minimax_h3/model.py @@ -7,6 +7,7 @@ from loguru import logger from safetensors import safe_open +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.base_model import BaseTransformerModel from lightx2v.models.networks.minimax_h3.infer.module_io import MiniMaxH3SequenceParallelState from lightx2v.models.networks.minimax_h3.infer.offload import MiniMaxH3OffloadTransformerInfer @@ -48,10 +49,6 @@ class MiniMaxH3Model(BaseTransformerModel): def __init__(self, model_path, config, device, lora_path=None, lora_strength=1.0, lora_alpha=None): self.lora_alpha = lora_alpha - self.block_offload = config.get("cpu_offload", False) and config.get("offload_granularity", "model") == "block" - # Model offload moves pre/blocks/post together. Pre/post residency only applies - # to block offload and is ignored otherwise. - self.prepost_resident = self.block_offload and config.get("dit_prepost_resident", False) if GET_DTYPE() != torch.bfloat16: raise ValueError( "MiniMax-H3 requires DTYPE=BF16. The native loader preserves the released checkpoint's 626 BF16 tensors and 12 FP32 projection/time/head tensors without dtype conversion." @@ -66,7 +63,7 @@ def __init__(self, model_path, config, device, lora_path=None, lora_strength=1.0 raise ValueError("MiniMax-H3 quantized inference requires dit_quantized_ckpt") elif config.get("dit_quant_scheme", "Default") != "Default": raise ValueError("MiniMax-H3 dit_quant_scheme requires a dit_quantized_ckpt") - if config.get("cpu_offload", False) and config.get("offload_granularity", "model") not in {"model", "block"}: + if config.get("cpu_offload", False) and get_offload_granularity(config) not in {"model", "block"}: raise NotImplementedError("MiniMax-H3 supports model and block CPU offload") if config.get("attn_type") == "sol_attn": reorder = str(config.get("sol_attn_setting", {}).get("reorder", "none")).lower() @@ -269,9 +266,7 @@ def _register_lora(self, lora_path, strength): self._register_dynamic_lora_weights(lora_weights, strength) self.lora_path = lora_path self.lora_strength = float(strength) - offload_manager = getattr(getattr(self, "transformer_infer", None), "offload_manager", None) - if offload_manager is not None: - offload_manager.need_init_first_buffer = True + self._reset_offload_staging_buffers() def _remove_lora(self): super()._remove_lora() @@ -287,9 +282,14 @@ def _update_lora(self, lora_path, strength, alpha=None): self._register_dynamic_lora_weights(lora_weights, strength) self.lora_path = lora_path self.lora_strength = float(strength) - offload_manager = getattr(getattr(self, "transformer_infer", None), "offload_manager", None) - if offload_manager is not None: - offload_manager.need_init_first_buffer = True + self._reset_offload_staging_buffers() + + def _reset_offload_staging_buffers(self): + transformer_infer = getattr(self, "transformer_infer", None) + if transformer_infer is None: + return + for manager in transformer_infer.get_offload_managers(): + manager.need_init_first_buffer = True def _validate_tensor_parallel_config(self): if not self.use_tp: @@ -476,8 +476,7 @@ def _init_infer(self): self.pre_infer = self.pre_infer_class(self.config) self.transformer_infer = self.transformer_infer_class(self.config) self.post_infer = self.post_infer_class(self.config) - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() + self._init_offload_manager() @torch.no_grad() def _infer_cond_uncond(self, inputs, infer_condition=True): @@ -494,16 +493,9 @@ def _infer_cond_uncond(self, inputs, infer_condition=True): @torch.no_grad() def infer(self, inputs): - prepost_offload = self.block_offload and not self.prepost_resident - if prepost_offload and self.scheduler.step_index == 0: - self.pre_weight.to_cuda() - self.post_weight.to_cuda() output = self._infer_cond_uncond(inputs, infer_condition=True) self.scheduler.video_noise_pred = output.video self.scheduler.audio_noise_pred = output.audio - if prepost_offload and self.scheduler.step_index == self.scheduler.infer_steps - 1: - self.pre_weight.to_cpu() - self.post_weight.to_cpu() @torch.no_grad() def _seq_parallel_pre_process(self, pre_infer_out): @@ -565,7 +557,4 @@ def _seq_parallel_post_process(self, output, pre_infer_out): def to_cpu(self): super().to_cpu() - if hasattr(self.transformer_infer, "offload_manager"): - # Full teardown moves the active aliases away from the persistent - # device buffers. Force buffer 0 to be populated again next run. - self.transformer_infer.offload_manager.need_init_first_buffer = True + self._reset_offload_staging_buffers() diff --git a/lightx2v/models/networks/minimax_h3/weights/transformer_weights.py b/lightx2v/models/networks/minimax_h3/weights/transformer_weights.py index 9634eb595..7e504b353 100644 --- a/lightx2v/models/networks/minimax_h3/weights/transformer_weights.py +++ b/lightx2v/models/networks/minimax_h3/weights/transformer_weights.py @@ -2,6 +2,7 @@ import torch.distributed as dist from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.minimax_h3.infer.triton_ops import MiniMaxH3TritonRope # noqa: F401 from lightx2v.utils.registry_factory import ATTN_WEIGHT_REGISTER, MM_WEIGHT_REGISTER, RMS_WEIGHT_REGISTER, ROPE_REGISTER @@ -134,10 +135,12 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): "MiniMax-H3 reads the official sharded checkpoint directly; disk lazy_load requires a converted block-sharded checkpoint and is not supported yet. Use lazy_load=false with model or block CPU offload." ) self.blocks = WeightModuleList([MiniMaxH3TransformerBlockWeights(i, config) for i in range(int(config.get("num_layers", 50)))]) - if config.get("cpu_offload", False) and config.get("offload_granularity", "model") == "block": - self.offload_block_cuda_buffers = WeightModuleList([MiniMaxH3TransformerBlockWeights(i, config, create_cuda_buffer=True) for i in range(2)]) - # Register device buffers before source blocks: buffer allocation - # needs checkpoint metadata that normal CPU loading consumes. - self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) + slot_count = self.register_offload_block_group(config, "blocks", self.blocks) + if config.get("cpu_offload", False) and get_offload_granularity(config) == "block": self.offload_phase_cuda_buffers = None + if slot_count: + self.offload_block_cuda_buffers = WeightModuleList([MiniMaxH3TransformerBlockWeights(i, config, create_cuda_buffer=True) for i in range(slot_count)]) + # Buffer allocation needs checkpoint metadata that normal CPU loading consumes. + self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) + self.register_offload_block_buffers("blocks", self.offload_block_cuda_buffers) self.add_module("blocks", self.blocks) diff --git a/lightx2v/models/runners/minimax_h3/minimax_h3_runner.py b/lightx2v/models/runners/minimax_h3/minimax_h3_runner.py index 7cd2698ea..9d1e603cb 100644 --- a/lightx2v/models/runners/minimax_h3/minimax_h3_runner.py +++ b/lightx2v/models/runners/minimax_h3/minimax_h3_runner.py @@ -7,6 +7,7 @@ from PIL import Image, ImageOps from loguru import logger +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.audio_encoders.hf.minimax_h3 import MiniMaxH3AudioVAE from lightx2v.models.input_encoders.hf.minimax_h3 import MiniMaxH3Qwen3VLTextEncoder from lightx2v.models.networks.minimax_h3.lora import MiniMaxH3LoraAdapter @@ -37,6 +38,7 @@ resolve_reference_image_size, trim_reference_num_frames, ) +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.default_runner import DefaultRunner from lightx2v.models.schedulers.minimax_h3 import MiniMaxH3Scheduler from lightx2v.models.video_encoders.hf.ltx2.audio_vae.ops import Audio @@ -108,10 +110,6 @@ def __init__(self, config): def init_modules(self): super().init_modules() self.run_input_encoder = self._run_input_encoder_local_h3 - if self.model.prepost_resident: - self.model.pre_weight.to_cuda() - self.model.post_weight.to_cuda() - logger.info("MiniMax-H3 pre/post weights will remain on the accelerator across requests") @ProfilingContext4DebugL1("Warmup") def run_warmup(self): @@ -135,10 +133,11 @@ def run_warmup(self): self.inputs = self._run_input_encoder_local_h3() self.init_run() - for step_index in range(min(self._WARMUP_STEP_COUNT, self.scheduler.infer_steps)): - self.scheduler.step_pre(step_index) - self.model.infer(self.inputs) - self.scheduler.step_post() + with self.transformer_offload_session(): + for step_index in range(min(self._WARMUP_STEP_COUNT, self.scheduler.infer_steps)): + self.scheduler.step_pre(step_index) + self.model.infer(self.inputs) + self.scheduler.step_post() video_rows = self.scheduler.video_latents audio_rows = self.scheduler.audio_latents @@ -561,13 +560,14 @@ def init_run(self): ) if not self.config.get("cpu_offload", False): logger.info("MiniMax-H3 transformer is resident on the accelerator") - elif self.config.get("offload_granularity", "model") == "model": + elif get_offload_granularity(self.config) == "model": logger.info("Moving the native MiniMax-H3 transformer to the accelerator") self.model.to_cuda() else: logger.info("MiniMax-H3 block offload enabled; keeping source blocks on CPU and using two accelerator buffers") torch_device_module.synchronize() + @keep_transformer_weights_loaded def run_segment(self, segment_idx=0): infer_steps = self.scheduler.infer_steps for step_index in range(infer_steps): @@ -593,11 +593,8 @@ def run_segment(self, segment_idx=0): def _offload_transformer(self): if not self.config.get("cpu_offload", False): return - if self.model.block_offload: - if not self.model.prepost_resident: - logger.info("Offloading MiniMax-H3 pre/post weights; retaining the two block-offload device buffers") - self.model.pre_weight.to_cpu() - self.model.post_weight.to_cpu() + if get_offload_granularity(self.config) == "block": + logger.info("MiniMax-H3 block-offload session released resident weights; retaining staging buffers") else: logger.info("Offloading MiniMax-H3 transformer before VAE decode") self.model.to_cpu() From 073ebbffe08d6544f0a0fd9dbd10d1f8f6e254cc Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:42:14 +0800 Subject: [PATCH 08/31] feat(offload): migrate Qwen Image block scheduling --- .../infer/offload/transformer_infer.py | 62 ++++++++++++++++--- lightx2v/models/networks/qwen_image/model.py | 7 +-- .../qwen_image/weights/transformer_weights.py | 27 +++++--- .../runners/qwen_image/qwen_image_runner.py | 15 +++-- 4 files changed, 83 insertions(+), 28 deletions(-) diff --git a/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py b/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py index 8dcab3830..79d183917 100755 --- a/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/qwen_image/infer/offload/transformer_infer.py @@ -1,5 +1,6 @@ import torch +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.common.offload.manager import WeightAsyncStreamManager from lightx2v.models.networks.qwen_image.infer.transformer_infer import ( QwenImageTransformerInfer, @@ -16,17 +17,16 @@ def __init__(self, config): self.phases_num = 4 if self.config.get("cpu_offload", False): self.offload_ratio = self.config.get("offload_ratio", 1) - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": 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) - if self.lazy_load: + if self.lazy_load and offload_granularity == "phase": self.offload_manager.init_lazy_load(num_workers=self.config.get("num_disk_workers", 4)) def get_compile_block_key(self, _block_idx, block): @@ -145,22 +145,64 @@ def infer_with_blocks_offload( image_rotary_positions, modulate_index, ): + if self.lazy_load: + return self.infer_with_lazy_blocks_offload( + blocks, + hidden_states, + encoder_hidden_states, + temb_img_silu, + temb_txt_silu, + image_rotary_emb, + image_rotary_positions, + modulate_index, + ) + + def run_qwen_block(block_idx, block): + nonlocal encoder_hidden_states, hidden_states + 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, + ) + return encoder_hidden_states, hidden_states + + self.run_blocks_with_offload(blocks, run_qwen_block) + return hidden_states + + def infer_with_lazy_blocks_offload( + self, + blocks, + hidden_states, + encoder_hidden_states, + temb_img_silu, + temb_txt_silu, + image_rotary_emb, + image_rotary_positions, + modulate_index, + ): + manager = self.get_block_offload_manager(blocks) for block_idx in range(self.num_blocks): if self.lazy_load: next_prefetch = (block_idx + 1) % self.num_blocks - self.offload_manager.start_prefetch_block(next_prefetch) + manager.start_prefetch_block(next_prefetch) if block_idx == 0: - self.offload_manager.init_first_buffer(blocks) + manager.init_first_buffer(blocks) if self.lazy_load: - self.offload_manager.swap_cpu_buffers() - self.offload_manager.prefetch_weights((block_idx + 1) % self.num_blocks, blocks) + manager.swap_cpu_buffers() + manager.prefetch_weights((block_idx + 1) % self.num_blocks, blocks) - with torch_device_module.stream(self.offload_manager.compute_stream): + with torch_device_module.stream(manager.compute_stream): encoder_hidden_states, hidden_states = self.run_block( block_idx, - self.offload_manager.cuda_buffers[0], + manager.cuda_buffers[0], hidden_states, encoder_hidden_states, temb_img_silu, @@ -170,6 +212,6 @@ def infer_with_blocks_offload( modulate_index, ) - self.offload_manager.swap_blocks() + manager.swap_blocks() return hidden_states diff --git a/lightx2v/models/networks/qwen_image/model.py b/lightx2v/models/networks/qwen_image/model.py index c8619b1fd..f9afa01b6 100755 --- a/lightx2v/models/networks/qwen_image/model.py +++ b/lightx2v/models/networks/qwen_image/model.py @@ -46,8 +46,7 @@ def _init_infer(self): img_rope=first_block.compute_phases[0].rope, txt_rope=first_block.compute_phases[1].rope, ) - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() + self._init_offload_manager() @torch.no_grad() def _infer_cond_uncond(self, latents_input, prompt_embeds, infer_condition=True): @@ -93,7 +92,7 @@ def _seq_parallel_post_process(self, noise_pred): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == 0: self.to_cuda() elif self.offload_granularity != "model": @@ -147,7 +146,7 @@ def infer(self, inputs): noise_pred = noise_pred[:, : latents.size(1)] self.scheduler.noise_pred = noise_pred - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1: self.to_cpu() elif self.offload_granularity != "model": diff --git a/lightx2v/models/networks/qwen_image/weights/transformer_weights.py b/lightx2v/models/networks/qwen_image/weights/transformer_weights.py index 229f2b839..254b372c1 100755 --- a/lightx2v/models/networks/qwen_image/weights/transformer_weights.py +++ b/lightx2v/models/networks/qwen_image/weights/transformer_weights.py @@ -1,6 +1,7 @@ import torch from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.utils.registry_factory import ( ATTN_WEIGHT_REGISTER, LN_WEIGHT_REGISTER, @@ -36,13 +37,21 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): ) for i in range(self.blocks_num) ) - self.register_offload_buffers(config, lazy_load_path, lora_path) + slot_count = self.register_offload_block_group(config, "blocks", blocks) + self.register_offload_buffers(config, lazy_load_path, lora_path, slot_count) self.add_module("blocks", blocks) - def register_offload_buffers(self, config, lazy_load_path, lora_path): + def register_offload_buffers(self, config, lazy_load_path, lora_path, slot_count): if config["cpu_offload"]: - if config["offload_granularity"] == "block": - self.offload_blocks_num = 2 + offload_granularity = get_offload_granularity(config) + if offload_granularity == "block": + self.offload_blocks_num = slot_count + self.offload_block_cuda_buffers = None + self.offload_block_cpu_buffers = None + self.offload_phase_cuda_buffers = None + self.offload_phase_cpu_buffers = None + if slot_count == 0: + return self.offload_block_cuda_buffers = WeightModuleList( [ QwenImageTransformerAttentionBlock( @@ -60,9 +69,7 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) - self.offload_phase_cuda_buffers = None if self.lazy_load: - self.offload_blocks_num = 2 self.offload_block_cpu_buffers = WeightModuleList( [ QwenImageTransformerAttentionBlock( @@ -81,9 +88,13 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cpu_buffers", self.offload_block_cpu_buffers) - self.offload_phase_cpu_buffers = None + self.register_offload_block_buffers( + "blocks", + self.offload_block_cuda_buffers, + self.offload_block_cpu_buffers, + ) - elif config["offload_granularity"] == "phase": + elif offload_granularity == "phase": self.offload_phase_cuda_buffers = QwenImageTransformerAttentionBlock( 0, self.task, diff --git a/lightx2v/models/runners/qwen_image/qwen_image_runner.py b/lightx2v/models/runners/qwen_image/qwen_image_runner.py index a659b1564..944dc3c87 100755 --- a/lightx2v/models/runners/qwen_image/qwen_image_runner.py +++ b/lightx2v/models/runners/qwen_image/qwen_image_runner.py @@ -6,10 +6,12 @@ from PIL import Image from loguru import logger +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.disagg.disagg_mixin import DisaggMixin from lightx2v.models.input_encoders.hf.qwen25.qwen25_vlforconditionalgeneration import Qwen25_VLForConditionalGeneration_TextEncoder from lightx2v.models.networks.lora_adapter import LoraAdapter from lightx2v.models.networks.qwen_image.model import QwenImageTransformerModel +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.default_runner import DefaultRunner from lightx2v.models.schedulers.qwen_image.scheduler import QwenImageScheduler from lightx2v.models.video_encoders.hf.qwen_image.vae import AutoencoderKLQwenImageVAE @@ -107,13 +109,14 @@ def _run_warmup(self): t2i_text_cache = self._prepare_warmup_inputs(height, width, t2i_text_cache) scheduler.generator = None scheduler.prepare(self.input_info) - scheduler.step_pre(step_index=0) - self.model.infer(self.inputs) - scheduler.step_post() + with self.transformer_offload_session(): + scheduler.step_pre(step_index=0) + self.model.infer(self.inputs) + scheduler.step_post() self.run_vae_decoder(scheduler.latents) torch_device_module.synchronize() finally: - if self.config.get("cpu_offload", False) and self.config.get("offload_granularity") == "model": + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "model": self.model.to_cpu() self.clear_warmup_state() self.input_info = None @@ -161,8 +164,7 @@ def clean_lazy_load_warmup(self): model = getattr(self, "model", None) if model is not None: torch_device_module.synchronize() - if hasattr(getattr(model, "transformer_infer", None), "offload_manager"): - del model.transformer_infer.offload_manager + model.transformer_infer.clear_offload_managers() self.scheduler.transformer_infer = None self.model = None @@ -413,6 +415,7 @@ def run_vae_decoder(self, latents): self.maybe_empty_cache() return images + @keep_transformer_weights_loaded def run(self, total_steps=None): if total_steps is None: total_steps = self.model.scheduler.infer_steps From 797b1b273e72de4958fac08b1c0d4b5d54b6b9b7 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:42:26 +0800 Subject: [PATCH 09/31] feat(offload): migrate SeedVR block scheduling --- .../seedvr/infer/offload/transformer_infer.py | 58 +++++-------------- lightx2v/models/networks/seedvr/model.py | 20 +------ .../seedvr/weights/transformer_weights.py | 11 ++-- .../models/runners/seedvr/seedvr_runner.py | 5 +- 4 files changed, 29 insertions(+), 65 deletions(-) diff --git a/lightx2v/models/networks/seedvr/infer/offload/transformer_infer.py b/lightx2v/models/networks/seedvr/infer/offload/transformer_infer.py index 1ca25d42e..1e5be65d9 100644 --- a/lightx2v/models/networks/seedvr/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/seedvr/infer/offload/transformer_infer.py @@ -1,17 +1,9 @@ import torch -from lightx2v.common.offload.manager import WeightAsyncStreamManager from lightx2v.models.networks.seedvr.infer.transformer_infer import SeedVRTransformerInfer -from lightx2v_platform.base.global_var import AI_DEVICE - -torch_device_module = getattr(torch, AI_DEVICE) class SeedVROffloadTransformerInfer(SeedVRTransformerInfer): - def __init__(self, config): - super().__init__(config) - self.offload_manager = WeightAsyncStreamManager(offload_granularity="block") - @torch.no_grad() def infer(self, block_weights, pre_infer_out): vid = pre_infer_out.vid @@ -21,40 +13,20 @@ def infer(self, block_weights, pre_infer_out): emb = pre_infer_out.emb cache = pre_infer_out.cache - # Pre-infer and segment input preparation run on the caller's current - # stream. Make the offload compute stream wait before consuming them. - current_stream = torch_device_module.current_stream() - self.offload_manager.compute_stream.wait_stream(current_stream) - - for block_idx in range(len(block_weights)): - if self.offload_manager.need_init_first_buffer: - self.offload_manager.init_first_buffer(block_weights) - - next_block_idx = (block_idx + 1) % len(block_weights) - self.offload_manager.prefetch_weights(next_block_idx, block_weights) - - if AI_DEVICE == "xpu": - vid, txt, vid_shape, txt_shape = self._infer_block( - self.offload_manager.cuda_buffers[0], - vid, - txt, - vid_shape, - txt_shape, - emb, - cache, - ) - else: - with torch_device_module.stream(self.offload_manager.compute_stream): - vid, txt, vid_shape, txt_shape = self._infer_block( - self.offload_manager.cuda_buffers[0], - vid, - txt, - vid_shape, - txt_shape, - emb, - cache, - ) - - self.offload_manager.swap_blocks() + def infer_block(block_idx, block_weight): + nonlocal vid, txt, vid_shape, txt_shape + self.block_idx = block_idx + vid, txt, vid_shape, txt_shape = self._infer_block( + block_weight, + vid, + txt, + vid_shape, + txt_shape, + emb, + cache, + ) + return vid, txt, vid_shape, txt_shape + + self.run_blocks_with_offload(block_weights, infer_block) return vid, txt, vid_shape, txt_shape diff --git a/lightx2v/models/networks/seedvr/model.py b/lightx2v/models/networks/seedvr/model.py index 85c62d69d..de3a9d9eb 100644 --- a/lightx2v/models/networks/seedvr/model.py +++ b/lightx2v/models/networks/seedvr/model.py @@ -106,8 +106,7 @@ def _init_infer(self): self.pre_infer = self.pre_infer_class(self.config) self.transformer_infer = self.transformer_infer_class(self.config) self.post_infer = self.post_infer_class(self.config) - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() + self._init_offload_manager() def _load_ckpt(self, unified_dtype, sensitive_layer): # SeedVR weights are typically in .pth/.pt format, not safetensors. @@ -196,17 +195,9 @@ def infer(self, inputs): texts_pos = inputs["text_encoder_output"]["texts_pos"] texts_neg = inputs["text_encoder_output"]["texts_neg"] - if self.cpu_offload and self.offload_granularity == "block": - # Each request/segment must begin with block 0 even if a previous - # inference was interrupted before the final buffer swap. - self.transformer_infer.offload_manager.need_init_first_buffer = True - if self.cpu_offload: if self.offload_granularity == "model": self.to_cuda() - else: - self.pre_weight.to_cuda() - self.post_weight.to_cuda() texts_pos[0] = texts_pos[0].to(AI_DEVICE) texts_neg[0] = texts_neg[0].to(AI_DEVICE) @@ -244,10 +235,5 @@ def infer(self, inputs): latents_list = na_utils.unflatten(latents, latents_shapes) self.scheduler.latents = latents_list - if self.cpu_offload: - if self.offload_granularity == "model": - self.to_cpu() - else: - self.pre_weight.to_cpu() - self.post_weight.to_cpu() - return + if self.cpu_offload and self.offload_granularity == "model": + self.to_cpu() diff --git a/lightx2v/models/networks/seedvr/weights/transformer_weights.py b/lightx2v/models/networks/seedvr/weights/transformer_weights.py index f0304b428..5c3677c98 100755 --- a/lightx2v/models/networks/seedvr/weights/transformer_weights.py +++ b/lightx2v/models/networks/seedvr/weights/transformer_weights.py @@ -28,8 +28,10 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): self.mm_layers = config.get("mm_layers", self.blocks_num) self.last_layer_vid_only = bool(config.get("last_layer_vid_only")) - self.register_offload_buffers() blocks = WeightModuleList(self._make_block(i) for i in range(self.blocks_num)) + slot_count = self.register_offload_block_group(config, "blocks", blocks) + self.register_offload_buffers(slot_count) + # Buffer allocation consumes checkpoint metadata, so it must load before the source blocks. self.add_module("blocks", blocks) def _block_uses_shared_weights(self, block_index): @@ -54,10 +56,10 @@ def _make_block(self, block_index, *, create_cuda_buffer=False, branches=None, a alias_shared_to_vid=alias_shared_to_vid, ) - def register_offload_buffers(self): + def register_offload_buffers(self, slot_count): self.offload_block_cuda_buffers = None self.offload_phase_cuda_buffers = None - if not self.config.get("cpu_offload", False) or self.config.get("offload_granularity", "block") != "block": + if slot_count == 0: return split_block_indices = [i for i in range(self.blocks_num) if not self._block_uses_shared_weights(i)] @@ -77,9 +79,10 @@ def register_offload_buffers(self): branches=branches, alias_shared_to_vid=alias_shared_to_vid, ) - for _ in range(2) + for _ in range(slot_count) ) self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) + self.register_offload_block_buffers("blocks", self.offload_block_cuda_buffers) class SeedVRTransformerBlockWeights(WeightModule): diff --git a/lightx2v/models/runners/seedvr/seedvr_runner.py b/lightx2v/models/runners/seedvr/seedvr_runner.py index 687137877..249788d5e 100755 --- a/lightx2v/models/runners/seedvr/seedvr_runner.py +++ b/lightx2v/models/runners/seedvr/seedvr_runner.py @@ -22,6 +22,8 @@ from loguru import logger from torch import Tensor +from lightx2v.common.offload.config import get_offload_granularity +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.default_runner import DefaultRunner from lightx2v.models.schedulers.seedvr.scheduler import SeedVRScheduler from lightx2v.models.video_encoders.hf.seedvr import attn_video_vae_v3_s8_c16_t4_inflation_sd3_init @@ -355,6 +357,7 @@ def _run_sr_single_segment(self): self.input_info = cached_input_info return raw_video + @keep_transformer_weights_loaded def run_segment(self, segment_idx=0): """Run SeedVR diffusion steps under the single outer DiT profile.""" infer_steps = self.model.scheduler.infer_steps @@ -462,7 +465,7 @@ def load_transformer(self): logger.info( f"[SeedVRRunner] DiT config: model_size={self.config.get('model_size', '3b')}, " f"cpu_offload={self.config.get('cpu_offload', False)}, " - f"offload_granularity={self.config.get('offload_granularity', 'block')}, " + f"offload_granularity={get_offload_granularity(self.config)}, " f"quant_scheme={self.config.get('dit_quant_scheme', 'Default')}" ) model = SeedVRNaDiTModel( From 474c0f0b083a32979bbc2726107ace32e81ebdcd Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:42:33 +0800 Subject: [PATCH 10/31] feat(offload): migrate Z-Image block scheduling --- .../infer/offload/transformer_infer.py | 57 +++++++++++++++---- lightx2v/models/networks/z_image/model.py | 7 +-- .../z_image/weights/transformer_weights.py | 32 +++++++---- .../models/runners/z_image/z_image_runner.py | 2 + 4 files changed, 70 insertions(+), 28 deletions(-) diff --git a/lightx2v/models/networks/z_image/infer/offload/transformer_infer.py b/lightx2v/models/networks/z_image/infer/offload/transformer_infer.py index 3b3737ea0..487c694e6 100755 --- a/lightx2v/models/networks/z_image/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/z_image/infer/offload/transformer_infer.py @@ -1,6 +1,6 @@ import torch -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.z_image.infer.transformer_infer import ZImageTransformerInfer from lightx2v_platform.base.global_var import AI_DEVICE @@ -11,13 +11,10 @@ class ZImageOffloadTransformerInfer(ZImageTransformerInfer): def __init__(self, config): super().__init__(config) if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": - self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) self.lazy_load = self.config.get("lazy_load", False) self.infer_main_blocks = self.infer_main_blocks_offload - if self.lazy_load: - self.offload_manager.init_lazy_load(num_workers=self.config.get("num_disk_workers", 4)) elif offload_granularity == "phase": raise NotImplementedError("offload_granularity=phase not supported") @@ -30,24 +27,60 @@ def infer_with_blocks_offload( adaln_input, image_tokens_len, ): + if self.lazy_load: + return self.infer_with_lazy_blocks_offload( + main_blocks, + unified, + unified_freqs_cis, + unified_rope_positions, + adaln_input, + image_tokens_len, + ) + + def run_z_image_block(block_idx, block): + nonlocal unified + self.block_idx = block_idx + unified = self.infer_block( + block_weight=block, + hidden_states=unified, + freqs_cis=unified_freqs_cis, + rope_positions=unified_rope_positions, + adaln_input=adaln_input, + image_tokens_len=image_tokens_len, + ) + return unified + + self.run_blocks_with_offload(main_blocks, run_z_image_block) + return unified + + def infer_with_lazy_blocks_offload( + self, + main_blocks, + unified, + unified_freqs_cis, + unified_rope_positions, + adaln_input, + image_tokens_len, + ): + manager = self.get_block_offload_manager(main_blocks) num_blocks = len(main_blocks) for block_idx in range(num_blocks): self.block_idx = block_idx if self.lazy_load: next_prefetch = (block_idx + 1) % num_blocks - self.offload_manager.start_prefetch_block(next_prefetch) + manager.start_prefetch_block(next_prefetch) if block_idx == 0: - self.offload_manager.init_first_buffer(main_blocks) + manager.init_first_buffer(main_blocks) if self.lazy_load: - self.offload_manager.swap_cpu_buffers() - self.offload_manager.prefetch_weights((block_idx + 1) % num_blocks, main_blocks) + manager.swap_cpu_buffers() + manager.prefetch_weights((block_idx + 1) % num_blocks, main_blocks) - with torch_device_module.stream(self.offload_manager.compute_stream): + with torch_device_module.stream(manager.compute_stream): unified = self.infer_block( - block_weight=self.offload_manager.cuda_buffers[0], + block_weight=manager.cuda_buffers[0], hidden_states=unified, freqs_cis=unified_freqs_cis, rope_positions=unified_rope_positions, @@ -55,7 +88,7 @@ def infer_with_blocks_offload( image_tokens_len=image_tokens_len, ) - self.offload_manager.swap_blocks() + manager.swap_blocks() return unified diff --git a/lightx2v/models/networks/z_image/model.py b/lightx2v/models/networks/z_image/model.py index 03ed3752a..800ff186a 100755 --- a/lightx2v/models/networks/z_image/model.py +++ b/lightx2v/models/networks/z_image/model.py @@ -41,8 +41,7 @@ def _init_infer(self): self.pre_infer = self.pre_infer_class(self.config) self.post_infer = self.post_infer_class(self.config) self.pre_infer.set_rope(self.transformer_weights.blocks[0].compute_phases[1].rope) - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() + self._init_offload_manager() @torch.no_grad() def _infer_cond_uncond(self, latents_input, prompt_embeds, infer_condition=True): @@ -97,7 +96,7 @@ def _seq_parallel_post_process(self, hidden_states, image_shard_len): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == 0: self.to_cuda() elif self.offload_granularity != "model": @@ -150,7 +149,7 @@ def infer(self, inputs): self.scheduler.noise_pred = noise_pred - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1 and "wan2.2_moe" not in self.config["model_cls"]: self.to_cpu() elif self.offload_granularity != "model": diff --git a/lightx2v/models/networks/z_image/weights/transformer_weights.py b/lightx2v/models/networks/z_image/weights/transformer_weights.py index 5e64da9da..ef4f69d45 100755 --- a/lightx2v/models/networks/z_image/weights/transformer_weights.py +++ b/lightx2v/models/networks/z_image/weights/transformer_weights.py @@ -1,6 +1,7 @@ import torch from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.utils.registry_factory import ( ATTN_WEIGHT_REGISTER, MM_WEIGHT_REGISTER, @@ -20,13 +21,12 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): assert config.get("dit_quantized") is True self.lazy_load = self.config.get("lazy_load", False) self.n_refiner_layers = config.get("n_refiner_layers", 0) - self.register_offload_buffers(config, lazy_load_path, lora_path) - self.add_module( - "blocks", - WeightModuleList( - ZImageTransformerBlock(i, self.task, self.mm_type, self.config, False, False, "layers", lazy_load=self.lazy_load, lazy_load_path=lazy_load_path) for i in range(self.blocks_num) - ), + self.blocks = WeightModuleList( + ZImageTransformerBlock(i, self.task, self.mm_type, self.config, False, False, "layers", lazy_load=self.lazy_load, lazy_load_path=lazy_load_path) for i in range(self.blocks_num) ) + slot_count = self.register_offload_block_group(config, "blocks", self.blocks) + self.register_offload_buffers(config, lazy_load_path, lora_path, slot_count) + self.add_module("blocks", self.blocks) self.add_module( "noise_refiner", @@ -62,10 +62,16 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): ), ) - def register_offload_buffers(self, config, lazy_load_path, lora_path): + def register_offload_buffers(self, config, lazy_load_path, lora_path, slot_count): if config["cpu_offload"]: - if config["offload_granularity"] == "block": - self.offload_blocks_num = 2 + if get_offload_granularity(config) == "block": + self.offload_blocks_num = slot_count + self.offload_block_cuda_buffers = None + self.offload_block_cpu_buffers = None + self.offload_phase_cuda_buffers = None + self.offload_phase_cpu_buffers = None + if slot_count == 0: + return self.offload_block_cuda_buffers = WeightModuleList( [ ZImageTransformerBlock( @@ -83,9 +89,7 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) - self.offload_phase_cuda_buffers = None if self.lazy_load: - self.offload_blocks_num = 2 self.offload_block_cpu_buffers = WeightModuleList( [ ZImageTransformerBlock( @@ -104,7 +108,11 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cpu_buffers", self.offload_block_cpu_buffers) - self.offload_phase_cpu_buffers = None + self.register_offload_block_buffers( + "blocks", + self.offload_block_cuda_buffers, + self.offload_block_cpu_buffers, + ) def non_block_weights_to_cuda(self): self.noise_refiner.to_cuda() diff --git a/lightx2v/models/runners/z_image/z_image_runner.py b/lightx2v/models/runners/z_image/z_image_runner.py index 4c7eaf8ca..e3b6389d1 100755 --- a/lightx2v/models/runners/z_image/z_image_runner.py +++ b/lightx2v/models/runners/z_image/z_image_runner.py @@ -9,6 +9,7 @@ from lightx2v.models.input_encoders.hf.z_image.qwen3_model import Qwen3Model_TextEncoder from lightx2v.models.networks.lora_adapter import LoraAdapter from lightx2v.models.networks.z_image.model import ZImageTransformerModel +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.default_runner import DefaultRunner from lightx2v.models.schedulers.z_image.scheduler import ZImageScheduler from lightx2v.models.video_encoders.hf.z_image.vae import AutoencoderKLZImageVAE @@ -258,6 +259,7 @@ def run_vae_encoder(self, image): gc.collect() return {"image_latents": image_latents} + @keep_transformer_weights_loaded def run(self, total_steps=None): if total_steps is None: total_steps = self.model.scheduler.infer_steps From dfbb78cf7f0cf20b0358e7a34d3a7648ac0e7558 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:43:02 +0800 Subject: [PATCH 11/31] refactor(offload): read Bagel granularity from offload plan --- lightx2v/models/networks/bagel/model.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/lightx2v/models/networks/bagel/model.py b/lightx2v/models/networks/bagel/model.py index e58b6c561..a286bb59b 100644 --- a/lightx2v/models/networks/bagel/model.py +++ b/lightx2v/models/networks/bagel/model.py @@ -8,6 +8,7 @@ from loguru import logger from torch.nn import functional as F +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.bagel.data_utils import add_special_tokens, patchify from lightx2v.models.networks.bagel.infer.post_infer import BagelPostInfer from lightx2v.models.networks.bagel.infer.pre_infer import BagelPreInfer @@ -58,7 +59,7 @@ def __init__(self, config): self.enable_vision_context = config.get("enable_vision_context", config.get("task", "t2i") == "i2i") self.cpu_offload = config.get("cpu_offload", False) - self.offload_granularity = self.config.get("offload_granularity", "block") + self.offload_granularity = get_offload_granularity(self.config) self.device = torch.device("cpu") if self.cpu_offload else torch.device(AI_DEVICE) self._init_infer_class() self._init_weights() From 1c0e089a42a97b1017be48f020fb7d71670e84a6 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:43:14 +0800 Subject: [PATCH 12/31] refactor(offload): read HunyuanImage3 granularity from offload plan --- .../networks/hunyuan_image3/weights/transformer_weights.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/lightx2v/models/networks/hunyuan_image3/weights/transformer_weights.py b/lightx2v/models/networks/hunyuan_image3/weights/transformer_weights.py index af5f843e1..e76cd45ba 100644 --- a/lightx2v/models/networks/hunyuan_image3/weights/transformer_weights.py +++ b/lightx2v/models/networks/hunyuan_image3/weights/transformer_weights.py @@ -1,4 +1,5 @@ from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.hunyuan_image3.weights.common import ( HunyuanImage3AttentionWeights, HunyuanImage3MLPPhaseWeights, @@ -41,7 +42,8 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): if not config.get("cpu_offload", False): return - if config.get("offload_granularity", "block") == "block": + offload_granularity = get_offload_granularity(config) + if offload_granularity == "block": self.offload_blocks_num = 2 self.offload_block_cuda_buffers = WeightModuleList( [ @@ -78,7 +80,7 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cpu_buffers", self.offload_block_cpu_buffers) - elif config.get("offload_granularity") == "phase": + elif offload_granularity == "phase": self.offload_phase_cuda_buffers = HunyuanImage3TransformerBlock( 0, config, From 01b844c71bb7291e7f8e1f84cc71a09f49f56233 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:44:34 +0800 Subject: [PATCH 13/31] feat(offload): migrate Wan block scheduling --- configs/offload/block/wan_t2v_block.json | 8 ++- .../wan/infer/offload/transformer_infer.py | 56 +++++++++++----- lightx2v/models/networks/wan/model.py | 65 ++++++++++++------- .../wan/weights/transformer_weights.py | 26 +++++--- lightx2v/models/runners/wan/wan_runner.py | 54 +++++++++++---- 5 files changed, 144 insertions(+), 65 deletions(-) diff --git a/configs/offload/block/wan_t2v_block.json b/configs/offload/block/wan_t2v_block.json index 2f926bce9..cb87fc0a5 100755 --- a/configs/offload/block/wan_t2v_block.json +++ b/configs/offload/block/wan_t2v_block.json @@ -17,7 +17,13 @@ "clip_quantized": true, "clip_quant_scheme": "fp8-q8f", "cpu_offload": true, - "offload_granularity": "block", + "offload_plan": { + "offload_granularity": "block", + "resident_blocks": { + "blocks": 20 + }, + "use_event_offload": false + }, "t5_cpu_offload": false, "vae_cpu_offload": false, "clip_cpu_offload": false diff --git a/lightx2v/models/networks/wan/infer/offload/transformer_infer.py b/lightx2v/models/networks/wan/infer/offload/transformer_infer.py index bff53522b..d97731af7 100755 --- a/lightx2v/models/networks/wan/infer/offload/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/offload/transformer_infer.py @@ -1,5 +1,6 @@ import torch +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.common.offload.manager import WeightAsyncStreamManager from lightx2v.models.networks.wan.infer.transformer_infer import WanTransformerInfer from lightx2v_platform.base.global_var import AI_DEVICE @@ -11,7 +12,7 @@ class WanOffloadTransformerInfer(WanTransformerInfer): def __init__(self, config): super().__init__(config) if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": self.infer_func = self.infer_with_blocks_offload elif offload_granularity == "phase": @@ -31,45 +32,64 @@ def __init__(self, config): elif offload_granularity == "model": self.infer_func = self.infer_without_offload - if offload_granularity != "model": + if offload_granularity not in ("block", "model"): self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) self.lazy_load = self.config.get("lazy_load", False) - if self.lazy_load: + if self.lazy_load and offload_granularity == "phase": self.offload_manager.init_lazy_load(num_workers=self.config.get("num_disk_workers", 4)) + def infer(self, weights, pre_infer_out): + try: + return super().infer(weights, pre_infer_out) + finally: + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "block": + self.clear_block_offload_inputs(pre_infer_out) + def infer_with_blocks_offload(self, blocks, x, pre_infer_out): + if self.lazy_load: + return self.infer_with_lazy_blocks_offload(blocks, x, pre_infer_out) + + def run_wan_block(block_idx, block): + nonlocal x + self.block_idx = block_idx + x = self.run_block(block_idx, block, x, pre_infer_out) + return x + + self.run_blocks_with_offload(blocks, run_wan_block) + return x + + def infer_with_lazy_blocks_offload(self, blocks, x, pre_infer_out): + manager = self.get_block_offload_manager(blocks) for block_idx in range(len(blocks)): self.block_idx = block_idx if self.lazy_load: next_prefetch = (block_idx + 1) % len(blocks) - self.offload_manager.start_prefetch_block(next_prefetch) + manager.start_prefetch_block(next_prefetch) - if self.offload_manager.need_init_first_buffer: - self.offload_manager.init_first_buffer(blocks) + if manager.need_init_first_buffer: + manager.init_first_buffer(blocks) if self.lazy_load: - self.offload_manager.swap_cpu_buffers() + manager.swap_cpu_buffers() - self.offload_manager.prefetch_weights((block_idx + 1) % len(blocks), blocks) + manager.prefetch_weights((block_idx + 1) % len(blocks), blocks) if AI_DEVICE == "xpu": # XPU streams do not guarantee cross-stream memory visibility even # after a device-wide sync, so run compute on the default stream. - x = self.run_block(block_idx, self.offload_manager.cuda_buffers[0], x, pre_infer_out) + x = self.run_block(block_idx, manager.cuda_buffers[0], x, pre_infer_out) else: - with torch_device_module.stream(self.offload_manager.compute_stream): - x = self.run_block(block_idx, self.offload_manager.cuda_buffers[0], x, pre_infer_out) + with torch_device_module.stream(manager.compute_stream): + x = self.run_block(block_idx, manager.cuda_buffers[0], x, pre_infer_out) - self.offload_manager.swap_blocks() + manager.swap_blocks() + + return x + def clear_block_offload_inputs(self, pre_infer_out): if self.clean_cuda_cache: - del ( - pre_infer_out.embed0, - pre_infer_out.context, - ) + del pre_infer_out.embed0, pre_infer_out.context torch_device_module.empty_cache() - return x - def infer_with_phases_offload(self, blocks, x, pre_infer_out): for block_idx in range(len(blocks)): self.block_idx = block_idx diff --git a/lightx2v/models/networks/wan/model.py b/lightx2v/models/networks/wan/model.py index 9da2efd4d..74ee30f43 100755 --- a/lightx2v/models/networks/wan/model.py +++ b/lightx2v/models/networks/wan/model.py @@ -193,26 +193,38 @@ def _init_infer_class(self): self.pre_infer_class = WanPreInfer self.post_infer_class = WanPostInfer - if self.config["feature_caching"] == "NoCaching": + feature_caching = self.config["feature_caching"] + unsupported_block_caching = { + "TaylorSeer", + "Ada", + "Custom", + "FirstBlock", + "DualBlock", + "DynamicBlock", + } + if self.cpu_offload and self.offload_granularity == "block" and feature_caching in unsupported_block_caching: + raise NotImplementedError(f"Wan block offload does not support feature_caching={feature_caching!r}") + + if feature_caching == "NoCaching": self.transformer_infer_class = WanTransformerInfer if not self.cpu_offload else WanOffloadTransformerInfer - elif self.config["feature_caching"] == "Tea": + elif feature_caching == "Tea": self.transformer_infer_class = WanTransformerInferTeaCaching - elif self.config["feature_caching"] == "TaylorSeer": + elif feature_caching == "TaylorSeer": self.transformer_infer_class = WanTransformerInferTaylorCaching - elif self.config["feature_caching"] == "Ada": + elif feature_caching == "Ada": self.transformer_infer_class = WanTransformerInferAdaCaching - elif self.config["feature_caching"] == "Custom": + elif feature_caching == "Custom": self.transformer_infer_class = WanTransformerInferCustomCaching - elif self.config["feature_caching"] == "FirstBlock": + elif feature_caching == "FirstBlock": self.transformer_infer_class = WanTransformerInferFirstBlock - elif self.config["feature_caching"] == "DualBlock": + elif feature_caching == "DualBlock": self.transformer_infer_class = WanTransformerInferDualBlock - elif self.config["feature_caching"] == "DynamicBlock": + elif feature_caching == "DynamicBlock": self.transformer_infer_class = WanTransformerInferDynamicBlock - elif self.config["feature_caching"] == "Mag": + elif feature_caching == "Mag": self.transformer_infer_class = WanTransformerInferMagCaching else: - raise NotImplementedError(f"Unsupported feature_caching type: {self.config['feature_caching']}") + raise NotImplementedError(f"Unsupported feature_caching type: {feature_caching}") def _init_infer(self): self.pre_infer = self.pre_infer_class(self.config) @@ -224,8 +236,7 @@ def _init_infer(self): self.pre_infer.set_rope(rope if rope is not None else first_attn.rope) if hasattr(self.pre_infer, "set_audio_rope"): self.pre_infer.set_audio_rope(self.transformer_weights.blocks[0].compute_phases[2].rope_1d) - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() + self._init_offload_manager() def _should_init_empty_model(self): if self.config.get("lora_configs") and self.config["lora_configs"] and not self.config.get("lora_dynamic_apply", False): @@ -297,12 +308,23 @@ def _seq_parallel_post_process(self, x): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: - if self.offload_granularity == "model" and self.scheduler.step_index == 0 and "wan2.2_moe" not in self.config["model_cls"]: - self.to_cuda() - elif self.offload_granularity != "model": - self.pre_weight.to_cuda() - self.transformer_weights.non_block_weights_to_cuda() + try: + if self.cpu_offload and self.offload_granularity != "block": + if self.offload_granularity == "model" and self.scheduler.step_index == 0 and "wan2.2_moe" not in self.config["model_cls"]: + self.to_cuda() + elif self.offload_granularity != "model": + self.pre_weight.to_cuda() + self.transformer_weights.non_block_weights_to_cuda() + self._infer(inputs) + finally: + if self.cpu_offload and self.offload_granularity != "block": + if self.offload_granularity == "model" and self.scheduler.step_index == self.scheduler.infer_steps - 1 and "wan2.2_moe" not in self.config["model_cls"]: + self.to_cpu() + elif self.offload_granularity != "model": + self.pre_weight.to_cpu() + self.transformer_weights.non_block_weights_to_cpu() + + def _infer(self, inputs): if self.config["enable_cfg"]: if self.config["cfg_parallel"]: @@ -337,10 +359,3 @@ def infer(self, inputs): self.scheduler.noise_pred_uncond = None self.scheduler.noise_pred_guided = noise_pred 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 and "wan2.2_moe" not in self.config["model_cls"]: - self.to_cpu() - elif self.offload_granularity != "model": - self.pre_weight.to_cpu() - self.transformer_weights.non_block_weights_to_cpu() diff --git a/lightx2v/models/networks/wan/weights/transformer_weights.py b/lightx2v/models/networks/wan/weights/transformer_weights.py index 3d16dd6da..0f2d93441 100755 --- a/lightx2v/models/networks/wan/weights/transformer_weights.py +++ b/lightx2v/models/networks/wan/weights/transformer_weights.py @@ -2,6 +2,7 @@ import torch.distributed as dist from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.wan.infer.utils import WanCausalRope # noqa: F401 from lightx2v.utils.registry_factory import ( ATTN_WEIGHT_REGISTER, @@ -144,7 +145,8 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): for i in range(self.blocks_num) ] ) - self.register_offload_buffers(config, lazy_load_path, lora_path) + slot_count = self.register_offload_block_group(config, "blocks", self.blocks) + self._register_main_offload_buffers(config, lazy_load_path, lora_path, slot_count) self.add_module("blocks", self.blocks) # non blocks weights @@ -159,10 +161,16 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): ) self.register_parameter("head_modulation", TENSOR_REGISTER["Default"]("head.modulation")) - def register_offload_buffers(self, config, lazy_load_path, lora_path): + def _register_main_offload_buffers(self, config, lazy_load_path, lora_path, slot_count): if config["cpu_offload"]: - if config["offload_granularity"] == "block": - self.offload_blocks_num = 2 + if get_offload_granularity(config) == "block": + self.offload_blocks_num = slot_count + self.offload_phase_cuda_buffers = None + self.offload_block_cuda_buffers = None + self.offload_block_cpu_buffers = None + self.offload_phase_cpu_buffers = None + if slot_count == 0: + return self.offload_block_cuda_buffers = WeightModuleList( [ WanTransformerAttentionBlock( @@ -180,10 +188,8 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) - self.offload_phase_cuda_buffers = None if self.lazy_load: - self.offload_blocks_num = 2 self.offload_block_cpu_buffers = WeightModuleList( [ WanTransformerAttentionBlock( @@ -201,9 +207,13 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cpu_buffers", self.offload_block_cpu_buffers) - self.offload_phase_cpu_buffers = None + self.register_offload_block_buffers( + "blocks", + self.offload_block_cuda_buffers, + self.offload_block_cpu_buffers, + ) - elif config["offload_granularity"] == "phase": + elif get_offload_granularity(config) == "phase": self.offload_phase_cuda_buffers = WanTransformerAttentionBlock( block_index=0, task=self.task, diff --git a/lightx2v/models/runners/wan/wan_runner.py b/lightx2v/models/runners/wan/wan_runner.py index 1f2527cb2..0b26bf035 100755 --- a/lightx2v/models/runners/wan/wan_runner.py +++ b/lightx2v/models/runners/wan/wan_runner.py @@ -15,6 +15,7 @@ Rotation = None Slerp = None +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.disagg.disagg_mixin import DisaggMixin from lightx2v.models.input_encoders.hf.wan.t5.model import T5EncoderModel from lightx2v.models.input_encoders.hf.wan.xlm_roberta.model import CLIPModel @@ -132,16 +133,17 @@ def _run_warmup(self): if self.config.get("model_cls") == "wan2.2" and self.config["task"] == "i2v": inputs["image_encoder_output"]["vae_encoder_out"] = None try: - previous_step_index = None - for step_index in self.get_warmup_step_indices(scheduler): - if previous_step_index is not None and step_index != previous_step_index + 1: - scheduler.reset(seed=input_info.seed, latent_shape=latent_shape, step_index=step_index) - scheduler.step_pre(step_index=step_index) - self.model.infer(inputs) - scheduler.step_post() - previous_step_index = step_index + with self.transformer_offload_session(): + previous_step_index = None + for step_index in self.get_warmup_step_indices(scheduler): + if previous_step_index is not None and step_index != previous_step_index + 1: + scheduler.reset(seed=input_info.seed, latent_shape=latent_shape, step_index=step_index) + scheduler.step_pre(step_index=step_index) + self.model.infer(inputs) + scheduler.step_post() + previous_step_index = step_index finally: - if self.config.get("cpu_offload", False) and self.config.get("offload_granularity") == "model": + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "model": for model in filter(None, self.get_warmup_models()): model.to_cpu() self.run_vae_decoder(scheduler.latents) @@ -217,8 +219,7 @@ def clean_lazy_load_warmup(self): if models: torch_device_module.synchronize() for model in models: - if hasattr(getattr(model, "transformer_infer", None), "offload_manager"): - del model.transformer_infer.offload_manager + model.transformer_infer.clear_offload_managers() self.scheduler.transformer_infer = None self.model = None for name in ("text_encoders", "image_encoder", "vae_encoder", "vae_decoder"): @@ -918,6 +919,7 @@ def __init__(self, model_list, config, num_train_timesteps=1000): assert len(self.model) == 2, "MultiModelStruct only supports 2 models now." self.config = config self.cur_model_index = -1 + self._offload_weights_active = False self.distill_method = get_wan_distill_method(config) if self.distill_method not in (None, "dmd2"): raise NotImplementedError(f"MultiModelStruct does not support distill_method {self.distill_method!r}") @@ -953,7 +955,13 @@ def set_scheduler(self, shared_scheduler): model.set_scheduler(shared_scheduler) def infer(self, inputs): + previous_model_index = self.cur_model_index self.get_current_model_index() + if self._offload_weights_active and self.cur_model_index != previous_model_index: + if previous_model_index >= 0 and self.model[previous_model_index] is not None: + self.model[previous_model_index].cleanup_offload_weights() + if self.model[self.cur_model_index] is not None: + self.model[self.cur_model_index].prepare_offload_weights() if not self.config.get("lazy_load", False) and not self.config.get("unload_modules", False): self.model[self.cur_model_index].infer(inputs) else: @@ -975,6 +983,8 @@ def infer(self, inputs): high_noise_model = build_wan_model_with_lora(WanModel, self.config, high_model_kwargs, lora_configs, model_type="high_noise_model") high_noise_model.set_scheduler(self.scheduler) self.model[0] = high_noise_model + if self._offload_weights_active: + high_noise_model.prepare_offload_weights() self.model[0].infer(inputs) elif self.cur_model_index == 1: lora_configs = self.config.get("lora_configs") @@ -991,15 +1001,33 @@ def infer(self, inputs): low_noise_model = build_wan_model_with_lora(WanModel, self.config, low_model_kwargs, lora_configs, model_type="low_noise_model") low_noise_model.set_scheduler(self.scheduler) self.model[1] = low_noise_model + if self._offload_weights_active: + low_noise_model.prepare_offload_weights() self.model[1].infer(inputs) + def prepare_offload_weights(self): + if not self.config.get("cpu_offload", False) or get_offload_granularity(self.config) != "block" or self._offload_weights_active: + return False + self.cur_model_index = -1 + self._offload_weights_active = True + return True + + def cleanup_offload_weights(self): + if not self._offload_weights_active: + return + for model in self.model: + if model is not None: + model.cleanup_offload_weights() + self.cur_model_index = -1 + self._offload_weights_active = False + @ProfilingContext4DebugL2("Switch models in infer_main costs") def get_current_model_index(self): if self.uses_high_noise_model(): logger.info(f"using - HIGH - noise model at step_index {self.scheduler.step_index + 1}") if self.config["enable_cfg"]: self.scheduler.sample_guide_scale = self.config["sample_guide_scale"][0] - if self.config.get("cpu_offload", False) and self.config.get("offload_granularity", "block") == "model": + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "model": if self.cur_model_index == -1: self.to_cuda(model_index=0) elif self.cur_model_index == 1: # 1 -> 0 @@ -1010,7 +1038,7 @@ def get_current_model_index(self): logger.info(f"using - LOW - noise model at step_index {self.scheduler.step_index + 1}") if self.config["enable_cfg"]: self.scheduler.sample_guide_scale = self.config["sample_guide_scale"][1] - if self.config.get("cpu_offload", False) and self.config.get("offload_granularity", "block") == "model": + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "model": if self.cur_model_index == -1: self.to_cuda(model_index=1) elif self.cur_model_index == 0: # 0 -> 1 From cc0dd7160cbf1074c124bde0fee12f9b0804a97a Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:46:33 +0800 Subject: [PATCH 14/31] feat(offload): migrate Wan Animate block scheduling --- .../wan/infer/animate/transformer_infer.py | 21 ++++++++++++------- 1 file changed, 13 insertions(+), 8 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/animate/transformer_infer.py b/lightx2v/models/networks/wan/infer/animate/transformer_infer.py index 0c7bae276..ef94b2c28 100755 --- a/lightx2v/models/networks/wan/infer/animate/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/animate/transformer_infer.py @@ -15,17 +15,22 @@ def __init__(self, config): self.adapter_cu_seqlens_kv = None self._adapter_cu_seqlens_key = None + @staticmethod + def get_adapter_block_index(block_idx): + return block_idx // 5 + def infer_with_blocks_offload(self, blocks, x, pre_infer_out): - for block_idx in range(len(blocks)): + def run_animate_block(block_idx, block): + nonlocal x self.block_idx = block_idx - if block_idx == 0: - self.offload_manager.init_first_buffer(blocks, block_idx // 5) - if block_idx < len(blocks) - 1: - self.offload_manager.prefetch_weights(block_idx + 1, blocks, (block_idx + 1) // 5) + x = self.run_block(block_idx, block, x, pre_infer_out) + return x - with torch.cuda.stream(self.offload_manager.compute_stream): - x = self.infer_block(self.offload_manager.cuda_buffers[0], x, pre_infer_out) - self.offload_manager.swap_blocks() + self.run_blocks_with_offload( + blocks, + run_animate_block, + adapter_block_index=self.get_adapter_block_index, + ) return x def infer_phases(self, block_idx, blocks, x, pre_infer_out, lazy=None): From e3e8b2a520262c9779bcbe1bcfc6371a3caf3efe Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:46:44 +0800 Subject: [PATCH 15/31] feat(offload): migrate Wan Animate2 block scheduling --- lightx2v/models/networks/wan/animate2_model.py | 6 ++---- .../models/networks/wan/infer/animate2/transformer_infer.py | 3 ++- lightx2v/models/runners/wan/wan_animate2_runner.py | 6 ++++-- 3 files changed, 8 insertions(+), 7 deletions(-) diff --git a/lightx2v/models/networks/wan/animate2_model.py b/lightx2v/models/networks/wan/animate2_model.py index 48e91fab5..1e4620948 100644 --- a/lightx2v/models/networks/wan/animate2_model.py +++ b/lightx2v/models/networks/wan/animate2_model.py @@ -59,7 +59,7 @@ def _init_infer_class(self): @torch.no_grad() def prepare_reference(self, inputs): """Prefill every layer's immutable driving-reference K/V once per clip.""" - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model": self.to_cuda() else: @@ -71,9 +71,7 @@ def prepare_reference(self, inputs): pre_infer_out = self._seq_parallel_pre_process(pre_infer_out) self.transformer_infer.infer_reference(self.transformer_weights, pre_infer_out) finally: - # Resident runners must restore their requested offload state even - # when a prefill is cancelled or raises partway through a block. - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model": self.to_cpu() else: diff --git a/lightx2v/models/networks/wan/infer/animate2/transformer_infer.py b/lightx2v/models/networks/wan/infer/animate2/transformer_infer.py index d77d2a01a..dec2b8f26 100644 --- a/lightx2v/models/networks/wan/infer/animate2/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/animate2/transformer_infer.py @@ -1,5 +1,6 @@ import torch +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.common.ops.attn.flex_attn import FlexAttnWeight from lightx2v.common.ops.attn.utils.all2all import all2all_head2seq, all2all_seq2head from lightx2v.models.networks.wan.infer.offload.transformer_infer import WanOffloadTransformerInfer @@ -12,7 +13,7 @@ def __init__(self, config): super().__init__(config) if config.get("feature_caching", "NoCaching") != "NoCaching": raise NotImplementedError("Wan-Animate-2 does not support feature caching.") - if config.get("cpu_offload", False) and config.get("offload_granularity", "block") == "phase": + if config.get("cpu_offload", False) and get_offload_granularity(config) == "phase": raise NotImplementedError("Wan-Animate-2 supports model/block offload, not phase offload.") if config.get("use_compile", False): raise NotImplementedError("Wan-Animate-2 block compilation is disabled; its FlexAttention kernel is compiled internally.") diff --git a/lightx2v/models/runners/wan/wan_animate2_runner.py b/lightx2v/models/runners/wan/wan_animate2_runner.py index 0c71075dd..256c13cb2 100644 --- a/lightx2v/models/runners/wan/wan_animate2_runner.py +++ b/lightx2v/models/runners/wan/wan_animate2_runner.py @@ -16,6 +16,7 @@ VideoReader = None from lightx2v.common.kvcache import KVCacheManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.wan.animate2_identity import WAN_ANIMATE2_MODEL_ID from lightx2v.models.networks.wan.animate2_model import WanAnimate2Model from lightx2v.models.runners.wan.wan_runner import WanRunner, build_wan_model_with_lora @@ -184,7 +185,7 @@ def __init__(self, config): raise NotImplementedError("Wan-Animate-2 does not support disaggregated inference yet.") if self.config.get("lazy_load", False): raise NotImplementedError("Wan-Animate-2 does not support lazy loading yet.") - if self.config.get("cpu_offload", False) and self.config.get("offload_granularity", "block") == "phase": + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "phase": raise NotImplementedError("Wan-Animate-2 supports model/block offload, not phase offload.") if self.config.get("enable_reuse", False): raise NotImplementedError("Wan-Animate-2 request reuse is not implemented for autoregressive inputs.") @@ -650,7 +651,8 @@ def _build_segment_inputs(self, segment_idx): def init_run_segment(self, segment_idx): try: self._build_segment_inputs(segment_idx) - self.model.prepare_reference(self.inputs) + with self.transformer_offload_session(): + self.model.prepare_reference(self.inputs) if segment_idx == 0: self.model.scheduler.prepare(self.input_info.seed, self.input_info.latent_shape) From 0ebc62483b8a976d3ef1af24f0c5f4bed159b675 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:46:46 +0800 Subject: [PATCH 16/31] feat(offload): migrate Wan self-forcing block scheduling --- .../infer/self_forcing/transformer_infer.py | 36 +++++++------------ lightx2v/models/networks/wan/sf_model.py | 26 +++++++------- lightx2v/models/runners/wan/wan_sf_runner.py | 25 ++++++------- 3 files changed, 38 insertions(+), 49 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/self_forcing/transformer_infer.py b/lightx2v/models/networks/wan/infer/self_forcing/transformer_infer.py index 43e6024fe..425ea9c21 100755 --- a/lightx2v/models/networks/wan/infer/self_forcing/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/self_forcing/transformer_infer.py @@ -3,7 +3,7 @@ import torch.nn.functional as F from loguru import logger -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.common.ops.attn.utils.all2all import all2all_seq2head from lightx2v.models.networks.wan.infer.transformer_infer import WanTransformerInfer from lightx2v_platform.base.global_var import AI_DEVICE @@ -26,14 +26,8 @@ def __init__(self, config): # ``infer_block_func`` (KV cache CPU offload vs on-GPU). self._weight_offload_block_compute = False cpu_off = self.config.get("cpu_offload", False) - gran = self.config.get("offload_granularity", "block") + gran = get_offload_granularity(self.config) if cpu_off and gran == "block": - self.offload_manager = WeightAsyncStreamManager(offload_granularity="block") - self.lazy_load = self.config.get("lazy_load", False) - if self.lazy_load: - self.offload_manager.init_lazy_load( - num_workers=self.config.get("num_disk_workers", 4), - ) self.infer_func = self.infer_with_kvcache_blocks_offload self._weight_offload_block_compute = True elif cpu_off: @@ -106,39 +100,33 @@ def infer_with_kvcache(self, blocks, x, pre_infer_out): def infer_with_kvcache_blocks_offload(self, blocks, x, pre_infer_out): """Run transformer blocks with both weight offload and KV cache support.""" + block_manager = self.get_block_offload_manager(blocks) mgr = self.kv_cache_manager self.kv_cache_size = mgr.kv_size self.max_attention_size = mgr.max_attention_size self._kv_offload = self._ar_kv_offload kv_cache = mgr.self_attn_kv_cache - num_blocks = len(blocks) - for block_idx in range(num_blocks): + def run_kvcache_block(block_idx, block): + nonlocal x self.block_idx = block_idx self._set_layer_cache_limits(block_idx) if self._kv_offload: self._next_prefetch = None + x = self.infer_block_func(block, x, pre_infer_out) + return x - if self.offload_manager.need_init_first_buffer: - self.offload_manager.init_first_buffer(blocks) - - self.offload_manager.prefetch_weights((block_idx + 1) % num_blocks, blocks) - gpu_block = self.offload_manager.cuda_buffers[0] - if AI_DEVICE == "xpu": - x = self.infer_block_func(gpu_block, x, pre_infer_out) - else: - with torch_device_module.stream(self.offload_manager.compute_stream): - x = self.infer_block_func(gpu_block, x, pre_infer_out) - - self.offload_manager.swap_blocks() + self.run_blocks_with_offload(blocks, run_kvcache_block) if self.clean_cuda_cache: del pre_infer_out.embed0, pre_infer_out.context torch_device_module.empty_cache() if self._kv_offload: - if self._weight_offload_block_compute and AI_DEVICE == "cuda": - self.offload_manager.compute_stream.synchronize() + if self._weight_offload_block_compute and block_manager is not None and block_manager.uses_events: + torch_device_module.current_stream().synchronize() + elif self._weight_offload_block_compute and block_manager is not None and AI_DEVICE == "cuda": + block_manager.compute_stream.synchronize() else: comp = getattr(kv_cache, "compute_stream", None) if comp is not None: diff --git a/lightx2v/models/networks/wan/sf_model.py b/lightx2v/models/networks/wan/sf_model.py index 48eaaa99f..a3a28ee83 100755 --- a/lightx2v/models/networks/wan/sf_model.py +++ b/lightx2v/models/networks/wan/sf_model.py @@ -41,21 +41,21 @@ def _init_infer_class(self): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == 0: self.to_cuda() elif self.offload_granularity != "model": self.pre_weight.to_cuda() self.transformer_weights.non_block_weights_to_cuda() - - current_start_frame = self.scheduler.seg_index * self.scheduler.num_frame_per_chunk - current_end_frame = (self.scheduler.seg_index + 1) * self.scheduler.num_frame_per_chunk - noise_pred = self._infer_cond_uncond(inputs, infer_condition=True) - - self.scheduler.noise_pred[:, current_start_frame:current_end_frame] = noise_pred - if self.cpu_offload: - if 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() - self.transformer_weights.non_block_weights_to_cpu() + try: + current_start_frame = self.scheduler.seg_index * self.scheduler.num_frame_per_chunk + current_end_frame = (self.scheduler.seg_index + 1) * self.scheduler.num_frame_per_chunk + noise_pred = self._infer_cond_uncond(inputs, infer_condition=True) + self.scheduler.noise_pred[:, current_start_frame:current_end_frame] = noise_pred + finally: + if self.cpu_offload and self.offload_granularity != "block": + if 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() + self.transformer_weights.non_block_weights_to_cpu() diff --git a/lightx2v/models/runners/wan/wan_sf_runner.py b/lightx2v/models/runners/wan/wan_sf_runner.py index 591b0207a..1a32fc825 100755 --- a/lightx2v/models/runners/wan/wan_sf_runner.py +++ b/lightx2v/models/runners/wan/wan_sf_runner.py @@ -151,18 +151,19 @@ def run_main(self, total_steps=None): metrics_func=monitor_cli.lightx2v_run_segments_end2end_duration, metrics_labels=["DefaultRunner"], ): - self.check_stop() - self.init_run_segment(segment_idx) - latents = self.run_segment(segment_idx) - - with ProfilingContext4DebugL1("step_pre_in_rerun"): - self.model.scheduler.step_pre( - seg_index=segment_idx, - step_index=self.model.scheduler.infer_steps - 1, - is_rerun=True, - ) - with ProfilingContext4DebugL1("infer_main_in_rerun"): - self.model.infer(self.inputs) + with self.transformer_offload_session(): + self.check_stop() + self.init_run_segment(segment_idx) + latents = self.run_segment(segment_idx) + + with ProfilingContext4DebugL1("step_pre_in_rerun"): + self.model.scheduler.step_pre( + seg_index=segment_idx, + step_index=self.model.scheduler.infer_steps - 1, + is_rerun=True, + ) + with ProfilingContext4DebugL1("infer_main_in_rerun"): + self.model.infer(self.inputs) vae_decoder.submit(self.decode_segment_latents, segment_idx, latents) torch_device_module.empty_cache() From 80b68af4d1d08ffd6f7204ab85944a8aa723ce1f Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:46:57 +0800 Subject: [PATCH 17/31] refactor(offload): read Wan Dancer granularity from offload plan --- lightx2v/models/networks/wan/dancer_model.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/lightx2v/models/networks/wan/dancer_model.py b/lightx2v/models/networks/wan/dancer_model.py index 4e18e0a87..ac2e3bbfc 100644 --- a/lightx2v/models/networks/wan/dancer_model.py +++ b/lightx2v/models/networks/wan/dancer_model.py @@ -1,3 +1,4 @@ +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.lora_adapter import LoraAdapter from lightx2v.models.networks.wan.infer.dancer import ( WanDancerPostInfer, @@ -65,7 +66,7 @@ def apply_merged_lora(self, lora_configs): def _init_infer_class(self): if self.config.get("feature_caching", "NoCaching") != "NoCaching": raise NotImplementedError("Wan-Dancer parity mode requires feature_caching=NoCaching.") - if self.config.get("cpu_offload", False) and self.config.get("offload_granularity", "block") not in {"block", "model"}: + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) not in {"block", "model"}: raise NotImplementedError("Wan-Dancer supports block/model offload.") self.pre_infer_class = WanDancerPreInfer self.post_infer_class = WanDancerPostInfer From 45bb455235d333d26a9963d48b7d69443405523b Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:47:45 +0800 Subject: [PATCH 18/31] feat(offload): migrate Wan Audio block scheduling --- lightx2v/models/networks/wan/audio_model.py | 33 +++++---- .../models/runners/wan/wan_audio_runner.py | 67 ++++++++++--------- 2 files changed, 52 insertions(+), 48 deletions(-) diff --git a/lightx2v/models/networks/wan/audio_model.py b/lightx2v/models/networks/wan/audio_model.py index c6d9011b1..3515af7b1 100755 --- a/lightx2v/models/networks/wan/audio_model.py +++ b/lightx2v/models/networks/wan/audio_model.py @@ -116,31 +116,30 @@ def _init_infer_class(self): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == 0: self.to_cuda() elif self.offload_granularity != "model": self.pre_weight.to_cuda() self.transformer_weights.non_block_weights_to_cuda() + try: + pre_infer_out = self.pre_infer.infer(self.pre_weight, inputs) + if self.config["seq_parallel"] and not inputs.get("_ar_ref_prefill", False): + pre_infer_out = self._seq_parallel_pre_process(pre_infer_out) - pre_infer_out = self.pre_infer.infer(self.pre_weight, inputs) - if self.config["seq_parallel"] and not inputs.get("_ar_ref_prefill", False): - pre_infer_out = self._seq_parallel_pre_process(pre_infer_out) + x = self.transformer_infer.infer(self.transformer_weights, pre_infer_out) - x = self.transformer_infer.infer(self.transformer_weights, pre_infer_out) - - if inputs.get("_ar_ref_prefill", False): - noise_pred = None - else: + if inputs.get("_ar_ref_prefill", False): + return None if self.config["seq_parallel"]: x = self._seq_parallel_post_process(x) noise_pred = self.post_infer.infer(x, pre_infer_out)[0] 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: - self.to_cpu() - elif self.offload_granularity != "model": - self.pre_weight.to_cpu() - self.transformer_weights.non_block_weights_to_cpu() - return noise_pred + return noise_pred + finally: + if self.cpu_offload and self.offload_granularity != "block": + if 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() + self.transformer_weights.non_block_weights_to_cpu() diff --git a/lightx2v/models/runners/wan/wan_audio_runner.py b/lightx2v/models/runners/wan/wan_audio_runner.py index d1fb5e228..1e42baa01 100755 --- a/lightx2v/models/runners/wan/wan_audio_runner.py +++ b/lightx2v/models/runners/wan/wan_audio_runner.py @@ -20,6 +20,7 @@ from lightx2v.models.input_encoders.hf.seko_audio.audio_adapter import AudioAdapter, CausalAudioSlidingProcessor from lightx2v.models.input_encoders.hf.seko_audio.audio_encoder import SekoAudioEncoderModel from lightx2v.models.networks.wan.audio_model import WanAudioARModel, WanAudioModel +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.wan.wan_runner import WanRunner, build_wan_model_with_lora from lightx2v.models.schedulers.wan.audio.scheduler import EulerScheduler, WanAudioARScheduler from lightx2v.models.video_encoders.hf.wan.vae_2_2 import Wan2_2_VAE @@ -882,14 +883,15 @@ def run_vae_cached_decoder_withflag(self, latents, is_first: bool, is_last: bool def run_clip(self): infer_steps = self.model.scheduler.infer_steps - for step_index in range(infer_steps): - logger.info(f"==> step_index: {step_index + 1} / {infer_steps}") - with ProfilingContext4DebugL1("step_pre"): - self.model.scheduler.step_pre(step_index=step_index) - with ProfilingContext4DebugL1("🚀 infer_main"): - self.model.infer(self.inputs) - with ProfilingContext4DebugL1("step_post"): - self.model.scheduler.step_post() + with self.transformer_offload_session(): + for step_index in range(infer_steps): + logger.info(f"==> step_index: {step_index + 1} / {infer_steps}") + with ProfilingContext4DebugL1("step_pre"): + self.model.scheduler.step_pre(step_index=step_index) + with ProfilingContext4DebugL1("🚀 infer_main"): + self.model.infer(self.inputs) + with ProfilingContext4DebugL1("step_post"): + self.model.scheduler.step_post() return self.model.scheduler.latents @@ -1326,13 +1328,15 @@ def prefill_reference_kv(self): ref_frames = 1 if ref_latents is None else int(ref_latents.shape[1]) self.inputs["_ar_ref_prefill"] = True try: - for step_index in range(self.model.scheduler.infer_steps): - self.model.kv_cache_manager.current_step = step_index - self.model.scheduler.step_pre_ref(step_index, ref_frames) - self.model.infer(self.inputs) + with self.transformer_offload_session(): + for step_index in range(self.model.scheduler.infer_steps): + self.model.kv_cache_manager.current_step = step_index + self.model.scheduler.step_pre_ref(step_index, ref_frames) + self.model.infer(self.inputs) finally: self.inputs.pop("_ar_ref_prefill", None) + @keep_transformer_weights_loaded def run_segment(self, segment_idx=0): infer_steps = self.model.scheduler.infer_steps chunk_size = int(self.model.scheduler.chunk_size) @@ -1360,25 +1364,26 @@ def run_segment(self, segment_idx=0): return xt def run_stream_segment(self, segment_idx=0): - infer_steps = self.model.scheduler.infer_steps - self.model.scheduler.set_timesteps(infer_steps, device=AI_DEVICE) - chunk_noise = self._make_ar_chunk_noise(segment_idx) - if chunk_noise is None: - xt = self.model.scheduler.noise.to(AI_DEVICE) - else: - xt = chunk_noise.to(AI_DEVICE) - for step_index in range(infer_steps): - # logger.info(f"==> stream chunk: {segment_idx + 1}, step_index: {step_index + 1} / {infer_steps}") - self.model.kv_cache_manager.current_step = step_index - with ProfilingContext4DebugL1("step_pre"): - self.model.scheduler.step_pre(segment_idx, step_index, xt) - with ProfilingContext4DebugL1("🚀 infer_main"): - self.model.infer(self.inputs) - with ProfilingContext4DebugL1("step_post"): - xt = self.model.scheduler.step_post(xt) - if self.progress_callback: - self.progress_callback(((step_index + 1) / infer_steps) * 100, 100) - return xt + with self.transformer_offload_session(): + infer_steps = self.model.scheduler.infer_steps + self.model.scheduler.set_timesteps(infer_steps, device=AI_DEVICE) + chunk_noise = self._make_ar_chunk_noise(segment_idx) + if chunk_noise is None: + xt = self.model.scheduler.noise.to(AI_DEVICE) + else: + xt = chunk_noise.to(AI_DEVICE) + for step_index in range(infer_steps): + # logger.info(f"==> stream chunk: {segment_idx + 1}, step_index: {step_index + 1} / {infer_steps}") + self.model.kv_cache_manager.current_step = step_index + with ProfilingContext4DebugL1("step_pre"): + self.model.scheduler.step_pre(segment_idx, step_index, xt) + with ProfilingContext4DebugL1("🚀 infer_main"): + self.model.infer(self.inputs) + with ProfilingContext4DebugL1("step_post"): + xt = self.model.scheduler.step_post(xt) + if self.progress_callback: + self.progress_callback(((step_index + 1) / infer_steps) * 100, 100) + return xt def decode_segment_latents(self, segment_idx: int, segment_latents: torch.Tensor) -> torch.Tensor: is_first = segment_idx == 0 From d9a0477e8dfcc4559e657eb17da2d0bfc7ca6fca Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:47:56 +0800 Subject: [PATCH 19/31] feat(offload): migrate Wan DreamZero block scheduling --- .../wan/infer/dreamzero/transformer_infer.py | 13 +++++++++---- lightx2v/models/runners/wan/wan_dreamzero_runner.py | 11 +++++++++-- 2 files changed, 18 insertions(+), 6 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/dreamzero/transformer_infer.py b/lightx2v/models/networks/wan/infer/dreamzero/transformer_infer.py index 509e7c3a1..bd17d1b94 100644 --- a/lightx2v/models/networks/wan/infer/dreamzero/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/dreamzero/transformer_infer.py @@ -5,7 +5,7 @@ from lightx2v.common.ops.rope import FlashInferRope, RopeTemplate, TorchRealRope from lightx2v.models.networks.wan.infer.dreamzero.pre_infer import _category_linear -from lightx2v.models.networks.wan.infer.transformer_infer import WanTransformerInfer +from lightx2v.models.networks.wan.infer.offload.transformer_infer import WanOffloadTransformerInfer from lightx2v.models.networks.wan.infer.triton_ops import apply_rotary_embedding from lightx2v.utils.envs import GET_DTYPE from lightx2v.utils.registry_factory import ROPE_REGISTER @@ -63,7 +63,7 @@ def apply_single(self, x, freqs, **kwargs): return output -class DreamZeroTransformerInfer(WanTransformerInfer): +class DreamZeroTransformerInfer(WanOffloadTransformerInfer): def __init__(self, config): super().__init__(config) self.num_action_per_block = int(config.get("num_action_per_block", config.get("action_horizon", 24))) @@ -362,9 +362,14 @@ def infer_block(self, block, x, pre_infer_out, kv_cache): def infer_main_blocks(self, blocks, pre_infer_out, kv_cache): x = pre_infer_out.x - for block_idx in range(len(blocks)): + + def run_dreamzero_block(block_idx, block): + nonlocal x self.block_idx = block_idx - x = self.infer_block(blocks[block_idx], x, pre_infer_out, kv_cache) + x = self.infer_block(block, x, pre_infer_out, kv_cache) + return x + + self.run_blocks_with_offload(blocks, run_dreamzero_block) return x def infer_action_decoder(self, weights, x, pre_infer_out): diff --git a/lightx2v/models/runners/wan/wan_dreamzero_runner.py b/lightx2v/models/runners/wan/wan_dreamzero_runner.py index 612ab13cc..c56ed49c8 100644 --- a/lightx2v/models/runners/wan/wan_dreamzero_runner.py +++ b/lightx2v/models/runners/wan/wan_dreamzero_runner.py @@ -571,13 +571,20 @@ def run_chunk(self, frame_indices): self.model.clear_cache(self.cache_name, clear_pre_infer=False) observed_latents = None + image_latent = None if self.current_start_frame == 0: self.clip_feas, self.ys, image_latent = self._encode_first_frame_condition(videos) - self._run_cache_warmup(image_latent, 0, 1) - self.current_start_frame += 1 else: observed_latents = self._encode_observed_latents(videos) + with self.transformer_offload_session(): + return self._run_chunk_dit(image_latent, observed_latents) + + def _run_chunk_dit(self, image_latent, observed_latents): + if self.current_start_frame == 0: + self._run_cache_warmup(image_latent, 0, 1) + self.current_start_frame += 1 + if self.current_start_frame != 1 and observed_latents is not None: current_ref_latents = observed_latents[:, :, -self.num_frame_per_block :] self._run_cache_warmup(current_ref_latents, self.current_start_frame - self.num_frame_per_block, self.num_frame_per_block) From 95aa86111c127ee623f004ff46bd89dd8eadf5f8 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:48:02 +0800 Subject: [PATCH 20/31] feat(offload): migrate Wan InfiniteTalk block scheduling --- .../infer/infinitetalk/transformer_infer.py | 3 +- .../models/networks/wan/infinitetalk_model.py | 34 +++++++++++-------- .../infinitetalk/transformer_weights.py | 24 +++++++++---- .../runners/wan/wan_infinitetalk_runner.py | 2 ++ 4 files changed, 41 insertions(+), 22 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/infinitetalk/transformer_infer.py b/lightx2v/models/networks/wan/infer/infinitetalk/transformer_infer.py index c76b34e16..50e671fcd 100755 --- a/lightx2v/models/networks/wan/infer/infinitetalk/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/infinitetalk/transformer_infer.py @@ -3,6 +3,7 @@ import torch.nn.functional as F from einops import rearrange +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.wan.infer.offload.transformer_infer import WanOffloadTransformerInfer from lightx2v.utils.envs import GET_DTYPE @@ -24,7 +25,7 @@ def normalize_and_scale(column, source_range, target_range, epsilon=1e-8): class WanInfiniteTalkTransformerInfer(WanOffloadTransformerInfer): def __init__(self, config): - offload_granularity = config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(config) if config.get("cpu_offload", False) and offload_granularity not in {"block", "model"}: raise NotImplementedError(f"InfiniteTalk currently supports block/model offload, not {offload_granularity} offload.") super().__init__(config) diff --git a/lightx2v/models/networks/wan/infinitetalk_model.py b/lightx2v/models/networks/wan/infinitetalk_model.py index d1efe8c29..238ef976b 100755 --- a/lightx2v/models/networks/wan/infinitetalk_model.py +++ b/lightx2v/models/networks/wan/infinitetalk_model.py @@ -3,6 +3,7 @@ import torch.nn.functional as F from loguru import logger +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.wan.infer.infinitetalk.pre_infer import WanInfiniteTalkPreInfer from lightx2v.models.networks.wan.infer.infinitetalk.transformer_infer import WanInfiniteTalkTransformerInfer from lightx2v.models.networks.wan.infer.post_infer import WanPostInfer @@ -24,7 +25,7 @@ def __init__(self, model_path, config, device, lora_path=None, lora_strength=1.0 def _init_infer_class(self): if self.config.get("feature_caching", "NoCaching") != "NoCaching": raise NotImplementedError("InfiniteTalk parity path requires feature_caching=NoCaching.") - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if self.config.get("cpu_offload", False) and offload_granularity not in {"block", "model"}: raise NotImplementedError(f"InfiniteTalk currently supports block/model offload, not {offload_granularity} offload.") self.pre_infer_class = WanInfiniteTalkPreInfer @@ -116,16 +117,26 @@ def _run_infinitetalk_cfg_parallel(self, inputs, branches): @torch.no_grad() def infer(self, inputs): + try: + if self.cpu_offload and self.offload_granularity != "block": + if self.offload_granularity == "model" and self.scheduler.step_index == 0: + self.to_cuda() + elif self.offload_granularity != "model": + self.pre_weight.to_cuda() + self.transformer_weights.non_block_weights_to_cuda() + self._infer(inputs) + finally: + if self.cpu_offload and self.offload_granularity != "block": + if 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() + self.transformer_weights.non_block_weights_to_cpu() + + def _infer(self, inputs): if self.config.get("use_apg", False): raise NotImplementedError("InfiniteTalk APG is not implemented in the LightX2V parity path yet.") - if self.cpu_offload: - if self.offload_granularity == "model" and self.scheduler.step_index == 0: - self.to_cuda() - elif self.offload_granularity != "model": - self.pre_weight.to_cuda() - self.transformer_weights.non_block_weights_to_cuda() - if not self.config["enable_cfg"]: noise_pred_cond = self._infer_infinitetalk_branch(inputs, infer_condition=True, use_audio=True) noise_pred_guided = noise_pred_cond @@ -174,13 +185,6 @@ def infer(self, inputs): self.scheduler.noise_pred_guided = noise_pred_guided self.scheduler.noise_pred = -noise_pred_guided - if self.cpu_offload: - if 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() - self.transformer_weights.non_block_weights_to_cpu() - @torch.no_grad() def _seq_parallel_pre_process(self, pre_infer_out): x = pre_infer_out.x diff --git a/lightx2v/models/networks/wan/weights/infinitetalk/transformer_weights.py b/lightx2v/models/networks/wan/weights/infinitetalk/transformer_weights.py index 1cdba30c2..f148b2d0f 100755 --- a/lightx2v/models/networks/wan/weights/infinitetalk/transformer_weights.py +++ b/lightx2v/models/networks/wan/weights/infinitetalk/transformer_weights.py @@ -1,6 +1,7 @@ import torch from lightx2v.common.modules.weight_module import WeightModule, WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.wan.weights.transformer_weights import WanCrossAttention, WanFFN, WanSelfAttention from lightx2v.utils.registry_factory import ATTN_WEIGHT_REGISTER, LN_WEIGHT_REGISTER, MM_WEIGHT_REGISTER, ROPE_REGISTER, TENSOR_REGISTER @@ -28,7 +29,8 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): for i in range(self.blocks_num) ] ) - self.register_offload_buffers(config, lazy_load_path, lora_path) + slot_count = self.register_offload_block_group(config, "blocks", self.blocks) + self.register_offload_buffers(config, lazy_load_path, lora_path, slot_count) self.add_module("blocks", self.blocks) self.register_parameter("norm", LN_WEIGHT_REGISTER[config.get("layer_norm_type", "torch")]()) @@ -42,17 +44,24 @@ def __init__(self, config, lazy_load_path=None, lora_path=None): ) self.register_parameter("head_modulation", TENSOR_REGISTER["Default"]("head.modulation")) - def register_offload_buffers(self, config, lazy_load_path, lora_path): + def register_offload_buffers(self, config, lazy_load_path, lora_path, slot_count): if not config.get("cpu_offload", False): return - offload_granularity = config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(config) if offload_granularity == "model": return if offload_granularity != "block": raise NotImplementedError(f"InfiniteTalk currently supports block/model offload, not {offload_granularity} offload.") - self.offload_blocks_num = 2 + self.offload_blocks_num = slot_count + self.offload_block_cuda_buffers = None + self.offload_block_cpu_buffers = None + self.offload_phase_cuda_buffers = None + self.offload_phase_cpu_buffers = None + if slot_count == 0: + return + self.offload_block_cuda_buffers = WeightModuleList( [ WanInfiniteTalkTransformerAttentionBlock( @@ -71,7 +80,6 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cuda_buffers", self.offload_block_cuda_buffers) - self.offload_phase_cuda_buffers = None if self.lazy_load: self.offload_block_cpu_buffers = WeightModuleList( @@ -92,7 +100,11 @@ def register_offload_buffers(self, config, lazy_load_path, lora_path): ] ) self.add_module("offload_block_cpu_buffers", self.offload_block_cpu_buffers) - self.offload_phase_cpu_buffers = None + self.register_offload_block_buffers( + "blocks", + self.offload_block_cuda_buffers, + self.offload_block_cpu_buffers, + ) def non_block_weights_to_cuda(self): self.norm.to_cuda() diff --git a/lightx2v/models/runners/wan/wan_infinitetalk_runner.py b/lightx2v/models/runners/wan/wan_infinitetalk_runner.py index 782557de9..c37c44f3f 100644 --- a/lightx2v/models/runners/wan/wan_infinitetalk_runner.py +++ b/lightx2v/models/runners/wan/wan_infinitetalk_runner.py @@ -16,6 +16,7 @@ from lightx2v.models.input_encoders.hf.infinitetalk.audio_encoder import InfiniteTalkAudioEncoder from lightx2v.models.networks.wan.infinitetalk_model import WanInfiniteTalkModel +from lightx2v.models.runners.base_runner import keep_transformer_weights_loaded from lightx2v.models.runners.wan.wan_runner import WanRunner from lightx2v.models.schedulers.wan.infinitetalk.scheduler import InfiniteTalkScheduler from lightx2v.server.metrics import monitor_cli @@ -998,6 +999,7 @@ def init_run_segment(self, segment_idx): "ref_target_masks": ref_target_masks, } + @keep_transformer_weights_loaded def run_segment(self, segment_idx=0): self._run_dit_clip(self.dit_inputs) return self.scheduler.latents From d25a6857d48e96e3b85a3a59a623cf9732585462 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:48:14 +0800 Subject: [PATCH 21/31] feat(offload): migrate Wan Lingbot block scheduling --- .../models/networks/wan/infer/lingbot/transformer_infer.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/lingbot/transformer_infer.py b/lightx2v/models/networks/wan/infer/lingbot/transformer_infer.py index 701ad1c08..0c18c1596 100755 --- a/lightx2v/models/networks/wan/infer/lingbot/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/lingbot/transformer_infer.py @@ -1,12 +1,12 @@ import torch -from lightx2v.models.networks.wan.infer.transformer_infer import WanTransformerInfer +from lightx2v.models.networks.wan.infer.offload.transformer_infer import WanOffloadTransformerInfer from lightx2v_platform.base.global_var import AI_DEVICE torch_device_module = getattr(torch, AI_DEVICE) -class WanLingbotTransformerInfer(WanTransformerInfer): +class WanLingbotTransformerInfer(WanOffloadTransformerInfer): def infer_block(self, block, x, pre_infer_out): if hasattr(block.compute_phases[0], "before_proj") and block.compute_phases[0].before_proj.weight is not None: x = block.compute_phases[0].before_proj.apply(x) + pre_infer_out.x From c8cfa9c3decc334ef66bd20a70c61aa0f49ec0b8 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:48:26 +0800 Subject: [PATCH 22/31] feat(offload): migrate Wan Lingbot Fast block scheduling --- .../models/networks/wan/lingbot_fast_model.py | 27 +++++++++---------- .../runners/wan/wan_lingbot_fast_runner.py | 25 ++++++++--------- 2 files changed, 26 insertions(+), 26 deletions(-) diff --git a/lightx2v/models/networks/wan/lingbot_fast_model.py b/lightx2v/models/networks/wan/lingbot_fast_model.py index a085bb014..e5cecec86 100644 --- a/lightx2v/models/networks/wan/lingbot_fast_model.py +++ b/lightx2v/models/networks/wan/lingbot_fast_model.py @@ -17,22 +17,21 @@ def _init_infer_class(self): @torch.no_grad() def infer(self, inputs): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model" and self.scheduler.step_index == 0: self.to_cuda() elif self.offload_granularity != "model": self.pre_weight.to_cuda() self.transformer_weights.non_block_weights_to_cuda() - - current_start_frame = self.scheduler.seg_index * self.scheduler.num_frame_per_chunk - current_end_frame = (self.scheduler.seg_index + 1) * self.scheduler.num_frame_per_chunk - noise_pred = self._infer_cond_uncond(inputs, infer_condition=True) - - self.scheduler.noise_pred[:, current_start_frame:current_end_frame] = noise_pred - - if self.cpu_offload: - if 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() - self.transformer_weights.non_block_weights_to_cpu() + try: + current_start_frame = self.scheduler.seg_index * self.scheduler.num_frame_per_chunk + current_end_frame = (self.scheduler.seg_index + 1) * self.scheduler.num_frame_per_chunk + noise_pred = self._infer_cond_uncond(inputs, infer_condition=True) + self.scheduler.noise_pred[:, current_start_frame:current_end_frame] = noise_pred + finally: + if self.cpu_offload and self.offload_granularity != "block": + if 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() + self.transformer_weights.non_block_weights_to_cpu() diff --git a/lightx2v/models/runners/wan/wan_lingbot_fast_runner.py b/lightx2v/models/runners/wan/wan_lingbot_fast_runner.py index 458b48f20..fc4c5295d 100755 --- a/lightx2v/models/runners/wan/wan_lingbot_fast_runner.py +++ b/lightx2v/models/runners/wan/wan_lingbot_fast_runner.py @@ -127,18 +127,19 @@ def run_main(self, total_steps=None): metrics_func=monitor_cli.lightx2v_run_segments_end2end_duration, metrics_labels=["DefaultRunner"], ): - self.check_stop() - self.init_run_segment(segment_idx) - latents = self.run_segment(segment_idx) - - with ProfilingContext4DebugL1("step_pre_in_rerun"): - self.model.scheduler.step_pre( - seg_index=segment_idx, - step_index=self.model.scheduler.infer_steps - 1, - is_rerun=True, - ) - with ProfilingContext4DebugL1("infer_main_in_rerun"): - self.model.infer(self.inputs) + with self.transformer_offload_session(): + self.check_stop() + self.init_run_segment(segment_idx) + latents = self.run_segment(segment_idx) + + with ProfilingContext4DebugL1("step_pre_in_rerun"): + self.model.scheduler.step_pre( + seg_index=segment_idx, + step_index=self.model.scheduler.infer_steps - 1, + is_rerun=True, + ) + with ProfilingContext4DebugL1("infer_main_in_rerun"): + self.model.infer(self.inputs) vae_decoder.submit(self.decode_segment_latents, segment_idx, latents) torch.cuda.empty_cache() From 2373880e89990ef33eb02c2438ba391b8a55e3b2 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:48:32 +0800 Subject: [PATCH 23/31] feat(offload): migrate Wan Lingbot VA block scheduling --- .../wan/infer/lingbot_va/transformer_infer.py | 13 ++++-- .../models/networks/wan/lingbot_va_model.py | 4 +- .../runners/wan/wan_lingbot_va_runner.py | 40 +++++++++---------- 3 files changed, 31 insertions(+), 26 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/lingbot_va/transformer_infer.py b/lightx2v/models/networks/wan/infer/lingbot_va/transformer_infer.py index dad1329e3..245b56bc8 100644 --- a/lightx2v/models/networks/wan/infer/lingbot_va/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/lingbot_va/transformer_infer.py @@ -3,7 +3,7 @@ import torch import torch.nn.functional as F -from lightx2v.models.networks.wan.infer.transformer_infer import WanTransformerInfer +from lightx2v.models.networks.wan.infer.offload.transformer_infer import WanOffloadTransformerInfer from lightx2v.utils.envs import GET_DTYPE @@ -11,7 +11,7 @@ def _token_modulation(x: torch.Tensor) -> torch.Tensor: return x.reshape(-1, x.shape[-1]) -class LingbotVATransformerInfer(WanTransformerInfer): +class LingbotVATransformerInfer(WanOffloadTransformerInfer): def __init__(self, config): super().__init__(config) self.kv_cache_manager = None @@ -96,9 +96,14 @@ def infer_block(self, block, x, pre_infer_out, update_cache=0, cache_name="pos") def infer_main_blocks(self, blocks, pre_infer_out, update_cache=0, cache_name="pos"): x = pre_infer_out.x - for block_idx in range(len(blocks)): + + def run_lingbot_va_block(block_idx, block): + nonlocal x self.block_idx = block_idx - x = self.infer_block(blocks[block_idx], x, pre_infer_out, update_cache=update_cache, cache_name=cache_name) + x = self.infer_block(block, x, pre_infer_out, update_cache=update_cache, cache_name=cache_name) + return x + + self.run_blocks_with_offload(blocks, run_lingbot_va_block) return x def infer_non_blocks(self, weights, x, pre_infer_out, action_mode=False): diff --git a/lightx2v/models/networks/wan/lingbot_va_model.py b/lightx2v/models/networks/wan/lingbot_va_model.py index 931037f52..255a522dd 100644 --- a/lightx2v/models/networks/wan/lingbot_va_model.py +++ b/lightx2v/models/networks/wan/lingbot_va_model.py @@ -58,7 +58,7 @@ def _cache_names(cls, cache_name): return (cache_name, cls.cfg_cache_name(cache_name, True), cls.cfg_cache_name(cache_name, False)) def _to_cuda_for_lingbot_va(self): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model": self.to_cuda() else: @@ -66,7 +66,7 @@ def _to_cuda_for_lingbot_va(self): self.transformer_weights.non_block_weights_to_cuda() def _to_cpu_for_lingbot_va(self): - if self.cpu_offload: + if self.cpu_offload and self.offload_granularity != "block": if self.offload_granularity == "model": self.to_cpu() else: diff --git a/lightx2v/models/runners/wan/wan_lingbot_va_runner.py b/lightx2v/models/runners/wan/wan_lingbot_va_runner.py index 130904dd6..5a2127226 100644 --- a/lightx2v/models/runners/wan/wan_lingbot_va_runner.py +++ b/lightx2v/models/runners/wan/wan_lingbot_va_runner.py @@ -408,24 +408,25 @@ def run_segment(self, segment_idx=0): self.scheduler.bind_step_inputs(self.inputs, self._build_video_step_inputs) self.scheduler.bind_noise_pred_processor(self._postprocess_video_noise_pred) self.model.set_scheduler(self.scheduler) - self._run_scheduler_loop(self.scheduler) - latents = self.scheduler.latents - - action_cond = torch.zeros([1, self.config["action_dim"], 1, self.num_action_per_frame, 1], device=AI_DEVICE, dtype=GET_DTYPE()) if self.frame_st_id == 0 else None - self.action_scheduler.generator = self.scheduler.generator - self.action_scheduler.prepare_loop( - infer_steps=self.config["action_infer_steps"], - device=AI_DEVICE, - latent_shape=action_shape, - seed=self.input_info.seed, - dtype=GET_DTYPE(), - cond_latent=action_cond, - ) - self.action_scheduler.bind_step_inputs(self.inputs, self._build_action_step_inputs) - self.action_scheduler.bind_noise_pred_processor(self._postprocess_action_noise_pred) - self.model.set_scheduler(self.action_scheduler) - self._run_scheduler_loop(self.action_scheduler) - actions = self.action_scheduler.latents + with self.transformer_offload_session(): + self._run_scheduler_loop(self.scheduler) + latents = self.scheduler.latents + + action_cond = torch.zeros([1, self.config["action_dim"], 1, self.num_action_per_frame, 1], device=AI_DEVICE, dtype=GET_DTYPE()) if self.frame_st_id == 0 else None + self.action_scheduler.generator = self.scheduler.generator + self.action_scheduler.prepare_loop( + infer_steps=self.config["action_infer_steps"], + device=AI_DEVICE, + latent_shape=action_shape, + seed=self.input_info.seed, + dtype=GET_DTYPE(), + cond_latent=action_cond, + ) + self.action_scheduler.bind_step_inputs(self.inputs, self._build_action_step_inputs) + self.action_scheduler.bind_noise_pred_processor(self._postprocess_action_noise_pred) + self.model.set_scheduler(self.action_scheduler) + self._run_scheduler_loop(self.action_scheduler) + actions = self.action_scheduler.latents actions[:, ~self.action_mask] *= 0 self.model.set_scheduler(self.scheduler) @@ -511,8 +512,7 @@ def end_run(self): del self.inputs self.input_info = None if self.config.get("lazy_load", False) or self.config.get("unload_modules", False): - if hasattr(self.model.transformer_infer, "offload_manager"): - del self.model.transformer_infer.offload_manager + self.model.transformer_infer.clear_offload_managers() del self.model torch_device_module.empty_cache() gc.collect() From 9757757a3bc6abe42ee165b97c96f3f754504533 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:48:44 +0800 Subject: [PATCH 24/31] feat(offload): migrate Wan VACE block scheduling --- .../wan/infer/vace/transformer_infer.py | 36 ++++++++++++++----- lightx2v/models/networks/wan/vace_model.py | 16 --------- .../wan/weights/vace/transformer_weights.py | 26 ++++++++------ 3 files changed, 43 insertions(+), 35 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/vace/transformer_infer.py b/lightx2v/models/networks/wan/infer/vace/transformer_infer.py index 7c895158d..958cc35a7 100755 --- a/lightx2v/models/networks/wan/infer/vace/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/vace/transformer_infer.py @@ -2,6 +2,7 @@ import torch.distributed as dist import torch.nn.functional as F +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.wan.infer.offload.transformer_infer import WanOffloadTransformerInfer from lightx2v.utils.envs import * @@ -21,9 +22,13 @@ def infer(self, weights, pre_infer_out): if self.config.get("seq_parallel", False): pre_infer_out.c = self._chunk_c_for_seq_parallel(pre_infer_out.c, pre_infer_out.x) - self.infer_vace_blocks(weights.vace_blocks, pre_infer_out) - x = self.infer_main_blocks(weights.blocks, pre_infer_out) - return self.infer_non_blocks(weights, x, pre_infer_out.embed) + try: + self.infer_vace_blocks(weights.vace_blocks, pre_infer_out) + x = self.infer_main_blocks(weights.blocks, pre_infer_out) + return self.infer_non_blocks(weights, x, pre_infer_out.embed) + finally: + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "block": + self.clear_block_offload_inputs(pre_infer_out) def _chunk_c_for_seq_parallel(self, c, x): """Chunk c along sequence dimension to match x in seq parallel mode.""" @@ -53,12 +58,25 @@ def vace_pre_process(self, patch_embedding, vace_context): def infer_vace_blocks(self, vace_blocks, pre_infer_out): pre_infer_out.adapter_args["hints"] = [] self.infer_state = "vace" - if hasattr(self, "offload_manager"): - self.offload_manager.init_cuda_buffer(self.vace_offload_block_cuda_buffers, self.vace_offload_phase_cuda_buffers) - self.infer_func(vace_blocks, pre_infer_out.c, pre_infer_out) - self.infer_state = "base" - if hasattr(self, "offload_manager"): - self.offload_manager.init_cuda_buffer(self.offload_block_cuda_buffers, self.offload_phase_cuda_buffers) + try: + if not self.config.get("cpu_offload", False) or get_offload_granularity(self.config) != "block": + self.infer_func(vace_blocks, pre_infer_out.c, pre_infer_out) + return + + c = pre_infer_out.c + + def run_vace_block(block_idx, block): + nonlocal c + self.block_idx = block_idx + c = self.run_block(block_idx, block, c, pre_infer_out) + return c + + self.run_blocks_with_offload(vace_blocks, run_vace_block) + finally: + self.infer_state = "base" + + def get_compile_block_key(self, block_idx, block): + return self.infer_state, super().get_compile_block_key(block_idx, block) def post_process(self, x, y, c_gate_msa, pre_infer_out): x = super().post_process(x, y, c_gate_msa, pre_infer_out) diff --git a/lightx2v/models/networks/wan/vace_model.py b/lightx2v/models/networks/wan/vace_model.py index 454ba574a..535f41073 100755 --- a/lightx2v/models/networks/wan/vace_model.py +++ b/lightx2v/models/networks/wan/vace_model.py @@ -19,22 +19,6 @@ class WanVaceModel(WanModel): def __init__(self, model_path, config, device, model_type="wan2.1"): super().__init__(model_path, config, device, model_type) - def _init_infer(self): - super()._init_infer() - if hasattr(self.transformer_infer, "offload_manager"): - self._init_offload_manager() - - def _init_offload_manager(self): - self.transformer_infer.offload_block_cuda_buffers = self.transformer_weights.offload_block_cuda_buffers - self.transformer_infer.offload_phase_cuda_buffers = self.transformer_weights.offload_phase_cuda_buffers - self.transformer_infer.vace_offload_block_cuda_buffers = self.transformer_weights.vace_offload_block_cuda_buffers - self.transformer_infer.vace_offload_phase_cuda_buffers = self.transformer_weights.vace_offload_phase_cuda_buffers - if self.lazy_load: - self.transformer_infer.offload_block_cpu_buffers = self.transformer_weights.offload_block_cpu_buffers - self.transformer_infer.offload_phase_cpu_buffers = self.transformer_weights.offload_phase_cpu_buffers - self.transformer_infer.vace_offload_block_cpu_buffers = self.transformer_weights.vace_offload_block_cpu_buffers - self.transformer_infer.vace_offload_phase_cpu_buffers = self.transformer_weights.vace_offload_phase_cpu_buffers - def _init_infer_class(self): self.pre_infer_class = WanPreInfer self.post_infer_class = WanPostInfer diff --git a/lightx2v/models/networks/wan/weights/vace/transformer_weights.py b/lightx2v/models/networks/wan/weights/vace/transformer_weights.py index 3ee0ffe36..65971e425 100755 --- a/lightx2v/models/networks/wan/weights/vace/transformer_weights.py +++ b/lightx2v/models/networks/wan/weights/vace/transformer_weights.py @@ -1,4 +1,5 @@ from lightx2v.common.modules.weight_module import WeightModuleList +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.wan.weights.transformer_weights import ( WanTransformerAttentionBlock, WanTransformerWeights, @@ -11,31 +12,36 @@ class WanVaceTransformerWeights(WanTransformerWeights): def __init__(self, config, lazy_load_path=None, lora_path=None): + if config.get("lazy_load", False): + raise NotImplementedError("Wan VACE block offload does not support lazy_load") super().__init__(config, lazy_load_path, lora_path) self.patch_size = (1, 2, 2) - self.register_offload_buffers(config, lazy_load_path, lora_path) self.vace_blocks = WeightModuleList( [WanVaceTransformerAttentionBlock(self.config["vace_layers"][i], i, self.task, self.mm_type, self.config, False, False, "vace_blocks") for i in range(len(self.config["vace_layers"]))] ) + slot_count = self.register_offload_block_group(config, "vace_blocks", self.vace_blocks) + self._register_vace_offload_buffers(config, slot_count) self.add_module("vace_blocks", self.vace_blocks) self.add_module( "vace_patch_embedding", CONV3D_WEIGHT_REGISTER["Default"]("vace_patch_embedding.weight", "vace_patch_embedding.bias", stride=self.patch_size), ) - def register_offload_buffers(self, config, lazy_load_path, lora_path): - super().register_offload_buffers(config, lazy_load_path, lora_path) + def _register_vace_offload_buffers(self, config, slot_count): if config["cpu_offload"]: - if config["offload_granularity"] == "block": + if get_offload_granularity(config) == "block": + self.vace_offload_block_cuda_buffers = None + self.vace_offload_block_cpu_buffers = None + self.vace_offload_phase_cuda_buffers = None + self.vace_offload_phase_cpu_buffers = None + if slot_count == 0: + return self.vace_offload_block_cuda_buffers = WeightModuleList( - [ - WanVaceTransformerAttentionBlock(self.config["vace_layers"][0], 0, self.task, self.mm_type, self.config, True, False, "vace_blocks"), - WanVaceTransformerAttentionBlock(self.config["vace_layers"][0], 0, self.task, self.mm_type, self.config, True, False, "vace_blocks"), - ] + [WanVaceTransformerAttentionBlock(self.config["vace_layers"][0], 0, self.task, self.mm_type, self.config, True, False, "vace_blocks") for _ in range(slot_count)] ) self.add_module("vace_offload_block_cuda_buffers", self.vace_offload_block_cuda_buffers) - self.vace_offload_phase_cuda_buffers = None - elif config["offload_granularity"] == "phase": + self.register_offload_block_buffers("vace_blocks", self.vace_offload_block_cuda_buffers) + elif get_offload_granularity(config) == "phase": raise NotImplementedError def non_block_weights_to_cuda(self): From 119d842f2dc148d0515f3bc1eb57f55cd59d676e Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:50:02 +0800 Subject: [PATCH 25/31] fix(offload): reject unsupported Wan S2V block offload --- lightx2v/models/networks/wan/s2v_model.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/lightx2v/models/networks/wan/s2v_model.py b/lightx2v/models/networks/wan/s2v_model.py index bc81fce11..2780978ed 100644 --- a/lightx2v/models/networks/wan/s2v_model.py +++ b/lightx2v/models/networks/wan/s2v_model.py @@ -26,6 +26,8 @@ def __init__(self, model_path, config, device, lora_path=None, lora_strength=1.0 ) def _init_infer_class(self): + if self.cpu_offload and self.offload_granularity == "block": + raise NotImplementedError("Wan S2V does not support block offload") super()._init_infer_class() self.pre_infer_class = WanS2VPreInfer self.post_infer_class = WanS2VPostInfer From 5a49b98f1d45ee151278c4be21717decb6835f4f Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:50:33 +0800 Subject: [PATCH 26/31] feat(offload): migrate WorldPlay block scheduling --- .../worldplay/infer/transformer_infer.py | 23 +++++++------------ lightx2v/models/networks/worldplay/model.py | 1 + 2 files changed, 9 insertions(+), 15 deletions(-) diff --git a/lightx2v/models/networks/worldplay/infer/transformer_infer.py b/lightx2v/models/networks/worldplay/infer/transformer_infer.py index 6366cacf4..fe9a3d34c 100644 --- a/lightx2v/models/networks/worldplay/infer/transformer_infer.py +++ b/lightx2v/models/networks/worldplay/infer/transformer_infer.py @@ -2,7 +2,7 @@ import torch.nn.functional as F from einops import rearrange -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.hunyuan_video.infer.module_io import ( HunyuanVideo15ImgBranchOutput, HunyuanVideo15TxtBranchOutput, @@ -13,9 +13,6 @@ apply_gate, ) from lightx2v.models.networks.worldplay.prope.camera_rope import prope_qkv -from lightx2v_platform.base.global_var import AI_DEVICE - -torch_device_module = getattr(torch, AI_DEVICE) def modulate_per_token(x, scale, shift): @@ -55,15 +52,13 @@ def __init__(self, config): # Setup offload if enabled if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": self.infer_func = self.infer_with_blocks_offload elif offload_granularity == "model": self.infer_func = self.infer_without_offload else: raise NotImplementedError - if offload_granularity != "model": - self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) @property def _vec_is_per_token(self): @@ -395,15 +390,13 @@ def infer_without_offload(self, weights, infer_module_out): @torch.no_grad() def infer_with_blocks_offload(self, weights, infer_module_out): """Inference with block-level CPU offload.""" - for block_idx in range(self.double_blocks_num): + + def infer_block(block_idx, block_weights): self.block_idx = block_idx - if block_idx == 0: - self.offload_manager.init_first_buffer(weights.double_blocks) - if block_idx < self.double_blocks_num - 1: - self.offload_manager.prefetch_weights(block_idx + 1, weights.double_blocks) - with torch_device_module.stream(self.offload_manager.compute_stream): - infer_module_out.img, infer_module_out.txt = self.infer_double_block(self.offload_manager.cuda_buffers[0], infer_module_out, block_idx=block_idx) - self.offload_manager.swap_blocks() + infer_module_out.img, infer_module_out.txt = self.infer_double_block(block_weights, infer_module_out, block_idx=block_idx) + return infer_module_out.img, infer_module_out.txt + + self.run_blocks_with_offload(weights.double_blocks, infer_block) def set_action_weights(self, action_weights): """Set action weights for ProPE projection access.""" diff --git a/lightx2v/models/networks/worldplay/model.py b/lightx2v/models/networks/worldplay/model.py index 9998caebd..817d1b57f 100644 --- a/lightx2v/models/networks/worldplay/model.py +++ b/lightx2v/models/networks/worldplay/model.py @@ -67,6 +67,7 @@ def _init_weights(self): self.original_weight_dict = weight_dict self.pre_weight = WorldPlayPreWeights(self.config) self.transformer_weights = WorldPlayTransformerWeights(self.config) + self.transformer_weights.validate_offload_block_groups(self.config) self.post_weight = WorldPlayPostWeights(self.config) self._apply_weights() From 78312f4f0d83727121ff9277dc35e82a0ca491c1 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:50:44 +0800 Subject: [PATCH 27/31] feat(offload): migrate WorldPlay BI block scheduling --- .../models/networks/worldplay/bi_model.py | 1 + .../worldplay/infer/bi_transformer_infer.py | 23 +++++++------------ 2 files changed, 9 insertions(+), 15 deletions(-) diff --git a/lightx2v/models/networks/worldplay/bi_model.py b/lightx2v/models/networks/worldplay/bi_model.py index 855104a98..443c561b2 100644 --- a/lightx2v/models/networks/worldplay/bi_model.py +++ b/lightx2v/models/networks/worldplay/bi_model.py @@ -98,6 +98,7 @@ def _init_weights(self): self.original_weight_dict = weight_dict self.pre_weight = WorldPlayPreWeights(self.config) self.transformer_weights = WorldPlayTransformerWeights(self.config) + self.transformer_weights.validate_offload_block_groups(self.config) self.post_weight = WorldPlayPostWeights(self.config) self._apply_weights() diff --git a/lightx2v/models/networks/worldplay/infer/bi_transformer_infer.py b/lightx2v/models/networks/worldplay/infer/bi_transformer_infer.py index f9d3ff51d..2bd9d857a 100644 --- a/lightx2v/models/networks/worldplay/infer/bi_transformer_infer.py +++ b/lightx2v/models/networks/worldplay/infer/bi_transformer_infer.py @@ -2,7 +2,7 @@ import torch.nn.functional as F from einops import rearrange -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.hunyuan_video.infer.module_io import ( HunyuanVideo15ImgBranchOutput, HunyuanVideo15TxtBranchOutput, @@ -13,9 +13,6 @@ apply_gate, ) from lightx2v.models.networks.worldplay.prope.camera_rope import prope_qkv -from lightx2v_platform.base.global_var import AI_DEVICE - -torch_device_module = getattr(torch, AI_DEVICE) def modulate_per_token(x, scale, shift): @@ -65,15 +62,13 @@ def __init__(self, config): # Setup offload if enabled if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": self.infer_func = self.infer_with_blocks_offload elif offload_granularity == "model": self.infer_func = self.infer_without_offload else: raise NotImplementedError - if offload_granularity != "model": - self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) @property def _vec_is_per_token(self): @@ -408,15 +403,13 @@ def infer_without_offload(self, weights, infer_module_out): @torch.no_grad() def infer_with_blocks_offload(self, weights, infer_module_out): """Inference with block-level CPU offload.""" - for block_idx in range(self.double_blocks_num): + + def infer_block(block_idx, block_weights): self.block_idx = block_idx - if block_idx == 0: - self.offload_manager.init_first_buffer(weights.double_blocks) - if block_idx < self.double_blocks_num - 1: - self.offload_manager.prefetch_weights(block_idx + 1, weights.double_blocks) - with torch_device_module.stream(self.offload_manager.compute_stream): - infer_module_out.img, infer_module_out.txt = self.infer_double_block(self.offload_manager.cuda_buffers[0], infer_module_out, block_idx=block_idx) - self.offload_manager.swap_blocks() + infer_module_out.img, infer_module_out.txt = self.infer_double_block(block_weights, infer_module_out, block_idx=block_idx) + return infer_module_out.img, infer_module_out.txt + + self.run_blocks_with_offload(weights.double_blocks, infer_block) def set_action_weights(self, action_weights): """Set action weights for ProPE projection access.""" From 2646d5d9b5650fed1048a245b870a314f6e95d40 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:50:56 +0800 Subject: [PATCH 28/31] feat(offload): migrate WorldPlay AR block scheduling --- .../models/networks/worldplay/ar_model.py | 18 +- .../worldplay/infer/ar_transformer_infer.py | 254 ++++++++---------- .../runners/worldplay/worldplay_ar_runner.py | 3 +- 3 files changed, 121 insertions(+), 154 deletions(-) diff --git a/lightx2v/models/networks/worldplay/ar_model.py b/lightx2v/models/networks/worldplay/ar_model.py index 5609afafc..5905f1e32 100644 --- a/lightx2v/models/networks/worldplay/ar_model.py +++ b/lightx2v/models/networks/worldplay/ar_model.py @@ -65,6 +65,7 @@ def _init_weights(self): self.original_weight_dict = weight_dict self.pre_weight = WorldPlayPreWeights(self.config) self.transformer_weights = WorldPlayTransformerWeights(self.config) + self.transformer_weights.validate_offload_block_groups(self.config) self.post_weight = WorldPlayPostWeights(self.config) self._apply_weights() @@ -135,20 +136,12 @@ def infer_txt(self, inputs, cache_txt=True): if not hasattr(self.transformer_infer, "_kv_cache") or self.transformer_infer._kv_cache is None: self.init_kv_cache() - if self.cpu_offload and self.offload_granularity != "model": - self.pre_weight.to_cuda() - self.transformer_weights.non_block_weights_to_cuda() - # Run text-only pre-processing infer_module_out = self.pre_infer.infer_txt_only(self.pre_weight, inputs) # Cache text KV result = self.transformer_infer.infer_txt(self.transformer_weights, infer_module_out, cache_txt=cache_txt) - if self.cpu_offload and self.offload_granularity != "model": - self.pre_weight.to_cpu() - self.transformer_weights.non_block_weights_to_cpu() - return result @torch.no_grad() @@ -175,11 +168,6 @@ def infer_vision(self, inputs, cache_vision=False): if "action" in pose_output: self.scheduler.action = pose_output["action"] - # Run pre-inference (full, including image) - if self.cpu_offload and self.offload_granularity != "model": - self.pre_weight.to_cuda() - self.transformer_weights.non_block_weights_to_cuda() - infer_module_out = self.pre_infer.infer(self.pre_weight, inputs) if self.config["seq_parallel"]: @@ -193,10 +181,6 @@ def infer_vision(self, inputs, cache_vision=False): # Vision inference with KV cache output = self.transformer_infer.infer_vision(self.transformer_weights, infer_module_out, cache_vision=cache_vision) - if self.cpu_offload and self.offload_granularity != "model": - self.pre_weight.to_cpu() - self.transformer_weights.non_block_weights_to_cpu() - if self.config["seq_parallel"]: # Restore full chunk cos_sin so next denoising step starts from unsplit data self.scheduler.cos_sin = chunk_cos_sin diff --git a/lightx2v/models/networks/worldplay/infer/ar_transformer_infer.py b/lightx2v/models/networks/worldplay/infer/ar_transformer_infer.py index 98ba16ca7..609515508 100644 --- a/lightx2v/models/networks/worldplay/infer/ar_transformer_infer.py +++ b/lightx2v/models/networks/worldplay/infer/ar_transformer_infer.py @@ -5,7 +5,7 @@ import torch.nn.functional as F from einops import rearrange -from lightx2v.common.offload.manager import WeightAsyncStreamManager +from lightx2v.common.offload.config import get_offload_granularity from lightx2v.models.networks.hunyuan_video.infer.module_io import ( HunyuanVideo15ImgBranchOutput, HunyuanVideo15TxtBranchOutput, @@ -17,9 +17,6 @@ ) from lightx2v.models.networks.worldplay.prope.camera_rope import prope_qkv from lightx2v.utils.registry_factory import ATTN_WEIGHT_REGISTER -from lightx2v_platform.base.global_var import AI_DEVICE - -torch_device_module = getattr(torch, AI_DEVICE) class KVCache: @@ -145,15 +142,13 @@ def __init__(self, config): # Setup offload if enabled if self.config.get("cpu_offload", False): - offload_granularity = self.config.get("offload_granularity", "block") + offload_granularity = get_offload_granularity(self.config) if offload_granularity == "block": self.infer_func = self.infer_with_blocks_offload elif offload_granularity == "model": self.infer_func = self.infer_without_offload else: raise NotImplementedError - if offload_granularity != "model": - self.offload_manager = WeightAsyncStreamManager(offload_granularity=offload_granularity) @property def _vec_is_per_token(self): @@ -520,42 +515,37 @@ def infer_without_offload(self, weights, infer_module_out): @torch.no_grad() def infer_with_blocks_offload(self, weights, infer_module_out): """Inference with block-level CPU offload.""" - for block_idx in range(self.double_blocks_num): + + def infer_block(block_idx, block_weights): self.block_idx = block_idx - if block_idx == 0: - self.offload_manager.init_first_buffer(weights.double_blocks) - if block_idx < self.double_blocks_num - 1: - self.offload_manager.prefetch_weights(block_idx + 1, weights.double_blocks) - with torch_device_module.stream(self.offload_manager.compute_stream): - infer_module_out.img, infer_module_out.txt = self.infer_double_block(self.offload_manager.cuda_buffers[0], infer_module_out, block_idx=block_idx) - self.offload_manager.swap_blocks() + infer_module_out.img, infer_module_out.txt = self.infer_double_block(block_weights, infer_module_out, block_idx=block_idx) + return infer_module_out.img, infer_module_out.txt + + self.run_blocks_with_offload(weights.double_blocks, infer_block) @torch.no_grad() def infer_txt_with_offload(self, weights, infer_module_out, cache_txt=True): """Text KV caching inference using block-level CPU offload.""" - for block_idx in range(self.double_blocks_num): - if block_idx == 0: - self.offload_manager.init_first_buffer(weights.double_blocks) - if block_idx < self.double_blocks_num - 1: - self.offload_manager.prefetch_weights(block_idx + 1, weights.double_blocks) - with torch_device_module.stream(self.offload_manager.compute_stream): - block_weights = self.offload_manager.cuda_buffers[0] - txt_q, txt_k, txt_v, txt_branch_out = self._infer_txt_branch_before_attn(block_weights, infer_module_out) - txt_seqlen = txt_q.shape[1] - cu_seqlens_qkv = torch.tensor([0, txt_seqlen], dtype=torch.int32, device="cpu") - txt_attn = block_weights.self_attention.apply( - q=txt_q, - k=txt_k, - v=txt_v, - cu_seqlens_q=cu_seqlens_qkv, - cu_seqlens_kv=cu_seqlens_qkv, - max_seqlen_q=txt_seqlen, - max_seqlen_kv=txt_seqlen, - ) - if cache_txt and self._kv_cache is not None: - self._kv_cache.set_txt_cache(block_idx, txt_k.transpose(1, 2), txt_v.transpose(1, 2)) - infer_module_out.txt = self._infer_txt_branch_after_attn(block_weights, txt_attn, infer_module_out.txt, txt_branch_out) - self.offload_manager.swap_blocks() + + def infer_block(block_idx, block_weights): + txt_q, txt_k, txt_v, txt_branch_out = self._infer_txt_branch_before_attn(block_weights, infer_module_out) + txt_seqlen = txt_q.shape[1] + cu_seqlens_qkv = torch.tensor([0, txt_seqlen], dtype=torch.int32, device="cpu") + txt_attn = block_weights.self_attention.apply( + q=txt_q, + k=txt_k, + v=txt_v, + cu_seqlens_q=cu_seqlens_qkv, + cu_seqlens_kv=cu_seqlens_qkv, + max_seqlen_q=txt_seqlen, + max_seqlen_kv=txt_seqlen, + ) + if cache_txt and self._kv_cache is not None: + self._kv_cache.set_txt_cache(block_idx, txt_k.transpose(1, 2), txt_v.transpose(1, 2)) + infer_module_out.txt = self._infer_txt_branch_after_attn(block_weights, txt_attn, infer_module_out.txt, txt_branch_out) + return infer_module_out.img, infer_module_out.txt + + self.run_blocks_with_offload(weights.double_blocks, infer_block) return self._kv_cache @torch.no_grad() @@ -564,106 +554,98 @@ def infer_vision_with_offload(self, weights, infer_module_out, cache_vision=Fals use_prope = self.use_prope and hasattr(self.scheduler, "viewmats") and self.scheduler.viewmats is not None use_seq_parallel = self.seq_p_group is not None and self.config.get("seq_parallel", False) - for block_idx in range(self.double_blocks_num): - if block_idx == 0: - self.offload_manager.init_first_buffer(weights.double_blocks) - if block_idx < self.double_blocks_num - 1: - self.offload_manager.prefetch_weights(block_idx + 1, weights.double_blocks) - with torch_device_module.stream(self.offload_manager.compute_stream): - block_weights = self.offload_manager.cuda_buffers[0] - - (img_q, img_k, img_v, img_q_pre_rope, img_k_pre_rope, img_branch_out) = self._infer_img_branch_before_attn(block_weights, infer_module_out) - txt_k_cached, txt_v_cached = self._kv_cache.get_txt_cache(block_idx) - vision_k_cached, vision_v_cached = self._kv_cache.get_vision_cache(block_idx) - img_attn_prope = None - apply_fn_o = None - - if use_prope: - img_q_prope, img_k_prope, img_v_prope, apply_fn_o = self._apply_prope( - img_q_pre_rope, img_k_pre_rope, img_v, self.scheduler.viewmats, self.scheduler.Ks, infer_module_out.grid_sizes - ) - query = torch.cat([img_q, img_q_prope], dim=0) - key_current = torch.cat([img_k, img_k_prope], dim=0) - value_current = torch.cat([img_v, img_v_prope], dim=0) - if use_seq_parallel: - key_current = self._all_gather_seq(key_current, self.seq_p_group) - value_current = self._all_gather_seq(value_current, self.seq_p_group) - key_current_t = key_current.transpose(1, 2) - value_current_t = value_current.transpose(1, 2) - if cache_vision: - self._kv_cache.set_vision_cache(block_idx, key_current_t, value_current_t) - txt_k_repeated = txt_k_cached.repeat(2, 1, 1, 1) - txt_v_repeated = txt_v_cached.repeat(2, 1, 1, 1) - if vision_k_cached is not None and not cache_vision: - key_full = torch.cat([txt_k_repeated, vision_k_cached, key_current_t], dim=2) - value_full = torch.cat([txt_v_repeated, vision_v_cached, value_current_t], dim=2) - else: - key_full = torch.cat([txt_k_repeated, key_current_t], dim=2) - value_full = torch.cat([txt_v_repeated, value_current_t], dim=2) - key_full = key_full.transpose(1, 2) - value_full = value_full.transpose(1, 2) - img_seqlen = query.shape[1] - kv_seqlen = key_full.shape[1] - cu_seqlens_q = torch.tensor([0, img_seqlen, 2 * img_seqlen], dtype=torch.int32, device="cpu") - cu_seqlens_kv = torch.tensor([0, kv_seqlen, 2 * kv_seqlen], dtype=torch.int32, device="cpu") - attn_out = block_weights.self_attention.apply( - q=query, - k=key_full, - v=value_full, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_kv=cu_seqlens_kv, - max_seqlen_q=img_seqlen, - max_seqlen_kv=kv_seqlen, - ) - total_len = attn_out.shape[0] - img_attn = attn_out[: total_len // 2] - img_attn_prope = attn_out[total_len // 2 :] - if apply_fn_o is not None: - prope_proj_weight = getattr(self.action_weights, f"img_attn_prope_proj_{block_idx}", None) - if prope_proj_weight is not None: - L, C = img_attn_prope.shape - head_dim = C // self.heads_num - img_attn_prope_4d = img_attn_prope.reshape(1, L, self.heads_num, head_dim).transpose(1, 2) - img_attn_prope_transformed = apply_fn_o(img_attn_prope_4d) - img_attn_prope = img_attn_prope_transformed.transpose(1, 2).reshape(L, C) - img_attn_prope = prope_proj_weight.apply(img_attn_prope) + def infer_block(block_idx, block_weights): + (img_q, img_k, img_v, img_q_pre_rope, img_k_pre_rope, img_branch_out) = self._infer_img_branch_before_attn(block_weights, infer_module_out) + txt_k_cached, txt_v_cached = self._kv_cache.get_txt_cache(block_idx) + vision_k_cached, vision_v_cached = self._kv_cache.get_vision_cache(block_idx) + img_attn_prope = None + apply_fn_o = None + + if use_prope: + img_q_prope, img_k_prope, img_v_prope, apply_fn_o = self._apply_prope(img_q_pre_rope, img_k_pre_rope, img_v, self.scheduler.viewmats, self.scheduler.Ks, infer_module_out.grid_sizes) + query = torch.cat([img_q, img_q_prope], dim=0) + key_current = torch.cat([img_k, img_k_prope], dim=0) + value_current = torch.cat([img_v, img_v_prope], dim=0) + if use_seq_parallel: + key_current = self._all_gather_seq(key_current, self.seq_p_group) + value_current = self._all_gather_seq(value_current, self.seq_p_group) + key_current_t = key_current.transpose(1, 2) + value_current_t = value_current.transpose(1, 2) + if cache_vision: + self._kv_cache.set_vision_cache(block_idx, key_current_t, value_current_t) + txt_k_repeated = txt_k_cached.repeat(2, 1, 1, 1) + txt_v_repeated = txt_v_cached.repeat(2, 1, 1, 1) + if vision_k_cached is not None and not cache_vision: + key_full = torch.cat([txt_k_repeated, vision_k_cached, key_current_t], dim=2) + value_full = torch.cat([txt_v_repeated, vision_v_cached, value_current_t], dim=2) + else: + key_full = torch.cat([txt_k_repeated, key_current_t], dim=2) + value_full = torch.cat([txt_v_repeated, value_current_t], dim=2) + key_full = key_full.transpose(1, 2) + value_full = value_full.transpose(1, 2) + img_seqlen = query.shape[1] + kv_seqlen = key_full.shape[1] + cu_seqlens_q = torch.tensor([0, img_seqlen, 2 * img_seqlen], dtype=torch.int32, device="cpu") + cu_seqlens_kv = torch.tensor([0, kv_seqlen, 2 * kv_seqlen], dtype=torch.int32, device="cpu") + attn_out = block_weights.self_attention.apply( + q=query, + k=key_full, + v=value_full, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + max_seqlen_q=img_seqlen, + max_seqlen_kv=kv_seqlen, + ) + total_len = attn_out.shape[0] + img_attn = attn_out[: total_len // 2] + img_attn_prope = attn_out[total_len // 2 :] + if apply_fn_o is not None: + prope_proj_weight = getattr(self.action_weights, f"img_attn_prope_proj_{block_idx}", None) + if prope_proj_weight is not None: + L, C = img_attn_prope.shape + head_dim = C // self.heads_num + img_attn_prope_4d = img_attn_prope.reshape(1, L, self.heads_num, head_dim).transpose(1, 2) + img_attn_prope_transformed = apply_fn_o(img_attn_prope_4d) + img_attn_prope = img_attn_prope_transformed.transpose(1, 2).reshape(L, C) + img_attn_prope = prope_proj_weight.apply(img_attn_prope) + else: + if use_seq_parallel: + img_k = self._all_gather_seq(img_k, self.seq_p_group) + img_v = self._all_gather_seq(img_v, self.seq_p_group) + img_k_t = img_k.transpose(1, 2) + img_v_t = img_v.transpose(1, 2) + if cache_vision: + self._kv_cache.set_vision_cache(block_idx, img_k_t, img_v_t) + if vision_k_cached is not None and not cache_vision: + key = torch.cat([txt_k_cached, vision_k_cached, img_k_t], dim=2) + value = torch.cat([txt_v_cached, vision_v_cached, img_v_t], dim=2) else: - if use_seq_parallel: - img_k = self._all_gather_seq(img_k, self.seq_p_group) - img_v = self._all_gather_seq(img_v, self.seq_p_group) - img_k_t = img_k.transpose(1, 2) - img_v_t = img_v.transpose(1, 2) - if cache_vision: - self._kv_cache.set_vision_cache(block_idx, img_k_t, img_v_t) - if vision_k_cached is not None and not cache_vision: - key = torch.cat([txt_k_cached, vision_k_cached, img_k_t], dim=2) - value = torch.cat([txt_v_cached, vision_v_cached, img_v_t], dim=2) - else: - key = torch.cat([txt_k_cached, img_k_t], dim=2) - value = torch.cat([txt_v_cached, img_v_t], dim=2) - key = key.transpose(1, 2) - value = value.transpose(1, 2) - img_seqlen = img_q.shape[1] - kv_seqlen = key.shape[1] - cu_seqlens_q = torch.tensor([0, img_seqlen], dtype=torch.int32, device="cpu") - cu_seqlens_kv = torch.tensor([0, kv_seqlen], dtype=torch.int32, device="cpu") - img_attn = block_weights.self_attention.apply( - q=img_q, - k=key, - v=value, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_kv=cu_seqlens_kv, - max_seqlen_q=img_seqlen, - max_seqlen_kv=kv_seqlen, - ) - - infer_module_out.img = self._infer_img_branch_after_attn(block_weights, img_attn, infer_module_out.img, img_branch_out, img_attn_prope) - self.offload_manager.swap_blocks() + key = torch.cat([txt_k_cached, img_k_t], dim=2) + value = torch.cat([txt_v_cached, img_v_t], dim=2) + key = key.transpose(1, 2) + value = value.transpose(1, 2) + img_seqlen = img_q.shape[1] + kv_seqlen = key.shape[1] + cu_seqlens_q = torch.tensor([0, img_seqlen], dtype=torch.int32, device="cpu") + cu_seqlens_kv = torch.tensor([0, kv_seqlen], dtype=torch.int32, device="cpu") + img_attn = block_weights.self_attention.apply( + q=img_q, + k=key, + v=value, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_kv=cu_seqlens_kv, + max_seqlen_q=img_seqlen, + max_seqlen_kv=kv_seqlen, + ) + + infer_module_out.img = self._infer_img_branch_after_attn(block_weights, img_attn, infer_module_out.img, img_branch_out, img_attn_prope) + return infer_module_out.img, infer_module_out.txt + + self.run_blocks_with_offload(weights.double_blocks, infer_block) if cache_vision: return self._kv_cache - else: - return self.infer_final_layer(weights, infer_module_out) + return self.infer_final_layer(weights, infer_module_out) def set_action_weights(self, action_weights): """Set action weights for ProPE projection access.""" @@ -692,7 +674,7 @@ def infer_txt(self, weights, infer_module_out, cache_txt=True): Returns: KV cache reference """ - if hasattr(self, "offload_manager"): + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "block": return self.infer_txt_with_offload(weights, infer_module_out, cache_txt=cache_txt) for block_idx in range(self.double_blocks_num): @@ -745,7 +727,7 @@ def infer_vision(self, weights, infer_module_out, cache_vision=False): If cache_vision=True: KV cache reference If cache_vision=False: Final layer output (noise prediction) """ - if hasattr(self, "offload_manager"): + if self.config.get("cpu_offload", False) and get_offload_granularity(self.config) == "block": return self.infer_vision_with_offload(weights, infer_module_out, cache_vision=cache_vision) # Check if ProPE is enabled diff --git a/lightx2v/models/runners/worldplay/worldplay_ar_runner.py b/lightx2v/models/runners/worldplay/worldplay_ar_runner.py index f0082e532..910e2386a 100644 --- a/lightx2v/models/runners/worldplay/worldplay_ar_runner.py +++ b/lightx2v/models/runners/worldplay/worldplay_ar_runner.py @@ -199,7 +199,8 @@ def init_run(self): def run_main(self): """Override to use chunk-based AR generation instead of run_segment().""" self.init_run() - self.run_denoising_loop() + with self.transformer_offload_session(): + self.run_denoising_loop() latents = self.scheduler.latents if self.config.get("use_stream_vae", False): frames = [] From a09aec21a080b0fb9258ee776e9f53ac3f1eb7b9 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 20:52:27 +0800 Subject: [PATCH 29/31] configs(offload): add MetaX Wan resident example --- .../metax/wan_t2v_1_3b_resident.json | 27 +++++++++++++++++++ 1 file changed, 27 insertions(+) create mode 100644 configs/platforms/metax/wan_t2v_1_3b_resident.json diff --git a/configs/platforms/metax/wan_t2v_1_3b_resident.json b/configs/platforms/metax/wan_t2v_1_3b_resident.json new file mode 100644 index 000000000..5d9aaf3f7 --- /dev/null +++ b/configs/platforms/metax/wan_t2v_1_3b_resident.json @@ -0,0 +1,27 @@ +{ + "infer_steps": 4, + "target_video_length": 81, + "text_len": 512, + "target_height": 480, + "target_width": 832, + "self_attn_1_type": "metax_sage_attn2", + "cross_attn_1_type": "metax_sage_attn2", + "cross_attn_2_type": "metax_sage_attn2", + "sample_guide_scale": 6.0, + "sample_shift": 8, + "enable_cfg": true, + "cpu_offload": true, + "offload_plan": { + "offload_granularity": "block", + "resident_blocks": { + "blocks": 28 + }, + "use_event_offload": true + }, + "t5_cpu_offload": false, + "vae_cpu_offload": false, + "modulate_type": "torch", + "rope_type": "torch_complex_rope", + "layer_norm_type": "torch", + "rms_norm_type": "torch" +} From 4cec0fd7240520c57f885062f7a285ab07d0320e Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 21:09:02 +0800 Subject: [PATCH 30/31] fix(offload): expose complete offload plans in APIs --- lightx2v/disagg/utils.py | 11 ++++++----- lightx2v/pipeline.py | 10 +++++----- tests/common/offload/test_config.py | 23 +++++++++++++++++++++++ 3 files changed, 34 insertions(+), 10 deletions(-) diff --git a/lightx2v/disagg/utils.py b/lightx2v/disagg/utils.py index 1a6929318..437d7da91 100644 --- a/lightx2v/disagg/utils.py +++ b/lightx2v/disagg/utils.py @@ -59,6 +59,7 @@ def set_config( vae_offload=False, resident_blocks=None, use_event_offload=False, + offload_plan=None, **kwargs, ): """ @@ -128,11 +129,11 @@ def set_config( args_dict.update(kwargs) if "offload_plan" not in args_dict: - args_dict["offload_plan"] = { - "offload_granularity": args_dict.pop("offload_granularity", "block"), - "resident_blocks": {} if resident_blocks is None else resident_blocks, - "use_event_offload": args_dict.pop("use_event_offload", use_event_offload), - } + plan = {} if offload_plan is None else dict(offload_plan) + plan.setdefault("offload_granularity", args_dict.pop("offload_granularity", "block")) + plan.setdefault("resident_blocks", {} if resident_blocks is None else resident_blocks) + plan.setdefault("use_event_offload", args_dict.pop("use_event_offload", use_event_offload)) + args_dict["offload_plan"] = plan # Convert to object for set_config compatibility args = ConfigObj(**args_dict) diff --git a/lightx2v/pipeline.py b/lightx2v/pipeline.py index c68d9b674..ecbceb8eb 100755 --- a/lightx2v/pipeline.py +++ b/lightx2v/pipeline.py @@ -358,13 +358,13 @@ def enable_offload( vae_offload=False, resident_blocks=None, use_event_offload=False, + offload_plan=None, ): self.cpu_offload = cpu_offload - self.offload_plan = { - "offload_granularity": offload_granularity, - "resident_blocks": {} if resident_blocks is None else resident_blocks, - "use_event_offload": use_event_offload, - } + self.offload_plan = {} if offload_plan is None else dict(offload_plan) + self.offload_plan.setdefault("offload_granularity", offload_granularity) + self.offload_plan.setdefault("resident_blocks", {} if resident_blocks is None else resident_blocks) + self.offload_plan.setdefault("use_event_offload", use_event_offload) self.vae_cpu_offload = vae_offload if self.model_cls in [ "wan2.1", diff --git a/tests/common/offload/test_config.py b/tests/common/offload/test_config.py index 3abe7c529..b97f081c1 100644 --- a/tests/common/offload/test_config.py +++ b/tests/common/offload/test_config.py @@ -1,4 +1,5 @@ from lightx2v.common.offload.config import get_offload_granularity, normalize_offload_plan, use_event_offload +from lightx2v.pipeline import LightX2VPipeline def test_legacy_offload_keys_are_normalized_into_one_plan(): @@ -45,3 +46,25 @@ def test_offload_plan_defaults_are_added_without_changing_model_specific_keys(): "use_event_offload": False, "use_block_slab": True, } + + +def test_pipeline_accepts_model_specific_offload_plan(): + pipeline = LightX2VPipeline.__new__(LightX2VPipeline) + pipeline.model_cls = "flux2_dev" + + pipeline.enable_offload( + cpu_offload=True, + offload_plan={ + "offload_granularity": "block", + "resident_blocks": {"double_blocks": 4, "single_blocks": 8}, + "use_event_offload": True, + "use_block_slab": True, + }, + ) + + assert pipeline.offload_plan == { + "offload_granularity": "block", + "resident_blocks": {"double_blocks": 4, "single_blocks": 8}, + "use_event_offload": True, + "use_block_slab": True, + } From 2f36e92070b9fec297397f2f91c29a662dc24a57 Mon Sep 17 00:00:00 2001 From: Super User Date: Tue, 8 Sep 2026 21:16:45 +0800 Subject: [PATCH 31/31] fix(offload): keep compiled staging slots cached --- .../transformer_infer/transformer_infer.py | 2 +- tests/common/offload/test_offload_schedule.py | 20 +++++++++++++++++++ 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/lightx2v/common/transformer_infer/transformer_infer.py b/lightx2v/common/transformer_infer/transformer_infer.py index 76a9b11fa..76d99b775 100644 --- a/lightx2v/common/transformer_infer/transformer_infer.py +++ b/lightx2v/common/transformer_infer/transformer_infer.py @@ -61,7 +61,7 @@ def init_compile(self, config): logger.info(f"[Compile] Using torch.compile for {type(self).__name__}") def get_compiled_block(self, block_idx, block): - key = self.get_compile_block_key(block_idx, block) + key = self.get_compile_block_key(block_idx, block), id(block) cached = self.compiled_blocks.get(key) if cached is not None and cached[0] is block: return cached[1] diff --git a/tests/common/offload/test_offload_schedule.py b/tests/common/offload/test_offload_schedule.py index 1dcacee19..08fbaa1f8 100644 --- a/tests/common/offload/test_offload_schedule.py +++ b/tests/common/offload/test_offload_schedule.py @@ -203,6 +203,26 @@ def _block_offload_config(use_events=False): } +def test_compiled_blocks_are_cached_for_each_staging_slot(monkeypatch): + compiled = [] + + def compile_block(block_runner, dynamic): + compiled.append((block_runner, dynamic)) + return block_runner + + monkeypatch.setattr(transformer_infer_module.torch, "compile", compile_block) + infer = _Infer() + infer.compiled_blocks = {} + first_slot = object() + second_slot = object() + + first_compiled = infer.get_compiled_block(0, first_slot) + infer.get_compiled_block(0, second_slot) + + assert infer.get_compiled_block(0, first_slot) is first_compiled + assert len(compiled) == 2 + + def test_single_block_group_is_bound_to_its_buffers(): blocks = _make_blocks() buffers = [_Buffer(0), _Buffer(1)]