Make prompt extender use sampling and add seed option

This commit is contained in:
kijai
2025-08-05 21:17:11 +03:00
parent 92eac2a9e1
commit 69dd4689bf
2 changed files with 31 additions and 11 deletions
+15 -3
View File
@@ -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
View File
@@ -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,)