File size: 5,983 Bytes
3f96ccc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
"""
LoRA loading for HunyuanVideo 1.5 I2V.

Unlike the Wan pipelines, `HunyuanVideo15ImageToVideoPipeline` does not inherit a LoRA loader
mixin, so LoRAs go straight onto the transformer through `PeftAdapterMixin`
(`load_lora_adapter` / `set_adapters` / `delete_adapters` / `fuse_lora`).

Wan 2.2 LoRAs are NOT compatible: different architecture, different key names, and Wan's
high-noise/low-noise expert pairs have no counterpart here (HunyuanVideo 1.5 has a single
transformer). You need LoRAs trained for HunyuanVideo 1.5.

The catalog is data, not code. Put a `loras.json` next to this file:

    [
      {"label": "My Style",  "repo_id": "user/repo", "weight_name": "style.safetensors", "scale": 1.0},
      {"label": "Fused One", "repo_id": "user/repo2", "fuse_at_startup": true, "scale": 0.8}
    ]

or set LORA_CATALOG to the same JSON inline. Users can also type any repo into the
"Custom LoRA" box in the UI as `repo_id` or `repo_id:weight_name`.
"""
import json
import os

from huggingface_hub import hf_hub_download, snapshot_download

HF_TOKEN = os.environ.get("HF_TOKEN")  # authenticated downloads (covers private repos)
CATALOG_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "loras.json")

_LOADED_ADAPTERS = []


def _read_catalog():
    raw = os.environ.get("LORA_CATALOG")
    if raw:
        try:
            return json.loads(raw)
        except Exception as e:
            print("LORA_CATALOG is not valid JSON:", e)
            return []
    if os.path.exists(CATALOG_PATH):
        try:
            with open(CATALOG_PATH) as fh:
                return json.load(fh)
        except Exception as e:
            print(f"Could not read {CATALOG_PATH}:", e)
    return []


CATALOG = {}
for _entry in _read_catalog():
    if not isinstance(_entry, dict) or not _entry.get("repo_id"):
        continue
    _label = _entry.get("label") or _entry["repo_id"].split("/")[-1]
    CATALOG[_label] = _entry


def get_lora_choices():
    return sorted(CATALOG.keys())


def _resolve_path(repo_id, weight_name=None, revision=None):
    """Return a local path diffusers can load: a single file when weight_name is given,
    otherwise the whole snapshot (diffusers will find the LoRA inside)."""
    if weight_name:
        return hf_hub_download(repo_id, weight_name, token=HF_TOKEN, revision=revision)
    return snapshot_download(repo_id, token=HF_TOKEN, revision=revision)


def _parse_custom(spec):
    """'repo/name' or 'repo/name:file.safetensors' -> (repo_id, weight_name|None)"""
    spec = (spec or "").strip()
    if not spec:
        return None
    if ":" in spec:
        repo_id, weight_name = spec.split(":", 1)
        return repo_id.strip(), weight_name.strip() or None
    return spec, None


def load_loras_to_pipe(pipe, labels=None, custom=None, scale=1.0):
    """Load the selected catalog entries plus an optional custom LoRA onto pipe.transformer.

    Returns True if at least one adapter was attached. Note that stacking runtime LoRAs on an
    fp8-quantized transformer can fail depending on the torchao/peft versions; if that happens,
    use `fuse_at_startup` in the catalog instead (fusion runs before quantization).
    """
    unload_lora(pipe)

    requests = []
    for label in (labels or []):
        entry = CATALOG.get(label)
        if not entry or entry.get("fuse_at_startup"):
            continue
        requests.append((label, entry.get("repo_id"), entry.get("weight_name"),
                         entry.get("revision"), float(entry.get("scale", 1.0))))

    parsed = _parse_custom(custom)
    if parsed:
        requests.append(("custom", parsed[0], parsed[1], None, 1.0))

    if not requests:
        return False

    names, weights = [], []
    for idx, (label, repo_id, weight_name, revision, entry_scale) in enumerate(requests):
        path = _resolve_path(repo_id, weight_name, revision)
        adapter_name = f"lora_{idx}"
        pipe.transformer.load_lora_adapter(path, prefix="transformer", adapter_name=adapter_name)
        names.append(adapter_name)
        weights.append(entry_scale * float(scale))
        print(f"Loaded LoRA: {label} ({repo_id})")

    pipe.transformer.set_adapters(names, weights=weights)
    _LOADED_ADAPTERS[:] = names
    return True


def unload_lora(pipe):
    if not _LOADED_ADAPTERS:
        return
    try:
        pipe.transformer.delete_adapters(list(_LOADED_ADAPTERS))
    except Exception:
        try:
            pipe.transformer.unload_lora()
        except Exception:
            pass
    _LOADED_ADAPTERS.clear()


def fuse_startup_loras(pipe):
    """Fuse catalog entries marked `fuse_at_startup` into the transformer weights.

    Call this once at import time, before quantization: fused weights survive fp8 conversion,
    runtime adapters may not. Mirrors what the Wan reference space does with its Lightning LoRAs.
    """
    entries = [e for e in CATALOG.values() if e.get("fuse_at_startup")]
    if not entries:
        return
    for i, entry in enumerate(entries):
        adapter_name = f"startup_{i}"
        try:
            path = _resolve_path(entry["repo_id"], entry.get("weight_name"), entry.get("revision"))
            pipe.transformer.load_lora_adapter(path, prefix="transformer", adapter_name=adapter_name)
            pipe.transformer.set_adapters([adapter_name], weights=[1.0])
            pipe.transformer.fuse_lora(lora_scale=float(entry.get("scale", 1.0)),
                                       adapter_names=[adapter_name])
            pipe.transformer.delete_adapters([adapter_name])
            print(f"Fused LoRA at startup: {entry.get('label', entry['repo_id'])} "
                  f"(scale={entry.get('scale', 1.0)}), {i + 1}/{len(entries)}")
        except Exception as e:
            print("Error:", str(e))
            print("Failed LoRA:", entry.get("label", entry.get("repo_id")))
            try:
                pipe.transformer.delete_adapters([adapter_name])
            except Exception:
                pass