"""Strict, allocation-efficient loading of the existing VideoX VACE format.""" import json from pathlib import Path import torch def load_text_encoder(path, additional, dtype): from accelerate import init_empty_weights from safetensors.torch import load_file from videox_fun.models import WanT5EncoderModel from videox_fun.utils.utils import filter_kwargs with init_empty_weights(): model = WanT5EncoderModel(**filter_kwargs(WanT5EncoderModel,additional)) state = load_file(str(path)) if str(path).endswith('.safetensors') else torch.load( path,map_location='cpu',weights_only=True) model.load_state_dict(state,strict=True,assign=True) del state return model.to(dtype=dtype).eval() def load_vace(path, additional, dtype): from accelerate import init_empty_weights from diffusers.utils import WEIGHTS_NAME from safetensors.torch import load_file from videox_fun.models import VaceWanModel path = Path(path) config = json.loads((path/'config.json').read_text()) additional = dict(additional) for source,target in additional.get('dict_mapping',{}).items(): additional[target] = config[source] # Keep ordinary non-parameter RoPE tensors on CPU; only parameters are meta. with init_empty_weights(): model = VaceWanModel.from_config(config,**additional) binary = path/WEIGHTS_NAME safetensor = binary.with_suffix('.safetensors') if binary.is_file(): state = torch.load(binary,map_location='cpu',weights_only=True) else: files = [safetensor] if safetensor.is_file() else sorted(path.glob('*.safetensors')) if not files: raise FileNotFoundError(f'No VACE weights under {path}') state = {} for file in files: part = load_file(str(file)) overlap = set(state).intersection(part) if overlap: raise ValueError(f'Duplicate model tensors in {file}: {sorted(overlap)[:3]}') state.update(part) model.load_state_dict(state,strict=True,assign=True) del state return model.to(dtype=dtype)