From d01d637d9433abe0756bc71fae3f81ad820819d3 Mon Sep 17 00:00:00 2001 From: CalamitousFelicitousness Date: Mon, 25 May 2026 05:47:56 +0100 Subject: [PATCH] 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. --- pipelines/model_wanai.py | 56 +++++----------------------------------- 1 file changed, 7 insertions(+), 49 deletions(-) diff --git a/pipelines/model_wanai.py b/pipelines/model_wanai.py index f7a4413a8..e62fb165d 100644 --- a/pipelines/model_wanai.py +++ b/pipelines/model_wanai.py @@ -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)