Skip to content

Commit 65c24f6

Browse files
committed
cleanup
1 parent 88e8dd3 commit 65c24f6

2 files changed

Lines changed: 64 additions & 166 deletions

File tree

src/diffusers/models/transformers/transformer_hunyuan_video_framepack.py

Lines changed: 60 additions & 162 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@
1212
# See the License for the specific language governing permissions and
1313
# limitations under the License.
1414

15-
from typing import Any, Dict, Optional, Tuple
15+
from typing import Any, Dict, List, Optional, Tuple
1616

1717
import torch
1818
import torch.nn as nn
@@ -22,6 +22,7 @@
2222
from ...loaders import FromOriginalModelMixin, PeftAdapterMixin
2323
from ...utils import USE_PEFT_BACKEND, get_logger, scale_lora_layers, unscale_lora_layers
2424
from ..cache_utils import CacheMixin
25+
from ..embeddings import get_1d_rotary_pos_embed
2526
from ..modeling_outputs import Transformer2DModelOutput
2627
from ..modeling_utils import ModelMixin
2728
from ..normalization import AdaLayerNormContinuous
@@ -37,92 +38,49 @@
3738
logger = get_logger(__name__) # pylint: disable=invalid-name
3839

3940

40-
# class HunyuanVideoFramepackRotaryPosEmbed(nn.Module):
41-
# def __init__(self, patch_size: int, patch_size_t: int, rope_dim: List[int], theta: float = 256.0) -> None:
42-
# super().__init__()
43-
44-
# self.patch_size = patch_size
45-
# self.patch_size_t = patch_size_t
46-
# self.rope_dim = rope_dim
47-
# self.theta = theta
48-
49-
# def forward(self, frame_indices: torch.Tensor, height: int, width: int, device: torch.device):
50-
# frame_indices = frame_indices.unbind(0)
51-
# # This is from the original code. We don't call _forward for each batch index because we know that
52-
# # each batch has the same frame indices. However, it may be possible that the frame indices don't
53-
# # always be the same for every item in a batch (such as in training). We cannot use the original
54-
# # implementation because our `apply_rotary_emb` function broadcasts across the batch dim.
55-
# # freqs = [self._forward(f, height, width, device) for f in frame_indices]
56-
# # freqs_cos, freqs_sin = zip(*freqs)
57-
# # freqs_cos = torch.stack(freqs_cos, dim=0) # [B, W * H * T, D / 2]
58-
# # freqs_sin = torch.stack(freqs_sin, dim=0) # [B, W * H * T, D / 2]
59-
# # return freqs_cos, freqs_sin
60-
# return self._forward(frame_indices[0], height, width, device)
61-
62-
# def _forward(self, frame_indices, height, width, device):
63-
# height = height // self.patch_size
64-
# width = width // self.patch_size
65-
# grid = torch.meshgrid(
66-
# frame_indices.to(device=device, dtype=torch.float32),
67-
# torch.arange(0, height, device=device, dtype=torch.float32),
68-
# torch.arange(0, width, device=device, dtype=torch.float32),
69-
# indexing="ij",
70-
# ) # 3 * [W, H, T]
71-
# grid = torch.stack(grid, dim=0) # [3, W, H, T]
72-
73-
# freqs = []
74-
# for i in range(3):
75-
# freq = get_1d_rotary_pos_embed(self.rope_dim[i], grid[i].reshape(-1), self.theta, use_real=True)
76-
# freqs.append(freq)
77-
78-
# freqs_cos = torch.cat([f[0] for f in freqs], dim=1) # (W * H * T, D / 2)
79-
# freqs_sin = torch.cat([f[1] for f in freqs], dim=1) # (W * H * T, D / 2)
80-
81-
# return freqs_cos, freqs_sin
82-
83-
84-
class HunyuanVideoRotaryPosEmbed(nn.Module):
85-
def __init__(self, rope_dim, theta):
41+
class HunyuanVideoFramepackRotaryPosEmbed(nn.Module):
42+
def __init__(self, patch_size: int, patch_size_t: int, rope_dim: List[int], theta: float = 256.0) -> None:
8643
super().__init__()
87-
self.DT, self.DY, self.DX = rope_dim
44+
45+
self.patch_size = patch_size
46+
self.patch_size_t = patch_size_t
47+
self.rope_dim = rope_dim
8848
self.theta = theta
8949

