fix
This commit is contained in:
@@ -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":
|
||||
try:
|
||||
saved_state_dict = torch.load("sparge_wan.pt")
|
||||
for key in saved_state_dict.keys():
|
||||
print(key)
|
||||
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}")
|
||||
|
||||
|
||||
@@ -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',
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user