From 1061378b664513e04dd21c598bfde8a947027d72 Mon Sep 17 00:00:00 2001 From: Bubbliiiing <47347516+bubbliiiing@users.noreply.github.com> Date: Thu, 29 May 2025 11:28:00 +0800 Subject: [PATCH] Add bfloat16 norm && fix bug in cfg_optimization (#217) --- videox_fun/dist/fuser.py | 2 +- videox_fun/models/wan_transformer3d.py | 11 +++++------ videox_fun/utils/cfg_optimization.py | 6 +++--- 3 files changed, 9 insertions(+), 10 deletions(-) mode change 100644 => 100755 videox_fun/dist/fuser.py mode change 100644 => 100755 videox_fun/utils/cfg_optimization.py diff --git a/videox_fun/dist/fuser.py b/videox_fun/dist/fuser.py old mode 100644 new mode 100755 index bb437a7..de12a2b --- a/videox_fun/dist/fuser.py +++ b/videox_fun/dist/fuser.py @@ -13,7 +13,7 @@ try: initialize_model_parallel) from pai_fuser.core.long_ctx_attention import \ xFuserLongContextAttention - print("Enable PAI DiT Turbo") + print("Import PAI DiT Turbo") else: import xfuser from xfuser.core.distributed import (get_sequence_parallel_rank, diff --git a/videox_fun/models/wan_transformer3d.py b/videox_fun/models/wan_transformer3d.py index 2bc5f3d..41e6c07 100755 --- a/videox_fun/models/wan_transformer3d.py +++ b/videox_fun/models/wan_transformer3d.py @@ -367,7 +367,7 @@ def rope_apply(x, grid_sizes, freqs): # append to collection output.append(x_i) - return torch.stack(output).float() + return torch.stack(output).to(x.dtype) def rope_apply_qk(q, k, grid_sizes, freqs): @@ -389,10 +389,10 @@ class WanRMSNorm(nn.Module): Args: x(Tensor): Shape [B, L, C] """ - return self._norm(x.float()).type_as(x) * self.weight + return self._norm(x) * self.weight def _norm(self, x): - return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) + return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps).to(x.dtype) class WanLayerNorm(nn.LayerNorm): @@ -405,8 +405,7 @@ class WanLayerNorm(nn.LayerNorm): Args: x(Tensor): Shape [B, L, C] """ - return super().forward(x.float()).type_as(x) - + return super().forward(x) class WanSelfAttention(nn.Module): @@ -1132,7 +1131,7 @@ class WanTransformer3DModel(ModelMixin, ConfigMixin, FromOriginalModelMixin): # unpatchify x = self.unpatchify(x, grid_sizes) x = torch.stack(x) - if self.teacache is not None: + if self.teacache is not None and cond_flag: self.teacache.cnt += 1 if self.teacache.cnt == self.teacache.num_steps: self.teacache.reset() diff --git a/videox_fun/utils/cfg_optimization.py b/videox_fun/utils/cfg_optimization.py old mode 100644 new mode 100755 index 2d409cc..344a2ef --- a/videox_fun/utils/cfg_optimization.py +++ b/videox_fun/utils/cfg_optimization.py @@ -5,8 +5,8 @@ import torch def cfg_skip(): def decorator(func): def wrapper(self, x, *args, **kwargs): - if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio): - bs = len(x) + bs = len(x) + if bs >= 2 and self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio): bs_half = int(bs // 2) new_x = x[bs_half:] @@ -31,7 +31,7 @@ def cfg_skip(): result = func(self, new_x, *new_args, **new_kwargs) - if self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio): + if bs >= 2 and self.cfg_skip_ratio is not None and self.current_steps >= self.num_inference_steps * (1 - self.cfg_skip_ratio): result = torch.cat([result, result], dim=0) return result