mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
refactor(wanai): retire bespoke load_transformer, route through generic
The 40-line hand-rolled loader was generic.load_transformer plus a VACE/standard class switch and a hardcoded GGUF rejection. Switch moves to load_wan; GGUF rejection removed (generic handles it via GGUFQuantizationConfig). No native_spec passed: WanTransformer3DModel has a working diffusers converter, so from_single_file via generic stays correct.
This commit is contained in:
@@ -1,51 +1,8 @@
|
||||
import os
|
||||
import transformers
|
||||
import diffusers
|
||||
from modules import shared, devices, sd_models, model_quant, sd_hijack_te, sd_hijack_vae
|
||||
from modules.logger import log
|
||||
|
||||
|
||||
def load_transformer(repo_id, diffusers_load_config=None, subfolder='transformer'):
|
||||
if diffusers_load_config is None:
|
||||
diffusers_load_config = {}
|
||||
load_args, quant_args = model_quant.get_dit_args(diffusers_load_config, module='Model', device_map=True)
|
||||
fn = None
|
||||
|
||||
if 'VACE' in repo_id:
|
||||
transformer_cls = diffusers.WanVACETransformer3DModel
|
||||
else:
|
||||
transformer_cls = diffusers.WanTransformer3DModel
|
||||
|
||||
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):
|
||||
log.error(f'Load module: type=Transformer not found: {shared.opts.sd_unet}')
|
||||
return None
|
||||
fn = sd_unet.unet_dict[shared.opts.sd_unet] if os.path.exists(sd_unet.unet_dict[shared.opts.sd_unet]) else None
|
||||
|
||||
if fn is not None and 'gguf' in fn.lower():
|
||||
log.error('Load model: type=WanAI format="gguf" unsupported')
|
||||
transformer = None
|
||||
elif fn is not None and 'safetensors' in fn.lower():
|
||||
log.debug(f'Load model: type=WanAI {subfolder}="{fn}" quant="{model_quant.get_quant(repo_id)}" args={load_args}')
|
||||
transformer = transformer_cls.from_single_file(
|
||||
fn,
|
||||
cache_dir=shared.opts.hfcache_dir,
|
||||
**load_args,
|
||||
**quant_args,
|
||||
)
|
||||
else:
|
||||
log.debug(f'Load model: type=WanAI {subfolder}="{repo_id}" quant="{model_quant.get_quant_type(quant_args)}" args={load_args}')
|
||||
transformer = transformer_cls.from_pretrained(
|
||||
repo_id,
|
||||
subfolder=subfolder,
|
||||
cache_dir=shared.opts.hfcache_dir,
|
||||
**load_args,
|
||||
**quant_args,
|
||||
)
|
||||
if shared.opts.diffusers_offload_mode != 'none' and transformer is not None:
|
||||
sd_models.move_model(transformer, devices.cpu)
|
||||
return transformer
|
||||
from pipelines import generic
|
||||
|
||||
|
||||
def load_text_encoder(repo_id, diffusers_load_config=None):
|
||||
@@ -71,26 +28,27 @@ def load_wan(checkpoint_info, diffusers_load_config=None):
|
||||
diffusers_load_config = {}
|
||||
repo_id = sd_models.path_to_repo(checkpoint_info)
|
||||
sd_models.hf_auth_check(checkpoint_info)
|
||||
transformer_cls = diffusers.WanVACETransformer3DModel if 'VACE' in repo_id else diffusers.WanTransformer3DModel
|
||||
|
||||
boundary_ratio = None
|
||||
if 'a14b' in repo_id.lower() or 'fun-14b' in repo_id.lower():
|
||||
if shared.opts.model_wan_stage == 'high noise' or shared.opts.model_wan_stage == 'first':
|
||||
transformer = load_transformer(repo_id, diffusers_load_config, 'transformer')
|
||||
transformer = generic.load_transformer(repo_id, cls_name=transformer_cls, load_config=diffusers_load_config, subfolder='transformer')
|
||||
transformer_2 = None
|
||||
boundary_ratio = 0.0
|
||||
elif shared.opts.model_wan_stage == 'low noise' or shared.opts.model_wan_stage == 'second':
|
||||
transformer = None
|
||||
transformer_2 = load_transformer(repo_id, diffusers_load_config, 'transformer_2')
|
||||
transformer_2 = generic.load_transformer(repo_id, cls_name=transformer_cls, load_config=diffusers_load_config, subfolder='transformer_2')
|
||||
boundary_ratio = 1000.0
|
||||
elif shared.opts.model_wan_stage == 'combined' or 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')
|
||||
transformer = generic.load_transformer(repo_id, cls_name=transformer_cls, load_config=diffusers_load_config, subfolder='transformer')
|
||||
transformer_2 = generic.load_transformer(repo_id, cls_name=transformer_cls, load_config=diffusers_load_config, subfolder='transformer_2')
|
||||
boundary_ratio = shared.opts.model_wan_boundary
|
||||
else:
|
||||
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 = generic.load_transformer(repo_id, cls_name=transformer_cls, load_config=diffusers_load_config, subfolder='transformer')
|
||||
transformer_2 = None
|
||||
|
||||
text_encoder = load_text_encoder(repo_id, diffusers_load_config)
|
||||
|
||||
Reference in New Issue
Block a user