Make prompt extender use sampling and add seed option
This commit is contained in:
@@ -268,6 +268,7 @@ of the original Wan templates or a custom system prompt.
|
||||
device=device,
|
||||
force_offload=False,
|
||||
custom_system_prompt=extender_args["system_prompt"],
|
||||
seed=extender_args["seed"]
|
||||
)
|
||||
log.info(f"Extended positive prompt: {positive_prompt}")
|
||||
_extender_cache[extender_key] = positive_prompt
|
||||
@@ -1540,6 +1541,9 @@ class WanVideoSampler:
|
||||
|
||||
steps = len(timesteps)
|
||||
|
||||
if end_step != -1 and start_step >= end_step:
|
||||
raise ValueError("start_step must be less than end_step")
|
||||
|
||||
if denoise_strength < 1.0:
|
||||
if start_step != 0:
|
||||
raise ValueError("start_step must be 0 when denoise_strength is used")
|
||||
@@ -1742,9 +1746,6 @@ class WanVideoSampler:
|
||||
|
||||
latent_video_length = noise.shape[1]
|
||||
|
||||
#if noise.shape[2] % (vae_upscale_factor/4) != 0 or noise.shape[3] % (vae_upscale_factor/4) != 0:
|
||||
# raise ValueError(f"Width ({noise.shape[3] * vae_upscale_factor}) and height ({noise.shape[2] * vae_upscale_factor}) must be divisible by {vae_upscale_factor*2}. Got {noise.shape[3] * vae_upscale_factor}x{noise.shape[2] * vae_upscale_factor}.")
|
||||
|
||||
# Initialize FreeInit filter if enabled
|
||||
freq_filter = None
|
||||
if freeinit_args is not None:
|
||||
@@ -2466,11 +2467,22 @@ class WanVideoSampler:
|
||||
current_latent = latent
|
||||
|
||||
for iter_idx in range(iterations):
|
||||
|
||||
# FreeInit noise reinitialization (after first iteration)
|
||||
if freeinit_args is not None and iter_idx > 0:
|
||||
# restart scheduler for each iteration
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
|
||||
# Re-apply start_step and end_step logic to timesteps and sigmas
|
||||
if end_step != -1:
|
||||
timesteps = timesteps[:end_step]
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
|
||||
if start_step > 0:
|
||||
timesteps = timesteps[start_step:]
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[start_step:]
|
||||
if hasattr(sample_scheduler, 'timesteps'):
|
||||
sample_scheduler.timesteps = timesteps
|
||||
|
||||
# Diffuse current latent to t=999
|
||||
diffuse_timesteps = torch.full((noise.shape[0],), 999, device=device, dtype=torch.long)
|
||||
z_T = add_noise(
|
||||
|
||||
+16
-8
@@ -110,14 +110,15 @@ class WanVideoPromptExtender:
|
||||
},
|
||||
"optional": {
|
||||
"system_prompt": (SYSTEM_PROMPT_KEYS, {"tooltip": "System prompt to use for the model."}),
|
||||
"custom_system_prompt": ("STRING", {"default": "", "forceInput": True, "tooltip": "Custom system prompt to use instead of the predefined ones."})
|
||||
"custom_system_prompt": ("STRING", {"default": "", "forceInput": True, "tooltip": "Custom system prompt to use instead of the predefined ones."}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "generate"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def generate(self, qwen, prompt, device, force_offload, max_new_tokens, system_prompt=None, custom_system_prompt=None):
|
||||
def generate(self, qwen, prompt, device, force_offload, max_new_tokens, system_prompt=None, custom_system_prompt=None, seed=0):
|
||||
if device == "gpu":
|
||||
device = mm.get_torch_device()
|
||||
elif device == "cpu":
|
||||
@@ -135,14 +136,19 @@ class WanVideoPromptExtender:
|
||||
text = qwen.tokenizer.apply_chat_template(
|
||||
messages,
|
||||
tokenize=False,
|
||||
add_generation_prompt=True
|
||||
add_generation_prompt=True,
|
||||
)
|
||||
model_inputs = qwen.tokenizer([text], return_tensors="pt").to(device)
|
||||
|
||||
torch.manual_seed(seed)
|
||||
qwen.model.to(device)
|
||||
generated_ids = qwen.model.generate(
|
||||
**model_inputs,
|
||||
max_new_tokens=max_new_tokens
|
||||
max_new_tokens=max_new_tokens,
|
||||
do_sample=True,
|
||||
temperature=0.7,
|
||||
top_p=0.8,
|
||||
top_k=20,
|
||||
repetition_penalty=1.05,
|
||||
)
|
||||
if force_offload:
|
||||
qwen.model.to(offload_device)
|
||||
@@ -163,7 +169,8 @@ class WanVideoPromptExtenderSelect:
|
||||
"system_prompt": (SYSTEM_PROMPT_KEYS, {"tooltip": "System prompt to use for the model."}),
|
||||
},
|
||||
"optional": {
|
||||
"custom_system_prompt": ("STRING", {"default": "", "forceInput": True, "tooltip": "Custom system prompt to use instead of the predefined ones."})
|
||||
"custom_system_prompt": ("STRING", {"default": "", "forceInput": True, "tooltip": "Custom system prompt to use instead of the predefined ones."}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("WANVIDEOPROMPTEXTENDER_ARGS",)
|
||||
@@ -171,7 +178,7 @@ class WanVideoPromptExtenderSelect:
|
||||
FUNCTION = "set"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def set(self, model, system_prompt, max_new_tokens, custom_system_prompt=None):
|
||||
def set(self, model, system_prompt, max_new_tokens, custom_system_prompt=None, seed=0):
|
||||
|
||||
if custom_system_prompt is None:
|
||||
sys_prompt = next((item["prompt"] for item in SYSTEM_PROMPT_MAP if item["label"] == system_prompt), "")
|
||||
@@ -183,7 +190,8 @@ class WanVideoPromptExtenderSelect:
|
||||
"system_prompt": sys_prompt,
|
||||
"max_new_tokens": max_new_tokens,
|
||||
"device": "gpu",
|
||||
"force_offload": True
|
||||
"force_offload": True,
|
||||
"seed": seed
|
||||
}
|
||||
|
||||
return (extender_settings,)
|
||||
|
||||
Reference in New Issue
Block a user