Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions backend/api_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -445,6 +445,10 @@ class LoraEntry(BaseModel):
scale: float = Field(default=1.0, ge=0.0, le=4.0)


# Local Distilled offering cap. API pipelines stay at 0 via multi_keyframe=False.
LOCAL_MULTI_KEYFRAME_MAX_COUNT = 10


class KeyframeInput(BaseModel):
"""CamelCase of LTXV keyframe-edit `{image_uri, frame_index, strength}`."""

Expand Down Expand Up @@ -1083,8 +1087,10 @@ def _check_selection_is_mutually_exclusive(self) -> "EnhancePromptRequest":
if has_last and not has_first:
raise ValueError("Last frame requires a first-frame image")
if has_keyframes:
if len(self.keyframes) > 5:
raise ValueError("You can place up to 5 keyframes")
if len(self.keyframes) > LOCAL_MULTI_KEYFRAME_MAX_COUNT:
raise ValueError(
f"You can place up to {LOCAL_MULTI_KEYFRAME_MAX_COUNT} keyframes"
)
indices = [keyframe.frameIndex for keyframe in self.keyframes]
if len(set(indices)) != len(indices):
raise ValueError("Keyframe frame indices must be unique")
Expand Down
1 change: 1 addition & 0 deletions backend/handlers/video_generation_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -407,6 +407,7 @@ def generate_video(
frame_rate=fps,
images=images,
output_path=str(output_path),
guide_all_images=bool(keyframe_images),
)
t_inference_end = time.perf_counter()
logger.info("[%s] Inference: %.2fs", gen_mode, t_inference_end - t_inference_start)
Expand Down
23 changes: 20 additions & 3 deletions backend/ltx2_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,16 @@
# ignores it there. Must be set before importing torch; setdefault lets an explicit env override.
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")

# Before any mps-sdpa import. Calibration benches the slow pyobjc backend at
# L<=2048 and can cache fused_min_bytes=None ("always stock"). That then also
# disables mpsgraph_zc, so video-length attention hits stock MPS SDPA, which
# materializes S×S (~47 GiB at 720p/5s) and OOMs / wedged-U-state
# (https://github.com/Lightricks/LTX-Desktop/issues/161). setdefault so a
# developer override still wins. Remove once mps-sdpa treats None as "use
# defaults" for large sequences, or calibrates against mpsgraph_zc / video L.
if sys.platform == "darwin":
os.environ.setdefault("MPS_SDPA_SKIP_CALIBRATION", "1")

import torch

# macOS: stage mps-sdpa's prebuilt zero-copy attention extension cache before mps-sdpa
Expand Down Expand Up @@ -227,19 +237,22 @@ def _resolve_app_data_dir() -> Path:
from state import RuntimeConfig, build_initial_state
from runtime_config.runtime_policy import LocalGenerationMode, decide_local_generation_mode
from server_utils.model_layout_migration import migrate_legacy_models_layout
from services.gpu_info.gpu_info_impl import GpuInfoImpl
from services.gpu_info.gpu_info_impl import GpuInfoImpl, platform_label

migrate_legacy_models_layout(APP_DATA_DIR)

LTX_API_BASE_URL = "https://api.ltx.video"


def _resolve_local_generations_mode() -> LocalGenerationMode:
from runtime_config.accelerator import accelerator_backend

gpu_info = GpuInfoImpl()
system = platform.system()
cuda_available = gpu_info.get_cuda_available()
mps_available = gpu_info.get_mps_available()
vram_gb = gpu_info.get_vram_total_gb()
fp8_capable = accelerator_backend() == "cuda"
# On Darwin there's no discrete VRAM (unified memory), so gate on *available* RAM,
# not total — total overstates real headroom once the OS/Electron/app are running.
# See GpuInfoImpl.get_available_ram_gb.
Expand All @@ -252,16 +265,18 @@ def _resolve_local_generations_mode() -> LocalGenerationMode:
vram_gb=vram_gb,
mps_available=mps_available,
ram_gb=available_ram_gb,
fp8_capable=fp8_capable,
)
logger.info(
"Runtime policy local_generations_mode=%s (system=%s cuda_available=%s mps_available=%s "
"vram_gb=%s available_ram_gb=%s)",
"vram_gb=%s available_ram_gb=%s fp8_capable=%s)",
mode,
system,
cuda_available,
mps_available,
vram_gb,
available_ram_gb,
fp8_capable,
)
return mode

Expand Down Expand Up @@ -332,7 +347,7 @@ def log_hardware_info() -> None:
gpu_info = gpu.get_gpu_info()
vram_gb = gpu_info["vram"] // 1024 if gpu_info["vram"] else 0

