"""Reversible Forge-compatible LoRA injection for the DiffSynth Anima pipeline.
The base model is never fused or mutated. Standard low-rank branches are added
with temporary forward hooks and removed in ``finally``. This is important for
a shared, long-lived ZeroGPU process where requests can use different LoRAs.
"""
from __future__ import annotations
from contextlib import contextmanager
from dataclasses import dataclass, field
import gc
import math
from pathlib import Path
import re
from typing import Iterable, Iterator, Mapping, Sequence
import torch
import torch.nn.functional as F
from safetensors.torch import load_file
PAIR_RE = re.compile(
r"^(?P.+)\.(?Plora_A|lora_B|lora_down|lora_up)(?:\.default)?\.weight$",
re.IGNORECASE,
)
UNSUPPORTED_MARKERS = (
"dora_scale",
"hada_w1",
"hada_w2",
"lokr_",
"oft_blocks",
"boft_",
"ia3_",
)
class LoRACompatibilityError(RuntimeError):
"""Raised when a LoRA cannot be mapped safely to the loaded Anima model."""
@dataclass(frozen=True)
class LoRASource:
path: str
strength: float = 1.0
label: str = ""
@property
def display_name(self) -> str:
return self.label or Path(self.path).name
@dataclass
class RawLoRAPair:
base: str
down: torch.Tensor
up: torch.Tensor
alpha: float | None
@dataclass
class SourceReport:
source: str
total_pairs: int = 0
matched_pairs: int = 0
unmatched: list[str] = field(default_factory=list)
invalid_shapes: list[str] = field(default_factory=list)
incomplete: list[str] = field(default_factory=list)
unsupported_keys: list[str] = field(default_factory=list)
targets: list[str] = field(default_factory=list)
@property
def match_ratio(self) -> float:
return self.matched_pairs / self.total_pairs if self.total_pairs else 0.0
@dataclass(frozen=True)
class Target:
scope: str
name: str
module: torch.nn.Module
@property
def canonical(self) -> str:
return f"{self.scope}:{self.name}"
def _normalise_alias(value: str) -> str:
value = value.strip().replace("\\", ".").replace("/", ".")
value = value.replace(".processor.", ".")
value = re.sub(r"\.+", ".", value).strip(".")
for prefix in ("module.", "base_model.model.", "base_model."):
while value.lower().startswith(prefix):
value = value[len(prefix) :]
return value.lower()
def _linear_like(module: torch.nn.Module) -> bool:
if isinstance(module, torch.nn.Embedding):
return False
weight = getattr(module, "weight", None)
if not isinstance(weight, torch.Tensor) or weight.ndim != 2:
return False
if isinstance(module, torch.nn.Linear):
return True
class_name = module.__class__.__name__.lower()
return "linear" in class_name or (
hasattr(module, "in_features") and hasattr(module, "out_features")
)
class TargetIndex:
"""Map DiffSynth, PEFT, Comfy and Forge layer names to live modules."""
def __init__(self, pipe) -> None:
self.targets: list[Target] = []
self.aliases: dict[str, list[Target]] = {}
self._build(pipe)
def _record(self, alias: str, target: Target) -> None:
key = _normalise_alias(alias)
if not key:
return
bucket = self.aliases.setdefault(key, [])
if all(existing.canonical != target.canonical for existing in bucket):
bucket.append(target)
def _add_dit(self, name: str, module: torch.nn.Module) -> None:
target = Target("dit", name, module)
self.targets.append(target)
underscored = name.replace(".", "_")
aliases = {
name,
f"dit.{name}",
f"transformer.{name}",
f"diffusion_model.{name}",
f"model.diffusion_model.{name}",
f"lora_unet_{underscored}",
f"lora_transformer_{underscored}",
f"lycoris_{underscored}",
}
# Forge moves this module into qwen3_06b at runtime, whereas DiffSynth
# retains it under the DiT. Both namespaces must resolve to one target.
if name.startswith("llm_adapter."):
aliases.update(
{
f"text_encoders.qwen3_06b.{name}",
f"qwen3_06b.{name}",
f"text_encoder.{name}",
f"lora_te_{underscored}",
f"lora_te1_{underscored}",
}
)
for alias in aliases:
self._record(alias, target)
def _add_text_encoder(self, name: str, module: torch.nn.Module) -> None:
target = Target("text_encoder", name, module)
self.targets.append(target)
underscored = name.replace(".", "_")
aliases = {
name,
f"text_encoder.{name}",
f"text_encoders.qwen3_06b.{name}",
f"qwen3_06b.{name}",
f"lora_te_{underscored}",
f"lora_te1_{underscored}",
}
if name.startswith("model.layers."):
rest = name[len("model.layers.") :]
aliases.update(
{
f"lora_te_layers_{rest.replace('.', '_')}",
f"lora_te1_layers_{rest.replace('.', '_')}",
}
)
for alias in aliases:
self._record(alias, target)
def _build(self, pipe) -> None:
for name, module in pipe.dit.named_modules():
if name and _linear_like(module):
self._add_dit(name, module)
for name, module in pipe.text_encoder.named_modules():
if name and _linear_like(module):
self._add_text_encoder(name, module)
@staticmethod
def _scope_hint(key: str, target: Target) -> bool:
key = key.lower()
text_hint = any(marker in key for marker in ("lora_te", "text_encoder", "text_encoders", "qwen3_06b"))
model_hint = any(marker in key for marker in ("lora_unet", "diffusion_model", "transformer"))
if text_hint:
return target.scope == "text_encoder" or target.name.startswith("llm_adapter.")
if model_hint:
return target.scope == "dit"
return True
def resolve(self, raw_base: str) -> Target | None:
key = _normalise_alias(raw_base)
direct = self.aliases.get(key, [])
if len(direct) == 1:
return direct[0]
if len(direct) > 1:
hinted = [target for target in direct if self._scope_hint(key, target)]
if len(hinted) == 1:
return hinted[0]
# Conservative unique-suffix fallback for uncommon wrapper prefixes.
candidates: list[Target] = []
for target in self.targets:
dotted = _normalise_alias(target.name)
underscored = dotted.replace(".", "_")
if (key.endswith(dotted) or key.endswith(underscored)) and self._scope_hint(key, target):
candidates.append(target)
unique = {target.canonical: target for target in candidates}
return next(iter(unique.values())) if len(unique) == 1 else None
def _alpha_for_base(state_dict: Mapping[str, torch.Tensor], base: str) -> float | None:
candidates = (
f"{base}.alpha",
f"{base}.lora_alpha",
f"{base}.alpha.default",
)
for key in candidates:
value = state_dict.get(key)
if value is None:
continue
if not isinstance(value, torch.Tensor) or value.numel() != 1:
continue
alpha = float(value.detach().float().cpu().item())
if math.isfinite(alpha):
return alpha
return None
def parse_lora_state_dict(
state_dict: Mapping[str, torch.Tensor],
) -> tuple[list[RawLoRAPair], list[str], list[str]]:
grouped: dict[str, dict[str, torch.Tensor]] = {}
unsupported: list[str] = []
for key, value in state_dict.items():
lowered = key.lower()
if any(marker in lowered for marker in UNSUPPORTED_MARKERS):
unsupported.append(key)
match = PAIR_RE.match(key)
if match is None or not isinstance(value, torch.Tensor):
continue
base = match.group("base")
tag = match.group("tag").lower()
role = "down" if tag in {"lora_a", "lora_down"} else "up"
grouped.setdefault(base, {})[role] = value.detach().cpu()
pairs: list[RawLoRAPair] = []
incomplete: list[str] = []
for base, tensors in grouped.items():
if "down" not in tensors or "up" not in tensors:
incomplete.append(base)
continue
pairs.append(
RawLoRAPair(
base=base,
down=tensors["down"],
up=tensors["up"],
alpha=_alpha_for_base(state_dict, base),
)
)
return pairs, incomplete, unsupported
def _as_matrix(tensor: torch.Tensor) -> torch.Tensor | None:
if tensor.ndim == 2:
return tensor.contiguous()
if tensor.ndim in (3, 4, 5) and all(size == 1 for size in tensor.shape[2:]):
return tensor.reshape(tensor.shape[0], tensor.shape[1]).contiguous()
return None
def _orient_pair(
pair: RawLoRAPair, target: Target
) -> tuple[torch.Tensor, torch.Tensor, int] | None:
down = _as_matrix(pair.down)
up = _as_matrix(pair.up)
if down is None or up is None:
return None
weight = getattr(target.module, "weight")
out_features, in_features = int(weight.shape[0]), int(weight.shape[1])
down_options = (down, down.t().contiguous())
up_options = (up, up.t().contiguous())
for down_candidate in down_options:
rank, in_dim = int(down_candidate.shape[0]), int(down_candidate.shape[1])
if in_dim != in_features:
continue
for up_candidate in up_options:
out_dim, up_rank = int(up_candidate.shape[0]), int(up_candidate.shape[1])
if out_dim == out_features and up_rank == rank:
return down_candidate, up_candidate, rank
return None
@dataclass
class _LinearBranch:
down_cpu: torch.Tensor
up_cpu: torch.Tensor
scale: float
source: str
base: str
_cache: dict[tuple[str, int | None, torch.dtype], tuple[torch.Tensor, torch.Tensor]] = field(default_factory=dict)
def _materialize(self, device: torch.device, dtype: torch.dtype) -> tuple[torch.Tensor, torch.Tensor]:
key = (device.type, device.index, dtype)
tensors = self._cache.get(key)
if tensors is None:
tensors = (
self.down_cpu.to(device=device, dtype=dtype, non_blocking=True),
self.up_cpu.to(device=device, dtype=dtype, non_blocking=True),
)
self._cache[key] = tensors
return tensors
def project(self, x: torch.Tensor) -> torch.Tensor:
dtype = x.dtype if x.is_floating_point() else torch.float32
down, up = self._materialize(x.device, dtype)
return F.linear(F.linear(x.to(dtype=dtype), down), up) * self.scale
def clear(self) -> None:
self._cache.clear()
class LoRAHookSession:
"""Apply one request's LoRAs, then cleanly remove every hook."""
def __init__(self, pipe, *, minimum_match_ratio: float = 0.5) -> None:
self.pipe = pipe
self.index = TargetIndex(pipe)
self.minimum_match_ratio = float(minimum_match_ratio)
self._handles: list[torch.utils.hooks.RemovableHandle] = []
self._branches: list[_LinearBranch] = []
@staticmethod
def _load(path: str) -> dict[str, torch.Tensor]:
suffix = Path(path).suffix.lower()
if suffix != ".safetensors":
raise LoRACompatibilityError(
f"Only .safetensors LoRAs are accepted; got {Path(path).name!r}."
)
return load_file(path, device="cpu")
def apply(self, sources: Sequence[LoRASource]) -> list[SourceReport]:
if self._handles:
raise RuntimeError("This LoRA session is already active")
grouped: dict[int, tuple[Target, list[_LinearBranch]]] = {}
reports: list[SourceReport] = []
for source in sources:
if not math.isfinite(source.strength):
raise LoRACompatibilityError(f"Invalid strength for {source.display_name}")
state_dict = self._load(source.path)
pairs, incomplete, unsupported = parse_lora_state_dict(state_dict)
report = SourceReport(
source=source.display_name,
total_pairs=len(pairs),
incomplete=incomplete[:20],
unsupported_keys=unsupported[:20],
)
if unsupported:
raise LoRACompatibilityError(
f"{source.display_name}: DoRA/LyCORIS/OFT-style tensors were detected. "
"This ZeroGPU runtime intentionally supports standard linear LoRA only."
)
if not pairs:
raise LoRACompatibilityError(
f"{source.display_name}: no complete lora_A/lora_B or lora_down/lora_up pairs were found."
)
for pair in pairs:
target = self.index.resolve(pair.base)
if target is None:
report.unmatched.append(pair.base)
continue
oriented = _orient_pair(pair, target)
if oriented is None:
report.invalid_shapes.append(pair.base)
continue
down, up, rank = oriented
alpha = float(rank) if pair.alpha is None else float(pair.alpha)
effective_scale = float(source.strength) * alpha / float(rank)
branch = _LinearBranch(
down_cpu=down,
up_cpu=up,
scale=effective_scale,
source=source.display_name,
base=pair.base,
)
module_key = id(target.module)
if module_key not in grouped:
grouped[module_key] = (target, [])
grouped[module_key][1].append(branch)
self._branches.append(branch)
report.matched_pairs += 1
if len(report.targets) < 30:
report.targets.append(target.canonical)
reports.append(report)
del state_dict
if report.matched_pairs == 0 or report.match_ratio < self.minimum_match_ratio:
preview = ", ".join((report.unmatched + report.invalid_shapes)[:5]) or "no resolvable targets"
raise LoRACompatibilityError(
f"{source.display_name}: matched {report.matched_pairs}/{report.total_pairs} LoRA pairs "
f"({report.match_ratio:.0%}); expected at least {self.minimum_match_ratio:.0%}. "
f"Examples: {preview}"
)
for target, branches in grouped.values():
def hook(module, inputs, output, *, _branches=tuple(branches)):
if not inputs or not isinstance(inputs[0], torch.Tensor):
raise LoRACompatibilityError(
f"LoRA target {module.__class__.__name__} received no tensor input"
)
if not isinstance(output, torch.Tensor):
raise LoRACompatibilityError(
f"LoRA target {module.__class__.__name__} returned a non-tensor output"
)
x = inputs[0]
delta: torch.Tensor | None = None
for branch in _branches:
update = branch.project(x)
delta = update if delta is None else delta + update
if delta is None:
return output
return output + delta.to(dtype=output.dtype)
self._handles.append(target.module.register_forward_hook(hook))
return reports
def clear(self) -> None:
for handle in reversed(self._handles):
handle.remove()
self._handles.clear()
for branch in self._branches:
branch.clear()
self._branches.clear()
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def __enter__(self) -> "LoRAHookSession":
return self
def __exit__(self, exc_type, exc, traceback) -> None:
self.clear()
@contextmanager
def temporary_loras(
pipe,
sources: Sequence[LoRASource],
*,
minimum_match_ratio: float = 0.5,
) -> Iterator[list[SourceReport]]:
session = LoRAHookSession(pipe, minimum_match_ratio=minimum_match_ratio)
try:
reports = session.apply(sources) if sources else []
yield reports
finally:
session.clear()
def format_reports(reports: Sequence[SourceReport]) -> str:
if not reports:
return "No LoRA applied."
lines: list[str] = []
for report in reports:
scopes = sorted({target.split(":", 1)[0] for target in report.targets})
scope_text = ", ".join(scopes) if scopes else "none"
line = (
f"- **{report.source}**: {report.matched_pairs}/{report.total_pairs} pairs matched "
f"({report.match_ratio:.0%}); scopes: `{scope_text}`"
)
if report.unmatched:
line += f"; unmatched: {len(report.unmatched)}"
if report.invalid_shapes:
line += f"; shape mismatches: {len(report.invalid_shapes)}"
lines.append(line)
return "\n".join(lines)