better preview

This commit is contained in:
kijai
2025-02-27 11:13:30 +02:00
parent c6e6af841d
commit 9e957121be
3 changed files with 15 additions and 9 deletions
+9 -4
View File
@@ -102,14 +102,19 @@ class WanVideoModel(comfy.model_base.BaseModel):
def __setitem__(self, k, v):
self.pipeline[k] = v
from comfy.latent_formats import LatentFormat
try:
from comfy.latent_formats import Wan21
latent_format = Wan21
except: #for backwards compatibility
log.warning("Wan21 latent format not found, update ComfyUI for better livepreview")
from comfy.latent_formats import HunyuanVideo
latent_format = HunyuanVideo
class WanVideoModelConfig:
def __init__(self, dtype):
self.unet_config = {}
self.unet_extra_config = {}
self.latent_format = comfy.latent_formats.HunyuanVideo #todo better values
self.latent_format = latent_format
self.latent_format.latent_channels = 16
self.manual_cast_dtype = dtype
self.sampling_settings = {"multiplier": 1.0}
@@ -960,7 +965,7 @@ class WanVideoSampler:
mm.soft_empty_cache()
gc.collect()
try:
torch.cuda.reset_peak_memory_stats(device)
except:
+4 -4
View File
@@ -173,10 +173,10 @@ def attention(
version=fa_version,
)
elif attention_mode == 'sdpa':
if q_lens is not None or k_lens is not None:
warnings.warn(
'Padding mask is disabled when using scaled_dot_product_attention. It can have a significant impact on performance.'
)
# if q_lens is not None or k_lens is not None:
# warnings.warn(
# 'Padding mask is disabled when using scaled_dot_product_attention. It can have a significant impact on performance.'
# )
attn_mask = None
q = q.transpose(1, 2).to(dtype)
+2 -1
View File
@@ -273,7 +273,8 @@ class WanAttentionBlock(nn.Module):
num_heads,
(-1, -1),
qk_norm,
eps)
eps,#attention_mode=attention_mode sageattn doesn't seem faster here
)
self.norm2 = WanLayerNorm(dim, eps)
self.ffn = nn.Sequential(
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),