From c4db00609a93cf25be63a570b4752f21fa27a2a1 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 1 Dec 2025 17:39:50 +0200 Subject: [PATCH] Use torch.chunk for chunking --- wanvideo/modules/model.py | 35 ++++++++++++----------------------- 1 file changed, 12 insertions(+), 23 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 4cdb271..a31665d 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -414,29 +414,24 @@ class WanSelfAttention(nn.Module): def _qkv_fn_with_rope(self, x, linear_layer, norm_layer, freqs, num_chunks=1, is_longcat=False): b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim - + use_chunked = num_chunks > 1 if use_chunked: - chunk_sizes = [s // num_chunks + (1 if i < s % num_chunks else 0) - for i in range(num_chunks)] - out = torch.empty(b, s, n, d, dtype=x.dtype, device=x.device) - start_idx = 0 - for size in chunk_sizes: - end_idx = start_idx + size - - x_chunk = x[:, start_idx:end_idx] - + + for i, x_chunk in enumerate(x.chunk(num_chunks, dim=1)): + chunk_size = x_chunk.size(1) + start_idx = i * (s // num_chunks + (1 if i < s % num_chunks else 0)) + if is_longcat: - chunk = linear_layer(x_chunk).view(b, size, n, d) + chunk = linear_layer(x_chunk).view(b, chunk_size, n, d) chunk = norm_layer(chunk.float()).to(x.dtype) else: - chunk = norm_layer(linear_layer(x_chunk).to(norm_layer.weight.dtype)).to(x.dtype).view(b, size, n, d) + chunk = norm_layer(linear_layer(x_chunk).to(norm_layer.weight.dtype)).to(x.dtype).view(b, chunk_size, n, d) + + freqs_chunk = freqs[:, start_idx:start_idx + chunk_size] if freqs.shape[1] > 1 else freqs + out[:, start_idx:start_idx + chunk_size] = apply_rope_comfy1(chunk, freqs_chunk) - freqs_chunk = freqs[:, start_idx:end_idx] if freqs.shape[1] > 1 else freqs - out[:, start_idx:end_idx] = apply_rope_comfy1(chunk, freqs_chunk) - - start_idx = end_idx return out else: if is_longcat: @@ -982,13 +977,7 @@ class WanAttentionBlock(nn.Module): seq_len = mod_x.shape[1] if seq_len <= 8192 or num_chunks <= 1: return self.ffn(mod_x) - - chunk_size = (seq_len + num_chunks - 1) // num_chunks - for i in range(0, seq_len, chunk_size): - end_idx = min(i + chunk_size, seq_len) - mod_x[:, i:end_idx] = self.ffn(mod_x[:, i:end_idx].contiguous()) - - return mod_x + return torch.cat([self.ffn(chunk.contiguous()) for chunk in mod_x.chunk(num_chunks, dim=1)], dim=1) #region attention forward def forward(