Merge branch 'main' into speed

This commit is contained in:
bubbliiiing
2025-05-29 03:35:06 +00:00
3 changed files with 9 additions and 10 deletions
Vendored Regular → Executable
+1 -1
View File
@@ -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,
+5 -6
View File
@@ -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()
+3 -3
View File
@@ -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