From 4c44cf6e4543995568104ad72a1e43ce1f6dc62c Mon Sep 17 00:00:00 2001 From: wushuo <540295877@qq.com> Date: Wed, 22 Jul 2026 10:56:44 +0000 Subject: [PATCH 1/9] perf(z-image): scope RoPE cache to each request --- .../networks/z_image/infer/pre_infer.py | 32 ++-- .../models/schedulers/z_image/scheduler.py | 144 +----------------- 2 files changed, 26 insertions(+), 150 deletions(-) diff --git a/lightx2v/models/networks/z_image/infer/pre_infer.py b/lightx2v/models/networks/z_image/infer/pre_infer.py index 3b6d8a316..52e817824 100755 --- a/lightx2v/models/networks/z_image/infer/pre_infer.py +++ b/lightx2v/models/networks/z_image/infer/pre_infer.py @@ -17,15 +17,20 @@ def __init__(self, config): self.cpu_offload = config.get("cpu_offload", False) self.zero_cond_t = config.get("zero_cond_t", False) self.rope = None - self._rope_cache = {} + self.scheduler = None + self.clear_rope_cache() + + def clear_rope_cache(self): + self._cached_request_id = None + self._rope_cache = {True: None, False: None} def set_rope(self, rope): self.rope = rope - self._rope_cache.clear() + self.clear_rope_cache() def set_scheduler(self, scheduler): self.scheduler = scheduler - self._rope_cache.clear() + self.clear_rope_cache() @staticmethod def _device_key(device): @@ -41,6 +46,14 @@ def _prepare_rope_cache( x_padded_len, cap_padded_len, ): + if self.scheduler is None: + raise RuntimeError("ZImagePreInfer scheduler is not initialized.") + + request_id = self.scheduler.rope_request_id + if request_id != self._cached_request_id: + self._cached_request_id = request_id + self._rope_cache = {True: None, False: None} + if self.rope is None: raise RuntimeError("ZImagePreInfer RoPE is not initialized.") @@ -62,9 +75,10 @@ def _prepare_rope_cache( world_size, rank, ) - cached = self._rope_cache.get(cache_key) - if cached is not None: - return cached + branch = bool(self.scheduler.infer_condition) + cached = self._rope_cache[branch] + if cached is not None and cached[0] == cache_key: + return cached[1] cap_pos_ids = self.scheduler.create_coordinate_grid( size=(cap_padded_len, 1, 1), @@ -97,7 +111,7 @@ def _prepare_rope_cache( x_freqs_cis = self.rope.prepare_freqs(x_freqs_cis, rotary_dim=rotary_dim) cap_freqs_cis = self.rope.prepare_freqs(cap_freqs_cis, rotary_dim=rotary_dim) unified_freqs_cis = torch.cat([x_freqs_cis, cap_freqs_cis], dim=0) - cached = ( + value = ( x_freqs_cis, cap_freqs_cis, unified_freqs_cis, @@ -105,8 +119,8 @@ def _prepare_rope_cache( self.rope.prepare_positions(cap_freqs_cis), self.rope.prepare_positions(unified_freqs_cis), ) - self._rope_cache[cache_key] = cached - return cached + self._rope_cache[branch] = (cache_key, value) + return value def infer(self, weights, hidden_states, encoder_hidden_states): patch_size = self.config.get("patch_size", 2) diff --git a/lightx2v/models/schedulers/z_image/scheduler.py b/lightx2v/models/schedulers/z_image/scheduler.py index be1d4ae20..4603d1285 100755 --- a/lightx2v/models/schedulers/z_image/scheduler.py +++ b/lightx2v/models/schedulers/z_image/scheduler.py @@ -1,4 +1,3 @@ -import functools import inspect import json import math @@ -11,7 +10,6 @@ import torch.distributed as dist from diffusers.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler from loguru import logger -from torch import nn from torch.nn import functional as F from lightx2v.models.schedulers.scheduler import BaseScheduler @@ -216,104 +214,6 @@ def get_timestep_embedding( return emb -class ZEmbedRope(nn.Module): - def __init__(self, theta: int, axes_dim: List[int], scale_rope=False): - super().__init__() - self.theta = theta - self.axes_dim = axes_dim - pos_index = torch.arange(4096) - neg_index = torch.arange(4096).flip(0) * -1 - 1 - self.pos_freqs = torch.cat( - [ - self.rope_params(pos_index, self.axes_dim[0], self.theta), - self.rope_params(pos_index, self.axes_dim[1], self.theta), - self.rope_params(pos_index, self.axes_dim[2], self.theta), - ], - dim=1, - ) - self.neg_freqs = torch.cat( - [ - self.rope_params(neg_index, self.axes_dim[0], self.theta), - self.rope_params(neg_index, self.axes_dim[1], self.theta), - self.rope_params(neg_index, self.axes_dim[2], self.theta), - ], - dim=1, - ) - self.rope_cache = {} - - # DO NOT USING REGISTER BUFFER HERE, IT WILL CAUSE COMPLEX NUMBERS LOSE ITS IMAGINARY PART - self.scale_rope = scale_rope - - def rope_params(self, index, dim, theta=10000): - """ - Args: - index: [0, 1, 2, 3] 1D Tensor representing the position index of the token - """ - assert dim % 2 == 0 - freqs = torch.outer(index, 1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float32).div(dim))) - freqs = torch.polar(torch.ones_like(freqs), freqs) - return freqs - - def forward(self, video_fhw, txt_seq_lens, device): - """ - Args: video_fhw: [frame, height, width] a list of 3 integers representing the shape of the video Args: - txt_length: [bs] a list of 1 integers representing the length of the text - """ - if self.pos_freqs.device != device: - self.pos_freqs = self.pos_freqs.to(device) - self.neg_freqs = self.neg_freqs.to(device) - - if isinstance(video_fhw, list): - video_fhw = video_fhw[0] - if not isinstance(video_fhw, list): - video_fhw = [video_fhw] - - vid_freqs = [] - max_vid_index = 0 - for idx, fhw in enumerate(video_fhw): - frame, height, width = fhw - rope_key = f"{idx}_{height}_{width}" - - if not torch.compiler.is_compiling(): - if rope_key not in self.rope_cache: - self.rope_cache[rope_key] = self._compute_video_freqs(frame, height, width, idx) - video_freq = self.rope_cache[rope_key] - else: - video_freq = self._compute_video_freqs(frame, height, width, idx) - video_freq = video_freq.to(device) - vid_freqs.append(video_freq) - - if self.scale_rope: - max_vid_index = max(height // 2, width // 2, max_vid_index) - else: - max_vid_index = max(height, width, max_vid_index) - - max_len = txt_seq_lens - txt_freqs = self.pos_freqs[max_vid_index : max_vid_index + max_len, ...] - vid_freqs = torch.cat(vid_freqs, dim=0) - - return [vid_freqs, txt_freqs] - - @functools.lru_cache(maxsize=None) - def _compute_video_freqs(self, frame, height, width, idx=0): - seq_lens = frame * height * width - freqs_pos = self.pos_freqs.split([x // 2 for x in self.axes_dim], dim=1) - freqs_neg = self.neg_freqs.split([x // 2 for x in self.axes_dim], dim=1) - - freqs_frame = freqs_pos[0][idx : idx + frame].view(frame, 1, 1, -1).expand(frame, height, width, -1) - if self.scale_rope: - freqs_height = torch.cat([freqs_neg[1][-(height - height // 2) :], freqs_pos[1][: height // 2]], dim=0) - freqs_height = freqs_height.view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = torch.cat([freqs_neg[2][-(width - width // 2) :], freqs_pos[2][: width // 2]], dim=0) - freqs_width = freqs_width.view(1, 1, width, -1).expand(frame, height, width, -1) - else: - freqs_height = freqs_pos[1][:height].view(1, height, 1, -1).expand(frame, height, width, -1) - freqs_width = freqs_pos[2][:width].view(1, 1, width, -1).expand(frame, height, width, -1) - - freqs = torch.cat([freqs_frame, freqs_height, freqs_width], dim=-1).reshape(seq_lens, -1) - return freqs.clone().contiguous() - - class RopeEmbedder: def __init__( self, @@ -374,8 +274,6 @@ def __init__(self, config): self.seq_p_group = self.config.get("device_mesh").get_group(mesh_dim="seq_p") else: self.seq_p_group = None - self.pos_embed = ZEmbedRope(theta=10000, axes_dim=[16, 56, 56], scale_rope=True) - # Initialize RopeEmbedder for generating freqs_cis from position IDs (used in pre_infer) rope_theta = config.get("rope_theta", 256.0) axes_dims = config.get("axes_dims", [32, 48, 48]) @@ -386,6 +284,7 @@ def __init__(self, config): axes_lens=axes_lens, ) self.freqs_cis_cache = {} + self.rope_request_id = 0 @staticmethod def _pack_latents(latents, batch_size, num_channels_latents, height, width): @@ -410,18 +309,6 @@ def _unpack_latents(latents, height, width, vae_scale_factor): return latents - @staticmethod - def _prepare_latent_image_ids(batch_size, height, width, device, dtype): - latent_image_ids = torch.zeros(height, width, 3) - - latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None] - latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :] - - latent_image_id_height, latent_image_id_width, latent_image_id_channels = latent_image_ids.shape - latent_image_ids = latent_image_ids.reshape(latent_image_id_height * latent_image_id_width, latent_image_id_channels) - - return latent_image_ids.to(device=device, dtype=dtype) - @staticmethod def create_coordinate_grid(size, start=None, device=None): """Create a 3D coordinate grid.""" @@ -492,14 +379,9 @@ def prepare_latents(self, input_info): if len(shape) != 4: raise ValueError(f"target_shape must be 4D [B, C, H, W], got {len(shape)}D: {shape}") - batch_size, num_channels, height, width = shape - latents = randn_tensor(shape, generator=self.generator, device=AI_DEVICE, dtype=self.dtype) - latent_image_ids = self._prepare_latent_image_ids(1, height // 2, width // 2, AI_DEVICE, self.dtype) - self.latents = latents - self.latent_image_ids = latent_image_ids self.noise_pred = None def generate_freqs_cis_from_position_ids(self, position_ids: torch.Tensor, device: torch.device = None) -> torch.Tensor: @@ -558,7 +440,9 @@ def set_timesteps(self): self.num_warmup_steps = num_warmup_steps def prepare(self, input_info): + self.rope_request_id += 1 self.freqs_cis_cache = {} + if self.generator is None: self.generator = torch.Generator(device=AI_DEVICE).manual_seed(input_info.seed) else: @@ -569,28 +453,6 @@ def prepare(self, input_info): if self.config["task"] == "i2i" and strength is not None: self.prepare_i2i_denoise_strength_latents(input_info) - self.image_rotary_emb = self.pos_embed(self.input_info.image_shapes, input_info.txt_seq_lens[0], device=AI_DEVICE) - - if self.seq_p_group is not None: - world_size = dist.get_world_size(self.seq_p_group) - cur_rank = dist.get_rank(self.seq_p_group) - seqlen = self.image_rotary_emb[0].shape[0] - padding_size = (world_size - (seqlen % world_size)) % world_size - if padding_size > 0: - self.image_rotary_emb[0] = F.pad(self.image_rotary_emb[0], (0, 0, 0, padding_size)) - self.image_rotary_emb[0] = torch.chunk(self.image_rotary_emb[0], world_size, dim=0)[cur_rank] - - if self.config["enable_cfg"]: - self.negative_image_rotary_emb = self.pos_embed(self.input_info.image_shapes, input_info.txt_seq_lens[1], device=AI_DEVICE) - if self.seq_p_group is not None: - world_size = dist.get_world_size(self.seq_p_group) - cur_rank = dist.get_rank(self.seq_p_group) - seqlen = self.negative_image_rotary_emb[0].shape[0] - padding_size = (world_size - (seqlen % world_size)) % world_size - if padding_size > 0: - self.negative_image_rotary_emb[0] = F.pad(self.negative_image_rotary_emb[0], (0, 0, 0, padding_size)) - self.negative_image_rotary_emb[0] = torch.chunk(self.negative_image_rotary_emb[0], world_size, dim=0)[cur_rank] - if self.zero_cond_t: self.modulate_index = torch.tensor([[0] * prod(sample[0]) + [1] * sum([prod(s) for s in sample[1:]]) for sample in self.input_info.image_shapes], device=AI_DEVICE, dtype=torch.int) if self.seq_p_group is not None: From 9c0072d025ab5aa5392696156187efccccafef1b Mon Sep 17 00:00:00 2001 From: wushuo <540295877@qq.com> Date: Wed, 22 Jul 2026 11:11:32 +0000 Subject: [PATCH 2/9] perf(ernie-image): scope RoPE cache to each request --- .../networks/ernie_image/infer/pre_infer.py | 36 +++++++++++++++---- .../schedulers/ernie_image/scheduler.py | 2 ++ 2 files changed, 31 insertions(+), 7 deletions(-) diff --git a/lightx2v/models/networks/ernie_image/infer/pre_infer.py b/lightx2v/models/networks/ernie_image/infer/pre_infer.py index b802d0d46..7f3145340 100644 --- a/lightx2v/models/networks/ernie_image/infer/pre_infer.py +++ b/lightx2v/models/networks/ernie_image/infer/pre_infer.py @@ -50,10 +50,12 @@ def __init__(self, config): self.rope_axes_dim = config.get("rope_axes_dim", (32, 48, 48)) self.rotary_dim = sum(self.rope_axes_dim) self.rope = None + self.scheduler = None self.clear_rope_cache() def clear_rope_cache(self): - self._rope_cache = {} + self._rope_cache_request_id = None + self._rope_cache = {True: None, False: None} def set_rope(self, rope): self.rope = rope @@ -63,6 +65,10 @@ def set_scheduler(self, scheduler): self.scheduler = scheduler self.clear_rope_cache() + @staticmethod + def _device_key(device): + return device.type, device.index + def prepare_rope_cache(self, image_hw, text_len, device): if self.rope is None: raise RuntimeError("ErnieImagePreInfer RoPE is not initialized.") @@ -74,12 +80,28 @@ def prepare_rope_cache(self, image_hw, text_len, device): return rotary_freqs, rotary_positions def get_rope_cache(self, image_hw, text_len, device): - cache_key = (*image_hw, text_len, str(device)) - cached = self._rope_cache.get(cache_key) - if cached is None: - cached = self.prepare_rope_cache(image_hw, text_len, device) - self._rope_cache[cache_key] = cached - return cached + if self.scheduler is None: + raise RuntimeError("ErnieImagePreInfer scheduler is not initialized.") + + request_id = self.scheduler.rope_request_id + if request_id != self._rope_cache_request_id: + self._rope_cache_request_id = request_id + self._rope_cache = {True: None, False: None} + + cache_key = (self._device_key(device), *image_hw, text_len) + branch = bool(self.scheduler.infer_condition) + cached = self._rope_cache[branch] + if cached is not None and cached[0] == cache_key: + return cached[1] + + shared = self._rope_cache[not branch] + if shared is not None and shared[0] == cache_key: + self._rope_cache[branch] = shared + return shared[1] + + value = self.prepare_rope_cache(image_hw, text_len, device) + self._rope_cache[branch] = (cache_key, value) + return value def _pos_embed(self, ids: torch.Tensor) -> torch.Tensor: emb = torch.cat([_rope(ids[..., i], self.rope_axes_dim[i], self.rope_theta) for i in range(3)], dim=-1) diff --git a/lightx2v/models/schedulers/ernie_image/scheduler.py b/lightx2v/models/schedulers/ernie_image/scheduler.py index 2608d19b1..952780f9c 100644 --- a/lightx2v/models/schedulers/ernie_image/scheduler.py +++ b/lightx2v/models/schedulers/ernie_image/scheduler.py @@ -18,6 +18,7 @@ def __init__(self, config): self.sample_guide_scale = config.get("sample_guide_scale", 1.0) self.timestep = None self.noise_pred = None + self.rope_request_id = 0 def prepare_latents(self, input_info): shape = tuple(input_info.target_shape) @@ -41,6 +42,7 @@ def set_timesteps(self): self.infer_steps = len(self.timesteps) def prepare(self, input_info): + self.rope_request_id += 1 if self.generator is None: self.generator = torch.Generator(device=AI_DEVICE).manual_seed(input_info.seed) self.prepare_latents(input_info) From ddad62be25fd30aeb78ec255b210d786494238d0 Mon Sep 17 00:00:00 2001 From: wushuo <540295877@qq.com> Date: Wed, 22 Jul 2026 11:33:30 +0000 Subject: [PATCH 3/9] perf(motus): scope inference caches to each request --- .../models/networks/motus/infer/pre_infer.py | 38 +++++++++++++---- .../networks/motus/infer/transformer_infer.py | 42 +++++++++++++------ lightx2v/models/networks/motus/model.py | 26 ------------ lightx2v/models/schedulers/motus/scheduler.py | 2 + 4 files changed, 62 insertions(+), 46 deletions(-) diff --git a/lightx2v/models/networks/motus/infer/pre_infer.py b/lightx2v/models/networks/motus/infer/pre_infer.py index 0a264afc3..de30d1a1a 100644 --- a/lightx2v/models/networks/motus/infer/pre_infer.py +++ b/lightx2v/models/networks/motus/infer/pre_infer.py @@ -11,16 +11,45 @@ def __init__(self, model, config): super().__init__(config) self.model = model self.scheduler = None - self._grid_cache = {} + self.clear_request_cache() + + def clear_request_cache(self): + self._rope_cache_request_id = None + self._grid_cache = None + self.cos_sin = None + self.rope_positions = None + self.grid_sizes = (0, 0, 0) + + def set_rope(self, rope): + super().set_rope(rope) + self.clear_request_cache() def set_scheduler(self, scheduler): self.scheduler = scheduler + self.clear_request_cache() + + def _begin_request(self): + request_id = self.scheduler.rope_request_id + if request_id != self._rope_cache_request_id: + self.clear_request_cache() + self._rope_cache_request_id = request_id + + def _get_grid_output(self, batch_size, grid_tuple, device): + grid_key = (batch_size, grid_tuple, device.type, device.index) + if self._grid_cache is not None and self._grid_cache[0] == grid_key: + return self._grid_cache[1] + + grid_sizes = torch.tensor([grid_tuple], dtype=torch.long, device=device).expand(batch_size, -1) + grid_output = GridOutput(tensor=grid_sizes, tuple=grid_tuple) + self._grid_cache = (grid_key, grid_output) + return grid_output @torch.no_grad() def infer(self, weights, inputs, kv_start=0, kv_end=0): del weights, kv_start, kv_end if self.scheduler is None: raise RuntimeError("MotusPreInfer requires a scheduler before infer().") + self._begin_request() first_frame = inputs["motus_first_frame"] state = inputs["motus_state"] @@ -41,12 +70,7 @@ def infer(self, weights, inputs, kv_start=0, kv_end=0): latent_h // self.model.video_backbone.patch_size[1], latent_w // self.model.video_backbone.patch_size[2], ) - grid_key = (batch_size, grid_tuple, state.device.type, state.device.index) - grid_output = self._grid_cache.get(grid_key) - if grid_output is None: - grid_sizes = torch.tensor([grid_tuple], dtype=torch.long, device=state.device).expand(batch_size, -1) - grid_output = GridOutput(tensor=grid_sizes, tuple=grid_tuple) - self._grid_cache[grid_key] = grid_output + grid_output = self._get_grid_output(batch_size, grid_tuple, state.device) if self.cos_sin is None or self.grid_sizes != grid_output.tuple: self.grid_sizes = grid_output.tuple diff --git a/lightx2v/models/networks/motus/infer/transformer_infer.py b/lightx2v/models/networks/motus/infer/transformer_infer.py index b09f01c1c..dee33aa49 100644 --- a/lightx2v/models/networks/motus/infer/transformer_infer.py +++ b/lightx2v/models/networks/motus/infer/transformer_infer.py @@ -47,14 +47,25 @@ def __init__(self, model, config): self.cos_sin = None self.rope_positions = None self.weights = None - self._cu_seqlens_cache = {} - self.reset_infer_states() + self.clear_request_cache() + + def clear_request_cache(self): + self._cu_seqlens_cache_request_id = None + self._cu_seqlens_cache = { + "joint_attn_cu_seqlens_q": None, + "cross_attn_cu_seqlens_q": None, + "cross_attn_cu_seqlens_kv": None, + } + + def set_scheduler(self, scheduler): + super().set_scheduler(scheduler) + self.clear_request_cache() - def reset_infer_states(self): - self.joint_attn_cu_seqlens_q = None - self.joint_attn_cu_seqlens_kv = None - self.cross_attn_cu_seqlens_q = None - self.cross_attn_cu_seqlens_kv = None + def _begin_request(self): + request_id = self.scheduler.rope_request_id + if request_id != self._cu_seqlens_cache_request_id: + self.clear_request_cache() + self._cu_seqlens_cache_request_id = request_id def _maybe_empty_cache(self): if self.clean_cuda_cache: @@ -70,12 +81,17 @@ def _get_und_block(self, layer_idx): return self.weights.und.blocks[layer_idx] def _get_cu_seqlens(self, cache_name, batch, seq_len, device, attn_type): + if cache_name not in self._cu_seqlens_cache: + raise ValueError(f"Unsupported Motus attention metadata cache: {cache_name}") + cache_key = (cache_name, batch, seq_len, device.type, device.index, attn_type) - cu_seqlens = self._cu_seqlens_cache.get(cache_key) - if cu_seqlens is None: - tensor = torch.arange(0, (batch + 1) * seq_len, seq_len, dtype=torch.int32) - cu_seqlens = tensor.to(device, non_blocking=True) if attn_type in ["flash_attn2", "flash_attn3"] else tensor - self._cu_seqlens_cache[cache_key] = cu_seqlens + cached = self._cu_seqlens_cache[cache_name] + if cached is not None and cached[0] == cache_key: + return cached[1] + + tensor = torch.arange(0, (batch + 1) * seq_len, seq_len, dtype=torch.int32) + cu_seqlens = tensor.to(device, non_blocking=True) if attn_type in ["flash_attn2", "flash_attn3"] else tensor + self._cu_seqlens_cache[cache_name] = (cache_key, cu_seqlens) return cu_seqlens def _normalize_attention_dtype(self, tensor): @@ -332,7 +348,7 @@ def infer(self, weights, pre_infer_out): self.weights = weights self.cos_sin = pre_infer_out.cos_sin self.rope_positions = pre_infer_out.rope_positions - self.reset_infer_states() + self._begin_request() processed_t5_context = pre_infer_out.context und_tokens = pre_infer_out.und_tokens.clone() diff --git a/lightx2v/models/networks/motus/model.py b/lightx2v/models/networks/motus/model.py index dae7d045b..a48b77628 100644 --- a/lightx2v/models/networks/motus/model.py +++ b/lightx2v/models/networks/motus/model.py @@ -233,7 +233,6 @@ def __init__(self, config, device): logger.info("[Motus] Loading VLM processor") self.vlm_processor = AutoProcessor.from_pretrained(self.config["vlm_path"], trust_remote_code=True) self._load_normalization_stats() - self._rope_cos_sin_cache = {} self._patch_qwen3_vl_rope_index(self.vlm_model) logger.info("[Motus] Building Motus backbone helpers") self.video_backbone = MotusVideoBackbone(self.config, self.pre_weight.video, self.transformer_weights.video) @@ -502,31 +501,6 @@ def denormalize_actions(self, actions): restored = flat * self.action_range.unsqueeze(0) + self.action_min.unsqueeze(0) return restored.reshape(shape) - def get_wan_freqs(self): - return self.pre_infer.freqs - - def get_wan_rotary_cos_sin(self, grid_size): - if grid_size in self._rope_cos_sin_cache: - return self._rope_cos_sin_cache[grid_size] - freqs = self.get_wan_freqs() - head_dim_half = freqs.shape[1] - c_f = head_dim_half - 2 * (head_dim_half // 3) - c_h = head_dim_half // 3 - c_w = head_dim_half // 3 - fpart, hpart, wpart = freqs.split([c_f, c_h, c_w], dim=1) - f, h, w = grid_size - freq_grid = torch.cat( - [ - fpart[:f].view(f, 1, 1, -1).expand(f, h, w, -1), - hpart[:h].view(1, h, 1, -1).expand(f, h, w, -1), - wpart[:w].view(1, 1, w, -1).expand(f, h, w, -1), - ], - dim=-1, - ).reshape(f * h * w, -1) - cos_sin = (freq_grid.real.contiguous(), freq_grid.imag.contiguous()) - self._rope_cos_sin_cache[grid_size] = cos_sin - return cos_sin - def prepare_frame(self, image_path): image = Image.open(image_path).convert("RGB") image_np = np.asarray(image).astype(np.float32) / 255.0 diff --git a/lightx2v/models/schedulers/motus/scheduler.py b/lightx2v/models/schedulers/motus/scheduler.py index 69e58065f..b2e6d82b1 100644 --- a/lightx2v/models/schedulers/motus/scheduler.py +++ b/lightx2v/models/schedulers/motus/scheduler.py @@ -12,8 +12,10 @@ def __init__(self, config): self.action_latents = None self.action_noise_pred = None self.condition_frame_latent = None + self.rope_request_id = 0 def prepare(self, seed, latent_shape, image_encoder_output, action_shape): + self.rope_request_id += 1 self.vae_encoder_out = image_encoder_output["vae_encoder_out"] self.prepare_latents(seed, latent_shape, dtype=torch.float32) From dafad380cec73c5a1c2ff27c129edd2194c6f621 Mon Sep 17 00:00:00 2001 From: wushuo <540295877@qq.com> Date: Wed, 22 Jul 2026 11:33:57 +0000 Subject: [PATCH 4/9] perf(wan-causvid): scope inference caches to each request --- .../networks/wan/infer/causvid/pre_infer.py | 24 +++++++++++++++ .../wan/infer/causvid/transformer_infer.py | 29 +++++++++++++++++++ lightx2v/models/schedulers/wan/scheduler.py | 2 ++ 3 files changed, 55 insertions(+) diff --git a/lightx2v/models/networks/wan/infer/causvid/pre_infer.py b/lightx2v/models/networks/wan/infer/causvid/pre_infer.py index 524218e4e..7af65071b 100644 --- a/lightx2v/models/networks/wan/infer/causvid/pre_infer.py +++ b/lightx2v/models/networks/wan/infer/causvid/pre_infer.py @@ -13,11 +13,28 @@ class WanCausVidPreInfer(WanPreInfer): def __init__(self, config): super().__init__(config) self._causvid_start_frame = 0 + self._rope_cache_request_id = None self._causvid_rope_cache: Dict[tuple, Tuple[torch.Tensor, torch.Tensor | None]] = {} + def set_scheduler(self, scheduler): + super().set_scheduler(scheduler) + self.clear_rope_cache() + def set_rope(self, rope): super().set_rope(rope) + self.clear_rope_cache() + + def set_scheduler(self, scheduler): + super().set_scheduler(scheduler) + self.clear_rope_cache() + + def clear_rope_cache(self): + self._rope_cache_request_id = None self._causvid_rope_cache.clear() + self._causvid_start_frame = 0 + self.cos_sin = None + self.rope_positions = None + self.grid_sizes = (0, 0, 0) def _rope_cache_key(self, grid_sizes, start_frame): device = self.freqs.device @@ -44,6 +61,13 @@ def infer(self, weights, inputs, kv_start=0, kv_end=0): raise NotImplementedError("Sequence parallel inference is not implemented for CausVid.") if kv_start < 0 or kv_end <= kv_start: raise ValueError(f"Invalid CausVid KV range: [{kv_start}, {kv_end}).") + if not hasattr(self, "scheduler"): + raise RuntimeError("WanCausVidPreInfer scheduler is not initialized.") + + request_id = self.scheduler.rope_request_id + if request_id != self._rope_cache_request_id: + self.clear_rope_cache() + self._rope_cache_request_id = request_id # The previous local grid is the most accurate source of the number of # spatial tokens. On the first chunk, fall back to the configured value; diff --git a/lightx2v/models/networks/wan/infer/causvid/transformer_infer.py b/lightx2v/models/networks/wan/infer/causvid/transformer_infer.py index 98ef24072..7dd095bab 100755 --- a/lightx2v/models/networks/wan/infer/causvid/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/causvid/transformer_infer.py @@ -20,6 +20,12 @@ def __init__(self, config): self._kv_start = 0 self._kv_end = 0 self._cu_seqlens_cache = {} + self._cu_seqlens_cache_request_id = None + + def set_scheduler(self, scheduler): + super().set_scheduler(scheduler) + self._cu_seqlens_cache_request_id = None + self._clear_request_cache() def _init_kv_cache(self, dtype, device, kv_size=None): kv_size = kv_size or self.num_frames * self.frame_seq_length @@ -49,6 +55,22 @@ def _ensure_attention_caches(self, x): if self._cache_meta != expected_meta or self.kv_cache is None: self._init_kv_cache(x.dtype, x.device, kv_size) + def _clear_request_cache(self): + self._cu_seqlens_cache.clear() + if self.crossattn_cache is None: + return + for block_cache in self.crossattn_cache: + block_cache["k"] = None + block_cache["v"] = None + block_cache["k_img"] = None + block_cache["v_img"] = None + block_cache["is_init"] = False + + def set_scheduler(self, scheduler): + super().set_scheduler(scheduler) + self._cu_seqlens_cache_request_id = None + self._clear_request_cache() + def _get_cu_seqlens(self, q_len, kv_len): key = (int(q_len), int(kv_len)) cached = self._cu_seqlens_cache.get(key) @@ -77,9 +99,16 @@ def reset_infer_states(self, x, context): def infer(self, weights, pre_infer_out, kv_start, kv_end): if self.config["seq_parallel"]: raise NotImplementedError("Sequence parallel inference is not implemented for CausVid.") + if not hasattr(self, "scheduler"): + raise RuntimeError("WanTransformerInferCausVid scheduler is not initialized.") self._kv_start = int(kv_start) self._kv_end = int(kv_end) + request_id = self.scheduler.rope_request_id + if request_id != self._cu_seqlens_cache_request_id: + self._clear_request_cache() + self._cu_seqlens_cache_request_id = request_id + query_len = pre_infer_out.x.shape[0] if self._kv_start < 0 or self._kv_end <= self._kv_start: raise ValueError(f"Invalid CausVid KV range: [{self._kv_start}, {self._kv_end}).") diff --git a/lightx2v/models/schedulers/wan/scheduler.py b/lightx2v/models/schedulers/wan/scheduler.py index 46b7b0c03..a7de22a9f 100755 --- a/lightx2v/models/schedulers/wan/scheduler.py +++ b/lightx2v/models/schedulers/wan/scheduler.py @@ -28,6 +28,7 @@ def __init__(self, config): self.sample_guide_scale = self.config["sample_guide_scale"] self.caching_records_2 = [True] * self.config["infer_steps"] self.head_size = self.config["dim"] // self.config["num_heads"] + self.rope_request_id = 0 def refresh_from_config(self, config): self.config = config @@ -53,6 +54,7 @@ def _uses_conditioned_latent_prefix(self): return model_cls in {"wan2.2", "wan2.2_matrix_game3"} def prepare(self, seed, latent_shape, image_encoder_output=None): + self.rope_request_id += 1 if self._uses_conditioned_latent_prefix(): self.vae_encoder_out = image_encoder_output["vae_encoder_out"] if image_encoder_output is not None else None From 12150d029bf78af62863fb9b417bdde3b7825736 Mon Sep 17 00:00:00 2001 From: wushuo <540295877@qq.com> Date: Wed, 22 Jul 2026 11:34:25 +0000 Subject: [PATCH 5/9] perf(wan-dreamzero): scope inference caches to each request --- lightx2v/models/networks/wan/dreamzero_model.py | 5 +++-- lightx2v/models/networks/wan/infer/dreamzero/pre_infer.py | 3 +++ .../models/networks/wan/infer/dreamzero/transformer_infer.py | 2 +- lightx2v/models/runners/wan/wan_dreamzero_runner.py | 2 +- 4 files changed, 8 insertions(+), 4 deletions(-) diff --git a/lightx2v/models/networks/wan/dreamzero_model.py b/lightx2v/models/networks/wan/dreamzero_model.py index 7f736a946..9a54e6b6c 100644 --- a/lightx2v/models/networks/wan/dreamzero_model.py +++ b/lightx2v/models/networks/wan/dreamzero_model.py @@ -29,9 +29,10 @@ def _init_infer_class(self): def cfg_cache_name(cache_name, infer_condition): return f"{cache_name}_{'cond' if infer_condition else 'uncond'}" - def clear_cache(self, cache_name=None): + def clear_cache(self, cache_name=None, clear_pre_infer=True): self.transformer_infer.clear_cache(cache_name) - self.pre_infer.clear_cache() + if clear_pre_infer: + self.pre_infer.clear_cache() @staticmethod def _gather_cfg_tensor(tensor, group): diff --git a/lightx2v/models/networks/wan/infer/dreamzero/pre_infer.py b/lightx2v/models/networks/wan/infer/dreamzero/pre_infer.py index 66b275732..0dd687c14 100644 --- a/lightx2v/models/networks/wan/infer/dreamzero/pre_infer.py +++ b/lightx2v/models/networks/wan/infer/dreamzero/pre_infer.py @@ -75,6 +75,9 @@ def set_rope(self, rope): def clear_cache(self): self._context_projection_cache.clear() + self._freqs_cache.clear() + self._prepared_freqs_cache.clear() + self._time_embedding_cache.clear() @staticmethod def _rope_params(max_seq_len, dim, theta=10000): diff --git a/lightx2v/models/networks/wan/infer/dreamzero/transformer_infer.py b/lightx2v/models/networks/wan/infer/dreamzero/transformer_infer.py index 1a9db1a89..509e7c3a1 100644 --- a/lightx2v/models/networks/wan/infer/dreamzero/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/dreamzero/transformer_infer.py @@ -195,7 +195,6 @@ def clear_cache(self, cache_name=None): if cache_name is None: self.kv_caches.clear() self.cross_attn_kv_caches.clear() - self._cu_seqlens_cache.clear() else: cache_names = {cache_name, f"{cache_name}_cond", f"{cache_name}_uncond"} for name in cache_names: @@ -204,6 +203,7 @@ def clear_cache(self, cache_name=None): key_cache_name = key[0] if isinstance(key, tuple) else key if key_cache_name in cache_names: self.cross_attn_kv_caches.pop(key, None) + self._cu_seqlens_cache.clear() @staticmethod def _token_modulation(x): diff --git a/lightx2v/models/runners/wan/wan_dreamzero_runner.py b/lightx2v/models/runners/wan/wan_dreamzero_runner.py index 26ee9ebb3..fce471d91 100644 --- a/lightx2v/models/runners/wan/wan_dreamzero_runner.py +++ b/lightx2v/models/runners/wan/wan_dreamzero_runner.py @@ -568,7 +568,7 @@ def run_chunk(self, frame_indices): self.last_chunk_started_new_segment = len(frame_indices) == 1 or self.current_start_frame >= int(self.config.get("local_attn_size", 9)) if self.last_chunk_started_new_segment: self.current_start_frame = 0 - self.model.clear_cache(self.cache_name) + self.model.clear_cache(self.cache_name, clear_pre_infer=False) observed_latents = None if self.current_start_frame == 0: From 984e4bfa0509899e995633fb8f272ecc838a750c Mon Sep 17 00:00:00 2001 From: wushuo <540295877@qq.com> Date: Wed, 22 Jul 2026 11:34:52 +0000 Subject: [PATCH 6/9] perf(infinitetalk): scope RoPE caches to each request --- .../wan/infer/infinitetalk/pre_infer.py | 38 +++++++++++++++---- .../networks/wan/infer/infinitetalk/rope.py | 12 +++--- .../schedulers/wan/infinitetalk/scheduler.py | 3 ++ 3 files changed, 39 insertions(+), 14 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/infinitetalk/pre_infer.py b/lightx2v/models/networks/wan/infer/infinitetalk/pre_infer.py index d1fb99d32..bcb3a7fbc 100644 --- a/lightx2v/models/networks/wan/infer/infinitetalk/pre_infer.py +++ b/lightx2v/models/networks/wan/infer/infinitetalk/pre_infer.py @@ -18,14 +18,34 @@ def __init__(self, config): self.audio_output_dim = config.get("infinitetalk_audio_output_dim", 768) self.audio_rope = None self.rope_1d = RotaryPositionalEmbedding1D(self.head_size) - self._audio_rope_cache = {} + self.clear_request_rope_cache() + + def clear_request_rope_cache(self): + self._rope_cache_request_id = None + self._audio_rope_cache = None + self.cos_sin = None + self.rope_positions = None + self.grid_sizes = (0, 0, 0) def set_audio_rope(self, rope): self.audio_rope = rope - self._audio_rope_cache.clear() + self.clear_request_rope_cache() + + def set_scheduler(self, scheduler): + super().set_scheduler(scheduler) + self.clear_request_rope_cache() + + def _sync_request_rope_cache(self): + request_id = self.scheduler.rope_request_id + if request_id == self._rope_cache_request_id: + return + + self.clear_request_rope_cache() + self._rope_cache_request_id = request_id @torch.no_grad() def infer(self, weights, inputs, kv_start=0, kv_end=0): + self._sync_request_rope_cache() original_latents = self.scheduler.latents image_encoder_output = inputs.get("image_encoder_output", {}) original_vae_encoder_out = image_encoder_output.get("vae_encoder_out") @@ -134,8 +154,8 @@ def _append_audio_rope_cache(self, pre_out): device.index, dtype, ) - cached = self._audio_rope_cache.get(cache_key) - if cached is None: + cached = self._audio_rope_cache + if cached is None or cached[0] != cache_key: h1 = torch.tensor((0, self.config.get("infinitetalk_class_interval", 4)), dtype=dtype, device=device) class_range = self.config.get("infinitetalk_class_range", 24) h2 = torch.tensor( @@ -163,8 +183,10 @@ def _append_audio_rope_cache(self, pre_out): per_frame[token_edges[idx] : token_edges[idx + 1]] = rope_centers[idx] encoder_pos = per_frame.repeat(grid_t) encoder_freqs, encoder_positions = self.rope_1d.prepare(self.audio_rope, encoder_pos) - cached = (rope_ranges, encoder_freqs, encoder_positions) - self._audio_rope_cache[cache_key] = cached + value = (rope_ranges, encoder_freqs, encoder_positions) + self._audio_rope_cache = (cache_key, value) + else: + value = cached[1] - pre_out.adapter_args["audio_rope_ranges"] = cached[0] - pre_out.adapter_args["audio_encoder_rope"] = cached[1:] + pre_out.adapter_args["audio_rope_ranges"] = value[0] + pre_out.adapter_args["audio_encoder_rope"] = value[1:] diff --git a/lightx2v/models/networks/wan/infer/infinitetalk/rope.py b/lightx2v/models/networks/wan/infer/infinitetalk/rope.py index e6c5b7b8a..0b5d4f33e 100644 --- a/lightx2v/models/networks/wan/infer/infinitetalk/rope.py +++ b/lightx2v/models/networks/wan/infer/infinitetalk/rope.py @@ -6,15 +6,15 @@ class RotaryPositionalEmbedding1D: def __init__(self, head_dim, base=10000): self.head_dim = head_dim self.base = base - self._inv_freq_cache = {} + self._inv_freq_cache_key = None + self._inv_freq = None def _get_inv_freq(self, device): cache_key = (device.type, device.index) - inv_freq = self._inv_freq_cache.get(cache_key) - if inv_freq is None: - inv_freq = 1.0 / (self.base ** (torch.arange(0, self.head_dim, 2, device=device, dtype=torch.float32) / self.head_dim)) - self._inv_freq_cache[cache_key] = inv_freq - return inv_freq + if cache_key != self._inv_freq_cache_key: + self._inv_freq_cache_key = cache_key + self._inv_freq = 1.0 / (self.base ** (torch.arange(0, self.head_dim, 2, device=device, dtype=torch.float32) / self.head_dim)) + return self._inv_freq def prepare(self, rope, pos_indices): angles = torch.einsum("..., f -> ... f", pos_indices.float(), self._get_inv_freq(pos_indices.device)) diff --git a/lightx2v/models/schedulers/wan/infinitetalk/scheduler.py b/lightx2v/models/schedulers/wan/infinitetalk/scheduler.py index 8fc6abe8d..b0c57370f 100644 --- a/lightx2v/models/schedulers/wan/infinitetalk/scheduler.py +++ b/lightx2v/models/schedulers/wan/infinitetalk/scheduler.py @@ -25,6 +25,7 @@ def __init__(self, config): self.latent_motion_frames = None self.cur_motion_frames_latent_num = 0 self.is_first_clip = True + self.rope_request_id = 0 def seed_everything(self, seed): seed = seed if seed >= 0 else random.randint(0, 99999999) @@ -36,6 +37,8 @@ def seed_everything(self, seed): return seed def prepare(self, seed, latent_shape, latent_motion_frames=None, is_first_clip=True, cur_motion_frames_latent_num=1): + if is_first_clip: + self.rope_request_id += 1 self.latents = torch.randn(*latent_shape, dtype=torch.float32, device=AI_DEVICE) self.latent_motion_frames = latent_motion_frames self.is_first_clip = is_first_clip From 37a4ba8a60cad136a6446c4d2cb172b8eaaa5bd8 Mon Sep 17 00:00:00 2001 From: wushuo <540295877@qq.com> Date: Wed, 22 Jul 2026 11:35:24 +0000 Subject: [PATCH 7/9] perf(worldmirror): bound request-scoped positional caches --- .../networks/worldmirror/models/layers/block.py | 4 ++++ .../networks/worldmirror/models/layers/norm_rope.py | 4 ++++ .../networks/worldmirror/models/layers/rope.py | 8 ++++++++ .../worldmirror/models/models/visual_transformer.py | 12 +++++++++++- 4 files changed, 27 insertions(+), 1 deletion(-) diff --git a/lightx2v/models/networks/worldmirror/models/layers/block.py b/lightx2v/models/networks/worldmirror/models/layers/block.py index 9d678a3c8..8e8bcdf9d 100755 --- a/lightx2v/models/networks/worldmirror/models/layers/block.py +++ b/lightx2v/models/networks/worldmirror/models/layers/block.py @@ -172,6 +172,10 @@ def add_residual(x, brange, residual, residual_scale_factor, scaling_vector=None attn_bias_cache: Dict[Tuple, Any] = {} +def clear_attn_bias_cache() -> None: + attn_bias_cache.clear() + + def get_attn_bias_and_cat(x_list, branges=None): """ this will perform the index select, cat the tensors, and provide the attn_bias from cache diff --git a/lightx2v/models/networks/worldmirror/models/layers/norm_rope.py b/lightx2v/models/networks/worldmirror/models/layers/norm_rope.py index 92fd9508c..a1c0b8778 100644 --- a/lightx2v/models/networks/worldmirror/models/layers/norm_rope.py +++ b/lightx2v/models/networks/worldmirror/models/layers/norm_rope.py @@ -43,9 +43,13 @@ class PositionGetter: def __init__(self) -> None: self.position_cache: Dict[Tuple[int, int, torch.device], torch.Tensor] = {} + def clear_cache(self) -> None: + self.position_cache.clear() + def __call__(self, batch_size: int, height: int, width: int, device: torch.device) -> torch.Tensor: cache_key = (height, width, torch.device(device)) if cache_key not in self.position_cache: + self.position_cache.clear() y_coords = torch.arange(height, device=device) x_coords = torch.arange(width, device=device) self.position_cache[cache_key] = torch.cartesian_prod(y_coords, x_coords) diff --git a/lightx2v/models/networks/worldmirror/models/layers/rope.py b/lightx2v/models/networks/worldmirror/models/layers/rope.py index d425a2b22..2b24238d0 100755 --- a/lightx2v/models/networks/worldmirror/models/layers/rope.py +++ b/lightx2v/models/networks/worldmirror/models/layers/rope.py @@ -66,6 +66,9 @@ def __init__(self): """Initializes the position generator with an empty cache.""" self.position_cache: Dict[Tuple[int, int, torch.device], torch.Tensor] = {} + def clear_cache(self) -> None: + self.position_cache.clear() + def __call__(self, batch_size: int, height: int, width: int, device: torch.device) -> torch.Tensor: """Generates spatial positions for a batch of patches. @@ -81,6 +84,7 @@ def __call__(self, batch_size: int, height: int, width: int, device: torch.devic """ cache_key = (height, width, torch.device(device)) if cache_key not in self.position_cache: + self.position_cache.clear() y_coords = torch.arange(height, device=device) x_coords = torch.arange(width, device=device) positions = torch.cartesian_prod(y_coords, x_coords) @@ -118,6 +122,9 @@ def __init__( self.scaling_factor = scaling_factor self.frequency_cache: Dict[Tuple, Tuple[torch.Tensor, torch.Tensor]] = {} + def clear_cache(self) -> None: + self.frequency_cache.clear() + def _compute_frequency_components(self, dim: int, seq_len: int, device: torch.device, dtype: torch.dtype) -> Tuple[torch.Tensor, torch.Tensor]: """Computes frequency components for rotary embeddings. @@ -132,6 +139,7 @@ def _compute_frequency_components(self, dim: int, seq_len: int, device: torch.de """ cache_key = (dim, seq_len, device, dtype) if cache_key not in self.frequency_cache: + self.frequency_cache.clear() # Compute frequency bands exponents = torch.arange(0, dim, 2, device=device).float() / dim inv_freq = 1.0 / (self.base_frequency**exponents) diff --git a/lightx2v/models/networks/worldmirror/models/models/visual_transformer.py b/lightx2v/models/networks/worldmirror/models/models/visual_transformer.py index 2113e0247..7299e38fe 100755 --- a/lightx2v/models/networks/worldmirror/models/models/visual_transformer.py +++ b/lightx2v/models/networks/worldmirror/models/models/visual_transformer.py @@ -9,7 +9,7 @@ from ...comm.communication import _Allgather from ...comm.padding import depad_by_length, minimal_pad_to_divisible, pad_by_length from ..layers import PatchEmbed, PatchEmbed_Mlp -from ..layers.block import Block +from ..layers.block import Block, clear_attn_bias_cache from ..layers.vision_transformer import vit_base, vit_giant2, vit_large, vit_small logger = logging.getLogger(__name__) @@ -212,6 +212,15 @@ def _init_rotary_position_embedding(self, rope_base, normalized_rope, head_dim, ) self.pos_getter = PositionGetter() if self.rope is not None else None + def begin_cache_request(self) -> None: + clear_attn_bias_cache() + if self.pos_getter is None: + return + self.pos_getter.clear_cache() + clear_rope_cache = getattr(self.rope, "clear_cache", None) + if clear_rope_cache is not None: + clear_rope_cache() + def _init_transformer_blocks(self, block_fn, embed_dim, num_heads, mlp_ratio, qkv_bias, proj_bias, ffn_bias, init_values, qk_norm): self.frame_blocks = nn.ModuleList( [ @@ -264,6 +273,7 @@ def forward( Returns: (list[torch.Tensor], int): List of attention block outputs and patch_start_idx """ + self.begin_cache_request() depth_maps, ray_dirs, poses = priors if priors is not None else (None, None, None) # Slice to context frames if specified From 357d13d65f2eed0c196044fd941e154aa0a15a1b Mon Sep 17 00:00:00 2001 From: wushuo <540295877@qq.com> Date: Wed, 22 Jul 2026 12:44:12 +0000 Subject: [PATCH 8/9] fix(wan): isolate CausVid request caches --- .../networks/wan/infer/causvid/pre_infer.py | 32 ++++---- .../wan/infer/causvid/transformer_infer.py | 78 ++++++++++++------- 2 files changed, 70 insertions(+), 40 deletions(-) diff --git a/lightx2v/models/networks/wan/infer/causvid/pre_infer.py b/lightx2v/models/networks/wan/infer/causvid/pre_infer.py index 7af65071b..f28e431dc 100644 --- a/lightx2v/models/networks/wan/infer/causvid/pre_infer.py +++ b/lightx2v/models/networks/wan/infer/causvid/pre_infer.py @@ -12,23 +12,22 @@ class WanCausVidPreInfer(WanPreInfer): def __init__(self, config): super().__init__(config) - self._causvid_start_frame = 0 - self._rope_cache_request_id = None + self.scheduler = None self._causvid_rope_cache: Dict[tuple, Tuple[torch.Tensor, torch.Tensor | None]] = {} + self._reset_request_cache() def set_scheduler(self, scheduler): super().set_scheduler(scheduler) - self.clear_rope_cache() + self._reset_request_cache() + self._rope_cache_request_id = scheduler.rope_request_id def set_rope(self, rope): super().set_rope(rope) - self.clear_rope_cache() - - def set_scheduler(self, scheduler): - super().set_scheduler(scheduler) - self.clear_rope_cache() + self._reset_request_cache() + if self.scheduler is not None: + self._rope_cache_request_id = self.scheduler.rope_request_id - def clear_rope_cache(self): + def _reset_request_cache(self): self._rope_cache_request_id = None self._causvid_rope_cache.clear() self._causvid_start_frame = 0 @@ -36,6 +35,14 @@ def clear_rope_cache(self): self.rope_positions = None self.grid_sizes = (0, 0, 0) + def _sync_request_cache(self): + request_id = self.scheduler.rope_request_id + if request_id == self._rope_cache_request_id: + return + + self._reset_request_cache() + self._rope_cache_request_id = request_id + def _rope_cache_key(self, grid_sizes, start_frame): device = self.freqs.device return ( @@ -61,13 +68,10 @@ def infer(self, weights, inputs, kv_start=0, kv_end=0): raise NotImplementedError("Sequence parallel inference is not implemented for CausVid.") if kv_start < 0 or kv_end <= kv_start: raise ValueError(f"Invalid CausVid KV range: [{kv_start}, {kv_end}).") - if not hasattr(self, "scheduler"): + if self.scheduler is None: raise RuntimeError("WanCausVidPreInfer scheduler is not initialized.") - request_id = self.scheduler.rope_request_id - if request_id != self._rope_cache_request_id: - self.clear_rope_cache() - self._rope_cache_request_id = request_id + self._sync_request_cache() # The previous local grid is the most accurate source of the number of # spatial tokens. On the first chunk, fall back to the configured value; diff --git a/lightx2v/models/networks/wan/infer/causvid/transformer_infer.py b/lightx2v/models/networks/wan/infer/causvid/transformer_infer.py index 7dd095bab..bcc1152fd 100755 --- a/lightx2v/models/networks/wan/infer/causvid/transformer_infer.py +++ b/lightx2v/models/networks/wan/infer/causvid/transformer_infer.py @@ -12,20 +12,19 @@ class WanTransformerInferCausVid(WanOffloadTransformerInfer): def __init__(self, config): super().__init__(config) + self.scheduler = None self.num_frames = config["num_frames"] self.frame_seq_length = config["frame_seq_length"] self.kv_cache = None self.crossattn_cache = None self._cache_meta = None - self._kv_start = 0 - self._kv_end = 0 self._cu_seqlens_cache = {} - self._cu_seqlens_cache_request_id = None + self._reset_request_cache() def set_scheduler(self, scheduler): super().set_scheduler(scheduler) - self._cu_seqlens_cache_request_id = None - self._clear_request_cache() + self._reset_request_cache() + self._kv_cache_request_id = scheduler.rope_request_id def _init_kv_cache(self, dtype, device, kv_size=None): kv_size = kv_size or self.num_frames * self.frame_seq_length @@ -50,12 +49,30 @@ def _init_kv_cache(self, dtype, device, kv_size=None): def _ensure_attention_caches(self, x): configured_size = self.num_frames * self.frame_seq_length - kv_size = max(configured_size, self._kv_end) - expected_meta = (x.dtype, x.device, kv_size) - if self._cache_meta != expected_meta or self.kv_cache is None: - self._init_kv_cache(x.dtype, x.device, kv_size) + required_capacity = max(configured_size, self._kv_end, self._valid_kv_end) + if self._cache_meta is None or self.kv_cache is None: + self._init_kv_cache(x.dtype, x.device, required_capacity) + return + + cached_dtype, cached_device, cached_capacity = self._cache_meta + if cached_dtype == x.dtype and cached_device == x.device and cached_capacity >= required_capacity: + return + + old_kv_cache = self.kv_cache + valid_kv_end = self._valid_kv_end + self._init_kv_cache(x.dtype, x.device, required_capacity) + if valid_kv_end == 0: + return + + for old_block_cache, new_block_cache in zip(old_kv_cache, self.kv_cache, strict=True): + new_block_cache["k"][:valid_kv_end].copy_(old_block_cache["k"][:valid_kv_end]) + new_block_cache["v"][:valid_kv_end].copy_(old_block_cache["v"][:valid_kv_end]) - def _clear_request_cache(self): + def _reset_request_cache(self): + self._kv_cache_request_id = None + self._kv_start = 0 + self._kv_end = 0 + self._valid_kv_end = 0 self._cu_seqlens_cache.clear() if self.crossattn_cache is None: return @@ -66,10 +83,14 @@ def _clear_request_cache(self): block_cache["v_img"] = None block_cache["is_init"] = False - def set_scheduler(self, scheduler): - super().set_scheduler(scheduler) - self._cu_seqlens_cache_request_id = None - self._clear_request_cache() + def _sync_request_cache(self): + request_id = self.scheduler.rope_request_id + if request_id == self._kv_cache_request_id: + return request_id + + self._reset_request_cache() + self._kv_cache_request_id = request_id + return request_id def _get_cu_seqlens(self, q_len, kv_len): key = (int(q_len), int(kv_len)) @@ -99,24 +120,29 @@ def reset_infer_states(self, x, context): def infer(self, weights, pre_infer_out, kv_start, kv_end): if self.config["seq_parallel"]: raise NotImplementedError("Sequence parallel inference is not implemented for CausVid.") - if not hasattr(self, "scheduler"): + if self.scheduler is None: raise RuntimeError("WanTransformerInferCausVid scheduler is not initialized.") - self._kv_start = int(kv_start) - self._kv_end = int(kv_end) - request_id = self.scheduler.rope_request_id - if request_id != self._cu_seqlens_cache_request_id: - self._clear_request_cache() - self._cu_seqlens_cache_request_id = request_id + request_id = self._sync_request_cache() + kv_start = int(kv_start) + kv_end = int(kv_end) query_len = pre_infer_out.x.shape[0] - if self._kv_start < 0 or self._kv_end <= self._kv_start: - raise ValueError(f"Invalid CausVid KV range: [{self._kv_start}, {self._kv_end}).") - if self._kv_end - self._kv_start != query_len: - raise ValueError(f"CausVid query length must match its KV cache range: query_len={query_len}, range=[{self._kv_start}, {self._kv_end}).") + if kv_start < 0 or kv_end <= kv_start: + raise ValueError(f"Invalid CausVid KV range: [{kv_start}, {kv_end}).") + if self._valid_kv_end == 0 and kv_start != 0: + raise ValueError(f"CausVid request_id={request_id} must start at kv_start=0; got kv_start={kv_start}, kv_end={kv_end}.") + if kv_start > self._valid_kv_end: + raise ValueError(f"CausVid request_id={request_id} leaves an uninitialized KV gap: valid_kv_end={self._valid_kv_end}, kv_start={kv_start}, kv_end={kv_end}.") + if kv_end - kv_start != query_len: + raise ValueError(f"CausVid query length must match its KV cache range: query_len={query_len}, range=[{kv_start}, {kv_end}).") + self._kv_start = kv_start + self._kv_end = kv_end self._ensure_attention_caches(pre_infer_out.x) - return super().infer(weights, pre_infer_out) + output = super().infer(weights, pre_infer_out) + self._valid_kv_end = max(self._valid_kv_end, kv_end) + return output def infer_self_attn(self, phase, x, shift_msa, scale_msa, grid_sizes=None): norm1_quant = None From cc5c967106e064d6054d607f866b79009d4ccb1b Mon Sep 17 00:00:00 2001 From: Watebear Date: Thu, 23 Jul 2026 14:20:55 +0800 Subject: [PATCH 9/9] Add Comment --- lightx2v/models/networks/ernie_image/infer/pre_infer.py | 1 + 1 file changed, 1 insertion(+) diff --git a/lightx2v/models/networks/ernie_image/infer/pre_infer.py b/lightx2v/models/networks/ernie_image/infer/pre_infer.py index 7f3145340..099202186 100644 --- a/lightx2v/models/networks/ernie_image/infer/pre_infer.py +++ b/lightx2v/models/networks/ernie_image/infer/pre_infer.py @@ -55,6 +55,7 @@ def __init__(self, config): def clear_rope_cache(self): self._rope_cache_request_id = None + # True:infer_condition=True; False:infer_condition=False self._rope_cache = {True: None, False: None} def set_rope(self, rope):