improve handling of wan22 stages

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-07-30 11:22:08 -04:00
parent c36e8efd38
commit d8e03bb855
6 changed files with 31 additions and 12 deletions
+6
View File
@@ -2,7 +2,13 @@
## Update for 2025-07-30
- **Feature**
- Wan select which stage to run: *first/second/both* with configurable *boundary ration* when running both stages
in settings -> model options
- **UI**
- modernui checkbox/radio styling
- **Fixes**
- fix Wan2.2 5B I2V workflow
- fix inpaint image metadata
- fix processing image save loop
- fix api progress reporting endpoint
+1 -1
View File
@@ -593,7 +593,7 @@ def check_diffusers():
t_start = time.time()
if args.skip_all or args.skip_git:
return
sha = '56d438727036b0918b30bbe3110c5fe1634ed19d' # diffusers commit hash
sha = 'c052791b5fe29ce8a308bf63dda97aa205b729be' # diffusers commit hash
pkg = pkg_resources.working_set.by_key.get('diffusers', None)
minor = int(pkg.version.split('.')[1] if pkg is not None else -1)
cur = opts.get('diffusers_version', '') if minor > -1 else ''
+1
View File
@@ -9,6 +9,7 @@ def interrogate(image):
if isinstance(image, dict) and 'name' in image:
image = Image.open(image['name'])
if image is None:
shared.log.error('Interrogate: no image provided')
return ''
t0 = time.time()
if shared.opts.interrogate_default_type == 'OpenCLiP':
+2 -1
View File
@@ -200,7 +200,8 @@ options_templates.update(options_section(('model_options', "Models Options"), {
"model_h1_sep": OptionInfo("<h2>HiDream</h2>", "", gr.HTML),
"model_h1_llama_repo": OptionInfo("Default", "LLama repo", gr.Textbox),
"model_wan_sep": OptionInfo("<h2>WanAI</h2>", "", gr.HTML),
"model_wan_disable_t2": OptionInfo(True, "Disable second stage"),
"model_wan_stage": OptionInfo("first", "Processing stage", gr.Radio, {"choices": ['first', 'second', 'both'] }),
"model_wan_boundary": OptionInfo(0.85, "Stage boundary ratio", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.05 }),
}))
options_templates.update(options_section(('vae_encoder', "Variational Auto Encoder"), {
+20 -9
View File
@@ -8,11 +8,6 @@ def load_transformer(repo_id, diffusers_load_config={}, subfolder='transformer')
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True)
fn = None
if subfolder == 'transformer_2' and 'a14b' not in repo_id.lower():
return None
if subfolder == 'transformer_2' and shared.opts.model_wan_disable_t2:
return None
if shared.opts.sd_unet is not None and shared.opts.sd_unet != 'Default':
from modules import sd_unet
if shared.opts.sd_unet not in list(sd_unet.unet_dict):
@@ -63,13 +58,28 @@ def load_wan(checkpoint_info, diffusers_load_config={}):
repo_id = sd_models.path_to_repo(checkpoint_info)
sd_models.hf_auth_check(checkpoint_info)
transformer = load_transformer(repo_id, diffusers_load_config, 'transformer')
transformer_2 = load_transformer(repo_id, diffusers_load_config, 'transformer_2')
if 'a14b' in repo_id.lower():
if shared.opts.model_wan_stage == 'first':
transformer = load_transformer(repo_id, diffusers_load_config, 'transformer')
transformer_2 = None
elif shared.opts.model_wan_stage == 'second':
transformer = load_transformer(repo_id, diffusers_load_config, 'transformer_2')
transformer_2 = None
elif shared.opts.model_wan_stage == 'both':
transformer = load_transformer(repo_id, diffusers_load_config, 'transformer')
transformer_2 = load_transformer(repo_id, diffusers_load_config, 'transformer_2')
else:
shared.log.error(f'Load model: type=WanAI stage="{shared.opts.model_wan_stage}" unsupported')
return None
else:
transformer = load_transformer(repo_id, diffusers_load_config, 'transformer')
transformer_2 = None
text_encoder = load_text_encoder(repo_id, diffusers_load_config)
load_args, _quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model')
boundary_ratio = 0.8 if transformer_2 is not None else None
shared.log.debug(f'Load model: type=WanAI model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args} boundary={boundary_ratio}')
boundary_ratio = shared.opts.model_wan_boundary if transformer_2 is not None else None
shared.log.debug(f'Load model: type=WanAI model="{checkpoint_info.name}" repo="{repo_id}" offload={shared.opts.diffusers_offload_mode} dtype={devices.dtype} args={load_args} stage={shared.opts.model_wan_stage} boundary={boundary_ratio}')
cls = diffusers.WanPipeline
pipe = cls.from_pretrained(
@@ -88,6 +98,7 @@ def load_wan(checkpoint_info, diffusers_load_config={}):
del text_encoder
del transformer
del transformer_2
sd_hijack_te.init_hijack(pipe)
from modules.video_models import video_vae