logger.info(f"Platform: {platform.system()} ({platform.machine()})")
logger.info(f"Platform: {platform_label()}")
logger.info(f"Device: {DEVICE} | Dtype: {DTYPE}")
gpu_line = f"GPU: {gpu_info['name']} | VRAM: {vram_gb} GB"
# On Apple Silicon there's no discrete VRAM — the figure above is total unified
Expand All @@ -344,6 +359,8 @@ def log_hardware_info() -> None:
"LTX 2.5 decode uses eager SDPA on Mac (no Triton; slower than Linux/Windows)."
)
logger.info(gpu_line)
from runtime_config.accelerator import accelerator_backend
logger.info(f"Accelerator: {accelerator_backend()} | HIP: {getattr(torch.version, 'hip', None)}")
logger.info(f"SageAttention: {'enabled' if use_sage_attention else 'disabled'}")
logger.info(f"Python: {sys.version.split()[0]} | Torch: {torch.__version__}")

Expand Down
111 changes: 109 additions & 2 deletions backend/mps_prebuilt_ext.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,13 +26,23 @@
import os
import sys
from pathlib import Path
from typing import Any
from typing import Any, Callable, cast

from server_utils.units import gib

logger = logging.getLogger(__name__)

_EXT_NAME = "mps_sdpa_zc_ext"
_GUARD_ATTR = "_ltx_coerces_none_fused_min"

# Mirror of mps-sdpa's M4-tuned defaults (backends/_calibrate._DEFAULT_THRESHOLDS).
# Used when calibration stored fused_min_bytes=None ("always stock") — that sentinel
# is a speed conclusion at L<=2048, not a memory-safe choice for video attention.
_FALLBACK_FUSED_MIN_BYTES = {
"bf16": 4 * 1024**2,
"fp16": 4 * 1024**2,
"fp32": 8 * 1024**2,
}


def _prebuilt_so() -> Path | None:
Expand All @@ -44,19 +54,94 @@ def _prebuilt_so() -> Path | None:
return so if so.is_file() else None


def coerce_mps_sdpa_thresholds(
thresholds: dict[str, Any],
*,
defaults: dict[str, int] | None = None,
) -> dict[str, Any]:
"""Replace fused_min_bytes=None ("always stock") with safe defaults.
mps-sdpa v0.2.0 calibration benches pyobjc mpsgraph vs stock at L<=2048. If
pyobjc never wins by 5%, it caches ``null``. Both ``mpsgraph_zc`` and
``mpsgraph`` then treat ``fused_min is None`` as "use stock for every
shape", including ~14k-token video self-attention. Stock MPS SDPA
materializes the S×S score matrix (~47 GiB) and OOMs
(https://github.com/Lightricks/LTX-Desktop/issues/161).
"""
fused_raw = thresholds.get("fused_min_bytes")
if not isinstance(fused_raw, dict):
return thresholds
fused = cast(dict[str, Any], fused_raw)
replacements = defaults or _FALLBACK_FUSED_MIN_BYTES
out_fused: dict[str, Any] = dict(fused)
changed = False
for key, default in replacements.items():
if out_fused.get(key) is None:
out_fused[key] = default
changed = True
if not changed:
return thresholds
return {**thresholds, "fused_min_bytes": out_fused}


def install_mps_sdpa_threshold_guard() -> None:
"""Wrap mps-sdpa ``get_thresholds()`` so None fused_min never reaches dispatch.
No-op off Darwin or when mps-sdpa isn't installed. Idempotent.
"""
if sys.platform != "darwin":
return
try:
from mps_sdpa.backends import _calibrate # noqa: PLC0415 # type: ignore[reportMissingModuleSource]
except ImportError:
return

current = cast(
Callable[[], dict[str, Any]],
_calibrate.get_thresholds, # pyright: ignore[reportUnknownMemberType]
)
if getattr(current, _GUARD_ATTR, False):
return

def _set_cached(value: dict[str, Any]) -> None:
setattr(_calibrate, "_cached_thresholds", value)

def _guarded() -> dict[str, Any]:
raw = current()
thresholds = coerce_mps_sdpa_thresholds(raw)
if thresholds is not raw:
_set_cached(thresholds)
logger.warning(
"mps-sdpa fused_min_bytes had None (always-stock); using defaults %s "
"(raw=%s). Stock MPS SDPA OOMs video attention (LTX-Desktop#161).",
thresholds.get("fused_min_bytes"),
raw.get("fused_min_bytes"),
)
return thresholds

setattr(_guarded, _GUARD_ATTR, True)
_calibrate.get_thresholds = _guarded # pyright: ignore[reportUnknownMemberType]
cached = getattr(_calibrate, "_cached_thresholds", None)
if isinstance(cached, dict):
_set_cached(coerce_mps_sdpa_thresholds(cast(dict[str, Any], cached)))
logger.info("mps_prebuilt_ext: installed fused_min_bytes=None → defaults guard")


def setup_prebuilt_mps_extension() -> None:
if sys.platform != "darwin":
return

so = _prebuilt_so()
if so is None:
logger.info("mps_prebuilt_ext: no bundled prebuilt .so; torch will JIT-build if a compiler is present")
install_mps_sdpa_threshold_guard()
return

try:
from torch.utils import cpp_extension as _cppext # noqa: PLC0415
except Exception:
logger.warning("mps_prebuilt_ext: torch.utils.cpp_extension unavailable; skipping", exc_info=True)
install_mps_sdpa_threshold_guard()
return

