Download scripts/convert_cached_wan_diffusers_to_original.py from Cccccz/Causal-Forcing-a: direct link, hf CLI and curl.
- Browser
- Download file 3.42 kB
-
https://huggingface.co/Cccccz/Causal-Forcing-a/resolve/main/scripts/convert_cached_wan_diffusers_to_original.py
- Command line
-
hf download hf://Cccccz/Causal-Forcing-a/scripts/convert_cached_wan_diffusers_to_original.py
-
curl -L -o convert_cached_wan_diffusers_to_original.py https://huggingface.co/Cccccz/Causal-Forcing-a/resolve/main/scripts/convert_cached_wan_diffusers_to_original.py
3.42 kB
| #!/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() | |