mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
support configurable multi-stage models in video tab
Signed-off-by: Vladimir Mandic <mandic00@live.com>
This commit is contained in:
@@ -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',
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user