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
|