support configurable multi-stage models in video tab

Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
Vladimir Mandic
2025-09-22 13:18:29 -04:00
parent 8b47d72610
commit 41bb446697
7 changed files with 156 additions and 30 deletions
+3 -2
View File
@@ -170,11 +170,12 @@ models = {
dit_cls=diffusers.WanTransformer3DModel,
dit_folder=("transformer", "transformer_2")),
Model(name='WAN 2.2 14B VACE',
url='https://huggingface.co/Wan-AI/Wan2.2-14B-VACE-T2V-Diffusers',
url='https://huggingface.co/linoyts/Wan2.2-VACE-Fun-14B-diffusers',
repo='linoyts/Wan2.2-VACE-Fun-14B-diffusers',
repo_cls=diffusers.WanVACEPipeline,
te_cls=transformers.T5EncoderModel,
dit_cls=diffusers.WanVACETransformer3DModel),
dit_cls=diffusers.WanVACETransformer3DModel,
dit_folder=("transformer", "transformer_2")),
Model(name='WAN 2.1 1.3B T2V',
url='https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers',
repo='Wan-AI/Wan2.1-T2V-1.3B-Diffusers',
+17 -24
View File
@@ -40,31 +40,24 @@ def load_model(selected: models_def.Model):
# transformer
try:
if selected.dit_folder is None:
selected.dit_folder = ['transformer']
if isinstance(selected.dit_folder, list) or isinstance(selected.dit_folder, tuple):
# wan a14b has transformer and transformer_2
for dit_folder in selected.dit_folder:
# get a new quant arg on every loop to prevent the quant config classes getting entangled
load_args, quant_args = model_quant.get_dit_args({}, module='Model', device_map=True)
shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{dit_folder}" cls={selected.dit_cls.__name__} quant={model_quant.get_quant_type(quant_args)}')
kwargs[dit_folder] = selected.dit_cls.from_pretrained(
pretrained_model_name_or_path=selected.dit or selected.repo,
subfolder=dit_folder,
revision=selected.dit_revision or selected.repo_revision,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args
)
else:
load_args, quant_args = model_quant.get_dit_args({}, module='Model', device_map=True)
shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" folder="{selected.dit_folder}" cls={selected.dit_cls.__name__} quant={model_quant.get_quant_type(quant_args)}')
kwargs["transformer"] = selected.dit_cls.from_pretrained(
pretrained_model_name_or_path=selected.dit or selected.repo,
subfolder=selected.dit_folder,
revision=selected.dit_revision or selected.repo_revision,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args
)
for dit_folder in selected.dit_folder: # wan a14b has transformer and transformer_2
if dit_folder is not None and dit_folder not in kwargs:
# get a new quant arg on every loop to prevent the quant config classes getting entangled
load_args, quant_args = model_quant.get_dit_args({}, module='Model', device_map=True)
shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" module="{dit_folder}" folder="{dit_folder}" cls={selected.dit_cls.__name__} quant={model_quant.get_quant_type(quant_args)}')
kwargs[dit_folder] = selected.dit_cls.from_pretrained(
pretrained_model_name_or_path=selected.dit or selected.repo,
subfolder=dit_folder,
revision=selected.dit_revision or selected.repo_revision,
cache_dir=shared.opts.hfcache_dir,
**load_args,
**quant_args
)
else:
shared.log.debug(f'Video load: module=transformer repo="{selected.dit or selected.repo}" module="{dit_folder}" folder="{dit_folder}" cls={selected.dit_cls.__name__} skip')
except Exception as e:
shared.log.error(f'video load: module=transformer cls={selected.dit_cls.__name__} {e}')
errors.display(e, 'video')
+9
View File
@@ -19,6 +19,15 @@ def load_override(selected: Model):
# WAN
if 'WAN 2.1 14B' in selected.name:
kwargs['vae'] = diffusers.AutoencoderKLWan.from_pretrained(selected.repo, subfolder="vae", torch_dtype=torch.float32, cache_dir=shared.opts.hfcache_dir)
if 'A14B' in selected.name or '14B VACE' in selected.name:
if shared.opts.model_wan_stage == 'combined':
kwargs['boundary_ratio'] = shared.opts.model_wan_boundary
elif shared.opts.model_wan_stage == 'high noise':
kwargs['transformer_2'] = None
kwargs['boundary_ratio'] = 0.0
elif shared.opts.model_wan_stage == 'low noise':
kwargs['boundary_ratio'] = 1.0
kwargs['transformer'] = None
debug(f'Video overrides: model="{selected.name}" kwargs={list(kwargs)}')
return kwargs