Merge branch 'main' into vap

This commit is contained in:
kijai
2025-11-04 10:39:35 +02:00
9 changed files with 146 additions and 104 deletions
+42 -28
View File
@@ -3,15 +3,20 @@ import torch.nn as nn
from accelerate import init_empty_weights
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None):
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None, compile_args=None):
has_children = list(model.children())
if not has_children:
return
allow_compile = False
for name, module in model.named_children():
if compile_args is not None:
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
module_prefix = prefix + name + "."
module_prefix = module_prefix.replace("_orig_mod.", "")
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights)
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args)
if isinstance(module, nn.Linear) and "loras" not in module_prefix:
in_features = state_dict[module_prefix + "weight"].shape[1]
@@ -25,7 +30,8 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
out_features,
module.bias is not None,
compute_dtype=compute_dtype,
scale_weight=scale_weights.get(scale_key) if scale_weights else None
scale_weight=scale_weights.get(scale_key) if scale_weights else None,
allow_compile=allow_compile
)
model._modules[name].source_cls = type(module)
model._modules[name].requires_grad_(False)
@@ -77,7 +83,8 @@ class CustomLinear(nn.Linear):
bias=False,
compute_dtype=None,
device=None,
scale_weight=None
scale_weight=None,
allow_compile=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
@@ -85,38 +92,28 @@ class CustomLinear(nn.Linear):
self.step = 0
self.scale_weight = scale_weight
self.lora_strengths = []
self.allow_compile = allow_compile
if not allow_compile:
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
self.lora_diffs = []
for i, diff in enumerate(lora_diffs):
if isinstance(diff, tuple):
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device))
if len(diff) > 1:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
setattr(self, f"lora_diff_{i}_2", diff[2])
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
else:
self.register_buffer(f"lora_diff_{i}", diff.to(device))
self.lora_diffs.append(f"lora_diff_{i}")
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.lora_diffs.append(f"lora_diff_{i}_0")
def forward(self, input):
if self.bias is not None:
bias = self.bias.to(input)
else:
bias = None
weight = self.weight.to(input)
if self.scale_weight is not None:
if weight.numel() < input.numel():
weight = weight * self.scale_weight
else:
input = input * self.scale_weight
if hasattr(self, f"lora_diff_0_0"):
weight = self.apply_lora(weight).to(self.compute_dtype)
return torch.nn.functional.linear(input, weight, bias)
def apply_lora(self, weight):
def _get_weight_with_lora(self, weight):
"""Apply LoRA outside compiled region"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
if isinstance(lora_strength, list):
lora_strength = lora_strength[self.step]
@@ -139,6 +136,23 @@ class CustomLinear(nn.Linear):
lora_diff = getattr(self, lora_diff_names)
weight = weight.add(lora_diff, alpha=lora_strength)
return weight
def forward(self, input):
if self.bias is not None:
bias = self.bias.to(input)
else:
bias = None
weight = self.weight.to(input)
if self.scale_weight is not None:
if weight.numel() < input.numel():
weight = weight * self.scale_weight
else:
input = input * self.scale_weight
weight = self._get_weight_with_lora(weight)
return torch.nn.functional.linear(input, weight, bias)
def remove_lora_from_module(module):
for name, submodule in module.named_modules():
+24 -11
View File
@@ -19,7 +19,7 @@ def load_gguf(model_path):
return parsed_parameters, reader
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None):
def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modules_to_not_convert=[], patches=None, compile_args=None):
def _should_convert_to_gguf(state_dict, prefix):
weight_key = prefix + "weight"
return weight_key in state_dict and isinstance(state_dict[weight_key], GGUFParameter)
@@ -27,10 +27,14 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
has_children = list(model.children())
if not has_children:
return
allow_compile = False
for name, module in model.named_children():
if compile_args is not None:
allow_compile = compile_args.get("allow_unmerged_lora_compile", False)
module_prefix = prefix + name + "."
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches)
_replace_with_gguf_linear(module, compute_dtype, state_dict, module_prefix, modules_to_not_convert, patches, compile_args)
if (
isinstance(module, nn.Linear)
@@ -46,7 +50,8 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
in_features,
out_features,
module.bias is not None,
compute_dtype=compute_dtype
compute_dtype=compute_dtype,
allow_compile=allow_compile
)
model._modules[name].source_cls = type(module)
@@ -95,35 +100,43 @@ class GGUFLinear(nn.Linear):
bias=False,
compute_dtype=None,
device=None,
allow_compile=False
) -> None:
super().__init__(in_features, out_features, bias, device)
self.compute_dtype = compute_dtype
self.lora_diffs = []
self.lora_strengths = []
self.step = 0
self.allow_compile = allow_compile
if not allow_compile:
self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora)
def forward(self, inputs):
weight = dequantize_gguf_tensor(self.weight).to(self.compute_dtype)
bias = self.bias.to(self.compute_dtype) if self.bias is not None else None
if hasattr(self, f"lora_diff_0_0"):
weight = self.apply_lora(weight).to(self.compute_dtype)
weight = self._get_weight_with_lora(weight)#.to(self.compute_dtype)
return torch.nn.functional.linear(inputs, weight, bias)
def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")):
self.lora_diffs = []
for i, diff in enumerate(lora_diffs):
if isinstance(diff, tuple):
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device))
if len(diff) > 1:
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.register_buffer(f"lora_diff_{i}_1", diff[1].to(device, self.compute_dtype))
setattr(self, f"lora_diff_{i}_2", diff[2])
self.lora_diffs.append((f"lora_diff_{i}_0", f"lora_diff_{i}_1", f"lora_diff_{i}_2"))
else:
self.register_buffer(f"lora_diff_{i}", diff.to(device))
self.lora_diffs.append(f"lora_diff_{i}")
self.register_buffer(f"lora_diff_{i}_0", diff[0].to(device, self.compute_dtype))
self.lora_diffs.append(f"lora_diff_{i}_0")
def apply_lora(self, weight):
def _get_weight_with_lora(self, weight):
"""Apply LoRA outside compiled region"""
if not hasattr(self, "lora_diff_0_0"):
return weight
for lora_diff_names, lora_strength in zip(self.lora_diffs, self.lora_strengths):
if isinstance(lora_strength, list):
lora_strength = lora_strength[self.step]
+1 -1
View File
@@ -198,7 +198,7 @@ class AudioProjModel(nn.Module):
context_tokens = self.proj3(audio_embeds_c).reshape(batch_size_c*N_t, self.context_tokens, self.output_dim)
# normalization and reshape
context_tokens = self.norm(context_tokens)
context_tokens = self.norm(context_tokens.to(self.norm.weight.dtype)).to(context_tokens.dtype)
context_tokens = rearrange(context_tokens, "(bz f) m c -> bz f m c", f=video_length)
return context_tokens
+7 -4
View File
@@ -330,6 +330,7 @@ class WanVideoTorchCompileSettings:
"optional": {
"dynamo_recompile_limit": ("INT", {"default": 128, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.recompile_limit"}),
"force_parameter_static_shapes": ("BOOLEAN", {"default": False, "tooltip": "torch._dynamo.config.force_parameter_static_shapes"}),
"allow_unmerged_lora_compile": ("BOOLEAN", {"default": False, "tooltip": "Allow LoRA application to be compiled with torch.compile to avoid graph breaks, causes issues with some LoRAs, mostly dynamic ones"}),
},
}
RETURN_TYPES = ("WANCOMPILEARGS",)
@@ -338,7 +339,8 @@ class WanVideoTorchCompileSettings:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch > 2.7.0 is recommended"
def set_args(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_transformer_blocks_only, dynamo_recompile_limit=128, force_parameter_static_shapes=True):
def set_args(self, backend, fullgraph, mode, dynamic, dynamo_cache_size_limit, compile_transformer_blocks_only, dynamo_recompile_limit=128,
force_parameter_static_shapes=True, allow_unmerged_lora_compile=False):
compile_args = {
"backend": backend,
@@ -349,6 +351,7 @@ class WanVideoTorchCompileSettings:
"dynamo_recompile_limit": dynamo_recompile_limit,
"compile_transformer_blocks_only": compile_transformer_blocks_only,
"force_parameter_static_shapes": force_parameter_static_shapes,
"allow_unmerged_lora_compile": allow_unmerged_lora_compile,
}
return (compile_args, )
@@ -782,7 +785,7 @@ def rename_fuser_block(name):
return new_name
def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None):
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None, compile_args=None):
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "face_encoder", "fuser_block"}
param_count = sum(1 for _ in transformer.named_parameters())
@@ -837,7 +840,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
if not getattr(transformer, "gguf_patched", False):
transformer = _replace_with_gguf_linear(
transformer, base_dtype, sd, patches=patcher.patches
transformer, base_dtype, sd, patches=patcher.patches, compile_args=compile_args
)
transformer.gguf_patched = True
else:
@@ -1535,7 +1538,7 @@ class WanVideoModelLoader:
transformer.patched_linear = False
sd = None
elif "scaled" in quantization or lora is not None:
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights, compile_args=compile_args)
transformer.patched_linear = True
if "fast" in quantization:
+11 -4
View File
@@ -189,14 +189,21 @@ class WanVideoSampler:
vae_upscale_factor = 16 if is_5b else 8
# Load weights
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0:
transformer = _replace_linear(transformer, dtype, patcher.model["sd"])
if transformer.audio_model is not None:
for block in transformer.blocks:
if hasattr(block, 'audio_block'):
block.audio_block = None
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0 and gguf_reader is None:
transformer = _replace_linear(transformer, dtype, patcher.model["sd"], model["compile_args"])
transformer.patched_linear = True
if patcher.model["sd"] is not None and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device,
block_swap_args=block_swap_args, compile_args=model["compile_args"])
if gguf_reader is not None: #handle GGUF
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True,
reader=gguf_reader, block_swap_args=block_swap_args, compile_args=model["compile_args"])
set_lora_params_gguf(transformer, patcher.patches)
transformer.patched_linear = True
elif len(patcher.patches) != 0: #handle patched linear layers (unmerged loras, fp8 scaled)
+7
View File
@@ -74,6 +74,13 @@ HuMo: https://github.com/Phantom-video/HuMo
WanAnimate: https://github.com/Wan-Video/Wan2.2/tree/main/wan/modules/animate
Lynx: https://github.com/bytedance/lynx
Not exactly Wan model, but close enough to work with the code base:
LongCat-Video: https://meituan-longcat.github.io/LongCat-Video/
Examples:
---
+1 -1
View File
@@ -172,7 +172,7 @@ class WanVideoDiffusionForcingSampler:
# Load weights
if not transformer.patched_linear and patcher.model["sd"] is not None and len(patcher.patches) != 0:
transformer = _replace_linear(transformer, dtype, patcher.model["sd"])
transformer = _replace_linear(transformer, dtype, patcher.model["sd"], model["compile_args"])
transformer.patched_linear = True
if patcher.model["sd"] is not None and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
+52 -54
View File
@@ -29,14 +29,14 @@ from comfy import model_management as mm
__all__ = ['WanModel']
class AdaLayerNorm(nn.Module):
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5, dtype=None, device=None, operations=None):
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5):
super().__init__()
output_dim = output_dim or embedding_dim * 2
self.silu = nn.SiLU()
self.linear = operations.Linear(embedding_dim, output_dim, dtype=dtype, device=device)
self.norm = operations.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine, dtype=dtype, device=device)
self.linear = nn.Linear(embedding_dim, output_dim)
self.norm = nn.LayerNorm(output_dim // 2, norm_eps, norm_elementwise_affine)
def forward(self, x, temb):
temb = self.linear(self.silu(temb))
@@ -672,7 +672,7 @@ class WanT2VCrossAttention(WanSelfAttention):
if is_longcat:
if num_cond_latents is not None and num_cond_latents > 0:
num_cond_latents_thw = num_cond_latents * (s // grid_sizes[0][0])
num_cond_latents_thw = num_cond_latents * (s // num_latent_frames)
x = x[:, num_cond_latents_thw:]
q = self.norm_q(self.q(x).view(b, -1, n, d).to(self.norm_q.weight.dtype)).to(x.dtype)
else:
@@ -1042,7 +1042,7 @@ class WanAttentionBlock(nn.Module):
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
def ffn_chunked(self, x, shift_mlp, scale_mlp, num_chunks=4):
modulated_input = torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp)
modulated_input = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp).to(x.dtype)
result = torch.empty_like(x)
seq_len = modulated_input.shape[1]
@@ -1241,7 +1241,7 @@ class WanAttentionBlock(nn.Module):
full_v = torch.cat([v, v_ip], dim=1)
y = self.self_attn.forward(q, full_k, full_v, seq_lens)
elif is_longcat and num_cond_latents is not None and num_cond_latents > 0:
num_cond_latents_thw = num_cond_latents * (N // grid_sizes[0][0])
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
# process the condition tokens
x_cond = self.self_attn.forward(
q[:, :num_cond_latents_thw].contiguous(),
@@ -1338,6 +1338,7 @@ class WanAttentionBlock(nn.Module):
x_mot_ref = x_mot_ref.to(input_dtype)
# MultiTalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
x = x.add(x_audio, alpha=audio_scale)
@@ -1352,39 +1353,39 @@ class WanAttentionBlock(nn.Module):
x = self.audio_cross_attn_wrapper(x, humo_audio_input, grid_sizes, humo_audio_scale)
# ffn
if self.rope_func == "comfy_chunked":
x_ffn = self.ffn_chunked(x, shift_mlp, scale_mlp)
else:
if zero_timestep:
norm2_x = self.norm2(x.to(self.norm2.weight.dtype)).to(input_dtype)
parts = []
for i in range(2):
parts.append(norm2_x[:, self.seg_idx[i]:self.seg_idx[i + 1]] *
(1 + scale_mlp[:, i:i + 1]) + shift_mlp[:, i:i + 1])
norm2_x = torch.cat(parts, dim=1)
x_ffn = self.ffn(norm2_x)
else:
if not is_longcat:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
else:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C)
x_ffn = self.ffn(mod_x.to(input_dtype))
del shift_mlp, scale_mlp
# gate_mlp
# ffn
if self.rope_func == "comfy_chunked":
x_ffn = self.ffn_chunked(x, shift_mlp, scale_mlp)
else:
if zero_timestep:
z = []
norm2_x = self.norm2(x)
parts = []
for i in range(2):
z.append(x_ffn[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_mlp[:, i:i + 1])
x_ffn = torch.cat(z, dim=1)
x = x.add(x_ffn)
parts.append(norm2_x[:, self.seg_idx[i]:self.seg_idx[i + 1]] *
(1 + scale_mlp[:, i:i + 1]) + shift_mlp[:, i:i + 1])
norm2_x = torch.cat(parts, dim=1)
x_ffn = self.ffn(norm2_x)
else:
if not is_longcat:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
else:
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
del gate_mlp
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C)
x_ffn = self.ffn(mod_x.to(input_dtype))
del shift_mlp, scale_mlp
# gate_mlp
if zero_timestep:
z = []
for i in range(2):
z.append(x_ffn[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_mlp[:, i:i + 1])
x_ffn = torch.cat(z, dim=1)
x = x.add(x_ffn)
else:
if not is_longcat:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
else:
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
del gate_mlp
if x_ip is not None: #stand-in
x_ip = x_ip.addcmul(y_ip, gate_msa_ip)
@@ -1699,7 +1700,7 @@ class AudioInjector_WAN(nn.Module):
if enable_adain:
self.injector_adain_layers = nn.ModuleList([
AdaLayerNorm(
output_dim=dim * 2, embedding_dim=adain_dim, chunk_dim=1)
output_dim=dim * 2, embedding_dim=adain_dim)
for _ in range(audio_injector_id)
])
if need_adain_ont:
@@ -2662,7 +2663,17 @@ class WanModel(torch.nn.Module):
e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
# clip vision embedding
clip_embed = None
if clip_fea is not None and hasattr(self, "img_emb"):
clip_fea = clip_fea.to(self.main_device)
if self.offload_img_emb:
self.img_emb.to(self.main_device)
clip_embed = self.img_emb(clip_fea) # bs x 257 x dim
#context = torch.concat([context_clip, context], dim=1)
if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
#context (text embedding)
if hasattr(self, "text_embedding") and context != []:
text_embed_dtype = self.text_embedding[0].weight.dtype
@@ -2707,24 +2718,13 @@ class WanModel(torch.nn.Module):
if self.offload_txt_emb:
self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking)
seq_chunks = max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0)
chunked_self_attention = seq_chunks > 1 and current_step in self.video_attention_split_steps
else:
context = None
clip_embed = clip_embed_mot_ref = None
if clip_fea is not None and hasattr(self, "img_emb"):
clip_fea = clip_fea.to(self.main_device)
if self.offload_img_emb:
self.img_emb.to(self.main_device)
clip_embed = self.img_emb(clip_fea) # bs x 257 x dim
if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
if mot_ref_clip_embeds is not None:
mot_ref_clip_embeds = mot_ref_clip_embeds.to(self.main_device)
if self.offload_img_emb:
self.img_emb.to(self.main_device)
clip_embed_mot_ref = self.img_emb_mot_ref(mot_ref_clip_embeds) # bs x 257 x dim
if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
chunked_self_attention = False
seq_chunks = 0
# MultiTalk
if multitalk_audio is not None:
@@ -2902,8 +2902,6 @@ class WanModel(torch.nn.Module):
dwpose_emb = rearrange(unianim_data['dwpose'], 'b c f h w -> b (f h w) c').contiguous()
x.add_(dwpose_emb, alpha=unianim_data['strength'])
seq_chunks = max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0)
chunked_self_attention = seq_chunks > 1 and current_step in self.video_attention_split_steps
# arguments
kwargs = dict(
e=e0,
+1 -1
View File
@@ -60,7 +60,7 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
if flowedit_args: #seems to work better
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift))
else:
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas)
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
elif 'dpm' in scheduler:
if 'sde' in scheduler:
algorithm_type = "sde-dpmsolver++"