This commit is contained in:
kijai
2025-02-28 18:00:31 +02:00
parent b6503bc248
commit 889efe7371
3 changed files with 7 additions and 6 deletions
+6 -4
View File
@@ -1061,7 +1061,7 @@ class WanVideoSampler:
set_enhance_weight(feta_args["weight"])
feta_start_percent = feta_args["start_percent"]
feta_end_percent = feta_args["end_percent"]
set_num_frames(latent.shape[2])
set_num_frames(latent_video_length)
enable_enhance()
else:
disable_enhance()
@@ -1080,11 +1080,13 @@ class WanVideoSampler:
block.self_attn.verbose = True
block.self_attn.inner_attention = SparseAttentionMeansim(l1=0.06, pv_l1=0.065)
if transformer.attention_mode == "spargeattn":
saved_state_dict = torch.load("sparge_wan.pt")
for key in saved_state_dict.keys():
print(key)
try:
saved_state_dict = torch.load("sparge_wan.pt")
except:
raise ValueError("No saved parameters found for sparse attention, tuning is required first")
load_sparse_attention_state_dict(transformer, saved_state_dict, verbose = True)
#for idx, block in enumerate(transformer.blocks):
# print(f"Block {idx} attn1: {block}")
-2
View File
@@ -1,4 +1,3 @@
from .attention import flash_attention
from .model import WanModel
from .t5 import T5Decoder, T5Encoder, T5EncoderModel, T5Model
from .tokenizers import HuggingfaceTokenizer
@@ -10,5 +9,4 @@ __all__ = [
'T5Decoder',
'T5EncoderModel',
'HuggingfaceTokenizer',
'flash_attention',
]
+1
View File
@@ -35,6 +35,7 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=81, k=0):
1.0 / torch.pow(theta,
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
if k > 0:
print(f"RifleX: Using {k}th freq")
freqs[k-1] = 0.9 * 2 * torch.pi / L_test
freqs = torch.polar(torch.ones_like(freqs), freqs)
return freqs