LTX-2.5-diffusers-mini / ltx /model_loading.py
John6666's picture
Upload 39 files
e8b6587 verified
Raw
History Blame Contribute Delete
15.4 kB
from __future__ import annotations
import json
from pathlib import Path
from typing import Callable
import torch
from huggingface_hub import hf_hub_download
from transformers import (
AutoModelForImageTextToText,
AutoProcessor,
BitsAndBytesConfig as TransformersBitsAndBytesConfig,
)
from diffusers import (
BitsAndBytesConfig as DiffusersBitsAndBytesConfig,
GGUFQuantizationConfig,
LTX2Pipeline,
LTX2VideoTransformer3DModel,
)
from diffusers.quantizers import PipelineQuantizationConfig
LogFn = Callable[[str], None]
def maybe_enable_attention_backend(pipe, attention_backend: str) -> dict:
state = {"requested": attention_backend, "active": "sdpa", "status": "default"}
if attention_backend in {"", "sdpa", "default"}:
return state
if attention_backend not in {"flash3", "flash3_hub", "_flash_3_hub"}:
raise RuntimeError(f"Unsupported LTX25_ATTENTION_BACKEND={attention_backend!r}")
from kernels import get_kernel
get_kernel("kernels-community/flash-attn3", version=1)
pipe.transformer.set_attention_backend("_flash_3_hub")
state.update(status="enabled", active="_flash_3_hub")
return state
def component_config_path(model_dir: str | Path, subfolder: str) -> Path:
return Path(model_dir) / subfolder / "config.json"
def read_quantization_config(config_path: str | Path) -> dict | None:
path = Path(config_path)
if not path.exists():
raise FileNotFoundError(f"Component config not found: {path}")
data = json.loads(path.read_text(encoding="utf-8"))
value = data.get("quantization_config")
return value if isinstance(value, dict) else None
def quantization_kind(config: dict | None) -> str:
if not config:
return "unquantized"
quant_method = str(config.get("quant_method") or "").strip().lower()
load_in_4bit = bool(config.get("load_in_4bit") or config.get("_load_in_4bit"))
quant_type = str(config.get("bnb_4bit_quant_type") or "").strip().lower()
if (load_in_4bit or quant_method in {"bitsandbytes_4bit", "bitsandbytes"}) and quant_type == "nf4":
return "bitsandbytes_nf4"
if load_in_4bit or quant_method in {"bitsandbytes_4bit", "bitsandbytes"}:
return f"bitsandbytes_4bit:{quant_type or 'unknown'}"
return quant_method or "unknown_prequantized"
def nf4_config(component: str):
if component == "transformer":
return DiffusersBitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=False,
)
if component in {"text_encoder", "prompt_enhancer"}:
return TransformersBitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=False,
)
raise ValueError(f"Unknown NF4 component: {component}")
def quantization_decision(config_path: str | Path, component: str, policy: str):
qconfig = read_quantization_config(config_path)
kind = quantization_kind(qconfig)
if policy == "repo_native":
return None, f"repo-native ({kind})"
if kind == "unquantized":
return nf4_config(component), "on-load BitsAndBytes NF4 / BF16 compute"
if kind == "bitsandbytes_nf4":
return None, "repo-prequantized BitsAndBytes NF4"
raise RuntimeError(
f"{component} uses unsupported prequantized format {kind!r} under quantization policy 'nf4_auto'. "
"Use a BitsAndBytes NF4 checkpoint, an unquantized checkpoint, or explicitly opt into repo_native."
)
def remote_component_config(
repo_id: str,
subfolder: str | None,
revision: str | None,
*,
token: str | None,
) -> Path:
filename = f"{subfolder.strip('/')}/config.json" if subfolder else "config.json"
return Path(
hf_hub_download(
repo_id=repo_id,
filename=filename,
revision=revision or None,
token=token,
)
)
def load_prompt_enhancer_cpu(
record: dict,
*,
repo_id: str,
revision: str | None,
enabled: bool,
policy: str,
token: str | None,
log_fn: LogFn,
):
repo_id = str(repo_id or "").strip()
revision = str(revision or "").strip() or None
if not enabled:
record["effective"] = {"kind": "disabled", "reason": "disabled by space_config.py"}
return None, None
if not repo_id:
raise RuntimeError("PROMPT_ENHANCER_REPO_ID must not be empty when prompt enhancement is enabled.")
log_fn(
f"[D1R8P3] prompt enhancer CPU/NF4 prepare repo_id={repo_id} revision={revision!r}"
)
try:
import torchvision # noqa: F401
from transformers import Gemma4Processor # noqa: F401
except Exception as exc:
raise RuntimeError(
"Prompt enhancer Gemma4Processor vision preflight failed. "
"This Space requires torch==2.11.0 with torchvision==0.26.0. "
f"Underlying error: {type(exc).__name__}: {exc}"
) from exc
processor = AutoProcessor.from_pretrained(repo_id, revision=revision, token=token)
config_path = remote_component_config(repo_id, None, revision, token=token)
quant_config, quant_desc = quantization_decision(config_path, "prompt_enhancer", policy)
kwargs = {
"revision": revision,
"token": token,
"dtype": torch.bfloat16,
"device_map": {"": "cpu"},
"low_cpu_mem_usage": True,
}
if quant_config is not None:
kwargs["quantization_config"] = quant_config
model = AutoModelForImageTextToText.from_pretrained(repo_id, **kwargs)
model.eval()
record["effective"] = {
"kind": "dedicated_gemma4",
"repo_id": repo_id,
"revision": revision,
"quantization": quant_desc,
"dtype": "bfloat16 compute",
"residency": "startup RAM-ready / separate callback GPU-lazy",
"runtime_dependency": "Transformers only; no Unsloth package/runtime",
}
log_fn(f"[D1R8P3] prompt enhancer startup RAM-ready quantization={quant_desc}")
return model, processor
def load_transformer_override(
model_dir: str,
record: dict,
policy: str,
*,
repo_id: str,
path: str | None,
revision: str | None,
token: str | None,
log_fn: LogFn,
mark_fallback: Callable[[dict, str], None],
):
repo_id = str(repo_id or "").strip()
path = str(path or "").strip() or None
revision = str(revision or "").strip() or None
if not repo_id:
return None
log_fn(
f"[MODEL_OVERRIDE] transformer requested repo_id={repo_id} path={path!r} revision={revision!r}"
)
try:
if path and path.lower().endswith(".gguf"):
local = hf_hub_download(
repo_id=repo_id,
filename=path,
revision=revision,
token=token,
)
model = LTX2VideoTransformer3DModel.from_single_file(
local,
config=str(model_dir),
subfolder="transformer",
quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16),
dtype=torch.bfloat16,
)
record["effective"] = {
"kind": "override_gguf",
"repo_id": repo_id,
"path": path,
"revision": revision,
"quantization": "GGUF / BF16 compute",
}
log_fn("[MODEL_OVERRIDE] transformer override loaded as GGUF")
return model
config_path = remote_component_config(repo_id, path, revision, token=token)
quant_config, quant_desc = quantization_decision(config_path, "transformer", policy)
kwargs = {
"revision": revision,
"token": token,
"dtype": torch.bfloat16,
}
if path:
kwargs["subfolder"] = path
if quant_config is not None:
kwargs["quantization_config"] = quant_config
model = LTX2VideoTransformer3DModel.from_pretrained(repo_id, **kwargs)
record["effective"] = {
"kind": "override_pretrained",
"repo_id": repo_id,
"path": path,
"revision": revision,
"quantization": quant_desc,
}
log_fn(f"[MODEL_OVERRIDE] transformer override loaded quantization={quant_desc}")
return model
except Exception as exc:
reason = f"{type(exc).__name__}: {exc}"
log_fn(f"[MODEL_OVERRIDE] transformer override FAILED: {reason}")
mark_fallback(record, reason)
return None
def load_text_encoder_override(
record: dict,
policy: str,
*,
repo_id: str,
path: str | None,
revision: str | None,
token: str | None,
log_fn: LogFn,
mark_fallback: Callable[[dict, str], None],
):
repo_id = str(repo_id or "").strip()
path = str(path or "").strip() or None
revision = str(revision or "").strip() or None
if not repo_id:
return None
log_fn(
f"[MODEL_OVERRIDE] text_encoder requested repo_id={repo_id} path={path!r} revision={revision!r}"
)
try:
config_path = remote_component_config(repo_id, path, revision, token=token)
quant_config, quant_desc = quantization_decision(config_path, "text_encoder", policy)
kwargs = {
"revision": revision,
"token": token,
"dtype": torch.bfloat16,
}
if path:
kwargs["subfolder"] = path
if quant_config is not None:
kwargs["quantization_config"] = quant_config
model = AutoModelForImageTextToText.from_pretrained(repo_id, **kwargs)
record["effective"] = {
"kind": "override_pretrained_experimental",
"repo_id": repo_id,
"path": path,
"revision": revision,
"quantization": quant_desc,
}
log_fn(f"[MODEL_OVERRIDE] text_encoder override loaded quantization={quant_desc}")
return model
except Exception as exc:
reason = f"{type(exc).__name__}: {exc}"
log_fn(f"[MODEL_OVERRIDE] text_encoder override FAILED: {reason}")
mark_fallback(record, reason)
return None
def base_component_quantization(model_dir: str, component: str, policy: str):
config_path = component_config_path(model_dir, component)
return quantization_decision(config_path, component, policy)
def load_full_sft_transformer(
model_dir: str,
record: dict,
policy: str,
base_repo_id: str,
base_revision: str | None,
*,
path: str,
source_repo: str,
source_revision: str | None,
token: str | None,
runtime_profile: str,
):
path = str(path or "transformer_full").strip().strip("/")
use_base_snapshot = (
source_repo == str(base_repo_id).strip()
and source_revision == (str(base_revision or "").strip() or None)
)
if use_base_snapshot:
source_root = model_dir
config_path = component_config_path(model_dir, path)
kwargs = {"subfolder": path, "dtype": torch.bfloat16}
else:
source_root = source_repo
config_path = remote_component_config(source_repo, path, source_revision, token=token)
kwargs = {
"subfolder": path,
"revision": source_revision,
"token": token,
"dtype": torch.bfloat16,
}
quant_config, quant_desc = quantization_decision(config_path, "transformer", policy)
if quant_config is not None:
kwargs["quantization_config"] = quant_config
model = LTX2VideoTransformer3DModel.from_pretrained(source_root, **kwargs)
record["requested"] = {
"repo_id": source_repo,
"path": path,
"revision": source_revision,
}
record["effective"] = {
"kind": "full_sft_base_component",
"repo_id": source_repo,
"path": path,
"revision": source_revision,
"quantization": quant_desc,
"profile": runtime_profile,
"source_transport": "base_snapshot_component" if use_base_snapshot else "component_native_from_pretrained",
}
return model
def prepare_full_sft_stage2_lora(
model_dir: str,
base_repo_id: str,
base_revision: str | None,
*,
repo_id: str,
revision: str | None,
weight_name: str,
token: str | None,
) -> tuple[Path, str | None]:
weight_name = str(weight_name or "ltx-2.5-22b-distilled-lora-450-bf16.safetensors").strip()
if not weight_name:
raise RuntimeError("FULL_SFT_STAGE2_LORA_WEIGHT_NAME must not be empty in full_sft_nf4.")
same_base = (
repo_id == str(base_repo_id).strip()
and revision == (str(base_revision or "").strip() or None)
)
local = Path(model_dir) / weight_name if same_base else Path()
if not (same_base and local.is_file()):
local = Path(
hf_hub_download(
repo_id=repo_id,
filename=weight_name,
revision=revision,
token=token,
)
)
resolved_revision = None
parts = list(local.parts)
if "snapshots" in parts:
idx = parts.index("snapshots")
if idx + 1 < len(parts):
resolved_revision = parts[idx + 1]
return local, resolved_revision
def build_base_pipeline(
model_dir: str,
*,
transformer_override=None,
text_encoder_override=None,
policy: str = "nf4_auto",
auto_duration_enabled: bool,
):
quant_mapping = {}
component_quantization = {}
kwargs = {
"processor": None,
"prompt_enhancer": None,
"dtype": torch.bfloat16,
}
if not auto_duration_enabled:
kwargs["duration_head"] = None
if transformer_override is not None:
kwargs["transformer"] = transformer_override
else:
transformer_quant, desc = base_component_quantization(model_dir, "transformer", policy)
component_quantization["transformer"] = desc
if transformer_quant is not None:
quant_mapping["transformer"] = transformer_quant
if text_encoder_override is not None:
kwargs["text_encoder"] = text_encoder_override
else:
text_quant, desc = base_component_quantization(model_dir, "text_encoder", policy)
component_quantization["text_encoder"] = desc
if text_quant is not None:
quant_mapping["text_encoder"] = text_quant
if quant_mapping:
kwargs["quantization_config"] = PipelineQuantizationConfig(quant_mapping=quant_mapping)
pipe = LTX2Pipeline.from_pretrained(model_dir, **kwargs)
return pipe, component_quantization
def try_build_base_pipeline(
model_dir: str,
*,
transformer_override=None,
text_encoder_override=None,
policy: str = "nf4_auto",
auto_duration_enabled: bool,
):
try:
return (
build_base_pipeline(
model_dir,
transformer_override=transformer_override,
text_encoder_override=text_encoder_override,
policy=policy,
auto_duration_enabled=auto_duration_enabled,
),
None,
)
except Exception as exc:
return None, f"{type(exc).__name__}: {exc}"