cc / lora_loader.py
sdfdsfsf32e3's picture
Upload 6 files
3f96ccc verified
Raw
History Blame Contribute Delete
5.98 kB
"""
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