"""SLA Turbo LoRA support for the diffusers MiniMax-H3 transformer. The SLA LoRA from lightx2v/Minimax-h3-Turbo-SLA uses the same PEFT/diffusers key format as lightx2v/Minimax-h3-Turbo: keys like ``transformer_blocks.0.attn.to_q.lora_A.default.weight``, rank 128, alpha 128 (scale = alpha/rank = 1.0). The LoRA is applied by folding ``scale * (B @ A)`` into the bf16 weights rather than as runtime wrappers, matching the approach in the reference MiniMax-H3 Turbo LoRA Space. ``H3_LORA`` selects the SLA file (``off`` skips it), ``H3_LORA_REPO`` picks the Hub repo, ``H3_LORA_STRENGTH`` scales the update (sharpness/artifact dial). """ from __future__ import annotations import os import torch SLA_REPO = os.environ.get("H3_LORA_REPO", "lightx2v/Minimax-h3-Turbo-SLA") SLA_FILE = os.environ.get("H3_LORA", "minimax_h3_fl2v_turbo_4step_v0.1_768p_sla_bf16.safetensors") # alpha=128, rank=128 => scale=1.0, matching set_adapters(weights=1.0) SLA_ALPHA = int(os.environ.get("H3_LORA_ALPHA", "128")) SLA_STRENGTH = float(os.environ.get("H3_LORA_STRENGTH", "1.0")) _PIPE_TRANSFORMER = None def _load_sla() -> dict: """Load the SLA LoRA from the Hub and return a spec dict.""" from huggingface_hub import hf_hub_download from safetensors.torch import load_file lora = load_file(hf_hub_download(SLA_REPO, SLA_FILE)) suffix_a, suffix_b = ".lora_A.default.weight", ".lora_B.default.weight" bases = sorted({key[: -len(suffix_a)] for key in lora if key.endswith(suffix_a)}) ranks = {lora[f"{name}{suffix_a}"].shape[0] for name in bases} if len(ranks) != 1: raise ValueError(f"Mixed LoRA ranks in {SLA_FILE}: {sorted(ranks)}") rank = ranks.pop() entries = [(f"{name}.weight", lora[f"{name}{suffix_a}"], lora[f"{name}{suffix_b}"]) for name in bases] scale = SLA_ALPHA / rank return { "label": f"{SLA_REPO}/{SLA_FILE}", "scale": scale * SLA_STRENGTH, "entries": entries, } def _apply(entries, params, sign: float) -> None: """Fold (sign * scale * B @ A) into the named parameters in-place.""" for key, a, b in entries: param = params.get(key) if param is None: raise KeyError(f"LoRA target `{key}` not found in the transformer") delta = sign * (b.to(torch.float32) @ a.to(torch.float32)) param.data = (param.data.float() + delta.to(param.device)).to(param.dtype) def apply_lora(transformer) -> str | None: """Load the SLA LoRA and fold it into ``transformer``'s bf16 weights. Returns a status line, or ``None`` when disabled. """ global _PIPE_TRANSFORMER _PIPE_TRANSFORMER = transformer if SLA_FILE.lower() in ("", "off", "none"): return None spec = _load_sla() params = dict(transformer.named_parameters()) _apply(spec["entries"], params, spec["scale"]) transformer._lora_state = {"active": "sla", "sets": {"sla": spec}} return ( f"SLA LoRA loaded: `{spec['label']}`, {len(spec['entries'])} weights, " f"scale {spec['scale']:.4f}" ) def set_active(transformer, name: str) -> str: """Switch the folded LoRA in place. Only 'sla' and 'off' are supported.""" state = getattr(transformer, "_lora_state", None) if state is None: return "off" name = name if name in state["sets"] else "off" if state["active"] == name: return name params = dict(transformer.named_parameters()) if state["active"] != "off": old = state["sets"][state["active"]] _apply(old["entries"], params, -old["scale"]) if name != "off": _apply(state["sets"][name]["entries"], params, state["sets"][name]["scale"]) state["active"] = name return name def available() -> list[str]: """The LoRA sets loaded at startup, plus 'off'.""" state = getattr(_PIPE_TRANSFORMER, "_lora_state", None) if _PIPE_TRANSFORMER is not None else None return sorted(state["sets"]) + ["off"] if state else ["off"]