File size: 2,097 Bytes
9264c1c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
"""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)