diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 03ba53f..5f58f29 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -597,9 +597,12 @@ class HunyuanVideoPipeline(DiffusionPipeline): freqs_cos, freqs_sin = get_rotary_pos_embed( self.transformer, latent_video_length, height, width ) - - freqs_cos = freqs_cos.to(self.base_dtype).to(device) - freqs_sin = freqs_sin.to(self.base_dtype).to(device) + if not self.transformer.upcast_rope: + freqs_cos = freqs_cos.to(self.base_dtype).to(device) + freqs_sin = freqs_sin.to(self.base_dtype).to(device) + else: + freqs_cos = freqs_cos.to(device) + freqs_sin = freqs_sin.to(device) # 5. Prepare latent variables diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index e6d5c3f..ad4bbdf 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -196,6 +196,7 @@ class MMDoubleStreamBlock(nn.Module): max_seqlen_kv: Optional[int] = None, freqs_cis: tuple = None, attn_mask: Optional[torch.Tensor] = None, + upcast_rope: bool = True, ) -> Tuple[torch.Tensor, torch.Tensor]: ( img_mod1_shift, @@ -229,7 +230,7 @@ class MMDoubleStreamBlock(nn.Module): # Apply RoPE if needed. if freqs_cis is not None: - img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False) + img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope) # Prepare txt for attention. txt_modulated = self.txt_norm1(txt) @@ -380,6 +381,7 @@ class MMSingleStreamBlock(nn.Module): max_seqlen_kv: Optional[int] = None, freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None, attn_mask: Optional[torch.Tensor] = None, + upcast_rope: bool = True, stg_mode: Optional[str] = None, ) -> torch.Tensor: mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1) @@ -398,7 +400,7 @@ class MMSingleStreamBlock(nn.Module): if freqs_cis is not None: img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :] img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :] - img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False) + img_q, img_k = apply_rotary_emb(img_q, img_k, freqs_cis, upcast=upcast_rope) # assert ( # img_qq.shape == img_q.shape and img_kk.shape == img_k.shape # ), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}" @@ -467,6 +469,7 @@ class MMSingleStreamBlock(nn.Module): ) if is_enhance_enabled_single(): attn *= feta_scores + #attn[:, :-txt_len, :] *= feta_scores # Compute activation in mlp stream, cat again and run second linear layer. output = self.linear2(torch.cat((attn, self.mlp_act(mlp)), 2)) @@ -672,6 +675,9 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): get_activation_layer("silu"), **factory_kwargs, ) + + self.upcast_rope = True + #init block swap variables self.double_blocks_to_swap = -1 self.single_blocks_to_swap = -1 @@ -984,7 +990,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None - block_args = [cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv, freqs_cis, attn_mask] + block_args = [cu_seqlens_q, cu_seqlens_kv, max_seqlen_q, max_seqlen_kv, freqs_cis, attn_mask, self.upcast_rope] #tea_cache if self.enable_teacache: diff --git a/hyvideo/modules/posemb_layers.py b/hyvideo/modules/posemb_layers.py index dfce82c..fa42dbe 100644 --- a/hyvideo/modules/posemb_layers.py +++ b/hyvideo/modules/posemb_layers.py @@ -61,87 +61,18 @@ def get_meshgrid_nd(start, *args, dim=2): ################################################################################# # https://github.com/meta-llama/llama/blob/be327c427cc5e89cc1d3ab3d3fec4484df771245/llama/model.py#L80 - -def reshape_for_broadcast( - freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor]], - x: torch.Tensor, - head_first=False, -): - """ - Reshape frequency tensor for broadcasting it with another tensor. - - This function reshapes the frequency tensor to have the same shape as the target tensor 'x' - for the purpose of broadcasting the frequency tensor during element-wise operations. - - Notes: - When using FlashMHAModified, head_first should be False. - When using Attention, head_first should be True. - - Args: - freqs_cis (Union[torch.Tensor, Tuple[torch.Tensor]]): Frequency tensor to be reshaped. - x (torch.Tensor): Target tensor for broadcasting compatibility. - head_first (bool): head dimension first (except batch dim) or not. - - Returns: - torch.Tensor: Reshaped frequency tensor. - - Raises: - AssertionError: If the frequency tensor doesn't match the expected shape. - AssertionError: If the target tensor 'x' doesn't have the expected number of dimensions. - """ - ndim = x.ndim - assert 0 <= 1 < ndim - - if isinstance(freqs_cis, tuple): - # freqs_cis: (cos, sin) in real space - if head_first: - assert freqs_cis[0].shape == ( - x.shape[-2], - x.shape[-1], - ), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}" - shape = [ - d if i == ndim - 2 or i == ndim - 1 else 1 - for i, d in enumerate(x.shape) - ] - else: - assert freqs_cis[0].shape == ( - x.shape[1], - x.shape[-1], - ), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}" - shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)] - return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape) - else: - # freqs_cis: values in complex space - if head_first: - assert freqs_cis.shape == ( - x.shape[-2], - x.shape[-1], - ), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}" - shape = [ - d if i == ndim - 2 or i == ndim - 1 else 1 - for i, d in enumerate(x.shape) - ] - else: - assert freqs_cis.shape == ( - x.shape[1], - x.shape[-1], - ), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}" - shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)] - return freqs_cis.view(*shape) - - -def rotate_half(x): - x_real, x_imag = ( - x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1) - ) # [B, S, H, D//2] - return torch.stack([-x_imag, x_real], dim=-1).flatten(3) - + +def apply_rotary(x, cos, sin): + x_reshaped = x.view(*x.shape[:-1], -1, 2) + x1, x2 = x_reshaped.unbind(-1) + x_rotated = torch.stack([-x2, x1], dim=-1).flatten(3) + return (x * cos) + (x_rotated * sin) def apply_rotary_emb( xq: torch.Tensor, xk: torch.Tensor, freqs_cis: Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]], - head_first: bool = False, + upcast: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: """ Apply rotary embeddings to input tensors using the given frequency tensor. @@ -155,35 +86,20 @@ def apply_rotary_emb( xq (torch.Tensor): Query tensor to apply rotary embeddings. [B, S, H, D] xk (torch.Tensor): Key tensor to apply rotary embeddings. [B, S, H, D] freqs_cis (torch.Tensor or tuple): Precomputed frequency tensor for complex exponential. - head_first (bool): head dimension first (except batch dim) or not. Returns: Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings. """ - xk_out = None - if isinstance(freqs_cis, tuple): - cos, sin = reshape_for_broadcast(freqs_cis, xq, head_first) # [S, D] - cos, sin = cos.to(xq.device), sin.to(xq.device) - # real * cos - imag * sin - # imag * cos + real * sin - xq_out = (xq.float() * cos + rotate_half(xq.float()) * sin).type_as(xq) - xk_out = (xk.float() * cos + rotate_half(xk.float()) * sin).type_as(xk) + + cos, sin = [f.view(*xq.shape[:2], 1, xq.shape[3]) for f in freqs_cis] + + if upcast: + xq_out = apply_rotary(xq.float(), cos, sin).to(xq.dtype) + xk_out = apply_rotary(xk.float(), cos, sin).to(xk.dtype) else: - # view_as_complex will pack [..., D/2, 2](real) to [..., D/2](complex) - xq_ = torch.view_as_complex( - xq.float().reshape(*xq.shape[:-1], -1, 2) - ) # [B, S, H, D//2] - freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to( - xq.device - ) # [S, D//2] --> [1, S, 1, D//2] - # (real, imag) * (cos, sin) = (real * cos - imag * sin, imag * cos + real * sin) - # view_as_real will expand [..., D/2](complex) to [..., D/2, 2](real) - xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3).type_as(xq) - xk_ = torch.view_as_complex( - xk.float().reshape(*xk.shape[:-1], -1, 2) - ) # [B, S, H, D//2] - xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk) + xq_out = apply_rotary(xq, cos, sin) + xk_out = apply_rotary(xk, cos, sin) return xq_out, xk_out diff --git a/nodes.py b/nodes.py index 15deec4..f3f0383 100644 --- a/nodes.py +++ b/nodes.py @@ -275,6 +275,7 @@ class HyVideoModelLoader: "block_swap_args": ("BLOCKSWAPARGS", ), "lora": ("HYVIDLORA", {"default": None}), "auto_cpu_offload": ("BOOLEAN", {"default": False, "tooltip": "Enable auto offloading for reduced VRAM usage, implementation from DiffSynth-Studio, slightly different from block swapping and uses even less VRAM, but can be slower as you can't define how much VRAM to use"}), + "upcast_rope": ("BOOLEAN", {"default": True, "tooltip": "Upcast RoPE to fp32 for better accuracy, this is the default behaviour, disabling can improve speed and reduce memory use slightly"}), } } @@ -284,7 +285,7 @@ class HyVideoModelLoader: CATEGORY = "HunyuanVideoWrapper" def loadmodel(self, model, base_precision, load_device, quantization, - compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, auto_cpu_offload=False): + compile_args=None, attention_mode="sdpa", block_swap_args=None, lora=None, auto_cpu_offload=False, upcast_rope=True): transformer = None #mm.unload_all_models() mm.soft_empty_cache() @@ -328,6 +329,8 @@ class HyVideoModelLoader: ) transformer.eval() + transformer.upcast_rope = upcast_rope + comfy_model = HyVideoModel( HyVideoModelConfig(base_dtype), model_type=comfy.model_base.ModelType.FLOW,