File size: 3,415 Bytes
ae8ade0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
#!/usr/bin/env python3
"""Convert official Diffusers Wan transformer shards to WanModel key names.

No tensor is numerically changed.  Shards are processed one at a time so the
conversion does not need to materialize the 14B state dict twice in memory.
"""

from __future__ import annotations

import argparse
import json
import os
from pathlib import Path

from safetensors.torch import load_file, save_file


REPLACEMENTS = (
    ("attn1.to_q", "self_attn.q"),
    ("attn1.to_k", "self_attn.k"),
    ("attn1.to_v", "self_attn.v"),
    ("attn1.to_out.0", "self_attn.o"),
    ("attn1.norm_q", "self_attn.norm_q"),
    ("attn1.norm_k", "self_attn.norm_k"),
    ("attn2.to_q", "cross_attn.q"),
    ("attn2.to_k", "cross_attn.k"),
    ("attn2.to_v", "cross_attn.v"),
    ("attn2.to_out.0", "cross_attn.o"),
    ("attn2.norm_q", "cross_attn.norm_q"),
    ("attn2.norm_k", "cross_attn.norm_k"),
    ("ffn.net.0.proj", "ffn.0"),
    ("ffn.net.2", "ffn.2"),
    (".norm2.", ".norm3."),
    (".scale_shift_table", ".modulation"),
    ("condition_embedder.text_embedder.linear_1", "text_embedding.0"),
    ("condition_embedder.text_embedder.linear_2", "text_embedding.2"),
    ("condition_embedder.time_embedder.linear_1", "time_embedding.0"),
    ("condition_embedder.time_embedder.linear_2", "time_embedding.2"),
    ("condition_embedder.time_proj", "time_projection.1"),
    ("proj_out", "head.head"),
    ("scale_shift_table", "head.modulation"),
)


def rename(key: str) -> str:
    for source, target in REPLACEMENTS:
        key = key.replace(source, target)
    return key


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--source", type=Path, required=True)
    parser.add_argument("--output", type=Path, required=True)
    args = parser.parse_args()
    source = args.source.resolve()
    output = args.output.resolve()
    output.mkdir(parents=True, exist_ok=True)
    source_index = json.loads((source / "diffusion_pytorch_model.safetensors.index.json").read_text())
    source_map: dict[str, str] = source_index["weight_map"]
    output_map = {rename(key): filename for key, filename in source_map.items()}
    if len(output_map) != len(source_map):
        raise RuntimeError("Wan key conversion produced a collision")

    for shard in sorted(set(source_map.values())):
        destination = output / shard
        if destination.exists() and destination.stat().st_size == (source / shard).stat().st_size:
            print(f"[skip] {shard}", flush=True)
            continue
        print(f"[convert] {shard}", flush=True)
        tensors = load_file(str(source / shard), device="cpu")
        converted = {rename(key): value.contiguous() for key, value in tensors.items()}
        temporary = destination.with_suffix(destination.suffix + ".tmp")
        save_file(converted, str(temporary))
        os.replace(temporary, destination)
        del tensors, converted

    index = {"metadata": source_index.get("metadata", {}), "weight_map": output_map}
    temporary_index = output / "diffusion_pytorch_model.safetensors.index.json.tmp"
    temporary_index.write_text(json.dumps(index, indent=2, sort_keys=True) + "\n")
    os.replace(temporary_index, output / "diffusion_pytorch_model.safetensors.index.json")
    print(f"[complete] tensors={len(output_map)} shards={len(set(output_map.values()))}", flush=True)


if __name__ == "__main__":
    main()