John6666's picture
Upload 39 files
e8b6587 verified
Raw
History Blame Contribute Delete
31.6 kB
"""Ordinary user-LoRA preparation and request-scoped adapter lifecycle.
This module owns backend LoRA policy only. Gradio widgets and event wiring remain in
``app.py``. The internal Full/SFT Stage-2 adapter and IC-Colorizer adapter remain
separate domains; this module deliberately does not generalize them into one manager.
"""
from __future__ import annotations
import gc
import time
import urllib.parse
from pathlib import Path
from typing import Callable
import torch
from huggingface_hub import HfApi, hf_hub_download, parse_hf_uri
from huggingface_hub.errors import (
GatedRepoError,
HfHubHTTPError,
LocalEntryNotFoundError,
RemoteEntryNotFoundError,
RepositoryNotFoundError,
RevisionNotFoundError,
)
from safetensors import safe_open
from safetensors.torch import load_file as load_safetensors_file, save_file as save_safetensors_file
from ltx import civitai as civitai_backend
from ltx.runtime_utils import resolved_revision_from_hub_path
class LoraSelectionError(ValueError):
"""User-facing ordinary LoRA selection/preparation error."""
def sanitize_defs(items, source: str = "configured") -> list[dict]:
"""Normalize portable LoRA definitions without persisting transport secrets."""
clean: list[dict] = []
seen: set[str] = set()
for raw in items or []:
if not isinstance(raw, dict):
continue
label = str(raw.get("label") or "").strip()
weight_name = str(raw.get("weight_name") or "").strip()
item_source = str(raw.get("source") or source).strip() or source
if not label or not weight_name or label in seen:
continue
if item_source == "civitai":
try:
model_id = int(raw.get("civitai_model_id"))
version_id = int(raw.get("civitai_version_id"))
file_id = int(raw.get("civitai_file_id"))
except Exception:
continue
clean.append(
{
"label": label,
"source": "civitai",
"repo_id": "",
"weight_name": weight_name,
"revision": "",
"civitai_model_id": model_id,
"civitai_version_id": version_id,
"civitai_file_id": file_id,
"civitai_sha256": str(raw.get("civitai_sha256") or "").strip().lower(),
"civitai_page_url": str(raw.get("civitai_page_url") or "").strip(),
"trained_words": [str(x) for x in (raw.get("trained_words") or []) if str(x).strip()],
}
)
seen.add(label)
continue
repo_id = str(raw.get("repo_id") or "").strip()
revision = str(raw.get("revision") or "").strip()
if repo_id:
clean.append(
{
"label": label,
"repo_id": repo_id,
"weight_name": weight_name,
"revision": revision,
"source": item_source,
}
)
seen.add(label)
return clean
def all_defs(builtin_loras, custom_loras=None) -> list[dict]:
builtins = sanitize_defs(builtin_loras, source="built_in")
sessions = sanitize_defs(custom_loras, source="session")
builtin_labels = {item["label"] for item in builtins}
sessions = [item for item in sessions if item["label"] not in builtin_labels]
return builtins + sessions
def validate_repo_id(repo_id: str) -> str:
repo_id = str(repo_id or "").strip()
if repo_id.count("/") != 1 or repo_id.startswith("/") or repo_id.endswith("/"):
raise LoraSelectionError("HF repo ID must look like owner/repo.")
return repo_id
def validate_weight_name(weight_name: str) -> str:
weight_name = str(weight_name or "").strip().replace("\\", "/")
if not weight_name.endswith(".safetensors"):
raise LoraSelectionError("LoRA weight filename must end with .safetensors.")
if weight_name.startswith("/") or ".." in Path(weight_name).parts:
raise LoraSelectionError("LoRA weight filename must be a repository-relative path.")
return weight_name
def normalize_hf_lora_source(source: str, revision: str = "", weight_name: str = "") -> dict:
"""Canonicalize a Hub repo ID, model URL, file URL, or hf:// model URI."""
raw = str(source or "").strip()
if not raw:
raise LoraSelectionError("Enter an HF model repo ID or URL.")
revision = str(revision or "").strip()
weight_name = str(weight_name or "").strip().replace("\\", "/")
url_revision = ""
url_weight = ""
source_kind = "repo_id"
is_hf_uri = raw.startswith("hf://")
is_hf_web = "://" in raw or raw.startswith(("huggingface.co/", "www.huggingface.co/", "hf.co/"))
if is_hf_uri:
source_kind = "hf_uri"
try:
uri = parse_hf_uri(raw)
except Exception as exc:
raise LoraSelectionError("Invalid Hugging Face hf:// URI.") from exc
if uri.type != "model":
raise LoraSelectionError("HF LoRA source must point to a model repository, not a dataset/Space/kernel/bucket.")
repo_id = validate_repo_id(uri.id)
url_revision = str(uri.revision or "").strip()
url_weight = str(uri.path_in_repo or "").strip().replace("\\", "/")
elif not is_hf_web:
repo_id = validate_repo_id(raw)
else:
source_kind = "url"
web_raw = raw if "://" in raw else f"https://{raw}"
# Newer huggingface_hub releases parse supported web URLs directly. Keep
# the small fallback below so the product remains tolerant of older
# environments while requirements/runtime converge.
uri = None
try:
uri = parse_hf_uri(web_raw)
except Exception:
uri = None
if uri is not None:
if uri.type != "model":
raise LoraSelectionError("HF LoRA source must point to a model repository, not a dataset/Space/kernel/bucket.")
repo_id = validate_repo_id(uri.id)
url_revision = str(uri.revision or "").strip()
url_weight = str(uri.path_in_repo or "").strip().replace("\\", "/")
else:
try:
parsed = urllib.parse.urlsplit(web_raw)
except Exception as exc:
raise LoraSelectionError("Invalid Hugging Face URL.") from exc
if (parsed.hostname or "").casefold() not in {"huggingface.co", "www.huggingface.co", "hf.co"}:
raise LoraSelectionError("HF LoRA URL must point to huggingface.co or hf.co.")
raw_parts = [x for x in (parsed.path or "").strip("/").split("/") if x]
parts = [urllib.parse.unquote(x) for x in raw_parts]
if parts and parts[0] == "models":
parts = parts[1:]
if parts and parts[0] in {"datasets", "spaces", "kernels", "buckets", "collections"}:
raise LoraSelectionError("HF LoRA source must point to a model repository.")
if len(parts) < 2:
raise LoraSelectionError("HF model URL must include owner/repo.")
repo_id = validate_repo_id(f"{parts[0]}/{parts[1]}")
if len(parts) > 2:
route = parts[2]
if route in {"blob", "resolve", "raw"}:
rest = parts[3:]
if len(rest) < 2:
raise LoraSelectionError("HF file URL must include revision and repository-relative file path.")
if len(rest) >= 4 and rest[0] == "refs" and rest[1] in {"pr", "convert"}:
url_revision = "/".join(rest[:3])
url_weight = "/".join(rest[3:])
else:
url_revision = rest[0]
url_weight = "/".join(rest[1:])
elif route == "tree":
rest = parts[3:]
if not rest:
raise LoraSelectionError("HF tree URL must include a revision.")
if len(rest) >= 3 and rest[0] == "refs" and rest[1] in {"pr", "convert"}:
url_revision = "/".join(rest[:3])
else:
url_revision = rest[0]
else:
raise LoraSelectionError("Unsupported Hugging Face model URL route. Use the repo page or a blob/resolve/raw file URL.")
if url_revision and revision and url_revision != revision:
raise LoraSelectionError(f"HF URL revision {url_revision!r} conflicts with Revision {revision!r}.")
resolved_revision = url_revision or revision
if url_weight:
url_weight = validate_weight_name(url_weight)
if weight_name and validate_weight_name(weight_name) != url_weight:
raise LoraSelectionError("HF file URL conflicts with the Safetensors file field.")
resolved_weight = url_weight
else:
resolved_weight = validate_weight_name(weight_name) if weight_name else ""
return {
"repo_id": repo_id,
"revision": resolved_revision,
"weight_name": resolved_weight,
"source_kind": source_kind,
}
def _download_hf_lora_file(item: dict, hf_token: str | None) -> Path:
"""Download one exact Hub LoRA file and normalize common Hub failures."""
repo_id = str(item.get("repo_id") or "").strip()
weight_name = str(item.get("weight_name") or "").strip()
revision = str(item.get("revision") or "").strip() or None
try:
return Path(
hf_hub_download(
repo_id=repo_id,
filename=weight_name,
revision=revision,
token=hf_token,
)
)
except GatedRepoError as exc:
raise LoraSelectionError(
f"HF LoRA repo {repo_id} is gated and the current Space token cannot access it. "
"Request access for the token owner or configure an authorized HF token."
) from exc
except RepositoryNotFoundError as exc:
raise LoraSelectionError(
f"HF LoRA repo {repo_id} was not found or is private to the current Space token."
) from exc
except RevisionNotFoundError as exc:
raise LoraSelectionError(
f"HF revision {revision!r} was not found in {repo_id}."
) from exc
except RemoteEntryNotFoundError as exc:
where = f" at revision {revision!r}" if revision else ""
raise LoraSelectionError(
f"HF LoRA file {weight_name!r} was not found in {repo_id}{where}. "
"Inspect the repo and choose an exact .safetensors file."
) from exc
except LocalEntryNotFoundError as exc:
raise LoraSelectionError(
f"HF LoRA file {weight_name!r} is not cached locally and the Hub could not be reached. "
"Check network/offline mode and try again."
) from exc
except HfHubHTTPError as exc:
status = getattr(getattr(exc, "response", None), "status_code", None)
if status in {401, 403}:
raise LoraSelectionError(
f"HF denied access while downloading {repo_id}/{weight_name} (HTTP {status}). "
"Check gated/private access and the Space HF token."
) from exc
if status == 429:
raise LoraSelectionError(
"HF Hub rate-limited the LoRA download (HTTP 429). Wait briefly and try again."
) from exc
suffix = f" (HTTP {status})" if status else ""
raise LoraSelectionError(
f"HF LoRA download failed for {repo_id}/{weight_name}{suffix}."
) from exc
except OSError as exc:
raise LoraSelectionError(
f"HF LoRA download could not complete for {repo_id}/{weight_name}: {type(exc).__name__}."
) from exc
def resolve_selected(selected, builtin_loras, custom_loras) -> list[dict]:
defs = {item["label"]: item for item in all_defs(builtin_loras, custom_loras)}
resolved = []
for label in selected or []:
if label not in defs:
raise LoraSelectionError(f"Selected LoRA is unavailable: {label}")
resolved.append(defs[label])
return resolved
def inspect_model_repo(repo_id: str, hf_token: str | None, revision: str | None = None) -> dict:
"""CPU/network-only Hub metadata inspection for the session-LoRA form."""
repo_id = validate_repo_id(repo_id)
revision = str(revision or "").strip() or None
try:
info = HfApi(token=hf_token).model_info(repo_id=repo_id, revision=revision, files_metadata=False)
except GatedRepoError as exc:
raise LoraSelectionError(f"HF repo {repo_id} is gated and the current Space token cannot access it.") from exc
except RepositoryNotFoundError as exc:
raise LoraSelectionError(f"HF model repo {repo_id} was not found or is private to the current Space token.") from exc
except RevisionNotFoundError as exc:
raise LoraSelectionError(f"HF revision {revision!r} was not found in {repo_id}.") from exc
except HfHubHTTPError as exc:
status = getattr(getattr(exc, "response", None), "status_code", None)
suffix = f" (HTTP {status})" if status else ""
raise LoraSelectionError(f"HF repo inspection failed for {repo_id}{suffix}.") from exc
except Exception as exc:
raise LoraSelectionError(f"HF repo inspection failed for {repo_id}: {type(exc).__name__}.") from exc
safetensors = sorted(
{
str(getattr(sibling, "rfilename", "") or "")
for sibling in (getattr(info, "siblings", None) or [])
if str(getattr(sibling, "rfilename", "") or "").lower().endswith(".safetensors")
}
)
return {
"repo_id": repo_id,
"requested_revision": revision,
"resolved_revision": str(getattr(info, "sha", "") or "") or None,
"private": bool(getattr(info, "private", False)),
"gated": getattr(info, "gated", None),
"safetensors": safetensors,
"tags": [str(x) for x in (getattr(info, "tags", None) or [])],
}
_LTX2_LORA_PREFIXES = ("diffusion_model.", "text_embedding_projection.", "transformer.", "connectors.")
def _checkpoint_key_sample(keys, limit: int = 6) -> list[str]:
return [str(key) for key in list(keys)[: max(1, int(limit))]]
def _mapped_lora_key(key: str) -> str:
return str(key).replace(".lora_down.weight", ".lora_A.weight").replace(".lora_up.weight", ".lora_B.weight")
def _inspect_civitai_ltx2_checkpoint(local: Path) -> dict:
"""Inspect a Civitai safetensors header without loading tensor payloads."""
try:
with safe_open(str(local), framework="pt", device="cpu") as handle:
keys = list(handle.keys())
except Exception as exc:
raise LoraSelectionError(f"Could not inspect Civitai LoRA safetensors header: {type(exc).__name__}: {exc}") from exc
if not keys:
raise LoraSelectionError("Civitai checkpoint is empty; refusing to request GPU quota.")
unsupported_prefixes = [key for key in keys if not str(key).startswith(_LTX2_LORA_PREFIXES)]
if unsupported_prefixes:
sample = ", ".join(_checkpoint_key_sample(unsupported_prefixes))
raise LoraSelectionError(
"Unsupported LTX-2 LoRA checkpoint dialect from Civitai. "
f"Unexpected parameter namespace(s): {sample}. Try another file/version. GPU quota was not requested."
)
alpha_keys = [key for key in keys if str(key).endswith(".alpha")]
down_up_keys = [key for key in keys if ".lora_down.weight" in str(key) or ".lora_up.weight" in str(key)]
unsupported_non_lora = [key for key in keys if "lora" not in str(key).lower() and not str(key).endswith(".alpha")]
if unsupported_non_lora:
sample = ", ".join(_checkpoint_key_sample(unsupported_non_lora))
raise LoraSelectionError(
"Civitai file is not a supported LTX-2 LoRA-only checkpoint. "
f"Unexpected parameter(s): {sample}. Try another exact .safetensors file/version. GPU quota was not requested."
)
if alpha_keys or down_up_keys:
dialect = "ltx2_comfy_alpha_or_down_up"
needs_normalization = True
elif any(str(key).startswith("diffusion_model.") or str(key).startswith("text_embedding_projection.") for key in keys):
dialect = "ltx2_comfy_native"
needs_normalization = False
else:
dialect = "ltx2_diffusers_peft"
needs_normalization = False
return {
"dialect": dialect,
"key_count": len(keys),
"key_sample": _checkpoint_key_sample(keys),
"alpha_key_count": len(alpha_keys),
"down_up_key_count": len(down_up_keys),
"needs_normalization": needs_normalization,
}
def _normalize_civitai_ltx2_checkpoint(local: Path, cache_root: Path, source_sha256: str) -> tuple[Path, dict]:
"""Normalize only known LTX-2 Comfy LoRA variants before GPU allocation.
Current pinned Diffusers accepts LTX2 Comfy-style ``diffusion_model.`` keys but
does not normalize per-module ``.alpha`` or generic ``lora_down/lora_up`` keys.
For that narrow, known dialect we rename down/up to A/B and fold alpha/rank
into lora_B. Unknown or mixed checkpoint dialects fail closed.
"""
evidence = _inspect_civitai_ltx2_checkpoint(local)
if not evidence["needs_normalization"]:
evidence.update(normalized=False, normalized_path=None, alpha_fold_count=0)
return local, evidence
try:
state = load_safetensors_file(str(local), device="cpu")
except Exception as exc:
raise LoraSelectionError(f"Could not read Civitai LoRA tensors for CPU normalization: {type(exc).__name__}: {exc}") from exc
normalized: dict[str, torch.Tensor] = {}
alpha_tensors: dict[str, torch.Tensor] = {}
for raw_key, tensor in state.items():
key = str(raw_key)
if key.endswith(".alpha"):
alpha_tensors[key[:-len(".alpha")]] = tensor
continue
mapped = _mapped_lora_key(key)
if mapped in normalized:
raise LoraSelectionError(f"Checkpoint normalization produced a duplicate LoRA key: {mapped}")
normalized[mapped] = tensor
a_keys = [key for key in normalized if key.endswith(".lora_A.weight")]
b_keys = [key for key in normalized if key.endswith(".lora_B.weight")]
if not a_keys or not b_keys:
raise LoraSelectionError(
"Unsupported Civitai LoRA checkpoint: no complete lora_A/lora_B pairs were found after known-format normalization. "
"Try another file/version. GPU quota was not requested."
)
pair_bases = set()
for key in a_keys:
base = key[:-len(".lora_A.weight")]
b_key = base + ".lora_B.weight"
if b_key not in normalized:
raise LoraSelectionError(f"Incomplete Civitai LoRA pair for {base}: lora_B is missing. GPU quota was not requested.")
pair_bases.add(base)
for key in b_keys:
base = key[:-len(".lora_B.weight")]
if base + ".lora_A.weight" not in normalized:
raise LoraSelectionError(f"Incomplete Civitai LoRA pair for {base}: lora_A is missing. GPU quota was not requested.")
pair_bases.add(base)
alpha_fold_count = 0
for base, alpha_tensor in alpha_tensors.items():
a_key = base + ".lora_A.weight"
b_key = base + ".lora_B.weight"
if a_key not in normalized or b_key not in normalized:
raise LoraSelectionError(
f"Civitai LoRA alpha has no matching A/B pair for {base}. Refusing unsafe fallback before GPU allocation."
)
if alpha_tensor.numel() != 1:
raise LoraSelectionError(f"Civitai LoRA alpha for {base} is not scalar; unsupported checkpoint dialect.")
a_tensor = normalized[a_key]
if a_tensor.ndim < 1 or int(a_tensor.shape[0]) <= 0:
raise LoraSelectionError(f"Could not derive LoRA rank for {base}; unsupported checkpoint dialect.")
rank = int(a_tensor.shape[0])
alpha = float(alpha_tensor.detach().float().item())
scale = alpha / float(rank)
original_b = normalized[b_key]
normalized[b_key] = (original_b.detach().float() * scale).to(dtype=original_b.dtype)
alpha_fold_count += 1
remaining_bad = [key for key in normalized if "lora" not in key.lower()]
if remaining_bad:
sample = ", ".join(_checkpoint_key_sample(remaining_bad))
raise LoraSelectionError(f"Unsupported non-LoRA parameter(s) remain after normalization: {sample}")
source_sha256 = str(source_sha256 or "").strip().lower()
if len(source_sha256) != 64:
import hashlib
digest = hashlib.sha256()
with local.open("rb") as fh:
for chunk in iter(lambda: fh.read(1024 * 1024), b""):
digest.update(chunk)
source_sha256 = digest.hexdigest()
out_dir = Path(cache_root) / "normalized_ltx2_lora"
out_dir.mkdir(parents=True, exist_ok=True)
out_path = out_dir / f"{source_sha256}.diffusers-compatible.safetensors"
try:
save_safetensors_file(normalized, str(out_path))
except Exception as exc:
raise LoraSelectionError(f"Could not write runtime-normalized Civitai LoRA: {type(exc).__name__}: {exc}") from exc
evidence.update(
normalized=True,
normalized_path=str(out_path),
normalized_key_count=len(normalized),
alpha_fold_count=alpha_fold_count,
source_sha256=source_sha256,
)
return out_path, evidence
def definition_key(item: dict) -> tuple[str, ...]:
source = str(item.get("source") or "").strip()
if source == "civitai":
return (
"civitai",
str(item.get("civitai_model_id") or ""),
str(item.get("civitai_version_id") or ""),
str(item.get("civitai_file_id") or ""),
str(item.get("civitai_sha256") or "").lower(),
)
return (
"hf",
str(item.get("repo_id") or ""),
str(item.get("weight_name") or ""),
str(item.get("revision") or item.get("requested_revision") or ""),
)
def prepare_selected(selected, builtin_loras, custom_loras, hf_token: str | None, civitai_cache_root: Path | None = None, civitai_api_key: str = "") -> tuple[list[dict], float]:
"""CPU-side source preparation. This function must never run inside a GPU callback."""
resolved = resolve_selected(selected, builtin_loras, custom_loras)
if not resolved:
return [], 0.0
prepared = []
total_started = time.monotonic()
for item in resolved:
item_started = time.monotonic()
if str(item.get("source") or "") == "civitai":
if civitai_cache_root is None:
raise LoraSelectionError("Civitai LoRA cache root is unavailable.")
try:
acquired = civitai_backend.download_exact_file(item, Path(civitai_cache_root), api_key=civitai_api_key)
except civitai_backend.CivitaiError as exc:
raise LoraSelectionError(str(exc)) from exc
local = Path(str(acquired["local_path"]))
source_sha256 = str(acquired.get("sha256") or item.get("civitai_sha256") or "").strip().lower()
local, checkpoint_compatibility = _normalize_civitai_ltx2_checkpoint(
local, Path(civitai_cache_root), source_sha256
)
record = {
"label": item["label"],
"source": "civitai",
"repo_id": None,
"weight_name": item["weight_name"],
"requested_revision": None,
"resolved_revision": None,
"civitai_model_id": item.get("civitai_model_id"),
"civitai_version_id": item.get("civitai_version_id"),
"civitai_file_id": item.get("civitai_file_id"),
"civitai_sha256": acquired.get("sha256") or item.get("civitai_sha256"),
"local_path": str(local),
"source_size_bytes": int(acquired.get("size_bytes") or 0),
"size_bytes": int(local.stat().st_size),
"checkpoint_compatibility": checkpoint_compatibility,
"cache_hit": bool(acquired.get("cache_hit")),
"prepare_seconds": time.monotonic() - item_started,
}
else:
local = _download_hf_lora_file(item, hf_token)
record = {
"label": item["label"],
"source": item.get("source"),
"repo_id": item["repo_id"],
"weight_name": item["weight_name"],
"requested_revision": item["revision"] or None,
"resolved_revision": resolved_revision_from_hub_path(local),
"local_path": str(local),
"size_bytes": int(local.stat().st_size),
"prepare_seconds": time.monotonic() - item_started,
}
record["definition_key"] = list(definition_key(item))
prepared.append(record)
return prepared, time.monotonic() - total_started
def prepared_lookup(prepared_loras) -> dict:
lookup = {}
for item in prepared_loras or []:
if not isinstance(item, dict):
continue
raw_key = item.get("definition_key")
if isinstance(raw_key, (list, tuple)) and raw_key:
key = tuple(str(x) for x in raw_key)
else:
key = definition_key(item)
lookup[key] = item
return lookup
def adapter_state(pipe) -> dict:
state = {}
for name in ("get_active_adapters", "get_list_adapters"):
fn = getattr(pipe, name, None)
if callable(fn):
try:
state[name] = fn()
except Exception as exc:
state[name] = {"error": f"{type(exc).__name__}: {exc}"}
return state
def load_request_loras(
*,
pipe,
pipe_i2v,
pipe_condition,
pipe_ic,
selected,
builtin_loras,
custom_loras,
prepared_loras,
strength: float,
request_id: str,
gpu_state: Callable[[], dict],
) -> tuple[list[dict], dict]:
metrics = {
"requested_labels": [str(x) for x in (selected or [])],
"requested_count": len(selected or []),
"strength": float(strength),
"hub_download_inside_gpu_callback": False,
"request_scoped": True,
"fused": False,
}
if not selected:
metrics["status"] = "not_requested"
return [], metrics
resolved = resolve_selected(selected, builtin_loras, custom_loras)
prepared = prepared_lookup(prepared_loras)
loaded = []
adapter_names = []
metrics["gpu_before_load"] = gpu_state()
total_started = time.monotonic()
try:
per_adapter = []
for idx, item in enumerate(resolved):
key = definition_key(item)
prep = prepared.get(key)
if not prep:
raise LoraSelectionError(
f"Selected LoRA is not CPU-prepared: {item['label']}. "
"Use Prepare selected LoRAs; Generate normally performs this pre-step automatically."
)
local = Path(str(prep.get("local_path") or ""))
if not local.is_file():
raise LoraSelectionError(f"Prepared LoRA file is no longer available locally: {item['label']}")
adapter_name = f"req_{request_id[:8]}_{idx}"
load_started = time.monotonic()
pipe.load_lora_weights(str(local.parent), weight_name=local.name, adapter_name=adapter_name)
load_seconds = time.monotonic() - load_started
adapter_names.append(adapter_name)
public_record = {
"label": item["label"],
"source": item.get("source"),
"repo_id": item.get("repo_id") or None,
"weight_name": item["weight_name"],
"requested_revision": item.get("revision") or None,
"resolved_revision": prep.get("resolved_revision"),
"size_bytes": prep.get("size_bytes"),
"adapter_name": adapter_name,
"load_seconds": load_seconds,
}
if item.get("source") == "civitai":
public_record.update(
civitai_model_id=item.get("civitai_model_id"),
civitai_version_id=item.get("civitai_version_id"),
civitai_file_id=item.get("civitai_file_id"),
civitai_sha256=prep.get("civitai_sha256") or item.get("civitai_sha256"),
checkpoint_compatibility=prep.get("checkpoint_compatibility"),
)
loaded.append(public_record)
per_adapter.append({k: public_record[k] for k in ("label", "adapter_name", "size_bytes", "load_seconds")})
set_started = time.monotonic()
pipe.set_adapters(adapter_names, adapter_weights=[float(strength)] * len(adapter_names))
pipe.enable_lora()
metrics["set_adapters_seconds"] = time.monotonic() - set_started
metrics["load_total_seconds"] = time.monotonic() - total_started
metrics["per_adapter"] = per_adapter
metrics["adapter_names"] = list(adapter_names)
metrics["adapter_state_after_set"] = adapter_state(pipe)
metrics["gpu_after_load"] = gpu_state()
metrics["shared_component_identity"] = {
"pipe_i2v_transformer_shared": bool(pipe_i2v is None or pipe_i2v.transformer is pipe.transformer),
"pipe_i2v_connectors_shared": bool(pipe_i2v is None or pipe_i2v.connectors is pipe.connectors),
"pipe_condition_transformer_shared": bool(pipe_condition is None or pipe_condition.transformer is pipe.transformer),
"pipe_condition_connectors_shared": bool(pipe_condition is None or pipe_condition.connectors is pipe.connectors),
"pipe_ic_transformer_shared": bool(pipe_ic is None or pipe_ic.transformer is pipe.transformer),
"pipe_ic_connectors_shared": bool(pipe_ic is None or pipe_ic.connectors is pipe.connectors),
}
metrics["status"] = "loaded_active"
return loaded, metrics
except Exception:
try:
pipe.disable_lora()
if adapter_names:
pipe.delete_adapters(adapter_names)
except Exception:
pass
raise
def cleanup_request_loras(
*,
pipe,
loaded,
metrics=None,
full_sft_profile: bool,
gpu_state: Callable[[], dict],
) -> dict:
metrics = metrics if isinstance(metrics, dict) else {}
if not loaded:
metrics.setdefault("cleanup_status", "not_needed")
return metrics
adapter_names = [str(item.get("adapter_name")) for item in loaded if item.get("adapter_name")]
metrics["gpu_before_cleanup"] = gpu_state()
metrics["cleanup_adapter_names"] = list(adapter_names)
metrics["internal_stage2_adapter_preserved"] = bool(full_sft_profile)
started = time.monotonic()
try:
pipe.disable_lora()
if adapter_names:
pipe.delete_adapters(adapter_names)
metrics["cleanup_status"] = "PASS"
finally:
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
metrics["cleanup_seconds"] = time.monotonic() - started
metrics["adapter_state_after_cleanup"] = adapter_state(pipe)
metrics["gpu_after_cleanup"] = gpu_state()
return metrics