diff --git a/backend/api_types.py b/backend/api_types.py index 79b439165..8f004fff1 100644 --- a/backend/api_types.py +++ b/backend/api_types.py @@ -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}`.""" @@ -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") diff --git a/backend/handlers/video_generation_handler.py b/backend/handlers/video_generation_handler.py index c7ed128dc..58634eb47 100644 --- a/backend/handlers/video_generation_handler.py +++ b/backend/handlers/video_generation_handler.py @@ -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) diff --git a/backend/ltx2_server.py b/backend/ltx2_server.py index f250fb7c1..6031d0e2a 100644 --- a/backend/ltx2_server.py +++ b/backend/ltx2_server.py @@ -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 @@ -227,7 +237,7 @@ 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) @@ -235,11 +245,14 @@ def _resolve_app_data_dir() -> Path: 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. @@ -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 @@ -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 @@ -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__}") diff --git a/backend/mps_prebuilt_ext.py b/backend/mps_prebuilt_ext.py index e36514e66..06f0f2850 100644 --- a/backend/mps_prebuilt_ext.py +++ b/backend/mps_prebuilt_ext.py @@ -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: @@ -44,6 +54,79 @@ 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 @@ -51,12 +134,14 @@ def setup_prebuilt_mps_extension() -> None: 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] @@ -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: @@ -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( @@ -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) @@ -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 "" @@ -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 diff --git a/backend/runtime_config/accelerator.py b/backend/runtime_config/accelerator.py new file mode 100644 index 000000000..06a9c579e --- /dev/null +++ b/backend/runtime_config/accelerator.py @@ -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" diff --git a/backend/runtime_config/ltx_capabilities.py b/backend/runtime_config/ltx_capabilities.py index d38f20c27..436995e96 100644 --- a/backend/runtime_config/ltx_capabilities.py +++ b/backend/runtime_config/ltx_capabilities.py @@ -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", @@ -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, @@ -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, diff --git a/backend/runtime_config/runtime_policy.py b/backend/runtime_config/runtime_policy.py index 91bbcce41..6d9ed90ab 100644 --- a/backend/runtime_config/runtime_policy.py +++ b/backend/runtime_config/runtime_policy.py @@ -48,6 +48,7 @@ def decide_local_generation_mode( vram_gb: int | None, mps_available: bool = False, ram_gb: int | None = None, + fp8_capable: bool = True, ) -> LocalGenerationMode: """Pick the local-generation mode for this runtime. @@ -59,11 +60,15 @@ def decide_local_generation_mode( On CUDA (Windows/Linux) the memory figure is discrete (total) VRAM, and "full_models_loading" (>=31 GB) holds the fp8-halved (~23 GB) transformer - resident. On Apple Silicon (Darwin) there is no discrete VRAM — the GPU shares - system RAM — so ``ram_gb`` is the *available* (free) RAM, not total: total RAM - overstates real headroom once the OS/Electron/app are already running. Gated on - MPS being available (i.e. Apple Silicon, not an Intel Mac, which has no MPS - backend and stays unsupported). + resident. Pass ``fp8_capable=False`` (ROCm today — see + ``runtime_config.accelerator.accelerator_backend``) to stay on the streaming + path regardless of VRAM: without fp8 the full ~42-46 GB bf16 transformer + would try to stay resident at that floor and OOM. On Apple Silicon (Darwin) + there is no discrete VRAM — the GPU shares system RAM — so ``ram_gb`` is the + *available* (free) RAM, not total: total RAM overstates real headroom once + the OS/Electron/app are already running. Gated on MPS being available + (i.e. Apple Silicon, not an Intel Mac, which has no MPS backend and stays + unsupported). ``fp8_capable`` is ignored on Darwin (no fp8 path on MPS). Darwin's "full_models_loading" (>=``DARWIN_FULL_RESIDENT_FLOOR_GB``) is NOT the same regime as CUDA's: MPS has no fp8 path, so it holds the *full* ~46 GB bf16 @@ -92,6 +97,15 @@ def decide_local_generation_mode( return "unsupported" if vram_gb < 15: return "unsupported" + # full_models_loading's >=31 GB floor assumes the fp8-halved (~23 GB) transformer + # (see module docstring). Without fp8 (ROCm today — see + # runtime_config.accelerator.accelerator_backend), holding the full bf16 + # (~42-46 GB) transformer resident instead would OOM at this floor, so stay on + # the streaming path regardless of VRAM until a real bf16-full-resident floor is + # established on non-CUDA hardware. + # Originally contributed by boxwrench in https://github.com/Lightricks/LTX-Desktop/pull/160 + if not fp8_capable: + return "streaming_models_loading" if vram_gb < 31: return "streaming_models_loading" return "full_models_loading" diff --git a/backend/services/fast_video_pipeline/distilled_keyframe_guiding.py b/backend/services/fast_video_pipeline/distilled_keyframe_guiding.py new file mode 100644 index 000000000..4fcf0a7a3 --- /dev/null +++ b/backend/services/fast_video_pipeline/distilled_keyframe_guiding.py @@ -0,0 +1,28 @@ +"""Per-call Distilled helper swap for multi-keyframe generate. + +DistilledPipeline hardcodes ``combined_image_conditionings`` (frame 0 replaces +the latent). MKF interpolation needs ``image_conditionings_by_adding_guiding_latent`` +instead. This is scoped to one Fast generate so first/last i2v keeps replace-at-0. +""" + +from __future__ import annotations + +import threading +from collections.abc import Iterator +from contextlib import contextmanager + +_SWAP_LOCK = threading.Lock() + + +@contextmanager +def distilled_keyframe_guiding() -> Iterator[None]: + import ltx_pipelines.distilled as distilled + from ltx_pipelines.utils.helpers import image_conditionings_by_adding_guiding_latent + + with _SWAP_LOCK: + original = distilled.combined_image_conditionings + distilled.combined_image_conditionings = image_conditionings_by_adding_guiding_latent + try: + yield + finally: + distilled.combined_image_conditionings = original diff --git a/backend/services/fast_video_pipeline/fast_video_pipeline.py b/backend/services/fast_video_pipeline/fast_video_pipeline.py index 1b2777f4f..bc1ac682d 100644 --- a/backend/services/fast_video_pipeline/fast_video_pipeline.py +++ b/backend/services/fast_video_pipeline/fast_video_pipeline.py @@ -39,6 +39,8 @@ def generate( frame_rate: float, images: list[ImageConditioningInput], output_path: str, + *, + guide_all_images: bool = False, ) -> None: ... diff --git a/backend/services/fast_video_pipeline/ltx_fast_video_pipeline.py b/backend/services/fast_video_pipeline/ltx_fast_video_pipeline.py index d06649ad6..36ace6e9a 100644 --- a/backend/services/fast_video_pipeline/ltx_fast_video_pipeline.py +++ b/backend/services/fast_video_pipeline/ltx_fast_video_pipeline.py @@ -107,9 +107,14 @@ def _run_inference( frame_rate: float, images: list[ImageConditioningInput], tiling_config: PipelineTilingType, + *, + guide_all_images: bool = False, ) -> tuple[torch.Tensor | Iterator[torch.Tensor], AudioOrNone, int, TilingConfigType | None]: + from contextlib import nullcontext + from ltx_pipelines.utils.args import ImageConditioningInput as _LtxImageInput from ltx_pipelines.utils.types import AutoDuration + from services.fast_video_pipeline.distilled_keyframe_guiding import distilled_keyframe_guiding pipeline_num_frames: int | AutoDuration = ( AutoDuration(min_seconds=num_frames.min_seconds, max_seconds=num_frames.max_seconds) @@ -117,16 +122,18 @@ def _run_inference( else num_frames ) - video, audio, resolved_frames, resolved_tiling = self.pipeline( - prompt=prompt, - seed=seed, - height=height, - width=width, - num_frames=pipeline_num_frames, - frame_rate=frame_rate, - images=[_LtxImageInput(img.path, img.frame_idx, img.strength) for img in images], - tiling_config=tiling_config, - ) + ctx = distilled_keyframe_guiding() if guide_all_images else nullcontext() + with ctx: + video, audio, resolved_frames, resolved_tiling = self.pipeline( + prompt=prompt, + seed=seed, + height=height, + width=width, + num_frames=pipeline_num_frames, + frame_rate=frame_rate, + images=[_LtxImageInput(img.path, img.frame_idx, img.strength) for img in images], + tiling_config=tiling_config, + ) return video, audio, resolved_frames, resolved_tiling @torch.inference_mode() @@ -140,6 +147,8 @@ def generate( frame_rate: float, images: list[ImageConditioningInput], output_path: str, + *, + guide_all_images: bool = False, ) -> None: video, audio, resolved_frames, resolved_tiling = self._run_inference( prompt=prompt, @@ -150,6 +159,7 @@ def generate( frame_rate=frame_rate, images=images, tiling_config=auto_tiling_config(), + guide_all_images=guide_all_images, ) chunks = video_chunks_number(resolved_frames, resolved_tiling) encode_video_output(video=video, audio=audio, fps=int(frame_rate), output_path=output_path, video_chunks_number_value=chunks) diff --git a/backend/services/gpu_info/gpu_info_impl.py b/backend/services/gpu_info/gpu_info_impl.py index b4e26032a..783ecdf9e 100644 --- a/backend/services/gpu_info/gpu_info_impl.py +++ b/backend/services/gpu_info/gpu_info_impl.py @@ -17,6 +17,18 @@ logger = logging.getLogger(__name__) +def platform_label() -> str: + """OS/arch for startup logs. Includes macOS version on Darwin (15+ is required + for fused MPSGraph SDPA; 14 falls back to stock and OOMs on video lengths).""" + system = platform.system() + machine = platform.machine() + if system == "Darwin": + mac_ver = platform.mac_ver()[0] + if mac_ver: + return f"{system} {mac_ver} ({machine})" + return f"{system} ({machine})" + + class _CudaDeviceProperties(Protocol): total_memory: int diff --git a/backend/services/services_utils.py b/backend/services/services_utils.py index e98784da2..42b34be00 100644 --- a/backend/services/services_utils.py +++ b/backend/services/services_utils.py @@ -90,7 +90,16 @@ def effective_edit_steps(num_inference_steps: int, strength: float) -> int: def device_supports_fp8(device: str | torch.device | object | None) -> bool: - return get_device_type(device) == "cuda" + # ROCm PyTorch also reports device.type == "cuda" (see runtime_config.accelerator), + # so this must not be a plain device-type check: ROCm has no fp8_cast kernel support + # here yet, and would otherwise be silently misidentified as CUDA/NVIDIA. + # Originally contributed by boxwrench in https://github.com/Lightricks/LTX-Desktop/pull/160 + if get_device_type(device) != "cuda": + return False + + from runtime_config.accelerator import accelerator_backend + + return accelerator_backend() == "cuda" def sync_device(device: str | torch.device | object | None) -> None: diff --git a/backend/tests/fakes/services.py b/backend/tests/fakes/services.py index 6ffd8a8bf..df00bd1d8 100644 --- a/backend/tests/fakes/services.py +++ b/backend/tests/fakes/services.py @@ -725,6 +725,8 @@ def generate( frame_rate: float, images: list[ImageConditioningInput], output_path: str, + *, + guide_all_images: bool = False, ) -> None: payload = { "prompt": prompt, @@ -735,6 +737,7 @@ def generate( "frame_rate": frame_rate, "images": images, "output_path": output_path, + "guide_all_images": guide_all_images, } if self.inference_steps: from services.generation_interrupt import raise_if_requested diff --git a/backend/tests/test_accelerator.py b/backend/tests/test_accelerator.py new file mode 100644 index 000000000..96a693b93 --- /dev/null +++ b/backend/tests/test_accelerator.py @@ -0,0 +1,69 @@ +"""Tests for ROCm vs CUDA accelerator detection.""" + +from __future__ import annotations + +from types import SimpleNamespace + +from runtime_config.accelerator import accelerator_backend +from services.services_utils import device_supports_fp8 + + +def test_accelerator_backend_rocm_when_hip_is_set(monkeypatch) -> None: + import runtime_config.accelerator as accel + + monkeypatch.setattr(accel.torch.version, "hip", "6.3.0", raising=False) + monkeypatch.setattr(accel.torch.cuda, "is_available", lambda: True) + assert accelerator_backend() == "rocm" + + +def test_accelerator_backend_cuda_when_hip_is_none(monkeypatch) -> None: + import runtime_config.accelerator as accel + + monkeypatch.setattr(accel.torch.version, "hip", None, raising=False) + monkeypatch.setattr(accel.torch.cuda, "is_available", lambda: True) + assert accelerator_backend() == "cuda" + + +def test_accelerator_backend_mps_when_no_cuda(monkeypatch) -> None: + import runtime_config.accelerator as accel + + monkeypatch.setattr(accel.torch.version, "hip", None, raising=False) + monkeypatch.setattr(accel.torch.cuda, "is_available", lambda: False) + mps = getattr(accel.torch.backends, "mps", None) + if mps is None: + monkeypatch.setattr( + accel.torch.backends, "mps", SimpleNamespace(is_available=lambda: True), raising=False + ) + else: + monkeypatch.setattr(mps, "is_available", lambda: True) + assert accelerator_backend() == "mps" + + +def test_accelerator_backend_cpu_fallback(monkeypatch) -> None: + import runtime_config.accelerator as accel + + monkeypatch.setattr(accel.torch.version, "hip", None, raising=False) + monkeypatch.setattr(accel.torch.cuda, "is_available", lambda: False) + monkeypatch.setattr(accel.torch.backends.mps, "is_available", lambda: False) + assert accelerator_backend() == "cpu" + + +def test_device_supports_fp8_false_on_mps() -> None: + assert device_supports_fp8("mps") is False + assert device_supports_fp8(SimpleNamespace(type="mps")) is False + + +def test_device_supports_fp8_false_on_cpu() -> None: + assert device_supports_fp8("cpu") is False + + +def test_device_supports_fp8_true_on_cuda_not_rocm(monkeypatch) -> None: + monkeypatch.setattr("runtime_config.accelerator.accelerator_backend", lambda: "cuda") + assert device_supports_fp8("cuda") is True + assert device_supports_fp8(SimpleNamespace(type="cuda")) is True + + +def test_device_supports_fp8_false_on_rocm_cuda_device(monkeypatch) -> None: + monkeypatch.setattr("runtime_config.accelerator.accelerator_backend", lambda: "rocm") + assert device_supports_fp8("cuda") is False + assert device_supports_fp8(SimpleNamespace(type="cuda")) is False diff --git a/backend/tests/test_distilled_keyframe_guiding.py b/backend/tests/test_distilled_keyframe_guiding.py new file mode 100644 index 000000000..1d2f4c186 --- /dev/null +++ b/backend/tests/test_distilled_keyframe_guiding.py @@ -0,0 +1,25 @@ +"""CPU-only checks for the Distilled MKF guiding-helper swap.""" + +from __future__ import annotations + +import pytest + +import ltx_pipelines.distilled as distilled +from ltx_pipelines.utils.helpers import image_conditionings_by_adding_guiding_latent +from services.fast_video_pipeline.distilled_keyframe_guiding import distilled_keyframe_guiding + + +def test_guiding_context_swaps_distilled_combined_helper() -> None: + original = distilled.combined_image_conditionings + with distilled_keyframe_guiding(): + assert distilled.combined_image_conditionings is image_conditionings_by_adding_guiding_latent + assert distilled.combined_image_conditionings is original + + +def test_guiding_context_restores_helper_after_exception() -> None: + original = distilled.combined_image_conditionings + with pytest.raises(RuntimeError, match="swap-failed"): + with distilled_keyframe_guiding(): + assert distilled.combined_image_conditionings is image_conditionings_by_adding_guiding_latent + raise RuntimeError("swap-failed") + assert distilled.combined_image_conditionings is original diff --git a/backend/tests/test_generation.py b/backend/tests/test_generation.py index 651712073..d9da81e0f 100644 --- a/backend/tests/test_generation.py +++ b/backend/tests/test_generation.py @@ -10,7 +10,11 @@ import pytest from _routes._errors import HTTPError -from api_types import GenerateImageRequest, GenerateVideoRequest +from api_types import ( + LOCAL_MULTI_KEYFRAME_MAX_COUNT, + GenerateImageRequest, + GenerateVideoRequest, +) from frame_math import AutoDurationSpec, compute_num_frames from runtime_config.model_download_specs import delete_cp_path, get_ltx_model_spec, resolve_model_path from services import generation_interrupt @@ -389,7 +393,9 @@ def test_i2v_last_frame_sends_two_conditionings( ) assert r.status_code == 200 - images = fake_services.fast_video_pipeline.generate_calls[0]["images"] + call = fake_services.fast_video_pipeline.generate_calls[0] + images = call["images"] + assert call["guide_all_images"] is False assert len(images) == 2 assert images[0].frame_idx == 0 assert images[0].strength == 1.0 @@ -474,7 +480,7 @@ def test_keyframes_cap_is_enforced(self, client): **_T2V_JSON, "keyframes": [ {"imagePath": f"/tmp/kf-{index}.png", "frameIndex": index} - for index in range(6) + for index in range(LOCAL_MULTI_KEYFRAME_MAX_COUNT + 1) ], }, ) @@ -482,9 +488,23 @@ def test_keyframes_cap_is_enforced(self, client): r, status_code=422, code="INVALID_VIDEO_GENERATION_SPEC", - message="You can place up to 5 keyframes", + message=f"You can place up to {LOCAL_MULTI_KEYFRAME_MAX_COUNT} keyframes", + ) + + def test_keyframes_cap_allows_exact_limit(self, client): + r = client.post( + "/api/generate", + json={ + **_T2V_JSON, + "keyframes": [ + {"imagePath": f"/tmp/kf-{index}.png", "frameIndex": index} + for index in range(LOCAL_MULTI_KEYFRAME_MAX_COUNT) + ], + }, ) + assert_http_error(r, status_code=409, code="NO_DOWNLOADED_LTX_MODEL") + def test_keyframes_send_conditionings_at_requested_frames( self, client, test_state, fake_services, create_fake_model_files, make_test_image, tmp_path ): @@ -507,7 +527,9 @@ def test_keyframes_send_conditionings_at_requested_frames( ) assert r.status_code == 200 - images = fake_services.fast_video_pipeline.generate_calls[0]["images"] + call = fake_services.fast_video_pipeline.generate_calls[0] + images = call["images"] + assert call["guide_all_images"] is True assert [(image.frame_idx, image.strength) for image in images] == [(0, 1.0), (80, 1.0)] assert images[0].path != images[1].path diff --git a/backend/tests/test_ltx_capabilities.py b/backend/tests/test_ltx_capabilities.py index 3ffd50f49..117f8b083 100644 --- a/backend/tests/test_ltx_capabilities.py +++ b/backend/tests/test_ltx_capabilities.py @@ -4,6 +4,7 @@ import pytest +from api_types import LOCAL_MULTI_KEYFRAME_MAX_COUNT from runtime_config.ltx_capabilities import ( LocalOfferingCapabilities, api_caps, @@ -55,7 +56,7 @@ def test_local_2_5_allows_ic_lora_and_user_loras(): assert supports(caps, "retake") is False assert supports(caps, "extend") is False assert supports(caps, "multi_keyframe") is True - assert caps.multi_keyframe_max_count == 5 + assert caps.multi_keyframe_max_count == LOCAL_MULTI_KEYFRAME_MAX_COUNT def test_local_2_3_allows_ic_lora_user_loras_retake(): @@ -65,7 +66,7 @@ def test_local_2_3_allows_ic_lora_user_loras_retake(): assert supports(caps, "retake") is True assert supports(caps, "extend") is True assert supports(caps, "multi_keyframe") is True - assert caps.multi_keyframe_max_count == 5 + assert caps.multi_keyframe_max_count == LOCAL_MULTI_KEYFRAME_MAX_COUNT assert supports(caps, "auto_duration") is False diff --git a/backend/tests/test_mps_prebuilt_ext.py b/backend/tests/test_mps_prebuilt_ext.py index e422f4f2a..34abb67b5 100644 --- a/backend/tests/test_mps_prebuilt_ext.py +++ b/backend/tests/test_mps_prebuilt_ext.py @@ -9,6 +9,8 @@ from __future__ import annotations from pathlib import Path +import sys +import types import mps_prebuilt_ext @@ -54,3 +56,57 @@ def test_mps_memory_sample_none_off_darwin(monkeypatch) -> None: def test_reset_mps_sdpa_stats_noop_off_darwin(monkeypatch) -> None: monkeypatch.setattr(mps_prebuilt_ext.sys, "platform", "linux") mps_prebuilt_ext.reset_mps_sdpa_stats() # must not raise + + +def test_coerce_replaces_none_fused_min_with_defaults() -> None: + raw: dict = {"fused_min_bytes": {"bf16": None, "fp16": 1, "fp32": None}, "calibrated": True} + out = mps_prebuilt_ext.coerce_mps_sdpa_thresholds(raw) + assert out["fused_min_bytes"]["bf16"] == 4 * 1024**2 + assert out["fused_min_bytes"]["fp16"] == 1 + assert out["fused_min_bytes"]["fp32"] == 8 * 1024**2 + assert raw["fused_min_bytes"]["bf16"] is None + + +def test_coerce_is_noop_when_all_fused_min_are_set() -> None: + raw: dict = {"fused_min_bytes": {"bf16": 8, "fp16": 8, "fp32": 16}, "calibrated": True} + assert mps_prebuilt_ext.coerce_mps_sdpa_thresholds(raw) is raw + + +def test_coerce_passthrough_when_fused_min_missing() -> None: + raw: dict = {"calibrated": False} + assert mps_prebuilt_ext.coerce_mps_sdpa_thresholds(raw) is raw + + +def test_install_threshold_guard_noop_off_darwin(monkeypatch) -> None: + monkeypatch.setattr(mps_prebuilt_ext.sys, "platform", "linux") + mps_prebuilt_ext.install_mps_sdpa_threshold_guard() # must not raise + + +def test_threshold_guard_coerces_none_on_get_thresholds(monkeypatch) -> None: + monkeypatch.setattr(mps_prebuilt_ext.sys, "platform", "darwin") + unsafe: dict = {"fused_min_bytes": {"bf16": None, "fp16": None, "fp32": None}, "calibrated": True} + + fake_cal = types.SimpleNamespace(get_thresholds=lambda: unsafe, _cached_thresholds=unsafe) + fake_backends = types.ModuleType("mps_sdpa.backends") + fake_backends._calibrate = fake_cal # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, "mps_sdpa", types.ModuleType("mps_sdpa")) + monkeypatch.setitem(sys.modules, "mps_sdpa.backends", fake_backends) + + mps_prebuilt_ext.install_mps_sdpa_threshold_guard() + out = fake_cal.get_thresholds() + assert out["fused_min_bytes"]["bf16"] == 4 * 1024**2 + assert fake_cal._cached_thresholds["fused_min_bytes"]["bf16"] == 4 * 1024**2 + + +def test_mps_sdpa_call_stats_includes_fallback_reasons(monkeypatch) -> None: + fake_api = types.ModuleType("mps_sdpa.api") + fake_api.get_call_stats = lambda: {"stock_fallback": 12} # type: ignore[attr-defined] + fake_api.get_fallback_stats = lambda: {"short-seq": 12} # type: ignore[attr-defined] + fake_pkg = types.ModuleType("mps_sdpa") + fake_pkg.api = fake_api # type: ignore[attr-defined] + monkeypatch.setitem(sys.modules, "mps_sdpa", fake_pkg) + monkeypatch.setitem(sys.modules, "mps_sdpa.api", fake_api) + + text = mps_prebuilt_ext._mps_sdpa_call_stats() + assert "stock_fallback=12" in text + assert "fb:short-seq=12" in text diff --git a/backend/tests/test_platform_label.py b/backend/tests/test_platform_label.py new file mode 100644 index 000000000..faf3e7dc3 --- /dev/null +++ b/backend/tests/test_platform_label.py @@ -0,0 +1,26 @@ +"""Tests for startup platform labeling (macOS version in logs).""" + +from __future__ import annotations + +from services.gpu_info.gpu_info_impl import platform_label + + +def test_platform_label_includes_macos_version(monkeypatch) -> None: + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.system", lambda: "Darwin") + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.machine", lambda: "arm64") + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.mac_ver", lambda: ("15.6.1", ("", "", ""), "arm64")) + assert platform_label() == "Darwin 15.6.1 (arm64)" + + +def test_platform_label_omits_empty_macos_version(monkeypatch) -> None: + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.system", lambda: "Darwin") + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.machine", lambda: "arm64") + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.mac_ver", lambda: ("", ("", "", ""), "arm64")) + assert platform_label() == "Darwin (arm64)" + + +def test_platform_label_non_darwin_has_no_mac_ver(monkeypatch) -> None: + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.system", lambda: "Linux") + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.machine", lambda: "x86_64") + monkeypatch.setattr("services.gpu_info.gpu_info_impl.platform.mac_ver", lambda: ("", ("", "", ""), "")) + assert platform_label() == "Linux (x86_64)" diff --git a/backend/tests/test_prompt_enhancement.py b/backend/tests/test_prompt_enhancement.py index 453ce2c6f..f70abafa3 100644 --- a/backend/tests/test_prompt_enhancement.py +++ b/backend/tests/test_prompt_enhancement.py @@ -11,6 +11,7 @@ IcLoraCatalogItem, InputSpec, InstructionSection, + LOCAL_MULTI_KEYFRAME_MAX_COUNT, LoraCatalogItem, PromptTemplatePlaceholder, PromptTemplateSpec, @@ -231,7 +232,10 @@ def test_keyframes_cap_is_enforced_on_enhance(self, client): "/api/enhance-prompt", json={ "prompt": "walk the hall", - "keyframes": [{"imagePath": f"/tmp/kf-{index}.png", "frameIndex": index} for index in range(6)], + "keyframes": [ + {"imagePath": f"/tmp/kf-{index}.png", "frameIndex": index} + for index in range(LOCAL_MULTI_KEYFRAME_MAX_COUNT + 1) + ], }, ) assert r.status_code == 422 diff --git a/backend/tests/test_runtime_policy_decision.py b/backend/tests/test_runtime_policy_decision.py index 5371a5425..397450edb 100644 --- a/backend/tests/test_runtime_policy_decision.py +++ b/backend/tests/test_runtime_policy_decision.py @@ -106,6 +106,65 @@ def test_linux_full_loading_range() -> None: assert decide_local_generation_mode(system="Linux", cuda_available=True, vram_gb=31) == "full_models_loading" +def test_linux_without_fp8_still_unsupported_below_streaming_floor() -> None: + assert ( + decide_local_generation_mode( + system="Linux", cuda_available=True, vram_gb=14, fp8_capable=False + ) + == "unsupported" + ) + + +def test_linux_without_fp8_streams_even_above_full_floor() -> None: + """ROCm reports as CUDA but has no fp8; the 31 GB floor assumes an fp8 transformer.""" + assert ( + decide_local_generation_mode( + system="Linux", cuda_available=True, vram_gb=31, fp8_capable=False + ) + == "streaming_models_loading" + ) + assert ( + decide_local_generation_mode( + system="Linux", cuda_available=True, vram_gb=96, fp8_capable=False + ) + == "streaming_models_loading" + ) + + +def test_windows_without_fp8_streams_even_above_full_floor() -> None: + assert ( + decide_local_generation_mode( + system="Windows", cuda_available=True, vram_gb=31, fp8_capable=False + ) + == "streaming_models_loading" + ) + + +def test_fp8_capable_default_preserves_cuda_full_loading() -> None: + assert decide_local_generation_mode(system="Linux", cuda_available=True, vram_gb=31) == "full_models_loading" + assert ( + decide_local_generation_mode( + system="Linux", cuda_available=True, vram_gb=31, fp8_capable=True + ) + == "full_models_loading" + ) + + +def test_darwin_ignores_fp8_capable() -> None: + assert ( + decide_local_generation_mode( + system="Darwin", cuda_available=False, vram_gb=None, mps_available=True, ram_gb=48, fp8_capable=False + ) + == "streaming_models_loading" + ) + assert ( + decide_local_generation_mode( + system="Darwin", cuda_available=False, vram_gb=None, mps_available=True, ram_gb=85, fp8_capable=False + ) + == "full_models_loading" + ) + + def test_other_systems_fail_closed() -> None: assert decide_local_generation_mode(system="FreeBSD", cuda_available=True, vram_gb=48) == "unsupported" diff --git a/electron/python-backend.ts b/electron/python-backend.ts index 40394cc91..9bd4db8fd 100644 --- a/electron/python-backend.ts +++ b/electron/python-backend.ts @@ -336,6 +336,11 @@ export async function startPythonBackend(): Promise { // can't shadow PATH entries for backend subprocesses on Windows/Linux. ...(process.platform === 'darwin' ? { PATH: `${path.dirname(pythonPath)}${path.delimiter}${process.env.PATH ?? ''}`, + // Skip mps-sdpa's import-time microbench: it can cache fused_min_bytes=None + // ("always stock"), which routes video-length attention through stock MPS + // SDPA and OOMs (https://github.com/Lightricks/LTX-Desktop/issues/161). + // Honor an explicit parent value (matches setdefault in ltx2_server.py). + MPS_SDPA_SKIP_CALIBRATION: process.env.MPS_SDPA_SKIP_CALIBRATION ?? '1', } : {}), // Only pass LTX_PORT when the developer explicitly set it ...(process.env.LTX_PORT ? { LTX_PORT: process.env.LTX_PORT } : {}), diff --git a/frontend/components/KeyframeStrengthRail.tsx b/frontend/components/KeyframeStrengthRail.tsx new file mode 100644 index 000000000..cd035785f --- /dev/null +++ b/frontend/components/KeyframeStrengthRail.tsx @@ -0,0 +1,85 @@ +import { useRef } from 'react' +import { + formatKeyframeStrength, + nudgeKeyframeStrength, + strengthFromPointer, +} from '../lib/keyframe-strength' + +interface KeyframeStrengthRailProps { + strength: number + label: string + onStrengthChange: (strength: number) => void +} + +export function KeyframeStrengthRail({ + strength, + label, + onStrengthChange, +}: KeyframeStrengthRailProps) { + const railRef = useRef(null) + const draggingRef = useRef(false) + const clamped = Math.min(1, Math.max(0, strength)) + const percent = formatKeyframeStrength(strength) + + const applyFromClientY = (clientY: number) => { + const rect = railRef.current?.getBoundingClientRect() + if (!rect) return + onStrengthChange(strengthFromPointer(clientY, rect)) + } + + return ( +
{ + event.preventDefault() + event.stopPropagation() + draggingRef.current = true + event.currentTarget.setPointerCapture(event.pointerId) + applyFromClientY(event.clientY) + }} + onPointerMove={(event) => { + if (!draggingRef.current) return + event.stopPropagation() + applyFromClientY(event.clientY) + }} + onPointerUp={(event) => { + draggingRef.current = false + event.stopPropagation() + }} + onPointerCancel={() => { + draggingRef.current = false + }} + onLostPointerCapture={() => { + draggingRef.current = false + }} + onKeyDown={(event) => { + if (event.key !== 'ArrowUp' && event.key !== 'ArrowDown') return + event.preventDefault() + event.stopPropagation() + onStrengthChange(nudgeKeyframeStrength(strength, event.key === 'ArrowUp' ? 1 : -1)) + }} + > +
+
+
+
+
+ ) +} diff --git a/frontend/components/KeyframeTimeline.tsx b/frontend/components/KeyframeTimeline.tsx index efbd400d6..eddb38843 100644 --- a/frontend/components/KeyframeTimeline.tsx +++ b/frontend/components/KeyframeTimeline.tsx @@ -1,6 +1,10 @@ import { useLayoutEffect, useRef, useState } from 'react' -import { Trash2 } from 'lucide-react' +import { RefreshCw, Trash2 } from 'lucide-react' import { applyTimecode, nudgeKeyframe } from '../lib/keyframe-controls' +import { + formatKeyframeStrength, + nudgeKeyframeStrength, +} from '../lib/keyframe-strength' import { findNearestFreeFrameIndex, formatKeyframeTimecode, @@ -11,6 +15,7 @@ import { } from '../lib/keyframe-timeline' import type { KeyframeItem } from '../lib/multi-keyframe' import { pathToFileUrl } from '../lib/file-url' +import { KeyframeStrengthRail } from './KeyframeStrengthRail' interface KeyframeTimelineProps { keyframes: readonly KeyframeItem[] @@ -20,6 +25,7 @@ interface KeyframeTimelineProps { onPlayheadChange: (frameIndex: number) => void onDragFrameChange?: (drag: DraggedFrame | null) => void onFrameChange: (id: string, frameIndex: number) => void + onStrengthChange: (id: string, strength: number) => void onReplaceRequest: (id: string) => void onDelete: (id: string) => void onImagesDrop: (dataTransfer: DataTransfer, replaceId: string | null) => void @@ -29,7 +35,6 @@ interface DragState { id: string frameIndex: number pointerId: number - startClientX: number } export function KeyframeTimeline({ @@ -40,12 +45,12 @@ export function KeyframeTimeline({ onPlayheadChange, onDragFrameChange, onFrameChange, + onStrengthChange, onReplaceRequest, onDelete, onImagesDrop, }: KeyframeTimelineProps) { const trackRef = useRef(null) - const suppressMarkerClickRef = useRef(false) const dragRef = useRef(null) const onDragFrameChangeRef = useRef(onDragFrameChange) onDragFrameChangeRef.current = onDragFrameChange @@ -95,7 +100,6 @@ export function KeyframeTimeline({ const finishDrag = (event: React.PointerEvent) => { const activeDrag = dragRef.current if (!activeDrag || activeDrag.pointerId !== event.pointerId) return - suppressMarkerClickRef.current = Math.abs(event.clientX - activeDrag.startClientX) > 3 const otherKeyframes = keyframes.filter((keyframe) => keyframe.id !== activeDrag.id) const frameIndex = findNearestFreeFrameIndex(otherKeyframes, frameAtPointer(event.clientX), lastFrame) if (frameIndex !== null) onFrameChange(activeDrag.id, frameIndex) @@ -139,6 +143,7 @@ export function KeyframeTimeline({ {displayed.map((keyframe) => { const markerFrame = keyframe.frameIndex const source = keyframes.find((item) => item.id === keyframe.id) ?? keyframe + const timecode = formatKeyframeTimecode(markerFrame, fps) return (
- + onStrengthChange(keyframe.id, strength)} + /> +
+ + +
{ setTimecodeDraft({ id: keyframe.id, value: event.currentTarget.value }) }} @@ -230,7 +258,7 @@ export function KeyframeTimeline({ event.stopPropagation() if (event.key === 'Enter') event.currentTarget.blur() if (event.key === 'Escape') { - event.currentTarget.value = formatKeyframeTimecode(markerFrame, fps) + event.currentTarget.value = timecode event.currentTarget.blur() } }} diff --git a/frontend/components/MultiKeyframePanel.tsx b/frontend/components/MultiKeyframePanel.tsx index 33f30ace0..29e6ee536 100644 --- a/frontend/components/MultiKeyframePanel.tsx +++ b/frontend/components/MultiKeyframePanel.tsx @@ -104,6 +104,11 @@ export function MultiKeyframePanel({ ))) onPlayheadChange(frameIndex) }} + onStrengthChange={(id, strength) => { + onChange(keyframes.map((keyframe) => ( + keyframe.id === id ? { ...keyframe, strength } : keyframe + ))) + }} onReplaceRequest={(id) => openFilePicker(id)} onDelete={(id) => onChange(keyframes.filter((keyframe) => keyframe.id !== id))} onImagesDrop={(dataTransfer, replaceId) => { diff --git a/frontend/lib/keyframe-strength.test.ts b/frontend/lib/keyframe-strength.test.ts new file mode 100644 index 000000000..8fbf2725c --- /dev/null +++ b/frontend/lib/keyframe-strength.test.ts @@ -0,0 +1,53 @@ +import assert from 'node:assert/strict' +import { describe, it } from 'node:test' +import { + DEFAULT_KEYFRAME_STRENGTH, + formatKeyframeStrength, + nudgeKeyframeStrength, + strengthFromPointer, +} from './keyframe-strength.ts' + +describe('strengthFromPointer', () => { + const rail = { top: 10, height: 100 } + + it('maps the top of the rail to full strength', () => { + assert.equal(strengthFromPointer(10, rail), 1) + }) + + it('maps the bottom of the rail to zero strength', () => { + assert.equal(strengthFromPointer(110, rail), 0) + }) + + it('maps the midpoint to 0.5', () => { + assert.equal(strengthFromPointer(60, rail), 0.5) + }) + + it('clamps above the rail to 1 and below it to 0', () => { + assert.equal(strengthFromPointer(0, rail), 1) + assert.equal(strengthFromPointer(200, rail), 0) + }) + + it('returns the new-still default when the rail has no height', () => { + assert.equal(strengthFromPointer(10, { top: 10, height: 0 }), DEFAULT_KEYFRAME_STRENGTH) + }) +}) + +describe('nudgeKeyframeStrength', () => { + it('steps by five percent', () => { + assert.equal(nudgeKeyframeStrength(0.7, 1), 0.75) + assert.equal(nudgeKeyframeStrength(0.7, -1), 0.65) + }) + + it('clamps at 0 and 1', () => { + assert.equal(nudgeKeyframeStrength(0.02, -1), 0) + assert.equal(nudgeKeyframeStrength(0.98, 1), 1) + }) +}) + +describe('formatKeyframeStrength', () => { + it('renders owned strength as a percent', () => { + assert.equal(formatKeyframeStrength(0.7), '70%') + assert.equal(formatKeyframeStrength(0), '0%') + assert.equal(formatKeyframeStrength(1), '100%') + }) +}) diff --git a/frontend/lib/keyframe-strength.ts b/frontend/lib/keyframe-strength.ts index f3f5651f8..00eecb317 100644 --- a/frontend/lib/keyframe-strength.ts +++ b/frontend/lib/keyframe-strength.ts @@ -1,8 +1,45 @@ -/** Full i2v lock. Matches the backend KeyframeInput default; the UI owns the value. */ -export const DEFAULT_KEYFRAME_STRENGTH = 1 +/** Newly placed GenSpace stills. Loose enough for Distilled all-guide interpolation. */ +export const DEFAULT_KEYFRAME_STRENGTH = 0.7 + +/** Omitted/invalid persist and HTTP. Matches the backend KeyframeInput Field default. */ +export const MISSING_KEYFRAME_STRENGTH = 1 + +/** ArrowUp / ArrowDown and the strength rail step in 5% increments. */ +export const KEYFRAME_STRENGTH_STEP = 0.05 /** Persist, restore, and HTTP all go through this so a bad float cannot 422 or fail project parse. */ export function clampKeyframeStrength(value: unknown): number { - if (typeof value !== 'number' || !Number.isFinite(value)) return DEFAULT_KEYFRAME_STRENGTH + if (typeof value !== 'number' || !Number.isFinite(value)) return MISSING_KEYFRAME_STRENGTH return Math.min(1, Math.max(0, value)) } + +function clampUnit(value: number): number { + return Math.min(1, Math.max(0, value)) +} + +function roundStrength(value: number): number { + return Math.round(clampUnit(value) * 100) / 100 +} + +/** Top of the rail is 1 (full lock), bottom is 0 (no lock). */ +export function strengthFromPointer( + clientY: number, + railRect: { top: number; height: number }, +): number { + if (!(railRect.height > 0) || !Number.isFinite(clientY)) return DEFAULT_KEYFRAME_STRENGTH + return roundStrength(1 - (clientY - railRect.top) / railRect.height) +} + +export function nudgeKeyframeStrength(strength: number, direction: -1 | 1): number { + const current = typeof strength === 'number' && Number.isFinite(strength) + ? clampUnit(strength) + : DEFAULT_KEYFRAME_STRENGTH + return roundStrength(current + direction * KEYFRAME_STRENGTH_STEP) +} + +export function formatKeyframeStrength(strength: number): string { + const current = typeof strength === 'number' && Number.isFinite(strength) + ? clampUnit(strength) + : DEFAULT_KEYFRAME_STRENGTH + return `${Math.round(current * 100)}%` +} diff --git a/frontend/lib/multi-keyframe.test.ts b/frontend/lib/multi-keyframe.test.ts index f55b18de7..eb7ad6749 100644 --- a/frontend/lib/multi-keyframe.test.ts +++ b/frontend/lib/multi-keyframe.test.ts @@ -6,6 +6,7 @@ import { appendKeyframePaths, applyKeyframeImagePaths, DEFAULT_KEYFRAME_STRENGTH, + MISSING_KEYFRAME_STRENGTH, fromPersistedKeyframes, toPersistedKeyframes, videoGenerationModeFromInputs, @@ -152,13 +153,14 @@ describe('persisted keyframes', () => { assert.deepEqual(restored, [item('id-0', '/opening.png', 0, 0.7)]) }) - it('fills missing persisted strength with the default lock', () => { + it('fills missing persisted strength with a full lock, not the new-still default', () => { let nextId = 0 const restored = fromPersistedKeyframes( [{ path: '/opening.png', frameIndex: 0 }], () => `id-${nextId++}`, ) - assert.deepEqual(restored, [item('id-0', '/opening.png', 0)]) + assert.deepEqual(restored, [item('id-0', '/opening.png', 0, MISSING_KEYFRAME_STRENGTH)]) + assert.notEqual(MISSING_KEYFRAME_STRENGTH, DEFAULT_KEYFRAME_STRENGTH) }) it('persists and restores a zero lock instead of treating it as missing', () => { @@ -266,7 +268,7 @@ describe('persistedKeyframeSchema', () => { it('defaults missing strength to a full lock', () => { assert.deepEqual( persistedKeyframeSchema.parse({ path: '/opening.png', frameIndex: 0 }), - { path: '/opening.png', frameIndex: 0, strength: DEFAULT_KEYFRAME_STRENGTH }, + { path: '/opening.png', frameIndex: 0, strength: MISSING_KEYFRAME_STRENGTH }, ) }) diff --git a/frontend/lib/multi-keyframe.ts b/frontend/lib/multi-keyframe.ts index 8159f5dce..c6490453a 100644 --- a/frontend/lib/multi-keyframe.ts +++ b/frontend/lib/multi-keyframe.ts @@ -1,7 +1,14 @@ import { clampKeyframeStrength, DEFAULT_KEYFRAME_STRENGTH } from './keyframe-strength.ts' import { pickFreeFrameIndex } from './keyframe-timeline.ts' -export { clampKeyframeStrength, DEFAULT_KEYFRAME_STRENGTH } from './keyframe-strength.ts' +export { + clampKeyframeStrength, + DEFAULT_KEYFRAME_STRENGTH, + MISSING_KEYFRAME_STRENGTH, +} from './keyframe-strength.ts' + +/** Local Distilled cap. Must match backend LOCAL_MULTI_KEYFRAME_MAX_COUNT. API stays 0. */ +export const LOCAL_MULTI_KEYFRAME_MAX_COUNT = 10 export interface KeyframeItem { id: string diff --git a/frontend/views/GenSpace.tsx b/frontend/views/GenSpace.tsx index 31b638a24..e2ec46ddf 100644 --- a/frontend/views/GenSpace.tsx +++ b/frontend/views/GenSpace.tsx @@ -84,7 +84,15 @@ import { modeOptionValues, type GenSpaceMode, } from '../lib/genspace-multi-keyframe' -import { applyKeyframeImagePaths, enhanceKeyframesPayload, fromPersistedKeyframes, toPersistedKeyframes, videoGenerationModeFromInputs, type KeyframeItem } from '../lib/multi-keyframe' +import { + applyKeyframeImagePaths, + enhanceKeyframesPayload, + fromPersistedKeyframes, + LOCAL_MULTI_KEYFRAME_MAX_COUNT, + toPersistedKeyframes, + videoGenerationModeFromInputs, + type KeyframeItem, +} from '../lib/multi-keyframe' import { lastFrameFromDuration, previewKeyframeForPlayhead, retimeKeyframesForSettings, sameDraggedFrame, type DraggedFrame } from '../lib/keyframe-timeline' import { GenSpaceFilterEmptyState } from './genspace/GenSpaceFilterEmptyState' import { GenSpaceGalleryToolbar } from './genspace/GenSpaceGalleryToolbar' @@ -1489,7 +1497,8 @@ export function GenSpace() { RETAKE_EXTEND_MODELS[0], ) const canUseUserLoras = isLocalMode && Boolean(localCaps?.user_loras) - const multiKeyframeMaxCount = localCaps?.multi_keyframe_max_count ?? 5 + const multiKeyframeMaxCount = + localCaps?.multi_keyframe_max_count ?? LOCAL_MULTI_KEYFRAME_MAX_COUNT const canUseIcLora = !forceApiGenerations && Boolean(localCaps?.ic_lora) const canUseRetake = isLocalMode ? Boolean(localCaps?.retake) : Boolean(apiCaps?.retake) const canUseExtend = isLocalMode ? Boolean(localCaps?.extend) : Boolean(apiCaps?.extend) diff --git a/package.json b/package.json index 04076e72a..ce8e9d1c9 100644 --- a/package.json +++ b/package.json @@ -1,6 +1,6 @@ { "name": "ltx-desktop", - "version": "1.2.6", + "version": "1.2.7", "description": "LTX-2 Video Generation - Desktop App", "type": "module", "main": "dist-electron/main.js", @@ -26,7 +26,7 @@ "openapi:generate": "pnpm openapi:export && pnpm openapi:types", "openapi:check": "pnpm openapi:generate && git diff --exit-code -- frontend/generated/backend-openapi.json frontend/generated/backend-openapi.ts", "backend:test": "cd backend && uv sync --frozen --extra test --extra dev && uv run pytest -v --tb=short", - "scripts:test": "node --test scripts/copy-koffi-native.test.cjs && node --experimental-strip-types --test frontend/lib/build-generate-video-body.test.ts frontend/lib/format.test.ts frontend/lib/genspace-layout.test.ts frontend/lib/genspace-last-frame.test.ts frontend/lib/genspace-multi-keyframe.test.ts frontend/lib/keyframe-controls.test.ts frontend/lib/keyframe-drop.test.ts frontend/lib/keyframe-timeline.test.ts frontend/lib/multi-keyframe.test.ts frontend/lib/fixed-menu-position.test.ts frontend/lib/enhance-gemini-key.test.ts electron/free-disk-space.test.ts", + "scripts:test": "node --test scripts/copy-koffi-native.test.cjs && node --experimental-strip-types --test frontend/lib/build-generate-video-body.test.ts frontend/lib/format.test.ts frontend/lib/genspace-layout.test.ts frontend/lib/genspace-last-frame.test.ts frontend/lib/genspace-multi-keyframe.test.ts frontend/lib/keyframe-controls.test.ts frontend/lib/keyframe-drop.test.ts frontend/lib/keyframe-strength.test.ts frontend/lib/keyframe-timeline.test.ts frontend/lib/multi-keyframe.test.ts frontend/lib/fixed-menu-position.test.ts frontend/lib/enhance-gemini-key.test.ts electron/free-disk-space.test.ts", "build": "node scripts/run-script.js scripts/local-build", "build:skip-python": "node scripts/run-script.js scripts/local-build --skip-python", "build:fast": "node scripts/run-script.js scripts/local-build --unpack --skip-python",