Kry4ta1's picture
Add files using upload-large-folder tool
9264c1c verified
Raw History Blame Contribute Delete
2.1 kB
"""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)