90-
@torch.no_grad()
91-
def get_frequency(self, dim, pos):
92-
T, H, W = pos.shape
93-
freqs = 1.0 / (
94-
self.theta ** (torch.arange(0, dim, 2, dtype=torch.float32, device=pos.device)[: (dim // 2)] / dim)
95-
)
96-
freqs = torch.outer(freqs, pos.reshape(-1)).unflatten(-1, (T, H, W)).repeat_interleave(2, dim=0)
97-
return freqs.cos(), freqs.sin()
98-
99-
@torch.no_grad()
100-
def forward_inner(self, frame_indices, height, width, device):
101-
# TODO(aryan)
102-
height = height // 2
103-
width = width // 2
104-
GT, GY, GX = torch.meshgrid(
50+
def forward(self, frame_indices: torch.Tensor, height: int, width: int, device: torch.device):
51+
# This is from the original code. We don't call _forward for each batch index because we know that
52+
# each batch has the same frame indices. However, it may be possible that the frame indices don't
53+
# always be the same for every item in a batch (such as in training). We cannot use the original
54+
# implementation because our `apply_rotary_emb` function broadcasts across the batch dim, so we'd
55+
# need to first implement another attention processor or modify the existing one with different apply_rotary_emb
56+
# frame_indices = frame_indices.unbind(0)
57+
# freqs = [self._forward(f, height, width, device) for f in frame_indices]
58+
# freqs_cos, freqs_sin = zip(*freqs)
59+
# freqs_cos = torch.stack(freqs_cos, dim=0) # [B, W * H * T, D / 2]
60+
# freqs_sin = torch.stack(freqs_sin, dim=0) # [B, W * H * T, D / 2]
61+
# return freqs_cos, freqs_sin
62+
return self._forward(frame_indices, height, width, device)
63+
64+
def _forward(self, frame_indices, height, width, device):
65+
height = height // self.patch_size
66+
width = width // self.patch_size
67+
grid = torch.meshgrid(
10568
frame_indices.to(device=device, dtype=torch.float32),
10669
torch.arange(0, height, device=device, dtype=torch.float32),
10770
torch.arange(0, width, device=device, dtype=torch.float32),
10871
indexing="ij",
109-
)
110-
111-
FCT, FST = self.get_frequency(self.DT, GT)
112-
FCY, FSY = self.get_frequency(self.DY, GY)
113-
FCX, FSX = self.get_frequency(self.DX, GX)
72+
) # 3 * [W, H, T]
73+
grid = torch.stack(grid, dim=0) # [3, W, H, T]
11474

115-
result = torch.cat([FCT, FCY, FCX, FST, FSY, FSX], dim=0)
75+
freqs = []
76+
for i in range(3):
77+
freq = get_1d_rotary_pos_embed(self.rope_dim[i], grid[i].reshape(-1), self.theta, use_real=True)
78+
freqs.append(freq)
11679

117-
return result.to(device)
80+
freqs_cos = torch.cat([f[0] for f in freqs], dim=1) # (W * H * T, D / 2)
81+
freqs_sin = torch.cat([f[1] for f in freqs], dim=1) # (W * H * T, D / 2)
11882

119-
@torch.no_grad()
120-
def forward(self, frame_indices, height, width, device):
121-
return self.forward_inner(frame_indices[0], height, width, device).unsqueeze(0)
122-
# frame_indices = frame_indices.unbind(0)
123-
# results = [self.forward_inner(f, height, width, device) for f in frame_indices]
124-
# results = torch.stack(results, dim=0)
125-
# return results
83+
return freqs_cos, freqs_sin
12684

12785

12886
class FramepackClipVisionProjection(nn.Module):
@@ -216,8 +174,7 @@ def __init__(
216174
)
217175

218176
# 2. RoPE
219-
# self.rope = HunyuanVideoFramepackRotaryPosEmbed(patch_size, patch_size_t, rope_axes_dim, rope_theta)
220-
self.rope = HunyuanVideoRotaryPosEmbed(rope_axes_dim, rope_theta)
177+
self.rope = HunyuanVideoFramepackRotaryPosEmbed(patch_size, patch_size_t, rope_axes_dim, rope_theta)
221178

222179
# 3. Dual stream transformer blocks
223180
self.transformer_blocks = nn.ModuleList(
@@ -320,18 +277,17 @@ def forward(
320277
attention_mask = torch.zeros(
321278
batch_size, sequence_length, device=hidden_states.device, dtype=torch.bool
322279
) # [B, N]
323-
324280
effective_condition_sequence_length = encoder_attention_mask.sum(dim=1, dtype=torch.int) # [B,]
325281
effective_sequence_length = latent_sequence_length + effective_condition_sequence_length
326282

327-
if batch_size == 1:
328-
encoder_hidden_states = encoder_hidden_states[:, : effective_condition_sequence_length[0]]
329-
attention_mask = None
330-
else:
331-
for i in range(batch_size):
332-
attention_mask[i, : effective_sequence_length[i]] = True
333-
# [B, 1, 1, N], for broadcasting across attention heads
334-
attention_mask = attention_mask.unsqueeze(1).unsqueeze(1)
283+
# if batch_size == 1:
284+
# encoder_hidden_states = encoder_hidden_states[:, : effective_condition_sequence_length[0]]
285+
# attention_mask = None
286+
# else:
287+
for i in range(batch_size):
288+
attention_mask[i, : effective_sequence_length[i]] = True
289+
# [B, 1, 1, N], for broadcasting across attention heads
290+
attention_mask = attention_mask.unsqueeze(1).unsqueeze(1)
335291

336292
if torch.is_grad_enabled() and self.gradient_checkpointing:
337293
for block in self.transformer_blocks:
@@ -393,7 +349,8 @@ def _pack_history_states(
393349
image_rotary_emb = self.rope(
394350
frame_indices=indices_latents, height=height, width=width, device=hidden_states.device
395351
)
396-
image_rotary_emb = image_rotary_emb.flatten(2).transpose(1, 2)
352+
image_rotary_emb = list(image_rotary_emb) # convert tuple to list for in-place modification
353+
pph, ppw = height // self.config.patch_size, width // self.config.patch_size
397354

398355
latents_clean, latents_history_2x, latents_history_4x = self.clean_x_embedder(
399356
latents_clean, latents_history_2x, latents_history_4x
@@ -405,93 +362,34 @@ def _pack_history_states(
405362
image_rotary_emb_clean = self.rope(
406363
frame_indices=indices_latents_clean, height=height, width=width, device=latents_clean.device
407364
)
408-
image_rotary_emb_clean = image_rotary_emb_clean.flatten(2).transpose(1, 2)
409-
image_rotary_emb = torch.cat([image_rotary_emb_clean, image_rotary_emb], dim=1)
365+
image_rotary_emb[0] = torch.cat([image_rotary_emb_clean[0], image_rotary_emb[0]], dim=0)
366+
image_rotary_emb[1] = torch.cat([image_rotary_emb_clean[1], image_rotary_emb[1]], dim=0)
410367

411368
if latents_history_2x is not None and indices_latents_history_2x is not None:
412369
hidden_states = torch.cat([latents_history_2x, hidden_states], dim=1)
413370

414371
image_rotary_emb_history_2x = self.rope(
415372
frame_indices=indices_latents_history_2x, height=height, width=width, device=latents_history_2x.device
416373
)
417-
image_rotary_emb_history_2x = _pad_for_3d_conv(image_rotary_emb_history_2x, (2, 2, 2))
418-
image_rotary_emb_history_2x = _center_down_sample_3d(image_rotary_emb_history_2x, (2, 2, 2))
419-
image_rotary_emb_history_2x = image_rotary_emb_history_2x.flatten(2).transpose(1, 2)
420-
image_rotary_emb = torch.cat([image_rotary_emb_history_2x, image_rotary_emb], dim=1)
374+
image_rotary_emb_history_2x = self._pad_rotary_emb(
375+
image_rotary_emb_history_2x, indices_latents_history_2x.size(0), pph, ppw, (2, 2, 2)
376+
)
377+
image_rotary_emb[0] = torch.cat([image_rotary_emb_history_2x[0], image_rotary_emb[0]], dim=0)
378+
image_rotary_emb[1] = torch.cat([image_rotary_emb_history_2x[1], image_rotary_emb[1]], dim=0)
421379

422380
if latents_history_4x is not None and indices_latents_history_4x is not None:
423381
hidden_states = torch.cat([latents_history_4x, hidden_states], dim=1)
424382

425383
image_rotary_emb_history_4x = self.rope(
426384
frame_indices=indices_latents_history_4x, height=height, width=width, device=latents_history_4x.device
427385
)
428-
image_rotary_emb_history_4x = _pad_for_3d_conv(image_rotary_emb_history_4x, (4, 4, 4))
429-
image_rotary_emb_history_4x = _center_down_sample_3d(image_rotary_emb_history_4x, (4, 4, 4))
430-
image_rotary_emb_history_4x = image_rotary_emb_history_4x.flatten(2).transpose(1, 2)
431-
image_rotary_emb = torch.cat([image_rotary_emb_history_4x, image_rotary_emb], dim=1)
432-
433-
return hidden_states, image_rotary_emb.squeeze(0).chunk(2, dim=-1)
434-
435-
# def _pack_history_states(
436-
# self,
437-
# hidden_states: torch.Tensor,
438-
# indices_latents: torch.Tensor,
439-
# latents_clean: Optional[torch.Tensor] = None,
440-
# latents_history_2x: Optional[torch.Tensor] = None,
441-
# latents_history_4x: Optional[torch.Tensor] = None,
442-
# indices_latents_clean: Optional[torch.Tensor] = None,
443-
# indices_latents_history_2x: Optional[torch.Tensor] = None,
444-
# indices_latents_history_4x: Optional[torch.Tensor] = None,
445-
# ):
446-
# batch_size, num_channels, num_frames, height, width = hidden_states.shape
447-
# if indices_latents is None:
448-
# indices_latents = torch.arange(0, num_frames).unsqueeze(0).expand(batch_size, -1)
449-
450-
# hidden_states = self.x_embedder(hidden_states)
451-
# image_rotary_emb = self.rope(
452-
# frame_indices=indices_latents, height=height, width=width, device=hidden_states.device
453-
# )
454-
# image_rotary_emb = list(image_rotary_emb) # convert tuple to list for in-place modification
455-
# pph, ppw = height // self.config.patch_size, width // self.config.patch_size
456-
457-
# latents_clean, latents_history_2x, latents_history_4x = self.clean_x_embedder(
458-
# latents_clean, latents_history_2x, latents_history_4x
459-
# )
460-
461-
# if latents_clean is not None:
462-
# hidden_states = torch.cat([latents_clean, hidden_states], dim=1)
463-
464-
# image_rotary_emb_clean = self.rope(
465-
# frame_indices=indices_latents_clean, height=height, width=width, device=latents_clean.device
466-
# )
467-
# image_rotary_emb[0] = torch.cat([image_rotary_emb_clean[0], image_rotary_emb[0]], dim=0)
468-
# image_rotary_emb[1] = torch.cat([image_rotary_emb_clean[1], image_rotary_emb[1]], dim=0)
469-
470-
# if latents_history_2x is not None and indices_latents_history_2x is not None:
471-
# hidden_states = torch.cat([latents_history_2x, hidden_states], dim=1)
472-
473-
# image_rotary_emb_history_2x = self.rope(
474-
# frame_indices=indices_latents_history_2x, height=height, width=width, device=latents_history_2x.device
475-
# )
476-
# image_rotary_emb_history_2x = self._pad_rotary_emb(
477-
# image_rotary_emb_history_2x, indices_latents_history_2x.size(1), pph, ppw, (2, 2, 2)
478-
# )
479-
# image_rotary_emb[0] = torch.cat([image_rotary_emb_history_2x[0], image_rotary_emb[0]], dim=0)
480-
# image_rotary_emb[1] = torch.cat([image_rotary_emb_history_2x[1], image_rotary_emb[1]], dim=0)
481-
482-
# if latents_history_4x is not None and indices_latents_history_4x is not None:
483-
# hidden_states = torch.cat([latents_history_4x, hidden_states], dim=1)
484-
485-
# image_rotary_emb_history_4x = self.rope(
486-
# frame_indices=indices_latents_history_4x, height=height, width=width, device=latents_history_4x.device
487-
# )
488-
# image_rotary_emb_history_4x = self._pad_rotary_emb(
489-
# image_rotary_emb_history_4x, indices_latents_history_4x.size(1), pph, ppw, (4, 4, 4)
490-
# )
491-
# image_rotary_emb[0] = torch.cat([image_rotary_emb_history_4x[0], image_rotary_emb[0]], dim=0)
492-
# image_rotary_emb[1] = torch.cat([image_rotary_emb_history_4x[1], image_rotary_emb[1]], dim=0)
493-
494-
# return hidden_states, image_rotary_emb
386+
image_rotary_emb_history_4x = self._pad_rotary_emb(
387+
image_rotary_emb_history_4x, indices_latents_history_4x.size(0), pph, ppw, (4, 4, 4)
388+
)
389+
image_rotary_emb[0] = torch.cat([image_rotary_emb_history_4x[0], image_rotary_emb[0]], dim=0)
390+
image_rotary_emb[1] = torch.cat([image_rotary_emb_history_4x[1], image_rotary_emb[1]], dim=0)
391+
392+
return hidden_states, image_rotary_emb
495393

496394
def _pad_rotary_emb(
497395
self,

src/diffusers/pipelines/hunyuan_video/pipeline_hunyuan_video_framepack.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -768,17 +768,17 @@ def __call__(
768768
is_last_section = latent_paddings[k] == 0
769769
latent_padding_size = latent_paddings[k] * latent_window_size
770770

771-
indices = torch.arange(0, sum([1, latent_padding_size, latent_window_size, *history_sizes])).unsqueeze(0)
771+
indices = torch.arange(0, sum([1, latent_padding_size, latent_window_size, *history_sizes]))
772772
(
773773
indices_prefix,
774774
indices_padding,
775775
indices_latents,
776776
indices_postfix,
777777
indices_latents_history_2x,
778778
indices_latents_history_4x,
779-
) = indices.split([1, latent_padding_size, latent_window_size, *history_sizes], dim=1)
779+
) = indices.split([1, latent_padding_size, latent_window_size, *history_sizes], dim=0)
780780
# Inverted anti-drifting sampling: Figure 2(c) in the paper
781-
indices_clean_latents = torch.cat([indices_prefix, indices_postfix], dim=1)
781+
indices_clean_latents = torch.cat([indices_prefix, indices_postfix], dim=0)
782782

783783
latents_prefix = image_latents
784784
latents_postfix, latents_history_2x, latents_history_4x = history_latents[
@@ -883,7 +883,7 @@ def __call__(
883883

884884
if XLA_AVAILABLE:
885885
xm.mark_step()
886-
886+
887887
if is_last_section:
888888
latents = torch.cat([image_latents, latents], dim=2)
889889

0 commit comments

Comments
 (0)