_orig_load = _cppext.load # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
Expand Down Expand Up @@ -90,6 +175,7 @@ def _patched_load(name: Any = None, *args: Any, **kwargs: Any) -> Any: # noqa:

_cppext.load = _patched_load # pyright: ignore[reportUnknownMemberType]
logger.info("mps_prebuilt_ext: cpp_extension.load patched to direct-import bundled %s", so)
install_mps_sdpa_threshold_guard()


def log_mps_backend_status() -> None:
Expand All @@ -103,8 +189,10 @@ def log_mps_backend_status() -> None:
"""
if sys.platform != "darwin":
return
install_mps_sdpa_threshold_guard()
try:
from mps_sdpa import api # noqa: PLC0415 # type: ignore[reportMissingModuleSource]
from mps_sdpa.backends import _calibrate # noqa: PLC0415 # type: ignore[reportMissingModuleSource]

st = api.backend_status(backend="auto", device="mps")
logger.info(
Expand All @@ -117,6 +205,16 @@ def log_mps_backend_status() -> None:
"attention may leak Metal memory. unavailable=%s",
st["picked"], st.get("unavailable"),
)
get_thresholds = cast(
Callable[[], dict[str, Any]],
_calibrate.get_thresholds, # pyright: ignore[reportUnknownMemberType]
)
thresholds = get_thresholds()
logger.info(
"mps-sdpa fused_min_bytes=%s calibrated=%s",
thresholds.get("fused_min_bytes"),
thresholds.get("calibrated"),
)
except Exception:
logger.warning("mps-sdpa: could not determine active attention backend", exc_info=True)

Expand All @@ -132,7 +230,13 @@ def _mps_sdpa_call_stats() -> str:
from mps_sdpa import api # noqa: PLC0415 # type: ignore[reportMissingModuleSource]

stats = api.get_call_stats()
return " ".join(f"{k}={v}" for k, v in sorted(stats.items())) or "(no calls yet)"
parts = [" ".join(f"{k}={v}" for k, v in sorted(stats.items()))] if stats else []
get_fallback = getattr(api, "get_fallback_stats", None)
fallback = get_fallback() if callable(get_fallback) else None
if isinstance(fallback, dict) and fallback:
fb_items = cast(dict[str, Any], fallback)
parts.append("fb:" + " ".join(f"{k}={v}" for k, v in sorted(fb_items.items())))
return " ".join(parts) or "(no calls yet)"
except Exception:
return ""

Expand All @@ -147,6 +251,9 @@ def reset_mps_sdpa_stats() -> None:
from mps_sdpa import api # noqa: PLC0415 # type: ignore[reportMissingModuleSource]

api.reset_call_stats()
reset_fallback = getattr(api, "reset_fallback_stats", None)
if callable(reset_fallback):
reset_fallback()
except Exception:
pass

Expand Down
32 changes: 32 additions & 0 deletions backend/runtime_config/accelerator.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
"""Accelerator backend detection.

ROCm PyTorch reports itself through the same `torch.cuda` interface as NVIDIA
CUDA builds (`torch.cuda.is_available()`, `device.type == "cuda"`, etc.), so
code that checks `device.type == "cuda"` to mean "this is an NVIDIA GPU" is
wrong under ROCm. `torch.version.hip` is the actual discriminator: it is set
on ROCm builds and `None` on CUDA builds.

Originally contributed by boxwrench in
https://github.com/Lightricks/LTX-Desktop/pull/160
"""

from __future__ import annotations

from typing import Literal

import torch

AcceleratorBackend = Literal["rocm", "cuda", "mps", "cpu"]


def accelerator_backend() -> AcceleratorBackend:
if getattr(torch.version, "hip", None):
return "rocm"

if torch.cuda.is_available():
return "cuda"

if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return "mps"

return "cpu"
11 changes: 8 additions & 3 deletions backend/runtime_config/ltx_capabilities.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,12 @@
from dataclasses import dataclass, replace
from typing import Literal, assert_never

from api_types import LTXLocalModelId, LTXVideoGenPipeline, LTXVideoGenResolution
from api_types import (
LOCAL_MULTI_KEYFRAME_MAX_COUNT,
LTXLocalModelId,
LTXVideoGenPipeline,
LTXVideoGenResolution,
)

LtxCapabilityFeature = Literal[
"t2v",
Expand Down Expand Up @@ -80,7 +85,7 @@ class ApiOfferingCapabilities(LtxOfferingCapabilities):
retake=True,
extend=True,
multi_keyframe=True,
multi_keyframe_max_count=5,
multi_keyframe_max_count=LOCAL_MULTI_KEYFRAME_MAX_COUNT,
user_loras=True,
camera_motion=True,
auto_duration=False,
Expand All @@ -99,7 +104,7 @@ class ApiOfferingCapabilities(LtxOfferingCapabilities):
retake=False,
extend=False,
multi_keyframe=True,
multi_keyframe_max_count=5,
multi_keyframe_max_count=LOCAL_MULTI_KEYFRAME_MAX_COUNT,
user_loras=True,
camera_motion=True,
auto_duration=True,
Expand Down
Loading
Loading