diff --git a/.gitignore b/.gitignore index 4f9ea8e73..3fb318acc 100644 --- a/.gitignore +++ b/.gitignore @@ -84,6 +84,10 @@ backend/generated/ # Build artifacts release/ python-embed/ +# Generated at build time (sha256 of `uv export`); bundled into the app so the +# first-run downloader can match the hosted python-embed archive. Regenerated +# each build — must not be committed, or it would go stale vs. the lockfile. +python-deps-hash.txt # Downloaded during CI for Windows NSIS installer resources/vc_redist.x64.exe diff --git a/FORK.md b/FORK.md new file mode 100644 index 000000000..338008015 --- /dev/null +++ b/FORK.md @@ -0,0 +1,401 @@ +# Fork Maintenance Guide + +This fork of [Lightricks/LTX-Desktop](https://github.com/Lightricks/LTX-Desktop) carries +substantial custom work. This document exists so that **merging a new upstream release is a +checklist, not an archaeology expedition.** + +When you hit a merge conflict, find the file in the map below to learn which feature owns it +and what must survive. + +Last updated: 2026-07-21 · Fork base: upstream **LTX Desktop 1.1.0** +(`sync-public-preview/2026-07-19`, merged in `cd9395e`) + +Branches: `feature/upstream-1.1.0` is the current line. +`feature/21-9-plus-save-video` is frozen at `5b5efbc` as the pre-1.1.0 backup — +falling back to it **requires** `cd backend && uv sync --extra test --extra dev`, +or you'll run 1.0.5 code against 1.1.0 libraries. + +--- + +## 1. Fork shape (why merges are survivable) + +| | Count | +|---|---| +| New files we added (never conflict) | 108 | +| Upstream files we modified (merge debt) | 69 | + +Most of the fork is **additive**, which is deliberate — keep it that way. When adding a +feature, put the logic in a **new file** and touch upstream files with as few lines as +possible (ideally: one import, one call). The Save-video feature is the model to copy; the +color-palette feature is the one to be careful with, because it edits shared config. + +--- + +## 2. Upstream tracking + +```bash +git remote -v # 'upstream' should point at Lightricks/LTX-Desktop +git fetch upstream # do this weekly-ish +git log --oneline upstream/main -1 # current upstream release +git branch -r | grep sync-public-preview # early warning: previews land here BEFORE main +``` + +Lightricks publishes `sync-public-preview/-` branches ahead of tagging `main`, so +you generally get advance notice of a release. Also worth enabling GitHub **Watch → Releases +only** on the upstream repo. + +**Merge every release, promptly.** Conflict pain is super-linear: two releases merged +separately is far cheaper than one two-release gap. + +### Merge recipe + +Never merge upstream straight into your working branch. Use a throwaway: + +```bash +git fetch upstream +git checkout -b merge/upstream-X.Y.Z +git merge upstream/main # resolve conflicts using the map below +pnpm run typecheck:ts # safety net 1 +cd backend && uv run pytest -q # safety net 2 +# then the manual smoke test in section 5 +git checkout && git merge merge/upstream-X.Y.Z +``` + +If it goes wrong: `git merge --abort`, delete the branch. Your working branch never moved. + +**Merge, don't rebase.** Rebasing replays every fork commit and makes you re-resolve the same +conflicts repeatedly. + +### ⚠️ Dependency pins: do NOT blindly take upstream's + +Learned the hard way in the 1.1.0 merge. **Upstream being newer overall does not mean every +pin is newer.** Their `diffusers` rev was a strict *ancestor* of ours — 380 commits behind — +and predated `Krea2Pipeline`, so accepting it broke Krea 2 at import time. + +For every conflicting pin in `backend/pyproject.toml`, check the ancestry before choosing: + +``` +https://api.github.com/repos///compare/... +``` +`status: ahead, behind_by: 0` means ours is strictly newer — keep ours. Only take theirs when +they're genuinely ahead, or when upstream code needs a version-specific API. + +Pins we deliberately keep ahead of upstream (see the comments in `pyproject.toml`): +- **`diffusers`** — ours has `Krea2Pipeline`; upstream's does not. + +### ⚠️ Tests cannot validate a merge — run the real app + +Also learned in the 1.1.0 merge: 487 tests passed and the frontend typechecked while the app +**crashed on startup** (`Krea2Pipeline` import) and later blew up mid-generation +(`_hf_hook`). Neither path is covered by tests. Always work through section 5's manual +checklist before committing a merge, and prefer a **fresh app process per model** — chaining +several heavy pipelines in one session parks them all in RAM and produces misleading +slowness that looks like a regression. + +--- + +## 3. What we changed, and what must survive + +### A. Prompt Manager Pro (GPM) — the biggest addition +Ported prompt/camera/workflow manager mounted as a right-hand dock panel. + +- **New:** `frontend/components/gpm/**`, `frontend/gpm-core/**`, `public/gpm-scene/**`, + `electron/gpm-library-root.ts` +- **Modified:** `frontend/App.tsx` (dock mount), `electron/csp.ts`, `electron/config.ts`, + `electron/main.ts`, `electron/preload.ts`, `shared/electron-api-schema.ts` +- **Must survive:** the CSP relaxations GPM needs; the library-root IPC; the dock mount point + in `App.tsx`. +- **Note:** the panel is named **"Prompt Manager Pro"**. Its component/state names use + `promptManagerPro*`. Don't let a merge revert it to "Studio Pro" — that name now collides + with the app name. + +### A2. SettingsDropdown — shared component carrying our fix +Upstream 1.1.0 extracted `SettingsDropdown` out of `GenSpace.tsx` into +`frontend/components/SettingsDropdown.tsx`, but **their version lacked our tag-popup +clipping fix**. We adopted their extraction and ported our fix into it, so the IC-LoRA panels +inherit it too. + +- **Must survive:** the `createPortal` render into `document.body` (an `overflow-hidden` + ancestor otherwise clips the popup invisible) and the viewport-aware + `maxHeight: Math.max(120, rect.top - 16)` (the panel grows upward, so a long list can + otherwise run off the top of the screen). Upstream's static `max-h-80` does **not** cover + either case. +- **Test:** open a tag dropdown with 5+ tags — it must stay fully on screen and scroll. + +### B. Krea 2 Turbo (second image model) +Self-hosted Krea 2 Turbo alongside Z-Image-Turbo, NF4-quantized with a disk cache. + +- **New:** `backend/services/image_generation_pipeline/krea2_*` +- **Modified:** `backend/api_types.py` (`ModelCheckpointID`, `ImageGenerationModelCheckpointID`), + `backend/runtime_config/model_download_specs.py`, `backend/handlers/image_generation_handler.py`, + `backend/handlers/pipelines_handler.py`, `backend/pyproject.toml` (+`bitsandbytes`), `uv.lock` +- **Must survive:** the NF4 quantization **and its disk cache** (~3x speedup — expensive to + regenerate), and the transformers-5.x compatibility fix. If image generation suddenly gets + slow after a merge, the NF4 cache path is the first suspect. +- **Must survive:** our `diffusers` pin (see the dependency-pin warning in section 2) — it is + the only rev that has `Krea2Pipeline`. +- **Must survive (both image pipelines):** `to()` installs accelerate hooks via + `enable_model_cpu_offload()` **once per move-to-accelerator**. Calling it again while + offload is already active leaves modules with accelerate's wrapped `forward` but no + `_hf_hook`, which fails at inference (`'Qwen3Model' object has no attribute '_hf_hook'`) — + and also causes needless CPU↔GPU shuttling. Don't "simplify" that guard away. + +### B2. Free (don't park) image pipelines when loading video ⚠️ upstream divergence +`backend/handlers/pipelines_handler.py` — `_evict_gpu_pipeline_for_swap`. + +Upstream parks an active image pipeline in host RAM (`park_image_generation_pipeline_on_cpu` +→ `CpuSlot`) so an image↔image switch stays warm. We changed it to **free** the image +pipeline (and drop any already-parked `cpu_slot`) instead, because every caller of that method +is loading a memory-hungry video/IC-LoRA/a2v pipeline, and a parked image model then thrashes +a 64 GB box: the video model's bf16-read + fp8-pin working set (~70 GB) plus a parked image +model exceeds RAM and spills to the pagefile. Measured: image→video went 5:04 (parked) → +3:42 (freed). + +- **Must survive:** the free-not-park behavior. Guardrail: + `tests/test_state_actions.py::test_image_pipeline_freed_not_parked_when_loading_video` — it + fails if a merge silently restores upstream's parking (`cpu_slot` becomes a `CpuSlot`). +- `park_image_generation_pipeline_on_cpu` is now **unused** but intentionally left in place to + keep the diff small and reduce merge friction; don't be surprised it has no callers. +- **This is a candidate to upstream** (it's a latent bug that hurts Z-Image too — Z-Image is + bf16/unquantized, so parking it is *worse* than our NF4 Krea 2). If Lightricks takes a PR, + drop this divergence. Root cause is that 64 GB RAM is below what upstream's design assumes; + see the perf notes in section 6. + +### B3. Diffusion stage cache — session-scoped streaming transformer +`backend/services/patches/diffusion_stage_cache.py` (monkey-patch on +`DiffusionStage._transformer_ctx`), plus eviction hooks in +`backend/handlers/generation_handler.py` (`evict_for_generation_start`) and +`backend/handlers/pipelines_handler.py` (unload/swap/image-load), gated from +`backend/ltx2_server.py` (`set_streaming_enabled`, streaming mode only). + +Upstream rebuilds the transformer from the checkpoint on **every** DiffusionStage call +(43GB read + fp8 cast + pin, ~90-100s each on a 3090 — ~190s of a ~226s generation). +Two cached kinds: **resident** (32GB+ cards, generation-scoped — a VRAM-resident cache +surviving the generation double-books VRAM with the next gen's text-encoder/VAE builds, +measured ~42GB on a 32GB card) and **streaming** (24GB cards, ~23GB fp8 pinned host RAM, +**session-scoped** — survives across generations ComfyUI-style; warm generations skip the +model load entirely). + +- **Must survive:** the `set_streaming_enabled` gate (IC-LoRA's `use_lora_in_stage_2` + forces CPU-mode streaming even on 5090s — must NOT session-cache there); the + builder-type + `cpu_slots_count` discriminants in the cache key; the eviction calls in + `pipelines_handler` (video↔image swaps need the pinned RAM back); `teardown()` before + the meta-swap when evicting a streaming entry. +- **Tests:** `tests/test_diffusion_stage_cache.py`. +- The patch reads DiffusionStage/StreamingModelBuilder privates — re-verify against + `ltx_pipelines.utils.blocks` on every rev bump (the docstring lists the exact surface). + +### B4. Aux-model session cache +`backend/services/patches/aux_block_cache.py` — Phase B companion to B3: caches the small +per-generation builds (VAE encoder — shared between ImageConditioner and VideoUpsampler — +spatial upsampler, video decoder, audio decoder, vocoder; ~3-4s/warm gen). Multi-slot, +VRAM-resident (~2-4GB), streaming-mode-gated, evicted at the same `pipelines_handler` +unload/swap/image-load sites (paired with `diffusion_stage_cache.evict()`), governed by the +same Settings toggle. + +- **Must survive:** the streaming gate; the `SingleGPUModelBuilder` isinstance exclusion + (multi-GPU custom decoder builders); the VideoDecoder cached iterator running WITHOUT + `gpu_model` (its teardown would meta-swap the cached decoder) with checkout-at-first-next; + the vocoder's effective-dtype key (fp32 on MPS); the paired eviction calls. +- **Tests:** `tests/test_aux_block_cache.py`. Patch reads `blocks.py` privates — re-verify + each `__call__` body on rev bumps (the docstring lists the surface). + +### B5. fp8 block-weight sidecar +`backend/services/patches/fp8_sidecar_cache.py` — persists the streaming transformer's +POST-downcast block weights to `.fp8-blocks-cache.safetensors` beside the +checkpoint (~22GB; Krea-2 NF4 placement precedent). Cold session builds drop ~95s → ~40s +(23GB identity mmap read, zero cast compute); one-time synchronous write on the first-ever +build. Sits beneath both diffusion-stage-cache paths via a patched +`StreamingModelBuilder._build_pinned_source` + wrapping `StateDictLoader`. + +- **Must survive:** the `__blocks` sd_ops-name intercept contract; the LoRA-free capture + point (pre-fusion — LoRA changes never invalidate the sidecar); atomic writes (tmp + + `os.replace`) with free-space preflight; stamp validation (original path/size/mtime/model-id + + sd_ops chain + fp8 suffix list + exact key-set match) with delete-on-invalid. +- Stale sidecars for deleted checkpoints are acceptable user-deletable residue next to the + models. Kill switch: `FP8_SIDECAR_CACHE=0`. **Tests:** `tests/test_fp8_sidecar_cache.py`. +- Patch reads `builder.py` privates (`_build_pinned_source` signature, `_model_loader`, + `_filtered_sd_ops` naming) — re-verify on rev bumps. + +### C. Qwen Multi-Angle +Camera-angle tool using Qwen-Image-Edit + GGUF + angle/Lightning LoRAs. + +- **New:** `backend/services/qwen_multiangle_pipeline/**`, `backend/handlers/qwen_multiangle_handler.py`, + `backend/_routes/qwen*`, `frontend/components/gpm/QwenMultiAnglePanel.tsx` +- **Modified:** `backend/api_types.py`, `backend/app_factory.py`, `backend/app_handler.py`, + `backend/handlers/__init__.py`, `backend/pyproject.toml` (+`gguf`, `torchvision`) +- **Must survive:** the **SageAttention NaN workaround** (silent black/NaN output without it), + and the state-lift fix that prevents losing state on panel enlarge/shrink. + +### D. Video editor: export transitions & timeline +Color correction, fade-to-black/white, wipes, and dissolve now export correctly (matching +preview), plus timeline in/out persistence and playback fixes. + +- **Modified:** `electron/export/export-handler.ts`, `electron/export/timeline.ts`, + `electron/export/video-filter.ts`, `electron/export/ffmpeg-utils.ts`, + `frontend/views/editor/**` (many) +- **Must survive:** the ffmpeg filter-graph construction for all four transition phases. + Dissolve specifically needs the correct `xfade` type — a regression here is silent + (exports look wrong rather than failing). + +### E. Gen Space: tags/folders & asset safety +Tag/folder (bins) organization, plus fixes for real data-loss bugs. + +- **Modified:** `frontend/views/GenSpace.tsx`, `frontend/contexts/ProjectContext.tsx`, + `frontend/types/project-model.ts` +- **Must survive:** the **stale-autosave fix** — Gen Space and Video Editor stay mounted + together, and a stale editor snapshot could clobber/resurrect assets (bins map, deleted + assets, `binId`/`favorite` fields). This one caused real data loss; treat it as critical. + +### F. 21:9 ultrawide video (local-only) +- **Modified:** `backend/api_types.py` (`aspectRatio` Literal includes `"21:9"`), + `backend/handlers/video_generation_handler.py` (**two** maps: fast path + a2v path), + `frontend/lib/video-generation-model-specs.ts` (`allowedAspectRatios`), + `frontend/views/GenSpace.tsx`, `frontend/generated/backend-openapi.{ts,json}` +- **Must survive:** both `RESOLUTION_MAP_21_9` maps (540p/720p only; 1080p intentionally + rejected), and the local-only gating (forced-API path must keep rejecting 21:9). +- **After merge:** regenerate OpenAPI (`pnpm openapi:generate`) if `api_types.py` changed. + +### G. Save video / Save video frame (right-click menu) +- **New:** `frontend/lib/video-save-actions.ts`, `frontend/components/useVideoSaveMenu.tsx` +- **Modified:** `shared/electron-api-schema.ts` (`copyFileToPath`; `extractVideoFrame` gained + `outputPath`), `electron/ipc/file-handlers.ts`, `electron/ipc/video-processing-handlers.ts`, + `electron/export/ffmpeg-utils.ts` (`accurate` seek), `frontend/lib/file-url.ts` + (`fileUrlToPath`), `frontend/views/GenSpace.tsx`, `frontend/components/gpm/media-preview.ts` +- **Must survive:** the `accurate` seek flag (frame-exactness depends on `-ss` *after* `-i`), + and frame extraction happening in the **main process via ffmpeg** — a canvas grab would + taint under production `webSecurity`. + +### H. Render timer + cancel button +- **Modified:** `frontend/views/GenSpace.tsx`, `frontend/hooks/use-generation.ts`, + `frontend/types/project-model.ts` (`renderMs`) +- **Must survive:** the `abortControllerRef.current?.signal.aborted` check in **both** the + video and image paths. Without it, cancelling shows a false "Generation Failed" dialog, + because the API client converts an aborted fetch into a synthetic error *result* rather + than throwing `AbortError`. + +### I. Color palette themes ⚠️ most merge-fragile +Six curated palettes, picker in Settings → Appearance, persisted in `localStorage`. + +- **New:** `frontend/lib/theme.ts` +- **Modified:** `tailwind.config.js`, `frontend/index.css`, `frontend/main.tsx` (`initTheme()`), + `frontend/components/SettingsModal.tsx` (Appearance tab) +- **Must survive:** the **Tailwind `zinc` scale redirect to CSS variables** in + `tailwind.config.js`, and the matching `--zinc-50..950` defaults in `index.css`. The entire + theming system depends on that indirection — if an upstream change reverts the zinc block, + themes silently stop working (the app still renders, just always dark-default). +- These are shared config files upstream also edits. **Check this feature first after any merge.** + +### J. Dev/runtime infrastructure +- **Modified:** `backend/handlers/base.py` (`LTX_MODELS_DIR` env override), + `electron/app-paths.ts` (`APP_FOLDER_NAME`, `app.setName`), `frontend/App.tsx` + (`requiredModelsGate`), `electron/python-backend.ts` (`LTX_HF_HOME` → `HF_HOME`), + `scripts/launch-fork-dev.cmd` (new) +- **Must survive:** + - **`LTX_MODELS_DIR` env override** — takes priority over persisted `models_dir`, because + the settings.json round-trip proved unreliable. Also set in `launch-fork-dev.cmd`; the two + must stay in sync. + - **`LTX_HF_HOME` env override** — when set, `python-backend.ts` injects it as `HF_HOME` + for the backend, redirecting the HuggingFace cache (Qwen Multi-Angle's ~17GB GGUF load + obeys it) to a fast drive without touching the global env or other HF tools. App-scoped, + portable (no-op when unset). `launch-fork-dev.cmd` sets it to `E:\hf-cache` because the + user's global `HF_HOME` pointed at a SATA HDD (`D:\hf-cache`) that pinned at 99% and made + cold Qwen loads take minutes. + - **Startup gate fix** — don't block startup on missing local models when an LTX API key is + present (otherwise a bogus ~25GB download prompt reappears). + - **Dev userData isolation** — dev builds use a separate app folder so the fork doesn't + collide with an installed LTX Desktop (shared single-instance lock). + +### K. Branding +Partially rebranded; a full rebrand is still pending (see section 6). + +- **Modified:** `index.html`, `package.json`, `frontend/App.tsx`, `frontend/views/Home.tsx`, + `frontend/components/PythonSetup.tsx`, `frontend/components/SettingsModal.tsx` + +### L. NOTICES.md +Third-party notices extended for our additions (Krea 2, the four Qwen artifacts, +`torchvision`/`bitsandbytes`/`gguf`, and several npm deps). **Re-check after every merge** — +upstream will edit this file too, and dropping our entries is a licensing problem, not just +a cosmetic one. + +--- + +## 4. Hot-spot files (highest conflict probability) + +These are files **both** we and upstream change regularly. They're mostly registration / +wiring points — unavoidable, since adding a feature means registering it. + +``` +backend/api_types.py backend/app_factory.py +backend/app_handler.py backend/handlers/__init__.py +backend/handlers/pipelines_handler.py backend/handlers/video_generation_handler.py +backend/runtime_config/model_download_specs.py +backend/pyproject.toml + uv.lock shared/electron-api-schema.ts +electron/preload.ts electron/ipc/file-handlers.ts +frontend/App.tsx frontend/views/GenSpace.tsx +frontend/components/SettingsModal.tsx tailwind.config.js + frontend/index.css +frontend/generated/backend-openapi.{ts,json} +``` + +`backend/generated/*` and `uv.lock` are **generated** — don't hand-merge them. Take either +side, then regenerate: +```bash +pnpm openapi:generate # regenerates backend-openapi.{json,ts} +cd backend && uv sync # regenerates uv.lock +``` + +--- + +## 5. Post-merge verification checklist + +Automated (must pass): +- [ ] `pnpm run typecheck:ts` +- [ ] `cd backend && uv run pytest -q` + +Manual smoke test (each maps to a feature above): +- [ ] App launches; no first-run/model-download gate (J) +- [ ] Gen Space: generate a video; **timer** counts, **Stop** cancels cleanly with no error dialog (H) +- [ ] Aspect ratio dropdown offers **21:9**; generates at 540p/720p (F) +- [ ] Right-click a video → **Save video** / **Save video frame** both work (G) +- [ ] Settings → **Appearance**: palettes apply and persist across restart (I) ⚠️ +- [ ] Prompt Manager Pro panel opens; library loads (A) +- [ ] Image mode: **Krea 2 Turbo** generates, and is fast (NF4 cache intact) (B) +- [ ] Qwen Multi-Angle generates without NaN/black output (C) +- [ ] Video editor: export a clip with a dissolve/wipe; transitions match preview (D) +- [ ] Gen Space tags/folders persist after switching to the editor and back (E) +- [ ] `NOTICES.md` still lists our models/deps (L) + +--- + +## 6. Known outstanding work + +- **Full rebrand** — the app is still named/branded around "LTX", and `electron-builder.yml` + still carries `appId: com.lightricks.ltx-desktop` plus Lightricks' Azure code-signing block. + Apache-2.0 grants no trademark rights, so this must change before any distribution. Note + that renaming the userData folder will orphan projects (they live in `localStorage` inside + it) — the folder must be renamed/migrated, not just repointed. +- ~~Upstream 1.1.0~~ — **merged** in `cd9395e`. Brings the LoRA / IC-LoRA catalog, video + Extend, outpainting, a Models tab, generation recovery, and heartbeat instrumentation. + Largely unexercised so far: the **LoRA catalog** and **Extend** got a code review but only + a light smoke test. +- **New checkpoint requirement:** 1.1.0 repoints *both* model variants at + `ltx-2.3-spatial-upscaler-x2-1.1`, so that ~950MB file is required even when staying on + the 1.0 transformer. The Models tab offers a whole-bundle download but skips files already + on disk, so it only fetches what's missing. +- **Ingredients IC-LoRA** is CLOSED — needs 32GB+ VRAM and the full (non-distilled) 22B + model. Not viable on a 3090; don't reopen without new hardware. +- **Possible perf win (untested):** on a 24GB card the runtime policy selects + `streaming_models_loading`, and the heartbeats show it streaming weights off disk while + ~20GB of VRAM sits free (peak usage only ~4.6GB). That band was presumably tuned for the + full model, not our fp8-cast distilled one. Forcing `full_models_loading` might cut render + times substantially. Pre-existing behaviour, not a merge issue — worth an A/B using the new + heartbeat instrumentation. + +--- + +## 7. Licensing reminders when merging + +- App code is **Apache 2.0** — retain copyright/attribution notices, and mark modified files. +- Model licenses are separate: **LTX-2 Community License** (commercial use permitted below + $10M annual revenue) and **Krea 2 Community License** (requires deployer-side content + filtering). Keep `NOTICES.md` accurate. diff --git a/NOTICES.md b/NOTICES.md index 62320b238..ef228d443 100644 --- a/NOTICES.md +++ b/NOTICES.md @@ -37,12 +37,38 @@ used by LTX Desktop. License: Apache License 2.0 https://huggingface.co/hr16/yolox-onnx +- **Krea 2 Turbo** + Copyright (c) Krea AI. + License: Krea 2 Community License + https://huggingface.co/krea/Krea-2-Turbo + +- **Qwen-Image-Edit-2511** + Copyright (c) Alibaba Cloud / Qwen Team. + License: Apache License 2.0 + https://huggingface.co/Qwen/Qwen-Image-Edit-2511 + +- **Qwen-Image-Edit-2511 (GGUF quantization)** + Copyright (c) Unsloth AI; base model (c) Alibaba Cloud / Qwen Team. + License: Apache License 2.0 + https://huggingface.co/unsloth/Qwen-Image-Edit-2511-GGUF + +- **Qwen-Image-Edit-2511 Multiple Angles LoRA** + Copyright (c) fal.ai. + License: Apache License 2.0 + https://huggingface.co/fal/Qwen-Image-Edit-2511-Multiple-Angles-LoRA + +- **Qwen-Image-Edit-2511 Lightning LoRA** + Copyright (c) LightX2V contributors. + License: Apache License 2.0 + https://huggingface.co/lightx2v/Qwen-Image-Edit-2511-Lightning + --- ## Python Dependencies - **pillow** — Copyright (c) Jeffrey A. Clark and Pillow contributors — HPND License - **torch** — Copyright (c) Meta Platforms, Inc. — BSD 3-Clause License +- **torchvision** — Copyright (c) Soumith Chintala and torchvision contributors — BSD 3-Clause License - **huggingface-hub** — Copyright (c) Hugging Face — Apache License 2.0 - **tqdm** — Copyright (c) Casper da Costa-Luis — Mozilla Public License 2.0 - **pynvml** — Copyright (c) NVIDIA Corporation — BSD 3-Clause License @@ -57,6 +83,8 @@ used by LTX Desktop. - **protobuf** — Copyright (c) Google LLC — BSD 3-Clause License - **transformers** — Copyright (c) Hugging Face — Apache License 2.0 - **sentencepiece** — Copyright (c) Google LLC — Apache License 2.0 +- **bitsandbytes** — Copyright (c) Facebook, Inc. and its affiliates — MIT License +- **gguf** — Copyright (c) GGML (ggml.ai) — MIT License - **sageattention** — Copyright (c) Jintao Zhang et al. — Apache License 2.0 - **natten** — Copyright (c) Ali Hassani, Steven Walton, et al. (SHI Labs) — Apache License 2.0 - **opencv-python-headless** — Copyright (c) OpenCV team — Apache License 2.0 @@ -77,8 +105,11 @@ used by LTX Desktop. - **electron-updater** — Copyright (c) electron-userland — MIT License - **js-yaml** — Copyright (c) Vitaly Puzrin — MIT License - **lucide-react** — Copyright (c) Lucide contributors — ISC License -- **react-dropzone** — Copyright (c) Param Aggarwal — MIT License +- **react-resizable-panels** — Copyright (c) Brian Vaughn — MIT License - **tailwind-merge** — Copyright (c) dcastil — MIT License +- **use-sync-external-store** — Copyright (c) Meta Platforms, Inc. and affiliates — MIT License +- **zod** — Copyright (c) Colin McDonnell — MIT License +- **zustand** — Copyright (c) Paul Henschel and Poimandres contributors — MIT License --- diff --git a/README.md b/README.md index 92f2d3f6f..60ff3ca72 100644 --- a/README.md +++ b/README.md @@ -148,6 +148,45 @@ Your models folder is the location you chose during setup, or the default for yo Open the selector to tick one or more LoRAs and adjust each one's strength. If you added files while the app was already open, reopen the generation view (or restart the app) — the installed‑LoRA list is read on load. +## LoRA/IC-LoRA Library + +LoRA adapters let you steer **local** video generation toward a specific style or subject. They apply to local generation (Windows/Linux NVIDIA hardware or Apple Silicon Macs), for text-to-video, image-to-video, and audio-to-video — not in API/cloud mode. + +In local mode, **Browse LoRAs** (styles/subjects for text/image/audio-to-video) and **Browse IC-LoRAs** (in-context effects, e.g. video-to-video) open a built-in library of ready-to-download LoRAs — no manual file placement needed. Each entry shows a preview, its instructions, and a link to its Hugging Face page; some gated models require signing in with Hugging Face first. Entries authored by **LTX** are official; everything else is community-contributed and flagged with a disclaimer that LTX doesn't endorse or take responsibility for it. + +### Custom LoRAs + +You can also use your own `.safetensors` LoRA files instead of (or alongside) the library. + +**Which LoRAs are supported.** The app generates locally with the **LTX‑2.3 (22B)** model. "LTX‑2" is the model generation and "2.3" is its current release, so LoRAs labeled **LTX‑2** or **LTX‑2.3** both target this model and are supported. This includes LoRAs exported from ComfyUI for LTX‑Video (their key names are remapped automatically). A LoRA whose tensors don't match the model is simply skipped, so it has no effect rather than producing an error. LoRAs built for other base models (e.g. SDXL, Wan, Hunyuan, or older LTXV 0.9.x) target a different architecture and will not take effect. + +**Where to put the files.** Library downloads land in their own subfolder under `loras/` (one per LoRA), so it's safe to drop your own `.safetensors` files straight into the `loras/` (or `lora/`) subfolder of your models folder: + +``` +models/ +└── loras/ + ├── cinematic.safetensors + ├── claymation.safetensors + └── / ← library download, one subfolder per LoRA + └── model.safetensors +``` + +> **Use the `loras/` subfolder.** A file is detected as a LoRA only if it lives in a folder named `loras`/`lora` **or** its filename contains `lora` — and many LoRAs aren't named that way. Dropping them straight into `loras/` is the reliable option; it won't collide with library downloads, which each get their own subfolder. Subfolders are scanned recursively, so further nesting is fine. + +Your models folder is the location you chose during setup, or the default for your platform: + +- **Windows:** `%LOCALAPPDATA%\LTXDesktop\models\` +- **macOS:** `~/Library/Application Support/LTXDesktop/models/` +- **Linux:** `$XDG_DATA_HOME/LTXDesktop/models/` (default: `~/.local/share/LTXDesktop/models/`) + +**When the LoRA button appears.** In the generation view, a **LoRA** selector appears next to the other video settings (FPS, aspect ratio) only when **all** of these are true: + +1. You're in **Video** generation mode (not Image, Retake, or IC‑LoRA). +2. You're generating **locally** — not in API/cloud generation mode. +3. **At least one** LoRA file is detected in the models folder. + +Open the selector to tick one or more LoRAs and adjust each one's strength. If you added files while the app was already open, reopen the generation view (or restart the app) — the installed‑LoRA list is read on load. + ## API keys, cost, and privacy ### LTX API key diff --git a/backend/_routes/qwen_multiangle.py b/backend/_routes/qwen_multiangle.py new file mode 100644 index 000000000..8b6c6a3c6 --- /dev/null +++ b/backend/_routes/qwen_multiangle.py @@ -0,0 +1,18 @@ +"""Route handler for POST /api/qwen-multiangle/generate.""" + +from __future__ import annotations + +from fastapi import APIRouter, Depends + +from api_types import QwenMultiAngleGenerateRequest, QwenMultiAngleGenerateResponse +from state import get_state_service +from app_handler import AppHandler + +router = APIRouter(prefix="/api", tags=["qwen-multiangle"]) + + +@router.post("/qwen-multiangle/generate", response_model=QwenMultiAngleGenerateResponse) +def route_qwen_multiangle_generate( + req: QwenMultiAngleGenerateRequest, handler: AppHandler = Depends(get_state_service) +) -> QwenMultiAngleGenerateResponse: + return handler.qwen_multiangle.generate(req) diff --git a/backend/api_types.py b/backend/api_types.py index 8f004fff1..eefd8bf23 100644 --- a/backend/api_types.py +++ b/backend/api_types.py @@ -30,7 +30,9 @@ "gemma4-12b-with-proj-ltx-2.5", "gemma-4-e2b-it", "z-image-turbo", + "krea-2-turbo", ] +ImageGenerationModelCheckpointID = Literal["z-image-turbo", "krea-2-turbo"] LTXLocalModelId = Literal[ "ltx-2.5-22b-distilled", "ltx-2.3-22b-distilled-1.1", @@ -161,6 +163,8 @@ class SuggestGapPromptResponse(BaseModel): class GenerateVideoCompleteResponse(BaseModel): status: Literal["complete"] video_path: str + # Seed actually used (None for API providers that don't expose it). + seed: int | None = None class GenerateVideoCancelledResponse(BaseModel): @@ -173,6 +177,8 @@ class GenerateVideoCancelledResponse(BaseModel): class GenerateImageCompleteResponse(BaseModel): status: Literal["complete"] image_paths: list[str] + # Base seed; image i in the batch used seed + i. + seed: int | None = None class GenerateImageCancelledResponse(BaseModel): @@ -475,7 +481,7 @@ class GenerateVideoRequest(BaseModel): lastImagePath: str | None = None keyframes: list[KeyframeInput] = Field(default_factory=list[KeyframeInput]) audioPath: str | None = None - aspectRatio: Literal["16:9", "9:16"] = "16:9" + aspectRatio: Literal["16:9", "9:16", "21:9"] = "16:9" seed: int | None = None loras: list[LoraEntry] = Field(default_factory=list[LoraEntry]) @@ -484,12 +490,15 @@ class GenerateImageRequest(BaseModel): model_config = ConfigDict(strict=True) prompt: NonEmptyPrompt + model: ImageGenerationModelCheckpointID = "z-image-turbo" width: int = Field(default=1024, ge=16) height: int = Field(default=1024, ge=16) numSteps: int = Field(default=4, ge=1) numImages: int = Field(default=1, ge=1) imagePath: str | None = None strength: float = Field(default=0.6, ge=0.0, le=1.0) + # Z-Image text-to-image only: 0 = off, 1 = max composition variety across seeds. + variation: float = Field(default=0.0, ge=0.0, le=1.0) def _default_model_types() -> set[ModelCheckpointID]: @@ -558,6 +567,7 @@ class TargetResolution(BaseModel): height: int = Field(gt=0) + class RetakeRequest(BaseModel): model_config = ConfigDict(strict=True) @@ -592,6 +602,42 @@ class ExtendRequest(BaseModel): ExtendResponse: TypeAlias = RetakeResponse +class QwenMultiAngleGenerateRequest(BaseModel): + model_config = ConfigDict(strict=True) + + image_data_url: str + # Compose mode: additional reference images (prop, location, wardrobe) fed to + # Qwen alongside the subject. image_data_url is the subject being re-angled; + # these are extra context the model composes from. Empty = classic single-image. + extra_image_data_urls: list[str] = Field(default_factory=list) + # Role of each extra ref, aligned index-for-index with extra_image_data_urls: + # "location" = a scene the subject is composited INTO, "prop" = an object the + # subject holds/uses. Empty (or length-mismatched) = classic compose, no + # explicit compositing instruction. Drives compose_prompt(). + extra_image_roles: list[Literal["location", "prop"]] = [] + azimuth_deg: float + elevation_deg: float + zoom: float + seed: int = 42 + randomize_seed: bool = False + # Sampling recipe: "fast" = 4-step Lightning (~30s, softer skin), "balanced" + # = 8-step Lightning (retains more texture, ~2x fast), "quality" = 28-step + # base, no distillation LoRA (sharpest skin, slowest). See the GGUF pipeline. + quality_mode: Literal["fast", "balanced", "quality"] = "fast" + # Optional skin-realism LoRA stacked on top of the active mode; skin_weight + # scales it (recommended ~0.6-1.2) so it doesn't overpower the angles LoRA. + use_skin: bool = False + skin_weight: float = 1.0 + extra_prompt: str = "" + + +class QwenMultiAngleGenerateResponse(BaseModel): + status: Literal["complete"] = "complete" + image_data_url: str + prompt: str + seed: int + + ConditioningType: TypeAlias = Literal["canny", "depth"] # Generation can additionally run a user-supplied IC-LoRA against a pre-rendered diff --git a/backend/app_factory.py b/backend/app_factory.py index 5073c3c0f..71d3231dc 100644 --- a/backend/app_factory.py +++ b/backend/app_factory.py @@ -23,6 +23,7 @@ from _routes.image_gen import router as image_gen_router from _routes.prompt_enhancement import router as prompt_enhancement_router from _routes.models import router as models_router +from _routes.qwen_multiangle import router as qwen_multiangle_router from _routes.suggest_gap_prompt import router as suggest_gap_prompt_router from _routes.retake import router as retake_router from _routes.extend import router as extend_router @@ -161,6 +162,7 @@ async def _route_generic_error_handler(request: Request, exc: Exception) -> JSON app.include_router(image_gen_router) app.include_router(suggest_gap_prompt_router) app.include_router(retake_router) + app.include_router(qwen_multiangle_router) app.include_router(extend_router) app.include_router(ic_lora_router) app.include_router(lora_catalog_router) diff --git a/backend/app_handler.py b/backend/app_handler.py index 5cfca0764..2ce6b1118 100644 --- a/backend/app_handler.py +++ b/backend/app_handler.py @@ -3,8 +3,10 @@ from __future__ import annotations import threading +from collections.abc import Mapping from dataclasses import dataclass +from api_types import ImageGenerationModelCheckpointID from state.app_settings import AppSettings from handlers import ( DownloadHandler, @@ -16,6 +18,7 @@ ModelsHandler, PipelinesHandler, LoraCatalogHandler, + QwenMultiAngleHandler, PromptEnhancementHandler, SuggestGapPromptHandler, RetakeHandler, @@ -39,6 +42,7 @@ LTXAPIClient, ModelDownloader, PoseProcessorPipeline, + QwenMultiAnglePipeline, PromptEnhancerPipeline, RetakePipeline, TaskRunner, @@ -68,13 +72,14 @@ def __init__( ltx_api_client: LTXAPIClient, zit_api_client: ZitAPIClient, fast_video_pipeline_class: type[FastVideoPipeline], - image_generation_pipeline_class: type[ImageGenerationPipeline], + image_generation_pipeline_classes: Mapping[ImageGenerationModelCheckpointID, type[ImageGenerationPipeline]], ic_lora_pipeline_class: type[IcLoraPipeline], depth_processor_pipeline_class: type[DepthProcessorPipeline], pose_processor_pipeline_class: type[PoseProcessorPipeline], a2v_pipeline_class: type[A2VPipeline], retake_pipeline_class: type[RetakePipeline], prompt_enhancer_pipeline_class: type[PromptEnhancerPipeline], + qwen_multiangle_pipeline_class: type[QwenMultiAnglePipeline] | None = None, ) -> None: self.config = config @@ -88,12 +93,13 @@ def __init__( self.ltx_api_client = ltx_api_client self.zit_api_client = zit_api_client self.fast_video_pipeline_class = fast_video_pipeline_class - self.image_generation_pipeline_class = image_generation_pipeline_class + self.image_generation_pipeline_classes = image_generation_pipeline_classes self.ic_lora_pipeline_class = ic_lora_pipeline_class self.depth_processor_pipeline_class = depth_processor_pipeline_class self.pose_processor_pipeline_class = pose_processor_pipeline_class self.a2v_pipeline_class = a2v_pipeline_class self.retake_pipeline_class = retake_pipeline_class + self.qwen_multiangle_pipeline_class = qwen_multiangle_pipeline_class self.prompt_enhancer_pipeline_class = prompt_enhancer_pipeline_class self._lock = threading.RLock() @@ -161,12 +167,13 @@ def __init__( text_handler=self.text, gpu_cleaner=gpu_cleaner, fast_video_pipeline_class=fast_video_pipeline_class, - image_generation_pipeline_class=image_generation_pipeline_class, + image_generation_pipeline_classes=image_generation_pipeline_classes, ic_lora_pipeline_class=ic_lora_pipeline_class, depth_processor_pipeline_class=depth_processor_pipeline_class, pose_processor_pipeline_class=pose_processor_pipeline_class, a2v_pipeline_class=a2v_pipeline_class, retake_pipeline_class=retake_pipeline_class, + qwen_multiangle_pipeline_class=qwen_multiangle_pipeline_class, config=config, ) @@ -233,6 +240,14 @@ def __init__( text_handler=self.text, ) + self.qwen_multiangle = QwenMultiAngleHandler( + state=self.state, + lock=self._lock, + config=config, + generation_handler=self.generation, + pipelines_handler=self.pipelines, + ) + self.extend = ExtendHandler( state=self.state, lock=self._lock, @@ -277,13 +292,14 @@ class ServiceBundle: ltx_api_client: LTXAPIClient zit_api_client: ZitAPIClient fast_video_pipeline_class: type[FastVideoPipeline] - image_generation_pipeline_class: type[ImageGenerationPipeline] + image_generation_pipeline_classes: Mapping[ImageGenerationModelCheckpointID, type[ImageGenerationPipeline]] ic_lora_pipeline_class: type[IcLoraPipeline] depth_processor_pipeline_class: type[DepthProcessorPipeline] pose_processor_pipeline_class: type[PoseProcessorPipeline] a2v_pipeline_class: type[A2VPipeline] retake_pipeline_class: type[RetakePipeline] prompt_enhancer_pipeline_class: type[PromptEnhancerPipeline] + qwen_multiangle_pipeline_class: type[QwenMultiAnglePipeline] | None = None def build_default_service_bundle(config: RuntimeConfig) -> ServiceBundle: @@ -296,10 +312,12 @@ def build_default_service_bundle(config: RuntimeConfig) -> ServiceBundle: from services.a2v_pipeline.ltx_a2v_pipeline import LTXa2vPipeline from services.depth_processor_pipeline.midas_dpt_pipeline import MidasDPTPipeline from services.ic_lora_pipeline.ltx_ic_lora_pipeline import LTXIcLoraPipeline + from services.image_generation_pipeline.krea2_image_generation_pipeline import Krea2ImageGenerationPipeline from services.image_generation_pipeline.zit_image_generation_pipeline import ZitImageGenerationPipeline from services.ltx_api_client.ltx_api_client_impl import LTXAPIClientImpl from services.model_downloader.hugging_face_downloader import HuggingFaceDownloader from services.retake_pipeline.ltx_retake_pipeline import LTXRetakePipeline + from services.qwen_multiangle_pipeline.gguf_qwen_multiangle_pipeline import GGUFQwenMultiAnglePipeline from services.prompt_enhancer_pipeline.ltx_prompt_enhancer_pipeline import LtxPromptEnhancerPipeline from services.pose_processor_pipeline.dw_pose_pipeline import DWPosePipeline from services.task_runner.threading_runner import ThreadingRunner @@ -327,12 +345,16 @@ def build_default_service_bundle(config: RuntimeConfig) -> ServiceBundle: ltx_api_client=LTXAPIClientImpl(http=http, ltx_api_base_url=config.ltx_api_base_url), zit_api_client=ZitAPIClientImpl(http=http), fast_video_pipeline_class=LTXFastVideoPipeline, - image_generation_pipeline_class=ZitImageGenerationPipeline, + image_generation_pipeline_classes={ + "z-image-turbo": ZitImageGenerationPipeline, + "krea-2-turbo": Krea2ImageGenerationPipeline, + }, ic_lora_pipeline_class=LTXIcLoraPipeline, depth_processor_pipeline_class=MidasDPTPipeline, pose_processor_pipeline_class=DWPosePipeline, a2v_pipeline_class=LTXa2vPipeline, retake_pipeline_class=LTXRetakePipeline, + qwen_multiangle_pipeline_class=GGUFQwenMultiAnglePipeline, prompt_enhancer_pipeline_class=LtxPromptEnhancerPipeline, ) @@ -358,11 +380,12 @@ def build_initial_state( ltx_api_client=bundle.ltx_api_client, zit_api_client=bundle.zit_api_client, fast_video_pipeline_class=bundle.fast_video_pipeline_class, - image_generation_pipeline_class=bundle.image_generation_pipeline_class, + image_generation_pipeline_classes=bundle.image_generation_pipeline_classes, ic_lora_pipeline_class=bundle.ic_lora_pipeline_class, depth_processor_pipeline_class=bundle.depth_processor_pipeline_class, pose_processor_pipeline_class=bundle.pose_processor_pipeline_class, a2v_pipeline_class=bundle.a2v_pipeline_class, retake_pipeline_class=bundle.retake_pipeline_class, + qwen_multiangle_pipeline_class=bundle.qwen_multiangle_pipeline_class, prompt_enhancer_pipeline_class=bundle.prompt_enhancer_pipeline_class, ) diff --git a/backend/export_openapi_schema.py b/backend/export_openapi_schema.py index 3cadbc2ad..d95ab219a 100644 --- a/backend/export_openapi_schema.py +++ b/backend/export_openapi_schema.py @@ -64,7 +64,13 @@ def _build_schema() -> dict[str, object]: ltx_api_client=cast(Any, fake.ltx_api_client), zit_api_client=cast(Any, fake.zit_api_client), fast_video_pipeline_class=cast(Any, type(fake.fast_video_pipeline)), - image_generation_pipeline_class=cast(Any, type(fake.image_generation_pipeline)), + image_generation_pipeline_classes=cast( + Any, + { + "z-image-turbo": type(fake.image_generation_pipeline), + "krea-2-turbo": type(fake.image_generation_pipeline), + }, + ), ic_lora_pipeline_class=cast(Any, type(fake.ic_lora_pipeline)), depth_processor_pipeline_class=cast(Any, type(fake.depth_processor_pipeline)), pose_processor_pipeline_class=cast(Any, type(fake.pose_processor_pipeline)), diff --git a/backend/handlers/__init__.py b/backend/handlers/__init__.py index 3f598943f..95b59c129 100644 --- a/backend/handlers/__init__.py +++ b/backend/handlers/__init__.py @@ -9,6 +9,7 @@ from handlers.models_handler import ModelsHandler from handlers.pipelines_handler import PipelinesHandler from handlers.lora_catalog_handler import LoraCatalogHandler +from handlers.qwen_multiangle_handler import QwenMultiAngleHandler from handlers.prompt_enhancement_handler import PromptEnhancementHandler from handlers.suggest_gap_prompt_handler import SuggestGapPromptHandler from handlers.retake_handler import RetakeHandler @@ -31,6 +32,7 @@ "SuggestGapPromptHandler", "RetakeHandler", "ExtendHandler", + "QwenMultiAngleHandler", "RuntimePolicyHandler", "IcLoraHandler", "HuggingFaceAuthHandler", diff --git a/backend/handlers/base.py b/backend/handlers/base.py index 470f4e3b9..1bf3aed78 100644 --- a/backend/handlers/base.py +++ b/backend/handlers/base.py @@ -3,7 +3,8 @@ from __future__ import annotations import logging -import time +import os +import secrets from collections.abc import Callable from functools import wraps from pathlib import Path @@ -47,19 +48,31 @@ def config(self) -> RuntimeConfig: @property def models_dir(self) -> Path: - """Effective models dir: custom from settings, or startup default.""" + """Effective models dir: env override, then custom from settings, then startup default. + + LTX_MODELS_DIR takes priority over the persisted setting because the + settings.json round-trip (app boot -> frontend sync -> save-on-exit) + has proven unreliable for this field in practice — an explicit env + var set by the launcher is a hard guarantee that doesn't depend on + that sync path working correctly. + """ + env_override = os.environ.get("LTX_MODELS_DIR") + if env_override: + return Path(env_override) custom = self._state.app_settings.models_dir return Path(custom) if custom else self._config.default_models_dir def _resolve_seed(self) -> int: - """Resolve the generation seed: locked seed, dev-mode constant, or time-based.""" + """Resolve the generation seed: the locked seed, else a fresh random one. + + (Dev mode used to pin unlocked gens to a constant 1000, which made "random" + silently non-random in every dev run — removed.) + """ settings = self.state.app_settings if settings.seed_locked: logger.info("Using locked seed: %s", settings.locked_seed) return settings.locked_seed - if self.config.dev_mode: - return 1000 - return int(time.time()) % _MAX_SEED + return secrets.randbelow(_MAX_SEED) def with_state_lock( diff --git a/backend/handlers/generation_handler.py b/backend/handlers/generation_handler.py index 499f9c81a..7ac78fff1 100644 --- a/backend/handlers/generation_handler.py +++ b/backend/handlers/generation_handler.py @@ -19,7 +19,7 @@ from handlers.base import StateHandlerBase, with_state_lock from services import generation_interrupt from services.generation_interrupt import GenerationCancelledError -from services.patches import diffusion_stage_cache +from services.patches import aux_block_cache, diffusion_stage_cache from state.app_state_types import ( ApiGeneration, AppState, @@ -118,13 +118,20 @@ def start_generation(self, generation_id: str) -> None: raise GenerationCancelledError() self.state.generation_starting_since = None - # EXPERIMENTAL: push the live Settings toggle, then drop any transformer - # cached from the previous generation before this one starts -- otherwise - # it stays resident while this generation's own text encoder/VAE/etc. - # build, double-booking VRAM. See that module's GENERATION-SCOPED - # docstring section for the RTX 5090 repro (~42GB reported on a 32GB card). + # EXPERIMENTAL: push the live Settings toggle, then apply the kind-scoped + # generation-start eviction: a VRAM-resident cached transformer is dropped + # (it would double-book VRAM with this generation's own text encoder/VAE + # builds -- see that module's GENERATION-SCOPED docstring section for the + # RTX 5090 repro, ~42GB reported on a 32GB card), while a streaming entry + # (pinned host RAM) is session-scoped and kept across generations unless + # system RAM is under pressure -- see the SESSION-SCOPED section. This is + # the win on the 3090 (streaming path): the transformer stays resident in + # pinned host RAM across generations instead of rebuilding every stage. diffusion_stage_cache.set_enabled(self.state.app_settings.diffusion_stage_cache_enabled) - diffusion_stage_cache.evict() + # The aux-model cache shares the same Settings toggle; its (VRAM, small) + # entries are kept across generations unconditionally -- no gen-start evict. + aux_block_cache.set_enabled(self.state.app_settings.diffusion_stage_cache_enabled) + diffusion_stage_cache.evict_for_generation_start() generation_interrupt.clear() self.state.active_generation = GpuGeneration( @@ -148,9 +155,14 @@ def start_api_generation(self, generation_id: str) -> None: self.state.generation_starting_since = None # EXPERIMENTAL: see start_generation -- an API generation doesn't build a - # local transformer itself, but evicting here still releases VRAM held by - # a previous local generation's cached build. - diffusion_stage_cache.evict() + # local transformer itself, but the kind-scoped eviction still releases + # VRAM held by a previous local generation's resident build, while keeping + # a session-scoped streaming entry warm for the next local generation. + # Push the toggle here too: a user who disables the setting but then only + # runs API generations must still reclaim the pinned host RAM. + diffusion_stage_cache.set_enabled(self.state.app_settings.diffusion_stage_cache_enabled) + aux_block_cache.set_enabled(self.state.app_settings.diffusion_stage_cache_enabled) + diffusion_stage_cache.evict_for_generation_start() generation_interrupt.clear() self.state.active_generation = ApiGeneration( diff --git a/backend/handlers/image_generation_handler.py b/backend/handlers/image_generation_handler.py index be543a438..e1eae5483 100644 --- a/backend/handlers/image_generation_handler.py +++ b/backend/handlers/image_generation_handler.py @@ -19,6 +19,7 @@ GenerateImageCompleteResponse, GenerateImageRequest, GenerateImageResponse, + ImageGenerationModelCheckpointID, ) from handlers.base import StateHandlerBase from handlers.generation_handler import GenerationHandler @@ -73,7 +74,18 @@ def generate(self, req: GenerateImageRequest) -> GenerateImageResponse: use_fal_api=use_fal_api, ) + # FORK: Krea 2 is dev-only. Its license (§4.2) requires deployers to run content + # filters and forbids circumventing its safety tuning; release builds ship + # Z-Image (Apache 2.0) only. Checked before the API branch so it always applies. + if req.model == "krea-2-turbo" and not self.config.dev_mode: + raise HTTPError(400, "KREA_2_DEV_ONLY") + if use_fal_api: + # FORK: Krea 2 is self-hosted only — it has no FAL/API path, so a + # forced-API request for it must fail rather than silently fall back + # to the API's z-image model. + if req.model == "krea-2-turbo": + raise HTTPError(400, "KREA_2_IS_SELF_HOSTED_ONLY") return self._generate_via_api( prompt=req.prompt, width=width, @@ -86,18 +98,20 @@ def generate(self, req: GenerateImageRequest) -> GenerateImageResponse: generation_id = uuid.uuid4().hex[:8] try: self._generation.raise_if_cancelled() - self._pipelines.load_image_generation_pipeline_to_gpu() + self._pipelines.load_image_generation_pipeline_to_gpu(req.model) self._generation.start_generation(generation_id) output_paths = self.generate_image( + model=req.model, prompt=req.prompt, width=width, height=height, num_inference_steps=req.numSteps, seed=seed, num_images=num_images, + variation=req.variation, ) self._generation.complete_generation(output_paths) - return GenerateImageCompleteResponse(status="complete", image_paths=output_paths) + return GenerateImageCompleteResponse(status="complete", image_paths=output_paths, seed=seed) except Exception as e: self._generation.fail_generation(str(e)) if is_cancel_exception(e): @@ -136,7 +150,7 @@ def _edit( generation_id = uuid.uuid4().hex[:8] try: self._generation.raise_if_cancelled() - self._pipelines.load_image_generation_pipeline_to_gpu() + self._pipelines.load_image_generation_pipeline_to_gpu("z-image-turbo") self._generation.start_generation(generation_id) output_paths = self.edit_image( prompt=req.prompt, @@ -147,7 +161,7 @@ def _edit( num_images=num_images, ) self._generation.complete_generation(output_paths) - return GenerateImageCompleteResponse(status="complete", image_paths=output_paths) + return GenerateImageCompleteResponse(status="complete", image_paths=output_paths, seed=seed) except Exception as e: self._generation.fail_generation(str(e)) if is_cancel_exception(e): @@ -167,7 +181,7 @@ def edit_image( self._generation.raise_if_cancelled() self._generation.update_progress("loading_model", 5, 0, num_inference_steps) - image_generation_pipeline = self._pipelines.load_image_generation_pipeline_to_gpu() + image_generation_pipeline = self._pipelines.load_image_generation_pipeline_to_gpu("z-image-turbo") self._generation.update_progress("inference", 15, 0, num_inference_steps) source = self._load_edit_source(image_path) @@ -187,17 +201,19 @@ def generate_one(seed_i: int) -> PILImage.Image: def generate_image( self, + model: ImageGenerationModelCheckpointID, prompt: str, width: int, height: int, num_inference_steps: int, seed: int, num_images: int, + variation: float = 0.0, ) -> list[str]: self._generation.raise_if_cancelled() self._generation.update_progress("loading_model", 5, 0, num_inference_steps) - image_generation_pipeline = self._pipelines.load_image_generation_pipeline_to_gpu() + image_generation_pipeline = self._pipelines.load_image_generation_pipeline_to_gpu(model) self._generation.update_progress("inference", 15, 0, num_inference_steps) def generate_one(seed_i: int) -> PILImage.Image: @@ -208,6 +224,7 @@ def generate_one(seed_i: int) -> PILImage.Image: guidance_scale=0.0, num_inference_steps=num_inference_steps, seed=seed_i, + variation=variation, ) return result.images[0] @@ -339,7 +356,7 @@ def _run_api_batch( self._generation.update_progress("complete", 100, None, None) self._generation.complete_generation([str(p) for p in output_paths]) - return GenerateImageCompleteResponse(status="complete", image_paths=[str(p) for p in output_paths]) + return GenerateImageCompleteResponse(status="complete", image_paths=[str(p) for p in output_paths], seed=seed) except HTTPError as e: self._generation.fail_generation(e.detail) for path in output_paths: diff --git a/backend/handlers/models_handler.py b/backend/handlers/models_handler.py index 4eaa25bee..a03d54df8 100644 --- a/backend/handlers/models_handler.py +++ b/backend/handlers/models_handler.py @@ -292,6 +292,17 @@ def get_ltx_recommendation(self) -> LtxRecommendationResponse: optional_cp_ids=self._ordered_cp_ids(self._get_optional_ltx_cp_ids(current_model_id)), ) + # A required checkpoint for the current model can be missing even when its base + # transformer is present — e.g. a hotfixed shared companion (the 2x upscaler) that + # superseded the version already on disk. Surface that download before offering any + # base upgrade: the current setup needs it regardless of whether the user upgrades, + # and routing it through the 'download' status lets the missing-models gate prompt it. + missing_current = self._ordered_cp_ids( + self._get_missing_cp_ids(self._get_required_ltx_cp_ids(current_model_id)) + ) + if missing_current: + return LtxDownloadRecommendationResponse(status="download", cps_to_download=missing_current) + if current_model_id == latest_model_id: return LtxOkRecommendationResponse(status="ok") diff --git a/backend/handlers/pipelines_handler.py b/backend/handlers/pipelines_handler.py index aff78e8ad..7314a4092 100644 --- a/backend/handlers/pipelines_handler.py +++ b/backend/handlers/pipelines_handler.py @@ -3,19 +3,20 @@ from __future__ import annotations import logging +from collections.abc import Mapping from threading import RLock from typing import TYPE_CHECKING from _routes._errors import HTTPError -from api_types import LTXLocalModelId +from api_types import ImageGenerationModelCheckpointID, LTXLocalModelId from handlers.base import StateHandlerBase from handlers.text_handler import TextHandler from runtime_config.model_download_specs import ( - IMG_GEN_MODEL_CP_ID, get_existing_cp_path, resolve_active_ltx_model_id, ) from runtime_config.runtime_policy import streaming_prefetch_count_for_mode +from services.patches import aux_block_cache, diffusion_stage_cache from services.interfaces import ( A2VPipeline, DepthProcessorPipeline, @@ -24,6 +25,7 @@ GpuCleaner, IcLoraPipeline, PoseProcessorPipeline, + QwenMultiAnglePipeline, RetakePipeline, VideoPipelineModelType, ) @@ -36,6 +38,7 @@ GenerationRunning, GpuSlot, ICLoraState, + QwenMultiAngleState, RetakePipelineState, VideoPipelineState, ) @@ -54,24 +57,27 @@ def __init__( text_handler: TextHandler, gpu_cleaner: GpuCleaner, fast_video_pipeline_class: type[FastVideoPipeline], - image_generation_pipeline_class: type[ImageGenerationPipeline], + image_generation_pipeline_classes: Mapping[ImageGenerationModelCheckpointID, type[ImageGenerationPipeline]], ic_lora_pipeline_class: type[IcLoraPipeline], depth_processor_pipeline_class: type[DepthProcessorPipeline], pose_processor_pipeline_class: type[PoseProcessorPipeline], a2v_pipeline_class: type[A2VPipeline], retake_pipeline_class: type[RetakePipeline], config: RuntimeConfig, + qwen_multiangle_pipeline_class: type[QwenMultiAnglePipeline] | None = None, ) -> None: super().__init__(state, lock, config) self._text_handler = text_handler self._gpu_cleaner = gpu_cleaner self._fast_video_pipeline_class = fast_video_pipeline_class - self._image_generation_pipeline_class = image_generation_pipeline_class + self._image_generation_pipeline_classes = image_generation_pipeline_classes + self._active_image_generation_cp_id: ImageGenerationModelCheckpointID | None = None self._ic_lora_pipeline_class = ic_lora_pipeline_class self._depth_processor_pipeline_class = depth_processor_pipeline_class self._pose_processor_pipeline_class = pose_processor_pipeline_class self._a2v_pipeline_class = a2v_pipeline_class self._retake_pipeline_class = retake_pipeline_class + self._qwen_multiangle_pipeline_class = qwen_multiangle_pipeline_class self._runtime_device = get_device_type(self.config.device) def _resolve_ltx_paths(self, model_id: LTXLocalModelId, gemma_root: str | None): @@ -188,6 +194,11 @@ def unload_gpu_pipeline(self) -> None: self._ensure_no_running_generation() self.state.gpu_slot = None self._assert_invariants() + # Drop any session-scoped cached transformer with its pipeline: unloading + # means giving the memory back, including the streaming entry's pinned host + # RAM. Before cleanup() so the pass reclaims what eviction released. + diffusion_stage_cache.evict() + aux_block_cache.evict() self._gpu_cleaner.cleanup() def park_image_generation_pipeline_on_cpu(self) -> None: @@ -217,27 +228,52 @@ def park_image_generation_pipeline_on_cpu(self) -> None: self.state.cpu_slot = CpuSlot(active_pipeline=image_generation_pipeline) self._assert_invariants() - def load_image_generation_pipeline_to_gpu(self) -> ImageGenerationPipeline: + def load_image_generation_pipeline_to_gpu( + self, cp_id: ImageGenerationModelCheckpointID + ) -> ImageGenerationPipeline: with self._lock: if self.state.gpu_slot is not None: active = self.state.gpu_slot.active_pipeline if isinstance(active, ImageGenerationPipeline): - return active - self._ensure_no_running_generation() + if self._active_image_generation_cp_id == cp_id: + return active + self._ensure_no_running_generation() + # Wrong model loaded — drop it rather than reuse, so we never + # silently generate with a different checkpoint than requested. + self.state.gpu_slot = None + else: + self._ensure_no_running_generation() + # A video-family pipeline holds the GPU — drop it too, so its + # weights don't occupy VRAM underneath the image pipeline. + self.state.gpu_slot = None + self._assert_invariants() image_generation_pipeline: ImageGenerationPipeline | None = None with self._lock: match self.state.cpu_slot: - case CpuSlot(active_pipeline=stored): + case CpuSlot(active_pipeline=stored) if self._active_image_generation_cp_id == cp_id: image_generation_pipeline = stored self.state.cpu_slot = None + case CpuSlot(): + # Parked pipeline is for a different model — discard it too. + self.state.cpu_slot = None case _: - image_generation_pipeline = None + pass + + # Release whatever was just dropped BEFORE loading the new model, so the + # image pipeline never has to fit alongside evicted weights. This includes + # any session-scoped cached transformer (this path drops a video pipeline + # inline above and does NOT route through _evict_gpu_pipeline_for_swap, so + # it needs its own cache evict -- a no-op when nothing is cached). + diffusion_stage_cache.evict() + aux_block_cache.evict() + self._gpu_cleaner.cleanup() if image_generation_pipeline is None: - zit_path = get_existing_cp_path(self.models_dir, IMG_GEN_MODEL_CP_ID) - image_generation_pipeline = self._image_generation_pipeline_class.create(str(zit_path), self._runtime_device) + pipeline_class = self._image_generation_pipeline_classes[cp_id] + model_path = get_existing_cp_path(self.models_dir, cp_id) + image_generation_pipeline = pipeline_class.create(str(model_path), self._runtime_device) else: image_generation_pipeline.to(self._runtime_device) @@ -245,31 +281,64 @@ def load_image_generation_pipeline_to_gpu(self) -> ImageGenerationPipeline: with self._lock: self.state.gpu_slot = GpuSlot(active_pipeline=image_generation_pipeline) + self._active_image_generation_cp_id = cp_id self._assert_invariants() return image_generation_pipeline def _evict_gpu_pipeline_for_swap(self) -> None: - should_park_image_generation_pipeline = False - should_cleanup = False - + # FORK CHANGE — diverges from upstream (see FORK.md "must survive"). + # Upstream parks an active image pipeline in host RAM here so a later + # image<->image switch stays warm (via park_image_generation_pipeline_on_cpu, + # now unused). But every caller of this method is loading a memory-hungry + # video-class pipeline (video / IC-LoRA / a2v / retake), for which a parked + # image model is pure dead weight: on a 64 GB box it pushes the video model's + # bf16-read + fp8-pin working set past physical RAM and into pagefile + # thrashing — and image->video (i2v) is the *common* workflow. So free the + # image pipeline outright instead of parking it, and also drop any pipeline + # parked by a previous swap. Returning to image gen reloads from its fast + # NF4/disk cache, far cheaper than thrashing every video render. with self._lock: self._ensure_no_running_generation() - if self.state.gpu_slot is None: + if self.state.gpu_slot is None and self.state.cpu_slot is None: return + self.state.gpu_slot = None + self.state.cpu_slot = None + self._active_image_generation_cp_id = None + self._assert_invariants() - active = self.state.gpu_slot.active_pipeline - if isinstance(active, ImageGenerationPipeline): - should_park_image_generation_pipeline = True - else: - self.state.gpu_slot = None - self._assert_invariants() - should_cleanup = True + # A pipeline swap invalidates the session-scoped transformer cache: the + # incoming pipeline (different checkpoint/LoRAs, or an image model) needs + # both the cached build's VRAM slice and, for streaming entries, its ~23GB + # of pinned host RAM. Evict before cleanup() so the pass reclaims it. + diffusion_stage_cache.evict() + aux_block_cache.evict() + self._gpu_cleaner.cleanup() + + def load_qwen_multiangle_pipeline(self) -> QwenMultiAngleState: + with self._lock: + match self.state.gpu_slot: + case GpuSlot(active_pipeline=QwenMultiAngleState() as state): + return state + case _: + pass + + if self._qwen_multiangle_pipeline_class is None: + raise HTTPError(500, "Qwen multi-angle pipeline is not configured") + + self._evict_gpu_pipeline_for_swap() + + import torch + + device = torch.device(self._runtime_device) + pipeline = self._qwen_multiangle_pipeline_class.create(device=device) + state = QwenMultiAngleState(pipeline=pipeline) + + with self._lock: + self.state.gpu_slot = GpuSlot(active_pipeline=state) + self._assert_invariants() + return state - if should_park_image_generation_pipeline: - self.park_image_generation_pipeline_on_cpu() - elif should_cleanup: - self._gpu_cleaner.cleanup() def evict_gpu_pipeline_for_prompt_enhancement(self) -> None: """Free whatever big pipeline is resident before a standalone Gemma enhance call. diff --git a/backend/handlers/qwen_multiangle_handler.py b/backend/handlers/qwen_multiangle_handler.py new file mode 100644 index 000000000..0c466b9d1 --- /dev/null +++ b/backend/handlers/qwen_multiangle_handler.py @@ -0,0 +1,137 @@ +"""Qwen multi-angle API orchestration handler.""" + +from __future__ import annotations + +import base64 +import io +import math +import time +import uuid +from threading import RLock + +from PIL import Image +from PIL.Image import Resampling + +from _routes._errors import HTTPError +from api_types import QwenMultiAngleGenerateRequest, QwenMultiAngleGenerateResponse +from handlers.base import StateHandlerBase +from handlers.generation_handler import GenerationHandler +from handlers.pipelines_handler import PipelinesHandler +from runtime_config.runtime_config import RuntimeConfig +from services.qwen_multiangle_pipeline.angle_mapping import compose_prompt, snap_pose +from state.app_state_types import AppState + +# The pipeline internally rescales condition images to ~1MP snapped to /32. +# Resizing here first keeps memory/latency predictable regardless of what the +# user dropped in (a raw camera photo, a screenshot, a prior GPM asset). +TARGET_AREA = 1024 * 1024 + + +class QwenMultiAngleHandler(StateHandlerBase): + def __init__( + self, + state: AppState, + lock: RLock, + config: RuntimeConfig, + generation_handler: GenerationHandler, + pipelines_handler: PipelinesHandler, + ) -> None: + super().__init__(state, lock, config) + self._generation = generation_handler + self._pipelines = pipelines_handler + + def generate(self, req: QwenMultiAngleGenerateRequest) -> QwenMultiAngleGenerateResponse: + if self._generation.is_generation_running(): + raise HTTPError(409, "Generation already in progress") + + source_image = _decode_data_url(req.image_data_url) + source_image = _model_resize(source_image) + extra_images = [_model_resize(_decode_data_url(u)) for u in req.extra_image_data_urls] + + # Roles only apply when aligned index-for-index with the refs; a mismatch + # (or none supplied) falls back to classic compose with no compositing clause. + extra_roles = list(req.extra_image_roles) if len(req.extra_image_roles) == len(extra_images) else [] + + seed = req.seed + if req.randomize_seed: + seed = int(time.time() * 1000) % 2147483647 + + pose = snap_pose(req.azimuth_deg, req.elevation_deg, req.zoom) + prompt = compose_prompt(pose, req.extra_prompt, extra_roles) + + generation_id = uuid.uuid4().hex[:8] + + try: + pipeline_state = self._pipelines.load_qwen_multiangle_pipeline() + self._generation.start_generation(generation_id) + self._generation.update_progress("loading_model", 5, 0, 1) + + def _on_step(step: int, total_steps: int) -> None: + progress = 10 + int((step / total_steps) * 85) + self._generation.update_progress("inference", progress, step, total_steps) + + result_image = pipeline_state.pipeline.generate( + image=source_image, + extra_images=extra_images, + extra_roles=extra_roles, + azimuth_deg=req.azimuth_deg, + elevation_deg=req.elevation_deg, + zoom=req.zoom, + seed=seed, + extra_prompt=req.extra_prompt, + quality_mode=req.quality_mode, + use_skin=req.use_skin, + skin_weight=req.skin_weight, + on_step=_on_step, + ) + + if self._generation.is_generation_cancelled(): + raise RuntimeError("Generation was cancelled") + + self._generation.update_progress("complete", 100, 1, 1) + result_data_url = _encode_png_data_url(result_image) + self._generation.complete_generation(result_data_url) + + return QwenMultiAngleGenerateResponse( + status="complete", + image_data_url=result_data_url, + prompt=prompt, + seed=seed, + ) + except HTTPError: + self._generation.fail_generation("Multi-angle generation failed") + raise + except Exception as exc: + self._generation.fail_generation(str(exc)) + raise HTTPError(500, f"Generation error: {exc}") from exc + + +def _decode_data_url(data_url: str) -> Image.Image: + _, _, encoded = data_url.partition(",") + if not encoded: + raise HTTPError(400, "image_data_url must be a data: URL") + try: + raw = base64.b64decode(encoded) + image = Image.open(io.BytesIO(raw)) + image.load() + except Exception as exc: + raise HTTPError(400, f"Not a readable image: {exc}") from exc + return image.convert("RGB") + + +def _encode_png_data_url(image: Image.Image) -> str: + buf = io.BytesIO() + image.save(buf, format="PNG") + encoded = base64.b64encode(buf.getvalue()).decode() + return f"data:image/png;base64,{encoded}" + + +def _model_resize(image: Image.Image) -> Image.Image: + """Preserve aspect, area-normalize to ~1MP, snap dimensions to /32.""" + width, height = image.size + scale = math.sqrt(TARGET_AREA / (width * height)) + new_width = max(32, round(width * scale / 32) * 32) + new_height = max(32, round(height * scale / 32) * 32) + if (new_width, new_height) == (width, height): + return image + return image.resize((new_width, new_height), Resampling.LANCZOS) diff --git a/backend/handlers/settings_handler.py b/backend/handlers/settings_handler.py index 0d68f21d6..2ccbaad75 100644 --- a/backend/handlers/settings_handler.py +++ b/backend/handlers/settings_handler.py @@ -65,6 +65,16 @@ def save_settings(self) -> None: payload = self.get_settings_snapshot().model_dump(by_alias=False) with open(self.config.settings_file, "w", encoding="utf-8") as f: json.dump(payload, f, indent=2) + # DIAGNOSTIC (settings-persistence bug): settings.json was observed + # untouched since 2026-07-11 despite successful POST /api/settings + # round-trips ("Applied settings patch" logged, no save warning). + # Log the resolved path + size on every save so the next session + # shows definitively whether this code runs and where it writes. + logger.info( + "Settings saved to %s (%d bytes)", + self.config.settings_file, + self.config.settings_file.stat().st_size, + ) except Exception as exc: logger.warning("Could not save settings: %s", exc, exc_info=True) @@ -98,6 +108,12 @@ def update_settings(self, patch: UpdateSettingsRequest) -> tuple[AppSettings, Ap # unrunnable active model (base present, companions missing). patch_payload.pop("active_ltx_model_id", None) + # active_ltx_model_id is patchable here only because the patch model is auto-derived + # from every AppSettings field — but it must go through set_active_ltx_model, which + # checks the full required bundle. Drop it so the generic PATCH can't persist an + # unrunnable active model (base present, companions missing). + patch_payload.pop("active_ltx_model_id", None) + before = self.state.app_settings.model_copy(deep=True) before_payload = ensure_json_object(before.model_dump(by_alias=False)) diff --git a/backend/handlers/video_generation_handler.py b/backend/handlers/video_generation_handler.py index 58634eb47..daa69874a 100644 --- a/backend/handlers/video_generation_handler.py +++ b/backend/handlers/video_generation_handler.py @@ -161,6 +161,14 @@ def _local_pixels( model_id = self._active_ltx_model_id() if model_id is None: raise HTTPError(409, "NO_DOWNLOADED_LTX_MODEL") + # FORK: 21:9 (~2.37:1) ultrawide is local-only and not modelled in ltx_capabilities. + # Map it directly, capped at 540p/720p (1080p ultrawide is too heavy). Dims are + # pre-rounded to /64 so generate rounding doesn't distort the ratio. + if aspect == "21:9": + ultrawide = {"540p": (1216, 512), "720p": (1664, 704)}.get(resolution) + if ultrawide is None: + raise HTTPError(400, "UNSUPPORTED_ULTRAWIDE_RESOLUTION") + return ultrawide try: return pixels_for(local_caps(model_id), resolution, aspect) except KeyError as exc: @@ -302,7 +310,7 @@ def generate(self, req: GenerateVideoRequest) -> GenerateVideoResponse: ) self._generation.complete_generation(output_path) - return GenerateVideoCompleteResponse(status="complete", video_path=output_path) + return GenerateVideoCompleteResponse(status="complete", video_path=output_path, seed=seed) except HTTPError as e: self._generation.fail_generation(e.detail) @@ -520,7 +528,7 @@ def _generate_a2v( self._generation.update_progress("complete", 100, total_steps, total_steps) self._generation.complete_generation(str(output_path)) - return GenerateVideoCompleteResponse(status="complete", video_path=str(output_path)) + return GenerateVideoCompleteResponse(status="complete", video_path=str(output_path), seed=seed) except HTTPError as e: self._generation.fail_generation(e.detail) diff --git a/backend/ltx2_server.py b/backend/ltx2_server.py index 6031d0e2a..c38675f54 100644 --- a/backend/ltx2_server.py +++ b/backend/ltx2_server.py @@ -73,6 +73,10 @@ del _ic_lora_stage2_lora import services.patches.diffusion_stage_cache as _diffusion_stage_cache # pyright: ignore[reportUnusedImport] # EXPERIMENTAL: remove once DiffusionStage caches/reuses identical builds upstream del _diffusion_stage_cache +import services.patches.aux_block_cache as _aux_block_cache_import # pyright: ignore[reportUnusedImport] # EXPERIMENTAL: remove once blocks.py caches/reuses identical aux builds upstream +del _aux_block_cache_import +import services.patches.fp8_sidecar_cache as _fp8_sidecar_cache # pyright: ignore[reportUnusedImport] # EXPERIMENTAL: remove if ltx-core caches post-sd_ops block weights on disk +del _fp8_sidecar_cache import services.patches.diffvae_mps_tiling_budget as _diffvae_mps_tiling_budget # pyright: ignore[reportUnusedImport] # Remove once ltx-pipelines queries MPS/unified free memory del _diffvae_mps_tiling_budget import services.patches.diffvae_decode_vram as _diffvae_decode_vram # pyright: ignore[reportUnusedImport] # Remove once ltx-pipelines offloads the transformer before DiffVAE decode @@ -214,7 +218,8 @@ def _resolve_app_data_dir() -> Path: APP_DATA_DIR = _resolve_app_data_dir() -DEFAULT_MODELS_DIR = APP_DATA_DIR / "models" +_env_models_dir = os.environ.get("LTX_MODELS_DIR") +DEFAULT_MODELS_DIR = Path(_env_models_dir) if _env_models_dir else (APP_DATA_DIR / "models") DEFAULT_MODELS_DIR.mkdir(parents=True, exist_ok=True) PROJECT_ROOT = Path(__file__).parent.parent @@ -283,6 +288,18 @@ def _resolve_local_generations_mode() -> LocalGenerationMode: LOCAL_GENERATIONS_MODE = _resolve_local_generations_mode() +# Opt the diffusion-stage and aux-block caches' session-scoped streaming kinds in +# only for the streaming runtime mode. Load-bearing gate: IC-LoRA's use_lora_in_stage_2 +# forces a CPU-mode streaming stage_2 even on full-loading (5090) cards, which must NOT +# get session-cached there -- see those modules' streaming-kind docstring sections. +import services.patches.aux_block_cache as _aux_block_cache_gate +import services.patches.diffusion_stage_cache as _diffusion_stage_cache_gate + +_streaming_mode = LOCAL_GENERATIONS_MODE == "streaming_models_loading" +_aux_block_cache_gate.set_streaming_enabled(_streaming_mode) +_diffusion_stage_cache_gate.set_streaming_enabled(_streaming_mode) +del _aux_block_cache_gate, _diffusion_stage_cache_gate + CAMERA_MOTION_PROMPTS = { "none": "", "static": ", static camera, locked off shot, no camera movement", diff --git a/backend/pyproject.toml b/backend/pyproject.toml index 7092a1f63..ab345b7be 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -11,9 +11,6 @@ dependencies = [ # torch to 2.11.x keeps torch+torchaudio in lockstep. win/linux use the cu128 index below. "torch>=2.3.0; sys_platform != 'darwin'", "torch>=2.3.0,<2.12; sys_platform == 'darwin'", - # transformers' Gemma 4 processor imports torchvision.transforms.v2 at module scope, so - # LTX 2.5's text encoder fails to build without it (ImportError on Gemma4UnifiedProcessor). - "torchvision>=0.18.0", "huggingface-hub>=0.23.0", "tqdm>=4.66.0", "pynvml>=11.5.0; sys_platform != 'darwin'", @@ -35,7 +32,17 @@ dependencies = [ # ltx-core 1.2 / Gemma 4 need transformers 5.8+; 5.15 makes Gemma 4 # attention dims per-layer (AmbiguousGlobalPerLayerAttributeError). "transformers>=5.14.1,<5.15", + # FORK: Krea 2 loads its transformer with a bitsandbytes NF4 config; without + # bitsandbytes the Krea2Pipeline quantized path fails. Upstream has no Krea 2. + "bitsandbytes>=0.48.0; sys_platform != 'darwin'", "sentencepiece>=0.1.99", + # Qwen multi-angle pipeline: GGUF checkpoint loading + the Qwen2.5-VL + # video-processor import path (raises ImportError without torchvision + # present, even though we never touch video). torchvision is also required + # by transformers' Gemma 4 processor (LTX 2.5 text encoder), which imports + # torchvision.transforms.v2 at module scope (Gemma4UnifiedProcessor). + "gguf>=0.10.0", + "torchvision>=0.18.0", "sageattention>=1.0.0; sys_platform != 'darwin'", # Official NATTEN wheels are Linux-only. GCS-hosted Windows wheel matches # the embed runtime (cp313, torch 2.10.0+cu128). Do not use ltx-core's @@ -77,6 +84,11 @@ torchvision = [ # Linux keeps resolving sageattention from PyPI (no reported Blackwell crashes there yet). sageattention = { url = "https://github.com/woct0rdho/SageAttention/releases/download/v2.2.0-windows.post5/sageattention-2.2.0+cu128torch2.10.0andhigher.post5-cp310-abi3-win_amd64.whl", marker = "sys_platform == 'win32'" } natten = { url = "https://storage.googleapis.com/ltx-desktop-artifacts/wheels/natten-0.21.6+torch2100cu128-cp313-cp313-win_amd64.whl", marker = "sys_platform == 'win32'" } +# FORK: keep our newer diffusers rev, NOT upstream's. Ours (0d32f805, 2026-04-15) is a +# strict descendant of upstream's (01de02e8, 2026-02-20) — 380 commits ahead, 0 behind — so +# it contains everything upstream needs, plus Krea2Pipeline, which theirs predates. +# Downgrading to upstream's pin breaks Krea 2 at import time. See FORK.md section B. +diffusers = { git = "https://github.com/huggingface/diffusers.git", rev = "0d32f8054438fd38204ae2d46155d2c971c90da5" } # Official LTX-2 v1.2.0 packages (not on PyPI). Same git+subdirectory pattern as # 2.3; the tag is the 2.5-capable release. We do not take ltx-core's own # tool.uv.sources (cu132) — this app's pytorch-cu128 pin above owns torch. diff --git a/backend/runtime_config/lora_catalog.json b/backend/runtime_config/lora_catalog.json index 2448e082e..82e0f49fb 100644 --- a/backend/runtime_config/lora_catalog.json +++ b/backend/runtime_config/lora_catalog.json @@ -1511,7 +1511,9 @@ ], "default_settings": { "skip_stage_2": true, - "resolution_factor": 1.5, + "use_lora_in_stage_2": true, + "resolution_factor": 2.0, + "lora_strength": 1.4, "audio_mode": "generated" }, "prompt_template": { diff --git a/backend/runtime_config/ltx_capabilities.py b/backend/runtime_config/ltx_capabilities.py index 436995e96..00c732f5b 100644 --- a/backend/runtime_config/ltx_capabilities.py +++ b/backend/runtime_config/ltx_capabilities.py @@ -29,7 +29,7 @@ "camera_motion", "auto_duration", ] -LtxAspectRatio = Literal["16:9", "9:16"] +LtxAspectRatio = Literal["16:9", "9:16", "21:9"] @dataclass(frozen=True) diff --git a/backend/runtime_config/model_download_specs.py b/backend/runtime_config/model_download_specs.py index 9bafad251..28cbe5ad0 100644 --- a/backend/runtime_config/model_download_specs.py +++ b/backend/runtime_config/model_download_specs.py @@ -339,6 +339,14 @@ def get_model_cp_spec(cp_id: ModelCheckpointID) -> ModelCheckpointSpec: repo_id="Tongyi-MAI/Z-Image-Turbo", description="Z-Image-Turbo model for text-to-image generation", ) + case "krea-2-turbo": + return ModelCheckpointSpec( + relative_path=Path("Krea-2-Turbo"), + expected_size_bytes=35_700_000_000, + is_folder=True, + repo_id="krea/Krea-2-Turbo", + description="Krea 2 Turbo model for text-to-image generation", + ) case _: assert_never(cp_id) diff --git a/backend/services/gpu_cleaner/torch_cleaner.py b/backend/services/gpu_cleaner/torch_cleaner.py index 3e22e73b8..74766a2b2 100644 --- a/backend/services/gpu_cleaner/torch_cleaner.py +++ b/backend/services/gpu_cleaner/torch_cleaner.py @@ -16,5 +16,8 @@ def __init__(self, device: str | torch.device = "cpu") -> None: self._device = device def cleanup(self) -> None: - empty_device_cache(self._device) + # Collect first: dropped pipelines hold reference cycles, so their + # tensors only return to the allocator during gc — emptying the cache + # before that releases nothing. gc.collect() + empty_device_cache(self._device) diff --git a/backend/services/ic_lora_pipeline/ltx_ic_lora_pipeline.py b/backend/services/ic_lora_pipeline/ltx_ic_lora_pipeline.py index fcd140f3f..4e31b1488 100644 --- a/backend/services/ic_lora_pipeline/ltx_ic_lora_pipeline.py +++ b/backend/services/ic_lora_pipeline/ltx_ic_lora_pipeline.py @@ -25,6 +25,8 @@ logger = logging.getLogger(__name__) +logger = logging.getLogger(__name__) + class LTXIcLoraPipeline: @staticmethod diff --git a/backend/services/image_generation_pipeline/image_generation_pipeline.py b/backend/services/image_generation_pipeline/image_generation_pipeline.py index 83178d738..9984aa443 100644 --- a/backend/services/image_generation_pipeline/image_generation_pipeline.py +++ b/backend/services/image_generation_pipeline/image_generation_pipeline.py @@ -24,6 +24,7 @@ def generate( guidance_scale: float, num_inference_steps: int, seed: int, + variation: float = 0.0, ) -> ImagePipelineOutputLike: ... diff --git a/backend/services/image_generation_pipeline/krea2_image_generation_pipeline.py b/backend/services/image_generation_pipeline/krea2_image_generation_pipeline.py new file mode 100644 index 000000000..595099076 --- /dev/null +++ b/backend/services/image_generation_pipeline/krea2_image_generation_pipeline.py @@ -0,0 +1,317 @@ +"""Krea 2 Turbo image generation pipeline wrapper.""" + +# Krea2Pipeline / Krea2Transformer2DModel are untyped in diffusers, so every value +# derived from them is "unknown" under pyright's strict mode. Suppress the unknown-type +# family for this file rather than scattering per-line ignores that can't fully resolve +# a third-party class with no stubs. (Fork-only file — see FORK.md section B.) +# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false +# pyright: reportUnknownParameterType=false, reportUnknownArgumentType=false + +from __future__ import annotations + +import logging +from collections.abc import Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Any, cast + +import torch +from diffusers import BitsAndBytesConfig # type: ignore[reportPrivateImportUsage] +from diffusers import Krea2Pipeline # type: ignore[reportUnknownVariableType] +from diffusers import Krea2Transformer2DModel # type: ignore[reportPrivateImportUsage] +from PIL.Image import Image as PILImage +from PIL.Image import Resampling +from transformers import BitsAndBytesConfig as TransformersBitsAndBytesConfig +from transformers.models.qwen3_vl.modeling_qwen3_vl import Qwen3VLModel + +from services.generation_interrupt import diffusers_step_callback +from services.image_generation_pipeline.variation_boost import ( + boost_step_count, + clamp_variation, + log_boost, + perturb_embeds, + restoring_callback, +) +from services.services_utils import ( + ImagePipelineOutputLike, + PILImageType, + empty_device_cache, + get_device_type, + sync_device, +) + +logger = logging.getLogger(__name__) + +# Cached beside the model files themselves (never touches the originals) so +# subsequent loads read the already-quantized weights instead of the full +# bf16 originals. Cuts cold-load time roughly 8x and makes the load far less +# likely to be hit by the original files getting evicted from the OS file +# cache under RAM pressure (e.g. from a large video model) - the text +# encoder is a single ~9 GB file with no shard-level progress reporting, so +# that eviction previously showed up as a multi-minute silent stall. +_NF4_TRANSFORMER_CACHE_DIRNAME = "_nf4_transformer_cache" +_NF4_TEXT_ENCODER_CACHE_DIRNAME = "_nf4_text_encoder_cache" + + +@dataclass(slots=True) +class _Krea2Output: + images: Sequence[PILImageType] + + +class Krea2ImageGenerationPipeline: + @staticmethod + def create( + model_path: str, + device: str | None = None, + ) -> "Krea2ImageGenerationPipeline": + return Krea2ImageGenerationPipeline(model_path=model_path, device=device) + + @staticmethod + def _load_quantized_transformer(model_path: str) -> Krea2Transformer2DModel: + cache_dir = Path(model_path) / _NF4_TRANSFORMER_CACHE_DIRNAME + if (cache_dir / "config.json").exists(): + try: + return Krea2Transformer2DModel.from_pretrained( # type: ignore[reportUnknownMemberType] + cache_dir, torch_dtype=torch.bfloat16 + ) + except Exception: + logger.warning( + "Failed to load cached NF4 Krea 2 transformer from %s; re-quantizing from source.", + cache_dir, + exc_info=True, + ) + + quantization_config = BitsAndBytesConfig( + load_in_4bit=True, + bnb_4bit_quant_type="nf4", + bnb_4bit_compute_dtype=torch.bfloat16, + ) + transformer = Krea2Transformer2DModel.from_pretrained( # type: ignore[reportUnknownMemberType] + model_path, + subfolder="transformer", + quantization_config=quantization_config, + torch_dtype=torch.bfloat16, + ) + try: + transformer.save_pretrained(cache_dir) # type: ignore[reportUnknownMemberType] + except OSError: + logger.warning("Failed to cache NF4 Krea 2 transformer to %s", cache_dir, exc_info=True) + return transformer + + @staticmethod + def _load_quantized_text_encoder(model_path: str) -> Qwen3VLModel: + cache_dir = Path(model_path) / _NF4_TEXT_ENCODER_CACHE_DIRNAME + if (cache_dir / "config.json").exists(): + try: + return Qwen3VLModel.from_pretrained(cache_dir, torch_dtype=torch.bfloat16) # type: ignore[reportUnknownMemberType] + except Exception: + logger.warning( + "Failed to load cached NF4 Krea 2 text encoder from %s; re-quantizing from source.", + cache_dir, + exc_info=True, + ) + + quantization_config = TransformersBitsAndBytesConfig( + load_in_4bit=True, + bnb_4bit_quant_type="nf4", + bnb_4bit_compute_dtype=torch.bfloat16, + ) + text_encoder = Qwen3VLModel.from_pretrained( # type: ignore[reportUnknownMemberType] + model_path, + subfolder="text_encoder", + quantization_config=quantization_config, + torch_dtype=torch.bfloat16, + ) + try: + text_encoder.save_pretrained(cache_dir) # type: ignore[reportUnknownMemberType] + except OSError: + logger.warning("Failed to cache NF4 Krea 2 text encoder to %s", cache_dir, exc_info=True) + return text_encoder + + def __init__(self, model_path: str, device: str | None = None) -> None: + self._device: str | None = None + self._cpu_offload_active = False + + # NF4 quantization only benefits CUDA: it shrinks the ~24 GB bf16 + # transformer to ~7 GB and the ~9 GB text encoder to ~3 GB, small + # enough to fit on a consumer GPU as whole components (no + # block-level streaming needed), which is both faster (no per-block + # copy overhead) and lighter to keep resident when parked between + # generations. bitsandbytes has no MPS backend, so other devices + # keep the plain bf16 components. + if get_device_type(device) == "cuda": + transformer = self._load_quantized_transformer(model_path) + text_encoder = self._load_quantized_text_encoder(model_path) + self.pipeline = Krea2Pipeline.from_pretrained( # type: ignore[reportUnknownMemberType] + model_path, + transformer=transformer, + text_encoder=text_encoder, + torch_dtype=torch.bfloat16, + ) + else: + self.pipeline = Krea2Pipeline.from_pretrained( # type: ignore[reportUnknownMemberType] + model_path, + torch_dtype=torch.bfloat16, + ) + + if device is not None: + self.to(device) + + def _resolve_generator_device(self) -> str: + if self._cpu_offload_active: + return "cuda" + if self._device is not None: + return self._device + + execution_device = getattr(self.pipeline, "_execution_device", None) + return get_device_type(execution_device) + + @staticmethod + def _normalize_output(output: object) -> ImagePipelineOutputLike: + images = getattr(output, "images", None) + if not isinstance(images, Sequence): + raise RuntimeError("Unexpected Krea 2 pipeline output format: missing images sequence") + + images_list = cast(Sequence[object], images) + validated_images: list[PILImageType] = [] + for image in images_list: + if not isinstance(image, PILImage): + raise RuntimeError("Unexpected Krea 2 pipeline output format: images must be PIL.Image instances") + validated_images.append(image) + + return _Krea2Output(images=validated_images) + + # Quantizing the transformer shrinks its *weights* (~24 GB -> ~7 GB), but + # activation memory still scales with pixel count regardless of + # quantization. Measured on a 24 GB card: 1344x768 reserves ~17 GB and + # runs at full speed; 1536x896 reserves ~24 GB (no safety margin left for + # other apps); 1728x992 reserves ~32 GB, overflowing into system RAM + # (~40x slower). Keep the same conservative cap as the pre-quantization + # implementation. + _MAX_GENERATION_PIXELS = 1344 * 768 + + @torch.inference_mode() + def generate( + self, + prompt: str, + height: int, + width: int, + guidance_scale: float, + num_inference_steps: int, + seed: int, + variation: float = 0.0, + ) -> ImagePipelineOutputLike: + requested_width, requested_height = width, height + if width * height > self._MAX_GENERATION_PIXELS: + scale = (self._MAX_GENERATION_PIXELS / (width * height)) ** 0.5 + width = max(16, int(width * scale) // 16 * 16) + height = max(16, int(height * scale) // 16 * 16) + + device = self._resolve_generator_device() + generator = torch.Generator(device=device).manual_seed(seed) + pipeline = cast(Any, self.pipeline) + prompt_kwargs, callback, callback_inputs = self._variation_inputs( + pipeline, prompt, seed, variation, num_inference_steps + ) + + try: + output = pipeline( + **prompt_kwargs, + height=height, + width=width, + guidance_scale=guidance_scale, + num_inference_steps=num_inference_steps, + generator=generator, + output_type="pil", + return_dict=True, + # Poll the cooperative cancel Event between denoising steps so Stop + # actually aborts a Krea 2 run (this was missing, unlike the Z-Image + # path, so cancel spun forever and only exiting the app freed the GPU). + callback_on_step_end=callback, + callback_on_step_end_tensor_inputs=callback_inputs, + ) + finally: + # Leftover allocator cache would starve the next run. + sync_device(device) + empty_device_cache(device) + + result = self._normalize_output(output) + if (width, height) != (requested_width, requested_height): + # Deliver the size the caller asked for; Lanczos upscale from the + # model's native ~1 MP is visually clean and keeps VRAM bounded. + result = _Krea2Output( + images=[ + img.resize((requested_width, requested_height), Resampling.LANCZOS) for img in result.images + ] + ) + return result + + # Variation boost (see variation_boost.py). Krea 2 embeds are + # [batch, seq, text_layers, dim] with a padding mask; noise on padded positions is + # harmless since attention masks them out. Not separately tuned: carries over the + # Z-Image finding that noising only the FIRST step keeps prompt adherence (2 noisy + # steps drifted off-prompt), so 0.15 x 6 steps => 1 step. + _VARIATION_MAX_NOISE = 1.0 + _VARIATION_MASK_FRACTION = 0.5 + _VARIATION_BOOST_FRACTION = 0.15 + + def _variation_inputs( + self, + pipeline: Any, + prompt: str, + seed: int, + variation: float, + num_inference_steps: int, + ) -> tuple[dict[str, Any], Any, list[str]]: + variation = clamp_variation(variation) + if variation <= 0.0: + return {"prompt": prompt}, diffusers_step_callback, ["latents"] + + clean, mask = pipeline.encode_prompt(prompt=prompt, device=pipeline._execution_device) + noisy = perturb_embeds( + clean, + seed=seed, + variation=variation, + max_noise=self._VARIATION_MAX_NOISE, + mask_fraction=self._VARIATION_MASK_FRACTION, + ) + boost_steps = boost_step_count(num_inference_steps, self._VARIATION_BOOST_FRACTION) + log_boost("Krea 2", variation, boost_steps, num_inference_steps) + return ( + {"prompt_embeds": noisy, "prompt_embeds_mask": mask}, + restoring_callback(clean, boost_steps), + ["prompt_embeds"], + ) + + def edit( + self, + prompt: str, + image: PILImageType, + strength: float, + num_inference_steps: int, + seed: int, + ) -> ImagePipelineOutputLike: + # Krea 2 Turbo is text-to-image only. Image editing is routed to Z-Image by + # ImageGenerationHandler._edit; this exists solely to satisfy the + # ImageGenerationPipeline protocol and must never be called for Krea 2. + raise NotImplementedError("Krea 2 Turbo does not support image editing") + + def to(self, device: str) -> None: + runtime_device = get_device_type(device) + if runtime_device in ("cuda", "mps"): + # The NF4-quantized transformer (~7 GB) fits on the accelerator as + # a whole component, so simple whole-model offloading is enough - + # no custom block-level streaming needed. + # + # Only install the accelerate hooks once per move-to-accelerator. Calling + # enable_model_cpu_offload() again while offload is already active leaves + # modules with accelerate's wrapped forward but no _hf_hook attribute, which + # fails at inference ("object has no attribute '_hf_hook'"). A park-to-CPU + # cycle takes the else-branch and clears the flag, so coming back re-installs. + if not (self._cpu_offload_active and self._device == runtime_device): + self.pipeline.enable_model_cpu_offload() # type: ignore[reportUnknownMemberType] + self._cpu_offload_active = True + else: + self._cpu_offload_active = False + self.pipeline.to(runtime_device) # type: ignore[reportUnknownMemberType] + self._device = runtime_device diff --git a/backend/services/image_generation_pipeline/variation_boost.py b/backend/services/image_generation_pipeline/variation_boost.py new file mode 100644 index 000000000..07b9119c9 --- /dev/null +++ b/backend/services/image_generation_pipeline/variation_boost.py @@ -0,0 +1,78 @@ +"""Variation boost for distilled, guidance-free image models (Z-Image Turbo, Krea 2 Turbo). + +Distilled turbo models let the prompt pin the composition, so different seeds barely move +it. Adding seeded noise to a random subset of the text-embedding values for the first +step(s) — where layout is decided — knocks each seed into a different composition; the +clean embeddings are restored afterwards so the remaining steps refine detail normally. +Costs nothing measurable (the prompt is encoded once, same as without the boost), and the +same seed + variation reproduces the same image. +""" + +from __future__ import annotations + +import logging +from collections.abc import Callable +from typing import Any, TypeVar + +import torch + +from services.generation_interrupt import diffusers_step_callback + +logger = logging.getLogger(__name__) + +StepCallback = Callable[[object, int, object, dict[str, Any]], dict[str, Any]] +_E = TypeVar("_E", torch.Tensor, list[torch.Tensor]) + + +def perturb_embeds( + clean: _E, + *, + seed: int, + variation: float, + max_noise: float, + mask_fraction: float = 0.5, +) -> _E: + """Noise with std = variation x max_noise x the embeddings' own std, on a random mask.""" + # Seeded on CPU so results don't depend on device RNG. + gen = torch.Generator(device="cpu").manual_seed((seed * 7919 + 1) % (2**63)) + + def one(emb: torch.Tensor) -> torch.Tensor: + shape = tuple(emb.shape) + noise = torch.randn(shape, generator=gen, dtype=torch.float32) + mask = torch.rand(shape, generator=gen) < mask_fraction + scale = float(emb.float().std()) * max_noise * variation + return emb + (noise * mask * scale).to(device=emb.device, dtype=emb.dtype) + + if isinstance(clean, torch.Tensor): + return one(clean) + return [one(e) for e in clean] + + +def boost_step_count(num_inference_steps: int, boost_fraction: float) -> int: + return max(1, round(num_inference_steps * boost_fraction)) + + +def restoring_callback(clean: object, boost_steps: int) -> StepCallback: + """Interrupt-aware step callback that swaps the clean embeds back in after boost_steps.""" + + def callback(pipe: object, step_index: int, timestep: object, kwargs: dict[str, Any]) -> dict[str, Any]: + kwargs = diffusers_step_callback(pipe, step_index, timestep, kwargs) + if step_index == boost_steps - 1: + kwargs["prompt_embeds"] = clean + return kwargs + + return callback + + +def clamp_variation(variation: float) -> float: + return max(0.0, min(1.0, variation)) + + +def log_boost(model: str, variation: float, boost_steps: int, num_inference_steps: int) -> None: + logger.info( + "%s variation boost %.2f: noisy prompt embeds for %d/%d steps", + model, + variation, + boost_steps, + num_inference_steps, + ) diff --git a/backend/services/image_generation_pipeline/zit_image_generation_pipeline.py b/backend/services/image_generation_pipeline/zit_image_generation_pipeline.py index 777a5d0e4..45a62f6c9 100644 --- a/backend/services/image_generation_pipeline/zit_image_generation_pipeline.py +++ b/backend/services/image_generation_pipeline/zit_image_generation_pipeline.py @@ -1,26 +1,55 @@ """Z-Image-Turbo image generation pipeline wrapper.""" +# ZImagePipeline / ZImageTransformer2DModel are untyped in diffusers, so values +# derived from them read as "unknown" under pyright's strict mode. Suppress the +# unknown-type family for this file rather than scattering per-line ignores. +# pyright: reportUnknownMemberType=false, reportUnknownVariableType=false +# pyright: reportUnknownParameterType=false, reportUnknownArgumentType=false + from __future__ import annotations import logging from collections.abc import Sequence from dataclasses import dataclass +from pathlib import Path from typing import Any, cast import torch +from diffusers import BitsAndBytesConfig # type: ignore[reportPrivateImportUsage] +from diffusers import ZImageTransformer2DModel # type: ignore[reportPrivateImportUsage] from diffusers.pipelines.auto_pipeline import ZImagePipeline # type: ignore[reportUnknownVariableType] from PIL.Image import Image as PILImage +from PIL.Image import Resampling +from transformers import BitsAndBytesConfig as TransformersBitsAndBytesConfig +from transformers.models.qwen3.modeling_qwen3 import Qwen3Model from services.generation_interrupt import diffusers_step_callback +from services.image_generation_pipeline.variation_boost import ( + boost_step_count, + clamp_variation, + log_boost, + perturb_embeds, + restoring_callback, +) from services.services_utils import ( ImagePipelineOutputLike, PILImageType, clamp_strength, + empty_device_cache, get_device_type, + sync_device, ) logger = logging.getLogger(__name__) +# Cached beside the model files themselves (never touches the originals) so +# subsequent loads read the already-quantized weights instead of re-quantizing +# the full bf16 originals. Cuts cold-load time roughly 8x and makes the load far +# less likely to stall on the OS file cache evicting the large source shards +# under RAM pressure (e.g. from a video model). +_NF4_TRANSFORMER_CACHE_DIRNAME = "_nf4_transformer_cache" +_NF4_TEXT_ENCODER_CACHE_DIRNAME = "_nf4_text_encoder_cache" + @dataclass(slots=True) class _ZImageOutput: @@ -35,22 +64,105 @@ def create( ) -> "ZitImageGenerationPipeline": return ZitImageGenerationPipeline(model_path=model_path, device=device) + @staticmethod + def _load_quantized_transformer(model_path: str) -> ZImageTransformer2DModel: + cache_dir = Path(model_path) / _NF4_TRANSFORMER_CACHE_DIRNAME + if (cache_dir / "config.json").exists(): + try: + return ZImageTransformer2DModel.from_pretrained( # type: ignore[reportUnknownMemberType] + cache_dir, torch_dtype=torch.bfloat16 + ) + except Exception: + logger.warning( + "Failed to load cached NF4 Z-Image transformer from %s; re-quantizing from source.", + cache_dir, + exc_info=True, + ) + + quantization_config = BitsAndBytesConfig( + load_in_4bit=True, + bnb_4bit_quant_type="nf4", + bnb_4bit_compute_dtype=torch.bfloat16, + ) + transformer = ZImageTransformer2DModel.from_pretrained( # type: ignore[reportUnknownMemberType] + model_path, + subfolder="transformer", + quantization_config=quantization_config, + torch_dtype=torch.bfloat16, + ) + try: + transformer.save_pretrained(cache_dir) # type: ignore[reportUnknownMemberType] + except OSError: + logger.warning("Failed to cache NF4 Z-Image transformer to %s", cache_dir, exc_info=True) + return transformer + + @staticmethod + def _load_quantized_text_encoder(model_path: str) -> Qwen3Model: + cache_dir = Path(model_path) / _NF4_TEXT_ENCODER_CACHE_DIRNAME + if (cache_dir / "config.json").exists(): + try: + return Qwen3Model.from_pretrained(cache_dir, torch_dtype=torch.bfloat16) # type: ignore[reportUnknownMemberType] + except Exception: + logger.warning( + "Failed to load cached NF4 Z-Image text encoder from %s; re-quantizing from source.", + cache_dir, + exc_info=True, + ) + + quantization_config = TransformersBitsAndBytesConfig( + load_in_4bit=True, + bnb_4bit_quant_type="nf4", + bnb_4bit_compute_dtype=torch.bfloat16, + ) + text_encoder = Qwen3Model.from_pretrained( # type: ignore[reportUnknownMemberType] + model_path, + subfolder="text_encoder", + quantization_config=quantization_config, + torch_dtype=torch.bfloat16, + ) + try: + text_encoder.save_pretrained(cache_dir) # type: ignore[reportUnknownMemberType] + except OSError: + logger.warning("Failed to cache NF4 Z-Image text encoder to %s", cache_dir, exc_info=True) + return text_encoder + def __init__(self, model_path: str, device: str | None = None) -> None: self._device: str | None = None self._cpu_offload_active = False self._img2img: Any = None - self.pipeline = ZImagePipeline.from_pretrained( # type: ignore[reportUnknownMemberType] - model_path, - torch_dtype=torch.bfloat16, - ) + + # NF4 quantization only benefits CUDA: it shrinks the ~24.6 GB bf16 + # transformer to ~7 GB and the ~8 GB text encoder to ~3 GB. Loaded as + # raw bf16 the transformer alone fills a 24 GB card, so + # enable_model_cpu_offload()'s whole-module move overflowed into shared + # system RAM and thrashed forever (uninterruptible — cancel is only + # polled between denoise steps). Quantized, both components fit as whole + # modules with room for activations. bitsandbytes has no MPS backend, so + # other devices keep the plain bf16 components. + if get_device_type(device) == "cuda": + transformer = self._load_quantized_transformer(model_path) + text_encoder = self._load_quantized_text_encoder(model_path) + self.pipeline = ZImagePipeline.from_pretrained( # type: ignore[reportUnknownMemberType] + model_path, + transformer=transformer, + text_encoder=text_encoder, + torch_dtype=torch.bfloat16, + ) + else: + self.pipeline = ZImagePipeline.from_pretrained( # type: ignore[reportUnknownMemberType] + model_path, + torch_dtype=torch.bfloat16, + ) + if device is not None: self.to(device) def _resolve_generator_device(self) -> str: # The configured runtime device is authoritative. With enable_model_cpu_offload() # the pipeline's _execution_device can read as "cpu", but the generator must live - # on the actual compute device. Offload is CUDA-only; if we ever lose `_device` - # while offload is active, CUDA is the only remaining compute device. + # on the actual compute device. This previously returned "cuda" whenever offload + # was active — but offload is enabled on MPS too, so on a Mac it built a CUDA + # generator and failed ("Cannot get CUDA generator without ATen_cuda library"). if self._device is not None: return self._device if self._cpu_offload_active: @@ -84,6 +196,14 @@ def _normalize_output(output: object) -> ImagePipelineOutputLike: return _ZImageOutput(images=validated_images) + # Quantizing the transformer shrinks its *weights* (~24.6 GB -> ~7 GB), but + # activation memory still scales with pixel count regardless of quantization. + # Cap generation at ~1 MP (matching the Krea 2 path) so a large requested + # resolution can't push the denoise activations back into shared system RAM + # and re-introduce the thrash this quantization work fixes; the result is + # Lanczos-upscaled to the requested size. + _MAX_GENERATION_PIXELS = 1344 * 768 + @torch.inference_mode() def generate( self, @@ -93,24 +213,82 @@ def generate( guidance_scale: float, num_inference_steps: int, seed: int, + variation: float = 0.0, ) -> ImagePipelineOutputLike: # ZImagePipeline ignores guidance_scale, so we drop it explicitly. _ = guidance_scale + requested_width, requested_height = width, height + if width * height > self._MAX_GENERATION_PIXELS: + scale = (self._MAX_GENERATION_PIXELS / (width * height)) ** 0.5 + width = max(16, int(width * scale) // 16 * 16) + height = max(16, int(height * scale) // 16 * 16) + self._log_run("generate", width=width, height=height) - generator = torch.Generator(device=self._resolve_generator_device()).manual_seed(seed) + device = self._resolve_generator_device() + generator = torch.Generator(device=device).manual_seed(seed) pipeline = cast(Any, self.pipeline) - output = pipeline( - prompt=prompt, - height=height, - width=width, - guidance_scale=0.0, - num_inference_steps=num_inference_steps, - generator=generator, - output_type="pil", - return_dict=True, - callback_on_step_end=diffusers_step_callback, + prompt_kwargs, callback, callback_inputs = self._variation_inputs( + pipeline, prompt, seed, variation, num_inference_steps ) - return self._normalize_output(output) + try: + output = pipeline( + **prompt_kwargs, + height=height, + width=width, + guidance_scale=0.0, + num_inference_steps=num_inference_steps, + generator=generator, + output_type="pil", + return_dict=True, + callback_on_step_end=callback, + callback_on_step_end_tensor_inputs=callback_inputs, + ) + finally: + # Leftover allocator cache would starve the next run. + sync_device(device) + empty_device_cache(device) + + result = self._normalize_output(output) + if (width, height) != (requested_width, requested_height): + result = _ZImageOutput( + images=[ + img.resize((requested_width, requested_height), Resampling.LANCZOS) for img in result.images + ] + ) + return result + + # Variation boost (see variation_boost.py). Tuned 2026-09-27 on the 3090 (4 steps): + # noising only the FIRST step keeps prompt adherence up to 1.0 with big + # layout/pose/camera changes; noising 2 of 4 steps broke the prompt from ~0.75 up. + _VARIATION_MAX_NOISE = 1.0 + _VARIATION_MASK_FRACTION = 0.5 + _VARIATION_BOOST_FRACTION = 0.25 + + def _variation_inputs( + self, + pipeline: Any, + prompt: str, + seed: int, + variation: float, + num_inference_steps: int, + ) -> tuple[dict[str, Any], Any, list[str]]: + variation = clamp_variation(variation) + if variation <= 0.0: + return {"prompt": prompt}, diffusers_step_callback, ["latents"] + + clean, _ = pipeline.encode_prompt( + prompt=prompt, device=pipeline._execution_device, do_classifier_free_guidance=False + ) + noisy = perturb_embeds( + list(clean), + seed=seed, + variation=variation, + max_noise=self._VARIATION_MAX_NOISE, + mask_fraction=self._VARIATION_MASK_FRACTION, + ) + boost_steps = boost_step_count(num_inference_steps, self._VARIATION_BOOST_FRACTION) + log_boost("Z-Image", variation, boost_steps, num_inference_steps) + return {"prompt_embeds": noisy}, restoring_callback(clean, boost_steps), ["prompt_embeds"] def _ensure_img2img_pipeline(self) -> Any: if self._img2img is not None: @@ -120,6 +298,7 @@ def _ensure_img2img_pipeline(self) -> Any: except Exception as e: raise RuntimeError("DIFFUSERS_IMG2IMG_UNAVAILABLE") from e # Reuse the loaded components so no second copy of the weights lands in VRAM. + # With the CUDA path these are the quantized transformer/text encoder. pipeline_any = cast(Any, self.pipeline) self._img2img = ZImageImg2ImgPipeline(**pipeline_any.components) # type: ignore[reportUnknownMemberType] return self._img2img @@ -136,25 +315,37 @@ def edit( img2img = self._ensure_img2img_pipeline() self._log_run("edit") - generator = torch.Generator(device=self._resolve_generator_device()).manual_seed(seed) - output = img2img( - prompt=prompt, - image=image, - strength=clamp_strength(strength), - num_inference_steps=num_inference_steps, - guidance_scale=0.0, # Turbo is guidance-free; img2img defaults to 5.0. - generator=generator, - output_type="pil", - return_dict=True, - callback_on_step_end=diffusers_step_callback, - ) + device = self._resolve_generator_device() + generator = torch.Generator(device=device).manual_seed(seed) + try: + output = img2img( + prompt=prompt, + image=image, + strength=clamp_strength(strength), + num_inference_steps=num_inference_steps, + guidance_scale=0.0, # Turbo is guidance-free; img2img defaults to 5.0. + generator=generator, + output_type="pil", + return_dict=True, + callback_on_step_end=diffusers_step_callback, + ) + finally: + sync_device(device) + empty_device_cache(device) return self._normalize_output(output) def to(self, device: str) -> None: runtime_device = get_device_type(device) if runtime_device == "cuda": - self.pipeline.enable_model_cpu_offload() # type: ignore[reportUnknownMemberType] - self._cpu_offload_active = True + # enable_model_cpu_offload() installs accelerate hooks. Re-invoking it while + # offload is already active for this device leaves modules holding + # accelerate's wrapped forward but no _hf_hook attribute, which then blows up + # at inference: "'Qwen3Model' object has no attribute '_hf_hook'". Install it + # once; a park-to-CPU cycle takes the else-branch below and clears the flag, + # so moving back to the accelerator re-installs correctly. + if not (self._cpu_offload_active and self._device == runtime_device): + self.pipeline.enable_model_cpu_offload() # type: ignore[reportUnknownMemberType] + self._cpu_offload_active = True else: # MPS is unified memory: cpu-offload keeps a host copy *and* a GPU copy of # the active module in the same RAM pool, which raised peak and jetsam'd Macs diff --git a/backend/services/interfaces.py b/backend/services/interfaces.py index 55a9dcdb3..c940cbe39 100644 --- a/backend/services/interfaces.py +++ b/backend/services/interfaces.py @@ -15,6 +15,7 @@ from services.image_generation_pipeline.image_generation_pipeline import ImageGenerationPipeline from services.ltx_api_client.ltx_api_client import LTXAPIClient from services.retake_pipeline.retake_pipeline import RetakePipeline +from services.qwen_multiangle_pipeline.qwen_multiangle_pipeline import QwenMultiAnglePipeline from services.model_downloader.model_downloader import ModelDownloader from services.pose_processor_pipeline.pose_processor_pipeline import PoseProcessorPipeline from services.prompt_enhancer_pipeline.prompt_enhancer_pipeline import PromptEnhancerPipeline @@ -48,6 +49,7 @@ "IcLoraPipeline", "LTXAPIClient", "RetakePipeline", + "QwenMultiAnglePipeline", "TextEncoder", "PromptEnhancerPipeline", ] diff --git a/backend/services/patches/aux_block_cache.py b/backend/services/patches/aux_block_cache.py new file mode 100644 index 000000000..378e53974 --- /dev/null +++ b/backend/services/patches/aux_block_cache.py @@ -0,0 +1,328 @@ +"""Monkey-patch (EXPERIMENTAL): session-scoped cache for the small aux models +rebuilt on every generation. + +Each generation rebuilds five small models from the checkpoint on EVERY call +(ltx_pipelines/utils/blocks.py, "builds on call, frees on exit"): the VAE +encoder (built twice per i2v gen: ImageConditioner at gen start and again +inside VideoUpsampler between stages), the spatial upsampler, the video +decoder, and the audio decoder + vocoder. Measured on the RTX 3090 (warm +35.2s gen, session log 2026-07-22): ~3-4s of a warm generation is this build +churn. This patch caches the BUILT models across calls and generations, the +Phase B companion to diffusion_stage_cache's transformer cache. + +Unlike the transformer cache this is a MULTI-SLOT dict keyed on builder +content: the entries are small bf16 modules held VRAM-RESIDENT (est. ~2-4GB +total; the 3090 shows 17-22GB free during denoise). ImageConditioner's and +VideoUpsampler's encoder builders are constructed from identical constants + +checkpoint, so they share ONE cache entry by content key. + +Gated to streaming mode (``set_streaming_enabled`` pushed by ltx2_server, +``streaming_models_loading`` only) exactly like the transformer cache: the +gate structurally excludes full-loading (5090-class) cards, so the documented +5090 VRAM-collision incident (a 23GB VRAM-resident transformer surviving into +the next generation's builds) cannot recur here -- and the aux set is an +order of magnitude smaller besides. Aux entries do NOT participate in the +generation-start RAM check (they are VRAM, not pinned host RAM) and are kept +across generations unconditionally; eviction happens at the same +PipelinesHandler unload/swap/image-load sites as the transformer cache +(paired explicit calls -- no facade, greppable), on the shared Settings +toggle, and on the streaming-gate kill switch. + +NO DIRTY TRACKING (divergence from diffusion_stage_cache): its dirty flag +exists because an abnormal unwind skips the streaming wrapper's forward +post-hooks and leaks BufferPool slots. Aux modules are plain stateless +nn.Modules -- no pools, no hooks, no cross-call state; an exception +mid-decode leaves them bit-identical. + +THE BYPASS BRANCH DOES NOT EVICT (second divergence): the transformer +cache's bypass evicts because a non-cacheable build needs the 23GB the cache +holds; evicting five warm ~GB-scale models because one call was a zombie +bypass would thrash for no memory benefit. + +CONCURRENCY: same in-use bypass as the transformer cache (zombie +generations observed live 2026-07-22): a checkout that finds the entry +checked out falls back to the original build-per-call path -- concurrent +sharing of one module across CUDA streams is unproven, and pre-cache each +call had a private model, so the bypass restores exactly that. Lookup and +in-use bump share ONE critical section (TOCTOU). + +VideoDecoder specifics: upstream returns ``_cleanup_iter(decoder.decode_video +(...), decoder)`` whose gpu_model teardown meta-swaps the decoder when the +chunk iterator is exhausted OR abandoned (GeneratorExit). The cached path +must therefore ``yield from decode_video`` WITHOUT gpu_model; the checkout +happens at first ``next()`` (a never-iterated generator must not leak +in_use), and the ``finally`` releases + ``cleanup_memory()`` (allocator trim, +NO meta-swap) on both exhaustion and abandonment. A caller-supplied custom +``decoder_builder`` (the multi-GPU path) is excluded via the same +``isinstance(builder, SingleGPUModelBuilder)`` check the transformer cache +uses. + +EXCLUDED: PromptEncoder -- in API-encoding mode the fork's text-encoder +patches return before any build, so caching buys nothing (future work for +local-Gemma users); AudioConditioner -- unused by the fast video pipeline. + +EXPERIMENTAL: depends on the private surfaces ``ImageConditioner +._encoder_builder/_dtype/_device``, ``VideoUpsampler._encoder_builder/ +_upsampler_builder``, ``VideoDecoder._decoder_builder``, ``AudioDecoder +._decoder_builder/_vocoder_builder`` (+ each ``__call__`` body, incl. the +vocoder's effective-dtype rule: fp32 on MPS, build dtype elsewhere) -- re- +verify against ltx_pipelines.utils.blocks on rev bumps. + +Usage: + import services.patches.aux_block_cache # noqa: F401 +""" + +from __future__ import annotations + +import gc +import logging +import os +import threading +from collections.abc import Callable, Iterator +from dataclasses import dataclass, field +from typing import TypeVar + +import torch + +from ltx_core.devices import synchronize_device +from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder +from ltx_pipelines.utils import blocks as _blocks +from ltx_pipelines.utils.blocks import ( + AudioDecoder, + ImageConditioner, + VideoDecoder, + VideoUpsampler, +) +from ltx_pipelines.utils.helpers import cleanup_memory + +logger = logging.getLogger(__name__) + +_T = TypeVar("_T") +_CacheKey = tuple[object, ...] + +_lock = threading.Lock() +_enabled = os.environ.get("AUX_BLOCK_CACHE_ENABLED", "1") != "0" +_streaming_enabled = False + + +@dataclass +class _Entry: + model: torch.nn.Module + in_use: int = field(default=0) + + +_cache: dict[_CacheKey, _Entry] = {} + + +def set_enabled(value: bool) -> None: + """Pushed alongside diffusion_stage_cache.set_enabled from the shared + Settings toggle at every generation start. Turning off evicts immediately.""" + global _enabled + with _lock: + if not value: + _evict_locked() + _enabled = value + + +def set_streaming_enabled(value: bool) -> None: + """Opted in by ltx2_server for streaming_models_loading only (same + load-bearing gate as the transformer cache). Env kill switch: + AUX_BLOCK_CACHE_STREAMING=0.""" + global _streaming_enabled + effective = value and os.environ.get("AUX_BLOCK_CACHE_STREAMING", "1") != "0" + with _lock: + if not effective: + _evict_locked() + _streaming_enabled = effective + logger.info("[aux-block-cache] session cache %s", "enabled" if effective else "disabled") + + +def _evict_locked() -> None: + busy = [key for key, entry in _cache.items() if entry.in_use > 0] + if busy: + raise RuntimeError( + "[aux-block-cache] evict requested while cached models are in use " + f"(mid-generation): {busy!r} -- overlapping generation detected." + ) + if _cache: + synchronize_device() + gc.collect() + for entry in _cache.values(): + entry.model.to("meta") + cleanup_memory() + _cache.clear() + + +def evict() -> None: + """Free and drop all cached aux models. + + Called by PipelinesHandler at the same unload/swap/image-load sites as + diffusion_stage_cache.evict(). Safe when empty; raises if any entry is + checked out (see _evict_locked).""" + with _lock: + _evict_locked() + + +def _cacheable(builder: object, device: torch.device) -> bool: + return ( + _enabled + and _streaming_enabled + and device.type == "cuda" # MPS unified memory: resident aux competes with system RAM + and isinstance(builder, SingleGPUModelBuilder) # excludes multi-GPU custom decoder builders + ) + + +def _key(builder: SingleGPUModelBuilder, dtype: torch.dtype, device: torch.device) -> _CacheKey: + # Content key in the transformer cache's style; dtype is the EFFECTIVE + # build dtype (the vocoder is fp32 on MPS). registry deliberately excluded + # (a caching layer, not content). + return ( + builder.model_path, + builder.model_sd_ops, + builder.module_ops, + builder.loras, + type(builder).__name__, + dtype, + device, + ) + + +def _checkout( + builder: SingleGPUModelBuilder, dtype: torch.dtype, device: torch.device, kind: str +) -> _Entry | None: + """Resolve-or-build and mark in use in ONE critical section (TOCTOU). + + Returns None when the entry is already checked out (zombie/overlapping + generation) -- the caller must fall back to the original build-per-call + path rather than share a module across concurrent CUDA streams.""" + key = _key(builder, dtype, device) + with _lock: + entry = _cache.get(key) + if entry is not None and entry.in_use > 0: + return None + if entry is None: + model = builder.build(device=device, dtype=dtype).eval() + entry = _Entry(model=model) + _cache[key] = entry + hit = False + else: + hit = True + entry.in_use += 1 + logger.info("[aux-block-cache] %s %s", "reusing" if hit else "built + cached", kind) + return entry + + +def _release(entry: _Entry) -> None: + with _lock: + entry.in_use = max(0, entry.in_use - 1) + + +_orig_image_conditioner_call = ImageConditioner.__call__ +_orig_video_upsampler_call = VideoUpsampler.__call__ +_orig_video_decoder_call = VideoDecoder.__call__ +_orig_audio_decoder_call = AudioDecoder.__call__ + + +def _cached_image_conditioner_call(self: ImageConditioner, fn: Callable[..., _T]) -> _T: + if not _cacheable(self._encoder_builder, self._device): # noqa: SLF001 + return _orig_image_conditioner_call(self, fn) + entry = _checkout(self._encoder_builder, self._dtype, self._device, "vae_encoder") # noqa: SLF001 + if entry is None: + logger.warning("[aux-block-cache] vae_encoder in use (overlapping generation?) -- isolated build") + return _orig_image_conditioner_call(self, fn) + try: + return fn(entry.model) + finally: + _release(entry) + + +def _cached_video_upsampler_call(self: VideoUpsampler, latent: torch.Tensor) -> torch.Tensor: + # All-or-nothing: never mix a cached half with a gpu_model-managed half. + if not ( + _cacheable(self._encoder_builder, self._device) # noqa: SLF001 + and _cacheable(self._upsampler_builder, self._device) # noqa: SLF001 + ): + return _orig_video_upsampler_call(self, latent) + encoder = _checkout(self._encoder_builder, self._dtype, self._device, "vae_encoder") # noqa: SLF001 + if encoder is None: + logger.warning("[aux-block-cache] vae_encoder in use (overlapping generation?) -- isolated build") + return _orig_video_upsampler_call(self, latent) + try: + upsampler = _checkout(self._upsampler_builder, self._dtype, self._device, "upsampler") # noqa: SLF001 + if upsampler is None: + logger.warning("[aux-block-cache] upsampler in use (overlapping generation?) -- isolated build") + return _orig_video_upsampler_call(self, latent) + try: + return _blocks.upsample_video( + latent=latent, video_encoder=encoder.model, upsampler=upsampler.model + ) + finally: + _release(upsampler) + finally: + _release(encoder) + + +def _cached_video_decoder_call( + self: VideoDecoder, + latent: torch.Tensor, + tiling_config: object = None, + generator: torch.Generator | None = None, + *, + dtype: torch.dtype | None = None, +) -> Iterator[torch.Tensor]: + # 1.2.0 added a keyword-only ``dtype`` override (float32 for HDR raw-in/raw-out); + # None keeps the constructor dtype. Mirror upstream: it selects the effective build + # dtype and casts the latent to it before decode. build_dtype flows into _checkout, + # which keys the cache on dtype, so an HDR float32 call caches separately from SDR. + build_dtype = self._dtype if dtype is None else dtype # noqa: SLF001 + latent = latent.to(dtype=build_dtype) + if not _cacheable(self._decoder_builder, self._device): # noqa: SLF001 + return _orig_video_decoder_call(self, latent, tiling_config, generator, dtype=dtype) # type: ignore[arg-type] + + def _decode() -> Iterator[torch.Tensor]: + # Checkout at first next(): a never-iterated generator must not leak in_use. + entry = _checkout(self._decoder_builder, build_dtype, self._device, "video_decoder") # noqa: SLF001 + if entry is None: + logger.warning("[aux-block-cache] video_decoder in use (overlapping generation?) -- isolated build") + yield from _orig_video_decoder_call(self, latent, tiling_config, generator, dtype=dtype) # type: ignore[arg-type] + return + try: + # No gpu_model wrapper: its teardown would meta-swap the cached + # decoder. The finally runs on exhaustion AND GeneratorExit + # (cancel/abandon), mirroring upstream _cleanup_iter's semantics. + yield from entry.model.decode_video(latent, tiling_config, generator) # type: ignore[operator] + finally: + _release(entry) + cleanup_memory() # allocator trim only -- the model stays resident + + return _decode() + + +def _cached_audio_decoder_call(self: AudioDecoder, latent: torch.Tensor) -> object: + vocoder_dtype = torch.float32 if self._device.type == "mps" else self._dtype # noqa: SLF001 + if not ( + _cacheable(self._decoder_builder, self._device) # noqa: SLF001 + and _cacheable(self._vocoder_builder, self._device) # noqa: SLF001 + ): + return _orig_audio_decoder_call(self, latent) + decoder = _checkout(self._decoder_builder, self._dtype, self._device, "audio_decoder") # noqa: SLF001 + if decoder is None: + logger.warning("[aux-block-cache] audio_decoder in use (overlapping generation?) -- isolated build") + return _orig_audio_decoder_call(self, latent) + try: + vocoder = _checkout(self._vocoder_builder, vocoder_dtype, self._device, "vocoder") # noqa: SLF001 + if vocoder is None: + logger.warning("[aux-block-cache] vocoder in use (overlapping generation?) -- isolated build") + return _orig_audio_decoder_call(self, latent) + try: + return _blocks.vae_decode_audio(latent, decoder.model, vocoder.model) + finally: + _release(vocoder) + finally: + _release(decoder) + + +ImageConditioner.__call__ = _cached_image_conditioner_call # type: ignore[method-assign] +VideoUpsampler.__call__ = _cached_video_upsampler_call # type: ignore[method-assign] +VideoDecoder.__call__ = _cached_video_decoder_call # type: ignore[method-assign] +AudioDecoder.__call__ = _cached_audio_decoder_call # type: ignore[method-assign] diff --git a/backend/services/patches/diffusion_stage_cache.py b/backend/services/patches/diffusion_stage_cache.py index 7882b813a..357342f25 100644 --- a/backend/services/patches/diffusion_stage_cache.py +++ b/backend/services/patches/diffusion_stage_cache.py @@ -1,5 +1,5 @@ """Monkey-patch (EXPERIMENTAL): cache the built transformer across DiffusionStage -calls within one generation when nothing about its config actually changed. +calls when nothing about its config actually changed. ``DiffusionStage`` "builds on each call, frees on exit" (ltx_pipelines/utils/ blocks.py) unconditionally -- it rebuilds the transformer from the checkpoint @@ -10,7 +10,8 @@ with the SAME LoRAs (the common "fast" t2v/i2v case), that means loading + fp8-casting a ~22B-param transformer from disk TWICE per generation. Measured on an RTX 5090 (32 GB): ~40s per rebuild, ~80s of a ~132s 540p/8s generation -spent rebuilding the identical transformer twice. +spent rebuilding the identical transformer twice. Measured on an RTX 3090 +(24 GB, streaming): ~90-100s per rebuild, ~190s of a ~226s generation. Patches ``_transformer_ctx``, not the public ``model_context()`` wrapper: ``DiffusionStage.__call__`` invokes ``self._transformer_ctx(video_tools=...)`` @@ -19,55 +20,85 @@ itself just forwards to ``_transformer_ctx()``, so patching the latter covers both call sites. -Scoped to the standard (non-streaming, non-multi-GPU) build path: caching only -applies when ``self._is_streaming`` is False (i.e. ``local_generations_mode == -"full_models_loading"`` -- high-VRAM cards where holding one build resident is -affordable) AND the prepared builder is a plain ``SingleGPUModelBuilder``. -Streaming instances (the low-VRAM path) are untouched -- forcing residency -there would defeat the point of streaming. Multi-GPU tiled builders are also -excluded: they use ``video_tools`` to shard identically-shaped work across -devices, whereas ``SingleGPUModelBuilder.build()`` ignores all extra kwargs -(``**kwargs: object, # noqa: ARG002`` in single_gpu_model_builder.py) -- -confirmed inert for the path we cache, not assumed. - -GENERATION-SCOPED, not session-scoped (found live on an RTX 5090): letting the -cache survive PAST the generation that built it collides with every other -component that builds fresh per call too (text encoder, VAE, upsampler, audio -decoder/vocoder) -- the next generation's text-encoder build then has to -coexist in VRAM with the still-resident transformer from the PREVIOUS -generation, instead of the transformer having already been freed by then. -Observed: a second generation's peak VRAM was reported at ~41.8 GB on a -31.82 GB card (Windows CUDA fell back to slow shared memory, backend liveness -probe failed, total generation time regressed to 143s -- worse than no cache -at all). Fix: ``handlers.generation_handler.GenerationHandler.start_generation`` -/``start_api_generation`` call :func:`evict` before marking a new generation as -running, so the cache never survives past the generation it was built for. +TWO CACHED KINDS with different scopes: + +- "resident" (``full_models_loading``, 32GB+ cards): the standard + ``SingleGPUModelBuilder`` build held resident in VRAM. GENERATION-SCOPED -- + see below. +- "streaming" (``streaming_models_loading``, 15-30GB CUDA cards): the + ``StreamingModelBuilder`` build in CPU/RAM mode (``cpu_slots_count is + None``, all blocks fp8-cast + pinned in host RAM, streamed to VRAM + per-step). SESSION-SCOPED -- survives stage_1 -> stage_2 AND across + generations, ComfyUI-style ("pay the cold build once, evict only under + pressure"). Gated by :func:`set_streaming_enabled`, pushed once by + ltx2_server after resolving ``LOCAL_GENERATIONS_MODE``. This gate is + load-bearing, not cosmetic: IC-LoRA's ``use_lora_in_stage_2`` forces a + CPU-mode streaming stage_2 even on full-loading (5090) cards + (``LTXIcLoraPipeline._ensure_stage_2_streams_for_lora``); without the gate + that stage_2 would silently session-cache ~23GB of pinned host RAM on + hardware this feature is not for. Multi-GPU tiled builders and MPS disk + streaming (``cpu_slots_count == DISK_CPU_SLOTS``) stay excluded. + +GENERATION-SCOPED for the resident kind (found live on an RTX 5090): letting a +VRAM-resident cache survive PAST the generation that built it collides with +every other component that builds fresh per call too (text encoder, VAE, +upsampler, audio decoder/vocoder) -- the next generation's text-encoder build +then has to coexist in VRAM with the still-resident transformer from the +PREVIOUS generation. Observed: a second generation's peak VRAM was reported at +~41.8 GB on a 31.82 GB card (Windows CUDA fell back to slow shared memory, +backend liveness probe failed, total generation time regressed to 143s -- +worse than no cache at all). Fix: ``handlers.generation_handler. +GenerationHandler.start_generation``/``start_api_generation`` call +:func:`evict_for_generation_start` before marking a new generation as running. + +SESSION-SCOPED for the streaming kind: that VRAM rationale does not apply -- +the weights live in *pinned host RAM* (~23GB for the fp8-cast 22B model), not +VRAM. The wrapper does still hold SOME VRAM while cached (non-block weights +resident on GPU + 2 GPU block slots in the BufferPool + a dedicated copy +stream, est. 2-4GB), which now coexists with the next generation's +text-encode/VAE phases -- watch peak VRAM there when validating on new +hardware. :func:`evict_for_generation_start` keeps a streaming entry unless +available system RAM has dropped below ``_MIN_FREE_RAM_GB`` (external memory +pressure; for a matching-key local generation, evicting a warm entry is +strictly worse than keeping it -- the rebuild's transient working set costs +more RAM than the pinned entry it replaces). Unconditional eviction sites for +the streaming kind: pipeline unload/swap (``PipelinesHandler``: video<->image +swaps need both the VRAM slice and the pinned RAM back), checkpoint/LoRA +changes (key mismatch), the Settings toggle, and the bypass branch below. +The streaming build is reusable across calls by construction: provider/pool/ +copy-stream state is shape-independent (stage_1 half-res vs stage_2 full-res +is fine -- the wrapper already runs N full forward passes per denoise run), +and RoPE caches key on model config, not resolution. torch.compile + streaming +cache is UNTESTED (the fork skips compile under SageAttention); +``_compilation_config`` is in the key, so a config change misses rather than +corrupts. NON-CACHEABLE TRANSITIONS also evict (found live on the same RTX 5090, IC-LoRA this time): IC-LoRA's ``use_lora_in_stage_2`` forces stage_2 onto the streaming -path (``LTXIcLoraPipeline._ensure_stage_2_streams_for_lora``, a deliberate -existing VRAM-safety mechanism -- streaming halves stage_2's resident -footprint since it conditions on the full-res reference video). That stage_2 -call correctly skips the cache (``_is_streaming`` is True), but skipping the +path (a deliberate existing VRAM-safety mechanism -- streaming halves stage_2's +resident footprint since it conditions on the full-res reference video). On a +full-loading card that stage_2 call correctly skips the cache, but skipping the cache branch also means the cache-key-mismatch eviction below never fires -- so stage_1's cached transformer stayed resident while stage_2 built its own streaming transformer AND did its tiled conditioning VAE encode on top of it. Observed: reserved VRAM climbed to 41.68 GB on the 31.82 GB card and hung there for 150+ seconds with no progress (denoising loop never started). -Fix: the non-cacheable bypass branch (streaming stages, disabled setting) +Fix: the non-cacheable bypass branch (disabled setting, excluded builders) evicts unconditionally before delegating to the original method, not just the cache-hit/miss branch. Single-slot cache: only the most recently built transformer stays resident. -Switching to a different checkpoint/LoRA/quantization config (or starting a -new generation, per above) evicts the old one (frees it, mirroring -gpu_model's own TRIM teardown: sync + dead-reference ``gc.collect()`` + +Switching to a different checkpoint/LoRA/quantization config evicts the old +one (frees it, mirroring the upstream teardown for its kind: sync + +dead-reference ``gc.collect()`` + [``teardown()`` for streaming] + ``.to("meta")`` + ``cleanup_memory()``) before building the new one -- never -holds two builds resident at once. The pre-free ``gc.collect()`` mirrors -ComfyUI's mandatory cleanup_models_gc: ``.to("meta")`` only frees storage once -no live reference (a compiled wrapper / cudagraph pool) remains, so collecting -dead references first is what turns a cumulative compile-path creep into a flat -reserved floor. +holds two builds resident at once. For streaming, ``teardown()`` releases the +forward hooks and drops the pinned-block references so their +``cudaHostUnregister`` finalizers fire (mirrors _streaming_model's finally, +blocks.py:154-158). The pre-free ``gc.collect()`` mirrors ComfyUI's mandatory +cleanup_models_gc: ``.to("meta")`` only frees storage once no live reference +(a compiled wrapper / cudagraph pool) remains, so collecting dead references +first is what turns a cumulative compile-path creep into a flat reserved floor. CONCURRENCY: the single-slot module-global cache is only correct under strict sequentiality. That is enforced upstream by @@ -81,10 +112,14 @@ the lock -- is what protects a mid-denoise transformer from a concurrent evict. Cache key is the *content* of the prepared builder (checkpoint path, sd_ops, -module_ops, LoRAs) plus quantization identity, compilation config, dtype, and -device -- not object identity -- so two independently-constructed -DiffusionStage instances (e.g. a pipeline's stage_1 and stage_2) that happen to -build the same thing correctly share the cache within one generation. +module_ops, LoRAs) plus builder type + cpu_slots_count, quantization identity, +compilation config, dtype, and device -- not object identity -- so two +independently-constructed DiffusionStage instances (e.g. a pipeline's stage_1 +and stage_2) that happen to build the same thing correctly share the cache. +The builder type + cpu_slots_count discriminants are REQUIRED, not +belt-and-suspenders: an IC-LoRA resident stage_1 and its forced-streaming +stage_2 can otherwise share an identical content key, and a streaming stage +must never "hit" a cached resident X0Model (or vice versa). ``SDOps``/``ModuleOps`` are a frozen dataclass / NamedTuple respectively (structural equality), so this is safe: a real config difference always produces a different key (cache miss, falls back to a normal rebuild), never @@ -96,13 +131,22 @@ every generation, so flipping it in Settings takes effect on the next generation without a restart. Also settable via env ``DIFFUSION_STAGE_CACHE_ENABLED`` (default "1") as the initial value before -any setting is pushed -- e.g. for headless/dev runs. +any setting is pushed -- e.g. for headless/dev runs. The streaming session +scope additionally requires :func:`set_streaming_enabled` (pushed by +ltx2_server from the runtime mode) and can be killed independently via env +``DIFFUSION_STAGE_CACHE_STREAMING=0``. ``DIFFUSION_STAGE_CACHE_MIN_FREE_RAM_GB`` +(default 4) tunes the generation-start RAM-pressure eviction threshold. EXPERIMENTAL: depends on DiffusionStage's private ``_is_streaming``, ``_prepared_builder()``, ``_build_transformer()``, ``_quantization``, -``_compilation_config``, ``_dtype``, ``_device`` staying as-is, and on -``__call__`` continuing to route through ``_transformer_ctx`` -- re-verify -against ltx_pipelines.utils.blocks.DiffusionStage on rev bumps. +``_compilation_config``, ``_dtype``, ``_device`` staying as-is, on +``__call__`` continuing to route through ``_transformer_ctx``, and on the +streaming contract of ``_streaming_model``/``BlockStreamingWrapper`` +(``build(device, dtype)`` -> wrapper; ``teardown()`` frees pins + ``dispose()`` +metas the storage on evict, mirroring ``_streaming_model``'s finally; a fresh +``X0Model(wrapper).eval()`` per checkout) -- re-verify against +ltx_pipelines.utils.blocks on rev bumps. (1.2.0: ``_streaming_model`` swaps its +finally from a bare ``.to("meta")`` to ``teardown()`` + ``dispose()``.) Usage: import services.patches.diffusion_stage_cache # noqa: F401 @@ -116,26 +160,48 @@ import threading from collections.abc import Iterator from contextlib import contextmanager +from typing import Literal +import psutil + +from ltx_core.block_streaming import StreamingModelBuilder from ltx_core.devices import synchronize_device from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder +from ltx_core.model.transformer.model import X0Model from ltx_pipelines.utils.blocks import DiffusionStage from ltx_pipelines.utils.helpers import cleanup_memory logger = logging.getLogger(__name__) _CacheKey = tuple[object, ...] +_CacheKind = Literal["resident", "streaming"] _lock = threading.Lock() _enabled = os.environ.get("DIFFUSION_STAGE_CACHE_ENABLED", "1") != "0" +# Opted in by ltx2_server once LOCAL_GENERATIONS_MODE is resolved -- see the +# module docstring's TWO CACHED KINDS section for why this gate is load-bearing. +_streaming_enabled = False _cached_key: _CacheKey | None = None _cached_model: object | None = None +_cached_kind: _CacheKind | None = None +# Set when a checkout exits abnormally (exception/cancel unwinding through the +# denoising loop). An abnormal unwind skips the streaming wrapper's forward +# post-hooks, which can leak GPU BufferPool slots ("BufferPool exhausted: all 2 +# buffers are in use" observed live after a zombie generation) -- so a dirty +# entry must never be reused: it is evicted and rebuilt on the next checkout +# (and at generation start). +_dirty = False # >0 while a cached transformer is checked out (yielded to a caller) and possibly # mid-denoise. The single-slot cache is only safe under strict sequentiality; this # counter lets _evict_locked() fail loud if something tries to free a model that is # still in use, instead of ``.to("meta")``-ing tensors another generation is reading. _in_use = 0 +# Generation-start RAM-pressure threshold for keeping a session-scoped streaming +# entry. Deliberately low: for a matching-key generation, evicting a warm entry is +# strictly worse than keeping it (see module docstring's SESSION-SCOPED section). +_MIN_FREE_RAM_GB = float(os.environ.get("DIFFUSION_STAGE_CACHE_MIN_FREE_RAM_GB", "4")) + def set_enabled(value: bool) -> None: """Turn the cache on/off, checked on every ``_transformer_ctx`` call. @@ -154,8 +220,39 @@ def set_enabled(value: bool) -> None: _enabled = value +def set_streaming_enabled(value: bool) -> None: + """Opt the streaming (session-scoped) kind in or out. + + Called once by ltx2_server after resolving ``LOCAL_GENERATIONS_MODE`` + (True only for ``streaming_models_loading``). ``DIFFUSION_STAGE_CACHE_STREAMING=0`` + is an independent kill switch for the streaming kind (the resident kind is + unaffected). Turning off with a streaming entry cached evicts it first -- + if that raises (in use), the flag stays un-applied. + """ + global _streaming_enabled + effective = value and os.environ.get("DIFFUSION_STAGE_CACHE_STREAMING", "1") != "0" + with _lock: + if not effective and _cached_kind == "streaming": + _evict_locked() + _streaming_enabled = effective + logger.info( + "[diffusion-stage-cache] streaming session cache %s", + "enabled" if effective else "disabled", + ) + + def _cacheable(stage: DiffusionStage) -> bool: - return not stage._is_streaming and isinstance(stage._prepared_builder(), SingleGPUModelBuilder) # noqa: SLF001 + builder = stage._prepared_builder() # noqa: SLF001 + if stage._is_streaming: # noqa: SLF001 + # Only the CPU/RAM streaming mode (all blocks pinned in host RAM) is + # session-cacheable; MPS disk streaming (small cpu_slots_count) keeps a + # different memory model and stays on the build-per-call path. + return ( + _streaming_enabled + and isinstance(builder, StreamingModelBuilder) + and builder.cpu_slots_count is None + ) + return isinstance(builder, SingleGPUModelBuilder) def _cache_key(stage: DiffusionStage) -> _CacheKey: @@ -165,6 +262,11 @@ def _cache_key(stage: DiffusionStage) -> _CacheKey: builder.model_sd_ops, builder.module_ops, builder.loras, + # Builder type + slot count discriminate resident vs streaming builds of + # otherwise identical content (e.g. IC-LoRA's forced-streaming stage_2 on a + # full-loading card) -- a streaming stage must never hit a resident entry. + type(builder).__name__, + getattr(builder, "cpu_slots_count", None), # Policy objects aren't structurally comparable; identity is safe here # since pipelines construct one QuantizationPolicy and share it by # reference across stage_1/stage_2. @@ -176,7 +278,7 @@ def _cache_key(stage: DiffusionStage) -> _CacheKey: def _evict_locked() -> None: - global _cached_key, _cached_model + global _cached_key, _cached_model, _cached_kind if _in_use > 0: # Overlapping/concurrent generation: someone is trying to free a transformer # that is currently checked out (mid-denoise). The single-slot module-global @@ -199,17 +301,29 @@ def _evict_locked() -> None: # the compile-on soak shows a flat reserved floor (reclaim) or a slow creep # (leak) -- see the module docstring's GENERATION-SCOPED repro. gc.collect() - _cached_model.to("meta") # type: ignore[attr-defined] + if _cached_kind == "streaming": + # Mirror _streaming_model's finally (blocks.py): teardown() releases the + # forward hooks / disk I/O worker thread / pinned-block references so their + # cudaHostUnregister finalizers can fire, before the storage is disposed. + _cached_model.teardown() # type: ignore[attr-defined] + # 1.2.0 frees model storage via ``Disposable.dispose()`` (both gpu_model and + # _streaming_model use it, replacing the bare ``.to("meta")``): it metas the + # parameter / persistent-buffer storage and is shell-safe so fused LoRA weights + # do not linger on a cached module. X0Model and BlockStreamingWrapper are both + # ``nn.Module, Disposable``. cleanup_memory() then returns cached blocks to the OS. + _cached_model.dispose() # type: ignore[attr-defined] cleanup_memory() - _cached_key, _cached_model = None, None + globals()["_dirty"] = False + _cached_key, _cached_model, _cached_kind = None, None, None def evict() -> None: - """Free and drop any resident cached transformer. + """Free and drop any resident cached transformer, regardless of kind. - Called from ``GenerationHandler.start_generation``/``start_api_generation`` - so the cache never survives past the generation it was built for -- see - the module docstring's GENERATION-SCOPED section for why that matters. + Unconditional eviction sites: the non-cacheable bypass branch, IC-LoRA's + stage boundary, ``set_enabled(False)``/``set_streaming_enabled(False)``, + and ``PipelinesHandler`` unload/swap paths (a video<->image swap needs both + the cached build's VRAM slice and, for streaming, its pinned host RAM). Safe to call even when nothing is cached (no-op) or when the patch is disabled (module-level cache is simply always empty). Raises if a cached transformer is currently in use (see :func:`_evict_locked`). @@ -218,6 +332,48 @@ def evict() -> None: _evict_locked() +def evict_for_generation_start() -> None: + """Generation-start eviction with kind-dependent scope. + + Resident (VRAM) entries are always evicted -- see the module docstring's + GENERATION-SCOPED section for the RTX 5090 repro. Streaming (pinned host + RAM) entries are session-scoped and kept, unless available system RAM has + fallen below ``_MIN_FREE_RAM_GB`` (external memory pressure -- e.g. the + user opened another model-hungry app between generations). + """ + with _lock: + if _cached_model is None: + return + if _dirty: + # A previous checkout exited abnormally -- never carry a possibly + # slot-leaked wrapper into a new generation. If a zombie checkout is + # still live (_in_use > 0), leave it alone: evicting would raise, and + # the checkout-time concurrency bypass isolates the new generation. + if _in_use == 0: + logger.warning( + "[diffusion-stage-cache] evicting dirty entry at generation start" + ) + _evict_locked() + return + if _cached_kind != "streaming": + _evict_locked() + return + available_gb = psutil.virtual_memory().available / 2**30 + if available_gb < _MIN_FREE_RAM_GB: + logger.warning( + "[diffusion-stage-cache] evicting streaming entry at generation start: " + "%.1f GB RAM available < %.1f GB threshold", + available_gb, + _MIN_FREE_RAM_GB, + ) + _evict_locked() + return + logger.info( + "[diffusion-stage-cache] keeping session-scoped streaming entry (%.1f GB RAM available)", + available_gb, + ) + + def _mark_free() -> None: """Mark that a caller is done with a cached transformer it checked out.""" global _in_use @@ -232,10 +388,11 @@ def _mark_free() -> None: def _cached_transformer_ctx(self: DiffusionStage, **kwargs: object) -> Iterator[object]: if not _enabled or not _cacheable(self): # A non-cacheable build (e.g. IC-LoRA's use_lora_in_stage_2 forcing stage_2 - # onto the streaming path -- see module docstring's NON-CACHEABLE TRANSITIONS - # section) needs the VRAM a still-resident cached transformer is holding. - # Evicting only on a cache-key mismatch (below) never fires for this case, - # since this path never touches the cache-key branch at all. + # onto the streaming path on a full-loading card -- see module docstring's + # NON-CACHEABLE TRANSITIONS section) needs the VRAM a still-resident cached + # transformer is holding. Evicting only on a cache-key mismatch (below) + # never fires for this case, since this path never touches the cache-key + # branch at all. evict() with _orig_transformer_ctx(self, **kwargs) as model: yield model @@ -244,24 +401,87 @@ def _cached_transformer_ctx(self: DiffusionStage, **kwargs: object) -> Iterator[ global _in_use key = _cache_key(self) with _lock: - if _cached_key == key and _cached_model is not None: - model = _cached_model - hit = True - else: - _evict_locked() - model = self._build_transformer(**kwargs) # noqa: SLF001 - globals()["_cached_key"], globals()["_cached_model"] = key, model - hit = False - # Mark in use BEFORE releasing the lock (not after): _evict_locked runs only - # under _lock, so bumping the counter here closes the window between resolving - # the model and marking it -- otherwise a concurrent evict could see _in_use==0 - # and free the model we're about to yield. This counter (not the lock, which is - # released before the yield) is what protects the model for the whole denoise. - _in_use += 1 - - logger.info("[diffusion-stage-cache] %s resident transformer", "reusing" if hit else "built + cached") + # Concurrency check and checkout must share ONE critical section: checking + # _in_use in a separate lock acquisition would let two threads both observe + # 0 and then serialize into a shared checkout anyway (TOCTOU). + concurrent = _in_use > 0 + if not concurrent: + if _cached_key == key and _cached_model is not None and not _dirty: + model = _cached_model + kind = _cached_kind + hit = True + else: + if _dirty and _cached_model is not None: + logger.warning( + "[diffusion-stage-cache] cached transformer is dirty (previous " + "checkout exited abnormally, may have leaked BufferPool slots) " + "-- evicting and rebuilding" + ) + elif _cached_key is not None: + logger.info( + "[diffusion-stage-cache] key mismatch (config changed) -- evicting before rebuild" + ) + _evict_locked() + if self._is_streaming: # noqa: SLF001 + # Build the streaming wrapper directly (mirror _streaming_model's + # build half, blocks.py:151). kwargs are deliberately dropped: + # upstream's _streaming_transformer_ctx ignores them too. Do NOT + # route through _build_transformer -- its .to(target) must not be + # applied to the meta-blocked streaming wrapper. + model = self._prepared_builder().build(device=self._device, dtype=self._dtype) # noqa: SLF001 + kind = "streaming" + else: + model = self._build_transformer(**kwargs) # noqa: SLF001 + kind = "resident" + globals()["_cached_key"] = key + globals()["_cached_model"] = model + globals()["_cached_kind"] = kind + hit = False + # Mark in use BEFORE releasing the lock (not after): _evict_locked runs only + # under _lock, so bumping the counter here closes the window between resolving + # the model and marking it -- otherwise a concurrent evict could see _in_use==0 + # and free the model we're about to yield. This counter (not the lock, which is + # released before the yield) is what protects the model for the whole denoise. + _in_use += 1 + + if concurrent: + # Another checkout is live (observed in the wild: a cancelled/failed + # generation's denoise thread still running as a zombie while a new + # generation starts). Sharing one wrapper between concurrent forward + # passes exhausts its 2-slot GPU BufferPool ("BufferPool exhausted: all 2 + # buffers are in use") -- pre-cache, each call had its own private + # wrapper, so concurrency was merely wasteful. Restore exactly that: + # bypass the cache with an isolated one-shot build and leave the cached + # entry alone. + logger.warning( + "[diffusion-stage-cache] checkout requested while cached transformer is " + "in use (overlapping generation?) -- bypassing cache with an isolated build" + ) + with _orig_transformer_ctx(self, **kwargs) as model: + yield model + return + + logger.info( + "[diffusion-stage-cache] %s %s transformer", + "reusing" if hit else "built + cached", + kind, + ) try: - yield model + if kind == "streaming": + # A fresh stateless X0Model wrapper per checkout, exactly as upstream's + # _streaming_transformer_ctx yields (blocks.py:402) -- the cached object + # is the BlockStreamingWrapper underneath. + yield X0Model(model).eval() # type: ignore[arg-type] + else: + yield model + except BaseException: + # Abnormal unwind (exception or cancellation) mid-denoise skips the + # wrapper's forward post-hooks, which can leak BufferPool slots. Mark the + # entry dirty so it is rebuilt instead of reused -- reusing a leaked-slot + # wrapper fails every subsequent generation until eviction. + with _lock: + globals()["_dirty"] = True + raise finally: _mark_free() @@ -270,9 +490,11 @@ def _cached_transformer_ctx(self: DiffusionStage, **kwargs: object) -> Iterator[ if __name__ == "__main__": - a = ("ckpt.safetensors", None, (), (), 1, None, "bf16", "cuda:0") - b = ("ckpt.safetensors", None, (), (), 1, None, "bf16", "cuda:0") - c = ("ckpt.safetensors", None, (), (), 2, None, "bf16", "cuda:0") + a = ("ckpt.safetensors", None, (), (), "SingleGPUModelBuilder", None, 1, None, "bf16", "cuda:0") + b = ("ckpt.safetensors", None, (), (), "SingleGPUModelBuilder", None, 1, None, "bf16", "cuda:0") + c = ("ckpt.safetensors", None, (), (), "SingleGPUModelBuilder", None, 2, None, "bf16", "cuda:0") + d = ("ckpt.safetensors", None, (), (), "StreamingModelBuilder", None, 1, None, "bf16", "cuda:0") assert a == b, "identical content must compare equal (cache hit path)" assert a != c, "different quantization identity must compare unequal (cache miss path)" + assert a != d, "resident and streaming builds of identical content must never share a key" print("diffusion_stage_cache: key-equality self-check OK") diff --git a/backend/services/patches/diffvae_decode_vram.py b/backend/services/patches/diffvae_decode_vram.py index beaa974e8..2f5556c7f 100644 --- a/backend/services/patches/diffvae_decode_vram.py +++ b/backend/services/patches/diffvae_decode_vram.py @@ -120,9 +120,28 @@ def _with_post_evict_cuda_tiling( return _replace_tiling_arg(args, kwargs, tiling) +def _decoder_checkpoint(decoder: Any) -> str | None: + checkpoint = getattr(decoder, "checkpoint_path", None) + if not isinstance(checkpoint, str): + checkpoint = getattr(decoder, "_checkpoint_path", None) + return checkpoint if isinstance(checkpoint, str) else None + + +def _is_diffusion_video_decoder(decoder: Any) -> bool: + checkpoint = _decoder_checkpoint(decoder) + return checkpoint is not None and is_diffusion_video_vae(checkpoint) + + def _patched_video_decoder_call(self: VideoDecoder, *args: Any, **kwargs: Any) -> Any: - _release_denoise_weights() - args, kwargs = _with_post_evict_cuda_tiling(self, args, kwargs) + # Only the DiffVAE (diffusion) decoder needs the resident transformer freed for its + # ~20GB tiled decode. The conv VAE ("fast decode", ~1.5GB) fits alongside the + # streaming transformer, so keep it resident for conv decode -- that is what lets + # the session-scoped diffusion_stage_cache survive across generations on the + # streaming path (the cross-gen win). Freeing it unconditionally would evict the + # cached transformer after every gen, forcing a stage_1 rebuild on the next one. + if _is_diffusion_video_decoder(self): + _release_denoise_weights() + args, kwargs = _with_post_evict_cuda_tiling(self, args, kwargs) return _orig_video_decoder_call(self, *args, **kwargs) diff --git a/backend/services/patches/fp8_sidecar_cache.py b/backend/services/patches/fp8_sidecar_cache.py new file mode 100644 index 000000000..4a6968014 --- /dev/null +++ b/backend/services/patches/fp8_sidecar_cache.py @@ -0,0 +1,277 @@ +"""Monkey-patch (EXPERIMENTAL): fp8 sidecar disk cache for the streaming +transformer's cold build. + +The once-per-session cold build of the 22B streaming transformer reads the +~43GB bf16 checkpoint and fp8-casts every covered tensor during load +(``StreamingModelBuilder._build_pinned_source`` -> ONE ``load_state_dict`` +call at block_streaming/builder.py:327 with the chained +rename+FP8_CAST_PREQUANT_AWARE sd_ops). Measured on the RTX 3090 (E: SATA +SSD): ~90-95s per cold build. This patch persists the POST-downcast block +weights to a ~22GB sidecar safetensors file beside the checkpoint +(``.fp8-blocks-cache.safetensors`` -- same next-to-the-model placement +as Krea-2's ``_nf4_transformer_cache``); later cold builds identity-load the +sidecar (~23GB read, zero cast compute, mmap via safetensors_loader_fix): +target ~40s. The first-ever build pays a one-time synchronous write (~30-45s, +clearly logged) -- a background writer would pin ~23GB of tensor refs alive +alongside the 23GB pinned copy on a 64GB box, the exact RAM regime the fork's +free-not-park divergence (FORK.md B2) exists to avoid; synchronous streaming +from the mmap-backed StateDict costs near-zero extra RAM. + +SEAM: ``StreamingModelBuilder._build_pinned_source`` is patched (not +``__init__``, not the loader class): it fires only on the CUDA pinned path +(MPS DISK streaming uses ``_build_disk_source`` and never enters), and it +sits beneath BOTH diffusion_stage_cache paths (the cache's miss-build and the +cache-off ``_streaming_model`` path both funnel into ``build()``). The patch +clones the builder (``copy.copy`` -- the builder's own ``with_*`` methods use +the same cloning) and swaps in a wrapping ``StateDictLoader`` that intercepts +only the ``...__blocks``-suffixed load; everything else (config metadata, key +scan, non-block weights, LoRA files) still reads the ORIGINAL checkpoint. + +SIDECAR CONTENT CONTRACT: +- Keys are POST-rename (``transformer_blocks.N....``) because the tee capture + happens on the loader's OUTPUT; the sidecar is therefore identity-loaded + (``sd_ops=None``), never passed back through the rename/downcast chain. +- LoRA-FREE by construction: the capture point is the load_state_dict return + value, BEFORE ``fuse_lora_weights`` runs -- LoRA fusion re-runs on every + build against the sidecar tensors (cheap fp8 delta fuse), so LoRA changes + never invalidate the sidecar. +- NEVER contains ``*_scale`` keys (asserted at write; the bf16 source has + none, and the prequant scale-fold is the one non-idempotent sd_op). + +INVALIDATION (all header-only reads -- 8-byte length + JSON header, the same +technique as safetensors_metadata_fix; that module is NOT imported because +its import has heavy side effects): stamp in the sidecar's ``__metadata__`` +must match the original checkpoint's resolved path + size + mtime_ns + +``encrypted_wandb_properties``, the full sd_ops chain name, and the fp8 +downcast suffix list (imported from ltx_core.quantization.fp8_cast -- an +upstream change to the cast set invalidates automatically). Additionally the +sidecar's tensor key set must EQUAL the load's ``allowed_keys`` (catches +config/layer-count drift and truncated files). Any validation or read +failure: warn, delete the sidecar, fall back to the original checkpoint and +rewrite. Any write failure (including disk-full): warn and continue -- the +generation is never at risk. Writes are atomic (tmp file + ``os.replace``) +with a free-space preflight, fixing the non-atomicity gap in the Krea-2 +pattern this is modeled on. + +Kill switch: env ``FP8_SIDECAR_CACHE=0``. Eligibility is structural +otherwise: single-file checkpoint path, fp8 policy present in the sd_ops +chain name (excludes Gemma's pinned streaming loads in local-encoding mode +and unquantized bf16 builds, where a sidecar would be a same-size copy), +CUDA available. + +EXPERIMENTAL: depends on ``StreamingModelBuilder._build_pinned_source``'s +signature and its single ``load_state_dict(self.model_path, self.model_loader, +...)`` call, ``_filtered_sd_ops``'s ``__blocks`` name suffix, the private +``_model_loader`` attribute, and ``SDOps.name``/``allowed_keys`` -- re-verify +against ltx_core.block_streaming.builder on rev bumps. + +Usage: + import services.patches.fp8_sidecar_cache # noqa: F401 +""" + +from __future__ import annotations + +import json +import logging +import os +import shutil +import struct +import time +from copy import copy +from pathlib import Path + +import torch + +from ltx_core.block_streaming.builder import StreamingModelBuilder +from ltx_core.loader.primitives import StateDict, StateDictLoader +from ltx_core.loader.sd_ops import SDOps +from ltx_core.quantization.fp8_cast import _FP8_CAST_LINEAR_SUFFIXES # noqa: PLC2701 + +logger = logging.getLogger(__name__) + +_SIDECAR_SUFFIX = ".fp8-blocks-cache.safetensors" +_SIDECAR_FORMAT = "1" +_ENABLED = os.environ.get("FP8_SIDECAR_CACHE", "1") != "0" + + +def _sidecar_path(checkpoint: str) -> Path: + p = Path(checkpoint) + return p.with_name(p.stem + _SIDECAR_SUFFIX) + + +def _read_header(path: Path) -> tuple[dict[str, str], set[str]]: + """Header-only safetensors read: (__metadata__, tensor key set). + + Same 8-byte-length + JSON-header technique as safetensors_metadata_fix's + ``_read_safetensors_metadata`` (duplicated on purpose -- importing that + patch module pulls in heavy side effects). + """ + with open(path, "rb") as f: + (header_len,) = struct.unpack(" int: + """Module-level for test patchability.""" + return shutil.disk_usage(directory).free + + +def _stamp(checkpoint: str, sd_ops_name: str) -> dict[str, str]: + resolved = Path(checkpoint).resolve() + st = os.stat(resolved) + try: + original_meta, _ = _read_header(resolved) + model_id = original_meta.get("encrypted_wandb_properties", "") + except Exception: + model_id = "" + return { + "sidecar_format": _SIDECAR_FORMAT, + "original_path": str(resolved), + "original_size": str(st.st_size), + "original_mtime_ns": str(st.st_mtime_ns), + "original_model_id": model_id, + "sd_ops_chain": sd_ops_name, + "fp8_downcast_suffixes": json.dumps(list(_FP8_CAST_LINEAR_SUFFIXES)), + } + + +def _valid(sidecar: Path, checkpoint: str, sd_ops: SDOps) -> bool: + """Validate the sidecar against the original + current transform identity. + + Any failure deletes the sidecar (Krea-2 corrupted-cache precedent) so the + fallback build rewrites a fresh one. + """ + try: + if not sidecar.exists(): + return False + metadata, keys = _read_header(sidecar) + expected = _stamp(checkpoint, sd_ops.name) + for field, value in expected.items(): + if metadata.get(field) != value: + logger.warning( + "[fp8-sidecar] %s: stamp mismatch on %r (checkpoint changed or " + "transform drifted) -- rebuilding from the original", + sidecar.name, + field, + ) + raise ValueError(field) + if sd_ops.allowed_keys is None or keys != set(sd_ops.allowed_keys): + logger.warning( + "[fp8-sidecar] %s: tensor key set does not match the expected block " + "partition -- rebuilding from the original", + sidecar.name, + ) + raise ValueError("key set") + if any(k.endswith("_scale") for k in keys): + raise ValueError("unexpected *_scale keys") + return True + except Exception: + try: + sidecar.unlink(missing_ok=True) + logger.warning("[fp8-sidecar] deleted invalid sidecar %s", sidecar) + except OSError: + logger.warning("[fp8-sidecar] could not delete invalid sidecar %s", sidecar, exc_info=True) + return False + + +def _write_sidecar(sidecar: Path, checkpoint: str, sd_ops: SDOps, result: StateDict) -> None: + """Best-effort atomic write; never raises out (the generation must proceed).""" + tmp = sidecar.with_name(sidecar.name + f".tmp-{os.getpid()}") + try: + from safetensors.torch import save_file + + tensors = { + key: (value if value.is_contiguous() else value.contiguous()) + for key, value in result.sd.items() + if isinstance(value, torch.Tensor) + } + if len(tensors) != len(result.sd): + logger.warning("[fp8-sidecar] state dict contains non-tensor entries; skipping write") + return + if any(key.endswith("_scale") for key in tensors): + logger.warning("[fp8-sidecar] refusing to write *_scale keys (non-idempotent transform)") + return + estimated = sum(t.numel() * t.element_size() for t in tensors.values()) + if _free_bytes(sidecar.parent) < estimated * 1.1: + logger.warning( + "[fp8-sidecar] skipping write: less than %.1f GB free next to the checkpoint", + estimated * 1.1 / 2**30, + ) + return + + start = time.time() + logger.info( + "[fp8-sidecar] writing %s (%.1f GB, one-time per checkpoint) ...", + sidecar.name, + estimated / 2**30, + ) + save_file(tensors, str(tmp), metadata=_stamp(checkpoint, sd_ops.name)) + os.replace(tmp, sidecar) + logger.info("[fp8-sidecar] wrote %s (%.1f GB) in %.1fs", sidecar.name, estimated / 2**30, time.time() - start) + except Exception: + logger.warning("[fp8-sidecar] write failed; continuing without a sidecar", exc_info=True) + try: + tmp.unlink(missing_ok=True) + except OSError: + pass + + +class _SidecarLoader: + """StateDictLoader wrapper: serve/capture the ``__blocks`` load via the sidecar.""" + + def __init__(self, inner: StateDictLoader, checkpoint: str) -> None: + self._inner = inner + self._checkpoint = checkpoint + + def metadata(self, path: str) -> dict: + return self._inner.metadata(path) + + def load( + self, + path: str | list[str], + sd_ops: SDOps | None = None, + device: torch.device | None = None, + ) -> StateDict: + if sd_ops is None or not sd_ops.name.endswith("__blocks") or sd_ops.allowed_keys is None: + return self._inner.load(path, sd_ops=sd_ops, device=device) + + sidecar = _sidecar_path(self._checkpoint) + if _valid(sidecar, self._checkpoint, sd_ops): + logger.info( + "[fp8-sidecar] loading pre-downcast block weights from %s (skips the " + "43GB bf16 read + fp8 cast)", + sidecar.name, + ) + # Identity load: keys are already post-rename, dtypes post-downcast. + return self._inner.load(str(sidecar), sd_ops=None, device=device) + + result = self._inner.load(path, sd_ops=sd_ops, device=device) + _write_sidecar(sidecar, self._checkpoint, sd_ops, result) + return result + + +def _eligible(builder: StreamingModelBuilder) -> bool: + return ( + _ENABLED + and isinstance(builder.model_path, str) + and builder.model_sd_ops is not None + and "FP8_CAST_PREQUANT_AWARE" in builder.model_sd_ops.name + and torch.cuda.is_available() + ) + + +_orig_build_pinned_source = StreamingModelBuilder._build_pinned_source # noqa: SLF001 + + +def _patched_build_pinned_source(self: StreamingModelBuilder, *args: object, **kwargs: object) -> object: + if not _eligible(self): + return _orig_build_pinned_source(self, *args, **kwargs) # type: ignore[arg-type] + # copy.copy mirrors the builder's own with_* cloning; only the loader differs. + clone = copy(self) + clone._model_loader = _SidecarLoader(self.model_loader, str(self.model_path)) # noqa: SLF001 + return _orig_build_pinned_source(clone, *args, **kwargs) # type: ignore[arg-type] + + +StreamingModelBuilder._build_pinned_source = _patched_build_pinned_source # type: ignore[method-assign] # noqa: SLF001 diff --git a/backend/services/patches/pinned_pool_fix.py b/backend/services/patches/pinned_pool_fix.py index 4d0f42e7c..237db812d 100644 --- a/backend/services/patches/pinned_pool_fix.py +++ b/backend/services/patches/pinned_pool_fix.py @@ -24,6 +24,13 @@ Remove once ltx-core ``alloc_buffer`` does not poison the CUDA context or mis-report pinned-host failure as VRAM OOM. +NOTE: the upstream MPS-support work rewrote the weight-streaming subsystem +(`ltx_core.layer_streaming` -> `ltx_core.block_streaming`). When that module is +absent this patch cleanly no-ops — it only ever applied to the old _LayerStore, +and it is irrelevant on MPS (DISK streaming). Whether the new block_streaming +pins all blocks upfront (the waste this patch fixed) must be re-evaluated on CUDA +before shipping there. + Usage: import services.patches.pinned_pool_fix # noqa: F401 """ diff --git a/backend/services/qwen_multiangle_pipeline/__init__.py b/backend/services/qwen_multiangle_pipeline/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/backend/services/qwen_multiangle_pipeline/angle_mapping.py b/backend/services/qwen_multiangle_pipeline/angle_mapping.py new file mode 100644 index 000000000..5f6883a7c --- /dev/null +++ b/backend/services/qwen_multiangle_pipeline/angle_mapping.py @@ -0,0 +1,178 @@ +"""Pure mapping from continuous camera-gizmo values to the LoRA's discrete pose vocabulary. + +The fal/Qwen-Image-Edit-2511-Multiple-Angles-LoRA is trained on exactly +96 named poses: 8 azimuths x 4 elevations x 3 distances. Prompts must use +the trigger format: + + {azimuth} {elevation} {distance} + +e.g. " back-right quarter view eye-level shot close-up" + +This module is the single source of truth for that vocabulary. The frontend +fetches the table via /api/mapping and never hardcodes it. + +Conventions: +- azimuth_deg: degrees clockwise from subject-front, viewed from above. + 0 = camera in front of subject, 90 = camera to subject's right side, + 180 = behind, 270 = left side. Any real value accepted (wraps mod 360). +- elevation_deg: camera height angle. Negative = below eye level looking up + (low-angle), positive = above looking down. Clamped to [-30, 60]. +- zoom: subject-distance multiplier. 0.6 = close-up, 1.0 = medium, 1.8 = wide. + Snapped in log space so the perceptual midpoints fall correctly. +""" + +from __future__ import annotations + +import math +from collections.abc import Sequence +from dataclasses import dataclass + +TRIGGER = "" + +# (bucket center in degrees, prompt phrase, camera-position label). +# Degrees are viewer-relative screen space: 45 = camera orbits to the +# viewer's right. The LoRA's phrases name the CAMERA'S position in the same +# viewer-relative sense — verified empirically 2026-07-10 with two eye-level +# renders: "left side view" and "front-left quarter view" both put the +# camera at the VIEWER'S left (subject's right cheek visible). So degrees +# and phrases align 1:1 with no mirroring. Beware the perceptual trap that +# caused a false bug report and a wrong "fix": in a quarter view the camera +# position and the subject's apparent facing point OPPOSITE ways — a camera +# at viewer-right shows the subject turned toward frame-left. Labels name +# the camera position; judge renders by where the camera moved, not which +# way the subject faces. +AZIMUTH_BUCKETS: tuple[tuple[float, str, str], ...] = ( + (0.0, "front view", "front"), + (45.0, "front-right quarter view", "front-right"), + (90.0, "right side view", "right side"), + (135.0, "back-right quarter view", "back-right"), + (180.0, "back view", "back"), + (225.0, "back-left quarter view", "back-left"), + (270.0, "left side view", "left side"), + (315.0, "front-left quarter view", "front-left"), +) + +# (bucket center in degrees, prompt phrase) +ELEVATION_BUCKETS: tuple[tuple[float, str], ...] = ( + (-30.0, "low-angle shot"), + (0.0, "eye-level shot"), + (30.0, "elevated shot"), + (60.0, "high-angle shot"), +) + +# (distance multiplier, prompt phrase) +DISTANCE_BUCKETS: tuple[tuple[float, str], ...] = ( + (0.6, "close-up"), + (1.0, "medium shot"), + (1.8, "wide shot"), +) + + +@dataclass(frozen=True, slots=True) +class Pose: + """A snapped camera pose: indices into the bucket tables plus the prompt.""" + + azimuth_index: int + elevation_index: int + distance_index: int + + @property + def azimuth_deg(self) -> float: + return AZIMUTH_BUCKETS[self.azimuth_index][0] + + @property + def elevation_deg(self) -> float: + return ELEVATION_BUCKETS[self.elevation_index][0] + + @property + def zoom(self) -> float: + return DISTANCE_BUCKETS[self.distance_index][0] + + @property + def prompt(self) -> str: + return ( + f"{TRIGGER} {AZIMUTH_BUCKETS[self.azimuth_index][1]}" + f" {ELEVATION_BUCKETS[self.elevation_index][1]}" + f" {DISTANCE_BUCKETS[self.distance_index][1]}" + ) + + +def snap_azimuth(azimuth_deg: float) -> int: + """Nearest 45°-spaced bucket, wrapping — e.g. 350° snaps to front (0°).""" + return round((azimuth_deg % 360.0) / 45.0) % len(AZIMUTH_BUCKETS) + + +def snap_elevation(elevation_deg: float) -> int: + return min( + range(len(ELEVATION_BUCKETS)), + key=lambda i: abs(ELEVATION_BUCKETS[i][0] - elevation_deg), + ) + + +def snap_distance(zoom: float) -> int: + """Snap in log space: the boundary between two zooms is their geometric mean.""" + z = math.log(max(zoom, 1e-6)) + return min( + range(len(DISTANCE_BUCKETS)), + key=lambda i: abs(math.log(DISTANCE_BUCKETS[i][0]) - z), + ) + + +def snap_pose(azimuth_deg: float, elevation_deg: float, zoom: float) -> Pose: + return Pose( + azimuth_index=snap_azimuth(azimuth_deg), + elevation_index=snap_elevation(elevation_deg), + distance_index=snap_distance(zoom), + ) + + +def compose_prompt(pose: Pose, extra_prompt: str = "", extra_roles: Sequence[str] = ()) -> str: + """The full generation prompt: optional compositing instruction, then the + pose trigger, then optional freeform style text. + + `extra_roles` aligns with the extra reference images the pipeline receives, + in the same order (the subject is image 1, so extra ref i is image i + 2). + A "location" ref makes the model composite the subject INTO that scene; each + "prop" ref is added as an object present with the subject (placement left to + the user's freeform text, so it doesn't fight e.g. "slung over her shoulder"). + Empty roles = the + classic pose-trigger-only behavior (no compositing clause), so callers that + pass nothing are unchanged. + + The compositing clause leads and the `` pose trigger follows; if the + angles LoRA ever responds better with the trigger first, this is the single + place to reorder it (both the live generation and the displayed prompt come + through here). + """ + location_imgs = [i + 2 for i, role in enumerate(extra_roles) if role == "location"] + prop_imgs = [i + 2 for i, role in enumerate(extra_roles) if role == "prop"] + + parts: list[str] = [] + if location_imgs: + parts.append( + f"Place the subject from image 1 into the scene shown in image {location_imgs[0]}, " + "keeping that location's architecture, lighting and background" + ) + if prop_imgs: + which = " and ".join(f"image {n}" for n in prop_imgs) + noun = "objects" if len(prop_imgs) > 1 else "object" + as_prop = "props" if len(prop_imgs) > 1 else "a prop" + parts.append(f"add the {noun} from {which} as {as_prop} with the subject") + + base = f"{'. '.join(parts)}. {pose.prompt}" if parts else pose.prompt + extra = extra_prompt.strip() + return f"{base}, {extra}" if extra else base + + +def mapping_table() -> dict[str, object]: + """The full vocabulary in one JSON-ready dict, for /api/mapping.""" + return { + "trigger": TRIGGER, + "prompt_format": f"{TRIGGER} {{azimuth}} {{elevation}} {{distance}}", + "azimuths": [ + {"deg": deg, "phrase": phrase, "label": label} + for deg, phrase, label in AZIMUTH_BUCKETS + ], + "elevations": [{"deg": deg, "phrase": phrase} for deg, phrase in ELEVATION_BUCKETS], + "distances": [{"zoom": z, "phrase": phrase} for z, phrase in DISTANCE_BUCKETS], + } diff --git a/backend/services/qwen_multiangle_pipeline/gguf_qwen_multiangle_pipeline.py b/backend/services/qwen_multiangle_pipeline/gguf_qwen_multiangle_pipeline.py new file mode 100644 index 000000000..198531949 --- /dev/null +++ b/backend/services/qwen_multiangle_pipeline/gguf_qwen_multiangle_pipeline.py @@ -0,0 +1,214 @@ +"""GGUF-quantized Qwen multi-angle pipeline. + +Ported from the standalone qwen-multiangle-studio app's server.py, GGUF path +only — bf16 and NF4 paths are deliberately not ported. NF4 was proven to +corrupt this model's texture reconstruction (blockwise quantization noise in +the transformer's attention over condition latents); bf16 works but takes +~40 minutes per generation on this hardware. GGUF K-quant is the only +validated-good tradeoff (~30-130s per generation, quality matching bf16) — +see the qwen-multiangle-integration-plan project notes for the full +debugging history if this ever needs revisiting. + +Loading takes several minutes cold (reads the ~54GB bf16 originals once to +locate/cache the GGUF file and text encoder); subsequent loads from a warm +HF cache on fast storage take ~30s. See PipelinesHandler.load_qwen_multiangle_pipeline +for why this pipeline is evicted-and-reloaded rather than parked on CPU. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING + +from services.qwen_multiangle_pipeline.angle_mapping import compose_prompt, snap_pose +from services.qwen_multiangle_pipeline.qwen_multiangle_pipeline import QwenMultiAnglePipeline + +if TYPE_CHECKING: + import torch + from collections.abc import Callable, Sequence + from PIL.Image import Image as PILImage + +logger = logging.getLogger(__name__) + +BASE_MODEL = "Qwen/Qwen-Image-Edit-2511" +GGUF_REPO = "unsloth/Qwen-Image-Edit-2511-GGUF" +GGUF_QUANT = "Q6_K" +ANGLES_LORA = "fal/Qwen-Image-Edit-2511-Multiple-Angles-LoRA" +LIGHTNING_LORA = "lightx2v/Qwen-Image-Edit-2511-Lightning" +LIGHTNING_4STEP_WEIGHT = "Qwen-Image-Edit-2511-Lightning-4steps-V1.0-bf16.safetensors" +LIGHTNING_8STEP_WEIGHT = "Qwen-Image-Edit-2511-Lightning-8steps-V1.0-bf16.safetensors" + +# Optional skin-realism adapter (prithivMLmods, native 2511 — pore-level skin +# texture). Header-verified same rank 16 / dim 3072 as the angles LoRA, so it +# stacks cleanly; the only difference is a PEFT ".default." infix and a missing +# "transformer." prefix in its keys, normalized at load time. Off by default, +# weighted at generate time so it never overpowers the angles LoRA. +SKIN_LORA_REPO = "prithivMLmods/Qwen-Image-Edit-2511-Hyper-Realistic-Portrait" +SKIN_LORA_WEIGHT = "HRP_5.safetensors" + +# Three sampling recipes selectable per generation (quality_mode). Both Lightning +# LoRAs are distillations at CFG 1.0; the 8-step ("balanced") keeps noticeably +# more high-frequency skin/texture detail than the 4-step ("fast") while staying +# ~2x faster than the un-distilled 28-step base ("quality"). +FAST_STEPS, FAST_CFG = 4, 1.0 +BALANCED_STEPS, BALANCED_CFG = 8, 1.0 +QUALITY_STEPS, QUALITY_CFG = 28, 4.0 + + +class GGUFQwenMultiAnglePipeline: + @staticmethod + def create(device: "torch.device") -> QwenMultiAnglePipeline: + return GGUFQwenMultiAnglePipeline(device=device) + + def __init__(self, device: "torch.device") -> None: + import torch + from diffusers import DiffusionPipeline # type: ignore[reportPrivateImportUsage] + from diffusers import GGUFQuantizationConfig # type: ignore[reportPrivateImportUsage] + from diffusers import QwenImageTransformer2DModel # type: ignore[reportPrivateImportUsage] + from huggingface_hub import hf_hub_download # type: ignore[reportUnknownVariableType] + from transformers import Qwen2_5_VLForConditionalGeneration + + self._device = device + + logger.info("Loading Qwen multi-angle GGUF transformer (%s)…", GGUF_QUANT) + gguf_path = hf_hub_download( + GGUF_REPO, + f"qwen-image-edit-2511-{GGUF_QUANT}.gguf", + local_files_only=True, + ) + transformer = QwenImageTransformer2DModel.from_single_file( # type: ignore[reportUnknownMemberType] + gguf_path, + quantization_config=GGUFQuantizationConfig(compute_dtype=torch.bfloat16), + torch_dtype=torch.bfloat16, + config=BASE_MODEL, + subfolder="transformer", + ) + + logger.info("Loading Qwen multi-angle bf16 text encoder…") + text_encoder = Qwen2_5_VLForConditionalGeneration.from_pretrained( # type: ignore[reportUnknownMemberType] + BASE_MODEL, + subfolder="text_encoder", + torch_dtype=torch.bfloat16, + local_files_only=True, + ) + + logger.info("Assembling Qwen multi-angle pipeline…") + pipe = DiffusionPipeline.from_pretrained( # type: ignore[reportUnknownMemberType] + BASE_MODEL, + transformer=transformer, + text_encoder=text_encoder, + torch_dtype=torch.bfloat16, + local_files_only=True, + ) + + logger.info("Loading Qwen multi-angle LoRAs…") + pipe.load_lora_weights(ANGLES_LORA, adapter_name="angles") + pipe.load_lora_weights(LIGHTNING_LORA, weight_name=LIGHTNING_4STEP_WEIGHT, adapter_name="lightning4") + pipe.load_lora_weights(LIGHTNING_LORA, weight_name=LIGHTNING_8STEP_WEIGHT, adapter_name="lightning8") + + # Skin-realism adapter: normalize its PEFT-style keys + # (transformer_blocks.N…lora_A.default.weight) to the diffusers + # convention the other adapters use (transformer.transformer_blocks.N… + # lora_A.weight) so load_lora_weights maps every key. + logger.info("Loading Qwen multi-angle skin-realism LoRA…") + from safetensors.torch import load_file # type: ignore[reportUnknownVariableType] # local: heavy import + + skin_path = hf_hub_download(SKIN_LORA_REPO, SKIN_LORA_WEIGHT) + skin_raw: dict[str, object] = load_file(skin_path) # type: ignore[reportUnknownMemberType] + skin_state: dict[str, object] = {} + for key, tensor in skin_raw.items(): + norm = key.replace(".default.", ".") + if not norm.startswith("transformer."): + norm = f"transformer.{norm}" + skin_state[norm] = tensor + pipe.load_lora_weights(skin_state, adapter_name="skin") + + # Load-bearing, not an optimization: the GGUF transformer (~17GB) and + # bf16 text encoder (~16GB) sum to more than 24GB VRAM. These hooks + # are what keep the two from ever being GPU-resident simultaneously + # (sequential per-forward-pass-stage swap). Never call `.to(device)` + # on this pipeline once these hooks are installed — see the protocol + # module's docstring for what breaks if something does. + pipe.enable_model_cpu_offload() + + self._pipe = pipe + + def generate( + self, + *, + image: "PILImage", + extra_images: "list[PILImage] | None" = None, + extra_roles: "Sequence[str] | None" = None, + azimuth_deg: float, + elevation_deg: float, + zoom: float, + seed: int, + extra_prompt: str = "", + quality_mode: str = "fast", + use_skin: bool = False, + skin_weight: float = 1.0, + on_step: "Callable[[int, int], None] | None" = None, + ) -> "PILImage": + import torch + import torch.nn.functional as F + + pose = snap_pose(azimuth_deg, elevation_deg, zoom) + prompt = compose_prompt(pose, extra_prompt, extra_roles or []) + + if quality_mode == "quality": + adapters, weights = ["angles"], [1.0] + steps, cfg = QUALITY_STEPS, QUALITY_CFG + elif quality_mode == "balanced": + adapters, weights = ["angles", "lightning8"], [1.0, 1.0] + steps, cfg = BALANCED_STEPS, BALANCED_CFG + else: # "fast" (default) + adapters, weights = ["angles", "lightning4"], [1.0, 1.0] + steps, cfg = FAST_STEPS, FAST_CFG + + # Skin-realism stacks on top of whatever mode is active (most useful on + # the softer Lightning modes). Weighted so it can't overpower angles. + if use_skin and skin_weight > 0: + adapters.append("skin") + weights.append(skin_weight) + self._pipe.set_adapters(adapters, adapter_weights=weights) + + generator = torch.Generator(device="cpu").manual_seed(seed) + + def _on_step_end(pipeline: object, step: int, timestep: object, kwargs: dict[str, object]) -> dict[str, object]: + del pipeline, timestep + if on_step is not None: + on_step(step + 1, steps) + return kwargs + + # ltx2_server.py globally monkeypatches F.scaled_dot_product_attention + # to route eligible shapes through SageAttention, tuned for LTX-2's + # video attention. Confirmed via a live end-to-end smoke test that it + # silently produces NaN output for this model (all-black generations) + # — force PyTorch's native SDPA for this call regardless of what's + # currently patched, then restore it so other pipelines in the same + # process (LTX video) keep their SageAttention speedup. + # Plus pipeline (2511) accepts a list; subject first, then the extra + # references (prop/location) it composes from. + call_kwargs: dict[str, object] = { + "image": [image, *extra_images] if extra_images else image, + "prompt": prompt, + "true_cfg_scale": cfg, + "num_inference_steps": steps, + "generator": generator, + "callback_on_step_end": _on_step_end, + } + # A negative prompt only does anything when CFG is on (true_cfg_scale > 1); + # the Lightning modes run at cfg 1.0, where diffusers ignores it and warns. + # Only pass it when it will actually be used (Quality mode). + if cfg > 1: + call_kwargs["negative_prompt"] = " " + + patched_sdpa = F.scaled_dot_product_attention + F.scaled_dot_product_attention = torch._C._nn.scaled_dot_product_attention # type: ignore[assignment] + try: + result = self._pipe(**call_kwargs) # type: ignore[reportCallIssue] + return result.images[0] # type: ignore[reportUnknownMemberType] + finally: + F.scaled_dot_product_attention = patched_sdpa + torch.cuda.synchronize() + torch.cuda.empty_cache() diff --git a/backend/services/qwen_multiangle_pipeline/qwen_multiangle_pipeline.py b/backend/services/qwen_multiangle_pipeline/qwen_multiangle_pipeline.py new file mode 100644 index 000000000..1176be642 --- /dev/null +++ b/backend/services/qwen_multiangle_pipeline/qwen_multiangle_pipeline.py @@ -0,0 +1,46 @@ +"""Qwen multi-angle pipeline protocol definition. + +Deliberately NOT the `ImageGenerationPipeline` protocol (services/interfaces.py) — +that protocol's park-on-CPU-eviction contract (a plain `.to(device)` call) +fights this model's own `enable_model_cpu_offload()` accelerate hooks, which +are load-bearing here (the ~17GB GGUF transformer + ~16GB bf16 text encoder +sum to more than the 24GB card; the hooks are what keeps them from ever being +GPU-resident simultaneously). This pipeline is evicted-and-reloaded like the +video-family pipelines instead — see PipelinesHandler.load_qwen_multiangle_pipeline. + +Also deliberately not `@runtime_checkable`, matching RetakePipeline — nothing +should ever isinstance-check against this; the state wrapper dataclass +(QwenMultiAngleState) is what pipelines_handler.py's eviction logic inspects, +never this protocol or a raw pipeline instance. +""" + +from __future__ import annotations + +from collections.abc import Callable, Sequence +from typing import TYPE_CHECKING, Protocol + +if TYPE_CHECKING: + import torch + from PIL.Image import Image as PILImage + + +class QwenMultiAnglePipeline(Protocol): + @staticmethod + def create(device: "torch.device") -> "QwenMultiAnglePipeline": ... + + def generate( + self, + *, + image: "PILImage", + extra_images: "list[PILImage] | None" = None, + extra_roles: "Sequence[str] | None" = None, + azimuth_deg: float, + elevation_deg: float, + zoom: float, + seed: int, + extra_prompt: str = "", + quality_mode: str = "fast", + use_skin: bool = False, + skin_weight: float = 1.0, + on_step: Callable[[int, int], None] | None = None, + ) -> "PILImage": ... diff --git a/backend/services/text_encoder/ltx_text_encoder.py b/backend/services/text_encoder/ltx_text_encoder.py index 805009a8b..14023f36d 100644 --- a/backend/services/text_encoder/ltx_text_encoder.py +++ b/backend/services/text_encoder/ltx_text_encoder.py @@ -35,6 +35,24 @@ ) + +def _is_sequence(value: object) -> TypeGuard[Sequence[object]]: + return isinstance(value, (list, tuple)) + + +def _first_embedding_tensor(conditioning: object) -> torch.Tensor: + """Pull ``conditioning[0][0]`` only if the nest is sequences of a tensor.""" + if not _is_sequence(conditioning) or len(conditioning) == 0: + raise pickle.UnpicklingError("unexpected conditioning container") + first = conditioning[0] + if not _is_sequence(first) or len(first) == 0: + raise pickle.UnpicklingError("unexpected conditioning row") + embeddings = first[0] + if not isinstance(embeddings, torch.Tensor): + raise pickle.UnpicklingError("conditioning is not a tensor") + return embeddings + + class _CpuUnpickler(pickle.Unpickler): """Unpickler that maps torch storages to CPU on load. @@ -57,23 +75,6 @@ def _load_from_bytes_cpu(b: bytes) -> Any: return super().find_class(module, name) -def _is_sequence(value: object) -> TypeGuard[Sequence[object]]: - return isinstance(value, (list, tuple)) - - -def _first_embedding_tensor(conditioning: object) -> torch.Tensor: - """Pull ``conditioning[0][0]`` only if the nest is sequences of a tensor.""" - if not _is_sequence(conditioning) or len(conditioning) == 0: - raise pickle.UnpicklingError("unexpected conditioning container") - first = conditioning[0] - if not _is_sequence(first) or len(first) == 0: - raise pickle.UnpicklingError("unexpected conditioning row") - embeddings = first[0] - if not isinstance(embeddings, torch.Tensor): - raise pickle.UnpicklingError("conditioning is not a tensor") - return embeddings - - class LTXTextEncoder: """Stateless text encoding operations with idempotent monkey-patching.""" @@ -239,6 +240,42 @@ def get_model_id_from_checkpoint(self, checkpoint_path: str) -> str | None: logger.warning("Could not extract model_id from checkpoint: %s", exc, exc_info=True) return None + @staticmethod + def _log_enhanced_prompt(conditioning: object) -> None: + """Best-effort transparency for the Prompt Enhancer. + + The user never sees the rewritten prompt the model actually follows -- which + is how the 2026-07-22 "i2v morphs the source character" incident stayed + invisible for hours. The /v1/prompt-embedding response is a pickled + conditioning structure that is primarily tensors; if the API includes the + enhanced TEXT anywhere in it, surface it. Otherwise at least record + unambiguously that a server-side rewrite happened. + """ + from typing import cast + + try: + found: list[str] = [] + stack: list[object] = [conditioning] + while stack and len(found) < 4: + item = stack.pop() + if isinstance(item, str): + if item.strip(): + found.append(item.strip()) + elif isinstance(item, (list, tuple)): + stack.extend(cast("list[object] | tuple[object, ...]", item)) + elif isinstance(item, dict): + stack.extend(cast("dict[object, object]", item).values()) + if found: + logger.info("Prompt as enhanced by the LTX API: %s", " | ".join(found)[:1500]) + else: + logger.info( + "Prompt was REWRITTEN server-side by the LTX API (Prompt Enhancer on); " + "the rewritten text was not included in the response -- the embeddings " + "encode the rewritten version, not the prompt as typed." + ) + except Exception: + logger.debug("Could not inspect conditioning payload for enhanced prompt text", exc_info=True) + def encode_via_api( self, prompt: str, @@ -286,6 +323,8 @@ def encode_via_api( # Map CUDA storages to CPU during unpickling so this works on non-CUDA hosts # (Apple Silicon / CPU-only); the tensors are moved to self.device below. conditioning = _CpuUnpickler(io.BytesIO(response.content)).load() # noqa: S301 + if enhance_prompt: + self._log_enhanced_prompt(conditioning) embeddings = _first_embedding_tensor(conditioning) video_dim = 4096 if embeddings.shape[-1] > video_dim: diff --git a/backend/state/app_settings.py b/backend/state/app_settings.py index ad56126a0..d1bd9b492 100644 --- a/backend/state/app_settings.py +++ b/backend/state/app_settings.py @@ -2,7 +2,6 @@ from __future__ import annotations -import sys from typing import Any, Literal, TypeGuard, TypeVar, cast, get_args from pydantic import BaseModel, ConfigDict, Field, create_model, field_validator @@ -50,7 +49,11 @@ class SettingsPatchModel(SettingsBaseModel): class AppSettings(SettingsBaseModel): use_torch_compile: bool = False - diffusion_stage_cache_enabled: bool = False + # Default ON: the session-scoped streaming cache is the fork's 6x video-gen + # speedup (226s -> 35s warm on the 3090) and is safe on both hardware tiers; + # settings persistence is currently broken (see save_settings diagnostic), + # so a False default would silently disable it on every launch. + diffusion_stage_cache_enabled: bool = True ltx_api_key: str = "" user_prefers_ltx_api_video_generations: bool = False fal_api_key: str = "" @@ -132,7 +135,8 @@ def _is_settings_model_annotation(annotation: object) -> TypeGuard[type[Settings class SettingsResponse(SettingsBaseModel): use_torch_compile: bool = False - diffusion_stage_cache_enabled: bool = False + # Keep in sync with AppSettings.diffusion_stage_cache_enabled (default ON). + diffusion_stage_cache_enabled: bool = True has_ltx_api_key: bool = False user_prefers_ltx_api_video_generations: bool = False has_fal_api_key: bool = False @@ -152,10 +156,16 @@ class SettingsResponse(SettingsBaseModel): def resolved_use_conv_vae(settings: AppSettings) -> bool: - """Effective Fast decode setting: user override, else Mac on / CUDA off.""" + """Effective Fast decode setting: user override, else on by default. + + Conv VAE ("fast decode") is the default on every platform. On a 24GB CUDA card + the diffusion VAE decode costs ~20s/gen; conv decode collapses that to a few + seconds at slightly lower fidelity, so it's the right default for the 3090-class + target. Was Mac-only-default before (CUDA fell back to the slow DiffVAE). + """ if settings.use_conv_vae is not None: return settings.use_conv_vae - return sys.platform == "darwin" + return True def to_settings_response(settings: AppSettings) -> SettingsResponse: diff --git a/backend/state/app_state_types.py b/backend/state/app_state_types.py index e2d23bd92..31b34a161 100644 --- a/backend/state/app_state_types.py +++ b/backend/state/app_state_types.py @@ -17,6 +17,7 @@ ImageGenerationPipeline, IcLoraPipeline, PoseProcessorPipeline, + QwenMultiAnglePipeline, RetakePipeline, TextEncoder, ) @@ -175,6 +176,15 @@ class RetakePipelineState: video_vae_path: str | None = None # cache key — see VideoPipelineState.video_vae_path +@dataclass +class QwenMultiAngleState: + """Evicted-and-reloaded on GPU swap, like the video-family states above — + deliberately NOT ImageGenerationPipeline's park-on-CPU protocol. See + services/qwen_multiangle_pipeline/qwen_multiangle_pipeline.py for why.""" + + pipeline: QwenMultiAnglePipeline + + # ============================================================ # Generation state # ============================================================ @@ -236,7 +246,7 @@ class ApiGeneration: @dataclass class GpuSlot: - active_pipeline: VideoPipelineState | ICLoraState | A2VPipelineState | RetakePipelineState | ImageGenerationPipeline + active_pipeline: VideoPipelineState | ICLoraState | A2VPipelineState | RetakePipelineState | QwenMultiAngleState | ImageGenerationPipeline @dataclass diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py index c5b539223..71e0bdd5a 100644 --- a/backend/tests/conftest.py +++ b/backend/tests/conftest.py @@ -87,12 +87,16 @@ def test_state(tmp_path: Path, fake_services: FakeServices): ltx_api_client=fake_services.ltx_api_client, zit_api_client=fake_services.zit_api_client, fast_video_pipeline_class=type(fake_services.fast_video_pipeline), - image_generation_pipeline_class=type(fake_services.image_generation_pipeline), + image_generation_pipeline_classes={ + "z-image-turbo": type(fake_services.image_generation_pipeline), + "krea-2-turbo": type(fake_services.image_generation_pipeline), + }, ic_lora_pipeline_class=type(fake_services.ic_lora_pipeline), depth_processor_pipeline_class=type(fake_services.depth_processor_pipeline), pose_processor_pipeline_class=type(fake_services.pose_processor_pipeline), a2v_pipeline_class=type(fake_services.a2v_pipeline), retake_pipeline_class=type(fake_services.retake_pipeline), + qwen_multiangle_pipeline_class=type(fake_services.qwen_multiangle_pipeline), prompt_enhancer_pipeline_class=type(fake_services.prompt_enhancer_pipeline), ) diff --git a/backend/tests/fakes/services.py b/backend/tests/fakes/services.py index df00bd1d8..02176be0a 100644 --- a/backend/tests/fakes/services.py +++ b/backend/tests/fakes/services.py @@ -1091,6 +1091,36 @@ def extend(self, **kwargs: Any) -> None: output_path.write_bytes(b"fake-extend-video") +class FakeQwenMultiAnglePipeline: + _singleton: ClassVar["FakeQwenMultiAnglePipeline | None"] = None + + @classmethod + def bind_singleton(cls, pipeline: "FakeQwenMultiAnglePipeline") -> None: + cls._singleton = pipeline + + @staticmethod + def create(device: str | object) -> "FakeQwenMultiAnglePipeline": + del device + pipeline = FakeQwenMultiAnglePipeline._singleton + if pipeline is None: + raise RuntimeError("FakeQwenMultiAnglePipeline singleton is not bound") + return pipeline + + def __init__(self) -> None: + self.generate_calls: list[dict[str, Any]] = [] + self.raise_on_generate: Exception | None = None + + def generate(self, **kwargs: Any) -> Image.Image: + self.generate_calls.append(kwargs) + if self.raise_on_generate is not None: + raise self.raise_on_generate + on_step = kwargs.get("on_step") + if on_step is not None: + on_step(1, 4) + on_step(4, 4) + return Image.new("RGB", (32, 32), "green") + + class FakeTextEncoder: def __init__(self) -> None: self.install_calls = 0 @@ -1141,6 +1171,7 @@ class FakeServices: pose_processor_pipeline: FakePoseProcessorPipeline = field(default_factory=FakePoseProcessorPipeline) a2v_pipeline: FakeA2VPipeline = field(default_factory=FakeA2VPipeline) retake_pipeline: FakeRetakePipeline = field(default_factory=FakeRetakePipeline) + qwen_multiangle_pipeline: FakeQwenMultiAnglePipeline = field(default_factory=FakeQwenMultiAnglePipeline) prompt_enhancer_pipeline: FakePromptEnhancerPipeline = field(default_factory=FakePromptEnhancerPipeline) def __post_init__(self) -> None: @@ -1151,4 +1182,5 @@ def __post_init__(self) -> None: FakePoseProcessorPipeline.bind_singleton(self.pose_processor_pipeline) FakeA2VPipeline.bind_singleton(self.a2v_pipeline) FakeRetakePipeline.bind_singleton(self.retake_pipeline) + FakeQwenMultiAnglePipeline.bind_singleton(self.qwen_multiangle_pipeline) FakePromptEnhancerPipeline.bind_singleton(self.prompt_enhancer_pipeline) diff --git a/backend/tests/test_api_calls.py b/backend/tests/test_api_calls.py index 7440df433..fd8912d42 100644 --- a/backend/tests/test_api_calls.py +++ b/backend/tests/test_api_calls.py @@ -1,6 +1,7 @@ """Integration-style tests for /api/suggest-gap-prompt, /api/retake, /api/extend.""" from __future__ import annotations +import base64 import logging import uuid @@ -423,6 +424,94 @@ def test_prefers_api_video_without_key_falls_back_to_local_retake( assert len(fake_services.retake_pipeline.generate_calls) == 1 +class TestQwenMultiAngle: + def _data_url(self, make_test_image) -> str: + png_bytes = make_test_image(64, 64, "red").getvalue() + return "data:image/png;base64," + base64.b64encode(png_bytes).decode() + + def _base_payload(self, make_test_image) -> dict[str, object]: + return { + "image_data_url": self._data_url(make_test_image), + "azimuth_deg": 135.0, + "elevation_deg": 0.0, + "zoom": 0.6, + } + + def test_happy_path(self, client, test_state, make_test_image, fake_services): + r = client.post("/api/qwen-multiangle/generate", json=self._base_payload(make_test_image)) + assert r.status_code == 200 + data = r.json() + assert data["status"] == "complete" + assert data["image_data_url"].startswith("data:image/png;base64,") + assert data["prompt"] == " back-right quarter view eye-level shot close-up" + assert data["seed"] == 42 + assert len(fake_services.qwen_multiangle_pipeline.generate_calls) == 1 + + def test_extra_prompt_appended(self, client, test_state, make_test_image): + payload = self._base_payload(make_test_image) + payload["extra_prompt"] = "dramatic lighting" + r = client.post("/api/qwen-multiangle/generate", json=payload) + assert r.status_code == 200 + assert r.json()["prompt"].endswith("dramatic lighting") + + def test_compositing_instruction_from_roles(self, client, test_state, make_test_image, fake_services): + # Location + prop refs => the prompt should instruct compositing, with the + # subject as image 1, location as image 2, prop as image 3. + payload = self._base_payload(make_test_image) + payload["extra_image_data_urls"] = [self._data_url(make_test_image), self._data_url(make_test_image)] + payload["extra_image_roles"] = ["location", "prop"] + r = client.post("/api/qwen-multiangle/generate", json=payload) + assert r.status_code == 200 + prompt = r.json()["prompt"] + assert "into the scene shown in image 2" in prompt + assert "image 3" in prompt + # roles must reach the pipeline aligned with the refs + call = fake_services.qwen_multiangle_pipeline.generate_calls[-1] + assert call["extra_roles"] == ["location", "prop"] + + def test_mismatched_roles_ignored(self, client, test_state, make_test_image): + # A roles list that doesn't line up with the refs is dropped, not applied. + payload = self._base_payload(make_test_image) + payload["extra_image_data_urls"] = [self._data_url(make_test_image)] + payload["extra_image_roles"] = ["location", "prop"] + r = client.post("/api/qwen-multiangle/generate", json=payload) + assert r.status_code == 200 + assert "into the scene" not in r.json()["prompt"] + + def test_skin_params_forwarded(self, client, test_state, make_test_image, fake_services): + payload = self._base_payload(make_test_image) + payload["use_skin"] = True + payload["skin_weight"] = 0.8 + r = client.post("/api/qwen-multiangle/generate", json=payload) + assert r.status_code == 200 + call = fake_services.qwen_multiangle_pipeline.generate_calls[-1] + assert call["use_skin"] is True + assert call["skin_weight"] == 0.8 + + def test_randomize_seed_returns_int(self, client, test_state, make_test_image): + payload = self._base_payload(make_test_image) + payload["seed"] = 42 + payload["randomize_seed"] = True + r = client.post("/api/qwen-multiangle/generate", json=payload) + assert r.status_code == 200 + assert isinstance(r.json()["seed"], int) + + def test_invalid_data_url_400(self, client, test_state): + payload = { + "image_data_url": "not-a-data-url", + "azimuth_deg": 0.0, + "elevation_deg": 0.0, + "zoom": 1.0, + } + r = client.post("/api/qwen-multiangle/generate", json=payload) + assert r.status_code == 400 + + def test_pipeline_failure_surfaces_500(self, client, test_state, make_test_image, fake_services): + fake_services.qwen_multiangle_pipeline.raise_on_generate = RuntimeError("boom") + r = client.post("/api/qwen-multiangle/generate", json=self._base_payload(make_test_image)) + assert r.status_code == 500 + + class TestExtend: def _make_video(self, test_state) -> str: video_file = test_state.config.outputs_dir / f"extend_input_{uuid.uuid4().hex[:6]}.mp4" diff --git a/backend/tests/test_aux_block_cache.py b/backend/tests/test_aux_block_cache.py new file mode 100644 index 000000000..be696b52f --- /dev/null +++ b/backend/tests/test_aux_block_cache.py @@ -0,0 +1,260 @@ +"""Tests for the aux_block_cache patch. + +Mirrors test_diffusion_stage_cache.py conventions: real (cheap, side-effect- +free) SingleGPUModelBuilder instances as data holders, duck-typed fake blocks +exercising the module's internals and patched callables directly, no mock +library, no GPU. The patched class methods are exercised through fake `self` +objects duck-typing only the private surface the patch reads. +""" + +from __future__ import annotations + +import pytest +import torch + +from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder +from services.patches import aux_block_cache as abc_ + + +class _FakeModel(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.events: list[str] = [] + + def to(self, device: str) -> "_FakeModel": # type: ignore[override] + self.events.append(f"to:{device}") + return self + + def eval(self) -> "_FakeModel": + return self + + def decode_video(self, latent: object, tiling_config: object, generator: object): + yield "chunk-0" + yield "chunk-1" + + +class _FakeBuilder(SingleGPUModelBuilder): + def __init__(self, model_path: str, loras: tuple[object, ...] = ()) -> None: + super().__init__(model_class_configurator=object, model_path=model_path, loras=loras) # type: ignore[arg-type] + self.build_count = 0 + self.built: list[_FakeModel] = [] + + def build(self, *args: object, **kwargs: object) -> _FakeModel: # type: ignore[override] + self.build_count += 1 + model = _FakeModel() + self.built.append(model) + return model + + +_CUDA = torch.device("cuda:0") + + +class _FakeImageConditioner: + __call__ = abc_._cached_image_conditioner_call + + def __init__(self, builder: _FakeBuilder, device: torch.device = _CUDA) -> None: + self._encoder_builder = builder + self._dtype = torch.bfloat16 + self._device = device + + +class _FakeVideoDecoder: + __call__ = abc_._cached_video_decoder_call + + def __init__(self, builder: object, device: torch.device = _CUDA) -> None: + self._decoder_builder = builder + self._dtype = torch.bfloat16 + self._device = device + self.orig_calls = 0 + + +@pytest.fixture(autouse=True) +def _reset_cache_state(): + abc_.set_enabled(True) + abc_.set_streaming_enabled(True) + abc_.evict() + yield + abc_.set_enabled(True) + abc_.set_streaming_enabled(False) + abc_.evict() + + +def test_hit_reuses_model_across_block_instances() -> None: + builder = _FakeBuilder("ckpt.safetensors") + cond_1 = _FakeImageConditioner(builder) + cond_2 = _FakeImageConditioner(builder) + + out_1 = cond_1(lambda m: m) + out_2 = cond_2(lambda m: m) + + assert builder.build_count == 1, "second call must reuse the cached build" + assert out_1 is out_2 + assert abc_._cache and next(iter(abc_._cache.values())).in_use == 0 + + +def test_identical_content_shares_one_entry() -> None: + """ImageConditioner's and VideoUpsampler's encoder builders have identical + content -> one shared entry (modeled with two same-content builders).""" + builder_a = _FakeBuilder("ckpt.safetensors") + builder_b = _FakeBuilder("ckpt.safetensors") + + _FakeImageConditioner(builder_a)(lambda m: m) + _FakeImageConditioner(builder_b)(lambda m: m) + + assert len(abc_._cache) == 1, "identical builder content must share one cache entry" + assert builder_a.build_count + builder_b.build_count == 1 + + +def test_multi_slot_two_checkpoints_coexist() -> None: + builder_a = _FakeBuilder("a.safetensors") + builder_b = _FakeBuilder("b.safetensors") + + _FakeImageConditioner(builder_a)(lambda m: m) + _FakeImageConditioner(builder_b)(lambda m: m) + + assert len(abc_._cache) == 2 + assert builder_a.built[0].events == [], "multi-slot cache must not cross-evict" + + +@pytest.mark.parametrize( + "setup_gate", + [ + lambda: abc_.set_enabled(False), + lambda: abc_.set_streaming_enabled(False), + ], + ids=["disabled", "streaming-off"], +) +def test_gates_delegate_to_original(monkeypatch: pytest.MonkeyPatch, setup_gate) -> None: + orig_calls: list[object] = [] + monkeypatch.setattr(abc_, "_orig_image_conditioner_call", lambda self, fn: orig_calls.append(self)) + setup_gate() + builder = _FakeBuilder("ckpt.safetensors") + cond = _FakeImageConditioner(builder) + + cond(lambda m: m) + + assert orig_calls == [cond], "gated-off call must delegate to the original" + assert builder.build_count == 0 + assert not abc_._cache + + +def test_non_cuda_device_and_non_single_gpu_builder_delegate(monkeypatch: pytest.MonkeyPatch) -> None: + orig_calls: list[str] = [] + monkeypatch.setattr(abc_, "_orig_image_conditioner_call", lambda self, fn: orig_calls.append("ic")) + monkeypatch.setattr( + abc_, + "_orig_video_decoder_call", + lambda self, latent, tc, gen, dtype=None: iter(orig_calls.append("vd") or []), + ) + + _FakeImageConditioner(_FakeBuilder("ckpt.safetensors"), device=torch.device("mps"))(lambda m: m) + + class _CustomBuilder: # multi-GPU style, not a SingleGPUModelBuilder + pass + + list(_FakeVideoDecoder(_CustomBuilder())(torch.zeros(1))) + + assert orig_calls == ["ic", "vd"] + assert not abc_._cache + + +def test_effective_dtype_keys_differ() -> None: + builder = _FakeBuilder("ckpt.safetensors") + key_bf16 = abc_._key(builder, torch.bfloat16, _CUDA) + key_fp32 = abc_._key(builder, torch.float32, _CUDA) + assert key_bf16 != key_fp32, "vocoder's effective dtype must produce a distinct key" + + +def test_evict_frees_all_and_clears() -> None: + builder_a = _FakeBuilder("a.safetensors") + builder_b = _FakeBuilder("b.safetensors") + _FakeImageConditioner(builder_a)(lambda m: m) + _FakeImageConditioner(builder_b)(lambda m: m) + + abc_.evict() + + assert builder_a.built[0].events == ["to:meta"] + assert builder_b.built[0].events == ["to:meta"] + assert not abc_._cache + abc_.evict() # safe when empty + + +def test_evict_raises_while_checked_out() -> None: + builder = _FakeBuilder("ckpt.safetensors") + cond = _FakeImageConditioner(builder) + + def _inside(model: object) -> object: + with pytest.raises(RuntimeError, match="in use"): + abc_.evict() + assert builder.built[0].events == [], "model must NOT be freed while checked out" + return model + + cond(_inside) + abc_.evict() # released -> eviction succeeds + assert builder.built[0].events == ["to:meta"] + + +def test_in_use_bypass_delegates_and_preserves_entry(monkeypatch: pytest.MonkeyPatch) -> None: + orig_calls: list[object] = [] + monkeypatch.setattr(abc_, "_orig_image_conditioner_call", lambda self, fn: orig_calls.append(self)) + builder = _FakeBuilder("ckpt.safetensors") + outer = _FakeImageConditioner(builder) + inner = _FakeImageConditioner(builder) + + def _nested(model: object) -> object: + inner(lambda m: m) # same entry currently checked out -> must bypass + return model + + outer(_nested) + + assert orig_calls == [inner], "concurrent checkout must delegate to the original" + assert builder.build_count == 1, "cached model must not be rebuilt or shared" + assert len(abc_._cache) == 1 and next(iter(abc_._cache.values())).in_use == 0 + + +def test_video_decoder_iterator_lifecycle() -> None: + builder = _FakeBuilder("ckpt.safetensors") + decoder_block = _FakeVideoDecoder(builder) + + # (a) never-started iterator leaves no checkout + _unused = decoder_block(torch.zeros(1)) + assert builder.build_count == 0, "checkout must happen at first next(), not at call" + assert all(e.in_use == 0 for e in abc_._cache.values()) + + # (b) full consumption releases + chunks = list(decoder_block(torch.zeros(1))) + assert chunks == ["chunk-0", "chunk-1"] + assert builder.build_count == 1 + assert next(iter(abc_._cache.values())).in_use == 0 + + # (c) close() after first next() releases (GeneratorExit path) + it = decoder_block(torch.zeros(1)) + assert next(it) == "chunk-0" + assert next(iter(abc_._cache.values())).in_use == 1 + it.close() + assert next(iter(abc_._cache.values())).in_use == 0 + assert builder.build_count == 1, "reuse across iterators" + assert builder.built[0].events == [], "cached decoder must never be meta-swapped by the iterator" + + +def test_set_enabled_false_evicts_and_delegates(monkeypatch: pytest.MonkeyPatch) -> None: + builder = _FakeBuilder("ckpt.safetensors") + _FakeImageConditioner(builder)(lambda m: m) + assert abc_._cache + + abc_.set_enabled(False) + + assert not abc_._cache + assert builder.built[0].events == ["to:meta"] + orig_calls: list[object] = [] + monkeypatch.setattr(abc_, "_orig_image_conditioner_call", lambda self, fn: orig_calls.append(self)) + _FakeImageConditioner(builder)(lambda m: m) + assert len(orig_calls) == 1 + + +def test_key_self_check() -> None: + a = ("ckpt", None, (), (), "SingleGPUModelBuilder", torch.bfloat16, _CUDA) + b = ("ckpt", None, (), (), "SingleGPUModelBuilder", torch.bfloat16, _CUDA) + c = ("ckpt", None, (), (), "SingleGPUModelBuilder", torch.float32, _CUDA) + assert a == b + assert a != c diff --git a/backend/tests/test_diffusion_stage_cache.py b/backend/tests/test_diffusion_stage_cache.py index 9f587a82a..9a8c0ca30 100644 --- a/backend/tests/test_diffusion_stage_cache.py +++ b/backend/tests/test_diffusion_stage_cache.py @@ -12,24 +12,28 @@ from __future__ import annotations +from types import SimpleNamespace + import pytest +import torch from ltx_core.block_streaming import StreamingModelBuilder from ltx_core.loader.single_gpu_model_builder import SingleGPUModelBuilder +from ltx_core.model.transformer.model import X0Model from ltx_core.allocator_trim_strategy import AllocatorTrimStrategy from services.patches import diffusion_stage_cache as dsc +# MPS disk streaming uses a small positive slot count (ltx_pipelines.utils.blocks +# passes DISK_CPU_SLOTS); any positive value exercises the same exclusion branch. +_DISK_SLOTS = 4 + class _FakeModel: def __init__(self) -> None: - self.freed_to: str | None = None self.disposed = False - def to(self, device: str) -> "_FakeModel": - self.freed_to = device - return self - def dispose(self) -> None: + # 1.2.0 frees resident (X0Model) storage via Disposable.dispose(), not .to("meta"). self.disposed = True @@ -67,12 +71,59 @@ def _single_gpu_builder(model_path: str, loras: tuple[object, ...] = ()) -> Sing return SingleGPUModelBuilder(model_class_configurator=object, model_path=model_path, loras=loras) +class _FakeStreamingWrapper(torch.nn.Module): + """Stands in for BlockStreamingWrapper: a real nn.Module (so X0Model accepts it) + with recording teardown()/to().""" + + def __init__(self) -> None: + super().__init__() + self.events: list[str] = [] + + def teardown(self) -> None: + self.events.append("teardown") + + def dispose(self) -> None: + # 1.2.0 streaming evict metas storage via Disposable.dispose() (not .to("meta")). + self.events.append("dispose") + + def to(self, device: str) -> "_FakeStreamingWrapper": # type: ignore[override] + self.events.append(f"to:{device}") + return self + + +class _FakeStreamingBuilder(StreamingModelBuilder): + """Real StreamingModelBuilder (so isinstance + the key's content properties work) + whose build() returns a fake wrapper instead of touching disk/GPU.""" + + def __init__(self, model_path: str, *, cpu_slots_count: int | None = None, loras: tuple[object, ...] = ()) -> None: + super().__init__( + model_class_configurator=object, # type: ignore[arg-type] + model_path=model_path, + loras=loras, # type: ignore[arg-type] + cpu_slots_count=cpu_slots_count, + ) + self.build_count = 0 + self.built: list[_FakeStreamingWrapper] = [] + + def build(self, *args: object, **kwargs: object) -> _FakeStreamingWrapper: # type: ignore[override] + self.build_count += 1 + wrapper = _FakeStreamingWrapper() + self.built.append(wrapper) + return wrapper + + +def _streaming_stage(builder: _FakeStreamingBuilder) -> _FakeStage: + return _FakeStage(builder, is_streaming=True) + + @pytest.fixture(autouse=True) def _reset_cache_state(): dsc.set_enabled(True) + dsc.set_streaming_enabled(False) dsc.evict() yield dsc.set_enabled(True) + dsc.set_streaming_enabled(False) dsc.evict() @@ -101,7 +152,7 @@ def test_cache_miss_on_different_checkpoint_evicts_old_model() -> None: assert stage_a.build_count == 1 assert stage_b.build_count == 1 - assert model_a.freed_to == "meta", "old model must be freed before building the new one" + assert model_a.disposed, "old model must be freed (disposed) before building the new one" assert model_b is not model_a @@ -118,11 +169,25 @@ def test_different_loras_are_treated_as_a_different_config() -> None: assert stage_b.build_count == 1, "different LoRA set must miss the cache, not reuse stage_a's build" -def test_cacheable_returns_false_for_a_streaming_stage() -> None: - streaming_builder = StreamingModelBuilder(model_class_configurator=object, model_path="ckpt.safetensors") - streaming_stage = _FakeStage(streaming_builder, is_streaming=True) +def test_streaming_stage_not_cacheable_until_opted_in() -> None: + """The streaming kind is gated on set_streaming_enabled (pushed by ltx2_server + only in streaming_models_loading mode) -- load-bearing: IC-LoRA forces CPU-mode + streaming stages even on full-loading cards, which must NOT session-cache.""" + stage = _streaming_stage(_FakeStreamingBuilder("ckpt.safetensors")) - assert dsc._cacheable(streaming_stage) is False + assert dsc._cacheable(stage) is False + + dsc.set_streaming_enabled(True) + assert dsc._cacheable(stage) is True + + +def test_disk_mode_streaming_builder_is_not_cacheable() -> None: + """Only CPU/RAM streaming (cpu_slots_count None, all blocks pinned) is cacheable; + MPS disk streaming stays on the build-per-call path.""" + dsc.set_streaming_enabled(True) + disk_stage = _streaming_stage(_FakeStreamingBuilder("ckpt.safetensors", cpu_slots_count=_DISK_SLOTS)) + + assert dsc._cacheable(disk_stage) is False class _OtherBuilder: @@ -145,7 +210,7 @@ def test_non_single_gpu_builder_bypasses_cache_and_evicts_resident_model() -> No with dsc._cached_transformer_ctx(other_stage) as other_model: pass - assert cached_model.freed_to == "meta", "resident cache must be freed before the non-cacheable build" + assert cached_model.disposed, "resident cache must be freed before the non-cacheable build" assert dsc._cached_model is None assert other_stage.build_count == 1 assert other_model is not cached_model @@ -158,7 +223,7 @@ def test_disabled_bypasses_cache_and_evicts_immediately() -> None: assert dsc._cached_model is not None dsc.set_enabled(False) - assert cached_model.freed_to == "meta" + assert cached_model.disposed assert dsc._cached_model is None with dsc._cached_transformer_ctx(stage): @@ -177,7 +242,7 @@ def test_evict_frees_and_clears_cache() -> None: dsc.evict() - assert model.freed_to == "meta" + assert model.disposed assert dsc._cached_model is None assert dsc._cached_key is None @@ -194,27 +259,27 @@ def test_cache_key_self_check() -> None: # --- hardening: dead-reference gc pass + in-use concurrency guard ------------ # -def test_evict_runs_gc_collect_before_meta_swap(monkeypatch: pytest.MonkeyPatch) -> None: - """The pre-free gc.collect() must run BEFORE .to('meta') (ComfyUI ordering).""" +def test_evict_runs_gc_collect_before_dispose(monkeypatch: pytest.MonkeyPatch) -> None: + """The pre-free gc.collect() must run BEFORE dispose() metas the storage (ComfyUI ordering).""" stage = _FakeStage(_single_gpu_builder("ckpt.safetensors")) with dsc._cached_transformer_ctx(stage) as model: pass events: list[str] = [] monkeypatch.setattr(dsc.gc, "collect", lambda *a, **k: events.append("gc")) - original_to = model.to + original_dispose = model.dispose - def _record_to(device: str) -> "_FakeModel": - events.append(f"to:{device}") - return original_to(device) + def _record_dispose() -> None: + events.append("dispose") + original_dispose() - monkeypatch.setattr(model, "to", _record_to) + monkeypatch.setattr(model, "dispose", _record_dispose) dsc.evict() - # cleanup_memory() runs its own gc.collect() AFTER the meta-swap, so a trailing - # "gc" is expected; assert only that the pre-free gc precedes the meta-swap. - assert events[:2] == ["gc", "to:meta"], f"pre-free gc must precede the meta-swap, got {events}" + # cleanup_memory() runs its own gc.collect() AFTER the dispose, so a trailing + # "gc" is expected; assert only that the pre-free gc precedes the dispose. + assert events[:2] == ["gc", "dispose"], f"pre-free gc must precede dispose(), got {events}" def test_in_use_is_tracked_across_the_yield() -> None: @@ -232,12 +297,12 @@ def test_evict_while_in_use_raises_and_does_not_free() -> None: # transformer is mid-denoise. It must fail loud rather than free the model. with pytest.raises(RuntimeError, match="in use"): dsc.evict() - assert model.freed_to is None, "model must NOT be freed while in use" + assert not model.disposed, "model must NOT be freed while in use" assert dsc._cached_model is model, "cache must remain intact after a rejected evict" # Once the caller is done, eviction works normally again. dsc.evict() - assert model.freed_to == "meta" + assert model.disposed assert dsc._cached_model is None @@ -249,3 +314,223 @@ def test_set_enabled_false_while_in_use_raises() -> None: dsc.set_enabled(False) # Evict runs before the flag flips, so a rejected toggle stays un-applied. assert dsc._enabled is True, "failed toggle-off must not half-apply" + + +# --- streaming (session-scoped) kind ------------------------------------------ # + + +def test_streaming_hit_builds_once_and_rewraps_fresh_x0model() -> None: + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + stage_1 = _streaming_stage(builder) + stage_2 = _streaming_stage(builder) + + with dsc._cached_transformer_ctx(stage_1) as model_1: + pass + with dsc._cached_transformer_ctx(stage_2) as model_2: + pass + + assert builder.build_count == 1, "stage_2 must reuse stage_1's cached streaming build" + assert isinstance(model_1, X0Model) + assert isinstance(model_2, X0Model) + assert model_1 is not model_2, "each checkout gets a fresh stateless X0Model wrapper" + assert model_1.velocity_model is model_2.velocity_model, "both wrap the SAME cached wrapper" + assert dsc._cached_kind == "streaming" + + +def test_streaming_evict_tears_down_before_meta_swap() -> None: + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + with dsc._cached_transformer_ctx(_streaming_stage(builder)): + pass + wrapper = builder.built[0] + + dsc.evict() + + assert wrapper.events == ["teardown", "dispose"], ( + "eviction must mirror _streaming_model's finally: teardown() (frees pins) " + f"BEFORE the meta-swap, got {wrapper.events}" + ) + assert dsc._cached_model is None + assert dsc._cached_kind is None + + +def _fake_virtual_memory(available_gb: float) -> SimpleNamespace: + return SimpleNamespace(available=int(available_gb * 2**30)) + + +def test_generation_start_keeps_streaming_entry_when_ram_is_fine(monkeypatch: pytest.MonkeyPatch) -> None: + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + with dsc._cached_transformer_ctx(_streaming_stage(builder)): + pass + monkeypatch.setattr(dsc.psutil, "virtual_memory", lambda: _fake_virtual_memory(20.0)) + + dsc.evict_for_generation_start() + + assert dsc._cached_model is not None, "session-scoped streaming entry must survive generation start" + assert builder.built[0].events == [] + + +def test_generation_start_evicts_streaming_entry_under_ram_pressure(monkeypatch: pytest.MonkeyPatch) -> None: + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + with dsc._cached_transformer_ctx(_streaming_stage(builder)): + pass + monkeypatch.setattr(dsc.psutil, "virtual_memory", lambda: _fake_virtual_memory(dsc._MIN_FREE_RAM_GB - 1)) + + dsc.evict_for_generation_start() + + assert dsc._cached_model is None, "streaming entry must be evicted under external RAM pressure" + assert builder.built[0].events == ["teardown", "dispose"] + + +def test_generation_start_always_evicts_resident_entry(monkeypatch: pytest.MonkeyPatch) -> None: + """The 5090 VRAM-collision fix: a resident (VRAM) entry never survives into the + next generation, regardless of RAM headroom.""" + stage = _FakeStage(_single_gpu_builder("ckpt.safetensors")) + with dsc._cached_transformer_ctx(stage) as model: + pass + monkeypatch.setattr(dsc.psutil, "virtual_memory", lambda: _fake_virtual_memory(64.0)) + + dsc.evict_for_generation_start() + + assert model.disposed + assert dsc._cached_model is None + + +def test_resident_and_streaming_keys_never_collide() -> None: + """Identical content must still miss across builder types (e.g. IC-LoRA's + resident stage_1 vs forced-streaming stage_2 on a full-loading card).""" + dsc.set_streaming_enabled(True) + resident_stage = _FakeStage(_single_gpu_builder("ckpt.safetensors")) + streaming_stage = _streaming_stage(_FakeStreamingBuilder("ckpt.safetensors")) + + assert dsc._cache_key(resident_stage) != dsc._cache_key(streaming_stage) + + +def test_streaming_checkout_tracks_in_use_and_rejects_evict() -> None: + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + with dsc._cached_transformer_ctx(_streaming_stage(builder)): + assert dsc._in_use == 1 + with pytest.raises(RuntimeError, match="in use"): + dsc.evict() + assert builder.built[0].events == [], "wrapper must NOT be torn down while checked out" + assert dsc._in_use == 0 + + +def test_set_streaming_enabled_false_evicts_streaming_entry() -> None: + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + with dsc._cached_transformer_ctx(_streaming_stage(builder)): + pass + assert dsc._cached_kind == "streaming" + + dsc.set_streaming_enabled(False) + + assert dsc._cached_model is None + assert builder.built[0].events == ["teardown", "dispose"] + # A resident entry is unaffected by the streaming gate. + stage = _FakeStage(_single_gpu_builder("ckpt.safetensors")) + with dsc._cached_transformer_ctx(stage): + pass + assert dsc._cached_kind == "resident" + + +def test_abnormal_exit_marks_dirty_and_next_checkout_rebuilds() -> None: + """An exception unwinding through the denoise loop can leak BufferPool slots + (observed live: 'BufferPool exhausted: all 2 buffers are in use' on every + retry) -- the entry must be rebuilt, never reused.""" + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + with pytest.raises(RuntimeError, match="boom"): + with dsc._cached_transformer_ctx(_streaming_stage(builder)): + raise RuntimeError("boom") + + assert dsc._dirty is True + assert dsc._in_use == 0, "abnormal exit must still release the checkout" + assert builder.built[0].events == [], "dirty entry is kept (lazily evicted), not freed mid-unwind" + + with dsc._cached_transformer_ctx(_streaming_stage(builder)): + pass + + assert builder.build_count == 2, "dirty entry must be rebuilt, not reused" + assert builder.built[0].events == ["teardown", "dispose"], "dirty wrapper torn down before rebuild" + assert dsc._dirty is False, "rebuild clears the dirty flag" + + +def test_generation_start_evicts_dirty_streaming_entry_despite_free_ram(monkeypatch: pytest.MonkeyPatch) -> None: + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + with pytest.raises(RuntimeError, match="boom"): + with dsc._cached_transformer_ctx(_streaming_stage(builder)): + raise RuntimeError("boom") + monkeypatch.setattr(dsc.psutil, "virtual_memory", lambda: _fake_virtual_memory(64.0)) + + dsc.evict_for_generation_start() + + assert dsc._cached_model is None, "dirty entry must never survive into a new generation" + assert builder.built[0].events == ["teardown", "dispose"] + + +def test_concurrent_checkout_bypasses_cache_with_isolated_build(monkeypatch: pytest.MonkeyPatch) -> None: + """A zombie generation's checkout must not share the cached wrapper (2-slot + BufferPool) with a new generation -- the second checkout gets an isolated + one-shot build via the original ctx, mirroring pre-cache behavior.""" + from contextlib import contextmanager + + dsc.set_streaming_enabled(True) + builder = _FakeStreamingBuilder("ckpt.safetensors") + orig_calls: list[object] = [] + + @contextmanager + def _fake_orig(stage: object, **kwargs: object): + orig_calls.append(stage) + yield object() + + monkeypatch.setattr(dsc, "_orig_transformer_ctx", _fake_orig) + + outer_stage = _streaming_stage(builder) + inner_stage = _streaming_stage(builder) + with dsc._cached_transformer_ctx(outer_stage): + assert dsc._in_use == 1 + with dsc._cached_transformer_ctx(inner_stage): + pass + assert orig_calls == [inner_stage], "concurrent checkout must delegate to the original ctx" + assert builder.build_count == 1, "cached wrapper must not be rebuilt or shared" + assert dsc._in_use == 1, "bypassed checkout must not touch the in-use counter" + assert dsc._in_use == 0 + assert dsc._cached_model is not None, "cached entry survives a concurrent bypass" + + +def test_streaming_stage_bypasses_and_evicts_when_gate_is_off() -> None: + """With the gate off (full-loading cards), a streaming stage takes the bypass + branch and unconditionally evicts -- the original NON-CACHEABLE TRANSITIONS + behavior (IC-LoRA repro) must be preserved.""" + cacheable_stage = _FakeStage(_single_gpu_builder("ckpt.safetensors")) + with dsc._cached_transformer_ctx(cacheable_stage) as cached_model: + pass + assert dsc._cached_model is cached_model + + streaming_stage = _streaming_stage(_FakeStreamingBuilder("ckpt.safetensors")) + orig_calls: list[object] = [] + + from contextlib import contextmanager + + @contextmanager + def _fake_orig(stage: object, **kwargs: object): + orig_calls.append(stage) + yield object() + + original = dsc._orig_transformer_ctx + dsc._orig_transformer_ctx = _fake_orig # type: ignore[assignment] + try: + with dsc._cached_transformer_ctx(streaming_stage): + pass + finally: + dsc._orig_transformer_ctx = original # type: ignore[assignment] + + assert cached_model.disposed, "bypass must evict the resident entry first" + assert dsc._cached_model is None + assert orig_calls == [streaming_stage], "gate-off streaming must delegate to the original ctx" diff --git a/backend/tests/test_diffvae_decode_vram.py b/backend/tests/test_diffvae_decode_vram.py index 34a657cf1..38155f245 100644 --- a/backend/tests/test_diffvae_decode_vram.py +++ b/backend/tests/test_diffvae_decode_vram.py @@ -13,11 +13,11 @@ class _FakeModel: def __init__(self) -> None: - self.freed_to: str | None = None + self.disposed = False - def to(self, device: str) -> "_FakeModel": - self.freed_to = device - return self + def dispose(self) -> None: + # 1.2.0 diffusion_stage_cache evict frees storage via Disposable.dispose(). + self.disposed = True class _FakeDecoder: @@ -47,24 +47,42 @@ def test_patch_rebinds_video_decoder_call() -> None: def test_video_decoder_call_evicts_cached_transformer(monkeypatch) -> None: + # DiffVAE decode (is_diffusion_video_vae True) frees the resident transformer. model = _FakeModel() dsc._cached_model = model dsc._cached_key = ("planted",) + monkeypatch.setattr(patch, "is_diffusion_video_vae", lambda _path: True) monkeypatch.setattr(patch, "_orig_video_decoder_call", lambda *args, **kwargs: "ok") - assert patch._patched_video_decoder_call(object()) == "ok" - assert model.freed_to == "meta" + assert patch._patched_video_decoder_call(_FakeDecoder(torch.device("cuda"))) == "ok" + assert model.disposed assert dsc._cached_model is None +def test_conv_video_decoder_keeps_cached_transformer(monkeypatch) -> None: + # Conv VAE decode (is_diffusion_video_vae False) is small and must NOT evict the + # session-cached streaming transformer -- keeping it resident across generations is + # the cross-gen win. Only DiffVAE, which needs the ~20GB, evicts it. + model = _FakeModel() + dsc._cached_model = model + dsc._cached_key = ("planted",) + monkeypatch.setattr(patch, "is_diffusion_video_vae", lambda _path: False) + monkeypatch.setattr(patch, "_orig_video_decoder_call", lambda *args, **kwargs: "ok") + + assert patch._patched_video_decoder_call(_FakeDecoder(torch.device("cuda"))) == "ok" + assert not model.disposed, "conv decode must keep the cached transformer resident" + assert dsc._cached_model is model + + def test_video_decoder_call_cleans_allocator_even_when_cache_empty(monkeypatch) -> None: cleaned: list[bool] = [] dsc.set_enabled(False) dsc.evict() + monkeypatch.setattr(patch, "is_diffusion_video_vae", lambda _path: True) monkeypatch.setattr(patch, "_orig_video_decoder_call", lambda *args, **kwargs: "ok") monkeypatch.setattr(patch, "cleanup_memory", lambda: cleaned.append(True)) - assert patch._patched_video_decoder_call(object()) == "ok" + assert patch._patched_video_decoder_call(_FakeDecoder(torch.device("cuda"))) == "ok" assert cleaned == [True] @@ -141,6 +159,7 @@ def test_cuda_diffvae_caps_29gib_budget_before_recommend(monkeypatch) -> None: def test_non_cuda_decoder_keeps_pipeline_tiling(monkeypatch) -> None: + monkeypatch.setattr(patch, "is_diffusion_video_vae", lambda _path: True) monkeypatch.setattr(patch, "tiling_config_for_vae", lambda *_args, **_kwargs: pytest.fail("should not re-resolve")) monkeypatch.setattr(patch, "_orig_video_decoder_call", lambda _self, *args, **kwargs: (args, kwargs)) latent = torch.zeros(1, 4, 61, 18, 32) diff --git a/backend/tests/test_fp8_sidecar_cache.py b/backend/tests/test_fp8_sidecar_cache.py new file mode 100644 index 000000000..e617e7331 --- /dev/null +++ b/backend/tests/test_fp8_sidecar_cache.py @@ -0,0 +1,236 @@ +"""Tests for the fp8_sidecar_cache patch. + +Uses small real safetensors files (bf16 + float8_e4m3fn tensors) in tmp_path and +a recording fake inner loader implementing the StateDictLoader protocol -- no +GPU, no mock library. The patched StreamingModelBuilder._build_pinned_source is +not exercised end-to-end (it needs a real meta model); the loader wrapper and +the eligibility gate carry the behavior and are tested directly. +""" + +from __future__ import annotations + +import json +import os +from pathlib import Path + +import pytest +import torch +from safetensors.torch import save_file + +from ltx_core.loader.primitives import StateDict +from ltx_core.loader.sd_ops import SDOps +from services.patches import fp8_sidecar_cache as fsc + + +def _block_tensors() -> dict[str, torch.Tensor]: + return { + "transformer_blocks.0.attn1.to_q.weight": torch.zeros(4, 4, dtype=torch.float8_e4m3fn), + "transformer_blocks.0.norm1.weight": torch.ones(4, dtype=torch.bfloat16), + "transformer_blocks.1.attn1.to_q.weight": torch.zeros(4, 4, dtype=torch.float8_e4m3fn), + } + + +def _blocks_sd_ops(keys: frozenset[str]) -> SDOps: + return SDOps(name="sd_ops_chain_LTXV+FP8_CAST_PREQUANT_AWARE__blocks", allowed_keys=keys) + + +def _write_checkpoint(path: Path, model_id: str = "wandb-abc") -> None: + # The "original" 43GB checkpoint stand-in: content irrelevant, only its + # header metadata + stat identity are read by the sidecar machinery. + save_file( + {"model.diffusion_model.transformer_blocks.0.attn1.to_q.weight": torch.zeros(2, dtype=torch.bfloat16)}, + str(path), + metadata={"encrypted_wandb_properties": model_id}, + ) + + +class _FakeInnerLoader: + """Records load calls; serves 'original' loads from a fixed tensor dict and + sidecar identity-loads from the actual sidecar file.""" + + def __init__(self, original_tensors: dict[str, torch.Tensor]) -> None: + self._original = original_tensors + self.calls: list[tuple[str, str | None]] = [] # (path, sd_ops name) + + def metadata(self, path: str) -> dict: + return {} + + def load(self, path, sd_ops=None, device=None) -> StateDict: # noqa: ANN001 + path_str = path if isinstance(path, str) else path[0] + self.calls.append((path_str, sd_ops.name if sd_ops is not None else None)) + if sd_ops is None: + # Identity load of the sidecar file itself. + from safetensors.torch import load_file + + tensors = load_file(path_str) + else: + tensors = dict(self._original) + return StateDict( + sd=tensors, + device=torch.device("cpu"), + size=sum(t.numel() * t.element_size() for t in tensors.values()), + dtype={t.dtype for t in tensors.values()}, + ) + + +@pytest.fixture() +def setup(tmp_path: Path): + checkpoint = tmp_path / "model-1.1.safetensors" + _write_checkpoint(checkpoint) + tensors = _block_tensors() + inner = _FakeInnerLoader(tensors) + loader = fsc._SidecarLoader(inner, str(checkpoint)) + sd_ops = _blocks_sd_ops(frozenset(tensors)) + return checkpoint, tensors, inner, loader, sd_ops + + +def test_miss_delegates_writes_sidecar_then_hits(setup) -> None: + checkpoint, tensors, inner, loader, sd_ops = setup + sidecar = fsc._sidecar_path(str(checkpoint)) + + first = loader.load(str(checkpoint), sd_ops=sd_ops, device=torch.device("cpu")) + + assert inner.calls == [(str(checkpoint), sd_ops.name)], "miss must delegate with the ORIGINAL sd_ops" + assert sidecar.exists(), "miss must write the sidecar" + assert not list(sidecar.parent.glob("*.tmp-*")), "atomic write must leave no tmp residue" + metadata, keys = fsc._read_header(sidecar) + assert metadata["sidecar_format"] == fsc._SIDECAR_FORMAT + assert metadata["original_model_id"] == "wandb-abc" + assert metadata["sd_ops_chain"] == sd_ops.name + assert keys == set(tensors) + + second = loader.load(str(checkpoint), sd_ops=sd_ops, device=torch.device("cpu")) + + assert inner.calls[-1] == (str(sidecar), None), "hit must identity-load the sidecar (sd_ops=None)" + for key, tensor in first.sd.items(): + assert second.sd[key].dtype == tensor.dtype, f"dtype must round-trip for {key}" + assert torch.equal(second.sd[key].view(torch.uint8), tensor.view(torch.uint8)) + + +def test_non_blocks_and_none_sd_ops_pass_through(setup) -> None: + checkpoint, tensors, inner, loader, _ = setup + other = SDOps(name="sd_ops_chain_LTXV+FP8_CAST_PREQUANT_AWARE__non_block", allowed_keys=frozenset(tensors)) + + loader.load(str(checkpoint), sd_ops=other) + loader.load(str(checkpoint), sd_ops=None) + + assert not fsc._sidecar_path(str(checkpoint)).exists(), "only __blocks loads are intercepted" + assert [name for _, name in inner.calls] == [other.name, None] + + +@pytest.mark.parametrize( + "mutate", + [ + # +1s, not +1ns: NTFS stores timestamps in 100ns ticks, a sub-tick bump rounds away. + lambda cp: os.utime(cp, ns=(os.stat(cp).st_atime_ns, os.stat(cp).st_mtime_ns + 1_000_000_000)), + lambda cp: _write_checkpoint(cp, model_id="wandb-DIFFERENT"), # content/id + size/mtime change + ], + ids=["mtime", "model-id"], +) +def test_stamp_mismatch_invalidates_deletes_and_rebuilds(setup, mutate) -> None: + checkpoint, _tensors, inner, loader, sd_ops = setup + loader.load(str(checkpoint), sd_ops=sd_ops) + sidecar = fsc._sidecar_path(str(checkpoint)) + assert sidecar.exists() + + mutate(checkpoint) + loader.load(str(checkpoint), sd_ops=sd_ops) + + assert inner.calls[-1] == (str(checkpoint), sd_ops.name), "invalid sidecar must fall back to the original" + assert sidecar.exists(), "fallback must rewrite a fresh sidecar" + metadata, _ = fsc._read_header(sidecar) + assert metadata["original_mtime_ns"] == str(os.stat(checkpoint).st_mtime_ns), "rewritten stamp must be current" + + +def test_key_set_mismatch_invalidates(setup) -> None: + checkpoint, tensors, inner, loader, sd_ops = setup + loader.load(str(checkpoint), sd_ops=sd_ops) + + # Same stamp, different expected partition (e.g. layer-count drift). + smaller = _blocks_sd_ops(frozenset(list(tensors)[:1])) + loader.load(str(checkpoint), sd_ops=smaller) + + assert inner.calls[-1] == (str(checkpoint), smaller.name), "key-set mismatch must fall back" + + +def test_corrupted_sidecar_deleted_and_rebuilt(setup) -> None: + checkpoint, _tensors, inner, loader, sd_ops = setup + loader.load(str(checkpoint), sd_ops=sd_ops) + sidecar = fsc._sidecar_path(str(checkpoint)) + sidecar.write_bytes(b"\x00" * 16) # garbage header + + loader.load(str(checkpoint), sd_ops=sd_ops) + + assert inner.calls[-1] == (str(checkpoint), sd_ops.name) + metadata, _ = fsc._read_header(sidecar) + assert metadata["sidecar_format"] == fsc._SIDECAR_FORMAT, "corrupt sidecar must be replaced by a valid one" + + +def test_write_skipped_when_disk_space_low(setup, monkeypatch: pytest.MonkeyPatch) -> None: + checkpoint, _tensors, inner, loader, sd_ops = setup + monkeypatch.setattr(fsc, "_free_bytes", lambda _dir: 0) + + result = loader.load(str(checkpoint), sd_ops=sd_ops) + + assert result.sd, "the delegated load result must still be returned" + assert not fsc._sidecar_path(str(checkpoint)).exists(), "preflight must skip the write" + assert not list(checkpoint.parent.glob("*.tmp-*")) + + +def test_write_failure_is_swallowed(setup, monkeypatch: pytest.MonkeyPatch) -> None: + checkpoint, _tensors, _inner, loader, sd_ops = setup + + def _boom(*_a: object, **_k: object) -> None: + raise OSError("disk detached") + + monkeypatch.setattr(fsc, "_stamp", _boom) + + result = loader.load(str(checkpoint), sd_ops=sd_ops) + + assert result.sd, "write failure must never break the load" + assert not fsc._sidecar_path(str(checkpoint)).exists() + assert not list(checkpoint.parent.glob("*.tmp-*")), "failed write must clean up its tmp file" + + +def test_scale_keys_refused_on_write_and_read(setup) -> None: + checkpoint, tensors, inner, loader, _ = setup + scale_tensors = {**tensors, "transformer_blocks.0.attn1.to_q.weight_scale": torch.ones(1)} + inner_with_scales = _FakeInnerLoader(scale_tensors) + loader2 = fsc._SidecarLoader(inner_with_scales, str(checkpoint)) + sd_ops = _blocks_sd_ops(frozenset(scale_tensors)) + + loader2.load(str(checkpoint), sd_ops=sd_ops) + assert not fsc._sidecar_path(str(checkpoint)).exists(), "write must refuse *_scale keys" + + # Hand-built sidecar containing a scale key must be rejected on read. + sidecar = fsc._sidecar_path(str(checkpoint)) + save_file(scale_tensors, str(sidecar), metadata=fsc._stamp(str(checkpoint), sd_ops.name)) + assert fsc._valid(sidecar, str(checkpoint), sd_ops) is False + assert not sidecar.exists(), "invalid sidecar must be deleted" + + +def test_eligibility_gate() -> None: + class _FakeBuilder: + model_path: object = "ckpt.safetensors" + model_sd_ops: SDOps | None = SDOps(name="LTXV+FP8_CAST_PREQUANT_AWARE") + + builder = _FakeBuilder() + cuda = torch.cuda.is_available() + + assert fsc._eligible(builder) is cuda # type: ignore[arg-type] + + builder.model_path = ("a.safetensors", "b.safetensors") + assert fsc._eligible(builder) is False, "sharded checkpoints are not eligible" # type: ignore[arg-type] + + builder.model_path = "ckpt.safetensors" + builder.model_sd_ops = SDOps(name="GEMMA_LLM_KEY_OPS") + assert fsc._eligible(builder) is False, "non-fp8 chains (e.g. Gemma) are not eligible" # type: ignore[arg-type] + + builder.model_sd_ops = None + assert fsc._eligible(builder) is False # type: ignore[arg-type] + + +def test_sidecar_path_naming() -> None: + p = fsc._sidecar_path(r"E:\models\ltx-2.3-22b-distilled-1.1.safetensors") + assert p.name == "ltx-2.3-22b-distilled-1.1.fp8-blocks-cache.safetensors" + assert str(p.parent).endswith("models") diff --git a/backend/tests/test_generation.py b/backend/tests/test_generation.py index d9da81e0f..edf442dfa 100644 --- a/backend/tests/test_generation.py +++ b/backend/tests/test_generation.py @@ -594,6 +594,37 @@ def test_resolution_mapping_720p(self, client, test_state, fake_services, create assert call["width"] == 1280 assert call["height"] == 704 + def test_resolution_mapping_21_9_540p(self, client, test_state, fake_services, create_fake_model_files): + create_fake_model_files() + _enable_local_text_encoding(test_state) + + r = client.post("/api/generate", json={**_T2V_JSON, "aspectRatio": "21:9"}) + assert r.status_code == 200 + + pipeline = fake_services.fast_video_pipeline + call = pipeline.generate_calls[0] + assert call["width"] == 1216 + assert call["height"] == 512 + + def test_resolution_mapping_21_9_720p(self, client, test_state, fake_services, create_fake_model_files): + create_fake_model_files() + _enable_local_text_encoding(test_state) + + r = client.post("/api/generate", json={**_T2V_JSON, "resolution": "720p", "aspectRatio": "21:9"}) + assert r.status_code == 200 + + pipeline = fake_services.fast_video_pipeline + call = pipeline.generate_calls[0] + assert call["width"] == 1664 + assert call["height"] == 704 + + def test_21_9_rejected_at_1080p(self, client, test_state, fake_services, create_fake_model_files): + create_fake_model_files() + _enable_local_text_encoding(test_state) + + r = client.post("/api/generate", json={**_T2V_JSON, "resolution": "1080p", "aspectRatio": "21:9"}) + assert r.status_code == 400 + def test_locked_seed(self, client, test_state, fake_services, create_fake_model_files): create_fake_model_files() _enable_local_text_encoding(test_state) @@ -2367,6 +2398,35 @@ def test_dimension_clamping(self, client, fake_services, create_fake_model_files assert call["width"] == 1008 assert call["height"] == 1008 + def test_variation_forwarded_to_pipeline(self, client, fake_services, create_fake_model_files): + create_fake_model_files(include_zit=True) + r = client.post("/api/generate-image", json={"prompt": "test", "variation": 0.4}) + assert r.status_code == 200 + assert fake_services.image_generation_pipeline.generate_calls[0]["variation"] == 0.4 + + def test_variation_defaults_off_and_is_bounded(self, client, fake_services, create_fake_model_files): + create_fake_model_files(include_zit=True) + assert client.post("/api/generate-image", json={"prompt": "test"}).status_code == 200 + assert fake_services.image_generation_pipeline.generate_calls[0]["variation"] == 0.0 + assert client.post("/api/generate-image", json={"prompt": "test", "variation": 1.5}).status_code == 422 + + def test_krea2_rejected_in_release_builds(self, client, fake_services, create_fake_model_files): + create_fake_model_files(include_zit=True) + r = client.post("/api/generate-image", json={"prompt": "test", "model": "krea-2-turbo"}) + assert r.status_code == 400 + assert "KREA_2_DEV_ONLY" in r.text + assert fake_services.image_generation_pipeline.generate_calls == [] + + def test_krea2_allowed_in_dev_mode(self, client, test_state, fake_services, create_fake_model_files): + create_fake_model_files(include_zit=True) + krea_dir = resolve_model_path(test_state.config.default_models_dir, "krea-2-turbo") + krea_dir.mkdir(parents=True, exist_ok=True) + (krea_dir / "model.safetensors").write_bytes(b"\x00" * 1024) + test_state.config.dev_mode = True + r = client.post("/api/generate-image", json={"prompt": "test", "model": "krea-2-turbo"}) + assert r.status_code == 200 + assert len(fake_services.image_generation_pipeline.generate_calls) == 1 + def test_num_images_clamped(self, client, fake_services, create_fake_model_files): create_fake_model_files(include_zit=True) r = client.post( diff --git a/backend/tests/test_generation_interrupt.py b/backend/tests/test_generation_interrupt.py index 31ebe3be6..f38a619cb 100644 --- a/backend/tests/test_generation_interrupt.py +++ b/backend/tests/test_generation_interrupt.py @@ -86,12 +86,66 @@ def test_zit_pipeline_passes_step_end_callback() -> None: ZitImageGenerationPipeline, ) + # generate() takes its callback from _variation_inputs (covered by the tests below). generate_src = inspect.getsource(ZitImageGenerationPipeline.generate) edit_src = inspect.getsource(ZitImageGenerationPipeline.edit) - assert "callback_on_step_end=diffusers_step_callback" in generate_src + assert "callback_on_step_end=callback" in generate_src assert "callback_on_step_end=diffusers_step_callback" in edit_src +class _StubZImagePipeline: + _execution_device = "cpu" + + def __init__(self) -> None: + import torch + + self.clean = [torch.randn(5, 8, generator=torch.Generator().manual_seed(0))] + + def encode_prompt(self, **_: object) -> tuple[list[object], list[object]]: + return self.clean, [] + + +def _zit_without_init() -> object: + from services.image_generation_pipeline.zit_image_generation_pipeline import ( + ZitImageGenerationPipeline, + ) + + return ZitImageGenerationPipeline.__new__(ZitImageGenerationPipeline) + + +def test_zit_variation_off_uses_plain_interrupt_callback() -> None: + zit = _zit_without_init() + kwargs, callback, inputs = zit._variation_inputs(_StubZImagePipeline(), "p", 1, 0.0, 4) # type: ignore[attr-defined] + assert kwargs == {"prompt": "p"} + assert callback is diffusers_step_callback + assert inputs == ["latents"] + + +def test_zit_variation_boost_perturbs_then_restores_and_stays_interruptible() -> None: + import torch + + clear() + stub = _StubZImagePipeline() + zit = _zit_without_init() + kwargs, callback, inputs = zit._variation_inputs(stub, "p", 7, 0.5, 4) # type: ignore[attr-defined] + noisy = kwargs["prompt_embeds"] + assert inputs == ["prompt_embeds"] + assert not torch.equal(noisy[0], stub.clean[0]) + + # Deterministic for the same seed + strength. + again, _, _ = zit._variation_inputs(stub, "p", 7, 0.5, 4) # type: ignore[attr-defined] + assert torch.equal(again["prompt_embeds"][0], noisy[0]) + + # 4 steps x 0.25 boost => clean embeds restored at the end of step 0. + out = callback(None, 0, 0, {"prompt_embeds": noisy}) + assert out["prompt_embeds"] is stub.clean + + request() + with pytest.raises(GenerationCancelledError): + callback(None, 1, 0, {"prompt_embeds": stub.clean}) + clear() + + def test_image_callback_stops_later_steps() -> None: from tests.fakes.services import FakeImageGenerationPipeline diff --git a/backend/tests/test_ic_lora.py b/backend/tests/test_ic_lora.py index cd483c969..cbb0a0da7 100644 --- a/backend/tests/test_ic_lora.py +++ b/backend/tests/test_ic_lora.py @@ -27,6 +27,15 @@ def _install_ic_lora_capable_model(create_fake_model_files, create_fake_ic_lora_ create_fake_ic_lora_files(include_depth=include_depth) +def _write_ic_lora_file(path: Path) -> None: + """Header-only safetensors carrying the IC-LoRA reference marker.""" + path.parent.mkdir(parents=True, exist_ok=True) + blob = json.dumps({"__metadata__": {"reference_downscale_factor": "1"}}).encode("utf-8") + with open(path, "wb") as f: + f.write(struct.pack(" 0 + # info copy moved to the frontend (keyed off role); backend only ships name + role. + assert cp["name"] and cp["role"] + assert by_id[spec.model_cp]["downloaded"] is True + assert by_id[IMG_GEN_MODEL_CP_ID]["downloaded"] is False + + def test_describe_classifies_any_base_version_as_base(self, client): + # An OLDER base transformer must be classified as "base", not "support" — _cp_role + # has to match every version's model_cp, not only the latest. + older = get_ltx_model_spec("ltx-2.3-22b-distilled") + assert older.model_cp != _current_ltx_spec().model_cp + response = client.post("/api/models/describe", json={"cp_ids": [older.model_cp]}) + assert response.status_code == 200 + cp = response.json()["checkpoints"][0] + assert cp["cp_id"] == older.model_cp + assert cp["role"] == "base" + def test_ic_lora_recommendation(self, client, create_fake_model_files, create_fake_ic_lora_files): create_fake_model_files() response = client.get("/api/models/ltx-ic-lora-recommendation") diff --git a/backend/tests/test_qwen_angle_mapping.py b/backend/tests/test_qwen_angle_mapping.py new file mode 100644 index 000000000..9479d47a5 --- /dev/null +++ b/backend/tests/test_qwen_angle_mapping.py @@ -0,0 +1,118 @@ +"""Unit tests for the gizmo-to-prompt bucket mapping.""" + +import pytest + +from services.qwen_multiangle_pipeline.angle_mapping import ( + AZIMUTH_BUCKETS, + DISTANCE_BUCKETS, + ELEVATION_BUCKETS, + mapping_table, + snap_azimuth, + snap_distance, + snap_elevation, + snap_pose, +) + + +def test_vocabulary_is_96_poses(): + assert len(AZIMUTH_BUCKETS) * len(ELEVATION_BUCKETS) * len(DISTANCE_BUCKETS) == 96 + + +# Degrees AND phrases are both viewer-relative camera positions (verified +# empirically 2026-07-10: "left side view" renders the camera at the +# viewer's left). 45 = camera to viewer's right = "front-right quarter view". +@pytest.mark.parametrize( + ("deg", "expected_phrase"), + [ + (0, "front view"), + (45, "front-right quarter view"), + (90, "right side view"), + (131, "back-right quarter view"), + (180, "back view"), + (225, "back-left quarter view"), + (270, "left side view"), + (315, "front-left quarter view"), + (350, "front view"), # wraps to nearest + (-45, "front-left quarter view"), # negative wraps + (360, "front view"), + (22.4, "front view"), # just inside front's half-sector + (22.6, "front-right quarter view"), # just past the boundary + ], +) +def test_azimuth_snapping(deg, expected_phrase): + assert AZIMUTH_BUCKETS[snap_azimuth(deg)][1] == expected_phrase + + +@pytest.mark.parametrize( + ("deg", "expected_label"), + [ + (0, "front"), + (45, "front-right"), + (90, "right side"), + (135, "back-right"), + (180, "back"), + (225, "back-left"), + (270, "left side"), + (315, "front-left"), + ], +) +def test_azimuth_labels_are_viewer_relative(deg, expected_label): + assert AZIMUTH_BUCKETS[snap_azimuth(deg)][2] == expected_label + + +@pytest.mark.parametrize( + ("deg", "expected_phrase"), + [ + (-90, "low-angle shot"), # clamped by nearest + (-30, "low-angle shot"), + (-5, "eye-level shot"), # the ComfyUI screenshot's exact value + (0, "eye-level shot"), + (14, "eye-level shot"), + (16, "elevated shot"), + (44, "elevated shot"), + (46, "high-angle shot"), + (60, "high-angle shot"), + (90, "high-angle shot"), + ], +) +def test_elevation_snapping(deg, expected_phrase): + assert ELEVATION_BUCKETS[snap_elevation(deg)][1] == expected_phrase + + +@pytest.mark.parametrize( + ("zoom", "expected_phrase"), + [ + (0.3, "close-up"), + (0.6, "close-up"), + (0.77, "close-up"), # geometric mean of 0.6 and 1.0 is ~0.775 + (0.78, "medium shot"), + (1.0, "medium shot"), + (1.33, "medium shot"), # geometric mean of 1.0 and 1.8 is ~1.342 + (1.35, "wide shot"), + (1.8, "wide shot"), + (8.0, "wide shot"), # the ComfyUI screenshot's exact value + ], +) +def test_distance_snapping(zoom, expected_phrase): + assert DISTANCE_BUCKETS[snap_distance(zoom)][1] == expected_phrase + + +def test_screenshot_pose_roundtrip(): + """The exact pose from the reference ComfyUI screenshot reproduces its prompt + (ComfyUI's H-angle convention matches ours: viewer-relative camera position).""" + pose = snap_pose(azimuth_deg=131, elevation_deg=-5, zoom=0.6) + assert pose.prompt == " back-right quarter view eye-level shot close-up" + + +def test_prompt_format(): + pose = snap_pose(0, 0, 1.0) + assert pose.prompt == " front view eye-level shot medium shot" + + +def test_mapping_table_shape(): + table = mapping_table() + assert table["trigger"] == "" + assert len(table["azimuths"]) == 8 + assert len(table["elevations"]) == 4 + assert len(table["distances"]) == 3 + assert all("label" in a and "phrase" in a and "deg" in a for a in table["azimuths"]) diff --git a/backend/tests/test_settings.py b/backend/tests/test_settings.py index 8186da85b..e84799a3d 100644 --- a/backend/tests/test_settings.py +++ b/backend/tests/test_settings.py @@ -259,7 +259,10 @@ def test_models_dir_persists_and_loads(self, client, test_state, default_app_set ltx_api_client=fake_services.ltx_api_client, zit_api_client=fake_services.zit_api_client, fast_video_pipeline_class=type(fake_services.fast_video_pipeline), - image_generation_pipeline_class=type(fake_services.image_generation_pipeline), + image_generation_pipeline_classes={ + "z-image-turbo": type(fake_services.image_generation_pipeline), + "krea-2-turbo": type(fake_services.image_generation_pipeline), + }, ic_lora_pipeline_class=type(fake_services.ic_lora_pipeline), depth_processor_pipeline_class=type(fake_services.depth_processor_pipeline), pose_processor_pipeline_class=type(fake_services.pose_processor_pipeline), @@ -287,7 +290,10 @@ def _new_state(self, test_state, default_app_settings): ltx_api_client=fake_services.ltx_api_client, zit_api_client=fake_services.zit_api_client, fast_video_pipeline_class=type(fake_services.fast_video_pipeline), - image_generation_pipeline_class=type(fake_services.image_generation_pipeline), + image_generation_pipeline_classes={ + "z-image-turbo": type(fake_services.image_generation_pipeline), + "krea-2-turbo": type(fake_services.image_generation_pipeline), + }, ic_lora_pipeline_class=type(fake_services.ic_lora_pipeline), depth_processor_pipeline_class=type(fake_services.depth_processor_pipeline), pose_processor_pipeline_class=type(fake_services.pose_processor_pipeline), @@ -349,24 +355,15 @@ def test_update_request_tracks_app_settings_fields(self): class TestResolvedUseConvVae: - def test_none_defaults_on_for_darwin(self, monkeypatch): - monkeypatch.setattr("state.app_settings.sys.platform", "darwin") + def test_none_defaults_on(self): + # Conv VAE ("fast decode") is the default on every platform now; the diffusion + # VAE decode is opt-out via an explicit False, not the CUDA default it once was. assert resolved_use_conv_vae(AppSettings()) is True - def test_none_defaults_off_for_linux(self, monkeypatch): - monkeypatch.setattr("state.app_settings.sys.platform", "linux") - assert resolved_use_conv_vae(AppSettings()) is False - - def test_none_defaults_off_for_windows(self, monkeypatch): - monkeypatch.setattr("state.app_settings.sys.platform", "win32") - assert resolved_use_conv_vae(AppSettings()) is False - - def test_explicit_true_overrides_linux_default(self, monkeypatch): - monkeypatch.setattr("state.app_settings.sys.platform", "linux") + def test_explicit_true(self): assert resolved_use_conv_vae(AppSettings(use_conv_vae=True)) is True - def test_explicit_false_overrides_darwin_default(self, monkeypatch): - monkeypatch.setattr("state.app_settings.sys.platform", "darwin") + def test_explicit_false_overrides_default(self): assert resolved_use_conv_vae(AppSettings(use_conv_vae=False)) is False diff --git a/backend/tests/test_state_actions.py b/backend/tests/test_state_actions.py index dfc0992f1..a2bc1e685 100644 --- a/backend/tests/test_state_actions.py +++ b/backend/tests/test_state_actions.py @@ -237,6 +237,25 @@ def test_retake_pipeline_eviction(test_state, create_fake_model_files): assert isinstance(test_state.state.gpu_slot.active_pipeline, VideoPipelineState) +def test_image_pipeline_freed_not_parked_when_loading_video(test_state, create_fake_model_files): + """FORK behavior (see FORK.md): loading a video pipeline FREES a resident image + pipeline outright rather than parking it in host RAM. Parking an image model + through a video generation is dead weight that thrashes a 64 GB machine + (bf16-read + fp8-pin working set exceeds RAM). A regression to upstream's + parking would leave ``cpu_slot`` populated here instead of None. + """ + create_fake_model_files(include_zit=True) + image_pipeline = test_state.pipelines.load_image_generation_pipeline_to_gpu("z-image-turbo") + assert isinstance(test_state.state.gpu_slot, GpuSlot) + assert test_state.state.gpu_slot.active_pipeline is image_pipeline + assert test_state.state.cpu_slot is None + + # Swapping to a video pipeline must free the image pipeline, not park it. + test_state.pipelines.load_gpu_pipeline("fast") + assert isinstance(test_state.state.gpu_slot.active_pipeline, VideoPipelineState) + assert test_state.state.cpu_slot is None + + def test_ic_lora_load_includes_depth_resources(test_state, fake_services, create_fake_model_files, create_fake_ic_lora_files): create_fake_model_files(model_id=_IC_LORA_MODEL_ID) create_fake_ic_lora_files() diff --git a/backend/uv.lock b/backend/uv.lock index c0d283408..cb6bc3a24 100644 --- a/backend/uv.lock +++ b/backend/uv.lock @@ -114,6 +114,23 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/83/41/7f13361db54d7e02f11552575c0384dadaf0918138f4eaa82ea03a9f9580/av-16.1.0-cp314-cp314t-win_amd64.whl", hash = "sha256:6f90dc082ff2068ddbe77618400b44d698d25d9c4edac57459e250c16b33d700", size = 31948164, upload-time = "2026-01-11T09:59:19.501Z" }, ] +[[package]] +name = "bitsandbytes" +version = "0.50.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "numpy", marker = "sys_platform != 'darwin'" }, + { name = "packaging", marker = "sys_platform != 'darwin'" }, + { name = "torch", version = "2.10.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" }, marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "torch", version = "2.12.1", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/88/d5/b2cb5b5a9daf7349a02b1af2c49b6a044fda2702c9cc5dc296f648358327/bitsandbytes-0.50.2-py3-none-manylinux_2_24_aarch64.whl", hash = "sha256:d5772560dd94c4d9c57f50c9b017450a1707f7687bfd4b3dc86f7342aafe721e", size = 23777467, upload-time = "2026-08-27T00:10:50.92Z" }, + { url = "https://files.pythonhosted.org/packages/a5/6e/e4e8b75716dbe5e50964f070266e06f4e6806ce051bfb97f52ee162b9310/bitsandbytes-0.50.2-py3-none-manylinux_2_24_x86_64.whl", hash = "sha256:55348a9a4a21bfd99cf8c7b32fe67b4030ae5c2a05738e03c1747f65fa6ec283", size = 43139553, upload-time = "2026-08-27T00:10:54.751Z" }, + { url = "https://files.pythonhosted.org/packages/72/82/742dc27a1feab90c8f87f2ed14e6d72d05f9e1cf764b4d2ba30aa9b4a2cb/bitsandbytes-0.50.2-py3-none-win_amd64.whl", hash = "sha256:c697963c8fda3dcd0d7ebd9b5211ae4067feef7cd06e0350d4e816a434fe683d", size = 39096375, upload-time = "2026-08-27T00:10:58.297Z" }, + { url = "https://files.pythonhosted.org/packages/a2/57/61636c5b11b0a32e505127a6dce6fa8fcbf73978babe8fa37082ab547f1c/bitsandbytes-0.50.2-py3-none-win_arm64.whl", hash = "sha256:8437ab68a04ea56daf1d6ecb54230fb1d88be4b89fe2d79bc399bc0203b487cf", size = 1058684, upload-time = "2026-08-27T00:11:00.664Z" }, +] + [[package]] name = "certifi" version = "2026.1.4" @@ -215,7 +232,7 @@ name = "cuda-bindings" version = "12.9.4" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "cuda-pathfinder" }, + { name = "cuda-pathfinder", marker = "sys_platform == 'linux'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/0c/c2/65bfd79292b8ff18be4dd7f7442cea37bcbc1a228c1886f1dea515c45b67/cuda_bindings-12.9.4-cp312-cp312-manylinux_2_24_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:694ba35023846625ef471257e6b5a4bc8af690f961d197d77d34b1d1db393f56", size = 11760260, upload-time = "2025-10-21T14:51:40.79Z" }, @@ -261,8 +278,8 @@ wheels = [ [[package]] name = "diffusers" -version = "0.39.0" -source = { registry = "https://pypi.org/simple" } +version = "0.39.0.dev0" +source = { git = "https://github.com/huggingface/diffusers.git?rev=0d32f8054438fd38204ae2d46155d2c971c90da5#0d32f8054438fd38204ae2d46155d2c971c90da5" } dependencies = [ { name = "filelock" }, { name = "httpx" }, @@ -274,10 +291,6 @@ dependencies = [ { name = "requests" }, { name = "safetensors" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/1a/81/6095237b86a3116c4789f28c4435d5296c00c0fc74ffde99008fd6b3a36c/diffusers-0.39.0.tar.gz", hash = "sha256:14bb1d98c85a0e463d734c99aaa73b480a7bc9bad22af30fbf730ef8f09c1d67", size = 4651240, upload-time = "2026-07-03T08:48:47.904Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/3f/3f/7469c46e9d22307ea686bab687d70e6bf328722952f9d10339f5e913e608/diffusers-0.39.0-py3-none-any.whl", hash = "sha256:912aca51b5787365110806e984d5555735bf8a461073bb8459029d0bca7870ef", size = 5631176, upload-time = "2026-07-03T08:48:45.337Z" }, -] [[package]] name = "einops" @@ -334,6 +347,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/ab/6e/81d47999aebc1b155f81eca4477a616a70f238a2549848c38983f3c22a82/ftfy-6.3.1-py3-none-any.whl", hash = "sha256:7c70eb532015cd2f9adb53f101fb6c7945988d023a085d127d1573dc49dd0083", size = 44821, upload-time = "2024-10-26T00:50:33.425Z" }, ] +[[package]] +name = "gguf" +version = "0.19.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "numpy" }, + { name = "pyyaml" }, + { name = "requests" }, + { name = "tqdm" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/48/ae/17f1308ae45cd7b08ebb521747d5b23f4efc4d172038a4e228dd5106c3ff/gguf-0.19.0.tar.gz", hash = "sha256:dbadcd6cc7ccd44256f2229fe7c2dff5e8aa5cf0612ab987fd2b1a57e428923f", size = 111220, upload-time = "2026-05-06T13:04:03.667Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/b3/bb/d71d6da82763528c2c2ed6b59a9d6142c6595545a4c448e2085d155e88c2/gguf-0.19.0-py3-none-any.whl", hash = "sha256:70bcd10edfe697fb2dad6e40af2234b9d8ece9a41a99761405121ebda1c3c1cd", size = 118475, upload-time = "2026-05-06T13:04:02.588Z" }, +] + [[package]] name = "h11" version = "0.16.0" @@ -538,9 +566,11 @@ name = "ltx-desktop-backend" version = "1.0.0" source = { virtual = "." } dependencies = [ + { name = "bitsandbytes", marker = "sys_platform != 'darwin'" }, { name = "diffusers" }, { name = "fastapi" }, { name = "ftfy" }, + { name = "gguf" }, { name = "huggingface-hub" }, { name = "imageio" }, { name = "imageio-ffmpeg" }, @@ -586,10 +616,12 @@ test = [ [package.metadata] requires-dist = [ + { name = "bitsandbytes", marker = "sys_platform != 'darwin'", specifier = ">=0.48.0" }, { name = "debugpy", marker = "extra == 'dev'", specifier = ">=1.8" }, - { name = "diffusers", specifier = ">=0.39.0" }, + { name = "diffusers", git = "https://github.com/huggingface/diffusers.git?rev=0d32f8054438fd38204ae2d46155d2c971c90da5" }, { name = "fastapi", specifier = ">=0.141.1" }, { name = "ftfy", specifier = ">=6.0.0" }, + { name = "gguf", specifier = ">=0.10.0" }, { name = "httpx", marker = "extra == 'test'", specifier = ">=0.27" }, { name = "huggingface-hub", specifier = ">=0.23.0" }, { name = "imageio", specifier = ">=2.37.2" }, @@ -737,13 +769,13 @@ name = "mps-sdpa" version = "0.2.0" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "ninja" }, - { name = "numpy" }, - { name = "psutil" }, - { name = "pyobjc-core" }, - { name = "pyobjc-framework-metal" }, - { name = "pyobjc-framework-metalperformanceshadersgraph" }, - { name = "torch", version = "2.11.0", source = { registry = "https://pypi.org/simple" } }, + { name = "ninja", marker = "sys_platform == 'darwin'" }, + { name = "numpy", marker = "sys_platform == 'darwin'" }, + { name = "psutil", marker = "sys_platform == 'darwin'" }, + { name = "pyobjc-core", marker = "sys_platform == 'darwin'" }, + { name = "pyobjc-framework-metal", marker = "sys_platform == 'darwin'" }, + { name = "pyobjc-framework-metalperformanceshadersgraph", marker = "sys_platform == 'darwin'" }, + { name = "torch", version = "2.11.0", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform == 'darwin'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/f0/c2/cf52e4364c4a5f9db84238a34974a9498e5922720349ec6f6a3700cd492c/mps_sdpa-0.2.0.tar.gz", hash = "sha256:1d6b6024801657bbdc4a6420ec7a0b52191ab0b88d847143a3ea1728ba3f93b5", size = 81037, upload-time = "2026-04-30T16:12:53.581Z" } wheels = [ @@ -887,7 +919,7 @@ name = "nvidia-cudnn-cu12" version = "9.10.2.21" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cublas-cu12" }, + { name = "nvidia-cublas-cu12", marker = "sys_platform == 'linux'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/fa/41/e79269ce215c857c935fd86bcfe91a451a584dfc27f1e068f568b9ad1ab7/nvidia_cudnn_cu12-9.10.2.21-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:c9132cc3f8958447b4910a1720036d9eff5928cc3179b0a51fb6d167c6cc87d8", size = 705026878, upload-time = "2025-06-06T21:52:51.348Z" }, @@ -899,7 +931,7 @@ name = "nvidia-cufft-cu12" version = "11.3.3.83" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-nvjitlink-cu12" }, + { name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/60/bc/7771846d3a0272026c416fbb7e5f4c1f146d6d80704534d0b187dd6f4800/nvidia_cufft_cu12-11.3.3.83-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:848ef7224d6305cdb2a4df928759dca7b1201874787083b6e7550dd6765ce69a", size = 193109211, upload-time = "2025-03-07T01:44:56.873Z" }, @@ -929,9 +961,9 @@ name = "nvidia-cusolver-cu12" version = "11.7.3.90" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-cublas-cu12" }, - { name = "nvidia-cusparse-cu12" }, - { name = "nvidia-nvjitlink-cu12" }, + { name = "nvidia-cublas-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cusparse-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/c8/32/f7cd6ce8a7690544d084ea21c26e910a97e077c9b7f07bf5de623ee19981/nvidia_cusolver_cu12-11.7.3.90-py3-none-manylinux_2_27_aarch64.whl", hash = "sha256:db9ed69dbef9715071232caa9b69c52ac7de3a95773c2db65bdba85916e4e5c0", size = 267229841, upload-time = "2025-03-07T01:46:54.356Z" }, @@ -943,7 +975,7 @@ name = "nvidia-cusparse-cu12" version = "12.5.8.93" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-nvjitlink-cu12" }, + { name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/bc/f7/cd777c4109681367721b00a106f491e0d0d15cfa1fd59672ce580ce42a97/nvidia_cusparse_cu12-12.5.8.93-py3-none-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:9b6c161cb130be1a07a27ea6923df8141f3c295852f4b260c65f18f3e0a091dc", size = 288117129, upload-time = "2025-03-07T01:47:40.407Z" }, @@ -1303,7 +1335,7 @@ name = "pynvml" version = "13.0.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "nvidia-ml-py" }, + { name = "nvidia-ml-py", marker = "sys_platform != 'darwin'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/5c/57/da7dc63a79f59e082e26a66ac02d87d69ea316b35b35b7a00d82f3ce3d2f/pynvml-13.0.1.tar.gz", hash = "sha256:1245991d9db786b4d2f277ce66869bd58f38ac654e38c9397d18f243c8f6e48f", size = 35226, upload-time = "2025-09-05T20:33:25.377Z" } wheels = [ @@ -1330,7 +1362,7 @@ name = "pyobjc-framework-cocoa" version = "12.2.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "pyobjc-core" }, + { name = "pyobjc-core", marker = "sys_platform == 'darwin'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/51/34/fbe38a204643aa4e1b91391cdce07a34da565a69171ebcad08de7438a556/pyobjc_framework_cocoa-12.2.1.tar.gz", hash = "sha256:b94b37fe5730e5ae1fb0052912cd174e6ec329b0bfba4a012ae5db1014b5864b", size = 3125751, upload-time = "2026-06-19T16:20:05.159Z" } wheels = [ @@ -1348,8 +1380,8 @@ name = "pyobjc-framework-metal" version = "12.2.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "pyobjc-core" }, - { name = "pyobjc-framework-cocoa" }, + { name = "pyobjc-core", marker = "sys_platform == 'darwin'" }, + { name = "pyobjc-framework-cocoa", marker = "sys_platform == 'darwin'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/15/46/5920d6cb66cbbe298744889b10b3266b1408ad823855f55cdcb967c0d51d/pyobjc_framework_metal-12.2.1.tar.gz", hash = "sha256:cd362194bdb7fd2a9116b8dc1e6b14ce19629136304cdf6b88d105a969fda72c", size = 238139, upload-time = "2026-06-19T16:21:06.897Z" } wheels = [ @@ -1367,8 +1399,8 @@ name = "pyobjc-framework-metalperformanceshaders" version = "12.2.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "pyobjc-core" }, - { name = "pyobjc-framework-metal" }, + { name = "pyobjc-core", marker = "sys_platform == 'darwin'" }, + { name = "pyobjc-framework-metal", marker = "sys_platform == 'darwin'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/50/ff/6938291dd5a71e39f6948037dbb271993d86a3ecd7706e7cc38034feeaea/pyobjc_framework_metalperformanceshaders-12.2.1.tar.gz", hash = "sha256:a4395f8619ad6f1d382aab5cf116e058b18d3646bec6b730c77daa8f692b5de4", size = 190474, upload-time = "2026-06-19T16:21:09.743Z" } wheels = [ @@ -1386,8 +1418,8 @@ name = "pyobjc-framework-metalperformanceshadersgraph" version = "12.2.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "pyobjc-core" }, - { name = "pyobjc-framework-metalperformanceshaders" }, + { name = "pyobjc-core", marker = "sys_platform == 'darwin'" }, + { name = "pyobjc-framework-metalperformanceshaders", marker = "sys_platform == 'darwin'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/da/11/15e3acf0636f9418384cd3e213296ad9882f546889865913cfaee2339ed6/pyobjc_framework_metalperformanceshadersgraph-12.2.1.tar.gz", hash = "sha256:656e70c86645814ef1d02bac74933eebc5ff6427100bee4a3bbffde921020ab6", size = 60199, upload-time = "2026-06-19T16:21:10.716Z" } wheels = [ @@ -1855,30 +1887,30 @@ resolution-markers = [ "sys_platform == 'linux'", ] dependencies = [ - { name = "cuda-bindings", marker = "sys_platform != 'win32'" }, - { name = "filelock" }, - { name = "fsspec" }, - { name = "jinja2" }, - { name = "networkx" }, - { name = "nvidia-cublas-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cuda-cupti-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cuda-nvrtc-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cuda-runtime-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cudnn-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cufft-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cufile-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-curand-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cusolver-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cusparse-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-cusparselt-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-nccl-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-nvjitlink-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-nvshmem-cu12", marker = "sys_platform != 'win32'" }, - { name = "nvidia-nvtx-cu12", marker = "sys_platform != 'win32'" }, - { name = "setuptools", version = "82.0.0", source = { registry = "https://pypi.org/simple" } }, - { name = "sympy" }, - { name = "triton", marker = "sys_platform != 'win32'" }, - { name = "typing-extensions" }, + { name = "cuda-bindings", marker = "sys_platform == 'linux'" }, + { name = "filelock", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "fsspec", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "jinja2", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "networkx", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "nvidia-cublas-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cuda-cupti-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cuda-nvrtc-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cuda-runtime-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cudnn-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cufft-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cufile-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-curand-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cusolver-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cusparse-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-cusparselt-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-nccl-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-nvjitlink-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-nvshmem-cu12", marker = "sys_platform == 'linux'" }, + { name = "nvidia-nvtx-cu12", marker = "sys_platform == 'linux'" }, + { name = "setuptools", version = "82.0.0", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "sympy", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "triton", marker = "sys_platform == 'linux'" }, + { name = "typing-extensions", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, ] wheels = [ { url = "https://download-r2.pytorch.org/whl/cu128/torch-2.10.0%2Bcu128-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:6f09cdf2415516be028ae82e6b985bcfc3eac37bc52ab401142689f6224516ca", upload-time = "2026-01-21T15:22:03Z" }, @@ -1906,13 +1938,13 @@ resolution-markers = [ "sys_platform == 'darwin'", ] dependencies = [ - { name = "filelock" }, - { name = "fsspec" }, - { name = "jinja2" }, - { name = "networkx" }, - { name = "setuptools", version = "81.0.0", source = { registry = "https://pypi.org/simple" } }, - { name = "sympy" }, - { name = "typing-extensions" }, + { name = "filelock", marker = "sys_platform == 'darwin'" }, + { name = "fsspec", marker = "sys_platform == 'darwin'" }, + { name = "jinja2", marker = "sys_platform == 'darwin'" }, + { name = "networkx", marker = "sys_platform == 'darwin'" }, + { name = "setuptools", version = "81.0.0", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform == 'darwin'" }, + { name = "sympy", marker = "sys_platform == 'darwin'" }, + { name = "typing-extensions", marker = "sys_platform == 'darwin'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/6f/8b/69e3008d78e5cee2b30183340cc425081b78afc5eff3d080daab0adda9aa/torch-2.11.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:4b5866312ee6e52ea625cd211dcb97d6a2cdc1131a5f15cc0d87eec948f6dd34", size = 80606338, upload-time = "2026-03-23T18:11:34.781Z" }, @@ -1930,13 +1962,13 @@ resolution-markers = [ "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'", ] dependencies = [ - { name = "filelock" }, - { name = "fsspec" }, - { name = "jinja2" }, - { name = "networkx" }, - { name = "setuptools", version = "81.0.0", source = { registry = "https://pypi.org/simple" } }, - { name = "sympy" }, - { name = "typing-extensions" }, + { name = "filelock", marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, + { name = "fsspec", marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, + { name = "jinja2", marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, + { name = "networkx", marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, + { name = "setuptools", version = "81.0.0", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, + { name = "sympy", marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, + { name = "typing-extensions", marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, ] [[package]] @@ -1948,7 +1980,7 @@ resolution-markers = [ "sys_platform == 'linux'", ] dependencies = [ - { name = "torch", version = "2.10.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" } }, + { name = "torch", version = "2.10.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" }, marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/ea/3f/df620439a76ece170472d41438d11a1545d5db5dc9f1eaeab8c6e055a328/torchaudio-2.10.0-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:42b148a0921a3721abd1f6ae098b1ec9f89703e555c4f7a0d44da87b8decbcb9", size = 391973, upload-time = "2026-01-21T16:28:39.732Z" }, @@ -1993,9 +2025,9 @@ resolution-markers = [ "sys_platform == 'linux'", ] dependencies = [ - { name = "numpy" }, - { name = "pillow" }, - { name = "torch", version = "2.10.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" } }, + { name = "numpy", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "pillow", marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, + { name = "torch", version = "2.10.0+cu128", source = { registry = "https://download.pytorch.org/whl/cu128" }, marker = "sys_platform == 'linux' or sys_platform == 'win32'" }, ] wheels = [ { url = "https://download-r2.pytorch.org/whl/cu128/torchvision-0.25.0%2Bcu128-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:8623e534ef6a815bd6407d4b52dd70c7154e2eda626ad4b9cb895d36c5a3305b", upload-time = "2026-01-21T22:32:23Z" }, @@ -2023,9 +2055,9 @@ resolution-markers = [ "sys_platform == 'darwin'", ] dependencies = [ - { name = "numpy" }, - { name = "pillow" }, - { name = "torch", version = "2.11.0", source = { registry = "https://pypi.org/simple" } }, + { name = "numpy", marker = "sys_platform == 'darwin'" }, + { name = "pillow", marker = "sys_platform == 'darwin'" }, + { name = "torch", version = "2.11.0", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform == 'darwin'" }, ] wheels = [ { url = "https://files.pythonhosted.org/packages/ae/e7/56b47cc3b132aea90ccce22bcb8975dec688b002150012acc842846039d0/torchvision-0.26.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:c409e1c3fdebec7a3834465086dbda8bf7680eff79abf7fd2f10c6b59520a7a4", size = 1863502, upload-time = "2026-03-23T18:12:57.326Z" }, @@ -2043,9 +2075,9 @@ resolution-markers = [ "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'", ] dependencies = [ - { name = "numpy" }, - { name = "pillow" }, - { name = "torch", version = "2.12.1", source = { registry = "https://pypi.org/simple" } }, + { name = "numpy", marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, + { name = "pillow", marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, + { name = "torch", version = "2.12.1", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform != 'darwin' and sys_platform != 'linux' and sys_platform != 'win32'" }, ] [[package]] diff --git a/electron-builder.yml b/electron-builder.yml index 17b983ae6..ed59dd565 100644 --- a/electron-builder.yml +++ b/electron-builder.yml @@ -1,6 +1,6 @@ -appId: com.lightricks.ltx-desktop -productName: LTX Desktop -copyright: Copyright © 2026 Lightricks +appId: com.michaelricks.rix-desktop-studio-pro +productName: RiX Desktop Studio Pro +copyright: Copyright © 2026 Michael Ricks directories: output: release @@ -39,6 +39,15 @@ extraResources: - "!tests/**" - from: python-deps-hash.txt to: python-deps-hash.txt + # Bundle the app icon into the packaged app's resources dir so the RUNNING + # window/taskbar shows the RiX icon. electron/window.ts loads + # /icon. at runtime (win32 -> icon.ico, else icon.png); + # without this it isn't packaged and the open app falls back to the default + # Electron icon. (The exe/shortcut icon is handled separately by win.icon.) + - from: resources/icon.ico + to: icon.ico + - from: resources/icon.png + to: icon.png win: target: @@ -47,12 +56,9 @@ win: - x64 icon: resources/icon.ico artifactName: ${productName}-Setup.${ext} - azureSignOptions: - publisherName: Lightricks US Inc - endpoint: https://eus.codesigning.azure.net/ - certificateProfileName: ltx-desktop-trusted - codeSigningAccountName: ltx-trusted - excludeEnvironmentCredential: true + # Unsigned for the tester phase — Windows SmartScreen will warn ("More info -> + # Run anyway"). Re-add signing (own cert / Azure Trusted Signing) before a wider + # release. The Lightricks Azure cert config was removed here. nsis: oneClick: false @@ -63,7 +69,7 @@ nsis: installerHeaderIcon: resources/icon.ico createDesktopShortcut: true createStartMenuShortcut: true - shortcutName: LTX Desktop + shortcutName: RiX Desktop Studio Pro mac: hardenedRuntime: true @@ -131,5 +137,5 @@ deb: publish: provider: github - owner: Lightricks - repo: ltx-desktop + owner: MichaelRicks + repo: LTX-Desktop-Mike diff --git a/electron/app-paths.ts b/electron/app-paths.ts index 40a361155..3621e0aaf 100644 --- a/electron/app-paths.ts +++ b/electron/app-paths.ts @@ -2,7 +2,15 @@ import { app } from 'electron' import path from 'path' import os from 'os' -export const APP_FOLDER_NAME = 'LTXDesktop' +// Dev runs use a separate folder so the fork never collides with an installed +// LTX Desktop (independent single-instance lock + isolated settings). The dev +// userData's `models/` is junctioned to the installed app's models, so nothing +// re-downloads. Packaged builds keep the original folder name. +export const APP_FOLDER_NAME = app.isPackaged ? 'LTXDesktop' : 'LTXDesktopMikeDev' + +if (!app.isPackaged) { + app.setName('LTX Desktop Studio Pro') +} function resolveUserDataPath(): string { if (process.platform === 'win32') { diff --git a/electron/config.ts b/electron/config.ts index f636f7b7a..36573360b 100644 --- a/electron/config.ts +++ b/electron/config.ts @@ -2,6 +2,7 @@ import { app } from 'electron' import path from 'path' import os from 'os' import { getProjectAssetsPath } from './app-state' +import { resolveLibraryRoot } from './gpm-library-root' export const isDev = !app.isPackaged @@ -26,5 +27,6 @@ export function getAllowedRoots(): string[] { roots.push(process.resourcesPath) } roots.push(getProjectAssetsPath()) + roots.push(resolveLibraryRoot()) return roots } diff --git a/electron/csp.ts b/electron/csp.ts index 9d83969e6..0022ee2cb 100644 --- a/electron/csp.ts +++ b/electron/csp.ts @@ -15,11 +15,13 @@ export function setupCSP(): void { "font-src 'self' https://fonts.gstatic.com", "connect-src 'self' http://localhost:* http://127.0.0.1:* ws://localhost:* ws://127.0.0.1:*", "img-src 'self' data: blob: file: https://storage.googleapis.com", - "media-src 'self' blob: file: https://videos.ltx.io https://storage.googleapis.com", + // 'data:' in media-src is ours (Prompt Manager Pro serves media as data URLs); + // the remote hosts are upstream's, for LoRA-catalog preview media. + "media-src 'self' blob: file: data: https://videos.ltx.io https://storage.googleapis.com", "object-src 'none'", "base-uri 'self'", "form-action 'self'", - "frame-ancestors 'none'", + "frame-ancestors 'self'", ].join('; ') : [ "default-src 'self'", @@ -28,11 +30,13 @@ export function setupCSP(): void { "font-src 'self' https://fonts.gstatic.com", "connect-src 'self' http://localhost:* http://127.0.0.1:* ws://localhost:* ws://127.0.0.1:*", "img-src 'self' data: blob: file: https://storage.googleapis.com", - "media-src 'self' blob: file: https://videos.ltx.io https://storage.googleapis.com", + // 'data:' in media-src is ours (Prompt Manager Pro serves media as data URLs); + // the remote hosts are upstream's, for LoRA-catalog preview media. + "media-src 'self' blob: file: data: https://videos.ltx.io https://storage.googleapis.com", "object-src 'none'", "base-uri 'self'", "form-action 'self'", - "frame-ancestors 'none'", + "frame-ancestors 'self'", ].join('; ') callback({ diff --git a/electron/export/audio-mix.ts b/electron/export/audio-mix.ts index 6493cf398..16e0d7cd6 100644 --- a/electron/export/audio-mix.ts +++ b/electron/export/audio-mix.ts @@ -52,6 +52,24 @@ function extractPcmBuffer( interface AudioSource { filePath: string; trimStart: number; trimEnd: number; timelineStart: number; speed: number; reversed: boolean; volume: number; + audioFadeIn: number; audioFadeOut: number; + volumeKeyframes?: { t: number; value: number }[]; +} + +/** Sample a pre-sorted piecewise-linear volume envelope at time `t` (seconds + * from clip start); clamps to the first/last keyframe outside the range. */ +function sampleVolumeEnvelope(ks: { t: number; value: number }[], t: number): number { + if (t <= ks[0].t) return ks[0].value + const last = ks[ks.length - 1] + if (t >= last.t) return last.value + for (let i = 0; i < ks.length - 1; i++) { + const a = ks[i], b = ks[i + 1] + if (t >= a.t && t <= b.t) { + const span = b.t - a.t + return span <= 0 ? b.value : a.value + (b.value - a.value) * ((t - a.t) / span) + } + } + return last.value } /** @@ -68,7 +86,9 @@ export async function mixAudioToPcm( const audioSources: AudioSource[] = [] for (const c of clips) { - if (c.muted || c.volume <= 0) continue + const hasKeyframes = !!(c.volumeKeyframes && c.volumeKeyframes.length > 0) + // A keyframed clip can be audible even if its flat volume is 0. + if (c.muted || (c.volume <= 0 && !hasKeyframes)) continue const fp = c.path if (!fp || !fs.existsSync(fp)) continue @@ -81,6 +101,9 @@ export async function mixAudioToPcm( speed: c.speed, reversed: c.reversed, volume: c.volume, + audioFadeIn: c.audioFadeIn ?? 0, + audioFadeOut: c.audioFadeOut ?? 0, + volumeKeyframes: c.volumeKeyframes, }) } else if (c.type === 'video') { if (!audioProbeCache.has(fp)) { @@ -95,6 +118,9 @@ export async function mixAudioToPcm( speed: c.speed, reversed: c.reversed, volume: c.volume, + audioFadeIn: c.audioFadeIn ?? 0, + audioFadeOut: c.audioFadeOut ?? 0, + volumeKeyframes: c.volumeKeyframes, }) } } @@ -116,11 +142,37 @@ export async function mixAudioToPcm( const startSample = startFrame * NUM_CHANNELS const numPcmSamples = Math.floor(pcm.length / BYTES_PER_SAMPLE) + // Linear fade in/out envelope, in frames (a frame = NUM_CHANNELS samples). + const clipFrames = Math.floor(numPcmSamples / NUM_CHANNELS) + const fadeInFrames = Math.max(0, Math.round((src.audioFadeIn || 0) * SAMPLE_RATE)) + const fadeOutFrames = Math.max(0, Math.round((src.audioFadeOut || 0) * SAMPLE_RATE)) + const hasFade = fadeInFrames > 0 || fadeOutFrames > 0 + + // Volume automation envelope (pre-sorted once); sampled per frame below. + const kfs = (src.volumeKeyframes && src.volumeKeyframes.length > 0) + ? [...src.volumeKeyframes].sort((a, b) => a.t - b.t) + : null + let baseGain = src.volume + for (let s = 0; s < numPcmSamples; s++) { const destIdx = startSample + s if (destIdx < 0 || destIdx >= totalSamples) continue + const frame = (s / NUM_CHANNELS) | 0 + // Recompute the automated base gain once per stereo frame (channel 0). + if (kfs && (s % NUM_CHANNELS) === 0) { + baseGain = sampleVolumeEnvelope(kfs, frame / SAMPLE_RATE) + } const value = pcm.readInt16LE(s * BYTES_PER_SAMPLE) - mixBuffer[destIdx] += value * src.volume + let gain = kfs ? baseGain : src.volume + if (hasFade) { + if (fadeInFrames > 0 && frame < fadeInFrames) { + gain *= frame / fadeInFrames + } + if (fadeOutFrames > 0 && frame >= clipFrames - fadeOutFrames) { + gain *= Math.max(0, (clipFrames - frame) / fadeOutFrames) + } + } + mixBuffer[destIdx] += value * gain } logger.info( `[Export] Audio ${i + 1}: mixed ${numPcmSamples} samples (${(numPcmSamples / SAMPLE_RATE / NUM_CHANNELS).toFixed(2)}s) at offset frame ${startFrame}`) } catch (err: any) { diff --git a/electron/export/export-handler.ts b/electron/export/export-handler.ts index 5a9014984..5a446aeab 100644 --- a/electron/export/export-handler.ts +++ b/electron/export/export-handler.ts @@ -2,118 +2,183 @@ import path from 'path' import fs from 'fs' import os from 'os' import { getAllowedRoots } from '../config' +import { getMainWindow } from '../window' import { logger } from '../logger' import { validatePath } from '../path-validation' import { findFfmpegPath, runFfmpeg, stopExportProcess } from './ffmpeg-utils' -import { flattenTimeline } from './timeline' +import { buildDissolveTimeRemap, computeFinalVideoDuration, flattenTimeline } from './timeline' import { buildVideoFilterGraph } from './video-filter' import { mixAudioToPcm } from './audio-mix' import { handle } from '../ipc/typed-handle' +import type { z } from 'zod' +import type { electronAPISchemas } from '../../shared/electron-api-schema' + +/** First existing system SANS font, so exported text overlays match the editor + * preview's sans look instead of ffmpeg's default serif. Arial first (it's the + * preview font stack's fallback). Returns undefined if none found — drawtext + * then omits fontfile and falls back to ffmpeg's default (still renders). */ +function resolveExportFont(): string | undefined { + const candidates = process.platform === 'win32' + ? ['C:/Windows/Fonts/arial.ttf', 'C:/Windows/Fonts/segoeui.ttf', 'C:/Windows/Fonts/calibri.ttf', 'C:/Windows/Fonts/tahoma.ttf'] + : process.platform === 'darwin' + ? ['/System/Library/Fonts/Supplemental/Arial.ttf', '/Library/Fonts/Arial.ttf', '/System/Library/Fonts/Helvetica.ttc'] + : ['/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf', '/usr/share/fonts/truetype/liberation/LiberationSans-Regular.ttf'] + for (const f of candidates) { + try { if (fs.existsSync(f)) return f } catch { /* ignore */ } + } + return undefined +} -export function registerExportHandlers(): void { - handle('exportNative', async ({ clips, outputPath, codec, width, height, fps, quality, letterbox, subtitles }) => { - const ffmpegPath = findFfmpegPath() - if (!ffmpegPath) return { success: false, error: 'FFmpeg not found' } - - try { - validatePath(outputPath, getAllowedRoots()) - for (const clip of clips) { - const fp = clip.path - if (fp) validatePath(fp, getAllowedRoots()) - } - } catch (err) { - return { success: false, error: String(err) } +export type ExportNativeInput = z.infer +export type ExportNativeResult = { success: true } | { success: false; error: string } + +/** Render the timeline to a file. `preview` trades quality for speed (ultrafast, + * high CRF) — used by the RiX MCP server to render frames for review. */ +export async function exportTimelineNative( + { clips, outputPath, codec, width, height, fps, quality, letterbox, subtitles, textOverlays }: ExportNativeInput, + { preview = false }: { preview?: boolean } = {}, +): Promise { + const ffmpegPath = findFfmpegPath() + if (!ffmpegPath) return { success: false, error: 'FFmpeg not found' } + + try { + validatePath(outputPath, getAllowedRoots()) + for (const clip of clips) { + const fp = clip.path + if (fp) validatePath(fp, getAllowedRoots()) } + } catch (err) { + return { success: false, error: String(err) } + } - const segments = flattenTimeline(clips) - if (segments.length === 0) return { success: false, error: 'No clips to export' } + const segments = flattenTimeline(clips) + if (segments.length === 0) return { success: false, error: 'No clips to export' } - for (const seg of segments) { - if (seg.filePath && !fs.existsSync(seg.filePath)) { - return { success: false, error: `Source file not found: ${path.basename(seg.filePath)}` } - } + for (const seg of segments) { + if (seg.filePath && !fs.existsSync(seg.filePath)) { + return { success: false, error: `Source file not found: ${path.basename(seg.filePath)}` } + } + } + + // Total program duration drives the progress percentage: ffmpeg reports the + // encoded position (`time=`), which we divide by this to get a fraction. + const totalDur = computeFinalVideoDuration(segments) + const emitProgress = (percent: number, stage: string) => { + getMainWindow()?.webContents.send('export-progress', { + percent: Math.max(0, Math.min(100, Math.round(percent))), + stage, + }) + } + // The video encode (step 1) is by far the longest, so it owns most of the + // bar; audio + mux share the tail. (An h264 mux is a stream copy and flies.) + const VIDEO_SHARE = 85 + + const tmpDir = os.tmpdir() + const ts = Date.now() + const tmpVideo = path.join(tmpDir, `ltx-export-video-${ts}.mkv`) + const tmpAudio = path.join(tmpDir, `ltx-export-audio-${ts}.wav`) + const cleanup = () => { + try { fs.unlinkSync(tmpVideo) } catch {} + try { fs.unlinkSync(tmpAudio) } catch {} + } + + try { + logger.info( `[Export] Step 1: Video-only export (${segments.length} segments)`) + { + const fontFile = resolveExportFont() + const { inputs, filterScript } = buildVideoFilterGraph(segments, { width, height, fps, letterbox, subtitles, textOverlays, fontFile }) + + const filterFile = path.join(tmpDir, `ltx-filter-v-${ts}.txt`) + fs.writeFileSync(filterFile, filterScript, 'utf8') + + emitProgress(0, 'Encoding video') + const r = await runFfmpeg(ffmpegPath, [ + '-y', ...inputs, '-filter_complex_script', filterFile, + '-map', '[outv]', '-an', '-c:v', 'libx264', + ...(preview ? ['-preset', 'ultrafast', '-crf', '28'] : ['-preset', 'fast', '-crf', '16']), + '-pix_fmt', 'yuv420p', tmpVideo + ], (t) => { + const frac = totalDur > 0 ? t / totalDur : 0 + emitProgress(frac * VIDEO_SHARE, 'Encoding video') + }) + try { fs.unlinkSync(filterFile) } catch {} + if (!r.success) { cleanup(); return { success: false, error: r.error } } } - const tmpDir = os.tmpdir() - const ts = Date.now() - const tmpVideo = path.join(tmpDir, `ltx-export-video-${ts}.mkv`) - const tmpAudio = path.join(tmpDir, `ltx-export-audio-${ts}.wav`) - const cleanup = () => { - try { fs.unlinkSync(tmpVideo) } catch {} - try { fs.unlinkSync(tmpAudio) } catch {} + emitProgress(VIDEO_SHARE, 'Mixing audio') + logger.info( '[Export] Step 2: Audio mixdown (PCM buffer approach)') + // A dissolve overlaps two clips, shrinking the program's real duration + // below the naive sum of clip lengths (see buildDissolveTimeRemap) - + // every clip's nominal startTime needs the same conversion applied to + // video, or its audio drifts later relative to the picture with every + // dissolve that came before it. + const remapTime = buildDissolveTimeRemap(segments) + const remappedClips = clips.map(c => ({ ...c, startTime: remapTime(c.startTime) })) + + let totalDuration = computeFinalVideoDuration(segments) + for (const c of remappedClips) { + totalDuration = Math.max(totalDuration, c.startTime + c.duration) } - try { - logger.info( `[Export] Step 1: Video-only export (${segments.length} segments)`) - { - const { inputs, filterScript } = buildVideoFilterGraph(segments, { width, height, fps, letterbox, subtitles }) - - const filterFile = path.join(tmpDir, `ltx-filter-v-${ts}.txt`) - fs.writeFileSync(filterFile, filterScript, 'utf8') - - const r = await runFfmpeg(ffmpegPath, [ - '-y', ...inputs, '-filter_complex_script', filterFile, - '-map', '[outv]', '-an', '-c:v', 'libx264', '-preset', 'fast', '-crf', '16', '-pix_fmt', 'yuv420p', tmpVideo - ]) - try { fs.unlinkSync(filterFile) } catch {} - if (!r.success) { cleanup(); return { success: false, error: r.error } } - } - - logger.info( '[Export] Step 2: Audio mixdown (PCM buffer approach)') - let totalDuration = segments.reduce((max, s) => Math.max(max, s.startTime + s.duration), 0) - for (const c of clips) { - totalDuration = Math.max(totalDuration, c.startTime + c.duration) - } - - const { pcmBuffer, sampleRate, channels: audioChannels } = await mixAudioToPcm(clips, totalDuration, ffmpegPath) - - const tmpRawPcm = path.join(tmpDir, `ltx-pcm-${ts}.raw`) - fs.writeFileSync(tmpRawPcm, pcmBuffer) - logger.info( `[Export] Wrote raw PCM: ${pcmBuffer.length} bytes (${totalDuration.toFixed(2)}s)`) - - { - const r = await runFfmpeg(ffmpegPath, [ - '-y', '-f', 's16le', '-ar', String(sampleRate), '-ac', String(audioChannels), - '-i', tmpRawPcm, '-c:a', 'pcm_s16le', tmpAudio, - ]) - try { fs.unlinkSync(tmpRawPcm) } catch {} - if (!r.success) { cleanup(); return { success: false, error: r.error } } - } - - logger.info( '[Export] Step 3: Combining video + audio') - let videoCodecArgs: string[] - let audioCodecArgs: string[] - if (codec === 'h264') { - videoCodecArgs = ['-c:v', 'libx264', '-preset', 'medium', '-crf', String(quality || 18), '-pix_fmt', 'yuv420p', '-movflags', '+faststart'] - audioCodecArgs = ['-c:a', 'aac', '-b:a', '192k'] - } else if (codec === 'prores') { - videoCodecArgs = ['-c:v', 'prores_ks', '-profile:v', String(quality || 3), '-pix_fmt', 'yuva444p10le'] - audioCodecArgs = ['-c:a', 'pcm_s16le'] - } else if (codec === 'vp9') { - videoCodecArgs = ['-c:v', 'libvpx-vp9', '-b:v', `${quality || 8}M`, '-pix_fmt', 'yuv420p'] - audioCodecArgs = ['-c:a', 'libopus', '-b:a', '128k'] - } else { - cleanup() - return { success: false, error: `Unknown codec: ${codec}` } - } - - const canCopyVideo = codec === 'h264' + const { pcmBuffer, sampleRate, channels: audioChannels } = await mixAudioToPcm(remappedClips, totalDuration, ffmpegPath) + + const tmpRawPcm = path.join(tmpDir, `ltx-pcm-${ts}.raw`) + fs.writeFileSync(tmpRawPcm, pcmBuffer) + logger.info( `[Export] Wrote raw PCM: ${pcmBuffer.length} bytes (${totalDuration.toFixed(2)}s)`) + + { const r = await runFfmpeg(ffmpegPath, [ - '-y', '-i', tmpVideo, '-i', tmpAudio, - '-map', '0:v', '-map', '1:a', - ...(canCopyVideo ? ['-c:v', 'copy'] : videoCodecArgs), - ...audioCodecArgs, '-shortest', outputPath + '-y', '-f', 's16le', '-ar', String(sampleRate), '-ac', String(audioChannels), + '-i', tmpRawPcm, '-c:a', 'pcm_s16le', tmpAudio, ]) + try { fs.unlinkSync(tmpRawPcm) } catch {} + if (!r.success) { cleanup(); return { success: false, error: r.error } } + } + emitProgress(90, 'Finalizing') + logger.info( '[Export] Step 3: Combining video + audio') + let videoCodecArgs: string[] + let audioCodecArgs: string[] + if (codec === 'h264') { + videoCodecArgs = ['-c:v', 'libx264', '-preset', 'medium', '-crf', String(quality || 18), '-pix_fmt', 'yuv420p', '-movflags', '+faststart'] + audioCodecArgs = ['-c:a', 'aac', '-b:a', '192k'] + } else if (codec === 'prores') { + videoCodecArgs = ['-c:v', 'prores_ks', '-profile:v', String(quality || 3), '-pix_fmt', 'yuva444p10le'] + audioCodecArgs = ['-c:a', 'pcm_s16le'] + } else if (codec === 'vp9') { + videoCodecArgs = ['-c:v', 'libvpx-vp9', '-b:v', `${quality || 8}M`, '-pix_fmt', 'yuv420p'] + audioCodecArgs = ['-c:a', 'libopus', '-b:a', '128k'] + } else { cleanup() - if (!r.success) return { success: false, error: r.error } - logger.info( `[Export] Done: ${outputPath}`) - return { success: true } - } catch (err) { - cleanup() - return { success: false, error: String(err) } + return { success: false, error: `Unknown codec: ${codec}` } } - }) + + const canCopyVideo = codec === 'h264' + const r = await runFfmpeg(ffmpegPath, [ + '-y', '-i', tmpVideo, '-i', tmpAudio, + '-map', '0:v', '-map', '1:a', + ...(canCopyVideo ? ['-c:v', 'copy'] : videoCodecArgs), + ...audioCodecArgs, '-shortest', outputPath + ], (t) => { + // Re-encoding codecs (ProRes/VP9) spend real time here; map it to the + // last 10%. h264 stream-copies and finishes near-instantly. + const frac = totalDur > 0 ? t / totalDur : 0 + emitProgress(90 + frac * 10, 'Finalizing') + }) + + cleanup() + if (!r.success) return { success: false, error: r.error } + emitProgress(100, 'Done') + logger.info( `[Export] Done: ${outputPath}`) + return { success: true } + } catch (err) { + cleanup() + return { success: false, error: String(err) } + } +} + +export function registerExportHandlers(): void { + handle('exportNative', (input) => exportTimelineNative(input)) handle('exportCancel', () => { stopExportProcess() diff --git a/electron/export/ffmpeg-utils.ts b/electron/export/ffmpeg-utils.ts index 4c5d4dd61..7cd86cbcb 100644 --- a/electron/export/ffmpeg-utils.ts +++ b/electron/export/ffmpeg-utils.ts @@ -53,8 +53,14 @@ export function fileHasAudio(ffmpegPath: string, filePath: string): boolean { } -/** Run an ffmpeg command and return a promise. Logs stderr and sets activeExportProcess. */ -export function runFfmpeg(ffmpegPath: string, args: string[]): Promise<{ success: boolean; error?: string }> { +/** Run an ffmpeg command and return a promise. Logs stderr and sets activeExportProcess. + * `onProgress`, if given, is called with the encoded output position in seconds + * (parsed from ffmpeg's `time=` field) so callers can drive a progress bar. */ +export function runFfmpeg( + ffmpegPath: string, + args: string[], + onProgress?: (outTimeSec: number) => void, +): Promise<{ success: boolean; error?: string }> { return new Promise((resolve) => { logger.info( `[ffmpeg] spawn: ${args.join(' ').slice(0, 400)}`) const proc = spawn(ffmpegPath, args, { stdio: ['pipe', 'pipe', 'pipe'] }) @@ -63,11 +69,20 @@ export function runFfmpeg(ffmpegPath: string, args: string[]): Promise<{ success proc.stderr?.on('data', (chunk: Buffer) => { const text = chunk.toString() stderrLog += text - const lines = text.trim().split('\n') + // ffmpeg rewrites its progress line in place with carriage returns, so + // split on both \r and \n to see each update, not one giant line. + const lines = text.split(/[\r\n]+/) for (const line of lines) { if (line.includes('frame=') || line.includes('Error') || line.includes('error')) { logger.info( `[ffmpeg] ${line.trim().slice(0, 200)}`) } + if (onProgress) { + const m = line.match(/time=\s*(\d+):(\d+):(\d+(?:\.\d+)?)/) + if (m) { + const sec = Number(m[1]) * 3600 + Number(m[2]) * 60 + parseFloat(m[3]) + if (Number.isFinite(sec)) onProgress(sec) + } + } } }) proc.on('close', (code) => { @@ -101,6 +116,7 @@ export function extractVideoFrameToFile({ width, quality, outputPath, + accurate = false, timeoutMs = 10000, }: { videoPath: string @@ -108,6 +124,13 @@ export function extractVideoFrameToFile({ width?: number quality?: number outputPath?: string + /** + * When true, seek *after* -i (frame-accurate but slower, decodes from 0 to + * seekTime). Default false uses the fast keyframe seek before -i, which is + * fine for thumbnails but can land a few frames off — not acceptable when the + * user is saving the exact frame they paused on. + */ + accurate?: boolean timeoutMs?: number }): string { const ffmpegPath = findFfmpegPath() @@ -124,9 +147,9 @@ export function extractVideoFrameToFile({ `ltx_frame_${Date.now()}_${Math.random().toString(36).slice(2, 8)}.jpg`, ) + const seekArgs = ['-ss', String(Math.max(0, seekTime))] const args: string[] = [ - '-ss', String(Math.max(0, seekTime)), - '-i', videoPath, + ...(accurate ? ['-i', videoPath, ...seekArgs] : [...seekArgs, '-i', videoPath]), ...(width ? ['-vf', `scale=${width}:-2`] : []), '-frames:v', '1', ...(quality !== undefined ? ['-q:v', String(quality)] : []), @@ -174,6 +197,130 @@ export function getVideoDimensions(videoPath: string): { width: number; height: return { width, height } } +/** Parse the video stream's fps from ffmpeg -i output; falls back to 24. */ +export function getVideoFps(ffmpegPath: string, videoPath: string): number { + try { + const result = spawnSync(ffmpegPath, ['-i', videoPath, '-hide_banner'], { encoding: 'utf8', timeout: 5000 }) + const output = (result.stdout || '') + (result.stderr || '') + const m = output.match(/(\d+(?:\.\d+)?)\s*fps/) + const fps = m ? Number(m[1]) : NaN + return Number.isFinite(fps) && fps > 0 ? fps : 24 + } catch { + return 24 + } +} + +/** + * Extract the final frame of a video to a full-resolution PNG. Reads only the + * last second (`-sseof -1`) and reverses it so `-frames:v 1` yields the true + * last frame — robust across variable frame counts, no fps math needed. + */ +export function extractLastFrameToFile({ videoPath, outputPath, timeoutMs = 15000 }: { + videoPath: string + outputPath: string + timeoutMs?: number +}): string { + const ffmpegPath = findFfmpegPath() + if (!ffmpegPath) throw new Error('ffmpeg not found') + if (!fs.existsSync(videoPath)) throw new Error(`Video file not found: ${videoPath}`) + + const args = ['-sseof', '-1', '-i', videoPath, '-vf', 'reverse', '-frames:v', '1', '-y', outputPath] + logger.info(`[extract-last-frame] ${args.join(' ').slice(0, 300)}`) + runFfmpegSyncOrThrow(ffmpegPath, args, timeoutMs) + if (!fs.existsSync(outputPath)) throw new Error('ffmpeg produced no output file') + return outputPath +} + +/** + * Average RGB of one frame, via a true box-average downscale to 1x1 read as raw + * rgb24 (3 bytes). `vf` selects/prepares the frame(s) before the scale (e.g. + * "trim=start_frame=1:end_frame=2" to grab a video's second frame). Returns null + * if ffmpeg produces nothing usable — callers must treat that as "skip matching". + */ +function averageRgb(ffmpegPath: string, inputPath: string, vf: string): [number, number, number] | null { + const result = spawnSync( + ffmpegPath, + ['-v', 'error', '-i', inputPath, '-vf', `${vf},scale=1:1:flags=area`, '-frames:v', '1', + '-f', 'rawvideo', '-pix_fmt', 'rgb24', '-'], + { maxBuffer: 1024 * 1024, timeout: 15000 }, + ) + const buf = result.stdout + if (!buf || buf.length < 3) return null + return [buf[0], buf[1], buf[2]] +} + +// Cap per-channel correction: the systematic VAE darkening is ~2-3/255, so a +// larger measured delta means the continuation legitimately changed grade (a +// light turned on, the camera moved to a brighter area) -- clamp so the match +// only cancels the bias and never fights real content. +const _MAX_COLOR_MATCH_OFFSET = 8 + +/** + * Per-channel additive offset (as an ffmpeg lutrgb filter fragment) that nudges the + * continuation's grade back to the source's, or null when no meaningful correction + * applies. `referencePath` is the seed frame (the source clip's last frame); the + * clip's second frame (index 1 -- the one that becomes the new lead after the trim) + * is what we match, since it is the boundary the viewer sees against the source. + */ +function colorMatchLutFilter(ffmpegPath: string, videoPath: string, referencePath: string): string | null { + const seed = averageRgb(ffmpegPath, referencePath, 'null') + const clip = averageRgb(ffmpegPath, videoPath, 'trim=start_frame=1:end_frame=2,setpts=PTS-STARTPTS') + if (!seed || !clip) return null + const clamp = (v: number) => Math.max(-_MAX_COLOR_MATCH_OFFSET, Math.min(_MAX_COLOR_MATCH_OFFSET, v)) + const [dr, dg, db] = [clamp(seed[0] - clip[0]), clamp(seed[1] - clip[1]), clamp(seed[2] - clip[2])] + // Sub-level offsets are below what encode rounding would preserve -- skip the + // extra filter pass rather than apply a no-op. + if (Math.abs(dr) < 0.5 && Math.abs(dg) < 0.5 && Math.abs(db) < 0.5) return null + const r = dr.toFixed(1), g = dg.toFixed(1), b = db.toFixed(1) + logger.info(`[trim-first-frame] color-match offset r=${r} g=${g} b=${b} (seed=${seed} clip1=${clip})`) + // lutrgb clamps expression output to [0,255] itself, so no explicit clip() needed. + return `lutrgb=r=val+(${r}):g=val+(${g}):b=val+(${b})` +} + +/** + * Re-encode a video with its first frame removed (and the matching audio slice + * dropped so A/V stays in sync). Used by "Continue as new shot" to delete the + * duplicate lead frame the i2v conditioning reproduces from the source's last + * frame, so the continuation butt-joins the source with no stutter. + * + * `colorMatchReferencePath` (the seed frame) enables a subtle per-channel grade + * match: the i2v VAE reproduces the conditioning frame ~2-3/255 darker, which + * compounds across chained continuations; matching the clip's lead frame back to + * the seed cancels that at the join. Measured/clamped, so it only removes the + * systematic bias. + */ +export function trimFirstFrameToFile({ videoPath, outputPath, colorMatchReferencePath, timeoutMs = 120000 }: { + videoPath: string + outputPath: string + colorMatchReferencePath?: string + timeoutMs?: number +}): string { + const ffmpegPath = findFfmpegPath() + if (!ffmpegPath) throw new Error('ffmpeg not found') + if (!fs.existsSync(videoPath)) throw new Error(`Video file not found: ${videoPath}`) + + const hasAudio = fileHasAudio(ffmpegPath, videoPath) + const fps = getVideoFps(ffmpegPath, videoPath) + const colorMatchLut = + colorMatchReferencePath && fs.existsSync(colorMatchReferencePath) + ? colorMatchLutFilter(ffmpegPath, videoPath, colorMatchReferencePath) + : null + const videoFilter = ['trim=start_frame=1', 'setpts=PTS-STARTPTS', ...(colorMatchLut ? [colorMatchLut] : [])].join(',') + const args: string[] = [ + '-i', videoPath, + '-vf', videoFilter, + ...(hasAudio + ? ['-af', `atrim=start=${(1 / fps).toFixed(6)},asetpts=PTS-STARTPTS`, '-c:a', 'aac', '-b:a', '192k'] + : ['-an']), + '-c:v', 'libx264', '-crf', '18', '-preset', 'veryfast', '-pix_fmt', 'yuv420p', + '-y', outputPath, + ] + logger.info(`[trim-first-frame] ${args.join(' ').slice(0, 300)}`) + runFfmpegSyncOrThrow(ffmpegPath, args, timeoutMs) + if (!fs.existsSync(outputPath)) throw new Error('ffmpeg produced no output file') + return outputPath +} + export function stopExportProcess(): void { if (activeExportProcess) { logger.info( 'Stopping active export process...') diff --git a/electron/export/timeline.ts b/electron/export/timeline.ts index fe3685106..a257f6bfc 100644 --- a/electron/export/timeline.ts +++ b/electron/export/timeline.ts @@ -1,13 +1,31 @@ +export interface ColorCorrection { + brightness: number; contrast: number; saturation: number; temperature: number; + tint: number; exposure: number; highlights: number; shadows: number; +} + +export interface ClipTransition { + type: string; duration: number; +} + export interface ExportClip { path: string; type: string; startTime: number; duration: number; trimStart: number; speed: number; reversed: boolean; flipH: boolean; flipV: boolean; opacity: number; trackIndex: number; - muted: boolean; volume: number; + muted: boolean; volume: number; audioFadeIn?: number; audioFadeOut?: number; + volumeKeyframes?: { t: number; value: number }[]; + colorCorrection?: ColorCorrection; transitionIn?: ClipTransition; transitionOut?: ClipTransition; } export interface FlatSegment { filePath: string; type: string; startTime: number; duration: number; trimStart: number; speed: number; reversed: boolean; flipH: boolean; flipV: boolean; opacity: number; muted: boolean; volume: number; + colorCorrection?: ColorCorrection; transitionIn?: ClipTransition; transitionOut?: ClipTransition; + // Position of this segment within its ORIGINAL clip's own timeline (not the + // overall program timeline) and that clip's total duration - needed so a + // clip fragmented by a higher-track overlay still only applies its + // transitionIn/Out at the true start/end of the original clip, not at + // every fragment boundary. + offsetInClip: number; clipDuration: number; } /** @@ -58,12 +76,18 @@ export function flattenTimeline(clips: ExportClip[]): FlatSegment[] { opacity: c.opacity, muted: c.muted, volume: c.volume, + colorCorrection: c.colorCorrection, + transitionIn: c.transitionIn, + transitionOut: c.transitionOut, + offsetInClip, + clipDuration: c.duration, }) } else { segments.push({ filePath: '', type: 'gap', startTime: t0, duration: segDur, trimStart: 0, speed: 1, reversed: false, flipH: false, flipV: false, opacity: 100, muted: true, volume: 0, + offsetInClip: 0, clipDuration: 0, }) } } @@ -85,3 +109,83 @@ export function flattenTimeline(clips: ExportClip[]): FlatSegment[] { return merged } + +export interface DissolveBoundary { + /** Dissolve sits between segments[index] and segments[index + 1]. */ + index: number + duration: number +} + +/** + * Which adjacent segment pairs should cross-fade instead of cut. Mirrors the + * editor's own preview logic (ProgramMonitor's getDissolveAtTime) exactly: + * the duration comes ONLY from the outgoing clip's transitionOut - the + * incoming clip's transitionIn.duration is never read, only its type is + * checked as a gate. Guarded the same way fades/wipes are (offsetInClip + * reaching the clip's true end / start) so a clip fragmented by a + * higher-track overlay doesn't dissolve at every internal fragment boundary. + */ +export function findDissolveBoundaries(segments: FlatSegment[]): DissolveBoundary[] { + const boundaries: DissolveBoundary[] = [] + for (let i = 0; i < segments.length - 1; i++) { + const a = segments[i] + const b = segments[i + 1] + if (a.transitionOut?.type !== 'dissolve' || !(a.transitionOut.duration > 0)) continue + if (b.transitionIn?.type !== 'dissolve') continue + + const aReachesClipEnd = Math.abs((a.offsetInClip + a.duration) - a.clipDuration) < 0.01 + const bIsClipStart = b.offsetInClip < 0.01 + if (!aReachesClipEnd || !bIsClipStart) continue + + const duration = Math.min(a.transitionOut.duration, a.duration, b.duration) + if (duration > 0) boundaries.push({ index: i, duration }) + } + return boundaries +} + +/** + * A dissolve overlaps the tail of one clip with the head of the next + * (matching the editor's own preview - see ProgramMonitor's + * getDissolveAtTime), shrinking the program's real presented duration below + * the naive sum of clip durations. Every clip's *nominal* startTime (as + * stored in the project) stays untouched by a dissolve, so anything that + * schedules events by nominal time - audio mixing in particular - needs to + * convert to the shrunk *actual* output time or it drifts out of sync with + * the video after every dissolve. Returns that conversion as a function + * rather than a lookup table since it needs to answer for arbitrary + * timestamps (e.g. an audio clip's own startTime), not just segment + * boundaries. + */ +export function buildDissolveTimeRemap(segments: FlatSegment[]): (nominalTime: number) => number { + const boundaries = findDissolveBoundaries(segments) + if (boundaries.length === 0) return (t) => t + + const shrinkPoints = boundaries + .map(b => ({ atTime: segments[b.index + 1].startTime, amount: b.duration })) + .sort((a, b) => a.atTime - b.atTime) + + return (nominalTime: number): number => { + let shrink = 0 + for (const point of shrinkPoints) { + if (point.atTime > nominalTime) break + shrink += point.amount + } + return nominalTime - shrink + } +} + +/** The program's real presented duration after dissolve overlaps shrink it + * below the naive sum of segment durations - the authoritative source for + * this is buildVideoFilterGraph's own accumulation (it has to compute the + * same running total to build correct xfade offsets), so this just mirrors + * that logic for callers that need the number without building the graph. */ +export function computeFinalVideoDuration(segments: FlatSegment[]): number { + const boundaries = findDissolveBoundaries(segments) + const shrinkAfter = new Map(boundaries.map(b => [b.index, b.duration])) + let total = segments.length > 0 ? segments[0].duration : 0 + for (let i = 1; i < segments.length; i++) { + const dissolveDuration = shrinkAfter.get(i - 1) + total += segments[i].duration - (dissolveDuration ?? 0) + } + return total +} diff --git a/electron/export/video-filter.ts b/electron/export/video-filter.ts index fccc9f66f..f17b75967 100644 --- a/electron/export/video-filter.ts +++ b/electron/export/video-filter.ts @@ -1,10 +1,226 @@ -import type { FlatSegment } from './timeline' +import { findDissolveBoundaries, type FlatSegment } from './timeline' export interface ExportSubtitle { text: string; startTime: number; endTime: number; style: { fontSize: number; fontFamily: string; fontWeight: string; color: string; backgroundColor: string; position: string; italic: boolean }; } +export interface ExportTextOverlay { + text: string; startTime: number; endTime: number; + fadeIn?: number; fadeOut?: number; + style: { + fontSize: number; color: string; backgroundColor: string; + positionX: number; positionY: number; + strokeColor: string; strokeWidth: number; + shadowColor: string; shadowOffsetX: number; shadowOffsetY: number; + opacity: number; padding: number; + textAlign?: string; + }; +} + +/** Build a drawtext `fontfile=` argument for a system font path. The path is + * single-quoted and its drive-letter colon escaped — the only form that parses + * inside a filter_complex_script on Windows (bare or backslash-only both fail + * with "No option name near ...", verified against the bundled ffmpeg). */ +function fontFileArg(p: string): string { + return `fontfile='${p.replace(/\\/g, '/').replace(/:/g, '\\:')}'` +} + +/** Escape a string for use inside an ffmpeg drawtext text='...' value. */ +function escapeDrawtext(text: string): string { + return text + .replace(/\\/g, '\\\\\\\\') + .replace(/'/g, "'\\\\\\''") + .replace(/:/g, '\\:') + .replace(/%/g, '%%') + // Newlines stay RAW: inside the quoted value a literal newline survives the + // filtergraph parser and drawtext breaks the line there. An escaped `\n` is + // un-escaped to a plain "n" ("TRICK ORnTREAT"). +} + +/** drawtext text_align letter for a CSS text-align (multi-line blocks). */ +function drawtextAlign(align: string | undefined): 'L' | 'C' | 'R' { + return align === 'left' ? 'L' : align === 'right' ? 'R' : 'C' +} + +/** Convert a CSS color (hex / rgb(a) / named / transparent) to an ffmpeg color + * token like `0xRRGGBB@0.500`. `extraAlpha` multiplies the resolved alpha (used + * for the overlay's global opacity). Returns null for transparent/none. */ +function cssColorToFfmpeg(css: string, extraAlpha = 1): string | null { + const c = (css || '').trim().toLowerCase() + if (!c || c === 'transparent' || c === 'none') return null + let r = 0, g = 0, b = 0, a = 1 + let m: RegExpMatchArray | null + if ((m = c.match(/^#([0-9a-f]{3})$/))) { + r = parseInt(m[1][0] + m[1][0], 16); g = parseInt(m[1][1] + m[1][1], 16); b = parseInt(m[1][2] + m[1][2], 16) + } else if ((m = c.match(/^#([0-9a-f]{6})([0-9a-f]{2})?$/))) { + r = parseInt(m[1].slice(0, 2), 16); g = parseInt(m[1].slice(2, 4), 16); b = parseInt(m[1].slice(4, 6), 16) + if (m[2]) a = parseInt(m[2], 16) / 255 + } else if ((m = c.match(/^rgba?\(([^)]+)\)$/))) { + const p = m[1].split(',').map(s => s.trim()) + r = parseInt(p[0], 10) || 0; g = parseInt(p[1], 10) || 0; b = parseInt(p[2], 10) || 0 + if (p[3] !== undefined) a = parseFloat(p[3]) + } else { + // A named color (white, black, red, ...) — ffmpeg understands these directly. + const alpha = Math.max(0, Math.min(1, extraAlpha)) + return alpha >= 0.999 ? c : `${c}@${alpha.toFixed(3)}` + } + const hex = (((r & 255) << 16) | ((g & 255) << 8) | (b & 255)).toString(16).padStart(6, '0') + const alpha = Math.max(0, Math.min(1, (isFinite(a) ? a : 1) * extraAlpha)) + return `0x${hex}@${alpha.toFixed(3)}` +} + +/** + * Color correction + fade-to-black/white + wipe filters for one segment, as + * an ffmpeg filter-chain suffix (leading comma, or '' if nothing applies). + * + * The eight color sliders don't map 1:1 onto ffmpeg's eq filter (which only + * exposes one brightness/contrast/saturation knob each), so brightness, + * exposure, and highlights combine into eq's additive brightness, and + * contrast and shadows combine into its multiplicative contrast - matching + * how the editor's own CSS preview groups them conceptually. This is a + * close creative match to the live preview, not colorimetric precision - + * consistent with the preview's own hue-rotate/sepia approximations. + */ +function buildGradingFilters(seg: FlatSegment, localDuration: number): string { + const parts: string[] = [] + const cc = seg.colorCorrection + + if (cc) { + const brightness = cc.brightness / 100 + cc.exposure / 200 + cc.highlights / 300 + const contrast = 1 + cc.contrast / 100 + cc.shadows / 300 + const saturation = Math.max(0, Math.min(3, 1 + cc.saturation / 100)) + if (brightness !== 0 || contrast !== 1 || saturation !== 1) { + parts.push(`eq=brightness=${brightness.toFixed(4)}:contrast=${contrast.toFixed(4)}:saturation=${saturation.toFixed(4)}`) + } + if (cc.tint !== 0) { + parts.push(`hue=h=${(cc.tint * 1.2).toFixed(2)}`) + } + if (cc.temperature !== 0) { + const kelvin = Math.max(1000, Math.min(40000, Math.round(6500 - cc.temperature * 35))) + parts.push(`colortemperature=temperature=${kelvin}`) + } + } + + // Fades are defined relative to the ORIGINAL clip's start/end, so only + // apply them to the fragment that actually touches that edge - a clip + // split by a higher-track overlay would otherwise fade every fragment. + const tIn = seg.transitionIn + if (tIn && (tIn.type === 'fade-to-black' || tIn.type === 'fade-to-white') && tIn.duration > 0 && seg.offsetInClip < 0.01) { + const d = Math.min(tIn.duration, localDuration) + parts.push(`fade=t=in:st=0:d=${d.toFixed(4)}:color=${tIn.type === 'fade-to-black' ? 'black' : 'white'}`) + } + const tOut = seg.transitionOut + if (tOut && (tOut.type === 'fade-to-black' || tOut.type === 'fade-to-white') && tOut.duration > 0) { + const reachesClipEnd = Math.abs((seg.offsetInClip + localDuration) - seg.clipDuration) < 0.01 + if (reachesClipEnd) { + const d = Math.min(tOut.duration, localDuration) + const st = Math.max(0, localDuration - d) + parts.push(`fade=t=out:st=${st.toFixed(4)}:d=${d.toFixed(4)}:color=${tOut.type === 'fade-to-black' ? 'black' : 'white'}`) + } + } + + return parts.length > 0 ? ',' + parts.join(',') : '' +} + +/** + * getWipeClipPath's direction names describe which way the wipe motion + * travels, matching ffmpeg's own xfade transition naming directly - this + * mapping is intentionally the identity (minus the hyphen), verified + * empirically by rendering each ffmpeg transition and comparing + * pixel-for-pixel against getWipeClipPath's output. Same mapping for both + * transitionIn and transitionOut. Kept as an explicit table rather than + * derived from the string (e.g. stripping the hyphen) so it stays correct + * and obvious if either naming scheme ever changes independently. + * + * Known limitation: ffmpeg's xfade completes its transition roughly one + * frame earlier than the nominal duration (verified with frame-by-frame + * contact sheets at 4fps and 24fps) - a fixed ~1-frame offset regardless of + * duration, not a scaling error. At typical transition lengths this is a + * fraction of a frame's worth of time (<50ms at 24fps) and isn't + * perceptible; not worth working around given it's inherent to xfade itself + * rather than this code's own math. + */ +const WIPE_TYPE_TO_XFADE: Record = { + 'wipe-left': 'wipeleft', + 'wipe-right': 'wiperight', + 'wipe-up': 'wipeup', + 'wipe-down': 'wipedown', +} + +/** + * Wipes are true two-layer composites (revealing/hiding against black), not + * a single-input filter like fade - ffmpeg has no per-pixel time-animated + * crop/drawbox (their w/h/x/y expressions are evaluated once at filter + * init, not per-frame), so this reuses ffmpeg's own tested xfade transition + * types against a synthetic black source instead of hand-rolling pixel math. + * + * Returns the new input args/filter lines to append and the label the wipe + * result ends up on, or null if no wipe applies to this segment. + */ +function applyWipeTransitions( + contentLabel: string, + seg: FlatSegment, + localDuration: number, + opts: { width: number; height: number; fps: number }, + nextIdx: number, +): { inputs: string[]; filterLines: string[]; label: string; nextIdx: number } | null { + const { width, height, fps } = opts + const inputs: string[] = [] + const filterLines: string[] = [] + let label = contentLabel + let idx = nextIdx + let applied = false + + const addBlackInput = (): number => { + inputs.push('-f', 'lavfi', '-i', `color=c=black:s=${width}x${height}:r=${fps}:d=${localDuration.toFixed(6)}`) + return idx++ + } + + const tIn = seg.transitionIn + if (tIn && tIn.type.startsWith('wipe-') && tIn.duration > 0 && seg.offsetInClip < 0.01) { + const xfadeType = WIPE_TYPE_TO_XFADE[tIn.type] + if (xfadeType) { + const d = Math.min(tIn.duration, localDuration) + const blackIdx = addBlackInput() + const nextLabel = `${contentLabel}wi` + // fps-align both sides: the content chain intentionally skips + // per-segment fps conversion (applied once after concat), but xfade + // blends frame-for-frame and needs matched timing to do that correctly. + filterLines.push( + `[${blackIdx}:v]fps=${fps},setsar=1[${nextLabel}blk];` + + `[${label}]fps=${fps}[${nextLabel}clip];` + + `[${nextLabel}blk][${nextLabel}clip]xfade=transition=${xfadeType}:duration=${d.toFixed(4)}:offset=0[${nextLabel}]`, + ) + label = nextLabel + applied = true + } + } + + const tOut = seg.transitionOut + if (tOut && tOut.type.startsWith('wipe-') && tOut.duration > 0) { + const reachesClipEnd = Math.abs((seg.offsetInClip + localDuration) - seg.clipDuration) < 0.01 + if (reachesClipEnd) { + const xfadeType = WIPE_TYPE_TO_XFADE[tOut.type] + if (xfadeType) { + const d = Math.min(tOut.duration, localDuration) + const st = Math.max(0, localDuration - d) + const blackIdx = addBlackInput() + const nextLabel = `${contentLabel}wo` + filterLines.push( + `[${label}]fps=${fps}[${nextLabel}clip];` + + `[${blackIdx}:v]fps=${fps},setsar=1[${nextLabel}blk];` + + `[${nextLabel}clip][${nextLabel}blk]xfade=transition=${xfadeType}:duration=${d.toFixed(4)}:offset=${st.toFixed(4)}[${nextLabel}]`, + ) + label = nextLabel + applied = true + } + } + } + + return applied ? { inputs, filterLines, label, nextIdx: idx } : null +} + /** * Build the ffmpeg filter_complex script and input arguments for the video-only pass. * Pure string building — zero I/O. @@ -15,9 +231,11 @@ export function buildVideoFilterGraph( width: number; height: number; fps: number; letterbox?: { ratio: number; color: string; opacity: number }; subtitles?: ExportSubtitle[]; + textOverlays?: ExportTextOverlay[]; + fontFile?: string; }, ): { inputs: string[]; filterScript: string } { - const { width, height, fps, letterbox, subtitles } = opts + const { width, height, fps, letterbox, subtitles, textOverlays, fontFile } = opts const inputs: string[] = [] const filterParts: string[] = [] let idx = 0 @@ -25,10 +243,12 @@ export function buildVideoFilterGraph( for (let i = 0; i < segments.length; i++) { const seg = segments[i] + const contentLabel = `v${i}c` + if (seg.type === 'gap') { // Gap: generate black frames at target fps (synthetic input) inputs.push('-f', 'lavfi', '-i', `color=c=black:s=${width}x${height}:r=${fps}:d=${seg.duration.toFixed(6)}`) - filterParts.push(`[${idx}:v]setsar=1[v${i}]`) + filterParts.push(`[${idx}:v]setsar=1[${contentLabel}]`) idx++ } else if (seg.type === 'image') { // Image: loop for exact duration, use target fps for frame generation @@ -36,7 +256,8 @@ export function buildVideoFilterGraph( let chain = `[${idx}:v]scale=${width}:${height}:force_original_aspect_ratio=decrease,pad=${width}:${height}:-1:-1:color=black,setsar=1` if (seg.flipH) chain += ',hflip' if (seg.flipV) chain += ',vflip' - chain += `[v${i}]` + chain += buildGradingFilters(seg, seg.duration) + chain += `[${contentLabel}]` filterParts.push(chain) idx++ } else { @@ -50,20 +271,80 @@ export function buildVideoFilterGraph( chain += `,scale=${width}:${height}:force_original_aspect_ratio=decrease,pad=${width}:${height}:-1:-1:color=black,setsar=1` if (seg.flipH) chain += ',hflip' if (seg.flipV) chain += ',vflip' - chain += `[v${i}]` + chain += buildGradingFilters(seg, seg.duration) + chain += `[${contentLabel}]` filterParts.push(chain) idx++ } + + const wipe = applyWipeTransitions(contentLabel, seg, seg.duration, { width, height, fps }, idx) + if (wipe) { + inputs.push(...wipe.inputs) + filterParts.push(...wipe.filterLines) + idx = wipe.nextIdx + filterParts.push(`[${wipe.label}]null[v${i}]`) + } else { + filterParts.push(`[${contentLabel}]null[v${i}]`) + } } - const concatInputs = segments.map((_, i) => `[v${i}]`).join('') + const dissolveBoundaries = findDissolveBoundaries(segments) + + let lastLabel: string + if (dissolveBoundaries.length === 0) { + // No dissolves: concat all segments in one pass, then apply fps ONCE to + // the entire output. This is how real NLEs work - frame rate conversion + // happens globally, not per-clip, so per-segment duration quantization + // doesn't accumulate. + const concatInputs = segments.map((_, i) => `[v${i}]`).join('') + lastLabel = 'fpsout' + filterParts.push(`${concatInputs}concat=n=${segments.length}:v=1:a=0[concatraw]`) + filterParts.push(`[concatraw]fps=${fps}[${lastLabel}]`) + } else { + // At least one dissolve: xfade blends two streams frame-for-frame, which + // needs matched timing, so every segment gets fps-normalized up front + // instead of once at the end. Segments combine left-to-right through an + // accumulator, using xfade at dissolve boundaries (which - like the + // editor's own preview - overlaps the tail of one clip with the head of + // the next, shrinking total duration by the dissolve's length) and plain + // concat everywhere else. + // + // xfade's "dissolve" transition is a misnomer for what this app (and + // every mainstream NLE) means by dissolve: it's a randomized pixel + // dither, not a smooth cross-fade - confirmed by blending solid red and + // blue test frames at 50%: "dissolve" produced visibly speckled + // red/blue pixels, not a uniform blend. xfade's "fade" is the one that + // actually does a plain linear alpha blend (verified: uniform purple at + // 50%), matching the editor's own opacity-based preview. + const XFADE_DISSOLVE_TYPE = 'fade' + const dissolveDurationAfter = new Map(dissolveBoundaries.map(b => [b.index, b.duration])) + + filterParts.push(`[v0]fps=${fps}[vacc0]`) + let accLabel = 'vacc0' + let accDuration = segments[0].duration + + for (let i = 1; i < segments.length; i++) { + const segFpsLabel = `v${i}fps` + filterParts.push(`[v${i}]fps=${fps}[${segFpsLabel}]`) + + const nextAccLabel = `vacc${i}` + const dissolveDuration = dissolveDurationAfter.get(i - 1) + if (dissolveDuration !== undefined) { + const offset = Math.max(0, accDuration - dissolveDuration) + filterParts.push( + `[${accLabel}][${segFpsLabel}]xfade=transition=${XFADE_DISSOLVE_TYPE}:duration=${dissolveDuration.toFixed(4)}:offset=${offset.toFixed(4)}[${nextAccLabel}]`, + ) + accDuration = accDuration + segments[i].duration - dissolveDuration + } else { + filterParts.push(`[${accLabel}][${segFpsLabel}]concat=n=2:v=1:a=0[${nextAccLabel}]`) + accDuration = accDuration + segments[i].duration + } + accLabel = nextAccLabel + } - // Concat all segments, then apply fps ONCE to the entire output. - // This is how real NLEs work: frame rate conversion happens globally, - // not per-clip, so per-segment duration quantization doesn't accumulate. - let lastLabel = 'fpsout' - filterParts.push(`${concatInputs}concat=n=${segments.length}:v=1:a=0[concatraw]`) - filterParts.push(`[concatraw]fps=${fps}[${lastLabel}]`) + lastLabel = 'fpsout' + filterParts.push(`[${accLabel}]null[${lastLabel}]`) + } // Letterbox overlay (drawbox) if (letterbox) { @@ -93,18 +374,89 @@ export function buildVideoFilterGraph( } } + // Text-overlay burn-in (drawtext) — the type:'text' clips with a textStyle. + // The live preview renders these as DOM; export mirrors position/size/color so + // the baked video matches. fontFamily/weight/style, letter-spacing, auto-wrap + // (maxWidth) and shadow blur aren't representable in drawtext and are dropped; + // explicit newlines in the text are preserved. + if (textOverlays && textOverlays.length > 0) { + const hf = height / 1080 // scale style px (authored against 1080p) to export height + for (let ti = 0; ti < textOverlays.length; ti++) { + const ov = textOverlays[ti] + const s = ov.style + const nextLabel = `txt${ti}` + const gA = Math.max(0, Math.min(1, (s.opacity ?? 100) / 100)) // global overlay opacity + const fontSize = Math.max(1, Math.round(s.fontSize * hf)) + const fontColor = cssColorToFfmpeg(s.color, gA) ?? `white@${gA.toFixed(3)}` + + // Preview positions the box's CENTER at (positionX%, positionY%); mirror that. + const px = Math.max(0, Math.min(1, (s.positionX ?? 50) / 100)) + const py = Math.max(0, Math.min(1, (s.positionY ?? 50) / 100)) + + const parts: string[] = [] + if (fontFile) parts.push(fontFileArg(fontFile)) + parts.push(`text='${escapeDrawtext(ov.text)}'`) + if (ov.text.includes('\n')) parts.push(`text_align=${drawtextAlign(s.textAlign)}`) + parts.push(`fontsize=${fontSize}`) + parts.push(`fontcolor=${fontColor}`) + parts.push(`x=(w*${px.toFixed(4)})-(text_w/2)`) + parts.push(`y=(h*${py.toFixed(4)})-(text_h/2)`) + + if ((s.strokeWidth ?? 0) > 0) { + const strokeColor = cssColorToFfmpeg(s.strokeColor, gA) + if (strokeColor) { + parts.push(`borderw=${Math.max(1, Math.round(s.strokeWidth * hf))}`) + parts.push(`bordercolor=${strokeColor}`) + } + } + + const shx = Math.round((s.shadowOffsetX ?? 0) * hf) + const shy = Math.round((s.shadowOffsetY ?? 0) * hf) + if (shx !== 0 || shy !== 0) { + const shadowColor = cssColorToFfmpeg(s.shadowColor, gA) + if (shadowColor) { + parts.push(`shadowx=${shx}`) + parts.push(`shadowy=${shy}`) + parts.push(`shadowcolor=${shadowColor}`) + } + } + + const boxColor = cssColorToFfmpeg(s.backgroundColor, gA) + if (boxColor) { + parts.push('box=1') + parts.push(`boxcolor=${boxColor}`) + parts.push(`boxborderw=${Math.max(0, Math.round((s.padding ?? 0) * hf))}`) + } + + // Opacity fade in/out via a time-based alpha expression (defaults 0.5s, + // each capped at half the overlay). Multiplies the drawtext alpha, so it + // rides on top of the style's own opacity. Commas escaped like `enable`. + const ovDur = Math.max(0.0001, ov.endTime - ov.startTime) + const fin = Math.min(ov.fadeIn ?? 0.5, ovDur / 2) + const fout = Math.min(ov.fadeOut ?? 0.5, ovDur / 2) + const ramps: string[] = [] + if (fin > 0.001) ramps.push(`(t-${ov.startTime.toFixed(3)})/${fin.toFixed(3)}`) + if (fout > 0.001) ramps.push(`(${ov.endTime.toFixed(3)}-t)/${fout.toFixed(3)}`) + if (ramps.length > 0) { + const inner = ramps.length === 2 ? `min(${ramps[0]}\\,${ramps[1]})` : ramps[0] + parts.push(`alpha='max(0\\,min(1\\,${inner}))'`) + } + + parts.push(`enable='between(t\\,${ov.startTime.toFixed(3)}\\,${ov.endTime.toFixed(3)})'`) + + filterParts.push(`[${lastLabel}]drawtext=${parts.join(':')}[${nextLabel}]`) + lastLabel = nextLabel + } + } + // Subtitle burn-in (drawtext) if (subtitles && subtitles.length > 0) { for (let si = 0; si < subtitles.length; si++) { const sub = subtitles[si] const nextLabel = `sub${si}` - // Escape text for ffmpeg drawtext: replace special chars - const escapedText = sub.text - .replace(/\\/g, '\\\\\\\\') - .replace(/'/g, "'\\\\\\''") - .replace(/:/g, '\\:') - .replace(/%/g, '%%') - .replace(/\n/g, '\\n') + // Escape text for ffmpeg drawtext (newlines stay raw, see escapeDrawtext). + const escapedText = escapeDrawtext(sub.text) + const alignPart = sub.text.includes('\n') ? ':text_align=C' : '' const fontSize = Math.round(sub.style.fontSize * (height / 1080)) // scale relative to export res const fontColor = sub.style.color.replace('#', '0x') @@ -129,7 +481,7 @@ export function buildVideoFilterGraph( boxPart = `:box=1:boxcolor=${bgColor}@${bgAlpha}:boxborderw=8` } - const dtFilter = `drawtext=text='${escapedText}':fontsize=${fontSize}:fontcolor=${fontColor}:x=(w-text_w)/2:y=${yExpr}${boxPart}:enable='between(t\\,${sub.startTime.toFixed(3)}\\,${sub.endTime.toFixed(3)})'` + const dtFilter = `drawtext=text='${escapedText}'${alignPart}:fontsize=${fontSize}:fontcolor=${fontColor}:x=(w-text_w)/2:y=${yExpr}${boxPart}:enable='between(t\\,${sub.startTime.toFixed(3)}\\,${sub.endTime.toFixed(3)})'` filterParts.push(`[${lastLabel}]${dtFilter}[${nextLabel}]`) lastLabel = nextLabel diff --git a/electron/gpm-library-root.ts b/electron/gpm-library-root.ts new file mode 100644 index 000000000..ad4933862 --- /dev/null +++ b/electron/gpm-library-root.ts @@ -0,0 +1,44 @@ +import { app } from 'electron' +import fs from 'fs' +import path from 'path' + +// Lets the user point the Prompt Manager Pro Downloads Browser at any folder +// on disk instead of the default `/PromptManagerPro`. The chosen +// path always comes from a native dialog (trusted), and is persisted here so +// it survives restarts and can be added to the allowed-roots list. + +const CONFIG_PATH = () => path.join(app.getPath('userData'), 'gpm-library-root.json') +let cached: string | null | undefined + +function load(): string | null { + if (cached !== undefined) return cached + try { + const raw = fs.readFileSync(CONFIG_PATH(), 'utf-8') + const parsed = JSON.parse(raw) as { root?: string } + cached = parsed.root && fs.existsSync(parsed.root) ? parsed.root : null + } catch { + cached = null + } + return cached +} + +export function getLibraryRootOverride(): string | null { + return load() +} + +export function setLibraryRootOverride(root: string | null): void { + cached = root + try { + fs.writeFileSync(CONFIG_PATH(), JSON.stringify({ root })) + } catch { + // Best effort — worst case the override doesn't survive a restart. + } +} + +export function defaultLibraryRoot(): string { + return path.join(app.getPath('downloads'), 'PromptManagerPro') +} + +export function resolveLibraryRoot(): string { + return getLibraryRootOverride() ?? defaultLibraryRoot() +} diff --git a/electron/ipc/file-handlers.ts b/electron/ipc/file-handlers.ts index 45d46b76e..1717134ce 100644 --- a/electron/ipc/file-handlers.ts +++ b/electron/ipc/file-handlers.ts @@ -207,6 +207,14 @@ export function registerFileHandlers(): void { return true }) + handle('openTwitterCompose', async ({ text }) => { + const { shell } = await import('electron') + const url = new URL('https://twitter.com/intent/tweet') + if (text) url.searchParams.set('text', text) + await shell.openExternal(url.toString()) + return true + }) + handle('openParentFolderOfFile', async ({ filePath }) => { const { shell } = await import('electron') const normalizedPath = validatePath(filePath, getAllowedRoots()) @@ -239,15 +247,26 @@ export function registerFileHandlers(): void { handle('showSaveDialog', async ({ title, defaultPath, filters }) => { const mainWindow = getMainWindow() - if (!mainWindow) return null - const result = await dialog.showSaveDialog(mainWindow, { - title: title || 'Save File', - defaultPath, - filters: filters || [], - }) - if (result.canceled || !result.filePath) return null - approvePath(result.filePath) - return result.filePath + if (!mainWindow) { + // Indistinguishable from "user cancelled" to the renderer (the IPC + // contract returns null either way), but this case is a real bug, not + // a cancellation — log it so it's diagnosable instead of silent. + logger.error('showSaveDialog: no main window available') + return null + } + try { + const result = await dialog.showSaveDialog(mainWindow, { + title: title || 'Save File', + defaultPath, + filters: filters || [], + }) + if (result.canceled || !result.filePath) return null + approvePath(result.filePath) + return result.filePath + } catch (error) { + logger.error(`showSaveDialog failed: ${error}`) + return null + } }) handle('saveFile', async ({ filePath, data, encoding }) => { @@ -276,6 +295,19 @@ export function registerFileHandlers(): void { } }) + handle('copyFileToPath', async ({ srcPath, destPath }) => { + try { + validatePath(srcPath, getAllowedRoots()) + // destPath is a user-chosen save location approved via showSaveDialog. + validatePath(destPath, getAllowedRoots()) + fs.copyFileSync(srcPath, destPath) + return { success: true, path: destPath } + } catch (error) { + logger.error(`Error copying file: ${error}`) + return { success: false, error: String(error) } + } + }) + handle('showOpenDirectoryDialog', async ({ title }) => { const mainWindow = getMainWindow() if (!mainWindow) return null diff --git a/electron/ipc/library-handlers.ts b/electron/ipc/library-handlers.ts new file mode 100644 index 000000000..a3c0d7306 --- /dev/null +++ b/electron/ipc/library-handlers.ts @@ -0,0 +1,242 @@ +import { dialog, shell } from 'electron' +import path from 'path' +import fs from 'fs' +import { getAllowedRoots } from '../config' +import { getMainWindow } from '../window' +import { validatePath } from '../path-validation' +import { logger } from '../logger' +import { handle } from './typed-handle' +import { defaultLibraryRoot, resolveLibraryRoot, setLibraryRootOverride } from '../gpm-library-root' + +// Real-filesystem media library for the Prompt Manager Pro "Downloads Browser". +// Lives under the OS Downloads dir (an allowed root) so files are real and the +// user can open the folder in their file manager. + +const VIDEO_EXT = new Set(['.mp4', '.webm', '.mkv', '.mov', '.avi']) +const AUDIO_EXT = new Set(['.mp3', '.wav', '.ogg', '.aac', '.flac', '.m4a']) +const IMAGE_EXT = new Set(['.png', '.jpg', '.jpeg', '.webp', '.gif', '.bmp']) +const MEDIA_EXT = new Set([...IMAGE_EXT, ...VIDEO_EXT, ...AUDIO_EXT]) +const MIME_BY_EXT: Record = { + '.png': 'image/png', '.jpg': 'image/jpeg', '.jpeg': 'image/jpeg', '.webp': 'image/webp', + '.gif': 'image/gif', '.bmp': 'image/bmp', + '.mp4': 'video/mp4', '.webm': 'video/webm', '.mkv': 'video/x-matroska', '.mov': 'video/quicktime', '.avi': 'video/x-msvideo', + '.mp3': 'audio/mpeg', '.wav': 'audio/wav', '.ogg': 'audio/ogg', '.aac': 'audio/aac', '.flac': 'audio/flac', '.m4a': 'audio/mp4', +} + +function libRoot(): string { + const root = resolveLibraryRoot() + fs.mkdirSync(root, { recursive: true }) + return root +} + +// Watch the library root so the Studio Assets panel refreshes the moment files +// are added/removed on disk (e.g. a "Save video" into the folder), instead of +// only on manual refresh. Debounced because fs.watch fires several events per +// operation; recursive so changes inside subfolders (Inbox, etc.) are seen. +let libWatcher: fs.FSWatcher | null = null +let watchDebounce: ReturnType | null = null + +function notifyLibChanged(): void { + if (watchDebounce) clearTimeout(watchDebounce) + watchDebounce = setTimeout(() => { + getMainWindow()?.webContents.send('gpm-lib-changed') + }, 250) +} + +function startLibWatch(): void { + if (libWatcher) { libWatcher.close(); libWatcher = null } + try { + libWatcher = fs.watch(libRoot(), { recursive: true }, () => notifyLibChanged()) + } catch (e) { + // recursive watch is unsupported on some platforms (e.g. Linux) — the panel + // still works via manual Refresh; just no live updates there. + logger.warn(`gpm library watch failed: ${e}`) + } +} + +function safeName(name: string): string { + if (!name || name.includes('/') || name.includes('\\') || name.includes('..')) { + throw new Error(`Invalid name: ${name}`) + } + return name.trim() +} + +/** Ensure a named library subfolder exists and return its absolute path. Used by + * features that write generated assets straight into the library (e.g. "Continue + * as new shot" saving the extracted frame + trimmed clip into Continuations). */ +export function ensureLibFolder(name: string): string { + const p = folderPath(name) + fs.mkdirSync(p, { recursive: true }) + return p +} + +function folderPath(name: string): string { + const p = path.join(libRoot(), safeName(name)) + validatePath(p, getAllowedRoots()) + return p +} + +function uniqueTarget(target: string): string { + if (!fs.existsSync(target)) return target + const dir = path.dirname(target) + const ext = path.extname(target) + const base = path.basename(target, ext) + let i = 1 + let candidate = path.join(dir, `${base}-${i}${ext}`) + while (fs.existsSync(candidate)) { i++; candidate = path.join(dir, `${base}-${i}${ext}`) } + return candidate +} + +function escapeRegExp(s: string): string { + return s.replace(/[.*+?^${}()|[\]\\]/g, '\\$&') +} + +// Content-aware auto-naming: normalize a caller-supplied subject to safe filename chars. +function sanitizeSubject(subject: string): string { + return subject + .toLowerCase() + .replace(/[^a-z0-9-]+/g, '-') + .replace(/-{2,}/g, '-') + .replace(/^-+|-+$/g, '') + .slice(0, 48) + .replace(/-+$/g, '') +} + +// Next "-NN." for a folder — the sequence source of truth is the folder itself +// (append-only: existing NN are never renumbered). Scanned per copy so a batch increments. +function nextSequencedName(dir: string, subject: string, ext: string): string { + let max = 0 + const re = new RegExp(`^${escapeRegExp(subject)}-(\\d+)${escapeRegExp(ext)}$`, 'i') + try { + for (const name of fs.readdirSync(dir)) { + const m = re.exec(name) + if (m) max = Math.max(max, parseInt(m[1], 10)) + } + } catch { /* folder is new / unreadable — start the sequence at 01 */ } + return `${subject}-${String(max + 1).padStart(2, '0')}${ext}` +} + +export type LibraryListing = { + root: string + folders: string[] + files: Array<{ folder: string; name: string; path: string; isVideo: boolean; isAudio: boolean; mtimeMs: number }> +} + +/** Snapshot of the Studio Assets library (folders + media files). Shared by the + * gpmLibList IPC and the RiX MCP server's library search. */ +export function listLibrary(): LibraryListing { + const root = libRoot() + let folders = fs.readdirSync(root, { withFileTypes: true }).filter((d) => d.isDirectory()).map((d) => d.name) + if (folders.length === 0) { + fs.mkdirSync(path.join(root, 'Inbox'), { recursive: true }) + folders = ['Inbox'] + } + const files: LibraryListing['files'] = [] + for (const folder of folders) { + const fp = path.join(root, folder) + for (const entry of fs.readdirSync(fp, { withFileTypes: true })) { + if (!entry.isFile()) continue + const ext = path.extname(entry.name).toLowerCase() + if (!MEDIA_EXT.has(ext)) continue + const full = path.join(fp, entry.name) + files.push({ folder, name: entry.name, path: full, isVideo: VIDEO_EXT.has(ext), isAudio: AUDIO_EXT.has(ext), mtimeMs: fs.statSync(full).mtimeMs }) + } + } + return { root, folders, files } +} + +export function registerLibraryHandlers(): void { + handle('gpmLibList', () => listLibrary()) + + handle('gpmLibCreateFolder', ({ name }) => { + try { fs.mkdirSync(folderPath(name)); return { success: true as const } } + catch (e) { return { success: false as const, error: e instanceof Error ? e.message : 'create failed' } } + }) + + handle('gpmLibRenameFolder', ({ from, to }) => { + try { fs.renameSync(folderPath(from), folderPath(to)); return { success: true as const } } + catch (e) { return { success: false as const, error: e instanceof Error ? e.message : 'rename failed' } } + }) + + handle('gpmLibDeleteFolder', ({ name }) => { + try { fs.rmSync(folderPath(name), { recursive: true, force: true }); return { success: true as const } } + catch (e) { return { success: false as const, error: e instanceof Error ? e.message : 'delete failed' } } + }) + + handle('gpmLibAddFiles', ({ folder, srcPaths, baseName }) => { + try { + const dest = folderPath(folder) + fs.mkdirSync(dest, { recursive: true }) // robust: quick-save target may have been deleted + // Content-aware auto-naming: with a subject, land as "-NN." (sequenced per + // folder); without one (imports / drag-drop), keep the source's original filename. + const subject = baseName ? sanitizeSubject(baseName) : '' + let added = 0 + for (const src of srcPaths) { + if (!fs.existsSync(src) || !fs.statSync(src).isFile()) continue + const ext = path.extname(src).toLowerCase() + if (!MEDIA_EXT.has(ext)) continue + const targetName = subject ? nextSequencedName(dest, subject, ext) : path.basename(src) + fs.copyFileSync(src, uniqueTarget(path.join(dest, targetName))) + added++ + } + return { success: true as const, added } + } catch (e) { + return { success: false as const, error: e instanceof Error ? e.message : 'add failed' } + } + }) + + handle('gpmLibMoveFile', ({ fromFolder, name, toFolder }) => { + try { + const src = path.join(folderPath(fromFolder), safeName(name)) + const target = uniqueTarget(path.join(folderPath(toFolder), safeName(name))) + fs.renameSync(src, target) + return { success: true as const } + } catch (e) { return { success: false as const, error: e instanceof Error ? e.message : 'move failed' } } + }) + + handle('gpmLibDeleteFile', ({ folder, name }) => { + try { fs.rmSync(path.join(folderPath(folder), safeName(name)), { force: true }); return { success: true as const } } + catch (e) { return { success: false as const, error: e instanceof Error ? e.message : 'delete failed' } } + }) + + handle('gpmLibReveal', ({ folder }) => { + try { void shell.openPath(folder ? folderPath(folder) : libRoot()); return { success: true as const } } + catch (e) { logger.warn(`gpmLibReveal failed: ${e}`); return { success: false as const, error: 'reveal failed' } } + }) + + handle('gpmLibReadAsDataUrl', ({ path: p }) => { + try { + validatePath(p, getAllowedRoots()) + const ext = path.extname(p).toLowerCase() + const mime = MIME_BY_EXT[ext] ?? 'application/octet-stream' + const b64 = fs.readFileSync(p).toString('base64') + return { success: true as const, dataUrl: `data:${mime};base64,${b64}` } + } catch (e) { + return { success: false as const, error: e instanceof Error ? e.message : 'read failed' } + } + }) + + handle('gpmLibGetRoot', () => ({ root: libRoot(), isDefault: resolveLibraryRoot() === defaultLibraryRoot() })) + + handle('gpmLibChooseRoot', async () => { + const win = getMainWindow() + const result = win + ? await dialog.showOpenDialog(win, { title: 'Choose Downloads Browser folder', properties: ['openDirectory', 'createDirectory'] }) + : await dialog.showOpenDialog({ title: 'Choose Downloads Browser folder', properties: ['openDirectory', 'createDirectory'] }) + if (result.canceled || result.filePaths.length === 0) return { root: null } + const chosen = result.filePaths[0] + setLibraryRootOverride(chosen) + fs.mkdirSync(chosen, { recursive: true }) + startLibWatch() // re-point the watcher at the newly chosen folder + return { root: chosen } + }) + + handle('gpmLibResetRoot', () => { + setLibraryRootOverride(null) + const root = libRoot() + startLibWatch() + return { root } + }) + + startLibWatch() +} diff --git a/electron/ipc/video-processing-handlers.ts b/electron/ipc/video-processing-handlers.ts index 7da9ce6fe..f381bc975 100644 --- a/electron/ipc/video-processing-handlers.ts +++ b/electron/ipc/video-processing-handlers.ts @@ -1,16 +1,63 @@ -import { extractVideoFrameToFile } from '../export/ffmpeg-utils' +import path from 'path' +import { + extractVideoFrameToFile, + extractLastFrameToFile, + trimFirstFrameToFile, + getVideoDimensions, + getVideoFps, + findFfmpegPath, +} from '../export/ffmpeg-utils' +import { ensureLibFolder } from './library-handlers' import { handle } from './typed-handle' +const CONTINUATIONS_FOLDER = 'Continuations' + +function stamp(): string { + return new Date().toISOString().replace(/[:.]/g, '-').slice(0, 19) +} + export function registerVideoProcessingHandlers(): void { - handle('extractVideoFrame', async ({ videoPath, seekTime, width, quality }) => { + handle('extractVideoFrame', async ({ videoPath, seekTime, width, quality, outputPath }) => { return { path: extractVideoFrameToFile({ videoPath, seekTime, width, - quality: quality ?? 2, + // When writing to a user-chosen file (Save frame), don't apply the + // thumbnail JPEG quality knob — ffmpeg infers the format (e.g. PNG) + // from the outputPath extension and -q:v would be meaningless. + quality: outputPath ? undefined : quality ?? 2, + outputPath, + // A user-chosen outputPath means "Save frame" — grab the exact frame. + accurate: Boolean(outputPath), timeoutMs: 10000, }), } }) + + // "Continue as new shot" — extract the clip's last frame into Continuations to + // seed an i2v continuation, and report the source dims/fps so the next gen matches. + handle('continuationExtractLastFrame', async ({ videoPath, seekTime }) => { + const dir = ensureLibFolder(CONTINUATIONS_FOLDER) + const framePath = path.join(dir, `continuation_frame_${stamp()}.png`) + if (seekTime != null) { + // Scrubbed "Continue from this frame": grab the exact frame (accurate seek). + extractVideoFrameToFile({ videoPath, seekTime, outputPath: framePath, accurate: true, timeoutMs: 15000 }) + } else { + extractLastFrameToFile({ videoPath, outputPath: framePath }) + } + const { width, height } = getVideoDimensions(videoPath) + const ffmpegPath = findFfmpegPath() + const fps = ffmpegPath ? getVideoFps(ffmpegPath, videoPath) : 24 + return { framePath, width, height, fps } + }) + + // Drop the duplicate lead frame from a freshly generated continuation and save + // the clean clip into Continuations. + handle('continuationSaveTrimmed', async ({ videoPath, colorMatchReference }) => { + const dir = ensureLibFolder(CONTINUATIONS_FOLDER) + const outPath = path.join(dir, `continuation_${stamp()}.mp4`) + trimFirstFrameToFile({ videoPath, outputPath: outPath, colorMatchReferencePath: colorMatchReference }) + return { path: outPath } + }) } diff --git a/electron/main.ts b/electron/main.ts index 68f614eb7..570dd8884 100644 --- a/electron/main.ts +++ b/electron/main.ts @@ -6,13 +6,16 @@ import { registerExportHandlers } from './export/export-handler' import { stopExportProcess } from './export/ffmpeg-utils' import { registerAppHandlers } from './ipc/app-handlers' import { registerFileHandlers } from './ipc/file-handlers' +import { registerLibraryHandlers } from './ipc/library-handlers' import { registerLogHandlers } from './ipc/log-handlers' import { registerVideoProcessingHandlers } from './ipc/video-processing-handlers' import { logger } from './logger' +import { startMcpServer } from './mcp/mcp-server' import { initSessionLog } from './logging-management' import { stopPythonBackend } from './python-backend' import { initAutoUpdater } from './updater' import { createWindow, getMainWindow } from './window' +import { createAppMenu, registerEditContextMenu } from './menu' import { sendAnalyticsEvent } from './analytics' function logAppVersion(): void { @@ -33,9 +36,11 @@ if (!gotLock) { registerAppHandlers() registerFileHandlers() + registerLibraryHandlers() registerLogHandlers() registerExportHandlers() registerVideoProcessingHandlers() + startMcpServer() app.on('second-instance', () => { const mainWindow = getMainWindow() @@ -50,13 +55,17 @@ if (!gotLock) { return } if (app.isReady()) { - createWindow() + const window = createWindow() + createAppMenu(window) + registerEditContextMenu(window) } }) app.whenReady().then(async () => { setupCSP() - createWindow() + const mainWindow = createWindow() + createAppMenu(mainWindow) + registerEditContextMenu(mainWindow) initAutoUpdater() // Python setup + backend start are now driven by the renderer via IPC @@ -73,7 +82,9 @@ if (!gotLock) { app.on('activate', () => { if (getMainWindow() === null) { - createWindow() + const window = createWindow() + createAppMenu(window) + registerEditContextMenu(window) } }) diff --git a/electron/mcp/editor-bridge.ts b/electron/mcp/editor-bridge.ts new file mode 100644 index 000000000..5e632cd53 --- /dev/null +++ b/electron/mcp/editor-bridge.ts @@ -0,0 +1,40 @@ +import { randomUUID } from 'crypto' +import { getMainWindow } from '../window' +import { handle } from '../ipc/typed-handle' + +// Main → renderer RPC for editor tools. The editor's state lives in a zustand +// store inside the renderer, so timeline reads/edits are forwarded there +// ('mcp-editor-request') and answered via the mcpEditorResponse IPC. Going +// through the live store (never writing project files directly) keeps edits +// on the same autosave/undo path as manual edits — no stale-snapshot clobbers. + +type Pending = { resolve: (v: unknown) => void; reject: (e: Error) => void; timer: ReturnType } +const pending = new Map() + +export function registerEditorBridge(): void { + handle('mcpEditorResponse', ({ id, ok, result, error }) => { + const p = pending.get(id) + if (!p) return + pending.delete(id) + clearTimeout(p.timer) + if (ok) p.resolve(result) + else p.reject(new Error(error || 'Editor request failed')) + }) +} + +export function callEditor(tool: string, args: Record = {}, timeoutMs = 60000): Promise { + const win = getMainWindow() + if (!win || win.isDestroyed()) return Promise.reject(new Error('RiX window is not open')) + const id = randomUUID() + return new Promise((resolve, reject) => { + const timer = setTimeout(() => { + pending.delete(id) + reject(new Error( + `RiX editor did not answer "${tool}" within ${Math.round(timeoutMs / 1000)}s. ` + + 'Open a project in RiX and visit the Video Editor tab once so the editor is loaded.', + )) + }, timeoutMs) + pending.set(id, { resolve: resolve as (v: unknown) => void, reject, timer }) + win.webContents.send('mcp-editor-request', { id, tool, args }) + }) +} diff --git a/electron/mcp/mcp-server.ts b/electron/mcp/mcp-server.ts new file mode 100644 index 000000000..57ecb7abf --- /dev/null +++ b/electron/mcp/mcp-server.ts @@ -0,0 +1,377 @@ +import http from 'http' +import fs from 'fs' +import os from 'os' +import path from 'path' +import { randomBytes } from 'crypto' +import { app } from 'electron' +import { isDev } from '../config' +import { logger } from '../logger' +import { listLibrary, ensureLibFolder } from '../ipc/library-handlers' +import { exportTimelineNative, type ExportNativeInput } from '../export/export-handler' +import { callEditor, registerEditorBridge } from './editor-bridge' +import { analyzeAudio, probeMedia, sampleFrames, type Frame } from './media-tools' + +/** + * RiX MCP server — lets an MCP client (Claude Code / Claude Desktop) assemble + * videos in the RiX Video Editor from the user's assets. + * + * Transport: MCP Streamable HTTP, stateless, JSON responses only, at + * http://127.0.0.1:/mcp. Guarded by a bearer token stored in + * /rix-mcp.json, loopback-only binding, a Host check (DNS rebinding) + * and refusal of browser-originated requests. Only started in dev builds or + * when RIX_MCP=1, so shipped installs never open a control port by default. + */ + +const DEFAULT_PORT = 47821 +const PROTOCOL_VERSIONS = ['2025-06-18', '2025-03-26', '2024-11-05'] + +type ToolContent = Array<{ type: 'text'; text: string } | { type: 'image'; data: string; mimeType: string }> +type Tool = { + name: string + description: string + inputSchema: Record + run: (args: Record) => Promise +} + +const text = (value: unknown): ToolContent => [{ type: 'text', text: typeof value === 'string' ? value : JSON.stringify(value, null, 2) }] +const withFrames = (summary: unknown, frames: Frame[]): ToolContent => [ + ...text(summary), + ...frames.flatMap(f => [ + ...(f.time !== undefined ? [{ type: 'text' as const, text: `t=${f.time}s` }] : []), + { type: 'image' as const, data: f.base64, mimeType: f.mimeType }, + ]), +] +const str = (v: unknown) => (typeof v === 'string' ? v : undefined) +const numArg = (v: unknown) => (typeof v === 'number' && Number.isFinite(v) ? v : undefined) + +let lastPreviewPath: string | null = null + +const PROJECT_ID_PROP = { type: 'string', description: 'editor.project.id from rix_status — guards against editing the wrong project' } + +async function renderTimeline(opts: { preview: boolean; width?: number; height?: number; fps?: number; quality?: number; outputPath?: string }) { + const { payload, aspect, duration, timelineName } = await callEditor<{ + payload: Omit + aspect: number + duration: number + timelineName: string + }>('export_payload') + const even = (n: number) => Math.max(2, Math.round(n / 2) * 2) + const long = opts.preview ? 640 : 1920 + let width = opts.width ?? (aspect >= 1 ? long : even(long * aspect)) + let height = opts.height ?? (aspect >= 1 ? even(long / aspect) : long) + if (opts.width && !opts.height) height = even(opts.width / aspect) + if (opts.height && !opts.width) width = even(opts.height * aspect) + const outputPath = opts.outputPath ?? path.join(os.tmpdir(), `rix-mcp-preview-${Date.now()}.mp4`) + const r = await exportTimelineNative( + { ...payload, outputPath, codec: 'h264', width: even(width), height: even(height), fps: opts.fps ?? 24, quality: opts.quality ?? 18 }, + { preview: opts.preview }, + ) + if (!r.success) throw new Error(`Render failed: ${r.error}`) + return { outputPath, duration, width: even(width), height: even(height), timelineName } +} + +const TOOLS: Tool[] = [ + { + name: 'rix_status', + description: 'Check that RiX is reachable and which project/timelines are open in the Video Editor. Call this first. ' + + 'Every editing tool requires the returned editor.project.id as "project_id" (edits are refused if the user has switched projects).', + inputSchema: { type: 'object', properties: {} }, + run: async () => { + let editor: unknown + try { editor = await callEditor('status', {}, 5000) } catch (e) { editor = { error: (e as Error).message } } + return text({ app: `RiX ${app.getVersion()}`, libraryRoot: listLibrary().root, editor }) + }, + }, + { + name: 'list_library', + description: 'Browse the Studio Assets library (the user\'s media folders on disk). With no folder, returns every folder with file counts. ' + + 'With a folder and/or query, returns matching files (newest first). File names are content-aware (e.g. storm-03.mp4), so query by subject.', + inputSchema: { + type: 'object', + properties: { + folder: { type: 'string', description: 'Folder name (case-insensitive), e.g. "Halloween"' }, + query: { type: 'string', description: 'Substring match on file name' }, + type: { type: 'string', enum: ['image', 'video', 'audio'] }, + limit: { type: 'number', description: 'Max files (default 100)' }, + }, + }, + run: async (args) => { + const lib = listLibrary() + const folder = str(args.folder)?.toLowerCase() + const query = str(args.query)?.toLowerCase() + const type = str(args.type) + if (!folder && !query && !type) { + const counts = new Map() + for (const f of lib.files) counts.set(f.folder, (counts.get(f.folder) ?? 0) + 1) + return text({ root: lib.root, folders: lib.folders.map(name => ({ name, files: counts.get(name) ?? 0 })) }) + } + const kind = (f: (typeof lib.files)[number]) => (f.isVideo ? 'video' : f.isAudio ? 'audio' : 'image') + const files = lib.files + .filter(f => !folder || f.folder.toLowerCase() === folder) + .filter(f => !query || f.name.toLowerCase().includes(query)) + .filter(f => !type || kind(f) === type) + .sort((a, b) => b.mtimeMs - a.mtimeMs) + const limit = Math.max(1, Math.min(numArg(args.limit) ?? 100, 1000)) + return text({ + total: files.length, + files: files.slice(0, limit).map(f => ({ folder: f.folder, name: f.name, type: kind(f), path: f.path, modified: new Date(f.mtimeMs).toISOString() })), + }) + }, + }, + { + name: 'list_project_assets', + description: 'List assets registered in the open RiX project (generations + imports), newest first, with their generation prompts. ' + + 'Prompts describe content — use them to pick shots.', + inputSchema: { + type: 'object', + properties: { + query: { type: 'string', description: 'Substring match on prompt, path or folder/tag name' }, + type: { type: 'string', enum: ['image', 'video', 'audio'] }, + limit: { type: 'number', description: 'Max assets (default 50)' }, + }, + }, + run: async (args) => text(await callEditor('list_project_assets', args)), + }, + { + name: 'inspect_media', + description: 'LOOK at a media file: returns metadata plus sampled frames as images (evenly spaced across a video, or the still itself). ' + + 'Use before choosing shots — prompts can lie; frames don\'t. Default 6 frames at 384px.', + inputSchema: { + type: 'object', + properties: { + path: { type: 'string', description: 'Absolute path to an image or video' }, + frames: { type: 'number', description: 'How many evenly spaced frames (1–16, default 6)' }, + times: { type: 'array', items: { type: 'number' }, description: 'Explicit timestamps in seconds (overrides frames)' }, + width: { type: 'number', description: 'Frame width in px (default 384)' }, + }, + required: ['path'], + }, + run: async (args) => { + const p = str(args.path) + if (!p) throw new Error('"path" is required') + const times = Array.isArray(args.times) ? args.times.filter((t): t is number => typeof t === 'number') : undefined + const { info, frames } = sampleFrames(p, { count: numArg(args.frames), times, width: numArg(args.width) }) + return withFrames({ path: p, ...info }, frames) + }, + }, + { + name: 'analyze_audio', + description: 'Hear a music/audio file: tempo (bpm), a beat grid, guessed downbeats (every 4th beat), the strongest hits, ' + + 'and loudness per second (0–1) to find intros, builds and drops. Use it to place cuts on the beat.', + inputSchema: { type: 'object', properties: { path: { type: 'string' } }, required: ['path'] }, + run: async (args) => { + const p = str(args.path) + if (!p) throw new Error('"path" is required') + return text(analyzeAudio(p)) + }, + }, + { + name: 'get_timeline', + description: 'Read the active timeline: tracks (index, kind) and every clip with id, track, start/end, source in-point, asset, text, transitions, audio settings.', + inputSchema: { type: 'object', properties: {} }, + run: async () => text(await callEditor('get_timeline')), + }, + { + name: 'apply_edits', + description: [ + 'Edit the active timeline with a batch of ops. The batch is ATOMIC and is ONE undo step (if any op fails, nothing changes).', + 'Times are seconds. Returns per-op results (new clip ids) and the resulting timeline.', + 'Ops:', + ' {"op":"import","path":"","ref":"name"} — register a library file in the project (ref usable later as "$name")', + ' {"op":"add_clip","asset":"","track":0,"start":0,"in":1.5,"duration":2.0,"speed":1,"with_audio":true,"ref":"c1"}', + ' start omitted = append at end of track; in = source in-point; duration capped to the source length. Placing over existing clips overwrites them.', + ' Video clips bring linked audio on an audio track; with_audio:false drops it; audio_track:N routes it (e.g. dialogue to A2 so it', + ' can\'t overwrite the music on A1). Images default to 5s unless duration is set.', + ' {"op":"duck","clip":"","ranges":[[t0,t1],...],"level":0.25,"attack":0.2,"release":0.4} — dip the music under', + ' dialogue (ranges = TIMELINE seconds of speech; take phrase times from inspect_media soundSegments + the clip\'s start/in).', + ' {"op":"add_text","text":"TITLE","start":0,"duration":3,"style":{"fontSize":96,"color":"#FFFFFF","positionX":50,"positionY":50},"fade_in":0.5,"fade_out":0.5}', + ' {"op":"update_clip","clip":"","set":{"start","duration","in","speed","volume"(0–2),"opacity"(0–100),"audio_fade_in","audio_fade_out",', + ' "text_fade_in","text_fade_out","muted","reversed","color":{brightness,contrast,saturation,temperature,tint,exposure,highlights,shadows},"text","style",', + ' "volume_keyframes":[{"t":clipSeconds,"value":0–2}]}}', + ' {"op":"transition","clip":"","in":{"type":"dissolve","duration":0.5},"out":{"type":"fade-to-black","duration":1}}', + ' types: none, dissolve, fade-to-black, fade-to-white, wipe-left, wipe-right, wipe-up, wipe-down. A dissolve needs out on the left clip AND in on the right clip.', + ' {"op":"delete","clips":[""]} (linked audio goes too) {"op":"clear"} (empty the active timeline)', + ' {"op":"new_timeline","name":"Trailer v1"} (becomes active) {"op":"add_track","kind":"video"|"audio"}', + ].join('\n'), + inputSchema: { + type: 'object', + properties: { + project_id: PROJECT_ID_PROP, + ops: { type: 'array', items: { type: 'object' }, description: 'Ordered list of edit ops (see description)' }, + }, + required: ['project_id', 'ops'], + }, + run: async (args) => text(await callEditor('apply_edits', args, 180000)), + }, + { + name: 'undo', + description: 'Revert YOUR last apply_edits batch (one step). Refuses if the user has changed anything since, or if the last ' + + 'change was not yours — it never undoes the user\'s work. In that case fix things with a corrective apply_edits batch.', + inputSchema: { type: 'object', properties: { project_id: PROJECT_ID_PROP }, required: ['project_id'] }, + run: async (args) => text(await callEditor('undo', args)), + }, + { + name: 'redo', + description: 'Re-apply a batch you just reverted with undo (same safety rules).', + inputSchema: { type: 'object', properties: { project_id: PROJECT_ID_PROP }, required: ['project_id'] }, + run: async (args) => text(await callEditor('redo', args)), + }, + { + name: 'render_frames', + description: 'Render the active timeline to a fast low-res preview (with transitions, text, color) and return frames at the given times, ' + + 'so you can review your own cut. Default: 8 evenly spaced frames. Also returns the preview mp4 path (usable with inspect_media).', + inputSchema: { + type: 'object', + properties: { + times: { type: 'array', items: { type: 'number' }, description: 'Timeline seconds to grab (max 16)' }, + count: { type: 'number', description: 'Evenly spaced frame count when times is omitted (default 8)' }, + }, + }, + run: async (args) => { + const render = await renderTimeline({ preview: true }) + if (lastPreviewPath && lastPreviewPath !== render.outputPath) { try { fs.unlinkSync(lastPreviewPath) } catch { /* temp */ } } + lastPreviewPath = render.outputPath + const times = Array.isArray(args.times) ? args.times.filter((t): t is number => typeof t === 'number') : undefined + const { frames } = sampleFrames(render.outputPath, { times, count: numArg(args.count) ?? 8, width: 480 }) + return withFrames({ previewPath: render.outputPath, duration: render.duration, size: `${render.width}x${render.height}` }, frames) + }, + }, + { + name: 'export_video', + description: 'Export the active timeline at full quality (H.264 MP4). Default output: Studio Assets/"RiX Edits"/-.mp4, ' + + 'sized from the first clip\'s aspect at 1080p. Takes a while for long timelines.', + inputSchema: { + type: 'object', + properties: { + output_path: { type: 'string', description: 'Absolute .mp4 path (must be inside the library, Downloads, or project folders)' }, + width: { type: 'number' }, + height: { type: 'number' }, + fps: { type: 'number', description: 'Default 24' }, + quality: { type: 'number', description: 'H.264 CRF, lower = better (default 18)' }, + }, + }, + run: async (args) => { + const stamp = new Date().toISOString().replace(/[:.]/g, '-').slice(0, 19) + let outputPath = str(args.output_path) + if (!outputPath) { + const { timelineName } = await callEditor<{ timelineName: string }>('export_payload') + const safe = timelineName.replace(/[^\w -]+/g, '').trim().replace(/\s+/g, '-') || 'timeline' + outputPath = path.join(ensureLibFolder('RiX Edits'), `${safe}-${stamp}.mp4`) + } + const r = await renderTimeline({ + preview: false, outputPath, width: numArg(args.width), height: numArg(args.height), fps: numArg(args.fps), quality: numArg(args.quality), + }) + return text({ exported: r.outputPath, duration: r.duration, size: `${r.width}x${r.height}`, media: probeMedia(r.outputPath) }) + }, + }, +] + +// ---------------------------------------------------------------- JSON-RPC + +type RpcMessage = { jsonrpc: '2.0'; id?: string | number | null; method?: string; params?: Record } +const rpcResult = (id: RpcMessage['id'], result: unknown) => ({ jsonrpc: '2.0', id, result }) +const rpcError = (id: RpcMessage['id'], code: number, message: string) => ({ jsonrpc: '2.0', id: id ?? null, error: { code, message } }) + +async function handleMessage(msg: RpcMessage): Promise { + if (msg.id === undefined || msg.id === null) return null // notification (e.g. notifications/initialized) + switch (msg.method) { + case 'initialize': { + const requested = str(msg.params?.protocolVersion) + return rpcResult(msg.id, { + protocolVersion: requested && PROTOCOL_VERSIONS.includes(requested) ? requested : PROTOCOL_VERSIONS[0], + capabilities: { tools: {} }, + serverInfo: { name: 'rix-editor', version: app.getVersion() }, + instructions: 'Tools for assembling videos in the RiX Video Editor from the user\'s generated assets. ' + + 'Start with rix_status. Look at media with inspect_media before cutting; review your cut with render_frames before exporting.', + }) + } + case 'ping': + return rpcResult(msg.id, {}) + case 'tools/list': + return rpcResult(msg.id, { tools: TOOLS.map(({ name, description, inputSchema }) => ({ name, description, inputSchema })) }) + case 'tools/call': { + const name = str(msg.params?.name) + const tool = TOOLS.find(t => t.name === name) + if (!tool) return rpcError(msg.id, -32602, `Unknown tool: ${name}`) + const args = (msg.params?.arguments ?? {}) as Record + const started = Date.now() + try { + const content = await tool.run(args) + logger.info(`[MCP] ${name} ok (${Date.now() - started}ms)`) + return rpcResult(msg.id, { content }) + } catch (e) { + const message = e instanceof Error ? e.message : String(e) + logger.warn(`[MCP] ${name} failed: ${message}`) + return rpcResult(msg.id, { content: text(`Error: ${message}`), isError: true }) + } + } + default: + return rpcError(msg.id, -32601, `Method not found: ${msg.method}`) + } +} + +// ---------------------------------------------------------------- HTTP + +type McpConfig = { port: number; token: string } + +function loadConfig(): McpConfig { + const file = path.join(app.getPath('userData'), 'rix-mcp.json') + let cfg: Partial = {} + try { cfg = JSON.parse(fs.readFileSync(file, 'utf8')) } catch { /* first run */ } + const envPort = Number(process.env.RIX_MCP_PORT) + const next: McpConfig = { + port: Number.isInteger(envPort) && envPort > 0 ? envPort : cfg.port ?? DEFAULT_PORT, + token: cfg.token ?? randomBytes(24).toString('hex'), + } + if (next.token !== cfg.token || next.port !== cfg.port) { + fs.writeFileSync(file, JSON.stringify(next, null, 2), { encoding: 'utf8', mode: 0o600 }) + } + return next +} + +function readBody(req: http.IncomingMessage, limit = 5 * 1024 * 1024): Promise { + return new Promise((resolve, reject) => { + let size = 0 + const chunks: Buffer[] = [] + req.on('data', (c: Buffer) => { + size += c.length + if (size > limit) { reject(new Error('Request too large')); req.destroy() } else chunks.push(c) + }) + req.on('end', () => resolve(Buffer.concat(chunks).toString('utf8'))) + req.on('error', reject) + }) +} + +export function startMcpServer(): void { + if (!isDev && process.env.RIX_MCP !== '1') return + registerEditorBridge() + const { port, token } = loadConfig() + const allowedHosts = new Set([`127.0.0.1:${port}`, `localhost:${port}`]) + + const server = http.createServer(async (req, res) => { + const send = (status: number, body?: unknown) => { + res.writeHead(status, body === undefined ? {} : { 'Content-Type': 'application/json' }) + res.end(body === undefined ? undefined : JSON.stringify(body)) + } + // Loopback-only + Host check (DNS rebinding) + no browser-originated calls + token. + if (!allowedHosts.has(req.headers.host ?? '')) return send(403, rpcError(null, -32000, 'Forbidden host')) + if (req.headers.origin) return send(403, rpcError(null, -32000, 'Browser requests are not allowed')) + if (req.headers.authorization !== `Bearer ${token}`) return send(401, rpcError(null, -32001, 'Missing or invalid bearer token')) + if (req.url?.split('?')[0] !== '/mcp') return send(404, rpcError(null, -32000, 'Not found')) + if (req.method === 'DELETE') return send(200) // stateless: no session to end + if (req.method !== 'POST') return send(405, rpcError(null, -32000, 'Use POST')) + + let parsed: RpcMessage | RpcMessage[] + try { parsed = JSON.parse(await readBody(req)) } catch { return send(400, rpcError(null, -32700, 'Parse error')) } + const batch = Array.isArray(parsed) + const replies = (await Promise.all((batch ? parsed : [parsed]).map(handleMessage))).filter((r): r is object => r !== null) + if (replies.length === 0) return send(202) + return send(200, batch ? replies : replies[0]) + }) + + server.on('error', (err) => logger.warn(`[MCP] server not started: ${err.message}`)) + server.listen(port, '127.0.0.1', () => { + logger.info(`[MCP] RiX editor MCP server listening on http://127.0.0.1:${port}/mcp (token in ${path.join(app.getPath('userData'), 'rix-mcp.json')})`) + }) + app.on('before-quit', () => server.close()) +} diff --git a/electron/mcp/media-tools.ts b/electron/mcp/media-tools.ts new file mode 100644 index 000000000..666a80a09 --- /dev/null +++ b/electron/mcp/media-tools.ts @@ -0,0 +1,195 @@ +import { spawnSync } from 'child_process' +import fs from 'fs' +import os from 'os' +import path from 'path' +import { extractVideoFrameToFile, findFfmpegPath } from '../export/ffmpeg-utils' + +// ffmpeg-backed perception helpers for the RiX MCP server: lets the agent +// actually look at clips (sampled frames) and hear music (beats/energy). + +const IMAGE_EXT = /\.(png|jpe?g|webp|gif|bmp)$/i +const AUDIO_EXT = /\.(mp3|wav|ogg|aac|flac|m4a)$/i + +function ffmpeg(): string { + const p = findFfmpegPath() + if (!p) throw new Error('ffmpeg not found (the Python backend environment must be installed)') + return p +} + +export type MediaInfo = { + kind: 'image' | 'video' | 'audio' + duration?: number + width?: number + height?: number + fps?: number + hasAudio: boolean +} + +/** Parse ffmpeg's -i banner — cheap, and avoids depending on a separate ffprobe binary. */ +export function probeMedia(filePath: string): MediaInfo { + if (!fs.existsSync(filePath)) throw new Error(`File not found: ${filePath}`) + const r = spawnSync(ffmpeg(), ['-hide_banner', '-i', filePath], { encoding: 'utf8', timeout: 10000 }) + const out = `${r.stdout || ''}\n${r.stderr || ''}` + const dur = out.match(/Duration:\s*(\d+):(\d+):(\d+(?:\.\d+)?)/) + const videoLine = out.split('\n').find(l => l.includes('Video:')) + const dims = videoLine?.match(/(\d{2,5})x(\d{2,5})(?:[,\s[]|$)/) + const fps = videoLine?.match(/(\d+(?:\.\d+)?)\s*fps/) + const kind: MediaInfo['kind'] = IMAGE_EXT.test(filePath) ? 'image' : AUDIO_EXT.test(filePath) || !videoLine ? 'audio' : 'video' + return { + kind, + duration: dur && kind !== 'image' ? Number(dur[1]) * 3600 + Number(dur[2]) * 60 + parseFloat(dur[3]) : undefined, + width: dims ? Number(dims[1]) : undefined, + height: dims ? Number(dims[2]) : undefined, + fps: fps ? Number(fps[1]) : undefined, + hasAudio: out.includes('Audio:'), + } +} + +/** Where a file has sound (speech, cackles, SFX) vs silence, via ffmpeg silencedetect: + * [start, end] seconds of non-silence at ~0.1s precision. For a dialogue clip these + * are the phrases — cut on their boundaries so lines are never clipped mid-word. */ +export function soundSegments(filePath: string, duration: number, opts: { noiseDb?: number; minSilence?: number } = {}): Array<[number, number]> { + const r = spawnSync(ffmpeg(), [ + '-hide_banner', '-i', filePath, + '-af', `silencedetect=noise=${opts.noiseDb ?? -32}dB:d=${opts.minSilence ?? 0.25}`, + '-f', 'null', '-', + ], { encoding: 'utf8', timeout: 60000 }) + const out = `${r.stdout || ''}\n${r.stderr || ''}` + const r2 = (x: number) => Math.round(x * 100) / 100 + // silencedetect logs silence_start / silence_end in pairs; a trailing silence has no end. + const starts = [...out.matchAll(/silence_start: (-?[\d.]+)/g)].map(x => Math.max(0, Number(x[1]))) + const ends = [...out.matchAll(/silence_end: ([\d.]+)/g)].map(x => Number(x[1])) + const silences = starts.map((a, i): [number, number] => [a, ends[i] ?? duration]) + const segments: Array<[number, number]> = [] + let cursor = 0 + for (const [a, b] of silences) { + if (a - cursor > 0.05) segments.push([r2(cursor), r2(a)]) + cursor = Math.max(cursor, b) + } + if (duration - cursor > 0.05) segments.push([r2(cursor), r2(duration)]) + return segments +} + +export type Frame = { time?: number; base64: string; mimeType: 'image/jpeg' } + +/** Sample frames from a video (evenly spaced, or at explicit times), or a + * downscaled copy of a still. JPEG, `width` px wide. */ +export function sampleFrames(filePath: string, opts: { count?: number; times?: number[]; width?: number } = {}): { info: MediaInfo & { soundSegments?: Array<[number, number]> }; frames: Frame[] } { + const probed = probeMedia(filePath) + const info = probed.kind === 'video' && probed.hasAudio && probed.duration + ? { ...probed, soundSegments: soundSegments(filePath, probed.duration) } + : probed + const width = Math.max(64, Math.min(opts.width ?? 384, 1280)) + const tmp = () => path.join(os.tmpdir(), `rix_mcp_${Date.now()}_${Math.random().toString(36).slice(2, 8)}.jpg`) + const read = (p: string) => { const b = fs.readFileSync(p).toString('base64'); try { fs.unlinkSync(p) } catch { /* temp */ } return b } + + if (info.kind === 'audio') throw new Error('This is an audio file — use analyze_audio instead') + if (info.kind === 'image') { + const out = tmp() + const r = spawnSync(ffmpeg(), ['-y', '-i', filePath, '-vf', `scale=${width}:-2`, '-frames:v', '1', '-q:v', '4', out], { timeout: 15000 }) + if (r.status !== 0 || !fs.existsSync(out)) throw new Error('Could not decode image') + return { info, frames: [{ base64: read(out), mimeType: 'image/jpeg' }] } + } + + const duration = info.duration ?? 0 + const count = Math.max(1, Math.min(opts.count ?? 6, 16)) + const times = (opts.times?.length ? opts.times : Array.from({ length: count }, (_, i) => duration * (i + 0.5) / count)) + .slice(0, 16) + .map(t => Math.max(0, Math.min(t, Math.max(0, duration - 0.05)))) + const frames = times.map(time => { + const out = extractVideoFrameToFile({ videoPath: filePath, seekTime: time, width, quality: 4, outputPath: tmp(), timeoutMs: 15000 }) + return { time: Math.round(time * 100) / 100, base64: read(out), mimeType: 'image/jpeg' as const } + }) + return { info, frames } +} + +/** + * Beat/energy analysis for cutting to music. Decodes to 11.025 kHz mono, builds a + * half-wave-rectified energy-flux onset envelope (~23 ms hops), estimates tempo by + * autocorrelation in 70–180 BPM, then phase-aligns a beat grid to the onsets. + * Heuristic, not a DAW — good enough to place cuts on the pulse. + */ +export function analyzeAudio(filePath: string) { + const info = probeMedia(filePath) + if (!info.hasAudio) throw new Error('No audio stream in this file') + const SR = 11025 + const r = spawnSync(ffmpeg(), ['-v', 'error', '-i', filePath, '-ac', '1', '-ar', String(SR), '-f', 'f32le', '-'], { + timeout: 60000, + maxBuffer: 1024 * 1024 * 200, + }) + if (r.status !== 0 || !r.stdout?.length) throw new Error('Could not decode audio') + const buf = r.stdout as Buffer + const samples = new Float32Array(buf.buffer, buf.byteOffset, Math.floor(buf.length / 4)) + + const HOP = 256 // ~23 ms + const n = Math.floor(samples.length / HOP) + const energy = new Float32Array(n) + for (let i = 0; i < n; i++) { + let s = 0 + for (let j = i * HOP; j < (i + 1) * HOP; j++) s += samples[j] * samples[j] + energy[i] = Math.sqrt(s / HOP) + } + const onset = new Float32Array(n) + for (let i = 1; i < n; i++) onset[i] = Math.max(0, Math.log1p(energy[i] * 100) - Math.log1p(energy[i - 1] * 100)) + + const hopSec = HOP / SR + const minLag = Math.round(60 / 180 / hopSec) + const maxLag = Math.round(60 / 70 / hopSec) + let bestLag = minLag + let bestScore = -Infinity + for (let lag = minLag; lag <= maxLag; lag++) { + let score = 0 + for (let i = lag; i < n; i++) score += onset[i] * onset[i - lag] + score /= (n - lag) + if (score > bestScore) { bestScore = score; bestLag = lag } + } + let bestPhase = 0 + let bestPhaseScore = -Infinity + for (let phase = 0; phase < bestLag; phase++) { + let score = 0 + for (let i = phase; i < n; i += bestLag) score += onset[i] + if (score > bestPhaseScore) { bestPhaseScore = score; bestPhase = phase } + } + const r2 = (x: number) => Math.round(x * 100) / 100 + const beats: number[] = [] + for (let i = bestPhase; i < n; i += bestLag) beats.push(r2(i * hopSec)) + + // Strongest individual hits (drops, accents): local maxima well above the mean. + let mean = 0 + for (const v of onset) mean += v + mean /= n || 1 + let sd = 0 + for (const v of onset) sd += (v - mean) ** 2 + sd = Math.sqrt(sd / (n || 1)) + const peaks: Array<{ t: number; s: number }> = [] + for (let i = 2; i < n - 2; i++) { + const v = onset[i] + if (v > mean + 2 * sd && v >= onset[i - 1] && v >= onset[i + 1] && v >= onset[i - 2] && v >= onset[i + 2]) { + peaks.push({ t: i * hopSec, s: v }) + } + } + const strong = peaks.sort((a, b) => b.s - a.s).slice(0, 40).map(p => r2(p.t)).sort((a, b) => a - b) + + // Loudness per second (0–1) to spot intros, builds and drops. + const perSec = Math.round(1 / hopSec) + const loud: number[] = [] + for (let i = 0; i < n; i += perSec) { + let s = 0 + const end = Math.min(n, i + perSec) + for (let j = i; j < end; j++) s += energy[j] + loud.push(s / (end - i)) + } + const peak = Math.max(...loud, 1e-9) + + const totalDur = info.duration ?? samples.length / SR + return { + duration: r2(totalDur), + soundSegments: totalDur <= 60 ? soundSegments(filePath, totalDur) : undefined, // phrases, for dialogue clips + bpm: r2(60 / (bestLag * hopSec)), + beatInterval: r2(bestLag * hopSec), + beats, + downbeatsGuess: beats.filter((_, i) => i % 4 === 0), + strongOnsets: strong, + loudnessPerSecond: loud.map(v => Math.round((v / peak) * 100) / 100), + } +} diff --git a/electron/menu.ts b/electron/menu.ts new file mode 100644 index 000000000..6aba4e814 --- /dev/null +++ b/electron/menu.ts @@ -0,0 +1,99 @@ +import { app, Menu, shell, type BrowserWindow, type MenuItemConstructorOptions } from 'electron' + +export type MenuAction = 'new-project' | 'save-project' | 'export-project' | 'import-project' | 'show-keyboard-shortcuts' + +function sendMenuAction(window: BrowserWindow, action: MenuAction): void { + window.webContents.send('menu-action', action) +} + +/** Replaces Electron's auto-generated default menu so File has real Project + * Save/Export/Import actions (mirroring the equivalent in-app buttons) for + * users who reach for the menu bar out of habit. The actions themselves run + * in the renderer (ProjectContext/Project.tsx own the project state), so + * each item just forwards an IPC message for the renderer to act on. */ +export function createAppMenu(window: BrowserWindow): void { + const template: MenuItemConstructorOptions[] = [ + { + label: 'File', + submenu: [ + { label: 'New Project', click: () => sendMenuAction(window, 'new-project') }, + { type: 'separator' }, + { label: 'Save Project', accelerator: 'CmdOrCtrl+S', click: () => sendMenuAction(window, 'save-project') }, + { label: 'Export Project...', click: () => sendMenuAction(window, 'export-project') }, + { type: 'separator' }, + { label: 'Import Project...', click: () => sendMenuAction(window, 'import-project') }, + { type: 'separator' }, + { role: process.platform === 'darwin' ? 'close' : 'quit' }, + ], + }, + { role: 'editMenu' }, + { role: 'viewMenu' }, + { role: 'windowMenu' }, + { + role: 'help', + submenu: [ + { label: `${app.getName()} ${app.getVersion()}`, enabled: false }, + { type: 'separator' }, + { label: 'Documentation', click: () => shell.openExternal('https://github.com/Lightricks/LTX-Desktop') }, + { label: 'Keyboard Shortcuts', click: () => sendMenuAction(window, 'show-keyboard-shortcuts') }, + ], + }, + ] + + Menu.setApplicationMenu(Menu.buildFromTemplate(template)) +} + +/** Electron shows no native right-click menu by default, even though the + * Edit menu's accelerators (Ctrl+V etc.) already work - `context-menu` fires + * on every right-click with everything needed to build one (isEditable, + * per-action editFlags), so no renderer/preload changes are required. */ +export function registerEditContextMenu(window: BrowserWindow): void { + // Ensure the spellchecker has a language so `context-menu` params carry + // misspelledWord + dictionarySuggestions. No-op / harmless on macOS, which + // uses the OS spellchecker and ignores an explicit language list. + try { + window.webContents.session.setSpellCheckerLanguages(['en-US']) + } catch { + // Some platforms reject an explicit list; the default language still works. + } + + window.webContents.on('context-menu', (_event, params) => { + if (!params.isEditable) return + + const { editFlags } = params + const template: MenuItemConstructorOptions[] = [] + + // Spelling suggestions for a misspelled word under the cursor come first, + // then "Add to Dictionary", then the standard edit actions. + if (params.misspelledWord) { + if (params.dictionarySuggestions.length > 0) { + for (const suggestion of params.dictionarySuggestions) { + template.push({ label: suggestion, click: () => window.webContents.replaceMisspelling(suggestion) }) + } + } else { + template.push({ label: 'No spelling suggestions', enabled: false }) + } + template.push( + { type: 'separator' }, + { + label: 'Add to Dictionary', + click: () => window.webContents.session.addWordToSpellCheckerDictionary(params.misspelledWord), + }, + { type: 'separator' }, + ) + } + + template.push( + { role: 'undo', enabled: editFlags.canUndo }, + { role: 'redo', enabled: editFlags.canRedo }, + { type: 'separator' }, + { role: 'cut', enabled: editFlags.canCut }, + { role: 'copy', enabled: editFlags.canCopy }, + { role: 'paste', enabled: editFlags.canPaste }, + { type: 'separator' }, + { role: 'selectAll', enabled: editFlags.canSelectAll }, + ) + + Menu.buildFromTemplate(template).popup() + }) +} diff --git a/electron/preload.ts b/electron/preload.ts index 5e8f6682e..2a5f85909 100644 --- a/electron/preload.ts +++ b/electron/preload.ts @@ -1,4 +1,4 @@ -import { electronAPISchemas, type BackendHealthStatus, type UpdateStatePayload } from '../shared/electron-api-schema' +import { electronAPISchemas, type BackendHealthStatus, type UpdateStatePayload, type ExportProgress, type McpEditorRequest } from '../shared/electron-api-schema' const { contextBridge, ipcRenderer, webUtils } = require('electron') @@ -24,7 +24,37 @@ api.onBackendHealthStatus = (cb: (data: BackendHealthStatus) => void) => { } } -api.onUpdateEvent = (cb: (data: UpdateStatePayload) => void) => { +api.onExportProgress = (cb: (data: ExportProgress) => void) => { + const listener = (_: unknown, data: ExportProgress) => cb(data) + ipcRenderer.on('export-progress', listener) + return () => { + ipcRenderer.removeListener('export-progress', listener) + } +} + +api.onMenuAction = (cb: (action: string) => void) => { + const listener = (_: unknown, action: string) => cb(action) + ipcRenderer.on('menu-action', listener) + return () => { + ipcRenderer.removeListener('menu-action', listener) + } +} + +api.onGpmLibChanged = (cb: () => void) => { + const listener = () => cb() + ipcRenderer.on('gpm-lib-changed', listener) + return () => { + ipcRenderer.removeListener('gpm-lib-changed', listener) + } +} + +api.onMcpEditorRequest = (cb: (req: McpEditorRequest) => void) => { + const listener = (_: unknown, req: McpEditorRequest) => cb(req) + ipcRenderer.on('mcp-editor-request', listener) + return () => ipcRenderer.removeListener('mcp-editor-request', listener) +} + +api.onUpdateEvent =(cb: (data: UpdateStatePayload) => void) => { const listener = (_: unknown, data: UpdateStatePayload) => cb(data) ipcRenderer.on('update-event', listener) return () => ipcRenderer.removeListener('update-event', listener) diff --git a/electron/python-backend.ts b/electron/python-backend.ts index 9bd4db8fd..5b70a912c 100644 --- a/electron/python-backend.ts +++ b/electron/python-backend.ts @@ -344,6 +344,14 @@ export async function startPythonBackend(): Promise { } : {}), // Only pass LTX_PORT when the developer explicitly set it ...(process.env.LTX_PORT ? { LTX_PORT: process.env.LTX_PORT } : {}), + // Redirect the HuggingFace cache (Qwen Multi-Angle loads ~17GB of GGUF + // weights via hf_hub_download/from_pretrained, which obey HF_HOME) to a + // fast drive. LTX_HF_HOME overrides an inherited HF_HOME that may point + // at a slow disk — Michael's user HF_HOME was D:\hf-cache (SATA HDD), + // pegging that drive at 99% and stretching cold Qwen loads to minutes. + // App-scoped (never touches the global env / other HF tools) and + // portable: a no-op when LTX_HF_HOME is unset, mirroring LTX_MODELS_DIR. + ...(process.env.LTX_HF_HOME ? { HF_HOME: process.env.LTX_HF_HOME } : {}), LTX_AUTH_TOKEN: authToken, LTX_ADMIN_TOKEN: adminToken, LTX_LOG_FILE: getCurrentLogFilename(), diff --git a/electron/python-setup.ts b/electron/python-setup.ts index 5bfc65efd..0c6107e2e 100644 --- a/electron/python-setup.ts +++ b/electron/python-setup.ts @@ -133,7 +133,7 @@ export async function preDownloadPythonForUpdate( } const baseUrl = (isDev && process.env.LTX_PYTHON_URL?.replace(/^["']+|["']+$/g, '')) - || `https://github.com/Lightricks/ltx-desktop/releases/download/v${newVersion}` + || `https://github.com/MichaelRicks/LTX-Desktop-Mike/releases/download/v${newVersion}` // Fetch the new version's deps hash const newHash = await fetchRemoteDepsHash(baseUrl, 'python-next-hash-check.txt') @@ -242,7 +242,7 @@ function getArchiveBase(): string { return process.env.LTX_PYTHON_URL.replace(/^["']+|["']+$/g, '') } const version = app.getVersion() - return `https://github.com/Lightricks/ltx-desktop/releases/download/v${version}` + return `https://github.com/MichaelRicks/LTX-Desktop-Mike/releases/download/v${version}` } function getFallbackArchiveUrl(): string | null { diff --git a/electron/updater.ts b/electron/updater.ts index fceeae188..7a4438fec 100644 --- a/electron/updater.ts +++ b/electron/updater.ts @@ -49,8 +49,35 @@ function hasRestorableOffer(): boolean { return Boolean(state.version && getSkippedUpdateVersion() !== state.version) } +// A beta whose only releases are pre-releases — and which carries no +// electron-updater metadata (latest.yml) yet — has no update feed to read. A +// failed check then is expected, not a fault: GitHub answers the releases lookup +// with 404/406, or electron-updater reports it can't find a production release. +// Detect that class so it's logged quietly and shown as "up to date", instead of +// dumping a multi-KB HttpError (headers and all) to the log and UI on every check. +function isBenignNoReleaseError(message: string): boolean { + const m = message.toLowerCase() + return m.includes('unable to find latest version') + || m.includes('ensure a production release exists') + || m.includes('cannot parse releases feed') + || m.includes('latest.yml') + || m.includes('httperror: 404') + || m.includes('httperror: 406') +} + // Network/feed errors must not drop a known offer or a finished download. function failUpdate(message: string): void { + // No update feed yet (beta pre-release / no metadata): treat as "nothing to + // update to" — quiet log, show up-to-date, don't surface an error. Only when + // there's no in-flight download or restorable offer to protect. + if (isBenignNoReleaseError(message) + && state.status !== 'downloading' && state.status !== 'downloaded' + && !hasRestorableOffer()) { + logger.info('[updater] No update feed published yet — treating as up to date') + endMacFlight() + setState({ status: 'not-available', message: undefined }) + return + } logger.error(`[updater] ${message}`) if (process.platform === 'darwin') { // No modal / Try again. Keep a finished download; otherwise idle so Check retries. diff --git a/frontend/App.tsx b/frontend/App.tsx index 4736ded02..d8bc606ef 100644 --- a/frontend/App.tsx +++ b/frontend/App.tsx @@ -3,7 +3,7 @@ import { Loader2, AlertCircle, Settings, FileText } from 'lucide-react' import { ApiClient, type ApiSuccessOf } from './lib/api-client' import { ProjectProvider } from './contexts/ProjectContext' import { ViewProvider, useView } from './contexts/ViewContext' -import { KeyboardShortcutsProvider } from './contexts/KeyboardShortcutsContext' +import { KeyboardShortcutsProvider, useKeyboardShortcuts } from './contexts/KeyboardShortcutsContext' import { AppSettingsProvider, useAppSettings } from './contexts/AppSettingsContext' import { DevFlagsProvider } from './contexts/DevFlagsContext' import { KeyboardShortcutsModal } from './components/KeyboardShortcutsModal' @@ -21,8 +21,13 @@ import { SettingsModal, type SettingsInitialReason, type SettingsTabId } from '. import { LogViewer } from './components/LogViewer' import { ApiGatewayModal, type ApiGatewaySection } from './components/ApiGatewayModal' import { Button } from './components/ui/button' +import { PromptManagerPro } from './components/gpm/PromptManagerPro' +import { DownloadsBrowser } from './components/gpm/DownloadsBrowser' +import { useDownloadsBrowserOpen, getDownloadsBrowserOpen, setDownloadsBrowserOpen } from './components/gpm/downloads-browser-store' +import { usePromptManagerProOpen, getPromptManagerProOpen, setPromptManagerProOpen } from './components/gpm/prompt-manager-pro-store' import { useAppUpdateModal } from './hooks/use-app-update' import { UpdateAvailableModal } from './components/UpdateAvailableModal' +import { dispatchMcpEditorRequest } from './views/editor/mcp-editor-tools' type SetupState = 'loading' | { needsSetup: boolean; needsLicense: boolean } type RequiredModelsGateState = 'checking' | 'missing' | 'ready' @@ -31,6 +36,43 @@ type LtxUpgradeRecommendation = Extract { + const onKeyDown = (e: KeyboardEvent) => { + if (e.key !== 'Tab') return + const el = document.activeElement as HTMLElement | null + const isEditable = !!el && (el.tagName === 'INPUT' || el.tagName === 'TEXTAREA' || el.tagName === 'SELECT' || el.isContentEditable) + if (isEditable) return + e.preventDefault() + if (!panelsCollapsedRef.current) { + savedPanelStateRef.current = { left: getDownloadsBrowserOpen(), right: getPromptManagerProOpen() } + setDownloadsBrowserOpen(false) + setPromptManagerProOpen(false) + panelsCollapsedRef.current = true + } else { + setDownloadsBrowserOpen(savedPanelStateRef.current.left) + setPromptManagerProOpen(savedPanelStateRef.current.right) + panelsCollapsedRef.current = false + } + } + window.addEventListener('keydown', onKeyDown) + return () => window.removeEventListener('keydown', onKeyDown) + }, []) + // RiX MCP server: answer editor tool calls forwarded from the main process. Lives + // here (always mounted) so a call with no editor open fails fast instead of timing out. + useEffect(() => window.electronAPI?.onMcpEditorRequest?.(({ id, tool, args }) => { + void dispatchMcpEditorRequest(tool, args) + .then(result => window.electronAPI.mcpEditorResponse({ id, ok: true, result })) + .catch((e: unknown) => window.electronAPI.mcpEditorResponse({ id, ok: false, error: e instanceof Error ? e.message : String(e) })) + }), []) const { connected, processStatus, isLoading: backendLoading } = useBackend() const { settings, saveLtxApiKey, saveFalApiKey, forceApiGenerations, isLoaded, runtimePolicyLoaded, notifyModelsChanged } = useAppSettings() // Always mounted here (unlike GenSpace, which unmounts on every view/tab switch) so a @@ -79,6 +121,20 @@ function AppContent() { return () => window.removeEventListener('open-settings', handler) }, []) + // Forward the native File menu's Project actions (Save/Export/Import/New) + // into the renderer as a window event, the same way other cross-component + // signals here are handled — Home.tsx/Project.tsx listen for the actions + // relevant to whichever of them is currently mounted. + useEffect(() => { + return window.electronAPI?.onMenuAction((action) => { + if (action === 'show-keyboard-shortcuts') { + setKbEditorOpen(true) + return + } + window.dispatchEvent(new CustomEvent('ltx:menu-action', { detail: action })) + }) + }, [setKbEditorOpen]) + useEffect(() => { const handler = (e: Event) => { const detail = (e as CustomEvent).detail ?? {} @@ -232,7 +288,11 @@ function AppContent() { return } - if (forceApiGenerations || setupState.needsLicense || setupState.needsSetup) { + // With an LTX API key, generation can always run via the (free) LTX API, so + // never block startup on missing local models — local downloads stay + // available in Settings. (Also avoids a cold-boot model-scan race wrongly + // demanding the local text encoder.) + if (forceApiGenerations || settings.hasLtxApiKey || setupState.needsLicense || setupState.needsSetup) { setRequiredModelsGate('ready') return } @@ -262,6 +322,7 @@ function AppContent() { areRequiredModelsDownloaded, backendLoading, forceApiGenerations, + settings.hasLtxApiKey, setupState, connected, waitingForRuntimePolicy, @@ -481,7 +542,7 @@ function AppContent() {
-

Starting LTX Desktop...

+

Starting RiX Desktop Studio Pro...

Initializing the inference engine

@@ -532,7 +593,15 @@ function AppContent() { return (
- {renderView()} +
+ {renderView()} +
{showGlobalControls && (
@@ -630,6 +699,9 @@ function AppContent() { )} {restartingOverlay} + + +
) } diff --git a/frontend/assets/gpm-camera/birdseye.jpg b/frontend/assets/gpm-camera/birdseye.jpg new file mode 100644 index 000000000..eae899630 Binary files /dev/null and b/frontend/assets/gpm-camera/birdseye.jpg differ diff --git a/frontend/assets/gpm-camera/butterfly.jpg b/frontend/assets/gpm-camera/butterfly.jpg new file mode 100644 index 000000000..039fa9e7d Binary files /dev/null and b/frontend/assets/gpm-camera/butterfly.jpg differ diff --git a/frontend/assets/gpm-camera/cantedroll.jpg b/frontend/assets/gpm-camera/cantedroll.jpg new file mode 100644 index 000000000..e058bd392 Binary files /dev/null and b/frontend/assets/gpm-camera/cantedroll.jpg differ diff --git a/frontend/assets/gpm-camera/centerframe.jpg b/frontend/assets/gpm-camera/centerframe.jpg new file mode 100644 index 000000000..6def7f3c5 Binary files /dev/null and b/frontend/assets/gpm-camera/centerframe.jpg differ diff --git a/frontend/assets/gpm-camera/cinestill800t.jpg b/frontend/assets/gpm-camera/cinestill800t.jpg new file mode 100644 index 000000000..8ec9bcde6 Binary files /dev/null and b/frontend/assets/gpm-camera/cinestill800t.jpg differ diff --git a/frontend/assets/gpm-camera/closeup.jpg b/frontend/assets/gpm-camera/closeup.jpg new file mode 100644 index 000000000..692e5bca3 Binary files /dev/null and b/frontend/assets/gpm-camera/closeup.jpg differ diff --git a/frontend/assets/gpm-camera/dutchangle.jpg b/frontend/assets/gpm-camera/dutchangle.jpg new file mode 100644 index 000000000..a0eddb7a1 Binary files /dev/null and b/frontend/assets/gpm-camera/dutchangle.jpg differ diff --git a/frontend/assets/gpm-camera/extremeclose.jpg b/frontend/assets/gpm-camera/extremeclose.jpg new file mode 100644 index 000000000..01ae88628 Binary files /dev/null and b/frontend/assets/gpm-camera/extremeclose.jpg differ diff --git a/frontend/assets/gpm-camera/extremewide.jpg b/frontend/assets/gpm-camera/extremewide.jpg new file mode 100644 index 000000000..9da86ee03 Binary files /dev/null and b/frontend/assets/gpm-camera/extremewide.jpg differ diff --git a/frontend/assets/gpm-camera/eyelevel.jpg b/frontend/assets/gpm-camera/eyelevel.jpg new file mode 100644 index 000000000..b66a20418 Binary files /dev/null and b/frontend/assets/gpm-camera/eyelevel.jpg differ diff --git a/frontend/assets/gpm-camera/firecandlelight.jpg b/frontend/assets/gpm-camera/firecandlelight.jpg new file mode 100644 index 000000000..770138cfd Binary files /dev/null and b/frontend/assets/gpm-camera/firecandlelight.jpg differ diff --git a/frontend/assets/gpm-camera/foregroundlead.jpg b/frontend/assets/gpm-camera/foregroundlead.jpg new file mode 100644 index 000000000..816e4a378 Binary files /dev/null and b/frontend/assets/gpm-camera/foregroundlead.jpg differ diff --git a/frontend/assets/gpm-camera/frameinframe.jpg b/frontend/assets/gpm-camera/frameinframe.jpg new file mode 100644 index 000000000..2cacf3453 Binary files /dev/null and b/frontend/assets/gpm-camera/frameinframe.jpg differ diff --git a/frontend/assets/gpm-camera/fujisuperia400.jpg b/frontend/assets/gpm-camera/fujisuperia400.jpg new file mode 100644 index 000000000..8348cdaf1 Binary files /dev/null and b/frontend/assets/gpm-camera/fujisuperia400.jpg differ diff --git a/frontend/assets/gpm-camera/fujivelvia50.jpg b/frontend/assets/gpm-camera/fujivelvia50.jpg new file mode 100644 index 000000000..7ee84770a Binary files /dev/null and b/frontend/assets/gpm-camera/fujivelvia50.jpg differ diff --git a/frontend/assets/gpm-camera/goldenhour.jpg b/frontend/assets/gpm-camera/goldenhour.jpg new file mode 100644 index 000000000..789e15391 Binary files /dev/null and b/frontend/assets/gpm-camera/goldenhour.jpg differ diff --git a/frontend/assets/gpm-camera/hardlight.jpg b/frontend/assets/gpm-camera/hardlight.jpg new file mode 100644 index 000000000..fb27d4c04 Binary files /dev/null and b/frontend/assets/gpm-camera/hardlight.jpg differ diff --git a/frontend/assets/gpm-camera/highangle.jpg b/frontend/assets/gpm-camera/highangle.jpg new file mode 100644 index 000000000..322f48bf5 Binary files /dev/null and b/frontend/assets/gpm-camera/highangle.jpg differ diff --git a/frontend/assets/gpm-camera/highkey.jpg b/frontend/assets/gpm-camera/highkey.jpg new file mode 100644 index 000000000..f79ab16b3 Binary files /dev/null and b/frontend/assets/gpm-camera/highkey.jpg differ diff --git a/frontend/assets/gpm-camera/ilfordhp5.jpg b/frontend/assets/gpm-camera/ilfordhp5.jpg new file mode 100644 index 000000000..98310fbc3 Binary files /dev/null and b/frontend/assets/gpm-camera/ilfordhp5.jpg differ diff --git a/frontend/assets/gpm-camera/kodakektar100.jpg b/frontend/assets/gpm-camera/kodakektar100.jpg new file mode 100644 index 000000000..5a41535a0 Binary files /dev/null and b/frontend/assets/gpm-camera/kodakektar100.jpg differ diff --git a/frontend/assets/gpm-camera/kodakportra400.jpg b/frontend/assets/gpm-camera/kodakportra400.jpg new file mode 100644 index 000000000..198301591 Binary files /dev/null and b/frontend/assets/gpm-camera/kodakportra400.jpg differ diff --git a/frontend/assets/gpm-camera/kodaktrix.jpg b/frontend/assets/gpm-camera/kodaktrix.jpg new file mode 100644 index 000000000..eed541443 Binary files /dev/null and b/frontend/assets/gpm-camera/kodaktrix.jpg differ diff --git a/frontend/assets/gpm-camera/lomography800.jpg b/frontend/assets/gpm-camera/lomography800.jpg new file mode 100644 index 000000000..d45a97cff Binary files /dev/null and b/frontend/assets/gpm-camera/lomography800.jpg differ diff --git a/frontend/assets/gpm-camera/lowangle.jpg b/frontend/assets/gpm-camera/lowangle.jpg new file mode 100644 index 000000000..fe6985dec Binary files /dev/null and b/frontend/assets/gpm-camera/lowangle.jpg differ diff --git a/frontend/assets/gpm-camera/lowkey.jpg b/frontend/assets/gpm-camera/lowkey.jpg new file mode 100644 index 000000000..a721c2550 Binary files /dev/null and b/frontend/assets/gpm-camera/lowkey.jpg differ diff --git a/frontend/assets/gpm-camera/medium.jpg b/frontend/assets/gpm-camera/medium.jpg new file mode 100644 index 000000000..68baaaa79 Binary files /dev/null and b/frontend/assets/gpm-camera/medium.jpg differ diff --git a/frontend/assets/gpm-camera/mediumclose.jpg b/frontend/assets/gpm-camera/mediumclose.jpg new file mode 100644 index 000000000..8b82c8328 Binary files /dev/null and b/frontend/assets/gpm-camera/mediumclose.jpg differ diff --git a/frontend/assets/gpm-camera/mediumwide.jpg b/frontend/assets/gpm-camera/mediumwide.jpg new file mode 100644 index 000000000..2b86930ad Binary files /dev/null and b/frontend/assets/gpm-camera/mediumwide.jpg differ diff --git a/frontend/assets/gpm-camera/mlapocalyptic.jpg b/frontend/assets/gpm-camera/mlapocalyptic.jpg new file mode 100644 index 000000000..21a463771 Binary files /dev/null and b/frontend/assets/gpm-camera/mlapocalyptic.jpg differ diff --git a/frontend/assets/gpm-camera/mlcandlelitperiod.jpg b/frontend/assets/gpm-camera/mlcandlelitperiod.jpg new file mode 100644 index 000000000..fe638be7f Binary files /dev/null and b/frontend/assets/gpm-camera/mlcandlelitperiod.jpg differ diff --git a/frontend/assets/gpm-camera/mlcoldminimalism.jpg b/frontend/assets/gpm-camera/mlcoldminimalism.jpg new file mode 100644 index 000000000..11d32b66b Binary files /dev/null and b/frontend/assets/gpm-camera/mlcoldminimalism.jpg differ diff --git a/frontend/assets/gpm-camera/mlcoldspace.jpg b/frontend/assets/gpm-camera/mlcoldspace.jpg new file mode 100644 index 000000000..bd86461e4 Binary files /dev/null and b/frontend/assets/gpm-camera/mlcoldspace.jpg differ diff --git a/frontend/assets/gpm-camera/mlcoldwilderness.jpg b/frontend/assets/gpm-camera/mlcoldwilderness.jpg new file mode 100644 index 000000000..1a6dbcc56 Binary files /dev/null and b/frontend/assets/gpm-camera/mlcoldwilderness.jpg differ diff --git a/frontend/assets/gpm-camera/mlcontemplativescifi.jpg b/frontend/assets/gpm-camera/mlcontemplativescifi.jpg new file mode 100644 index 000000000..45c8b025a Binary files /dev/null and b/frontend/assets/gpm-camera/mlcontemplativescifi.jpg differ diff --git a/frontend/assets/gpm-camera/mlcontrolledtension.jpg b/frontend/assets/gpm-camera/mlcontrolledtension.jpg new file mode 100644 index 000000000..f2240e27c Binary files /dev/null and b/frontend/assets/gpm-camera/mlcontrolledtension.jpg differ diff --git a/frontend/assets/gpm-camera/mldesaturateddread.jpg b/frontend/assets/gpm-camera/mldesaturateddread.jpg new file mode 100644 index 000000000..a2bb9cbc2 Binary files /dev/null and b/frontend/assets/gpm-camera/mldesaturateddread.jpg differ diff --git a/frontend/assets/gpm-camera/mldesaturatedtrenches.jpg b/frontend/assets/gpm-camera/mldesaturatedtrenches.jpg new file mode 100644 index 000000000..e791763b8 Binary files /dev/null and b/frontend/assets/gpm-camera/mldesaturatedtrenches.jpg differ diff --git a/frontend/assets/gpm-camera/mldesertgold.jpg b/frontend/assets/gpm-camera/mldesertgold.jpg new file mode 100644 index 000000000..ce2fd85dd Binary files /dev/null and b/frontend/assets/gpm-camera/mldesertgold.jpg differ diff --git a/frontend/assets/gpm-camera/mldigitalnightscape.jpg b/frontend/assets/gpm-camera/mldigitalnightscape.jpg new file mode 100644 index 000000000..cecfb7a92 Binary files /dev/null and b/frontend/assets/gpm-camera/mldigitalnightscape.jpg differ diff --git a/frontend/assets/gpm-camera/mldreamlike.jpg b/frontend/assets/gpm-camera/mldreamlike.jpg new file mode 100644 index 000000000..1297aaa9d Binary files /dev/null and b/frontend/assets/gpm-camera/mldreamlike.jpg differ diff --git a/frontend/assets/gpm-camera/mlfoggymelancholy.jpg b/frontend/assets/gpm-camera/mlfoggymelancholy.jpg new file mode 100644 index 000000000..9d7ba6820 Binary files /dev/null and b/frontend/assets/gpm-camera/mlfoggymelancholy.jpg differ diff --git a/frontend/assets/gpm-camera/mlfuturisticneon.jpg b/frontend/assets/gpm-camera/mlfuturisticneon.jpg new file mode 100644 index 000000000..6994f7880 Binary files /dev/null and b/frontend/assets/gpm-camera/mlfuturisticneon.jpg differ diff --git a/frontend/assets/gpm-camera/mlgoldenrome.jpg b/frontend/assets/gpm-camera/mlgoldenrome.jpg new file mode 100644 index 000000000..f7e9f5dad Binary files /dev/null and b/frontend/assets/gpm-camera/mlgoldenrome.jpg differ diff --git a/frontend/assets/gpm-camera/mlgreendigital.jpg b/frontend/assets/gpm-camera/mlgreendigital.jpg new file mode 100644 index 000000000..d1e4a8b88 Binary files /dev/null and b/frontend/assets/gpm-camera/mlgreendigital.jpg differ diff --git a/frontend/assets/gpm-camera/mlhighcontrastbw.jpg b/frontend/assets/gpm-camera/mlhighcontrastbw.jpg new file mode 100644 index 000000000..ab12675ec Binary files /dev/null and b/frontend/assets/gpm-camera/mlhighcontrastbw.jpg differ diff --git a/frontend/assets/gpm-camera/mlhorrorteal.jpg b/frontend/assets/gpm-camera/mlhorrorteal.jpg new file mode 100644 index 000000000..a5b934921 Binary files /dev/null and b/frontend/assets/gpm-camera/mlhorrorteal.jpg differ diff --git a/frontend/assets/gpm-camera/mlnearfuture.jpg b/frontend/assets/gpm-camera/mlnearfuture.jpg new file mode 100644 index 000000000..672e24ca9 Binary files /dev/null and b/frontend/assets/gpm-camera/mlnearfuture.jpg differ diff --git a/frontend/assets/gpm-camera/mlneoncyberpunk.jpg b/frontend/assets/gpm-camera/mlneoncyberpunk.jpg new file mode 100644 index 000000000..0669724a2 Binary files /dev/null and b/frontend/assets/gpm-camera/mlneoncyberpunk.jpg differ diff --git a/frontend/assets/gpm-camera/mlpastelsymmetrical.jpg b/frontend/assets/gpm-camera/mlpastelsymmetrical.jpg new file mode 100644 index 000000000..ae0dea1fa Binary files /dev/null and b/frontend/assets/gpm-camera/mlpastelsymmetrical.jpg differ diff --git a/frontend/assets/gpm-camera/mlromantichaze.jpg b/frontend/assets/gpm-camera/mlromantichaze.jpg new file mode 100644 index 000000000..3950b8f0e Binary files /dev/null and b/frontend/assets/gpm-camera/mlromantichaze.jpg differ diff --git a/frontend/assets/gpm-camera/mlsaturatedpop.jpg b/frontend/assets/gpm-camera/mlsaturatedpop.jpg new file mode 100644 index 000000000..8d9118d93 Binary files /dev/null and b/frontend/assets/gpm-camera/mlsaturatedpop.jpg differ diff --git a/frontend/assets/gpm-camera/mlsundrenched.jpg b/frontend/assets/gpm-camera/mlsundrenched.jpg new file mode 100644 index 000000000..0a7f55abd Binary files /dev/null and b/frontend/assets/gpm-camera/mlsundrenched.jpg differ diff --git a/frontend/assets/gpm-camera/mlwarmwhimsy.jpg b/frontend/assets/gpm-camera/mlwarmwhimsy.jpg new file mode 100644 index 000000000..f97b63402 Binary files /dev/null and b/frontend/assets/gpm-camera/mlwarmwhimsy.jpg differ diff --git a/frontend/assets/gpm-camera/negativespace.jpg b/frontend/assets/gpm-camera/negativespace.jpg new file mode 100644 index 000000000..368a05b8e Binary files /dev/null and b/frontend/assets/gpm-camera/negativespace.jpg differ diff --git a/frontend/assets/gpm-camera/neonnoir.jpg b/frontend/assets/gpm-camera/neonnoir.jpg new file mode 100644 index 000000000..42619a3dc Binary files /dev/null and b/frontend/assets/gpm-camera/neonnoir.jpg differ diff --git a/frontend/assets/gpm-camera/ots.jpg b/frontend/assets/gpm-camera/ots.jpg new file mode 100644 index 000000000..3118fd091 Binary files /dev/null and b/frontend/assets/gpm-camera/ots.jpg differ diff --git a/frontend/assets/gpm-camera/overheadflatlay.jpg b/frontend/assets/gpm-camera/overheadflatlay.jpg new file mode 100644 index 000000000..66c83b698 Binary files /dev/null and b/frontend/assets/gpm-camera/overheadflatlay.jpg differ diff --git a/frontend/assets/gpm-camera/pov.jpg b/frontend/assets/gpm-camera/pov.jpg new file mode 100644 index 000000000..0419f6161 Binary files /dev/null and b/frontend/assets/gpm-camera/pov.jpg differ diff --git a/frontend/assets/gpm-camera/practical.jpg b/frontend/assets/gpm-camera/practical.jpg new file mode 100644 index 000000000..893164aab Binary files /dev/null and b/frontend/assets/gpm-camera/practical.jpg differ diff --git a/frontend/assets/gpm-camera/rembrandt.jpg b/frontend/assets/gpm-camera/rembrandt.jpg new file mode 100644 index 000000000..d2d4b8122 Binary files /dev/null and b/frontend/assets/gpm-camera/rembrandt.jpg differ diff --git a/frontend/assets/gpm-camera/rim.jpg b/frontend/assets/gpm-camera/rim.jpg new file mode 100644 index 000000000..0e0e83df0 Binary files /dev/null and b/frontend/assets/gpm-camera/rim.jpg differ diff --git a/frontend/assets/gpm-camera/rulethirds.jpg b/frontend/assets/gpm-camera/rulethirds.jpg new file mode 100644 index 000000000..4fce881dd Binary files /dev/null and b/frontend/assets/gpm-camera/rulethirds.jpg differ diff --git a/frontend/assets/gpm-camera/silhouette.jpg b/frontend/assets/gpm-camera/silhouette.jpg new file mode 100644 index 000000000..d57e53b63 Binary files /dev/null and b/frontend/assets/gpm-camera/silhouette.jpg differ diff --git a/frontend/assets/gpm-camera/softdiffused.jpg b/frontend/assets/gpm-camera/softdiffused.jpg new file mode 100644 index 000000000..3ad3fcdce Binary files /dev/null and b/frontend/assets/gpm-camera/softdiffused.jpg differ diff --git a/frontend/assets/gpm-camera/split.jpg b/frontend/assets/gpm-camera/split.jpg new file mode 100644 index 000000000..a2f86c4f4 Binary files /dev/null and b/frontend/assets/gpm-camera/split.jpg differ diff --git a/frontend/assets/gpm-camera/threepoint.jpg b/frontend/assets/gpm-camera/threepoint.jpg new file mode 100644 index 000000000..83b3b7ec9 Binary files /dev/null and b/frontend/assets/gpm-camera/threepoint.jpg differ diff --git a/frontend/assets/gpm-camera/twoshot.jpg b/frontend/assets/gpm-camera/twoshot.jpg new file mode 100644 index 000000000..f69615c17 Binary files /dev/null and b/frontend/assets/gpm-camera/twoshot.jpg differ diff --git a/frontend/assets/gpm-camera/wideshot.jpg b/frontend/assets/gpm-camera/wideshot.jpg new file mode 100644 index 000000000..6767963bd Binary files /dev/null and b/frontend/assets/gpm-camera/wideshot.jpg differ diff --git a/frontend/assets/gpm-camera/wormseye.jpg b/frontend/assets/gpm-camera/wormseye.jpg new file mode 100644 index 000000000..138120cf8 Binary files /dev/null and b/frontend/assets/gpm-camera/wormseye.jpg differ diff --git a/frontend/components/AudioWaveform.tsx b/frontend/components/AudioWaveform.tsx index a6d432020..184d4dddb 100644 --- a/frontend/components/AudioWaveform.tsx +++ b/frontend/components/AudioWaveform.tsx @@ -15,9 +15,22 @@ interface AudioWaveformProps { isPlaying: boolean } -// Global waveform cache: URL → Float32Array of peak amplitudes (one per pixel-bucket) +// High-resolution amplitude envelope of a whole audio file, decoded once per URL. +// Every view (monitor, timeline clips at any zoom) derives what it draws from this. +export interface WaveformEnvelope { + peak: Float32Array // max |sample| per window + rms: Float32Array // root-mean-square per window (the "body" of the sound) + rate: number // windows per second of audio + duration: number // seconds +} + +const ENVELOPE_RATE = 200 // 5ms windows: smooth at any practical timeline zoom + +const envelopeCache = new Map() +const pendingEnvelopes = new Map>() + +// Global waveform cache: `${url}@${buckets}` → peak amplitudes resampled to `buckets` export const waveformCache = new Map() -const pendingDecodes = new Set() // Convert a base64 string to an ArrayBuffer function base64ToArrayBuffer(base64: string): ArrayBuffer { @@ -29,53 +42,81 @@ function base64ToArrayBuffer(base64: string): ArrayBuffer { return bytes.buffer } -// Decode audio file and extract amplitude envelope -export async function computeWaveform(url: string, buckets: number = 800): Promise { - if (waveformCache.has(url)) return waveformCache.get(url)! +async function decodeEnvelope(url: string): Promise { + let arrayBuffer: ArrayBuffer - if (pendingDecodes.has(url)) { - while (pendingDecodes.has(url)) { - await new Promise(r => setTimeout(r, 50)) - } - if (waveformCache.has(url)) return waveformCache.get(url)! + if (url.startsWith('file://') && (window as any).electronAPI?.readLocalFile) { + const { data } = await (window as any).electronAPI.readLocalFile({ filePath: url }) + arrayBuffer = base64ToArrayBuffer(data) + } else { + const response = await fetch(url) + arrayBuffer = await response.arrayBuffer() } - pendingDecodes.add(url) - try { - let arrayBuffer: ArrayBuffer - - if (url.startsWith('file://') && (window as any).electronAPI?.readLocalFile) { - const { data } = await (window as any).electronAPI.readLocalFile({ filePath: url }) - arrayBuffer = base64ToArrayBuffer(data) - } else { - const response = await fetch(url) - arrayBuffer = await response.arrayBuffer() + const audioCtx = new (window.AudioContext || (window as any).webkitAudioContext)() + const audioBuffer = await audioCtx.decodeAudioData(arrayBuffer) + audioCtx.close() + + // Mix all channels so a hard-panned stereo track still shows its full shape + const channels = Array.from({ length: audioBuffer.numberOfChannels }, (_, c) => audioBuffer.getChannelData(c)) + const length = audioBuffer.length + const samplesPerWindow = Math.max(1, Math.round(audioBuffer.sampleRate / ENVELOPE_RATE)) + const windows = Math.ceil(length / samplesPerWindow) + const peak = new Float32Array(windows) + const rms = new Float32Array(windows) + + for (let i = 0; i < windows; i++) { + const start = i * samplesPerWindow + const end = Math.min(start + samplesPerWindow, length) + let max = 0 + let sumSq = 0 + for (let j = start; j < end; j++) { + let s = 0 + for (const ch of channels) s += ch[j] + s /= channels.length + const abs = Math.abs(s) + if (abs > max) max = abs + sumSq += s * s } + peak[i] = max + rms[i] = Math.sqrt(sumSq / Math.max(1, end - start)) + } - const audioCtx = new (window.AudioContext || (window as any).webkitAudioContext)() - const audioBuffer = await audioCtx.decodeAudioData(arrayBuffer) - audioCtx.close() - - const channelData = audioBuffer.getChannelData(0) - const samplesPerBucket = Math.floor(channelData.length / buckets) - const peaks = new Float32Array(buckets) - - for (let i = 0; i < buckets; i++) { - let max = 0 - const start = i * samplesPerBucket - const end = Math.min(start + samplesPerBucket, channelData.length) - for (let j = start; j < end; j++) { - const abs = Math.abs(channelData[j]) - if (abs > max) max = abs - } - peaks[i] = max - } + return { peak, rms, rate: audioBuffer.sampleRate / samplesPerWindow, duration: audioBuffer.duration } +} + +// Decode (once) and return the high-resolution envelope for a file +export async function getWaveformEnvelope(url: string): Promise { + const cached = envelopeCache.get(url) + if (cached) return cached + let pending = pendingEnvelopes.get(url) + if (!pending) { + pending = decodeEnvelope(url) + .then(env => { envelopeCache.set(url, env); return env }) + .finally(() => pendingEnvelopes.delete(url)) + pendingEnvelopes.set(url, pending) + } + return pending +} - waveformCache.set(url, peaks) - return peaks - } finally { - pendingDecodes.delete(url) +// Peak amplitudes of the whole file resampled to `buckets` (max-pooled, so short hits survive) +export async function computeWaveform(url: string, buckets: number = 800): Promise { + const key = `${url}@${buckets}` + const cached = waveformCache.get(key) + if (cached) return cached + + const { peak } = await getWaveformEnvelope(url) + const peaks = new Float32Array(buckets) + for (let i = 0; i < buckets; i++) { + const start = Math.floor((i / buckets) * peak.length) + const end = Math.max(start + 1, Math.floor(((i + 1) / buckets) * peak.length)) + let max = 0 + for (let j = start; j < end && j < peak.length; j++) if (peak[j] > max) max = peak[j] + peaks[i] = max } + + waveformCache.set(key, peaks) + return peaks } export function AudioWaveform({ audioClips, currentTime, isPlaying }: AudioWaveformProps) { @@ -308,19 +349,57 @@ export function AudioWaveform({ audioClips, currentTime, isPlaying }: AudioWavef interface ClipWaveformProps { url: string className?: string + // Outer (peak) layer color?: string + // Inner (RMS) layer: the loudness body of the sound + bodyColor?: string + // Source window the clip plays. Omit to show the whole file. + trimStart?: number + duration?: number + speed?: number + reversed?: boolean +} + +// Envelope value over [a, b) window indices: max for peaks, power-mean for RMS. +// Sub-window spans are linearly interpolated so zoomed-in waveforms stay smooth. +function sampleEnvelope(data: Float32Array, a: number, b: number, mode: 'max' | 'rms'): number { + const n = data.length + if (n === 0) return 0 + if (b - a <= 1) { + const t = Math.min(Math.max((a + b) / 2 - 0.5, 0), n - 1) + const i = Math.floor(t) + const f = t - i + return data[i] * (1 - f) + data[Math.min(i + 1, n - 1)] * f + } + const start = Math.max(0, Math.floor(a)) + const end = Math.min(n, Math.ceil(b)) + let acc = 0 + for (let i = start; i < end; i++) { + const v = data[i] + if (mode === 'max') { if (v > acc) acc = v } else acc += v * v + } + return mode === 'max' ? acc : Math.sqrt(acc / Math.max(1, end - start)) } -export function ClipWaveform({ url, className = '', color = 'rgba(52, 211, 153, 0.7)' }: ClipWaveformProps) { +export function ClipWaveform({ + url, + className = '', + color = 'rgba(52, 211, 153, 0.45)', + bodyColor = 'rgba(110, 231, 183, 0.95)', + trimStart = 0, + duration, + speed = 1, + reversed = false, +}: ClipWaveformProps) { const canvasRef = useRef(null) const containerRef = useRef(null) - const [peaks, setPeaks] = useState(null) + const [envelope, setEnvelope] = useState(null) useEffect(() => { if (!url) return let cancelled = false - computeWaveform(url, 200).then(p => { - if (!cancelled) setPeaks(p) + getWaveformEnvelope(url).then(env => { + if (!cancelled) setEnvelope(env) }).catch(() => {}) return () => { cancelled = true } }, [url]) @@ -328,7 +407,7 @@ export function ClipWaveform({ url, className = '', color = 'rgba(52, 211, 153, const draw = useCallback(() => { const canvas = canvasRef.current const container = containerRef.current - if (!canvas || !container || !peaks) return + if (!canvas || !container || !envelope) return const dpr = window.devicePixelRatio || 1 const rect = container.getBoundingClientRect() @@ -349,24 +428,42 @@ export function ClipWaveform({ url, className = '', color = 'rgba(52, 211, 153, const centerY = h / 2 const maxAmp = h * 0.45 - ctx.fillStyle = color - ctx.beginPath() - for (let i = 0; i < w; i++) { - const peakIdx = Math.floor((i / w) * peaks.length) - const amp = peaks[Math.min(peakIdx, peaks.length - 1)] - const y = centerY - amp * maxAmp - if (i === 0) ctx.moveTo(i, y) - else ctx.lineTo(i, y) + // Source seconds this clip plays (in envelope-window units) + const srcLen = (duration ?? envelope.duration - trimStart) * speed + const srcStart = trimStart * envelope.rate + const srcSpan = Math.max(0, srcLen * envelope.rate) + const cols = Math.max(1, Math.ceil(w)) + const peakCol = new Float32Array(cols) + const rmsCol = new Float32Array(cols) + for (let x = 0; x < cols; x++) { + const fa = x / cols + const fb = (x + 1) / cols + const a = srcStart + (reversed ? 1 - fb : fa) * srcSpan + const b = srcStart + (reversed ? 1 - fa : fb) * srcSpan + peakCol[x] = sampleEnvelope(envelope.peak, a, b, 'max') + rmsCol[x] = Math.min(peakCol[x], sampleEnvelope(envelope.rms, a, b, 'rms')) } - for (let i = w - 1; i >= 0; i--) { - const peakIdx = Math.floor((i / w) * peaks.length) - const amp = peaks[Math.min(peakIdx, peaks.length - 1)] - const y = centerY + amp * maxAmp - ctx.lineTo(i, y) + + const fillMirrored = (amps: Float32Array, style: string) => { + ctx.fillStyle = style + ctx.beginPath() + ctx.moveTo(0, centerY) + for (let x = 0; x < cols; x++) ctx.lineTo(x + 0.5, centerY - amps[x] * maxAmp) + ctx.lineTo(cols, centerY) + for (let x = cols - 1; x >= 0; x--) ctx.lineTo(x + 0.5, centerY + amps[x] * maxAmp) + ctx.closePath() + ctx.fill() } - ctx.closePath() - ctx.fill() - }, [peaks, color]) + + fillMirrored(peakCol, color) + fillMirrored(rmsCol, bodyColor) + + // Hairline centre so silence still reads as "audio here" + ctx.fillStyle = bodyColor + ctx.globalAlpha = 0.35 + ctx.fillRect(0, Math.round(centerY) - 0.5, w, 1) + ctx.globalAlpha = 1 + }, [envelope, color, bodyColor, trimStart, duration, speed, reversed]) useEffect(() => { draw() diff --git a/frontend/components/ExportModal.tsx b/frontend/components/ExportModal.tsx index 28befde2b..567cdbbdb 100644 --- a/frontend/components/ExportModal.tsx +++ b/frontend/components/ExportModal.tsx @@ -1,12 +1,11 @@ import { useState, useRef, useCallback, useEffect, useMemo } from 'react' import { X, Download, FolderOpen, Film, Package, Loader2, Check, AlertCircle, ChevronDown } from 'lucide-react' import { Button } from './ui/button' -import { DEFAULT_SUBTITLE_STYLE } from '../types/project-model' import type { Track, TimelineClip } from '../types/project-model' +import { buildExportPayload } from '../views/editor/export-payload' import { selectActiveTimeline, selectAssets, - selectClipPathFromAssets, selectClips, selectShowExportModal, selectSubtitles, @@ -50,14 +49,6 @@ const PRORES_PROFILES = [ { value: 3, label: 'HQ' }, ] -const LETTERBOX_RATIO_MAP: Record = { - '2.35:1': 2.35, - '2.39:1': 2.39, - '2.76:1': 2.76, - '1.85:1': 1.85, - '4:3': 4 / 3, -} - // Generate FCPXML for Premiere / DaVinci function generateFCPXML( clips: TimelineClip[], @@ -161,63 +152,10 @@ export function ExportModal({ projectName }: ExportModalProps) { const tracks = useEditorStore(selectTracks) const subtitles = useEditorStore(selectSubtitles) - const exportClips = useMemo(() => ( - clips - .filter(clip => clip.type === 'video' || clip.type === 'image' || clip.type === 'audio') - .filter(clip => tracks[clip.trackIndex]?.enabled !== false) - .map(clip => ({ - path: selectClipPathFromAssets(assets, clip), - type: clip.type, - startTime: clip.startTime, - duration: clip.duration, - trimStart: clip.trimStart, - speed: clip.speed || 1, - reversed: clip.reversed || false, - flipH: clip.flipH || false, - flipV: clip.flipV || false, - opacity: clip.opacity ?? 100, - trackIndex: clip.trackIndex, - muted: clip.muted || false, - volume: clip.volume ?? 1, - })) - ), [assets, clips, tracks]) - - const subtitleData = useMemo(() => ( - subtitles.map(subtitle => { - const track = tracks[subtitle.trackIndex] - return { - text: subtitle.text, - startTime: subtitle.startTime, - endTime: subtitle.endTime, - style: { - ...DEFAULT_SUBTITLE_STYLE, - ...(track?.subtitleStyle || {}), - ...(subtitle.style || {}), - }, - } - }) - ), [subtitles, tracks]) - - const letterbox = useMemo(() => { - const adjustmentClips = clips.filter( - clip => - clip.type === 'adjustment' - && clip.letterbox?.enabled - && tracks[clip.trackIndex]?.enabled !== false, - ) - if (adjustmentClips.length === 0) return null - const best = adjustmentClips.reduce((currentBest, candidate) => ( - candidate.duration > currentBest.duration ? candidate : currentBest - )) - const config = best.letterbox! - return { - ratio: config.aspectRatio === 'custom' - ? (config.customRatio || 2.35) - : (LETTERBOX_RATIO_MAP[config.aspectRatio] || 2.35), - color: config.color || '#000000', - opacity: (config.opacity ?? 100) / 100, - } - }, [clips, tracks]) + const payload = useMemo( + () => buildExportPayload(assets, clips, tracks, subtitles), + [assets, clips, tracks, subtitles], + ) const [exportStatus, setExportStatus] = useState('idle') const [exportType, setExportType] = useState<'package' | 'video' | null>(null) @@ -241,7 +179,7 @@ export function ExportModal({ projectName }: ExportModalProps) { closeExportModal() }, [closeExportModal]) - const hasSubtitles = subtitleData.length > 0 + const hasSubtitles = subtitles.length > 0 useEffect(() => { if (!isOpen) return @@ -309,6 +247,7 @@ export function ExportModal({ projectName }: ExportModalProps) { setExportFrameInfo('Preparing...') abortRef.current = false + let unsubscribe: (() => void) | undefined try { const codecInfo = CODEC_INFO[settings.codec] @@ -329,16 +268,22 @@ export function ExportModal({ projectName }: ExportModalProps) { // Build clip data for ffmpeg native export (video/image + audio clips) setExportFrameInfo('Starting ffmpeg...') + // Live progress pushed from the main process as ffmpeg encodes. + unsubscribe = window.electronAPI?.onExportProgress?.(({ percent, stage }) => { + if (abortRef.current) return + setExportProgress(percent) + if (stage) setExportFrameInfo(stage) + }) + const result = await window.electronAPI?.exportNative({ - clips: exportClips, + ...payload, + subtitles: burnSubtitles ? payload.subtitles : undefined, outputPath: filePath, codec: settings.codec, width: settings.width, height: settings.height, fps: settings.fps, quality: settings.quality, - letterbox: letterbox || undefined, - subtitles: burnSubtitles && subtitleData.length > 0 ? subtitleData : undefined, }) if (result && !result.success) { @@ -352,8 +297,10 @@ export function ExportModal({ projectName }: ExportModalProps) { } catch (err) { setExportError(String(err)) setExportStatus('error') + } finally { + unsubscribe?.() } - }, [burnSubtitles, exportClips, letterbox, projectName, settings, subtitleData, timeline]) + }, [burnSubtitles, payload, projectName, settings, timeline]) const handleCancel = useCallback(async () => { abortRef.current = true diff --git a/frontend/components/FirstRunSetup.tsx b/frontend/components/FirstRunSetup.tsx index a0c345dac..8254f4d89 100644 --- a/frontend/components/FirstRunSetup.tsx +++ b/frontend/components/FirstRunSetup.tsx @@ -1171,7 +1171,7 @@ export function LaunchGate({ justifyContent: 'space-between', alignItems: 'center' }}> -
© 2026 Lightricks
+
© 2026 Michael Ricks
{/* Next/Install/Finish Button */} diff --git a/frontend/components/ICLoraPanel.tsx b/frontend/components/ICLoraPanel.tsx index 69cabc333..d5bceba10 100644 --- a/frontend/components/ICLoraPanel.tsx +++ b/frontend/components/ICLoraPanel.tsx @@ -1,12 +1,14 @@ import { useState, useRef, useEffect, useCallback } from 'react' import { Upload, Loader2, Film, Sparkles, Image as ImageIcon, - RefreshCw, Download, AlertCircle, Trash2, + RefreshCw, Download, AlertCircle, Trash2, Square, } from 'lucide-react' import { ApiClient, type ApiRequestBodyOf, type ApiSuccessOf } from '../lib/api-client' import { logger } from '../lib/logger' import { pathToFileUrl } from '../lib/file-url' import { OutpaintCanvasEditor, type OutpaintPads } from './OutpaintCanvasEditor' +import { saveDataUrlToTempFile, GPM_IMAGE_DND_TYPE, type GpmDndImage } from './gpm/gpm-image-file' +import { FILE_DND, type LibFile } from './gpm/DownloadsBrowser' export type ICLoraConditioningType = 'canny' | 'depth' | 'custom' @@ -16,6 +18,9 @@ interface ICLoraPanelProps { fillHeight?: boolean isProcessing?: boolean processingStatus?: string + // Abort the in-flight IC-LoRA generation (process-wide cancel). When provided, a Stop + // button shows in the Output panel while processing. Omitted = no Stop button. + onCancel?: () => void // IC-LoRA mode: 'image' accepts a still image and skips canny/depth preprocessing // (the IC-LoRA builds the control video server-side). Default 'video' = today's behavior. inputKind?: 'image' | 'video' @@ -67,6 +72,7 @@ export function ICLoraPanel({ fillHeight = false, isProcessing = false, processingStatus = '', + onCancel, inputKind = 'video', selectedIcLoraId = null, allowsReferenceImage = false, @@ -89,6 +95,14 @@ export function ICLoraPanel({ // Source pixel dimensions, read off the loaded media — drives the outpaint canvas editor. const [sourceDims, setSourceDims] = useState<{ w: number; h: number } | null>(null) + // Local "Stopping…" latch for the Output Stop button. Cancel is polled between denoise + // steps (which can be long for the 22B IC-LoRA path), so echo intent immediately, then + // clear it once the run actually ends (isProcessing → false). + const [stopRequested, setStopRequested] = useState(false) + useEffect(() => { + if (!isProcessing) setStopRequested(false) + }, [isProcessing]) + // Optional reference image for catalog IC-LoRAs that allow it (frame 0, strength 1.0). const [referenceImagePath, setReferenceImagePath] = useState(null) const referenceImageUrl = referenceImagePath ? pathToFileUrl(referenceImagePath) : null @@ -141,8 +155,15 @@ export function ICLoraPanel({ setSourceDims(null) setInternalCondType('canny') setInternalCondStrength(1.0) - onConditioningTypeChange?.('canny') - onConditioningStrengthChange?.(1.0) + // Only reset the PARENT's conditioning type/strength when no catalog IC-LoRA is selected. + // The panel remounts on every entry into IC-LoRA mode; when a catalog LoRA (e.g. Ingredients) + // is active it owns the selection + its default_settings, and notifying 'canny' here made the + // parent deselect it and wipe those settings — so it silently reverted to Canny Edges on + // re-entry. Canny/depth is meaningless for a catalog LoRA anyway. + if (!isCatalogIcLora) { + onConditioningTypeChange?.('canny') + onConditioningStrengthChange?.(1.0) + } setConditioningPreview(null) setExtractError(null) setReferenceImagePath(null) @@ -322,6 +343,28 @@ export function ICLoraPanel({ e.preventDefault() setIsDragOver(false) + const acceptInput = (path: string) => { + setInputVideoPath(path) + setConditioningPreview(null) + setExtractError(null) + } + + // Prompt Manager Pro image (data URL) — only meaningful for an image input. + const gpmData = e.dataTransfer.getData(GPM_IMAGE_DND_TYPE) + if (gpmData && isImage) { + const { name, dataUrl } = JSON.parse(gpmData) as GpmDndImage + void saveDataUrlToTempFile(dataUrl, name).then(acceptInput).catch(() => {}) + return + } + + // Studio Assets / Downloads Browser file (already a real path on disk). + const dlData = e.dataTransfer.getData(FILE_DND) + if (dlData) { + const f = JSON.parse(dlData) as LibFile + if (isImage ? (!f.isVideo && !f.isAudio) : f.isVideo) acceptInput(f.path) + return + } + const acceptedAssetType = isImage ? 'image' : 'video' const assetData = e.dataTransfer.getData('asset') if (assetData) { @@ -368,6 +411,22 @@ export function ICLoraPanel({ e.preventDefault() setIsReferenceDragOver(false) + // Prompt Manager Pro image (data URL) — persist to a temp file for a real path. + const gpmData = e.dataTransfer.getData(GPM_IMAGE_DND_TYPE) + if (gpmData) { + const { name, dataUrl } = JSON.parse(gpmData) as GpmDndImage + void saveDataUrlToTempFile(dataUrl, name).then(setReferenceImagePath).catch(() => {}) + return + } + + // Studio Assets / Downloads Browser image (already a real path on disk). + const dlData = e.dataTransfer.getData(FILE_DND) + if (dlData) { + const f = JSON.parse(dlData) as LibFile + if (!f.isVideo && !f.isAudio) setReferenceImagePath(f.path) + return + } + const assetData = e.dataTransfer.getData('asset') if (assetData) { try { @@ -727,6 +786,21 @@ export function ICLoraPanel({

{processingStatus || 'Generating...'}

+ {onCancel && ( + + )}
) : (
diff --git a/frontend/components/LtxLogo.tsx b/frontend/components/LtxLogo.tsx deleted file mode 100644 index 904164144..000000000 --- a/frontend/components/LtxLogo.tsx +++ /dev/null @@ -1,18 +0,0 @@ -interface LtxLogoProps { - className?: string -} - -export function LtxLogo({ className = "h-6" }: LtxLogoProps) { - return ( - - - - - - ) -} diff --git a/frontend/components/PythonSetup.tsx b/frontend/components/PythonSetup.tsx index 3b1746f7e..ca06ee0d2 100644 --- a/frontend/components/PythonSetup.tsx +++ b/frontend/components/PythonSetup.tsx @@ -107,7 +107,7 @@ export function PythonSetup({ onReady }: PythonSetupProps) { // @ts-expect-error - Electron-specific CSS property WebkitAppRegion: 'drag' }}> - LTX Desktop + RiX Desktop Studio Pro
{/* Main Container */} @@ -292,7 +292,7 @@ export function PythonSetup({ onReady }: PythonSetupProps) { justifyContent: 'space-between', alignItems: 'center' }}> -
© 2026 Lightricks
+
© 2026 Michael Ricks
diff --git a/frontend/components/RixLogo.tsx b/frontend/components/RixLogo.tsx new file mode 100644 index 000000000..42cf17140 --- /dev/null +++ b/frontend/components/RixLogo.tsx @@ -0,0 +1,30 @@ +interface RixLogoProps { + className?: string +} + +/** + * "RiX" app wordmark (RiX Desktop Studio Pro). The R and X are width-locked text + * (textLength) so it renders identically across fonts, and the lowercase "i" is a + * custom mark — a rounded stem with a dot — to give it a bit of identity. Single + * color via currentColor, so it inherits the caller's text color (e.g. text-white) + * and scales to the caller's height (h-5 w-auto). Swap for a richer asset later. + */ +export function RixLogo({ className = 'h-6' }: RixLogoProps) { + const font = "'Segoe UI', system-ui, -apple-system, sans-serif" + return ( + + R + {/* custom lowercase "i": stem + dot */} + + + X + + ) +} diff --git a/frontend/components/SettingsModal.tsx b/frontend/components/SettingsModal.tsx index 191254d69..81320dab6 100644 --- a/frontend/components/SettingsModal.tsx +++ b/frontend/components/SettingsModal.tsx @@ -1,5 +1,6 @@ -import { AlertCircle, Check, Download, Film, Folder, HardDrive, Info, KeyRound, Settings, Sparkles, X, Zap } from 'lucide-react' +import { AlertCircle, Check, Download, Film, Folder, HardDrive, Info, KeyRound, Palette, Settings, Sparkles, X, Zap } from 'lucide-react' import React, { useEffect, useMemo, useRef, useState } from 'react' +import { PALETTES, getStoredPaletteId, selectPalette } from '../lib/theme' import { Button } from './ui/button' import { BaseModelSection } from './settings/BaseModelSection' import { useAppSettings, type AppSettings, DEFAULT_GEMINI_MODEL } from '../contexts/AppSettingsContext' @@ -24,7 +25,7 @@ interface SettingsModalProps { onCheckForUpdates: () => void } -type TabId = 'general' | 'models' | 'apiKeys' | 'promptEnhancer' | 'about' +type TabId = 'general' | 'appearance' | 'models' | 'apiKeys' | 'promptEnhancer' | 'about' /** A checkpoint this modal can download: the text encoder or the optional prompt enhancer. */ type TextEncodingCp = NonNullable['cp_to_download']> @@ -268,6 +269,7 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda const [analyticsEnabled, setAnalyticsEnabled] = useState(false) const [autoCheckUpdates, setAutoCheckUpdatesState] = useState(true) const [projectAssetsPath, setProjectAssetsPath] = useState('') + const [paletteId, setPaletteId] = useState(getStoredPaletteId) const isMac = window.electronAPI.platform === 'darwin' const updateAction = aboutUpdateAction(update.state, onOpenUpdate, onCheckForUpdates, isMac, autoCheckUpdates) @@ -278,6 +280,14 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda } }, [isOpen, initialTab]) + // The Models tab is hidden in force-API mode; don't let the selection get stuck there + // (e.g. via initialTab or a stale value). + useEffect(() => { + if (forceApiGenerations && activeTab === 'models') { + setActiveTab('general') + } + }, [forceApiGenerations, activeTab]) + useEffect(() => { if (isOpen && initialReason === 'geminiKeyRequired') { geminiApiKey.openAndFocus() @@ -523,6 +533,7 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda const tabs = [ { id: 'general' as TabId, label: 'General', icon: Settings }, + { id: 'appearance' as TabId, label: 'Appearance', icon: Palette }, // The Models tab is local-model management — irrelevant (and non-functional) when all // generation is forced through the API, so hide it in that mode. ...(!forceApiGenerations ? [{ id: 'models' as TabId, label: 'Models', icon: HardDrive }] : []), @@ -531,6 +542,11 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda { id: 'about' as TabId, label: 'About', icon: Info }, ] + const handleSelectPalette = (id: string) => { + setPaletteId(id) + selectPalette(id) + } + return (
{/* Backdrop */} @@ -540,7 +556,7 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda /> {/* Modal */} -
+
{/* Header */}
@@ -557,8 +573,11 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda
- {/* Tabs */} -
+ {/* Tabs — the buttons are shrink-0/nowrap by design, so a row that outgrows the + modal overflows rather than compressing. Scroll instead of spilling past the + edge (as the 6th "Appearance" tab did). The wider max-w-3xl above means this + shouldn't normally engage; it's here so adding a 7th tab can't break the layout. */} +
{tabs.map((tab) => { const Icon = tab.icon return ( @@ -993,9 +1012,12 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda {cudaAvailable && ( Reuses an already-built transformer across stage 1/stage 2 within one generation - instead of reloading it from disk twice. Experimental: only - applies on high-VRAM cards (32GB+); no effect otherwise.} + description={<>Keeps the built transformer resident instead of reloading it from disk on every + stage. On 32GB+ cards it stays in VRAM for the generation; on 24GB cards it stays in + system RAM (~23GB pinned) for the whole session, so repeat generations skip the model + load entirely. Also keeps the small VAE/upsampler/audio models warm between generations + (~2-4GB VRAM). Experimental: turn off if other + apps need the RAM back.} enabled={settings.diffusionStageCacheEnabled} onToggle={handleToggleDiffusionStageCache} statusOn="Skipping redundant transformer reloads" @@ -1090,7 +1112,7 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda

- Share anonymous usage data to help improve LTX Desktop. + Share anonymous usage data to help improve RiX Desktop Studio Pro. Only basic technical information is collected — never personal data or generated content.

@@ -1114,6 +1136,51 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda )} + {activeTab === 'appearance' && ( +
+
+

Color palette

+

+ Changes the main app's colors. Prompt Manager Pro panels keep their own styling. +

+
+
+ {PALETTES.map((palette) => { + const selected = palette.id === paletteId + return ( + + ) + })} +
+
+ )} + {activeTab === 'models' && !forceApiGenerations && } {activeTab === 'apiKeys' && ( @@ -1394,7 +1461,7 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda

When enabled, Generate rewrites your prompt with visual detail, sound, and camera motion before the model sees it. Local generations use the on-device enhancer; - LTX API text encoding enhances on the server. The Enhance button in Gen Space is + LTX API text encoding enhances on the server. The Enhance button in Create is separate — it rewrites the prompt box so you can edit it first. Control independently for each generation type.

@@ -1442,6 +1509,11 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda

{settings.promptEnhancerEnabledI2V ? 'Prompts will be enhanced before I2V generation' : 'I2V prompts used as-is'}

+

+ ⚠️ Rewrites your prompt via the LTX API. On local generations the video may + follow the rewritten text instead of your source image (character/scene can + change a few frames in). Leave off for faithful image-to-video. +

{/* App Identity */}
-

LTX Desktop

+

RiX Desktop Studio Pro

Version {appVersion || '...'}

AI-Powered Video Editor

@@ -1630,7 +1702,7 @@ export function SettingsModal({ isOpen, onClose, initialTab, initialReason, upda {/* Copyright */}

- Copyright © 2026 Lightricks + Copyright © 2026 Michael Ricks

)} diff --git a/frontend/components/SettingsPanel.tsx b/frontend/components/SettingsPanel.tsx index 4986a8f81..231edeca3 100644 --- a/frontend/components/SettingsPanel.tsx +++ b/frontend/components/SettingsPanel.tsx @@ -34,8 +34,10 @@ export interface GenerationSettings { imageResolution: string imageAspectRatio: string imageSteps: number + imageModel?: 'z-image-turbo' | 'krea-2-turbo' variations?: number // Number of image variations to generate imageEditStrength?: number // Denoising strength when editing an existing image + imageVariation?: number // Z-Image text-to-image composition variety across seeds (0 = off, 1 = max) } interface SettingsPanelProps { diff --git a/frontend/components/gpm/DownloadsBrowser.tsx b/frontend/components/gpm/DownloadsBrowser.tsx new file mode 100644 index 000000000..bbf78d103 --- /dev/null +++ b/frontend/components/gpm/DownloadsBrowser.tsx @@ -0,0 +1,373 @@ +/** + * Downloads Browser — a left-side media library backed by REAL folders/files + * on disk (via Electron IPC under /PromptManagerPro). Create/rename/ + * delete folders, add images/videos/audio, move media between folders, and + * drag to reorder folders. Folders are an accordion (like the extension) — + * each expands inline to show its files. Unlike the browser extension + * there's no directory-handle permission dance — it's the real filesystem. + */ +import { useEffect, useRef, useState } from 'react' +import { ChevronLeft, ChevronsDownUp, FolderInput, FolderPlus, FolderOpen, GripVertical, Pencil, RefreshCw, Send, Trash2, Upload, X } from 'lucide-react' +import { pathToFileUrl } from '../../lib/file-url' +import { useProjects } from '../../contexts/ProjectContext' +import { MediaThumb } from './MediaThumb' +import { useDownloadsBrowserOpen, setDownloadsBrowserOpen } from './downloads-browser-store' +import { GPM_IMAGE_DND_TYPE, saveDataUrlToTempFile, type GpmDndImage } from './gpm-image-file' +import { setStudioAssetsTarget } from '../../lib/studio-assets-target' +import { C } from './gpm-theme' + +export interface LibFile { folder: string; name: string; path: string; isVideo: boolean; isAudio: boolean; mtimeMs: number } +type TypeFilter = 'all' | 'image' | 'video' | 'audio' +const ORDER_KEY = 'gpm_dl_order' +export const FILE_DND = 'application/x-gpm-dl-file' +const FOLDER_DND = 'application/x-gpm-dl-folder' + +function loadOrder(): string[] { try { return JSON.parse(localStorage.getItem(ORDER_KEY) || '[]') as string[] } catch { return [] } } +function saveOrder(o: string[]) { localStorage.setItem(ORDER_KEY, JSON.stringify(o)) } +function applyOrder(folders: string[]): string[] { + const ord = loadOrder().filter((f) => folders.includes(f)) + return [...ord, ...folders.filter((f) => !ord.includes(f))] +} + +/** A solid, thick triangle disclosure indicator (matches the extension's look better than a thin chevron stroke). */ +function Triangle({ open, color }: { open: boolean; color: string }) { + return ( + + + + ) +} + +function Dock({ onClose }: { onClose: () => void }) { + const { setGenSpaceInputImagePath, setCurrentTab } = useProjects() + const [folders, setFolders] = useState([]) + const [files, setFiles] = useState([]) + const [openFolders, setOpenFolders] = useState>(new Set()) + const [newName, setNewName] = useState('') + const newNameRef = useRef(null) + const [renaming, setRenaming] = useState(null) + const [renameVal, setRenameVal] = useState('') + const [toast, setToast] = useState(null) + const [moveTarget, setMoveTarget] = useState(null) + const [moveDest, setMoveDest] = useState('') + const [dropLine, setDropLine] = useState<{ target: string; before: boolean } | null>(null) + const [typeFilter, setTypeFilter] = useState('all') + const [rootPath, setRootPath] = useState('') + const dragFolder = useRef(null) + const flash = (m: string) => { setToast(m); setTimeout(() => setToast(null), 1600) } + + const api = window.electronAPI + const refresh = async () => { + if (!api) return + const d = await api.gpmLibList() + const ordered = applyOrder(d.folders) + setFolders(ordered) + setFiles(d.files) + setOpenFolders((prev) => new Set([...prev].filter((f) => ordered.includes(f)))) + const r = await api.gpmLibGetRoot() + setRootPath(r.root) + } + useEffect(() => { void refresh() }, []) + + // Live-refresh from a filesystem watch on the library folder: files saved into + // it (e.g. "Save video"/"Save frame") appear immediately, and only on real + // disk changes — no re-listing on every window focus. A ref keeps the + // subscription pinned to the latest refresh without re-subscribing each render. + const refreshRef = useRef(refresh) + refreshRef.current = refresh + useEffect(() => { + const unsubscribe = api?.onGpmLibChanged?.(() => void refreshRef.current()) + return () => unsubscribe?.() + }, []) + + const chooseFolder = async () => { + if (!api) return + const r = await api.gpmLibChooseRoot() + if (!r.root) return + setOpenFolders(new Set()) + flash(`Library set to ${r.root}`) + await refresh() + } + + const toggleFolder = (name: string) => { + setOpenFolders((prev) => { + const next = new Set(prev) + if (next.has(name)) next.delete(name) + else { next.add(name); setStudioAssetsTarget(name) } // opening = "the folder I'm working in" → quick-save target + return next + }) + } + const collapseAll = () => setOpenFolders(new Set()) + + const createFolder = async () => { + const n = newName.trim() + // The name comes from the inline box, not a popup — clicking the button with an + // empty box is a common "where do I type the name?" mistake, so nudge + focus it. + if (!n) { flash('Type a folder name in the box first, then click Create'); newNameRef.current?.focus(); return } + if (!api) return + const r = await api.gpmLibCreateFolder({ name: n }) + if (r.success) { + // New folder at the TOP of the list (not appended below an expanded Inbox of + // hundreds of files, where it looks like nothing happened). + saveOrder([n, ...loadOrder().filter((f) => f !== n)]); setNewName('') + setOpenFolders((prev) => new Set(prev).add(n)) + setStudioAssetsTarget(n) // route quick-saves into the folder just created + await refresh() + } else flash(r.error) + } + const commitRename = async (from: string) => { + const to = renameVal.trim(); setRenaming(null) + if (!to || to === from || !api) return + const r = await api.gpmLibRenameFolder({ from, to }) + if (r.success) { + saveOrder(loadOrder().map((f) => (f === from ? to : f))) + setOpenFolders((prev) => { const next = new Set(prev); if (next.delete(from)) next.add(to); return next }) + await refresh() + } else flash(r.error) + } + const deleteFolder = async (name: string) => { + if (!api || !window.confirm(`Delete folder "${name}" and its files?`)) return + const r = await api.gpmLibDeleteFolder({ name }) + if (r.success) { saveOrder(loadOrder().filter((f) => f !== name)); await refresh() } else flash(r.error) + } + const addFiles = async (folder: string) => { + if (!api || !folder) return + const paths = await api.showOpenFileDialog({ + title: `Add to ${folder}`, properties: ['openFile', 'multiSelections'], + filters: [{ name: 'Media', extensions: ['png', 'jpg', 'jpeg', 'webp', 'gif', 'bmp', 'mp4', 'webm', 'mkv', 'mov', 'avi', 'mp3', 'wav', 'ogg', 'aac', 'flac', 'm4a'] }], + }) + if (!paths || paths.length === 0) return + const r = await api.gpmLibAddFiles({ folder, srcPaths: paths }) + if (r.success) { flash(`Added ${r.added} file(s)`); await refresh() } else flash(r.error) + } + const moveFile = async (f: LibFile, toFolder: string) => { + if (!api || f.folder === toFolder) return + const r = await api.gpmLibMoveFile({ fromFolder: f.folder, name: f.name, toFolder }) + if (r.success) { await refresh() } else flash(r.error) + } + const deleteFile = async (f: LibFile) => { + if (!api) return + const r = await api.gpmLibDeleteFile({ folder: f.folder, name: f.name }) + if (r.success) await refresh(); else flash(r.error) + } + const sendToGenSpace = (f: LibFile) => { + if (f.isVideo) return + setGenSpaceInputImagePath(f.path) + setCurrentTab('gen-space') + flash('Sent to Create') + } + const openMove = (f: LibFile) => { + setMoveTarget(f) + setMoveDest(folders.find((x) => x !== f.folder) ?? '') + } + const confirmMove = async () => { + if (!moveTarget || !moveDest) return + await moveFile(moveTarget, moveDest) + setMoveTarget(null) + } + + // Drag-reorder folders, with a blue drop-line indicator like the extension. + const onFolderDrop = (target: string, before: boolean) => { + const dragged = dragFolder.current + dragFolder.current = null + setDropLine(null) + if (!dragged || dragged === target) return + const order = folders.filter((f) => f !== dragged) + let idx = order.indexOf(target) + if (!before) idx += 1 + order.splice(idx, 0, dragged) + saveOrder(order); setFolders(order) + } + + const filterFiles = (list: LibFile[]) => list.filter((f) => ( + typeFilter === 'all' ? true + : typeFilter === 'video' ? f.isVideo + : typeFilter === 'audio' ? f.isAudio + : !f.isVideo && !f.isAudio + )) + + return ( +
+
+ Studio Assets +
+ + + + +
+
+ +
+ {rootPath} + +
+ +
+ setNewName(e.target.value)} onKeyDown={(e) => { if (e.key === 'Enter') void createFolder() }} placeholder="Type a folder name…" className="flex-1 rounded-md px-2 py-1.5 text-xs outline-none" style={{ background: C.elev, color: C.text, border: `1px solid ${C.border}` }} /> + +
+ +
+ Filter + +
+ + {/* Folders — accordion: click a folder to expand its files inline. */} +
+ {folders.map((f) => { + // Newest first: most-recently-saved/added files appear at the top of + // each folder instead of the bottom (mtime desc), so fresh saves don't + // require scrolling past the whole folder to find. + const folderFiles = filterFiles(files.filter((x) => x.folder === f)).sort((a, b) => b.mtimeMs - a.mtimeMs) + const isOpen = openFolders.has(f) + return ( +
{ + if (!e.dataTransfer.types.includes(FOLDER_DND) && !e.dataTransfer.types.includes(FILE_DND) && !e.dataTransfer.types.includes(GPM_IMAGE_DND_TYPE)) return + e.preventDefault() + if (dragFolder.current && dragFolder.current !== f) { + const r = e.currentTarget.getBoundingClientRect() + const before = e.clientY - r.top < r.height / 2 + setDropLine({ target: f, before }) + } + }} + onDragLeave={() => setDropLine((d) => (d?.target === f ? null : d))} + onDrop={(e) => { + e.preventDefault() + if (e.dataTransfer.types.includes(FILE_DND)) { const d = JSON.parse(e.dataTransfer.getData(FILE_DND)) as LibFile; void moveFile(d, f) } + else if (e.dataTransfer.types.includes(GPM_IMAGE_DND_TYPE)) { + const raw = e.dataTransfer.getData(GPM_IMAGE_DND_TYPE) + if (raw && api) { + const { name, dataUrl } = JSON.parse(raw) as GpmDndImage + void saveDataUrlToTempFile(dataUrl, name) + .then((path) => api.gpmLibAddFiles({ folder: f, srcPaths: [path] })) + .then((r) => { if (r.success) { flash(`Added to ${f}`); void refresh() } else flash(r.error) }) + } + } + else if (dragFolder.current) onFolderDrop(f, dropLine?.before ?? true) + }} + > +
{ dragFolder.current = f; e.dataTransfer.setData(FOLDER_DND, f); e.dataTransfer.effectAllowed = 'move' }} + onDragEnd={() => { dragFolder.current = null; setDropLine(null) }} + onClick={() => toggleFolder(f)} + className="relative flex items-center gap-1.5 px-3 py-2.5 cursor-pointer group" + style={{ background: isOpen ? 'rgb(var(--accent) / 0.12)' : 'transparent', borderLeft: `2px solid ${isOpen ? C.blue : 'transparent'}`, borderBottom: `1px solid ${C.border}` }} + > + {dropLine?.target === f && ( +
+ )} + + + {renaming === f ? ( + setRenameVal(e.target.value)} onClick={(e) => e.stopPropagation()} onBlur={() => void commitRename(f)} onKeyDown={(e) => { if (e.key === 'Enter') void commitRename(f) }} className="flex-1 rounded px-1 py-0.5 text-xs outline-none" style={{ background: C.elev, color: C.text, border: `1px solid ${C.border}` }} /> + ) : ( + {f} + )} + {folderFiles.length} + { e.stopPropagation(); setRenaming(f); setRenameVal(f) }} /> + { e.stopPropagation(); void deleteFolder(f) }} /> +
+ + {isOpen && ( +
+
+ +
+ {folderFiles.length === 0 + ?

{files.filter((x) => x.folder === f).length === 0 ? 'Empty. Add files or drag media here.' : 'No files match this filter.'}

+ : ( +
+ {folderFiles.map((file) => ( + { e.dataTransfer.setData(FILE_DND, JSON.stringify(file)); e.dataTransfer.effectAllowed = 'copyMove' }} + > +
+ {!file.isVideo && !file.isAudio && ( + + )} + + +
+
+ ))} +
+ )} +
+ )} +
+ ) + })} + {folders.length === 0 &&

No folders yet.

} +
+ + {toast &&
{toast}
} + + + + {moveTarget && ( +
setMoveTarget(null)}> +
e.stopPropagation()} + > +

Move To Folder…

+

"{moveTarget.name}"

+

From: {moveTarget.folder}

+ +
+ + +
+
+
+ )} +
+ ) +} + +export function DownloadsBrowser() { + const open = useDownloadsBrowserOpen() + if (!open) { + return ( + + ) + } + return setDownloadsBrowserOpen(false)} /> +} diff --git a/frontend/components/gpm/MediaThumb.tsx b/frontend/components/gpm/MediaThumb.tsx new file mode 100644 index 000000000..b701066d7 --- /dev/null +++ b/frontend/components/gpm/MediaThumb.tsx @@ -0,0 +1,61 @@ +/** + * Widescreen (16:9) media thumbnail with hover side-preview and double-click + * lightbox. Renders a real