Compare commits
6 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 0d3d3a3855 | |||
| d6053703e0 | |||
| fe94199e5d | |||
| 7dc613f80c | |||
| 203d4842d4 | |||
| 1799988521 |
@@ -0,0 +1,277 @@
|
||||
---
|
||||
name: train-model
|
||||
description: End-to-end agent for training LTX-2 models. Probes filesystem and GPU, picks the right conditioning mode from the user's intent, prepares the dataset (scenes, captions, references), preprocesses, autotunes, launches, and monitors training. Use when the user wants to train, fine-tune, LoRA, or otherwise produce a custom LTX-2 model.
|
||||
argument-hint: [optional source path or run name]
|
||||
user-invocable: true
|
||||
allowed-tools: Bash, Read, Grep, Glob, Edit, Write, Agent, AskUserQuestion, TodoWrite
|
||||
---
|
||||
|
||||
# Train Model — Orchestrator
|
||||
|
||||
Take the user from "I want to train something" to a running, monitored training job — automating the mechanical glue (dataset layout, captioning, preprocessing, config patching, launch, monitoring) without making silent decisions on their behalf.
|
||||
|
||||
> **Source of truth for the trainer:** [`packages/ltx-trainer/docs/`](../../../packages/ltx-trainer/docs/). This skill orchestrates those scripts; it does not duplicate their reference content.
|
||||
|
||||
## Hard Invariants
|
||||
|
||||
These are non-negotiable. Re-read them before every action that touches the filesystem or starts a process.
|
||||
|
||||
1. **No file mutation outside the run workspace without explicit user approval.** The workspace is `./projects/<run-name>/`. Never overwrite, move, delete, or modify any user file or directory outside it without surfacing an explicit ask. Hours of dataset work must never be silently destroyed.
|
||||
2. **No heavy work before plan approval.** Captioning, preprocessing, training, and autotune do **not** start until the user approves `plan.md` (Phase 4). Probing the filesystem and running `nvidia-smi` is fine; encoding videos or downloading models is not.
|
||||
3. **No silent assumptions.** Every non-trivial default appears under "Assumptions" in `plan.md`.
|
||||
4. **No code changes to the trainer package without explicit consent.** If the user's intent doesn't map to a supported configuration, follow the **Escape Hatch** section below — do not unilaterally edit `packages/ltx-trainer/`.
|
||||
5. **No fabricated claims about training outcomes or data sufficiency.** Do not assert how well something *will* train, whether a dataset is "too small," how many samples/seconds of audio are "enough," which modality will "learn better," expected quality, or any similar prediction — you have no grounded basis for these, they're frequently wrong, and they mislead users. Stick to facts you can substantiate: what the trainer/docs actually say, observed numbers (loss, step time, VRAM), counts, and the user's own stated goals. If the user asks for a recommendation that depends on such judgment, you may share it **only** as an explicitly-flagged uncertainty ("I'm not sure — you'd have to try it"), never as authoritative fact. When in doubt, say less.
|
||||
|
||||
## Keep the User Informed
|
||||
|
||||
Most users don't know how this skill works under the hood — they don't know what "preprocessing," "a one-sample sanity check," or "autotune" mean or why they're happening. Narrate the run in plain language so it never feels like a black box:
|
||||
|
||||
- **Entering a phase:** one or two sentences on *what you're about to do and why* — in user terms, not jargon.
|
||||
- **Leaving a phase:** one line on *what came out of it* (e.g. "captioned 9 clips," "found the fastest stable config: batch 1, ~3s/step").
|
||||
- **Explain the non-obvious phases explicitly** — these are the ones that confuse people:
|
||||
- *Sanity check (Phase 6):* "Before the full run, I do a quick dry run on a single clip at your target resolution. It catches out-of-memory or config problems in a couple of minutes instead of failing hours into training."
|
||||
- *Autotune (Phase 6):* "Then I try a few configuration variants on that one clip to pick the fastest one that still fits your GPU — so the full run is as fast as it can be."
|
||||
- *Preprocess (Phase 7):* "I'm encoding your videos into the compressed latents the trainer reads. One-time step; the trained model never sees the raw videos directly."
|
||||
- Keep it concise — a sentence or two per transition, not walls of text or raw logs. This is running commentary, not a replacement for the upfront plan (Phase 4) or the status reports (Phase 8).
|
||||
- Long-running steps (preprocess, training): say roughly how long it'll take and that you'll report back, so silence doesn't read as "stuck."
|
||||
- **Describe what you're doing — don't editorialize about how it'll turn out.** Narration covers *what's happening*; it must not drift into unfounded predictions about training quality or data sufficiency (e.g. "26s of audio is too little," "voice won't learn well"). Those are fabricated claims — see Hard Invariant #5. State facts and the user's choices; leave the "will it be good?" judgment to the user watching the results.
|
||||
|
||||
## Phase 0 — Set Up
|
||||
|
||||
Create the workspace and todos.
|
||||
|
||||
1. Pick a workspace root in this order (use first writable):
|
||||
- `$LTX_TRAININGS_DIR`
|
||||
- `/data/ltx-trainings/`
|
||||
- `/workspace/ltx-trainings/`
|
||||
- `./projects/` (repo-relative — preferred default in this repo)
|
||||
2. Derive a tentative `<run-name>` from the user's words; finalise after Phase 1 once the mode is known. Format: `<mode>-<dataset-name>-<YYYYMMDD-HHMM>`.
|
||||
3. Create `<workspace>/<run-name>/` and seed empty subdirs: `dataset/`, `outputs/`, `overfit/`.
|
||||
4. Create a todo list covering Phases 1–9 so the user can see progress.
|
||||
|
||||
## Phase 1 — Intent
|
||||
|
||||
Ask one question, framed in user terms (not jargon):
|
||||
|
||||
> What do you want the model to learn? Examples: "generate videos from text," "make a LoRA of a specific style," "extend a video forward in time," "add sound effects to a silent video," "fill in masked regions of a video."
|
||||
|
||||
Map the answer to one or more conditioning modes via `references/mode-selector.md`. If the requested capability has no mapping, go to **Escape Hatch**.
|
||||
|
||||
### Plain concept/style LoRA → ask how it'll be used, default to I2V
|
||||
|
||||
A "train a LoRA on X" request (a character/style/concept LoRA, no specific conditioning task) maps to either T2V or I2V. These aren't locked to inference: LoRA weights are pipeline-agnostic (the same checkpoint loads in both T2V and I2V inference), and the `i2v_lora` config applies first-frame conditioning with **`probability: 0.5`** — so it learns **both** conditioned (I2V) and unconditioned (T2V) generation in one run, and the first frame is taken automatically from each training clip (no extra data prep). I2V is therefore a versatile superset.
|
||||
|
||||
Ask how they intend to use the result:
|
||||
|
||||
> Will you generate videos from **text alone** (T2V), from a **starting image** (I2V), or **both / not sure**?
|
||||
|
||||
- **Both / not sure (default):** use `i2v_lora` (probabilistic first-frame) — works for both at inference.
|
||||
- **I2V:** `i2v_lora`.
|
||||
- **Text only:** `t2v_lora`.
|
||||
|
||||
(This only applies to plain concept/style LoRAs. A specific task — extension, inpainting, foley, IC-LoRA, etc. — maps directly to its mode via `mode-selector.md`; no usage question needed.)
|
||||
|
||||
### Confirm the mode before proceeding
|
||||
|
||||
Once the mode is determined, **state it plainly and confirm it** before doing any probing or work — a wrong inference is cheap to fix here and expensive later:
|
||||
|
||||
> "Got it — I'll train an **I2V LoRA** (usable for both image-to-video and text-to-video at inference). Sound right?"
|
||||
|
||||
The mode also appears in the plan (Phase 4), but confirm it here so the rest of the flow isn't built on a wrong guess.
|
||||
|
||||
## Phase 2 — Probe
|
||||
|
||||
No questions in this phase. Inspect what's already there. Use `references/onboarding.md` as the source of truth for the prerequisite checklist and what to do when something is missing.
|
||||
|
||||
### Filesystem probe
|
||||
- If the user pointed at a path, classify: directory of raw videos, single long video, directory with a metadata file (CSV/JSON/JSONL), existing `.precomputed/`, partial outputs from a prior run.
|
||||
- For metadata files, identify columns: `video`/`media_path`, `caption`, `audio`, `reference_video`, `video_mask`, etc. (see `packages/ltx-trainer/docs/dataset-preparation.md`).
|
||||
- **Clip lengths (small datasets only):** for datasets up to a few hundred clips, `ffprobe` each clip's frame count and note the **minimum**. Clips shorter than the target frame bucket are silently skipped by `process_dataset.py`, so the shortest clip caps the achievable frame count — feed this into the Phase 3 resolution/frame-count choice (pick a bucket the clips support, or plan multi-bucket). Skip this per-clip probe for large datasets (too slow); rely instead on the post-preprocess reconciliation in Phase 7, which flags any dropped clips regardless of dataset size.
|
||||
- Check for an existing `<workspace>/<run-name>/` and whether `outputs/checkpoints/` contains a prior checkpoint (`lora_weights_step_*.safetensors` or `model_weights_step_*.safetensors`, plus a matching `training_state_step_*.pt` when resume state is enabled). This is a **resume candidate** — but note the trainer does *not* auto-resume from the output dir; resuming requires explicitly setting `model.load_checkpoint` in `config.yaml` to that checkpoint path. See `phases/launch-and-monitor.md` for the resume flow.
|
||||
- Check disk space at the workspace root. Preprocessed latents, checkpoints, and validation samples add up across a run; surface the available space alongside a rough sense of what one run consumes (one preprocessed bucket scales with sample count and resolution; each checkpoint is several GB), and warn the user if free space looks tight given their dataset size.
|
||||
|
||||
### Hardware probe
|
||||
- `nvidia-smi --query-gpu=name,memory.total,driver_version --format=csv,noheader` → GPU model, count, VRAM. Stop the run if no CUDA GPU.
|
||||
- W&B login state — use wandb's **own** credential resolution (source-agnostic: covers env var, netrc, and the wandb settings file), not a hand-rolled netrc grep and **not** `wandb status` (which misleadingly reports `api_key: null` even when logged in):
|
||||
|
||||
```bash
|
||||
uv run python -c "import wandb; print(bool(wandb.Api().api_key))" # True => logged in
|
||||
```
|
||||
|
||||
`True` → W&B is available, enable it. `False` → genuinely not logged in. If the check errors or is ambiguous, **ask the user** rather than silently disabling — a wrong "W&B off" assumption can incorrectly disable expected tracking.
|
||||
- Apply `references/hardware-profiles.md` to derive defaults (32GB / 40–60GB / 80GB+ VRAM tier).
|
||||
|
||||
### Prerequisite probe (first-run sanity)
|
||||
- `uv` installed (`command -v uv`).
|
||||
- Workspace synced (lockfile present + `ltx-trainer` import works).
|
||||
- LTX-2 `.safetensors` and Gemma text encoder dir present in `/models/`, `~/models/`, or `$LTX_MODELS_DIR`.
|
||||
- Captioner availability: Gemini auth (`GEMINI_API_KEY`/`GOOGLE_API_KEY` or gcloud/Vertex) OR a ≥40 GiB GPU to host the Qwen3-Omni-30B vLLM server (FP8; bf16 needs ≥66 GiB). On typical consumer GPUs (24/32 GB), Gemini is effectively the only local-free option — see `references/onboarding.md`.
|
||||
|
||||
For any missing prerequisite, **do not silently fail in a later phase**. Present the finding in chat with the specific next step from `references/onboarding.md`. The skill may offer to auto-install / auto-download missing pieces — but only ever after asking the user explicitly, one item at a time (model downloads are tens of GB each). Never auto-modify shell rc files or system config without consent.
|
||||
|
||||
When everything (or what the user agreed to set up) is in place, fold the resolved paths into the plan's Assumptions section.
|
||||
|
||||
### Pre-existing artifacts in the run dir
|
||||
|
||||
If the run dir (or a user-supplied path) already contains artifacts from a prior session, classify each into one of three buckets — only the third prompts the user:
|
||||
|
||||
1. **Deterministically verifiable → verify, then reuse silently (or stop).** `.precomputed/` latents: load a sample, check tensor shapes + modality coverage against the target (Phase 7). Match → reuse, no question. Mismatch/incomplete → stop and ask (reuse-at-old-spec / re-preprocess to a new dir / abort).
|
||||
2. **Cheap, fully-derived intermediates → regenerate silently.** `overfit/`, eval renders, sanity/temp configs, one-sample metadata. Delete and redo; don't ask.
|
||||
3. **Expensive AND not deterministically verifiable → ask.** Captions (`dataset.json`) and trained checkpoints/outputs. We can't programmatically decide whether existing captions or a half-finished run are what the user wants now, so surface what was found (counts, and how/when produced if knowable) and ask: reuse vs regenerate (for checkpoints: resume vs fresh).
|
||||
|
||||
Principle: only ask when reuse-vs-regenerate is a genuine judgment call with cost either way. Never silently delete user-supplied data (hard invariant #1).
|
||||
|
||||
## Phase 3 — Ask (minimum viable set)
|
||||
|
||||
Ask only what cannot be inferred. Use `AskUserQuestion` with multiple-choice when possible. Typical questions:
|
||||
|
||||
- **Target resolution / frame count** — propose a default per mode (e.g., `768x512x49` for T2V LoRA on consumer GPUs); offer overrides.
|
||||
- **Training steps** — if dataset size doesn't pin it, propose a default (e.g., 2000 for small LoRA datasets).
|
||||
- **LoRA trigger word / concept name** — only for style/concept LoRAs. Ask **only** for the word itself (or whether they want one). **Never** ask or mention *how* it's injected — it's always the `--lora-trigger` flag (passed to `process_dataset.py`, which forwards it to `process_captions.py` where the prepend happens); this is a fixed implementation detail. Presenting caption-injection as an option creates unnecessary confusion.
|
||||
- **Captioner backend** — only if more than one path is viable (e.g. a ≥40 GiB GPU can host the Qwen3-Omni-30B server *and* Gemini auth is available). On typical consumer GPUs, default to `gemini_flash` and surface that Gemini auth is required rather than asking.
|
||||
- **Model paths** — only if not found in the probe.
|
||||
|
||||
Never ask anything answerable by `ls`, `nvidia-smi`, or the W&B credential check above.
|
||||
|
||||
## Phase 4 — Plan
|
||||
|
||||
Write the plan to `<workspace>/<run-name>/plan.md` using `references/plan-template.md`. Present it to the user in chat (don't just dump the file path). Wait for explicit approval before proceeding.
|
||||
|
||||
If the user requests changes, edit the plan and re-present. Do not start Phase 5 until approval.
|
||||
|
||||
## Phase 5 — Prepare Dataset
|
||||
|
||||
If captioned metadata already exists with all required columns for the chosen mode, skip this phase. Otherwise follow `phases/prepare-dataset.md` — re-read it before acting.
|
||||
|
||||
**Captioning gate:** caption a 3-sample spot-check first, show the captions in full, and **STOP for explicit user approval** before captioning the full set. The user must approve or give tuning instructions — never auto-proceed to the full pass. (Details in `phases/prepare-dataset.md`.)
|
||||
|
||||
**Conditioning-inputs gate:** modes that need a reference (V2V/A2A/AV2AV IC-LoRA) or a mask (video/audio inpainting) require a per-sample input that encodes the user's specific idea. **Ask the user to provide it** — never invent the method (no defaulting to Canny/depth/generic masks). Only help generate it if the user explicitly asks, following *their* approach. Don't enter preprocessing for these modes without the input present. (Details in `phases/prepare-dataset.md` Step 4.)
|
||||
|
||||
## Phase 6 — Sanity Check + Autotune (always run)
|
||||
|
||||
**Tell the user what this phase is before starting it** — it's the most opaque to someone who doesn't know the design (see "Keep the User Informed"). In plain terms: a quick single-clip dry run at the target resolution to catch OOM/config errors cheaply, followed by trying a few config variants to pick the fastest stable one.
|
||||
|
||||
Run **at the full target resolution** on **one sample** before the full preprocess. Purpose:
|
||||
1. Catch OOM / config errors before paying the full preprocessing cost.
|
||||
2. Empirically pick the fastest stable config via a small sweep.
|
||||
|
||||
Steps:
|
||||
1. Pick one sample from the dataset metadata. Preprocess just that sample to `<workspace>/<run-name>/overfit/.precomputed/` (see `phases/preprocess-dataset.md` — use it in "one-sample" mode).
|
||||
2. Generate a temp config matching the planned full-run config but with `data.preprocessed_data_root: overfit/.precomputed`, `optimization.steps: 50`, `validation.interval: 50`, `checkpoints.interval: null`.
|
||||
3. Run the **baseline trial**: the conservative config from the matched VRAM tier (32GB tier = `t2v_lora_low_vram.yaml` defaults; 80GB+ tier = `t2v_lora.yaml` defaults — see `references/hardware-profiles.md`).
|
||||
4. **Success criteria** (all required):
|
||||
- No OOM, no NaN loss, no crash.
|
||||
- All 50 training steps complete.
|
||||
- Validation sample at step 50 generates successfully (validation pass is a real OOM risk — do not skip).
|
||||
- **For audio runs:** the one-sample `audio_latents/` is non-empty (the trainer log should report audio enabled). A joint/audio run that silently produced no audio latents is a failure even if steps complete — see the audio gate in `phases/preprocess-dataset.md`.
|
||||
- **Loss is NOT a success criterion.** Loss can be non-monotonic even when training is healthy.
|
||||
5. **Autotune sweep** — incremental, capped at 5 trials total. Each trial = current best + one change. Stop on OOM (revert), no step-time improvement, or 5 trials:
|
||||
- Trial 2: `quantization: null` (disable transformer quantization) if VRAM headroom allows.
|
||||
- Trial 3: `optimizer_type: adamw` (disable 8-bit optimizer) if headroom allows.
|
||||
- Trial 4: `batch_size` up (1 → 2 → 4). Adjust `gradient_accumulation_steps` proportionally to keep effective batch constant. **Note:** batch size can't be meaningfully tested on the one-sample set — defer it (test on the full set later, or just keep `batch_size: 1`, which is preferable for small concept-LoRA datasets anyway).
|
||||
- **Do not sweep:** resolution (user decision), `acceleration.load_text_encoder_in_8bit` (one-time, no step-time impact), `enable_gradient_checkpointing` (the trainer's example configs ship with it on; on the 80GB+ tier you *may* try it off, but for the 22B model it usually OOMs even with tens of GB of apparent headroom — don't expect a win; never turn it off on the 32GB tier).
|
||||
6. Collect per trial: step time and peak VRAM. **Prefer the trainer's own end-of-run stats** (it prints total time / step time and peak GPU memory) — no external timing tool is needed (`/usr/bin/time` is often not installed). Append results to `<workspace>/<run-name>/autotune.log`.
|
||||
7. Winning trial's deltas are patched into the main `config.yaml`. Summarise the sweep to the user (one line per trial + winner).
|
||||
|
||||
If the baseline trial fails, consult `references/troubleshooting.md`, propose a fix, re-run. Never push forward to the full preprocess after a failed sanity check.
|
||||
|
||||
## Phase 7 — Full Preprocess
|
||||
|
||||
Follow `phases/preprocess-dataset.md`. Re-read it before acting.
|
||||
|
||||
**If `.precomputed/` already exists at the target path** (user-supplied or prior run), the phase verifies shapes and modality coverage before reuse. On mismatch it stops and asks — never silently overwrites.
|
||||
|
||||
## Phase 8 — Launch & Monitor
|
||||
|
||||
Follow `phases/launch-and-monitor.md`. Re-read it before acting. Surface the W&B URL (if enabled) and produce periodic status reports. At completion, the phase writes `<workspace>/<run-name>/outputs/run-summary.md`.
|
||||
|
||||
## Phase 9 — Post-Train Validate
|
||||
|
||||
After training finishes, follow `phases/post-train-validate.md`. Re-read it before acting. The phase renders the final LoRA against in-distribution, out-of-distribution, and held-out prompts; surfaces the MP4 paths to the user and exits.
|
||||
|
||||
**Important constraint:** the post-train validate phase does not solicit a verdict and does not coach on causes for "soft" failures. Soft training quality has no reliable if/then rule book — the user inspects the renders and decides for themselves whether to ship, iterate, or change course. The orchestrator's job ends after Phase 9; iteration is a new invocation with a new run-name.
|
||||
|
||||
## Monitor-Only Entry
|
||||
|
||||
If invoked against an existing `<workspace>/<run-name>/` that already has a training process running or completed, **skip to Phase 8 in monitor-only mode** instead of restarting anything: report current step, recent loss, ETA, checkpoint list, W&B URL, and the resume command. Distinguish "live process" vs "stopped run with checkpoints" and offer the appropriate next action.
|
||||
|
||||
## Escape Hatch: Unsupported Modes
|
||||
|
||||
If Phase 1 intent doesn't map to any combination of supported `flexible`-strategy conditions:
|
||||
|
||||
1. Stop. Do not edit `packages/ltx-trainer/` on your own.
|
||||
2. Explain what's missing in concrete terms: "you want X. The trainer supports A, B, C via conditions D, E, F. X requires a new condition / strategy."
|
||||
3. Identify the code change needed (typically a new `Condition` subclass in `ltx_trainer/training_strategies/flexible.py` plus schema wiring in `config.py`).
|
||||
4. Ask the user explicitly: "Proceed with the code change, do it yourself, or abort?"
|
||||
5. Only on explicit consent, drop out of this skill's orchestrator and edit code as a normal agent task. After the change lands and is tested, return here.
|
||||
|
||||
## Ask-vs-Assume Cheat Sheet
|
||||
|
||||
| Decision | How |
|
||||
|----------|-----|
|
||||
| Precision, quantization, optimizer, grad checkpointing | Assume from matched VRAM tier. List in plan's "Assumptions". |
|
||||
| Checkpoint/validation interval, seed, W&B project name, output dir | Assume sensible defaults. List in "Assumptions". |
|
||||
| Target resolution / frame count | Ask if not given; propose mode-appropriate default. |
|
||||
| Step count | Ask if dataset size doesn't pin it. |
|
||||
| LoRA trigger word / concept name | Ask for the word only (style/concept LoRAs). Never ask *how* it's injected — always `--lora-trigger`. |
|
||||
| Captioner backend | Ask only if multiple backends are viable. |
|
||||
| Anything answerable by `ls`, `nvidia-smi`, or the W&B credential check | Never ask. Probe. |
|
||||
|
||||
## The Two `load_text_encoder_in_8bit` Flags
|
||||
|
||||
Same name, different layers — do not conflate:
|
||||
|
||||
| Flag | Layer | Effect |
|
||||
|------|-------|--------|
|
||||
| `process_dataset.py --load-text-encoder-in-8bit` | Preprocessing CLI | Memory during caption-embedding precompute (Phase 7). One-time per dataset. |
|
||||
| `acceleration.load_text_encoder_in_8bit` (trainer YAML) | Trainer config | Memory during validation-sample prompt-embedding caching at training start (Phase 8). One-time per run. |
|
||||
|
||||
Both are one-time costs and neither affects per-step training speed. Default per the trainer's shipped configs: **ON** on the 32GB tier (matches `t2v_lora_low_vram.yaml`), **OFF** on the 80GB+ tier (matches `t2v_lora.yaml`). No measured guidance for the 40–60GB tier beyond starting from the 32GB tier and autotuning.
|
||||
|
||||
## Workspace Layout
|
||||
|
||||
```
|
||||
<workspace>/<run-name>/
|
||||
plan.md # the approved plan
|
||||
config.yaml # generated training config (NOT in packages/ltx-trainer/configs/)
|
||||
autotune.log # per-trial sweep results
|
||||
dataset/
|
||||
dataset.json # captions + media paths (training split)
|
||||
holdout.jsonl # held-out split (if reserved)
|
||||
videos/ # source media copies (NO derived files here)
|
||||
.precomputed/ # latents/ audio_latents/ conditions/ (+ references/masks per mode)
|
||||
outputs/
|
||||
checkpoints/ # training checkpoints + states
|
||||
samples/ # in-training validation samples (step_*)
|
||||
eval/ # Phase 9: in-distribution/ out-of-distribution/ held-out/ + prompts.json
|
||||
run-summary.md # written at completion
|
||||
logs/ # all run logs
|
||||
overfit/ # Phase 6 scratch (one-sample preprocess + sanity/autotune runs)
|
||||
```
|
||||
|
||||
`<run-name>` default: `<mode>-<dataset-name>-<YYYYMMDD-HHMM>`. Surface in the plan; user may rename.
|
||||
|
||||
### Workspace hygiene (keep it clean)
|
||||
|
||||
- **Don't create undocumented directories** (e.g. an ad-hoc `scratch/`). Intermediates belong under `overfit/` (sanity/autotune scratch) or a `/tmp` tempdir — not loose in the run dir or the repo root.
|
||||
- **Never write derived files into `dataset/videos/`** (the source media dir). Latents go under `.precomputed/`; one-sample/eval metadata goes under `overfit/`, not `dataset/`.
|
||||
- **One canonical manifest per artifact** — don't leave duplicate `*_prompts.json` / metadata copies.
|
||||
- **Clean up phase byproducts:** the Phase 9 eval must delete its validate-only trainer cruft (see `phases/post-train-validate.md`); `overfit/` is scratch and may be removed after a successful run. The final tree should look like the layout above — no stray `.pt`/`.wav` files, no `eval/checkpoints/`, no duplicate manifests.
|
||||
|
||||
## References
|
||||
|
||||
- `references/mode-selector.md` — user intent → conditioning mode mapping, LoRA rank guidance (read in Phase 1).
|
||||
- `references/onboarding.md` — first-run prerequisite checklist, model download paths, captioner graceful degradation (read in Phase 2).
|
||||
- `references/hardware-profiles.md` — GPU VRAM → tier + config defaults (read in Phase 2).
|
||||
- `references/config-patching.md` — safe YAML edits + schema constraints (read whenever editing `config.yaml`).
|
||||
- `references/troubleshooting.md` — OOM, NaN, validation failures, resume (read on any failure).
|
||||
- `references/plan-template.md` — exact plan.md format (read in Phase 4).
|
||||
|
||||
## Phase Procedures
|
||||
|
||||
These are procedure documents the orchestrator reads when entering each phase. They are not standalone skills — they're never invoked by Claude's skill-discovery system. The orchestrator opens them via the `Read` tool and follows the instructions inline.
|
||||
|
||||
- `phases/prepare-dataset.md` — Phase 5: scenes, captioner iteration, IC-LoRA references, metadata, holdout split.
|
||||
- `phases/preprocess-dataset.md` — Phases 6 (one-sample) & 7 (full): `process_dataset.py` orchestration, existing-data verification.
|
||||
- `phases/launch-and-monitor.md` — Phase 8: launch command, accelerate, W&B, status reports, run-summary writing.
|
||||
- `phases/post-train-validate.md` — Phase 9: render the final LoRA against three prompt categories; surface paths only.
|
||||
@@ -0,0 +1,235 @@
|
||||
# Phase 8 — Launch & Monitor
|
||||
|
||||
Procedure document for the `train-model` orchestrator (Phase 8 + monitor-only re-entry). Read this file in full before acting on launch/monitor.
|
||||
|
||||
Goal: start the training job, surface the W&B URL, produce periodic status reports, write `run-summary.md` at completion.
|
||||
|
||||
The orchestrator's hard invariants apply (see `../SKILL.md`).
|
||||
|
||||
## Launch — Single GPU
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
uv run python scripts/train.py "<workspace>/<run-name>/config.yaml"
|
||||
```
|
||||
|
||||
## Launch — Multi-GPU
|
||||
|
||||
Use Accelerate. For LoRA, DDP (default) is fine. For full FT, use FSDP.
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
|
||||
# DDP (LoRA, multi-GPU)
|
||||
uv run accelerate launch scripts/train.py "<workspace>/<run-name>/config.yaml"
|
||||
|
||||
# FSDP (full FT, multi-GPU)
|
||||
uv run accelerate launch \
|
||||
--config_file configs/accelerate/fsdp.yaml \
|
||||
scripts/train.py "<workspace>/<run-name>/config.yaml"
|
||||
```
|
||||
|
||||
**Pass `--disable-progress-bars` to `train.py` whenever stdout is redirected to a log file** (every background run) or running multi-GPU. The Rich progress bar rewrites a single line with carriage returns and does **not** flush parseable newlines to a redirected log, so without this flag the log shows no step/loss lines and you're forced to poll `nvidia-smi`. With it, step/loss lines are written normally and the log is greppable.
|
||||
|
||||
## Pre-Launch Checks (every launch)
|
||||
|
||||
1. Run the self-check from `references/config-patching.md` (paths exist, frame/resolution constraints, generated modalities have matching latents dirs, references/masks dirs present for conditional modes).
|
||||
2. Confirm `nvidia-smi` shows expected GPUs available (not occupied by another process).
|
||||
3. If `wandb.enabled: true`, confirm credentials still resolve: `uv run python -c "import wandb; print(bool(wandb.Api().api_key))"` → `True`. (Don't use `wandb status`.) If `False`, surface to the user before launching — they may want to `wandb login` or run without tracking.
|
||||
|
||||
## Run In Background, Monitor Foreground
|
||||
|
||||
Long training runs should not block the agent's response loop, and they must survive past the launching turn.
|
||||
|
||||
**Prefer the agent's managed/native background-shell mechanism** (the harness facility for long-running background commands — output streaming + PID/exit tracking, survives across turns). It's the reliable way to launch training: it stays alive, streams to a log the agent can poll, and reports completion. Launch the training command through that mechanism, writing to `<workspace>/<run-name>/logs/train.log` with `--disable-progress-bars`, using **absolute paths** for the config and log (`uv run --directory packages/ltx-trainer` changes the cwd, so a relative config path won't resolve).
|
||||
|
||||
**`nohup ... &` is a last-resort fallback only.** Detached jobs can be harder to track and may not survive environment/session cleanup, so only use it if no managed background mechanism is available, and verify the PID is still alive afterward:
|
||||
|
||||
```bash
|
||||
# Fallback ONLY — prefer the managed background shell above.
|
||||
mkdir -p "<workspace>/<run-name>/logs"
|
||||
nohup uv run --directory packages/ltx-trainer python scripts/train.py \
|
||||
"<ABSOLUTE path>/<run-name>/config.yaml" --disable-progress-bars \
|
||||
> "<ABSOLUTE path>/<run-name>/logs/train.log" 2>&1 &
|
||||
echo $! > "<workspace>/<run-name>/logs/train.pid"
|
||||
```
|
||||
|
||||
## Status Report
|
||||
|
||||
Produce on user request (or every <interval> automatically). Pull from:
|
||||
|
||||
- **W&B run URL:** First lines of `train.log` after init, or `wandb.run.url` from a `wandb` Python snippet. Surface as a clickable URL.
|
||||
- **Latest step:** `tail -n 200 "<workspace>/<run-name>/logs/train.log" | grep -oE "step [0-9]+" | tail -1`.
|
||||
- **Recent loss:** `tail -n 200 "<workspace>/<run-name>/logs/train.log" | grep -oE "loss[: ]+[0-9.]+" | tail -5`.
|
||||
- **Checkpoints saved:** `ls -1t "<workspace>/<run-name>/outputs/checkpoints/" 2>/dev/null`.
|
||||
- **Validation samples:** `ls -1t "<workspace>/<run-name>/outputs/samples/" 2>/dev/null`.
|
||||
- **GPU utilization snapshot:** `nvidia-smi --query-gpu=name,utilization.gpu,memory.used,memory.total --format=csv,noheader`.
|
||||
- **ETA:** only report one **grounded in real numbers the trainer has actually emitted** — do not invent or "educated-guess" an ETA before the trainer has produced per-step timings. Compute as `(total_steps - current_step) * recent_avg_step_time`, where `recent_avg_step_time` is measured from **steady-state training steps** (the trainer's reported per-step time or log timestamps), **excluding** one-time setup that doesn't repeat per step: model loading, the step-0 validation pass, and periodic validation passes. The trainer's own early ETA projection is skewed high by the slow step-0 validation and settles after a few steps — wait for it to settle rather than quoting the inflated early figure. Until real step timings exist, say "measuring step time…" rather than guessing a duration.
|
||||
|
||||
Format the report tightly — one block, no fluff:
|
||||
|
||||
```
|
||||
Step 1240 / 2000 (62%) — loss ~0.072 (last 5: 0.071, 0.073, 0.069, 0.075, 0.072)
|
||||
ETA: ~1h 24m | GPU: 91% util, 39.2 / 48.0 GB
|
||||
Latest checkpoint: lora_weights_step_01000.safetensors
|
||||
Latest validation: samples/step_01200_*.mp4
|
||||
W&B: https://wandb.ai/<entity>/<project>/runs/<id>
|
||||
```
|
||||
|
||||
## Monitor-Only Mode
|
||||
|
||||
Invoked when the orchestrator detects an existing `<workspace>/<run-name>/` with checkpoints or a running process.
|
||||
|
||||
1. Check if the training process is live: `[ -f .../logs/train.pid ] && kill -0 $(cat .../logs/train.pid) 2>/dev/null && echo LIVE || echo STOPPED`.
|
||||
2. Produce the same status report as above.
|
||||
3. If STOPPED:
|
||||
- Compute step from last checkpoint.
|
||||
- Find the latest checkpoint pair (`lora_weights_step_*.safetensors` or `model_weights_step_*.safetensors`, plus a matching `training_state_step_*.pt` when resume state is enabled) under `<workspace>/<run-name>/outputs/checkpoints/`.
|
||||
- Patch `model.load_checkpoint` in `config.yaml` to point at that checkpoint file (this is the only way the trainer knows to resume — there's no auto-detection from `output_dir`).
|
||||
- Surface the resume command: `uv run python scripts/train.py "<workspace>/<run-name>/config.yaml"`.
|
||||
- Ask the user to confirm both the patch and the launch before applying.
|
||||
|
||||
## Resume
|
||||
|
||||
The trainer **does not auto-resume from `output_dir`**. To resume an interrupted run:
|
||||
|
||||
1. Set `model.load_checkpoint` in `config.yaml` to the latest checkpoint file (e.g. `<workspace>/<run-name>/outputs/checkpoints/lora_weights_step_02000.safetensors`).
|
||||
2. Launch normally. The trainer loads those weights, then looks for a matching `training_state_step_*.pt` **next to the loaded checkpoint** and restores optimizer/scheduler/step state from it. If the state file is missing, weights load but training starts from step 0.
|
||||
3. To load the weights but skip the state restore (e.g. for branching off into a new run from a known-good checkpoint), set `checkpoints.no_resume: true`.
|
||||
|
||||
Always ask the user before patching `model.load_checkpoint` or setting `no_resume` — checkpoints are precious.
|
||||
|
||||
## Failure During Training
|
||||
|
||||
If the training process exits non-zero:
|
||||
|
||||
1. Tail the log and identify the error type.
|
||||
2. Cross-reference `references/troubleshooting.md`.
|
||||
3. Propose a config fix (with the exact diff to `config.yaml`).
|
||||
4. Ask the user before applying. Then resume with the patched config.
|
||||
|
||||
Never silently restart a failed training run without acknowledging the failure to the user.
|
||||
|
||||
## After Training Completes
|
||||
|
||||
1. Write a **run summary** to `<workspace>/<run-name>/outputs/run-summary.md` (see below) so the user can find their bearings months later without rereading the trainer docs.
|
||||
2. Show final checkpoint path and step count.
|
||||
3. Show W&B URL if enabled.
|
||||
4. Return control to the orchestrator (Phase 9 — post-train validate runs next).
|
||||
|
||||
### Writing `run-summary.md`
|
||||
|
||||
The summary is the **landing page** for this run. Anyone (including the user months from now) should be able to read it and understand: what was trained, on what data, with what config, where everything lives, how to use the result, and how to continue. The trainer doesn't produce this — the skill does.
|
||||
|
||||
Template (fill from `plan.md`, `config.yaml`, `autotune.log`, dataset metadata, training log):
|
||||
|
||||
```markdown
|
||||
# <run-name>
|
||||
|
||||
**Trained:** <YYYY-MM-DD HH:MM> on <GPU(s)>
|
||||
**Final checkpoint:** `outputs/checkpoints/<filename>.safetensors` (step <N>)
|
||||
|
||||
## What this LoRA does
|
||||
|
||||
<One paragraph from the plan's Goal section — restating the user's intent.>
|
||||
|
||||
## Trigger word
|
||||
|
||||
`<trigger>` — include in prompts at inference time. (Omit this section if no trigger word.)
|
||||
|
||||
## Mode
|
||||
|
||||
<Mode name> (<lora|full>). Conditioning: <list of conditions, or "none">.
|
||||
|
||||
## Dataset
|
||||
|
||||
- Source: `<absolute path>`
|
||||
- Captioning: <`qwen_omni` (Qwen3-Omni-30B via vLLM) | `gemini_flash` | "user-supplied">
|
||||
- Captioner instruction used: <verbatim string, or "default">
|
||||
- Samples: <N training> + <K held-out> at <W>x<H>x<F>
|
||||
- Preprocessed to: `dataset/.precomputed/`
|
||||
|
||||
## Training config
|
||||
|
||||
Final values after autotune (deltas from baseline noted in `autotune.log`):
|
||||
|
||||
| Field | Value |
|
||||
|-------|-------|
|
||||
| Optimizer | <...> |
|
||||
| Mixed precision | <bf16/fp16> |
|
||||
| Quantization | <...> |
|
||||
| Gradient checkpointing | <on/off> |
|
||||
| Batch size × grad accum | <B> × <A> (effective <BxA>) |
|
||||
| LoRA rank / alpha | <R> / <A> (or "full FT") |
|
||||
| LoRA target modules | <list> (or "n/a") |
|
||||
| Steps | <N> |
|
||||
| Learning rate | <value> |
|
||||
| Step time (final) | ~<T>s |
|
||||
| Peak VRAM | ~<V> GB |
|
||||
|
||||
Full config: `<workspace>/<run-name>/config.yaml`
|
||||
|
||||
## Outputs
|
||||
|
||||
- Checkpoints: `outputs/checkpoints/`
|
||||
- Validation samples (during training): `outputs/samples/`
|
||||
- Post-train eval renders (if Phase 9 ran): `outputs/eval/`
|
||||
- W&B run: <url, or "(W&B not enabled)">
|
||||
|
||||
## How to use this checkpoint
|
||||
|
||||
For inference, point `packages/ltx-pipelines/` at the final checkpoint. Example invocation:
|
||||
|
||||
\`\`\`bash
|
||||
# (Minimal sketch — adapt to the pipeline you're using.)
|
||||
# load base LTX-2 model + apply this LoRA from outputs/checkpoints/<filename>.safetensors
|
||||
\`\`\`
|
||||
|
||||
## How to continue training
|
||||
|
||||
To resume from the final checkpoint (e.g. more steps, different LR), edit `config.yaml`:
|
||||
|
||||
\`\`\`yaml
|
||||
model:
|
||||
load_checkpoint: "<absolute path to outputs/checkpoints/<filename>.safetensors>"
|
||||
optimization:
|
||||
steps: <new total> # trainer resumes optimizer/scheduler/step from the training_state_step_*.pt sitting next to the checkpoint above
|
||||
\`\`\`
|
||||
|
||||
Then re-launch with the same command in the "Launched with" section below.
|
||||
|
||||
## How this was launched
|
||||
|
||||
\`\`\`
|
||||
<exact command used, with the workspace's absolute config path>
|
||||
\`\`\`
|
||||
|
||||
## Reproducibility
|
||||
|
||||
- Seed: <value>
|
||||
- Workspace: `<absolute path>`
|
||||
- Repo commit at launch: `<git rev-parse HEAD output>`
|
||||
- LTX-2 model: `<model.model_path from config>`
|
||||
- Text encoder: `<model.text_encoder_path from config>`
|
||||
```
|
||||
|
||||
Write the file using the `Write` tool. Don't embed it in a heredoc — the markdown nested in this skill is illustrative; fill the template with real values from the run's artifacts.
|
||||
|
||||
### Next steps (surface to user)
|
||||
|
||||
After writing the summary, point the user at:
|
||||
- The summary file path.
|
||||
- The final checkpoint path.
|
||||
- The W&B URL if enabled.
|
||||
- The upcoming Phase 9 (post-train validate) — the orchestrator handles the transition.
|
||||
|
||||
Suggest, but don't run:
|
||||
- Test inference with `packages/ltx-pipelines/`.
|
||||
- Push to HF Hub via the trainer's `hub.push_to_hub` config (a separate, lightweight re-launch).
|
||||
- Continue training from the final checkpoint (see summary's "How to continue training" section).
|
||||
|
||||
## Do Not
|
||||
|
||||
- Do not modify `output_dir` contents after a run completes — checkpoints belong to the user now.
|
||||
- Do not start a second training run into the same `output_dir` without explicit user approval. Resume requires patching `model.load_checkpoint` (the trainer does not auto-detect prior checkpoints); a true fresh-start from the same dir additionally needs `checkpoints.no_resume: true`.
|
||||
- Do not auto-restart a failed run without diagnosis and user approval.
|
||||
@@ -0,0 +1,171 @@
|
||||
# Phase 9 — Post-Train Validate
|
||||
|
||||
Procedure document for the `train-model` orchestrator (Phase 9). Read this file in full before acting on post-train validation.
|
||||
|
||||
Goal: render the final checkpoint against three prompt categories so the user can inspect the result and form their own judgement. Save outputs in an organized layout. **Do not** prompt for pass/fail verdicts, do not infer causes for failures, do not suggest fixes — soft training failures don't have a clean if/then rule book, and pretending otherwise wastes the user's time.
|
||||
|
||||
The orchestrator's hard invariants apply (see `../SKILL.md`).
|
||||
|
||||
## What this phase does
|
||||
|
||||
1. Collect prompts for three categories.
|
||||
2. Render the final LoRA against all collected prompts.
|
||||
3. Save outputs under `<workspace>/<run-name>/outputs/eval/<category>/`.
|
||||
4. Print the paths and exit.
|
||||
|
||||
## Categories
|
||||
|
||||
### 1 — In-distribution
|
||||
|
||||
A few captions from the training set itself. Tests whether the model learned what it was shown.
|
||||
|
||||
- Default: 3 random captions from the dataset metadata (seed 42 for reproducibility).
|
||||
- Source: `<workspace>/<run-name>/dataset/dataset.json` (the captions used for training, after the held-out split).
|
||||
|
||||
### 2 — Out-of-distribution
|
||||
|
||||
Prompts the model has never seen, but in the same domain. Tests whether the model generalizes the concept beyond memorized phrasings.
|
||||
|
||||
Ask the user once:
|
||||
|
||||
> "For out-of-distribution validation, paste 2–3 prompts you'd realistically want to generate at inference time. (If you don't have any specific ones in mind, reply 'default' — I'll use a few generic prompts that include the trigger word.)"
|
||||
|
||||
- On `default`: synthesize 3 short prompts using the LoRA's trigger word and a generic scene context (e.g., "<trigger> walking in a forest at dawn"). Note these are generic — they're better than nothing, but real user-style prompts make a stronger test.
|
||||
- Otherwise: use the user's prompts verbatim.
|
||||
- **For a run with a generated audio modality, the synthesized prompts must describe the audio** — matching how the training captions describe it (inspect a few from `dataset.json` first). A prompt with no audio direction leaves the audio branch unguided and the generated audio comes out poor. E.g. for the talking-head case include spoken-voice/room-tone direction; for music/ambience/foley describe the sound character. (Categories 1 and 3 reuse the real captions verbatim, so they already carry audio description — this only applies to the synthesized Category 2 prompts.) If the user pasted their own prompts and it's an audio run, and they omitted audio direction, note that the audio may be weak without it.
|
||||
|
||||
### 3 — Held-out
|
||||
|
||||
Captions from samples that were never seen during training. Tests true generalization, not memorization.
|
||||
|
||||
- Source: `<workspace>/<run-name>/dataset/holdout.jsonl` (written by `prepare-dataset` Step 5).
|
||||
- Use all entries if there are ≤5; otherwise sample 5 with seed 42.
|
||||
- **If holdout doesn't exist** (dataset was too small, or user-skipped during prepare): print a clear note in the output summary — *"Held-out evaluation skipped: no holdout set was reserved for this run. The post-train eval only covers Categories 1 and 2."* Don't synthesize substitutes.
|
||||
|
||||
## Rendering mechanism
|
||||
|
||||
Use the trainer's existing validation infrastructure rather than wiring up `ltx-pipelines` from scratch. Create a temporary "validate-only" config and run the trainer with it.
|
||||
|
||||
### Step 1 — Build eval config
|
||||
|
||||
Copy `<workspace>/<run-name>/config.yaml` to `<workspace>/<run-name>/eval-config.yaml`. Patch:
|
||||
|
||||
```yaml
|
||||
model:
|
||||
load_checkpoint: "<absolute path to outputs/checkpoints/<final-lora>.safetensors>"
|
||||
|
||||
optimization:
|
||||
steps: 1 # we don't want to train; we want validation to fire
|
||||
# Keep batch_size/grad_accum at the run's autotuned values to match its VRAM footprint.
|
||||
|
||||
validation:
|
||||
skip_initial_validation: false
|
||||
interval: 1 # run validation at step 0 (and at the only training step)
|
||||
samples:
|
||||
# Inject all collected prompts here, tagged by category in the prompt itself
|
||||
# so the output filenames make the category obvious.
|
||||
- prompt: "[CAT1-IND] <caption from dataset.json>"
|
||||
# ... repeat for each prompt in all three categories
|
||||
# Keep video_dims, frame_rate, guidance/STG settings as the trained config.
|
||||
# Keep generate_audio consistent with the trained modality config.
|
||||
|
||||
checkpoints:
|
||||
interval: null # do not save more checkpoints
|
||||
no_resume: true # load the LoRA's weights but do not restore optimizer/scheduler/step state
|
||||
|
||||
output_dir: "<workspace>/<run-name>/outputs/eval"
|
||||
```
|
||||
|
||||
The `[CAT1-IND]`, `[CAT2-OOD]`, `[CAT3-HELDOUT]` tags in the prompt strings make the output MP4 filenames self-describing in the trainer's validation sample directory.
|
||||
|
||||
**Attach the mode's conditions to each sample.** A bare `prompt` only validates a pure text-to-X mode (T2V, T2A). For any conditioned mode, the trained model expects the same conditioning at validation time — a prompt with no conditions tests a different task than what was trained, and conditioned modes may fail outright. Add the `conditions` list that matches the run's mode (the trained `config.yaml` `training_strategy` and the example config for the mode are the reference):
|
||||
|
||||
| Mode | Add to each sample |
|
||||
|------|--------------------|
|
||||
| I2V | `conditions: [{type: first_frame, image_or_video: <frame/clip path>}]` |
|
||||
| Video extension / suffix | `conditions: [{type: prefix|suffix, ...}]` |
|
||||
| V2V / AV2AV IC-LoRA | `conditions: [{type: reference, ...}]` (point at a held-out reference) |
|
||||
| V2A (foley) | `conditions: [{type: video_to_audio, ...}]` |
|
||||
| A2V | `conditions: [{type: audio_to_video, ...}]` |
|
||||
| Inpainting (video/audio) | `conditions: [{type: mask, ...}]` |
|
||||
| Outpainting | `conditions: [{type: spatial_crop, ...}]` |
|
||||
| A2A IC-LoRA | `conditions: [{type: reference, ...}]` (held-out reference audio) |
|
||||
| T2V, T2A | none — a bare `prompt` is correct |
|
||||
|
||||
Mirror the condition shapes used in the mode's example config under `packages/ltx-trainer/configs/`. For held-out (Category 3) and OOD (Category 2) samples on conditioned modes, draw the conditioning media from the held-out set so the eval stays out-of-distribution.
|
||||
|
||||
### Step 2 — Run the trainer in validate-only mode
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
uv run python scripts/train.py "<workspace>/<run-name>/eval-config.yaml"
|
||||
```
|
||||
|
||||
The trainer will load the LoRA, run initial validation against all the prompts, do one trivial training step (which we discard), and exit. Validation samples land in `<workspace>/<run-name>/outputs/eval/samples/`.
|
||||
|
||||
### Step 3 — Organize outputs and clean up trainer cruft
|
||||
|
||||
The validate-only run is a trainer run, so it inevitably writes throwaway artifacts: an indexed `samples/` dir, a forced final checkpoint (the trainer **always** saves one at the end, regardless of `checkpoints.interval`), and a `training_config.yaml`. Don't leave these around or duplicate the renders.
|
||||
|
||||
1. **Move** (don't copy) each generated MP4 from the trainer's indexed `samples/` dir into the category layout, naming by category + a slug of the prompt. Use the index→category mapping you built when constructing `validation.samples`.
|
||||
2. **Delete the trainer cruft** from the eval dir once the renders are moved: the indexed `samples/` dir, the forced `checkpoints/` dir, and `training_config.yaml`. (These are byproducts of the validate-only hack — there's no config flag to suppress the final-checkpoint save, so clean it up here.)
|
||||
3. Write **one** manifest, `outputs/eval/prompts.json` (filename → full prompt + category). Don't leave a second copy elsewhere.
|
||||
4. Remove the temporary `eval-config.yaml` (or keep it under the run's scratch, not in `outputs/`).
|
||||
|
||||
Final eval layout — exactly this, nothing else:
|
||||
|
||||
```
|
||||
<workspace>/<run-name>/outputs/eval/
|
||||
in-distribution/ <NN>_<prompt-slug>.mp4 ...
|
||||
out-of-distribution/ <NN>_<prompt-slug>.mp4 ...
|
||||
held-out/ <NN>_<prompt-slug>.mp4 ... # only if a holdout set existed
|
||||
prompts.json # filename -> full prompt + category (single manifest)
|
||||
```
|
||||
|
||||
No `eval/samples/`, no `eval/checkpoints/`, no `eval/training_config.yaml`, no duplicate manifest.
|
||||
|
||||
### Step 4 — Surface paths
|
||||
|
||||
Print a tight block, no judgement, no follow-up question:
|
||||
|
||||
```
|
||||
Post-train evaluation complete.
|
||||
|
||||
In-distribution renders (<K> samples):
|
||||
<workspace>/<run-name>/outputs/eval/in-distribution/
|
||||
|
||||
Out-of-distribution renders (<M> samples):
|
||||
<workspace>/<run-name>/outputs/eval/out-of-distribution/
|
||||
|
||||
Held-out renders (<N> samples):
|
||||
<workspace>/<run-name>/outputs/eval/held-out/ # or: "(skipped — no holdout set)"
|
||||
|
||||
Open the MP4s and decide for yourself whether the model is good. Soft
|
||||
training quality is judged by watching the videos, not by a checklist —
|
||||
there's no substitute for your own eyes here.
|
||||
```
|
||||
|
||||
Return control to the orchestrator. The orchestrator's run is now complete.
|
||||
|
||||
## What this phase does NOT do
|
||||
|
||||
- Does not ask "is this good?" / "pass / partial / fail?".
|
||||
- Does not infer failure causes.
|
||||
- Does not suggest fixes, follow-up runs, hyperparameter changes, dataset changes.
|
||||
- Does not write any verdict to `run-summary.md` or elsewhere.
|
||||
- Does not delete or modify training checkpoints.
|
||||
- Does not push to any remote / cloud / registry.
|
||||
|
||||
The user looks at the videos and makes their own call. If they want to iterate, they re-invoke the orchestrator with a new run-name.
|
||||
|
||||
## Failure modes
|
||||
|
||||
- **Final checkpoint missing.** Surface and stop. Don't render against an intermediate checkpoint without explicit user consent.
|
||||
- **`load_checkpoint` OOM at inference time.** Lower `validation.video_dims` in the eval config (smaller renders are still useful for a sanity look). Retry once. If still OOM, surface the failure and let the user run inference manually via `packages/ltx-pipelines/`.
|
||||
- **All renders look broken/black.** May be an inference-pipeline-side issue rather than a training failure. Mention in the output block: *"If renders look broken across all categories, try `packages/ltx-pipelines/` directly to rule out a pipeline issue."* Then exit. Do not investigate further.
|
||||
|
||||
## Do not
|
||||
|
||||
- Do not skip Category 1 or 2. They're cheap and informative.
|
||||
- Do not invent a held-out set if `holdout.jsonl` is missing — the prepare-dataset step decides that.
|
||||
- Do not coach the user on what "good" means for their use case.
|
||||
@@ -0,0 +1,245 @@
|
||||
# Phase 5 — Prepare Dataset
|
||||
|
||||
Procedure document for the `train-model` orchestrator (Phase 5). Read this file in full before acting on the prepare-dataset phase.
|
||||
|
||||
Goal: produce a captioned, complete dataset metadata file at `<workspace>/<run-name>/dataset/dataset.json` consumable by `process_dataset.py`. Idempotent — re-runs skip work already done.
|
||||
|
||||
The orchestrator's hard invariants apply (see `../SKILL.md`), especially: **no file mutation outside the workspace without explicit user approval.**
|
||||
|
||||
## Inputs
|
||||
|
||||
The orchestrator passes:
|
||||
- Source path (directory of videos / single video / pre-existing metadata file).
|
||||
- Target mode (T2V, I2V, V2V IC-LoRA, V2A, etc.) — determines which columns are required.
|
||||
- Captioner backend choice (Qwen3-Omni local vLLM server / Gemini Flash cloud / skip).
|
||||
- Workspace path `<workspace>/<run-name>/`.
|
||||
|
||||
## Required Columns by Mode
|
||||
|
||||
`process_dataset.py` detects columns by convention and resolves each to a role. The media column may be `video` **or** `audio`. When a dataset has a `video` column with an audio track and no separate `audio` column, audio is **auto-extracted** from the video (unless `--skip-audio`), so an explicit `audio` column is only needed when the audio lives in separate files.
|
||||
|
||||
**Video-generating modes** (need a `video` column):
|
||||
|
||||
| Mode | Required | Optional |
|
||||
|------|----------|----------|
|
||||
| T2V, I2V, video extension/suffix | `video`, `caption` | `audio` (else auto-extracted) |
|
||||
| Video outpainting | `video`, `caption` | |
|
||||
| Video inpainting | `video`, `caption`, `video_mask` | |
|
||||
| V2A (foley) | `video`, `caption` | `audio` (target; else auto-extracted) |
|
||||
| A2V | `video`, `caption` | `audio` (else auto-extracted from video) |
|
||||
| V2V IC-LoRA | `video`, `caption`, `reference_video` | |
|
||||
| AV2AV IC-LoRA | `video`, `caption`, `reference_video`, `reference_audio` | `audio` (else auto-extracted) |
|
||||
|
||||
**Audio-only modes** (no `video` column — the media column is `audio`):
|
||||
|
||||
| Mode | Required | Optional |
|
||||
|------|----------|----------|
|
||||
| T2A | `audio`, `caption` | |
|
||||
| Audio extension/suffix | `audio`, `caption` | |
|
||||
| Audio inpainting | `audio`, `caption`, `audio_mask` | |
|
||||
| A2A IC-LoRA | `audio`, `caption`, `reference_audio` | |
|
||||
|
||||
Aliases: `media_path` for `video`, `ref_media_path` for `reference_video`.
|
||||
|
||||
## Workflow
|
||||
|
||||
### Step 1 — Classify source
|
||||
|
||||
```bash
|
||||
# If source is a file, identify type:
|
||||
file "<source>"
|
||||
# If source is a directory, count media:
|
||||
find "<source>" -maxdepth 1 -type f \( -name "*.mp4" -o -name "*.mov" -o -name "*.webm" \) | wc -l
|
||||
```
|
||||
|
||||
Cases:
|
||||
- **Pre-existing metadata file** (CSV/JSON/JSONL) → copy to `<workspace>/<run-name>/dataset/dataset.json`, audit columns. Skip to Step 4.
|
||||
- **Directory of short scenes** → skip Step 2, go to Step 3.
|
||||
- **Directory containing long videos** → run Step 2.
|
||||
- **Single long video** → run Step 2.
|
||||
|
||||
**Stage the media under `dataset/` before captioning — don't discover this by failing.** Both `caption_videos.py` and `process_dataset.py` reference media by paths **relative to the metadata file's own directory**, so the media must live under `<workspace>/<run-name>/dataset/`. Stage it up front into `dataset/videos/` (and write metadata paths relative to `dataset/`, e.g. `videos/1.mp4`):
|
||||
- Prefer **symlinks** (instant, no disk cost): `ln -s <abs-source>/<clip> <workspace>/<run-name>/dataset/videos/<clip>`. Symlinks pointing at the original source location work correctly.
|
||||
- Use a **copy** instead if the workspace and source are on different filesystems or the source may move/change.
|
||||
- Never caption or preprocess directly against an out-of-tree source path (e.g. `/path/to/source-videos`) — it will fail the relative-path resolution. The original source is left untouched either way.
|
||||
(Scene-split output in Step 2 already lands under `dataset/scenes/`, which satisfies this.)
|
||||
|
||||
### Step 2 — Scene splitting (only for long videos)
|
||||
|
||||
`split_scenes.py` takes **one video file** at a time (`video_path` and `output_dir` are both positional arguments). When the source is a directory of long videos, iterate over each file. To drop scenes shorter than 2 seconds, use `--filter-shorter-than 2s` (the `--min-scene-length` option is an integer **frame** count, not seconds — don't pass a float).
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
# Single file:
|
||||
uv run python scripts/split_scenes.py \
|
||||
"<video-file>" \
|
||||
"<workspace>/<run-name>/dataset/scenes" \
|
||||
--filter-shorter-than 2s
|
||||
|
||||
# Directory of long videos — iterate:
|
||||
for f in "<source>"/*.mp4 "<source>"/*.mov "<source>"/*.webm; do
|
||||
[ -e "$f" ] || continue
|
||||
uv run python scripts/split_scenes.py "$f" \
|
||||
"<workspace>/<run-name>/dataset/scenes" \
|
||||
--filter-shorter-than 2s
|
||||
done
|
||||
```
|
||||
|
||||
Result: scenes saved to `<workspace>/<run-name>/dataset/scenes/`. Pass that directory to Step 3.
|
||||
|
||||
### Step 3 — Captioning
|
||||
|
||||
Skip entirely if a metadata file with all required `caption` entries already exists.
|
||||
|
||||
**Use the captioner's default instruction.** `caption_videos.py` ships a well-tuned default caption prompt — use it as-is (do **not** pass `--instruction`). Captioning runs in two phases: a small **spot-check pass** so the user can confirm the captions look sane, then a **full pass** on the rest.
|
||||
|
||||
A custom `--instruction` is the exception, not the norm. Only use one when:
|
||||
- the **nature of the dataset genuinely demands it** (e.g. a narrow domain the default prompt won't describe well), or
|
||||
- the **user, after seeing the spot-check captions, explicitly asks** for a change (e.g. "too much background detail").
|
||||
|
||||
Do not invent a custom instruction pre-emptively, and in particular **do not bake a subject name / trigger word into the captions via `--instruction`** — the trigger word is handled separately at preprocessing (see "Trigger word" below).
|
||||
|
||||
#### Choosing a backend
|
||||
|
||||
Two backends, with very different hardware needs:
|
||||
|
||||
- **`qwen_omni` (local, default):** Qwen3-Omni-30B-A3B-Thinking served by a local vLLM HTTP server (`serve_captioner.py`). ~65 GiB model download. Default **FP8** quantization uses ~31 GiB of weights and **fits on a 40 GiB GPU** (plus KV cache); **bf16** uses ~60 GiB and needs **≥66 GiB free VRAM**.
|
||||
- **`gemini_flash` (cloud):** Google `gemini-3.5-flash`. No local model, runs anywhere, parallelisable with `--num-workers`.
|
||||
|
||||
**Steer modest hardware to Gemini.** If the GPU is below ~40 GiB (i.e. typical consumer cards — 24 GB / 32 GB), it can't host even the FP8 server, so the local captioner isn't an option — recommend `gemini_flash` and tell the user they'll need Gemini auth: either a `GEMINI_API_KEY`/`GOOGLE_API_KEY` (get one at <https://aistudio.google.com/apikey>) or working gcloud/Vertex AI credentials. If they can't or won't set that up and the hardware can't run Qwen3, the only remaining path is bringing their own captions in the dataset metadata (skip captioning entirely). On a 40 GiB+ GPU the local server is viable (FP8); bf16 needs an 80GB-class card.
|
||||
|
||||
#### Qwen server prerequisite (qwen_omni only)
|
||||
|
||||
The local backend talks to a vLLM server that must already be running. Launch it once in a **separate terminal** (it stays loaded across captioning runs):
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
uv run python scripts/serve_captioner.py # FP8 by default, serves on http://127.0.0.1:8001/v1
|
||||
# bf16 (needs >= 66 GiB free VRAM): --quantization bf16
|
||||
# different port/interface: --port 9000 --host 0.0.0.0
|
||||
```
|
||||
|
||||
First launch downloads the model (~65 GiB). `caption_videos.py` reaches it via `--vllm-url` (default `http://127.0.0.1:8001/v1`). Skip this entirely when using `gemini_flash`.
|
||||
|
||||
#### 3a — Spot-check pass (3 samples)
|
||||
|
||||
Caption 3 samples with the **default prompt** (no `--instruction`) to confirm the captioner is producing sane output before committing to the whole set.
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
# qwen_omni (server from the previous step must be running):
|
||||
uv run python scripts/caption_videos.py \
|
||||
"<workspace>/<run-name>/dataset/videos/<one-staged-clip>" \
|
||||
--output "<workspace>/<run-name>/dataset/preview-captions.json" \
|
||||
--captioner-type qwen_omni
|
||||
# Point at 3 staged clips under dataset/videos/ (a small subdir or 3 explicit files) — not the out-of-tree source.
|
||||
# Optional: --vllm-url http://127.0.0.1:9000/v1 (if the server uses a non-default port)
|
||||
```
|
||||
|
||||
Print the 3 captions **in full** to the user, then **STOP and wait** for their explicit verdict:
|
||||
|
||||
> "Here are sample captions from the default prompt. Please review them — reply 'good' to caption the rest, or tell me what to change."
|
||||
|
||||
**This is a hard gate. Do NOT caption the full set until the user explicitly approves the samples.** Do not auto-proceed, do not assume "looks fine," do not batch this with other questions. The user must either approve or give tuning instructions first — the whole point of the spot-check is to let them judge caption quality and content before paying for the full pass.
|
||||
|
||||
If the user requests changes, introduce a custom `--instruction` (or switch backend), re-run the spot-check on the same 3 samples, show the new captions, and **stop for approval again**. Loop until the user approves. If a custom instruction still isn't converging after a few rounds, switch captioner backend or have the user supply a few manual captions as examples — but still don't proceed to the full set without their OK.
|
||||
|
||||
#### 3b — Full pass
|
||||
|
||||
Run on the **staged media dir** (`dataset/videos/` from Step 1, or `dataset/scenes/` from Step 2) with the **default prompt** (or the same `--instruction` only if one was explicitly agreed in 3a). Caption the staged in-tree media — not the original out-of-tree source path.
|
||||
|
||||
**Qwen3-Omni (local — server must be running):**
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
uv run python scripts/caption_videos.py \
|
||||
"<workspace>/<run-name>/dataset/videos" \
|
||||
--output "<workspace>/<run-name>/dataset/dataset.json" \
|
||||
--captioner-type qwen_omni
|
||||
```
|
||||
|
||||
**Gemini Flash (cloud — runs anywhere, parallelisable):**
|
||||
|
||||
```bash
|
||||
# Auth: GEMINI_API_KEY / GOOGLE_API_KEY env var, or gcloud / Vertex AI credentials.
|
||||
cd packages/ltx-trainer
|
||||
uv run python scripts/caption_videos.py \
|
||||
"<workspace>/<run-name>/dataset/videos" \
|
||||
--output "<workspace>/<run-name>/dataset/dataset.json" \
|
||||
--captioner-type gemini_flash \
|
||||
--num-workers 5
|
||||
```
|
||||
|
||||
Output: JSON list of `{caption, media_path}` with paths **relative to the output file location**. The 3 spot-check captions can be merged in to avoid re-captioning them.
|
||||
|
||||
#### Trigger word (handled at preprocessing, not in captions)
|
||||
|
||||
For style/concept LoRAs the trigger word is **not** written into the captions here. It is prepended to every caption at preprocessing: pass `--lora-trigger "<word>"` to `process_dataset.py` in Phase 7, which forwards it to the caption-processing step (`process_captions.py`, where the prepend actually happens). That is the canonical mechanism — keep the captions describing what's actually on screen (via the default prompt), and let the trigger flag bind the concept to the token. Record the chosen trigger word in the plan so Phase 7 passes it through. Do not also bake the word into captions (it would double up).
|
||||
|
||||
**The injection mechanism is a fixed implementation detail — never make it a user-facing question.** The *only* trigger-word thing to ask the user is the **word itself** (or whether they want a trigger word at all). Do **not** ask, mention, or present as an option *how* it gets injected (e.g. "inject into the caption vs via `process_dataset`") — it is always `--lora-trigger`, full stop. Surfacing the method as a choice creates unnecessary confusion.
|
||||
|
||||
### Step 4 — Conditioning inputs (modes that need references or masks)
|
||||
|
||||
Some modes need a per-sample conditioning input beyond the video/audio and caption:
|
||||
|
||||
| Mode | Required extra input | Column |
|
||||
|------|----------------------|--------|
|
||||
| V2V IC-LoRA | reference video | `reference_video` (alias `ref_media_path`) |
|
||||
| AV2AV IC-LoRA | reference video + reference audio | `reference_video`, `reference_audio` |
|
||||
| A2A IC-LoRA | reference audio | `reference_audio` |
|
||||
| Video inpainting | per-frame video mask | `video_mask` |
|
||||
| Audio inpainting | audio mask | `audio_mask` |
|
||||
|
||||
These inputs encode **the user's specific idea** for the LoRA (what the reference represents, which regions the mask covers). There is no universal recipe, so **do not invent or default to a particular method** (e.g. don't assume Canny edges, depth, pose, or some generic box/border mask). The agent must not pick the conditioning semantics for the user.
|
||||
|
||||
Workflow when the chosen mode needs one of these and the dataset doesn't already provide it:
|
||||
|
||||
1. **Check first** — if the user already supplied the column (and the files exist), use it as-is and move on.
|
||||
2. **Otherwise, ask the user to provide it**, explaining concretely what's needed: the column name, that it's one file per sample aligned to each clip, and that the *content/semantics are their call* (what the reference should depict, what the mask should cover). Make clear this reflects their specific use-case — you won't guess it.
|
||||
3. **Help generate only if the user asks.** If they say "can you generate the references/masks by doing X" (X = their described method), then help: write or run a small script for *their* approach, or use a repo tool if it fits. One such tool exists — `scripts/compute_reference.py` generates **Canny edge** reference videos — but only mention/use it if the user specifically wants Canny; never offer it as the default.
|
||||
4. **Hard gate:** do not proceed to preprocessing for a conditioning mode until the required column is present with real files. Surface clearly if it's missing.
|
||||
|
||||
**Column-naming note (if references are generated):** `compute_reference.py` writes a `reference_video` field, which
|
||||
`process_dataset.py` detects automatically. Legacy datasets using `ref_media_path` also work.
|
||||
|
||||
### Step 5 — Holdout split
|
||||
|
||||
Reserve a subset of samples as a **held-out set** never seen during training. This is what Phase 9 (post-train validate) renders against to test true generalization rather than memorization.
|
||||
|
||||
Decision tree:
|
||||
|
||||
1. **User already supplied a held-out set** (separate file or directory they explicitly nominated): do not split. Copy/reference their file to `<workspace>/<run-name>/dataset/holdout.jsonl` and leave `dataset.json` as-is.
|
||||
2. **Small dataset** (judge qualitatively; tens of samples or fewer): holding samples out meaningfully reduces training capacity. Ask the user:
|
||||
> "Dataset has <N> samples. Reserving a holdout meaningfully reduces what's available for training. Options: (a) reserve 1–2 for holdout, (b) skip holdout — post-train eval will only test in-distribution. Your call."
|
||||
3. **Otherwise:** auto-split. Reserve a small fraction (this skill's default: roughly 10% of samples, bounded so the holdout doesn't grow huge — a handful of held-out samples is usually enough). Use seed 42 for the split so it's reproducible. Surface the count and the picked IDs in the plan.
|
||||
|
||||
After splitting, write `<workspace>/<run-name>/dataset/holdout.jsonl` (one JSON object per line with the same columns as `dataset.json`). **Remove the held-out entries from `dataset.json`** so they don't enter preprocessing or training.
|
||||
|
||||
Always print:
|
||||
> "Held out <K> of <N> samples for post-train evaluation. Held-out IDs: <list>."
|
||||
|
||||
If holdout is skipped, surface in the plan: *"Skipping holdout — dataset is too small. Post-train eval will only render in-distribution prompts; true generalization isn't testable for this run."*
|
||||
|
||||
### Step 6 — Audit
|
||||
|
||||
Before returning to the orchestrator, verify the metadata file has all required columns for the chosen mode. Print a one-line summary:
|
||||
|
||||
> "Prepared <N> training samples (+ <K> held out) for <mode>. Columns: <list>. Saved to `<workspace>/<run-name>/dataset/dataset.json` (+ `holdout.jsonl`)."
|
||||
|
||||
## Idempotency
|
||||
|
||||
- If `dataset.json` already exists and all required columns are present: skip captioning. Confirm reuse with the user only if the file was supplied by them outside the workspace (per the orchestrator's file-safety invariant).
|
||||
- If captioning was partial (some entries missing `caption`), re-run captioning only on the missing entries by filtering the metadata file before passing to `caption_videos.py`.
|
||||
|
||||
## Failure Modes
|
||||
|
||||
- Qwen server won't start / OOMs on launch → use the default `--quantization fp8` (not `bf16`), lower `--gpu-memory-utilization`, or reduce `--max-model-len` on `serve_captioner.py`. If the GPU simply can't host a 30B model, switch to `gemini_flash`.
|
||||
- `caption_videos.py` can't connect (qwen_omni) → the vLLM server isn't running or `--vllm-url` is wrong. Start `serve_captioner.py` first and confirm the URL/port match.
|
||||
- Gemini rate-limit → reduce `--num-workers`, retry.
|
||||
- Scene splitter produces 0 scenes → the detector found no cuts, or `--filter-shorter-than` removed everything. Lower/remove `--filter-shorter-than`, or adjust the detector threshold. (`--min-scene-length` is an integer minimum-frames-per-scene for the detector, not a short-scene filter.)
|
||||
- IC-LoRA reference compute fails on some frames → script logs the failures; report counts to the user and ask whether to proceed with the remaining samples or stop.
|
||||
|
||||
## Do Not
|
||||
|
||||
- Do not move or rename the user's source files. The skill reads them in place; the workspace contains only **derived** artifacts.
|
||||
- Do not delete `scenes/` or partial captioning outputs without approval — they may be expensive to regenerate.
|
||||
@@ -0,0 +1,130 @@
|
||||
# Phases 6 (one-sample) & 7 (full) — Preprocess Dataset
|
||||
|
||||
Procedure document for the `train-model` orchestrator. Read this file in full before acting on the preprocess phase.
|
||||
|
||||
Goal: run `process_dataset.py` to produce VAE latents, audio latents, and text embeddings. Two modes:
|
||||
1. **One-sample** (Phase 6 sanity check) — preprocess a single sample to `<workspace>/<run-name>/overfit/.precomputed/`.
|
||||
2. **Full** (Phase 7) — preprocess the whole dataset to `<workspace>/<run-name>/dataset/.precomputed/`.
|
||||
|
||||
The orchestrator's hard invariants apply (see `../SKILL.md`), especially: **never silently overwrite existing user data.**
|
||||
|
||||
## Required Subdirectories by Mode
|
||||
|
||||
Under `.precomputed/`:
|
||||
|
||||
| Subdir | Required for |
|
||||
|--------|--------------|
|
||||
| `latents/` | Video-bearing modes only (T2V, I2V, video extend/inpaint/outpaint, V2V/AV2AV IC-LoRA, A2V, V2A). **Not** produced for audio-only modes. |
|
||||
| `conditions/` | Always (text embeddings) |
|
||||
| `audio_latents/` | Any mode with audio: video modes carrying audio, plus all audio-only modes (T2A, audio extend/suffix/inpaint, A2A IC-LoRA) |
|
||||
| `reference_latents/` | V2V IC-LoRA, AV2AV IC-LoRA |
|
||||
| `reference_audio_latents/` | A2A IC-LoRA, AV2AV IC-LoRA |
|
||||
| `video_masks/` | Video inpainting |
|
||||
| `audio_masks/` | Audio inpainting |
|
||||
|
||||
**Audio-only modes** (T2A, audio extend/suffix, audio inpainting, A2A IC-LoRA) produce `audio_latents/` + `conditions/` (plus `audio_masks/` or `reference_audio_latents/` as applicable) and **no `latents/`**. Do not flag a missing `latents/` as incomplete for these modes.
|
||||
|
||||
## Workflow — Full Preprocess
|
||||
|
||||
### Step 1 — Verify existing `.precomputed/` (if present)
|
||||
|
||||
If `<workspace>/<run-name>/dataset/.precomputed/` already exists:
|
||||
|
||||
1. List subdirectories present. Confirm all required for the chosen mode are present.
|
||||
2. Load one sample per modality and check tensor shapes:
|
||||
```bash
|
||||
uv run python -c "import torch; t = torch.load('<path>'); print(t.shape if hasattr(t, 'shape') else {k: v.shape for k, v in t.items()})"
|
||||
```
|
||||
3. Compare shapes against the target resolution from the plan.
|
||||
|
||||
**On any mismatch or missing subdirectory: STOP. Do not run `process_dataset.py`.** Ask the user via `AskUserQuestion`:
|
||||
- Reuse the existing data at its current resolution (update plan + config accordingly).
|
||||
- Re-preprocess to a new directory (`<workspace>/<run-name>/dataset/.precomputed-v2/` etc.) — preserves the existing data.
|
||||
- Abort.
|
||||
|
||||
**Never pass `--overwrite` without explicit user approval** for this exact action.
|
||||
|
||||
### Step 2 — Invoke `process_dataset.py`
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
uv run python scripts/process_dataset.py \
|
||||
"<workspace>/<run-name>/dataset/dataset.json" \
|
||||
--resolution-buckets "<W>x<H>x<F>" \
|
||||
--model-path "<absolute-model-path>" \
|
||||
--text-encoder-path "<absolute-gemma-path>" \
|
||||
--output-dir "<workspace>/<run-name>/dataset/.precomputed" \
|
||||
--load-text-encoder-in-8bit # on 32GB tier (low-VRAM config), per t2v_lora_low_vram.yaml
|
||||
```
|
||||
|
||||
Add as needed:
|
||||
- `--skip-audio` — if mode doesn't use audio (T2V video-only variants).
|
||||
- `--audio-durations "<list>"` — for T2A from a captions-only file.
|
||||
- `--lora-trigger "<trigger>"` — for style/concept LoRAs.
|
||||
- `--reference-downscale-factor <N>` — for IC-LoRA modes if downscaled references are desired.
|
||||
- `--video-column`, `--caption-column` — only if the metadata file uses non-standard column names.
|
||||
|
||||
**Do not pass `--overwrite`** unless re-preprocessing was explicitly approved in Step 1.
|
||||
|
||||
### Step 3 — Audit output
|
||||
|
||||
After completion, verify:
|
||||
|
||||
```bash
|
||||
ls "<workspace>/<run-name>/dataset/.precomputed/"
|
||||
# Expected subdirs per the mode table above.
|
||||
|
||||
# Count files in each:
|
||||
for d in latents conditions audio_latents reference_latents video_masks audio_masks; do
|
||||
if [ -d "<workspace>/<run-name>/dataset/.precomputed/$d" ]; then
|
||||
echo "$d: $(ls "<workspace>/<run-name>/dataset/.precomputed/$d" | wc -l)"
|
||||
fi
|
||||
done
|
||||
```
|
||||
|
||||
Counts in each required subdir should equal the dataset sample count.
|
||||
|
||||
**Reconcile counts — do this for every run, any dataset size.** Compare the `latents/` (and `audio_latents/`) count against the caption/sample count. If fewer latents were produced, `process_dataset.py` **silently skipped** clips — most commonly because they were **shorter than the target frame bucket** (it logs each skip). When counts don't match:
|
||||
|
||||
1. Identify which clips were dropped (grep the preprocess log for skip/"fewer frames" lines, or diff the produced `.pt` stems against the metadata).
|
||||
2. **Surface it to the user** with the count and the specific clips — never silently proceed on a shrunk dataset.
|
||||
3. Offer options: re-preprocess at a **smaller frame bucket** the clips support, add a **second (shorter) bucket** to keep the short clips (multi-bucket requires `batch_size: 1`), or accept the loss. Let the user decide.
|
||||
|
||||
**Audio gate (hard stop for audio runs).** For any run with an audio modality (joint audio+video, A2V, V2A, T2A, audio-only modes), verify `audio_latents/` is **present and non-empty** with one `.pt` per sample. `process_dataset.py` **swallows audio-decode errors and continues** — it logs "0 videos with audio" and produces empty `audio_latents/` rather than failing. If an audio run produced no audio latents, **stop** — do not proceed to training (it would silently train audio-free). The usual cause is a broken audio decode path (e.g. missing/incompatible `torchcodec`); confirm `uv run python -c "import torchaudio; torchaudio.load('<a clip>')"` works (see `references/troubleshooting.md`), fix it, then re-preprocess with `--overwrite`.
|
||||
|
||||
## Workflow — One-Sample (Phase 6)
|
||||
|
||||
Same as full preprocess, but operate on a single-sample metadata file:
|
||||
|
||||
1. Pick the first sample from `<workspace>/<run-name>/dataset/dataset.json` and write a one-sample metadata file **inside the dataset dir** — e.g. `<workspace>/<run-name>/dataset/_one_sample.json` — copying the entry **verbatim, keeping its relative `media_path`**. `process_dataset.py` resolves media paths relative to the metadata file's own directory, so the one-sample file must sit beside the real media (i.e. in `dataset/`, the same dir as `dataset.json`). **Do not** place it in `overfit/` and **do not** rewrite the path to an absolute one — an absolute path produces mirrored nested output dirs (`.precomputed/latents/absolute/path/.../x.pt`) instead of a clean `latents/x.pt`.
|
||||
2. Run `process_dataset.py` on that file with `--output-dir "<workspace>/<run-name>/overfit/.precomputed"` (output still goes to `overfit/`, only the metadata lives in `dataset/`).
|
||||
3. Use the **same `--resolution-buckets`** as the planned full run. Critical: a small-shape sanity check is misleading because resolution is the dominant memory factor.
|
||||
4. Clean up the temporary `dataset/_one_sample.json` afterward (it's scratch; don't leave it in the dataset dir).
|
||||
|
||||
## Decode-and-Verify (optional debug aid)
|
||||
|
||||
If the user reports validation samples look wrong or training diverges, decode one preprocessed sample back to media:
|
||||
|
||||
`decode_latents.py` takes the **latents directory** and an **output directory** as positional arguments (it decodes the whole directory, not a single `.pt` file). Add `--with-audio` and `--audio-latents-dir` if the dataset has audio.
|
||||
|
||||
```bash
|
||||
cd packages/ltx-trainer
|
||||
uv run python scripts/decode_latents.py \
|
||||
"<workspace>/<run-name>/dataset/.precomputed/latents" \
|
||||
"<workspace>/<run-name>/dataset/.precomputed/decoded_check" \
|
||||
--model-path "<absolute-model-path>"
|
||||
```
|
||||
|
||||
If decoded output is garbled, preprocessing itself is suspect (wrong model, wrong VAE).
|
||||
|
||||
## Failure Modes
|
||||
|
||||
- **"shape mismatch" on resume:** Step 1's verification check. Ask user before any mutation.
|
||||
- **`frames % 8 != 1`** error from process_dataset.py: the requested frame count is invalid; correct in the plan and re-launch.
|
||||
- **VRAM OOM during preprocessing:** add `--load-text-encoder-in-8bit`. If still OOM, reduce `--batch-size`.
|
||||
- **Disk full:** preprocessed latents can be large (especially audio). Surface to user with a `du -sh` summary of `.precomputed/`.
|
||||
|
||||
## Do Not
|
||||
|
||||
- Do not delete or overwrite existing `.precomputed/` data without explicit user approval for that exact action.
|
||||
- Do not preprocess at a smaller resolution to "save time" — the sanity check exists specifically to validate the planned resolution.
|
||||
@@ -0,0 +1,73 @@
|
||||
# Config Patching
|
||||
|
||||
How to safely produce `<workspace>/<run-name>/config.yaml` from an example in `packages/ltx-trainer/configs/`. The trainer's config schema is Pydantic with `extra="forbid"` — unknown fields are rejected. Full field reference: [`packages/ltx-trainer/docs/configuration-reference.md`](../../../../packages/ltx-trainer/docs/configuration-reference.md).
|
||||
|
||||
## Workflow
|
||||
|
||||
1. Copy the example config matching the selected mode (see `mode-selector.md`) to `<workspace>/<run-name>/config.yaml`.
|
||||
2. Patch fields as described below. Preserve YAML comments where possible — they help the user audit the run later.
|
||||
3. **Never** edit the example config in `packages/ltx-trainer/configs/`. That's the user's reference library.
|
||||
|
||||
## Required Patches (every run)
|
||||
|
||||
| Field | Value |
|
||||
|-------|-------|
|
||||
| `model.model_path` | Absolute path to local `.safetensors` (from probe or user). |
|
||||
| `model.text_encoder_path` | Absolute path to local Gemma directory (from probe or user). |
|
||||
| `data.preprocessed_data_root` | `<workspace>/<run-name>/dataset/.precomputed` (absolute). |
|
||||
| `output_dir` | `<workspace>/<run-name>/outputs` (absolute). |
|
||||
|
||||
## Hardware-Driven Patches
|
||||
|
||||
Apply per the matched VRAM tier in `references/hardware-profiles.md`. After autotune (Phase 6), patch the winning trial's deltas in.
|
||||
|
||||
## Schema Constraints (validate before launch)
|
||||
|
||||
These will cause Pydantic errors or runtime failures; check before invoking the trainer.
|
||||
|
||||
- **Frame count:** `validation.video_dims[2]` must satisfy `frames % 8 == 1` (1, 9, 17, 25, 33, 41, 49, 57, 65, 73, 81, 89, ...).
|
||||
- **Resolution:** `validation.video_dims[0]` and `[1]` must be divisible by 32.
|
||||
- **Multi-bucket training:** if dataset uses multiple resolution buckets, set `optimization.batch_size: 1`.
|
||||
- **At least one generated modality:** `training_strategy` must have at least one of `video.is_generated` or `audio.is_generated` set to `true`.
|
||||
- **Audio condition restrictions:** the audio modality cannot use `first_frame` or `spatial_crop` conditions.
|
||||
- **Strategy name:** prefer `training_strategy.name: "flexible"`. `text_to_video` and `video_to_video` still work but emit deprecation warnings.
|
||||
|
||||
## LoRA Patches
|
||||
|
||||
For style/concept LoRAs:
|
||||
|
||||
- `lora.rank` and `lora.alpha`: set from the matched VRAM tier and use case. **32GB tier** pins rank 16 per `t2v_lora_low_vram.yaml`; **80GB+ tier** uses rank 32 per `t2v_lora.yaml`. Keep `alpha == rank`. See `mode-selector.md` for use-case-driven rank guidance.
|
||||
- `lora.target_modules`: short patterns like `"to_k"`, `"to_q"`, `"to_v"`, `"to_out.0"` match all attention modules (video + audio + cross-modal). Add `"ff.net.0.proj"`, `"ff.net.2"` only if user explicitly wants higher capacity.
|
||||
- **Audio-only LoRA targets** (T2A, audio inpainting): use `"audio_attn1.to_*"`, `"audio_attn2.to_*"` patterns to avoid touching video weights. See `configs/t2a_lora.yaml` for the exact list.
|
||||
|
||||
## Validation Sample Prompts
|
||||
|
||||
The example configs ship with placeholder validation prompts. Validation condition fields are documented in
|
||||
[`configuration-reference.md#validation-condition-types`](../../../../packages/ltx-trainer/docs/configuration-reference.md#validation-condition-types).
|
||||
For style/concept LoRAs:
|
||||
- Replace at least one `validation.samples[].prompt` with a prompt that uses the user's trigger word or describes the target concept. Tells the user something useful at the first validation interval.
|
||||
- Keep `validation.video_dims` consistent with the training resolution to make samples comparable.
|
||||
- **Describe the audio, for any run with a generated audio modality** (joint audio+video, T2A, V2A, etc.). The validation prompts must describe the audio the **same way the training captions do** — if the training captions transcribe speech or characterise sound (e.g. *"he says: ‘…’"*, *"calm spoken voice, quiet room tone"*, *"upbeat acoustic guitar"*), the validation prompts must include comparable audio direction. A prompt with no audio description gives the model no guidance for the audio branch and the generated audio comes out poor. This is not speech-specific — any audio (music, ambience, foley) needs describing. Mirror the structure/level of audio detail found in the dataset captions (inspect a few before writing the prompts).
|
||||
|
||||
## W&B Patches
|
||||
|
||||
- If the W&B credential check passes (`uv run python -c "import wandb; print(bool(wandb.Api().api_key))"` → `True`): `wandb.enabled: true`, `wandb.project` = `ltx2-<mode>`, `wandb.tags` includes the mode. (Do not use `wandb status` — it falsely reports `api_key: null` when logged in via netrc.)
|
||||
- If `False`: `wandb.enabled: false`. Surface in plan: *"Not logged in to W&B — run `wandb login` before training to enable tracking."* If the check errored/was ambiguous, ask the user rather than assuming off.
|
||||
|
||||
## Output Dir Behaviour
|
||||
|
||||
The trainer resumes optimizer/scheduler/step state **only when `model.load_checkpoint` is set** to a checkpoint file; it then looks for a matching `training_state_step_*.pt` next to that file. It does **not** auto-detect prior checkpoints in `output_dir/checkpoints/`. To resume, patch `model.load_checkpoint` to the latest checkpoint. To load weights but skip state restore: `checkpoints.no_resume: true`.
|
||||
|
||||
For the orchestrator's resume flow: when resuming an interrupted run, patch `model.load_checkpoint` to the latest checkpoint under `output_dir/checkpoints/`. Leaving it unset starts a fresh run from step 0 even if checkpoints exist on disk.
|
||||
|
||||
## Self-Check Before Launch
|
||||
|
||||
Before any `python scripts/train.py` invocation:
|
||||
|
||||
1. All `model_path`, `text_encoder_path`, `preprocessed_data_root` exist on disk.
|
||||
2. Frame and resolution constraints satisfied (see above).
|
||||
3. Generated modalities have matching latents directories under `.precomputed/`.
|
||||
4. For modes with `reference` condition: `reference_latents/` (and/or `reference_audio_latents/`) exists.
|
||||
5. For modes with `mask` condition: `video_masks/` (and/or `audio_masks/`) exists.
|
||||
|
||||
A failed check at this point is much cheaper than a failed training start.
|
||||
@@ -0,0 +1,98 @@
|
||||
# VRAM Tiers
|
||||
|
||||
Map probed GPU(s) to a starting training config. The autotune sweep in Phase 6 then empirically improves on this baseline. Use these **tier names** in `plan.md` and user-facing chat — not letter codes.
|
||||
|
||||
Source of truth: the two configs shipped in the trainer repo —
|
||||
`packages/ltx-trainer/configs/t2v_lora.yaml` (standard) and
|
||||
`packages/ltx-trainer/configs/t2v_lora_low_vram.yaml` (low VRAM).
|
||||
Per `packages/ltx-trainer/docs/quick-start.md`, the trainer documents
|
||||
**80GB recommended** and **32GB minimum**. Anything below 32GB is
|
||||
unsupported by the project.
|
||||
|
||||
## Probe
|
||||
|
||||
```bash
|
||||
nvidia-smi --query-gpu=name,memory.total --format=csv,noheader
|
||||
```
|
||||
|
||||
Pick the **smallest** VRAM tier across the visible GPUs. Multi-GPU only adds throughput at the same per-GPU memory budget — it doesn't relax per-GPU limits.
|
||||
|
||||
## Minimum Gate
|
||||
|
||||
If per-GPU VRAM is **< 32 GB**, stop the run. Surface to the user:
|
||||
|
||||
> "This GPU has <N>GB VRAM. The LTX-2 trainer requires a minimum of 32GB (see `packages/ltx-trainer/docs/quick-start.md`). Training is unlikely to fit even with maximum memory savings, and we don't ship a tested config below 32GB. Options: (a) abort, (b) try anyway with the low-VRAM config and accept it may OOM — purely at your own risk."
|
||||
|
||||
Do not invent a sub-32GB tier. The trainer team doesn't ship one.
|
||||
|
||||
## 32GB tier — low-VRAM config
|
||||
|
||||
**VRAM range:** 32 GB per GPU (trainer minimum).
|
||||
|
||||
**Typical GPUs:** RTX 5090, V100 32GB.
|
||||
|
||||
Start from `packages/ltx-trainer/configs/t2v_lora_low_vram.yaml` verbatim. Key choices already in that file (do not re-specify in `<workspace>/<run-name>/config.yaml` — copy the file and patch only the paths from `references/config-patching.md`):
|
||||
|
||||
- `optimizer_type: "adamw8bit"`
|
||||
- `enable_gradient_checkpointing: true`
|
||||
- `batch_size: 1`, `gradient_accumulation_steps: 1`
|
||||
- `quantization: "int8-quanto"`
|
||||
- `load_text_encoder_in_8bit: true`
|
||||
- `offload_optimizer_during_validation: true`
|
||||
- `lora.rank: 16`, `lora.alpha: 16`
|
||||
|
||||
Autotune (Phase 6) will sweep `quantization` off, `optimizer_type` → adamw, and `batch_size` up — but at 32GB the sweep often hits OOM on trial 2 or 3. That's fine; the conservative baseline still works.
|
||||
|
||||
## 80GB+ tier — standard config
|
||||
|
||||
**VRAM range:** 80 GB per GPU and above (trainer recommended).
|
||||
|
||||
**Typical GPUs:** A100 80GB, H100 80GB, H200, B200.
|
||||
|
||||
Start from `packages/ltx-trainer/configs/t2v_lora.yaml` verbatim. Key choices already in that file:
|
||||
|
||||
- `optimizer_type: "adamw"`
|
||||
- `enable_gradient_checkpointing: true` (autotune may turn it off if headroom allows)
|
||||
- `batch_size: 1`, `gradient_accumulation_steps: 1`
|
||||
- `quantization: null`
|
||||
- `load_text_encoder_in_8bit: false`
|
||||
- `lora.rank: 32`, `lora.alpha: 32`
|
||||
|
||||
On the 80GB+ tier, the autotune baseline already equals this config (adamw, no quantization), so the quantization/optimizer trials are no-ops; the only real lever is gradient checkpointing off — but for the 22B model that **usually OOMs even with tens of GB of apparent headroom**, so treat a win there as unlikely. The trainer reports its own step-time and peak-VRAM at the end of each run — use those rather than an external timer.
|
||||
|
||||
For ≥140GB GPUs (H200, B200), the same 80GB+ tier baseline applies. FA3/FA4 attention backends are viable on Hopper/Blackwell and can speed up training, but they're optional — the trainer's defaults work on PyTorch SDPA without extra setup.
|
||||
|
||||
## 40–60GB tier — mid-range (autotune from low-VRAM)
|
||||
|
||||
**VRAM range:** 40–60 GB per GPU. The trainer doesn't ship a tested config for this range.
|
||||
|
||||
**Typical GPUs:** A40, A6000 48GB, L40, RTX 6000 Ada.
|
||||
|
||||
Start from the **32GB tier** (low-VRAM config) and let autotune relax `quantization`, `optimizer_type`, and `batch_size` based on actual headroom. Don't pre-bake intermediate YAML values that haven't been measured. Surface this as **40–60GB tier** in the plan.
|
||||
|
||||
## Multi-GPU
|
||||
|
||||
If `nvidia-smi` reports N ≥ 2 GPUs of the same model:
|
||||
|
||||
- Launch with `uv run accelerate launch scripts/train.py <config>`.
|
||||
- Use `packages/ltx-trainer/configs/accelerate/fsdp.yaml` for full fine-tune.
|
||||
- DDP (default `accelerate launch` without a config file) is fine for LoRA.
|
||||
- Effective batch = `batch_size * gradient_accumulation_steps * num_gpus`. Reduce `gradient_accumulation_steps` proportionally to keep the effective batch consistent with the plan.
|
||||
|
||||
## Full Fine-Tune
|
||||
|
||||
If the user chose full fine-tune (`model.training_mode: "full"`):
|
||||
|
||||
- Require multi-GPU + FSDP on 80GB+ tier GPUs. Otherwise warn in the plan that single-GPU full FT is unlikely to fit and propose LoRA instead.
|
||||
- Set `acceleration.offload_optimizer_during_validation: true` always (optimizer state is huge under full FT).
|
||||
|
||||
## Model Path Constraints
|
||||
|
||||
- `model.model_path`: local `.safetensors` only. No URLs.
|
||||
- `model.text_encoder_path`: local Gemma model directory. No URLs.
|
||||
|
||||
If probe didn't find these in conventional locations (`/models/`, `~/models/`, `$LTX_MODELS_DIR`), ask the user in Phase 3 (or offer to download per `references/onboarding.md`).
|
||||
|
||||
## Notes on Loss-of-Generality
|
||||
|
||||
The two anchor tiers (32GB and 80GB+) correspond directly to the two configs the trainer ships. The autotune sweep is the empirical layer — if a particular GPU consistently lands on a different stable config, **update the relevant trainer config first**, not this file. This skill follows the trainer's choices, not the other way around.
|
||||
@@ -0,0 +1,74 @@
|
||||
# Mode Selector
|
||||
|
||||
Map the user's stated intent to a `flexible`-strategy configuration. All modes are supported via a single strategy (`training_strategy.name: "flexible"`); the difference is which modality is generated and which `conditions` are attached.
|
||||
|
||||
> Full reference: [`packages/ltx-trainer/docs/training-modes.md`](../../../../packages/ltx-trainer/docs/training-modes.md). This file is the **lookup table** for translating user intent.
|
||||
|
||||
## Decision Table
|
||||
|
||||
| User says (roughly)... | Mode | Example config | Modalities | Conditions |
|
||||
|------------------------|------|----------------|------------|------------|
|
||||
| "generate videos from text", "T2V LoRA" | T2V | `configs/t2v_lora.yaml` | video gen, audio gen | none |
|
||||
| "generate videos from a starting image", "I2V" | I2V | `configs/i2v_lora.yaml` | video gen, audio gen | `first_frame` (video) |
|
||||
| **plain concept/style LoRA** ("train a LoRA on X", no specific task) | **I2V by default** (see note) | `configs/i2v_lora.yaml` | video gen, audio gen | `first_frame` (video), `probability: 0.5` |
|
||||
| "extend a video forward in time" | Video extension (prefix) | `configs/video_extend_lora.yaml` | video gen, audio gen | `prefix` (video) |
|
||||
| "extend a video backward in time" | Video extension (suffix) | `configs/video_suffix_lora.yaml` | video gen, audio gen | `suffix` (video) |
|
||||
| "fill in masked regions of a video" | Video inpainting | `configs/video_inpainting_lora.yaml` | video gen | `mask` (video) |
|
||||
| "expand a video beyond its borders" | Video outpainting | `configs/video_outpainting_lora.yaml` | video gen | `spatial_crop` (video) |
|
||||
| "style transfer from reference video", "IC-LoRA", "depth/pose/canny control" | V2V IC-LoRA | `configs/v2v_ic_lora.yaml` | video gen | `reference` (video) |
|
||||
| "generate video to match an audio track" | A2V | `configs/a2v_lora.yaml` | video gen, audio frozen | none (audio `is_generated: false`) |
|
||||
| "add sound effects to silent video", "foley", "V2A" | V2A | `configs/v2a_lora.yaml` | video frozen, audio gen | none (video `is_generated: false`) |
|
||||
| "generate audio from text", "T2A" | T2A | `configs/t2a_lora.yaml` | audio gen | none |
|
||||
| "extend audio forward / backward" | Audio extension | `configs/audio_extend_lora.yaml`, `configs/audio_suffix_lora.yaml` | audio gen | `prefix` / `suffix` (audio) |
|
||||
| "fill in masked regions of audio" | Audio inpainting | `configs/audio_inpainting_lora.yaml` | audio gen | `mask` (audio) |
|
||||
| "audio style transfer from reference", "A2A IC-LoRA" | A2A IC-LoRA | `configs/a2a_ic_lora.yaml` | audio gen | `reference` (audio) |
|
||||
| "joint video+audio reference control" | AV2AV IC-LoRA | `configs/av2av_ic_lora.yaml` | video gen, audio gen | `reference` (both) |
|
||||
| Any of the above with full fine-tune | Full FT variant | as above, set `model.training_mode: "full"` | (mode-specific) | (mode-specific) |
|
||||
|
||||
### Why I2V is the default for a plain concept/style LoRA
|
||||
|
||||
A "train a LoRA on X" request isn't tied to one inference mode: **LoRA weights are pipeline-agnostic** — the same checkpoint loads in both T2V and I2V inference (both use `TI2VidOneStagePipeline`/`TwoStages`). The `i2v_lora` config trains `first_frame` with **`probability: 0.5`**, so the model learns both first-frame-conditioned (I2V) and unconditioned (T2V) generation in one run, and the first frame comes from each training clip automatically (no extra data prep). That makes I2V a versatile **superset** — usable for both at inference at no extra cost — which is why it's the default for a plain LoRA. Ask the user how they'll use it (text-only / from an image / both) and only drop to `t2v_lora` if they're sure it's text-only. (See the orchestrator `SKILL.md` Phase 1.)
|
||||
|
||||
## Disambiguation Questions
|
||||
|
||||
When the user's first answer is ambiguous, ask **one** follow-up:
|
||||
|
||||
- "extend a video" → forward or backward in time?
|
||||
- "control with a reference" → video reference (depth/pose/canny/etc.) or audio reference?
|
||||
- "fill in regions" → video regions (masked frames) or audio regions (masked time)?
|
||||
- "T2V" → joint video+audio (default) or video-only?
|
||||
- LoRA or full fine-tune? Default to LoRA unless the user has multi-GPU + clear reason for full.
|
||||
|
||||
## Combining Modes
|
||||
|
||||
The flexible strategy allows stacking conditions. Common combinations:
|
||||
|
||||
- **I2V + V2A** (start frame + generate audio for the resulting video) — not directly expressible; would need two passes.
|
||||
- **Video extension + audio extension** — both modalities generate, both have `prefix` condition. Express in one config.
|
||||
- **IC-LoRA + I2V** — `first_frame` + `reference` conditions on the video modality.
|
||||
|
||||
If the user asks for a combination not listed, check `packages/ltx-trainer/src/ltx_trainer/training_strategies/flexible.py` for which conditions can co-exist on a modality. Audio modality cannot use `first_frame` or `spatial_crop`.
|
||||
|
||||
## When the Intent Doesn't Map
|
||||
|
||||
If after disambiguation there is no entry in the table and no combination of `flexible` conditions covers the user's request, go to the **Escape Hatch** section in the orchestrator `SKILL.md`. Do not silently pick the closest mode.
|
||||
|
||||
## LoRA Rank by Use Case
|
||||
|
||||
Once the mode is picked, choose `lora.rank` (and matching `lora.alpha`) based on what the LoRA is supposed to capture. These are starting points; autotune doesn't sweep rank because it's a quality knob, not a step-time one.
|
||||
|
||||
| Use case | Suggested rank | Notes |
|
||||
|----------|----------------|-------|
|
||||
| Single character, single object, single style | 32–64 | Default for most concept LoRAs. Start at 32; bump to 64 if validation samples underfit. |
|
||||
| Multi-character world, dense series, complex multi-concept | 96–128 | More capacity for distinguishing several concepts inside one LoRA. |
|
||||
| Camera move, motion, transition (i.e. behavioural, not visual) | 8–16 | Motion is a thin signal — high ranks just memorise frame content. |
|
||||
| IC-LoRA control (V2V depth/pose/Canny/etc., A2A audio reference) | 16–32 | Start at 16 for structural control (depth, pose, edges); 24–32 if the reference carries richer style/texture. Video IC-LoRA often lands lower than concept LoRAs. |
|
||||
| LTX-2 trainer's default if unsure | 32 | Safe baseline. |
|
||||
|
||||
On the **32GB tier** (`hardware-profiles.md`), the low-VRAM config already pins `lora.rank: 16`. If the user picks a higher-rank use case on 32GB, surface the trade-off in the plan but let them decide — the autotune sweep doesn't touch rank, so a too-high rank will simply OOM at training time.
|
||||
|
||||
Keep `alpha == rank` unless the user has a specific reason otherwise (effective scaling = `alpha / rank`).
|
||||
|
||||
## Starting Config
|
||||
|
||||
Copy the matching example config — the exact filename from the **Example config** column of the decision table above (e.g. T2V → `configs/t2v_lora.yaml`, I2V → `configs/i2v_lora.yaml`, A2A IC-LoRA → `configs/a2a_ic_lora.yaml`) — into `<workspace>/<run-name>/config.yaml` as the starting point. Then patch per `references/config-patching.md`.
|
||||
@@ -0,0 +1,115 @@
|
||||
# First-Run Onboarding
|
||||
|
||||
What the orchestrator's Phase 2 probe checks for, what to do when something is missing, and what the skill is allowed to set up automatically (with explicit user approval).
|
||||
|
||||
## Prerequisites Checked in Phase 2
|
||||
|
||||
| Prerequisite | How to detect | If missing |
|
||||
|--------------|---------------|------------|
|
||||
| CUDA GPU visible | `nvidia-smi` returns ≥1 GPU | Stop. Training requires CUDA — point the user at non-LTX-2 docs. |
|
||||
| Linux | `uname -s` returns `Linux` | Stop. Trainer uses Triton (Linux-only). |
|
||||
| `uv` installed | `command -v uv` | Offer to install (see "Auto-setup" below). |
|
||||
| Workspace synced | `[ -f uv.lock ] && uv pip list \| grep -q ltx-trainer` | Offer to run `uv sync` from repo root. |
|
||||
| LTX-2 model weights | Search `/models/`, `~/models/`, `$LTX_MODELS_DIR` for a `.safetensors` matching `*ltx*2*` | Offer to download (see "Model downloads"). |
|
||||
| Gemma text encoder dir | Search same locations for a directory containing Gemma config | Offer to download. |
|
||||
| Captioner backend | Gemini auth (`GEMINI_API_KEY`/`GOOGLE_API_KEY` or gcloud/Vertex), OR a ≥40 GiB GPU to host the Qwen3-Omni-30B vLLM server (FP8), OR captions already in the dataset | See "Captioner graceful degradation" below. **Check the HF cache for an already-downloaded Qwen model before assuming a download is needed** (see note below the table). |
|
||||
| W&B login (optional) | `uv run python -c "import wandb; print(bool(wandb.Api().api_key))"` → `True` means logged in. Uses wandb's own credential resolution (env/netrc/settings). **Don't** use `wandb status` (reports `api_key: null` even when logged in). | Not a blocker. If `False`: disabled in config + flagged in plan. If the check errors/ambiguous: ask the user, don't assume off. |
|
||||
| Disk space | `df -h $WORKSPACE` | Surface available space alongside what a run consumes (preprocessed latents, several-GB checkpoints, validation samples). Flag concerns to the user; don't enforce a hard threshold. |
|
||||
|
||||
Surface findings as a compact table in chat. For each missing item, present the user with a concrete next step (download command, install command, or "skip this — here's the consequence").
|
||||
|
||||
**Check the HF cache before declaring the local captioner unavailable.** The Qwen3-Omni model is served from the HuggingFace cache, not `~/models/`. Before concluding it must be downloaded (or ruling it out on free-disk grounds), check whether it's already cached:
|
||||
|
||||
```bash
|
||||
ls -d "${HF_HOME:-$HOME/.cache/huggingface}"/hub/models--Qwen--Qwen3-Omni* 2>/dev/null \
|
||||
&& du -sh "${HF_HOME:-$HOME/.cache/huggingface}"/hub/models--Qwen--Qwen3-Omni* 2>/dev/null
|
||||
```
|
||||
|
||||
If it's cached (~60 GiB, all shards present), no download is needed — don't rule out the local captioner because of low free disk on the *home* partition; the weights already exist. Only the GPU-VRAM constraint (≥40 GiB for FP8) then applies.
|
||||
|
||||
## Auto-Setup (with explicit user approval)
|
||||
|
||||
The skill may, only after the user explicitly says yes, do these setup actions. Each action is a single discrete question.
|
||||
|
||||
### Run `uv sync`
|
||||
|
||||
```bash
|
||||
cd <repo-root>
|
||||
uv sync
|
||||
```
|
||||
|
||||
Ask: *"Repo not synced. Run `uv sync` now? It will download the project's Python dependencies."*
|
||||
|
||||
### Install `uv`
|
||||
|
||||
Ask: *"`uv` not installed. Install via the official one-liner now?"*
|
||||
|
||||
```bash
|
||||
curl -LsSf https://astral.sh/uv/install.sh | sh
|
||||
```
|
||||
|
||||
After install, ask the user to restart their shell or `source ~/.bashrc` before continuing.
|
||||
|
||||
### Model downloads
|
||||
|
||||
Use `huggingface-cli` (comes with `huggingface-hub`, transitively pulled by `uv sync`). Default destination: `$LTX_MODELS_DIR` if set, else `~/models/` (create if missing). Surface destination in the prompt — never download into the repo or into the workspace.
|
||||
|
||||
**LTX-2 base model:**
|
||||
|
||||
```bash
|
||||
huggingface-cli download Lightricks/LTX-2.3 \
|
||||
ltx-2.3-22b-dev.safetensors \
|
||||
--local-dir ~/models/ltx-2.3
|
||||
```
|
||||
|
||||
Public reference: <https://huggingface.co/Lightricks/LTX-2.3>.
|
||||
|
||||
**Gemma text encoder:**
|
||||
|
||||
```bash
|
||||
huggingface-cli download google/gemma-3-12b-it-qat-q4_0-unquantized \
|
||||
--local-dir ~/models/gemma-3-12b
|
||||
```
|
||||
|
||||
**Qwen3-Omni captioner (only if the GPU can host it):**
|
||||
|
||||
The local captioner is Qwen3-Omni-30B-A3B-Thinking served by a vLLM server (`scripts/serve_captioner.py`), which downloads the model (~65 GiB) on first launch via `uvx vllm` — there's no separate `huggingface-cli` step. Default **FP8** (~31 GiB weights) fits on a **40 GiB** GPU; **bf16** (~60 GiB) needs **≥66 GiB free VRAM**. On a GPU below ~40 GiB (typical consumer 24/32 GB cards), don't use it — use Gemini instead (see "Captioner graceful degradation").
|
||||
|
||||
The base model + text encoder are large (multi-GB) downloads. Ask the user before each one — `huggingface-cli` reports the actual size at the start of the transfer. Do **not** batch them into a single "yes/no"; the user may want only what's missing.
|
||||
|
||||
### Hugging Face login (if any downloads fail with 401)
|
||||
|
||||
Ask: *"Hugging Face download requires login (some Lightricks models are gated). Run `huggingface-cli login` now? You'll need a token from <https://huggingface.co/settings/tokens>."*
|
||||
|
||||
## Captioner Graceful Degradation
|
||||
|
||||
The captioner is the trickiest prerequisite. The local backend (`qwen_omni`) is now a **30B model served by a vLLM server** — ~65 GiB download; FP8 fits on a 40 GiB GPU, bf16 needs ≥66 GiB. Typical consumer cards (24/32 GB) can't host it, so for most users **prefer Gemini**.
|
||||
|
||||
Decision tree:
|
||||
|
||||
1. **User already has captions in their dataset metadata** → skip the captioner entirely (Step 3 of `prepare-dataset` skips when captions are present).
|
||||
2. **GPU ≥40 GiB (FP8) / ≥66 GiB (bf16)** → `qwen_omni` is viable: launch `scripts/serve_captioner.py` first, then caption. Gemini is still fine here too.
|
||||
3. **GPU below ~40 GiB (the common consumer case)** → `qwen_omni` is not an option. Recommend **`gemini_flash`** and tell the user they need Gemini auth: a `GEMINI_API_KEY`/`GOOGLE_API_KEY` (get one at <https://aistudio.google.com/apikey>) **or** working gcloud/Vertex AI credentials.
|
||||
4. **No Gemini auth and can't run Qwen3** → the only remaining path is bringing their own captions: add a `caption` column to the dataset metadata, then re-invoke the skill.
|
||||
|
||||
Wait for the user's choice. Don't pick one automatically — but make the hardware reality explicit so they don't try to run the server on a card that can't host it.
|
||||
|
||||
## What Auto-Setup Does NOT Touch
|
||||
|
||||
- Does not modify the user's shell rc files except by explicit instruction (e.g., "add `export LTX_MODELS_DIR=...` to your `.bashrc`?" with the user agreeing).
|
||||
- Does not modify `~/.gitconfig`, `~/.ssh/`, or any auth-related files.
|
||||
- Does not install GPU drivers, CUDA, or system packages.
|
||||
- Does not delete or move existing model files. If a model is found but at an unexpected path, surface and let the user decide.
|
||||
- Does not download into the repo or workspace. Models live at `~/models/` (or `$LTX_MODELS_DIR`).
|
||||
|
||||
## Configuration Inheritance
|
||||
|
||||
After downloads, the skill records the resolved paths into the plan's Assumptions section and into `config.yaml`:
|
||||
|
||||
```yaml
|
||||
model:
|
||||
model_path: "/home/<user>/models/ltx-2.3/ltx-2.3-22b-dev.safetensors"
|
||||
text_encoder_path: "/home/<user>/models/gemma-3-12b"
|
||||
```
|
||||
|
||||
Suggest (don't enforce) setting `LTX_MODELS_DIR=~/models` in their shell rc for future runs.
|
||||
@@ -0,0 +1,107 @@
|
||||
# Plan Template
|
||||
|
||||
Write `<workspace>/<run-name>/plan.md` following this template. The plan is the user's contract with the agent — it's the gate before any heavy work runs.
|
||||
|
||||
Scale each section to its relevance. Skip subsections that don't apply (e.g., no captioning if data is already captioned), but never collapse "Assumptions" or "Cost/time estimate".
|
||||
|
||||
```markdown
|
||||
# Training Plan — <run-name>
|
||||
|
||||
## Goal
|
||||
|
||||
<One paragraph restating the user's intent in their own terms.>
|
||||
|
||||
## Mode
|
||||
|
||||
**<Mode name>** — <one-line rationale linking user intent to mode>.
|
||||
|
||||
- Config base: the concrete example config selected by the mode (e.g. `packages/ltx-trainer/configs/t2v_lora.yaml`) — use the actual filename, not a placeholder
|
||||
- Conditions: <list, or "none">
|
||||
- Training mode: <`lora` | `full`>
|
||||
|
||||
## Dataset
|
||||
|
||||
- Source: `<absolute path>` (<N samples>)
|
||||
- Captions: <`already present` | `will be generated with <backend>`>
|
||||
- Audio: <`present` | `absent — using --skip-audio` | `to be paired with --audio-durations`>
|
||||
- IC-LoRA references: <`present` | `to be generated via compute_reference.py` | `n/a`>
|
||||
|
||||
## Preprocessing
|
||||
|
||||
- Target resolution buckets: `<W>x<H>x<F>` (frames satisfy `frames % 8 == 1`; W,H divisible by 32)
|
||||
- Estimated time: ~<duration> on detected hardware
|
||||
- Output: `<workspace>/<run-name>/dataset/.precomputed/`
|
||||
|
||||
## Training Config
|
||||
|
||||
| Field | Value |
|
||||
|-------|-------|
|
||||
| Optimizer | <adamw / adamw8bit> |
|
||||
| Mixed precision | <bf16 / fp16> |
|
||||
| Quantization | <null / int8-quanto / ...> |
|
||||
| Gradient checkpointing | <on / off> |
|
||||
| Batch size | <N> |
|
||||
| Gradient accumulation | <N> (effective batch = <N>) |
|
||||
| Steps | <N> |
|
||||
| Learning rate | <value> |
|
||||
| LoRA rank / alpha | <N / N> (or "full FT") |
|
||||
| LoRA target modules | <list> (or "n/a") |
|
||||
| LoRA trigger word | <word> (or "n/a") |
|
||||
| Validation interval | every <N> steps |
|
||||
| Checkpoint interval | every <N> steps |
|
||||
|
||||
## Hardware
|
||||
|
||||
- GPU(s): <name> x <count>, <VRAM>GB per GPU
|
||||
- Launch: <`python scripts/train.py` (single) | `accelerate launch` (multi)>
|
||||
- VRAM tier: <32GB tier | 40–60GB tier | 80GB+ tier> — <low-VRAM config (`t2v_lora_low_vram.yaml`) | standard config (`t2v_lora.yaml`) | mid-range, autotuned from low-VRAM>
|
||||
|
||||
## Sanity Check + Autotune
|
||||
|
||||
*In plain terms: before committing to the full run, I do a quick dry run on a single clip at your target resolution to catch out-of-memory or config errors in ~2 minutes (rather than failing hours in), then try a few config variants to pick the fastest one that fits your GPU.*
|
||||
|
||||
Mechanics:
|
||||
- 1 sample, full target resolution, 50 steps + 1 validation pass.
|
||||
- Autotune sweep: up to 5 trials varying quantization / optimizer / batch size.
|
||||
- Stops at first OOM or no-improvement.
|
||||
|
||||
## Monitoring
|
||||
|
||||
<One of:>
|
||||
- W&B: enabled, project `<name>`, entity `<entity-or-default>`. URL will be surfaced once training starts.
|
||||
- W&B: **not logged in** — run `wandb login` before training to enable tracking. Otherwise training proceeds without remote logging.
|
||||
|
||||
## Outputs
|
||||
|
||||
- Training config: `<workspace>/<run-name>/config.yaml`
|
||||
- Checkpoints: `<workspace>/<run-name>/outputs/checkpoints/`
|
||||
- Validation samples: `<workspace>/<run-name>/outputs/samples/`
|
||||
- Autotune log: `<workspace>/<run-name>/autotune.log`
|
||||
|
||||
## Assumptions
|
||||
|
||||
Defaults the agent chose silently. Override any by replying with the new value.
|
||||
|
||||
- <list every non-trivial assumed value: precision, scheduler type, seed, validation prompts, checkpoint retention, etc.>
|
||||
|
||||
## Cost / Time Estimate
|
||||
|
||||
Give only estimates you can ground; label anything not yet measured as rough. Do **not** state a confident training duration before the sanity check has measured a real step time — say "training duration TBD until the sanity check measures step time" and fill it in afterward (per Hard Invariant #5: no fabricated predictions).
|
||||
|
||||
- Captioning: ~<duration> (rough)
|
||||
- Preprocessing: ~<duration> (rough)
|
||||
- Sanity check + autotune: ~<duration> (rough)
|
||||
- Full training: **measured after sanity check** — then `<measured step-time> × <steps>`
|
||||
- **Total wall-clock estimate:** rough until step time is measured; refine after the sanity check.
|
||||
|
||||
## Approve to Proceed
|
||||
|
||||
Reply "approve" (or with edits) to start. No captioning, preprocessing, autotune, or training will run before approval.
|
||||
```
|
||||
|
||||
## Notes on Writing the Plan
|
||||
|
||||
- Show numbers, not adjectives. "~3 hours" beats "fairly long."
|
||||
- Surface every assumption that, if wrong, would cost the user time. Better to over-list than under-list — the user can skim.
|
||||
- If a section reveals you need to ask another question, **stop and ask** before finalizing the plan. The plan is the last gate, not the first.
|
||||
- If the user's hardware can't reasonably support the requested mode (e.g., single-GPU full FT on a 32GB consumer card), say so plainly in the plan and propose the alternative (LoRA, multi-GPU, etc.), rather than silently downgrading.
|
||||
@@ -0,0 +1,84 @@
|
||||
# Troubleshooting
|
||||
|
||||
Quick lookup for failures during sanity check, preprocessing, or training. For deeper coverage see `packages/ltx-trainer/docs/troubleshooting.md`.
|
||||
|
||||
## OOM During Training Step
|
||||
|
||||
Order of operations (cheapest first):
|
||||
|
||||
1. `optimization.enable_gradient_checkpointing: true` (if not already on).
|
||||
2. `optimization.batch_size: 1` and increase `gradient_accumulation_steps` to preserve effective batch.
|
||||
3. `optimization.optimizer_type: "adamw8bit"`.
|
||||
4. `acceleration.quantization: "int8-quanto"`.
|
||||
5. Reduce `lora.rank` (32 → 16 → 8). Alpha follows rank.
|
||||
6. Reduce target resolution (`validation.video_dims` and re-preprocess the dataset at the new resolution).
|
||||
|
||||
The last option is expensive — flag it clearly to the user before re-preprocessing.
|
||||
|
||||
## OOM During Validation Sample Generation
|
||||
|
||||
The validation pass loads decoders + runs CFG/STG inference; it can OOM even when the training step fits.
|
||||
|
||||
1. `acceleration.load_text_encoder_in_8bit: true` (trainer config, not the dataset script).
|
||||
2. `acceleration.offload_optimizer_during_validation: true` (especially for full FT or high-rank LoRA).
|
||||
3. Reduce `validation.video_dims` (smaller validation than training is fine — it's only for visual feedback).
|
||||
4. Reduce `validation.inference_steps` (e.g. 30 → 20).
|
||||
5. Increase `validation.interval` to validate less often.
|
||||
|
||||
## NaN Loss
|
||||
|
||||
1. Check `acceleration.mixed_precision_mode`: prefer `"bf16"`. If `"fp16"`, switch.
|
||||
2. Verify dataset latents are well-formed: `uv run python scripts/decode_latents.py <latents-dir> <output-dir> --model-path <model>` (it decodes a whole latents directory, not a single `.pt`) should reconstruct sensibly.
|
||||
3. Lower `optimization.learning_rate` by 5x.
|
||||
4. Add `optimization.max_grad_norm: 1.0` (default; verify it's set).
|
||||
5. If using `quantization`, try `null` — INT8/INT4 quantization can interact badly with poorly-conditioned LoRA inits at high LR.
|
||||
|
||||
## Validation Samples Look Wrong but Loss Is Fine
|
||||
|
||||
Often not a bug — validation uses simplified inference. For real quality assessment, run a checkpoint through `packages/ltx-pipelines/` after training.
|
||||
|
||||
## Trainer Won't Start: Config Validation Error
|
||||
|
||||
Pydantic `extra="forbid"` means typos in field names fail loudly. Read the error carefully — it names the offending field and path. Fix and re-launch.
|
||||
|
||||
Common offenders:
|
||||
- `latents_dir` typo or wrong relative path.
|
||||
- A required field genuinely missing after copying an example (most fields have defaults; check the error message for the exact field path).
|
||||
- `target_modules` listed at wrong nesting level (must be under `lora:`).
|
||||
|
||||
## Trainer Won't Start: Missing Files
|
||||
|
||||
- `model_path` not found → re-probe `/models/`, `~/models/`, `$LTX_MODELS_DIR`, or ask the user.
|
||||
- `text_encoder_path` directory missing the Gemma config → ensure the path is to the Gemma model dir, not its parent.
|
||||
- `preprocessed_data_root` doesn't contain expected subdirs → re-verify Phase 7 ran for the chosen mode (see `phases/preprocess-dataset.md`).
|
||||
|
||||
## Resume Stops Working
|
||||
|
||||
The trainer **does not auto-resume from `output_dir`**. Resume happens only when `model.load_checkpoint` is explicitly set to a checkpoint file; the trainer then loads those weights and looks for a `training_state_step_*.pt` next to that file to restore optimizer/scheduler/step. Common pitfalls:
|
||||
- `model.load_checkpoint` not set or set to the wrong path → fresh run from step 0 even when `outputs/checkpoints/` is full of artifacts. Patch `model.load_checkpoint` to the latest checkpoint.
|
||||
- `checkpoints.no_resume: true` is set → weights load but state is discarded. Remove the flag if you want a proper resume.
|
||||
- `training_state_step_*.pt` missing from next to the loaded checkpoint → weights load but step counter resets. Make sure the state file accompanies the checkpoint.
|
||||
- `training_state_step_*.pt` corrupted (size 0, fails `torch.load`) → trainer falls back to step 0 with a warning.
|
||||
|
||||
## Autotune Trial Failed Mid-Sweep
|
||||
|
||||
- If trial 2 (quantization off) OOMs: revert to trial 1 and stop the sweep. The 32GB tier baseline is correctly aggressive.
|
||||
- If trial 3 (adamw) OOMs: revert to adamw8bit. Continue with trial 4 if VRAM headroom allows.
|
||||
- If trial 4 (batch_size up) OOMs: revert and stop. We've found the ceiling.
|
||||
|
||||
Never carry over a failing trial's deltas. Always revert to the last-known-good before the next change.
|
||||
|
||||
## Captioning Is Slow
|
||||
|
||||
- Local `qwen_omni` is a 30B model served by `serve_captioner.py` (vLLM). If the server won't start or OOMs on launch: keep the default `--quantization fp8` (don't use `bf16` unless ≥66 GiB free VRAM), lower `--gpu-memory-utilization`, or reduce `--max-model-len`. If the GPU can't host a 30B model at all, switch to `gemini_flash`.
|
||||
- `caption_videos.py --captioner-type qwen_omni` errors connecting → the vLLM server isn't running or `--vllm-url` doesn't match. Start `serve_captioner.py` first and confirm the port.
|
||||
- For most hardware and for larger datasets, prefer `--captioner-type gemini_flash --num-workers <N>` (needs Gemini auth: `GEMINI_API_KEY`/`GOOGLE_API_KEY` or gcloud/Vertex) — runs anywhere and parallelises; local Qwen needs a heavy GPU and a running server.
|
||||
|
||||
## Process_Dataset Errors
|
||||
|
||||
- "frames divisible by..." → the video doesn't have enough frames at the requested temporal resolution. Either shorten the requested frame count or use `split_scenes.py` to break long videos.
|
||||
- "shape mismatch" on existing `.precomputed/` → user requested a different resolution than the existing data. Per the invariants, **stop and ask** — do not overwrite. Offer: reuse at old resolution / re-preprocess to a new dir / abort.
|
||||
|
||||
## When To Give Up and Ask The User
|
||||
|
||||
If a fix isn't obvious from this file or `packages/ltx-trainer/docs/troubleshooting.md` within two attempts, stop and surface the full error + the steps already tried to the user. Don't loop indefinitely on autonomous fixes — the user has context the agent doesn't (which checkpoints are precious, what they care about preserving, etc.).
|
||||
@@ -2,7 +2,9 @@
|
||||
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
||||
*.sft filter=lfs diff=lfs merge=lfs -text
|
||||
*.pt filter=lfs diff=lfs merge=lfs -text
|
||||
*.mp3 filter=lfs diff=lfs merge=lfs -text
|
||||
*.mp4 filter=lfs diff=lfs merge=lfs -text
|
||||
packages/ltx-pipelines/tests/assets/*.wav filter=lfs diff=lfs merge=lfs -text
|
||||
*.png filter=lfs diff=lfs merge=lfs -text
|
||||
*.jpeg filter=lfs diff=lfs merge=lfs -text
|
||||
*.jpg filter=lfs diff=lfs merge=lfs -text
|
||||
|
||||
+21
-1
@@ -17,8 +17,10 @@ checkpoints/
|
||||
|
||||
# Other files
|
||||
.DS_Store
|
||||
tmp
|
||||
.wandb
|
||||
projects/
|
||||
tmp
|
||||
wandb/
|
||||
|
||||
# Model checkpoints
|
||||
*.ckpt
|
||||
@@ -36,6 +38,7 @@ tmp
|
||||
*.json
|
||||
*.m4a
|
||||
*.mov
|
||||
*.mp3
|
||||
*.mp4
|
||||
*.png
|
||||
*.wav
|
||||
@@ -45,5 +48,22 @@ tmp
|
||||
!packages/ltx-pipelines/tests/assets/expected_hdr_ic_lora_exr/frame_*.exr
|
||||
!packages/ltx-pipelines/tests/assets/hdr_ic_lora_test_input.mp4
|
||||
|
||||
# Text-to-audio (T2A) e2e test baseline (checked in via Git LFS)
|
||||
!packages/ltx-pipelines/tests/assets/expected_t2a_one_stage_ltx2_3.wav
|
||||
|
||||
# ltx-bench Grafana dashboards (source of truth in the repo)
|
||||
!packages/ltx-bench/grafana/*.json
|
||||
|
||||
# ltx-trainer E2E test dataset (committed via Git LFS).
|
||||
# The target_audio/ and reference_audio/ subdirectories are populated lazily at test runtime
|
||||
# (high-bitrate MP3 extracted from target_videos, plus a low-bitrate MP3 reference); their
|
||||
# contents stay gitignored via the root *.mp3 rule above.
|
||||
!packages/ltx-trainer/tests/assets/test_dataset/dataset.json
|
||||
!packages/ltx-trainer/tests/assets/test_dataset/target_videos/*.mp4
|
||||
!packages/ltx-trainer/tests/assets/test_dataset/reference_videos/*.mp4
|
||||
!packages/ltx-trainer/tests/assets/test_dataset/conditioning/first_frame.jpg
|
||||
!packages/ltx-trainer/tests/assets/test_dataset/conditioning/video_mask.png
|
||||
!packages/ltx-trainer/tests/assets/test_dataset/conditioning/audio_mask.pt
|
||||
|
||||
# Binary files
|
||||
*.so
|
||||
|
||||
@@ -77,7 +77,7 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
|
||||
|
||||
* **Use DistilledPipeline** - Fastest inference with only 8 predefined sigmas (8 steps stage 1, 4 steps stage 2)
|
||||
* **Enable FP8 quantization** - Enables lower memory footprint: `--quantization fp8-cast` (CLI) or `quantization=QuantizationPolicy.fp8_cast()` (Python). Fp8-cast should be used with bf16 checkpoints, it shall downcast them on the fly. For Hopper GPUs with TensorRT-LLM, use `--quantization fp8-scaled-mm` for FP8 scaled matrix multiplication. Fp8-scaled-mm should be used with fp8 checkpoints.
|
||||
* **Install attention optimizations** - Use xFormers (`uv sync --extra xformers`) or [Flash Attention 3](https://github.com/Dao-AILab/flash-attention) for Hopper GPUs
|
||||
* **Install attention optimizations** - On datacenter Blackwell GPUs (B200), install FlashAttention 4 manually: `uv pip install 'flash-attn-4==4.0.0b9'` (this specific revision is the one we have verified against torch 2.9.1+cu128; newer betas have known issues on consumer Blackwell). On other CUDA GPUs (including Hopper), use xFormers (`uv sync --extra xformers`).
|
||||
* **Use gradient estimation** - Reduce inference steps from 40 to 20-30 while maintaining quality (see [pipeline documentation](packages/ltx-pipelines/README.md#denoising-loop-optimization))
|
||||
* **Skip memory cleanup** - If you have sufficient VRAM, disable automatic memory cleanup between stages for faster processing
|
||||
* **Choose single-stage pipeline** - Use `TI2VidOneStagePipeline` for faster generation when high resolution isn't required
|
||||
|
||||
+48
-24
@@ -8,6 +8,7 @@ The foundational library for the LTX-2 Audio-Video generation model. This packag
|
||||
- **`conditioning/`**: Tools for preparing latent states and applying conditioning (image, video, keyframes)
|
||||
- **`guidance/`**: Perturbation system for fine-grained control over attention mechanisms
|
||||
- **`loader/`**: Utilities for loading weights from `.safetensors`, fusing LoRAs, and managing memory
|
||||
- **`block_streaming/`**: Memory-efficient inference that streams transformer blocks through the GPU one at a time (from pinned CPU buffers or directly from disk)
|
||||
- **`model/`**: PyTorch implementations of the LTX-2 Transformer, Video VAE, Audio VAE, Vocoder and Upscaler
|
||||
- **`text_encoders/gemma`**: Gemma text encoder implementation with tokenizers, feature extractors, and separate encoders for audio-video and video-only generation
|
||||
- **`quantization/`**: FP8 quantization backends (FP8-TensorRT-LLM scaled MM, FP8 cast) for reduced memory footprint.
|
||||
@@ -55,6 +56,7 @@ pip install -e packages/ltx-core
|
||||
|
||||
- **Loader** ([`loader/`](src/ltx_core/loader/)): Model loading from `.safetensors`, LoRA fusion, weight remapping, and memory management
|
||||
- **Quantization** ([`quantization/`](src/ltx_core/quantization/)): FP8 quantization backends for reduced memory footprint and faster inference
|
||||
- **Block Streaming** ([`block_streaming/`](src/ltx_core/block_streaming/)): Streams transformer blocks through the GPU one block at a time, so the full model runs on machines without enough memory to hold all its weights at once
|
||||
|
||||
### Loader
|
||||
|
||||
@@ -125,39 +127,26 @@ Uses NVIDIA TensorRT-LLM's `cublas_scaled_mm` for efficient FP8 matrix multiplic
|
||||
**Usage with QuantizationPolicy:**
|
||||
|
||||
```python
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.quantization.fp8_scaled_mm import build_policy as build_fp8_scaled_mm_policy
|
||||
|
||||
# Dynamic input quantization (no calibration needed)
|
||||
policy = QuantizationPolicy.fp8_scaled_mm()
|
||||
|
||||
# Static input quantization with calibration file
|
||||
policy = QuantizationPolicy.fp8_scaled_mm(calibration_amax_path="/path/to/amax.json")
|
||||
# Discovers the layer set from the checkpoint's .weight_scale tensors
|
||||
policy = build_fp8_scaled_mm_policy("/path/to/checkpoint.safetensors")
|
||||
```
|
||||
|
||||
The policy provides `sd_ops` and `module_ops` that can be passed to the model builder:
|
||||
The policy carries `sd_ops`, `module_ops`, and `fuse_rule` that are passed to the model builder:
|
||||
|
||||
```python
|
||||
import torch
|
||||
from ltx_core.loader import SingleGPUModelBuilder
|
||||
|
||||
builder = SingleGPUModelBuilder(
|
||||
model=model,
|
||||
device=device,
|
||||
sd_ops=policy.sd_ops,
|
||||
model_class_configurator=MyModelConfigurator,
|
||||
model_path="/path/to/checkpoint.safetensors",
|
||||
model_sd_ops=policy.sd_ops,
|
||||
module_ops=policy.module_ops,
|
||||
fuse_rule=policy.fuse_rule,
|
||||
)
|
||||
builder.load(checkpoint_path)
|
||||
```
|
||||
|
||||
**Calibration File Format** (for static input quantization):
|
||||
|
||||
```json
|
||||
{
|
||||
"amax_values": {
|
||||
"transformer_blocks.0.attn.to_q.input_quantizer": 12.5,
|
||||
"transformer_blocks.0.attn.to_k.input_quantizer": 8.3,
|
||||
...
|
||||
}
|
||||
}
|
||||
model = builder.build(device=torch.device("cuda"))
|
||||
```
|
||||
|
||||
#### FP8 Cast
|
||||
@@ -165,7 +154,42 @@ builder.load(checkpoint_path)
|
||||
A simpler approach that casts weights to FP8 for storage and upcasts during inference:
|
||||
|
||||
```python
|
||||
policy = QuantizationPolicy.fp8_cast()
|
||||
from ltx_core.quantization.fp8_cast import build_policy as build_fp8_cast_policy
|
||||
|
||||
policy = build_fp8_cast_policy("/path/to/checkpoint.safetensors")
|
||||
```
|
||||
|
||||
### Block Streaming
|
||||
|
||||
The `block_streaming/` module ([`src/ltx_core/block_streaming/`](src/ltx_core/block_streaming/)) lets the full model run on machines that lack the memory to hold all of its weights at once. It streams the transformer's blocks through a small rolling set of GPU buffers, loading each block's weights just before it runs and recycling them afterwards, so only a few blocks are resident on the GPU at any moment. Construct it with `StreamingModelBuilder`, which returns a `BlockStreamingWrapper` -- an `nn.Module` drop-in for the wrapped model.
|
||||
|
||||
#### Strategies
|
||||
|
||||
The strategy is chosen automatically from `cpu_slots_count` relative to the number of blocks:
|
||||
|
||||
- **RAM streaming** (default, `cpu_slots_count` omitted or `>= num_blocks`): all blocks are pre-loaded into pinned CPU buffers (with LoRA fusion) at build time, then copied to the GPU on demand. Fast; higher CPU memory.
|
||||
- **Disk streaming** (`cpu_slots_count < num_blocks`): blocks are read from the `.safetensors` file on demand on a background worker thread. Slower; lowest CPU memory.
|
||||
|
||||
#### Basic usage
|
||||
|
||||
```python
|
||||
import torch
|
||||
from ltx_core.block_streaming import StreamingModelBuilder
|
||||
|
||||
builder = StreamingModelBuilder(
|
||||
model_class_configurator=MyModelConfigurator,
|
||||
model_path="/path/to/model.safetensors",
|
||||
blocks_attr="transformer_blocks", # dotted path to the nn.ModuleList
|
||||
blocks_prefix="transformer_blocks", # state-dict key prefix for block weights
|
||||
)
|
||||
|
||||
# Omit cpu_slots_count for RAM streaming; pass a value < num_blocks for disk streaming.
|
||||
model = builder.build(
|
||||
device=torch.device("cuda"),
|
||||
dtype=torch.bfloat16,
|
||||
cpu_slots_count=4,
|
||||
gpu_slots_count=2,
|
||||
)
|
||||
```
|
||||
|
||||
For complete, production-ready pipeline implementations that combine these building blocks, see the [`ltx-pipelines`](../ltx-pipelines/) package.
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "ltx-core"
|
||||
version = "1.1.3"
|
||||
version = "v1.1.6"
|
||||
description = "Core implementation of Lightricks' LTX-2 model"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -5,8 +5,9 @@ CPU-to-GPU copies, caching, and stream synchronization. Two weight
|
||||
source strategies are available:
|
||||
- **RAM streaming** (default): all blocks pre-loaded into pinned CPU
|
||||
buffers with LoRA fusion at build time. Fast, higher CPU memory.
|
||||
- **Disk streaming** (``cpu_slots < num_blocks``): blocks read from
|
||||
disk on demand with FIFO eviction. Slower, lower CPU memory.
|
||||
- **Disk streaming** (``cpu_slots < blocks_number``): blocks are read from
|
||||
disk on demand by a :class:`DiskWeightSource`, on a background worker
|
||||
thread. Slower, lower CPU memory.
|
||||
"""
|
||||
|
||||
from ltx_core.block_streaming.builder import DISK_CPU_SLOTS, StreamingModelBuilder
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
"""BlockFetcher: async disk reads on a worker thread."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.block_streaming.disk import DiskBlockReader
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
Buffer = dict[str, torch.Tensor]
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class FetchHandle:
|
||||
"""Caller-facing handle for an outstanding read returned by :meth:`BlockFetcher.submit`.
|
||||
Carries only what the caller needs: a completion event the worker sets and the
|
||||
read's error. The worker updates these once the read finishes.
|
||||
"""
|
||||
|
||||
_done: threading.Event
|
||||
_error: BaseException | None = None
|
||||
|
||||
def wait(self) -> BaseException | None:
|
||||
"""Block until the read finishes; return its error, or ``None`` on success."""
|
||||
self._done.wait()
|
||||
return self._error
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _ReadRequest:
|
||||
"""One outstanding read, internal to :class:`BlockFetcher`.
|
||||
The fetcher's worker reads block ``idx`` into the caller-carved ``buffer`` and
|
||||
updates ``handle`` (its error, then its event) once the read has finished.
|
||||
"""
|
||||
|
||||
idx: int
|
||||
buffer: Buffer
|
||||
handle: FetchHandle
|
||||
|
||||
|
||||
class BlockFetcher:
|
||||
"""Fills caller-supplied buffers on a worker thread."""
|
||||
|
||||
def __init__(self, reader: DiskBlockReader) -> None:
|
||||
self._reader = reader
|
||||
self._request_queue: queue.SimpleQueue[_ReadRequest | None] = queue.SimpleQueue()
|
||||
self._worker = threading.Thread(target=self._run, name="BlockFetcher-IO", daemon=True)
|
||||
self._worker.start()
|
||||
|
||||
def submit(self, idx: int, buffer: Buffer) -> FetchHandle:
|
||||
"""Enqueue a read of block *idx* into the caller-carved *buffer*, return its handle."""
|
||||
handle = FetchHandle(_done=threading.Event())
|
||||
request = _ReadRequest(idx=idx, buffer=buffer, handle=handle)
|
||||
self._request_queue.put(request)
|
||||
return handle
|
||||
|
||||
def cleanup(self) -> None:
|
||||
"""Drain pending reads, join the worker, close the reader."""
|
||||
self._request_queue.put(None)
|
||||
self._worker.join()
|
||||
self._reader.cleanup()
|
||||
|
||||
def _run(self) -> None:
|
||||
# Pinned buffers are allocated under the caller's inference_mode, so
|
||||
# in-place copy_ from this thread requires inference_mode here too.
|
||||
with torch.inference_mode():
|
||||
while True:
|
||||
request = self._request_queue.get()
|
||||
if request is None:
|
||||
return
|
||||
|
||||
try:
|
||||
self._reader.read_into(request.buffer, request.idx)
|
||||
except Exception as exc:
|
||||
logger.exception("BlockFetcher: fetch failed for item %d", request.idx)
|
||||
request.handle._error = exc
|
||||
request.handle._done.set()
|
||||
@@ -2,46 +2,61 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field, replace
|
||||
from typing import Generic
|
||||
from dataclasses import replace
|
||||
from typing import TYPE_CHECKING, Final, Generic
|
||||
|
||||
import safetensors
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ltx_core.block_streaming import utils as bs_utils
|
||||
from ltx_core.block_streaming.block_fetcher import BlockFetcher
|
||||
from ltx_core.block_streaming.disk import DiskBlockReader, DiskTensorReader, LoraSource
|
||||
from ltx_core.block_streaming.pool import WeightPool
|
||||
from ltx_core.block_streaming.pool import BufferPool
|
||||
from ltx_core.block_streaming.provider import WeightsProvider
|
||||
from ltx_core.block_streaming.source import DiskWeightSource, PinnedWeightSource, WeightSource
|
||||
from ltx_core.block_streaming.utils import allocate_layout_views, derive_layout, make_block_key, resolve_attr
|
||||
from ltx_core.block_streaming.source import DiskWeightSource, PinnedBlock, PinnedWeightSource, WeightSource
|
||||
from ltx_core.block_streaming.utils import (
|
||||
carve_buffer,
|
||||
derive_layout,
|
||||
layout_nbytes,
|
||||
make_block_key,
|
||||
resolve_attr,
|
||||
)
|
||||
from ltx_core.block_streaming.wrapper import BlockStreamingWrapper
|
||||
from ltx_core.loader.fuse_loras import aggregate_lora_products, fuse_lora_weights
|
||||
from ltx_core.loader.fuse_loras import FuseRule, bf16_fuse_rule, fuse_lora_weights
|
||||
from ltx_core.loader.helpers import create_meta_model, load_state_dict, read_model_config
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import (
|
||||
LoraPathStrengthAndSDOps,
|
||||
LoraStateDictWithStrength,
|
||||
ModelBuilderProtocol,
|
||||
StateDict,
|
||||
StateDictLoader,
|
||||
TensorLayout,
|
||||
)
|
||||
from ltx_core.loader.registry import DummyRegistry, Registry
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
from ltx_core.model.model_protocol import ModelConfigurator, ModelType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from typing_extensions import Self
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DISK_CPU_SLOTS = 2
|
||||
_DEFAULT_GPU_SLOTS = 2
|
||||
_PREFETCH_DEPTH = 2
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType]):
|
||||
"""Immutable builder for :class:`BlockStreamingWrapper`.
|
||||
Reads block weights from safetensors on demand. ``cpu_slots`` and
|
||||
``gpu_slots`` control the memory/speed trade-off (see :meth:`build`).
|
||||
The builder is immutable (``with_*`` return modified copies) and exposes
|
||||
its state via read-only properties backed by private attributes.
|
||||
Args:
|
||||
model_class_configurator: Creates the model from a config dict.
|
||||
model_path: One or more ``.safetensors`` checkpoint paths.
|
||||
@@ -50,38 +65,108 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
|
||||
loras: LoRA adapters fused into weights at load time.
|
||||
model_loader: Strategy for reading checkpoint metadata.
|
||||
registry: Shared cache for loaded state dicts.
|
||||
fuse_rule: Per-policy LoRA merge rule. Defaults to ``bf16_fuse_rule``;
|
||||
use ``fp8_cast_fuse_rule`` for fp8_cast streaming so the pinned
|
||||
buffers receive correctly-quantized weights.
|
||||
blocks_attr: Dotted path to the ``nn.ModuleList`` (e.g.
|
||||
``"velocity_model.transformer_blocks"``).
|
||||
``"transformer_blocks"``).
|
||||
blocks_prefix: State-dict key prefix for block weights
|
||||
(e.g. ``"transformer_blocks"``).
|
||||
state_dict_prefix: Wrapper offset prepended to keys when loading into
|
||||
the meta model (e.g. ``"velocity_model."`` when wrapped by ``X0Model``).
|
||||
model_wrapper: Optional callable wrapping the model
|
||||
(e.g. ``X0Model``).
|
||||
"""
|
||||
|
||||
model_class_configurator: type[ModelConfigurator[ModelType]]
|
||||
model_path: str | tuple[str, ...]
|
||||
model_sd_ops: SDOps | None = None
|
||||
module_ops: tuple[ModuleOps, ...] = field(default_factory=tuple)
|
||||
loras: tuple[LoraPathStrengthAndSDOps, ...] = field(default_factory=tuple)
|
||||
model_loader: StateDictLoader = field(default_factory=SafetensorsModelStateDictLoader)
|
||||
registry: Registry = field(default_factory=DummyRegistry)
|
||||
def __init__(
|
||||
self,
|
||||
model_class_configurator: type[ModelConfigurator[ModelType]],
|
||||
model_path: str | tuple[str, ...],
|
||||
model_sd_ops: SDOps | None = None,
|
||||
module_ops: tuple[ModuleOps, ...] = (),
|
||||
loras: tuple[LoraPathStrengthAndSDOps, ...] = (),
|
||||
model_loader: StateDictLoader | None = None,
|
||||
registry: Registry | None = None,
|
||||
fuse_rule: FuseRule = bf16_fuse_rule,
|
||||
blocks_attr: str = "",
|
||||
blocks_prefix: str = "",
|
||||
) -> None:
|
||||
# Read-only: typed with the covariant ModelType, so it must not be a mutable attribute.
|
||||
self._model_class_configurator: Final = model_class_configurator
|
||||
self._model_path = model_path
|
||||
self._model_sd_ops = model_sd_ops
|
||||
self._module_ops = module_ops
|
||||
self._loras = loras
|
||||
self._model_loader = model_loader if model_loader is not None else SafetensorsModelStateDictLoader()
|
||||
self._registry = registry if registry is not None else DummyRegistry()
|
||||
self._fuse_rule = fuse_rule
|
||||
self._blocks_attr = blocks_attr
|
||||
self._blocks_prefix = blocks_prefix
|
||||
|
||||
# Streaming-specific
|
||||
blocks_attr: str = ""
|
||||
blocks_prefix: str = ""
|
||||
state_dict_prefix: str = ""
|
||||
model_wrapper: Callable[[ModelType], nn.Module] | None = None
|
||||
@property
|
||||
def model_class_configurator(self) -> type[ModelConfigurator[ModelType]]:
|
||||
return self._model_class_configurator
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> StreamingModelBuilder:
|
||||
return replace(self, model_sd_ops=sd_ops)
|
||||
@property
|
||||
def model_path(self) -> str | tuple[str, ...]:
|
||||
return self._model_path
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> StreamingModelBuilder:
|
||||
return replace(self, module_ops=module_ops)
|
||||
@property
|
||||
def model_sd_ops(self) -> SDOps | None:
|
||||
return self._model_sd_ops
|
||||
|
||||
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> StreamingModelBuilder:
|
||||
return replace(self, loras=loras)
|
||||
@property
|
||||
def module_ops(self) -> tuple[ModuleOps, ...]:
|
||||
return self._module_ops
|
||||
|
||||
@property
|
||||
def loras(self) -> tuple[LoraPathStrengthAndSDOps, ...]:
|
||||
return self._loras
|
||||
|
||||
@property
|
||||
def model_loader(self) -> StateDictLoader:
|
||||
return self._model_loader
|
||||
|
||||
@property
|
||||
def registry(self) -> Registry:
|
||||
return self._registry
|
||||
|
||||
@property
|
||||
def fuse_rule(self) -> FuseRule:
|
||||
return self._fuse_rule
|
||||
|
||||
@property
|
||||
def blocks_attr(self) -> str:
|
||||
return self._blocks_attr
|
||||
|
||||
@property
|
||||
def blocks_prefix(self) -> str:
|
||||
return self._blocks_prefix
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._model_sd_ops = sd_ops
|
||||
return clone
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._module_ops = module_ops
|
||||
return clone
|
||||
|
||||
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._loras = loras
|
||||
return clone
|
||||
|
||||
def with_registry(self, registry: Registry) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._registry = registry
|
||||
return clone
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> Self:
|
||||
# Streaming fuses LoRAs into pinned CPU buffers; no other staging device is meaningful.
|
||||
raise NotImplementedError("StreamingModelBuilder loads LoRA weights on CPU only.")
|
||||
|
||||
def with_fuse_rule(self, fuse_rule: FuseRule) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._fuse_rule = fuse_rule
|
||||
return clone
|
||||
|
||||
def model_config(self) -> dict:
|
||||
"""Read model configuration from the checkpoint metadata."""
|
||||
@@ -93,16 +178,16 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
|
||||
|
||||
def build(
|
||||
self,
|
||||
target_device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype | None = None,
|
||||
cpu_slots_count: int | None = None,
|
||||
gpu_slots_count: int | None = None,
|
||||
**_kwargs: object,
|
||||
) -> BlockStreamingWrapper:
|
||||
"""Build and return a ready-to-use :class:`BlockStreamingWrapper`.
|
||||
Args:
|
||||
target_device: GPU device for compute.
|
||||
dtype: Weight dtype (e.g. ``torch.bfloat16``).
|
||||
device: GPU device for compute. ``None`` defaults to ``cuda``.
|
||||
dtype: Weight dtype (e.g. ``torch.bfloat16``). Required.
|
||||
cpu_slots_count: Number of pinned CPU buffer slots.
|
||||
``None`` = RAM streaming (all blocks pre-loaded with LoRA fusion).
|
||||
gpu_slots_count: Number of GPU buffer slots.
|
||||
@@ -110,61 +195,71 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
|
||||
"""
|
||||
if not self.blocks_prefix:
|
||||
raise ValueError("blocks_prefix must be non-empty for streaming")
|
||||
if dtype is None:
|
||||
raise ValueError("StreamingModelBuilder.build requires an explicit dtype")
|
||||
device = device if device is not None else torch.device("cuda")
|
||||
|
||||
config = read_model_config(self.model_path, self.model_loader)
|
||||
meta_model: nn.Module = create_meta_model(self.model_class_configurator, config, self.module_ops)
|
||||
if self.model_wrapper is not None:
|
||||
meta_model = self.model_wrapper(meta_model)
|
||||
meta_model.eval()
|
||||
|
||||
blocks = resolve_attr(meta_model, self.blocks_attr)
|
||||
|
||||
checkpoint_paths = list(self.model_path) if isinstance(self.model_path, tuple) else [self.model_path]
|
||||
block_key_map, non_block_keys = _scan_checkpoint_keys(checkpoint_paths, self.model_sd_ops, self.blocks_prefix)
|
||||
expected_indices = set(range(len(blocks)))
|
||||
if set(block_key_map) != expected_indices:
|
||||
missing = sorted(expected_indices - set(block_key_map))
|
||||
extra = sorted(set(block_key_map) - expected_indices)
|
||||
raise ValueError(
|
||||
f"Block weights under prefix '{self.blocks_prefix}.' do not match the {len(blocks)} model blocks: "
|
||||
f"missing indices {missing}, unexpected indices {extra}"
|
||||
)
|
||||
|
||||
cpu_slots_count = cpu_slots_count if cpu_slots_count is not None else len(blocks)
|
||||
gpu_slots_count = gpu_slots_count if gpu_slots_count is not None else _DEFAULT_GPU_SLOTS
|
||||
|
||||
if cpu_slots_count >= len(blocks):
|
||||
lora_sd_and_strengths = self._load_lora_sds()
|
||||
source, lora_sources = self._build_pinned_source(
|
||||
meta_model, target_device, dtype, cpu_slots_count, block_key_map, non_block_keys
|
||||
blocks, dtype, cpu_slots_count, block_key_map, lora_sd_and_strengths
|
||||
)
|
||||
non_block_loras = lora_sd_and_strengths
|
||||
else:
|
||||
reader = DiskTensorReader(checkpoint_paths)
|
||||
source, lora_sources = self._build_disk_source(
|
||||
meta_model, target_device, dtype, cpu_slots_count, reader, block_key_map, non_block_keys
|
||||
blocks, dtype, cpu_slots_count, reader, block_key_map, prefetch_depth=_PREFETCH_DEPTH
|
||||
)
|
||||
non_block_loras = [src.as_state_dict_with_strength() for src in lora_sources]
|
||||
|
||||
copy_stream = torch.cuda.Stream(device=target_device)
|
||||
gpu_pool = WeightPool(
|
||||
source.block_layout,
|
||||
self._load_non_block_weights(meta_model, non_block_keys, device, dtype, non_block_loras)
|
||||
|
||||
copy_stream = torch.cuda.Stream(device=device)
|
||||
gpu_pool = BufferPool(
|
||||
source.slot_nbytes,
|
||||
gpu_slots_count,
|
||||
target_device,
|
||||
device,
|
||||
reuse_barrier=lambda event: copy_stream.wait_event(event),
|
||||
)
|
||||
provider = WeightsProvider(gpu_pool, copy_stream, target_device, source, lora_sources, self.blocks_prefix)
|
||||
provider = WeightsProvider(
|
||||
gpu_pool,
|
||||
copy_stream,
|
||||
device,
|
||||
source,
|
||||
lora_sources,
|
||||
self.blocks_prefix,
|
||||
fuse_rule=self.fuse_rule,
|
||||
)
|
||||
return BlockStreamingWrapper(
|
||||
model=meta_model,
|
||||
blocks=blocks,
|
||||
provider=provider,
|
||||
target_device=target_device,
|
||||
target_device=device,
|
||||
)
|
||||
|
||||
def _build_pinned_source(
|
||||
self,
|
||||
meta_model: nn.Module,
|
||||
target_device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
cpu_slots_count: int,
|
||||
block_key_map: dict[int, list[tuple[str, str]]],
|
||||
non_block_keys: list[tuple[str, str]],
|
||||
) -> tuple[WeightSource, list[LoraSource]]:
|
||||
"""Pre-load all blocks into pinned CPU buffers with LoRA fusion."""
|
||||
model_sd = load_state_dict(
|
||||
self.model_path, self.model_loader, self.registry, torch.device("cpu"), self.model_sd_ops
|
||||
)
|
||||
|
||||
lora_sd_and_strengths = [
|
||||
def _load_lora_sds(self) -> list[LoraStateDictWithStrength]:
|
||||
"""Load each configured LoRA into a state dict for fusion (pinned path)."""
|
||||
return [
|
||||
LoraStateDictWithStrength(
|
||||
load_state_dict([lora.path], self.model_loader, self.registry, torch.device("cpu"), lora.sd_ops),
|
||||
lora.strength,
|
||||
@@ -172,6 +267,25 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
|
||||
for lora in self.loras
|
||||
]
|
||||
|
||||
def _filtered_sd_ops(self, name_suffix: str, allowed_model_keys: frozenset[str]) -> SDOps:
|
||||
"""``model_sd_ops`` restricted to *allowed_model_keys* (post-rename keys).
|
||||
The loader skips keys filtered to None before reading them, so a restricted
|
||||
load never materializes the excluded partition. The distinct ``name`` avoids
|
||||
a registry cache-id collision with the other partition.
|
||||
"""
|
||||
base = self.model_sd_ops if self.model_sd_ops is not None else SDOps("streaming").with_matching()
|
||||
allowed = allowed_model_keys if base.allowed_keys is None else (allowed_model_keys & base.allowed_keys)
|
||||
return replace(base, name=f"{base.name}__{name_suffix}", allowed_keys=allowed)
|
||||
|
||||
def _build_pinned_source(
|
||||
self,
|
||||
blocks: nn.ModuleList,
|
||||
dtype: torch.dtype,
|
||||
cpu_slots_count: int,
|
||||
block_key_map: dict[int, list[tuple[str, str]]],
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
) -> tuple[WeightSource, list[LoraSource]]:
|
||||
"""Pre-load each block into its own contiguous pinned CPU buffer with LoRA fusion."""
|
||||
for block_idx in block_key_map:
|
||||
if block_idx >= cpu_slots_count:
|
||||
raise ValueError(
|
||||
@@ -179,85 +293,75 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
|
||||
f"got block index {block_idx} with only {cpu_slots_count} slots."
|
||||
)
|
||||
|
||||
blocks = resolve_attr(meta_model, self.blocks_attr)
|
||||
block_tensors: dict[str, torch.Tensor] = {}
|
||||
# One contiguous pinned buffer per block, carved into per-param views. The
|
||||
# views (flattened by full key) are filled in place; the source then keeps
|
||||
# only the contiguous buffer and the layout to re-carve it on read.
|
||||
pinned_buffers: dict[int, torch.Tensor] = {}
|
||||
block_layouts: dict[int, TensorLayout] = {}
|
||||
fill_views: dict[str, torch.Tensor] = {}
|
||||
for block_idx, entries in block_key_map.items():
|
||||
block_params = dict(blocks[block_idx].named_parameters())
|
||||
for _sft_key, param_name in entries:
|
||||
key = make_block_key(self.blocks_prefix, block_idx, param_name)
|
||||
block_tensors[key] = block_params[param_name]
|
||||
blocks_layout = derive_layout(block_tensors, dtype)
|
||||
pinned_blocks = allocate_layout_views(blocks_layout, pin_memory=True)
|
||||
block_state = _block_state(blocks[block_idx])
|
||||
layout = derive_layout({param_name: block_state[param_name] for _sft_key, param_name in entries}, dtype)
|
||||
buffer = bs_utils.alloc_buffer(layout_nbytes(layout), torch.device("cpu"), pin_memory=True)
|
||||
views = carve_buffer(buffer, layout)
|
||||
pinned_buffers[block_idx] = buffer
|
||||
block_layouts[block_idx] = layout
|
||||
for param_name, view in views.items():
|
||||
fill_views[make_block_key(self.blocks_prefix, block_idx, param_name)] = view
|
||||
|
||||
block_sd = load_state_dict(
|
||||
self.model_path,
|
||||
self.model_loader,
|
||||
self.registry,
|
||||
torch.device("cpu"),
|
||||
self._filtered_sd_ops("blocks", frozenset(fill_views)),
|
||||
)
|
||||
|
||||
should_sync = False
|
||||
for key, fused in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype=None, preserve_input_device=False):
|
||||
if key in pinned_blocks:
|
||||
pinned_blocks[key].copy_(fused, non_blocking=True)
|
||||
model_sd.sd[key] = None
|
||||
should_sync = True
|
||||
else:
|
||||
model_sd.sd[key] = fused
|
||||
for key, fused in fuse_lora_weights(
|
||||
block_sd, lora_sd_and_strengths, fuse_rule=self.fuse_rule, preserve_input_device=False
|
||||
):
|
||||
if key not in fill_views:
|
||||
raise ValueError(f"Block-restricted load produced {key!r}, which is not a pinned block weight")
|
||||
fill_views[key].copy_(fused, non_blocking=True)
|
||||
block_sd.sd[key] = None
|
||||
should_sync = True
|
||||
if should_sync:
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Fill remaining pinned keys from the source state dict.
|
||||
for key in blocks_layout:
|
||||
if model_sd.sd[key] is None:
|
||||
for key, view in fill_views.items():
|
||||
if block_sd.sd[key] is None:
|
||||
continue
|
||||
pinned_blocks[key].copy_(model_sd.sd[key])
|
||||
model_sd.sd[key] = None
|
||||
|
||||
pinned: dict[int, dict[str, torch.Tensor]] = {
|
||||
block_idx: {
|
||||
param_name: pinned_blocks[make_block_key(self.blocks_prefix, block_idx, param_name)]
|
||||
for _sft_key, param_name in entries
|
||||
}
|
||||
for block_idx, entries in block_key_map.items()
|
||||
}
|
||||
|
||||
non_block_sd: dict[str, torch.Tensor] = {
|
||||
self.state_dict_prefix + model_key: model_sd.sd[model_key].to(device=target_device, dtype=dtype)
|
||||
for _sft_key, model_key in non_block_keys
|
||||
}
|
||||
|
||||
meta_model.load_state_dict(non_block_sd, strict=False, assign=True)
|
||||
view.copy_(block_sd.sd[key])
|
||||
block_sd.sd[key] = None
|
||||
|
||||
pinned = {idx: PinnedBlock(pinned_buffers[idx], block_layouts[idx]) for idx in pinned_buffers}
|
||||
return PinnedWeightSource(pinned), []
|
||||
|
||||
def _build_disk_source(
|
||||
self,
|
||||
meta_model: nn.Module,
|
||||
target_device: torch.device,
|
||||
blocks: nn.ModuleList,
|
||||
dtype: torch.dtype,
|
||||
cpu_slots_count: int,
|
||||
reader: DiskTensorReader,
|
||||
block_key_map: dict[int, list[tuple[str, str]]],
|
||||
non_block_keys: list[tuple[str, str]],
|
||||
prefetch_depth: int,
|
||||
) -> tuple[WeightSource, list[LoraSource]]:
|
||||
"""Create a DiskWeightSource backed by a DiskBlockReader for lazy loading.
|
||||
Derives the shared pool layout from the meta model's block 0 — this
|
||||
relies on module_ops (e.g. fp8_cast) leaving the meta param dtype in
|
||||
sync with the post-sd_ops checkpoint dtype.
|
||||
"""Create a DiskWeightSource backed by a DiskBlockReader.
|
||||
Pool slots are sized to the largest block and carved per block on read, so
|
||||
heterogeneous blocks (e.g. layers with differing attention layouts) share
|
||||
one pool. Pool capacity is ``cpu_slots_count + prefetch_depth`` so the
|
||||
lookahead loop in ``DiskWeightSource.get`` never evicts its own target.
|
||||
Layouts come from the meta model; assumes module_ops keep the meta param
|
||||
dtype in sync with the post-sd_ops checkpoint dtype.
|
||||
"""
|
||||
lora_sources = [LoraSource(lora.path, lora.sd_ops, lora.strength) for lora in self.loras]
|
||||
block_layouts = _block_layouts(blocks, block_key_map, dtype)
|
||||
slot_nbytes = max(layout_nbytes(layout) for layout in block_layouts.values())
|
||||
|
||||
self._load_non_block_weights(
|
||||
reader,
|
||||
non_block_keys,
|
||||
meta_model,
|
||||
target_device,
|
||||
dtype,
|
||||
sd_ops=self.model_sd_ops,
|
||||
key_prefix=self.state_dict_prefix,
|
||||
lora_sources=lora_sources,
|
||||
)
|
||||
|
||||
blocks = resolve_attr(meta_model, self.blocks_attr)
|
||||
layout = derive_layout(dict(blocks[0].named_parameters()), dtype)
|
||||
|
||||
cpu_pool = WeightPool(
|
||||
layout,
|
||||
cpu_slots_count,
|
||||
cpu_pool = BufferPool(
|
||||
slot_nbytes,
|
||||
cpu_slots_count + prefetch_depth,
|
||||
torch.device("cpu"),
|
||||
reuse_barrier=lambda event: event.synchronize(),
|
||||
pin_memory=True,
|
||||
@@ -268,55 +372,76 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
|
||||
sd_ops=self.model_sd_ops,
|
||||
blocks_prefix=self.blocks_prefix,
|
||||
)
|
||||
source = DiskWeightSource(cpu_pool, block_reader)
|
||||
fetcher = BlockFetcher(block_reader)
|
||||
source = DiskWeightSource(
|
||||
cpu_pool,
|
||||
fetcher,
|
||||
block_layouts,
|
||||
blocks_number=len(blocks),
|
||||
prefetch_depth=prefetch_depth,
|
||||
)
|
||||
lora_sources = [LoraSource(lora.path, lora.sd_ops, lora.strength) for lora in self.loras]
|
||||
|
||||
return source, lora_sources
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _fuse_lora_delta(
|
||||
model_key: str,
|
||||
tensor: torch.Tensor,
|
||||
lora_sources: list[LoraSource],
|
||||
) -> torch.Tensor:
|
||||
"""Add all matching LoRA deltas to *tensor* in-place via ``addmm_``."""
|
||||
if not lora_sources or not model_key.endswith(".weight"):
|
||||
return tensor
|
||||
prefix = model_key[: -len(".weight")]
|
||||
products = (
|
||||
ab
|
||||
for ab in (s.get_ab(prefix, device=tensor.device, dtype=tensor.dtype) for s in lora_sources)
|
||||
if ab is not None
|
||||
)
|
||||
aggregate_lora_products(products, out=tensor)
|
||||
return tensor
|
||||
|
||||
@staticmethod
|
||||
@torch.inference_mode()
|
||||
def _load_non_block_weights(
|
||||
reader: DiskTensorReader,
|
||||
non_block_keys: list[tuple[str, str]],
|
||||
self,
|
||||
model: nn.Module,
|
||||
non_block_keys: list[tuple[str, str]],
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
sd_ops: SDOps | None = None,
|
||||
key_prefix: str = "",
|
||||
lora_sources: list[LoraSource] | None = None,
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
) -> None:
|
||||
"""Load non-block weights into *model* on *device*."""
|
||||
state_dict: dict[str, torch.Tensor] = {}
|
||||
sources = lora_sources or []
|
||||
for sft_key, model_key in non_block_keys:
|
||||
tensor = reader.get_tensor(sft_key).to(device=device, dtype=dtype)
|
||||
tensor = StreamingModelBuilder._fuse_lora_delta(model_key, tensor, sources)
|
||||
if sd_ops is not None:
|
||||
for kv in sd_ops.apply_to_key_value(model_key, tensor):
|
||||
state_dict[key_prefix + kv.new_key] = kv.new_value
|
||||
continue
|
||||
state_dict[key_prefix + model_key] = tensor
|
||||
model.load_state_dict(state_dict, strict=False, assign=True)
|
||||
"""Load the non-block weights onto *device* and fuse LoRAs -- both paths.
|
||||
Reads through the loader with ``model_sd_ops`` restricted to the
|
||||
non-block keys, so ``sd_ops`` (incl. kv-ops such as Gemma's ``lm_head``
|
||||
duplication) is applied exactly once and block tensors are never read.
|
||||
"""
|
||||
non_block_sd_ops = self._filtered_sd_ops("non_block", frozenset(mk for _sft_key, mk in non_block_keys))
|
||||
loaded = load_state_dict(self.model_path, self.model_loader, self.registry, device, non_block_sd_ops)
|
||||
non_block_sd = {key: tensor.to(dtype=dtype) for key, tensor in loaded.sd.items()}
|
||||
|
||||
if lora_sd_and_strengths:
|
||||
non_block_state = StateDict(sd=non_block_sd, device=device, size=0, dtype={dtype})
|
||||
for key, fused in fuse_lora_weights(
|
||||
non_block_state,
|
||||
lora_sd_and_strengths,
|
||||
fuse_rule=self.fuse_rule,
|
||||
preserve_input_device=True,
|
||||
):
|
||||
non_block_sd[key] = fused
|
||||
|
||||
model.load_state_dict(non_block_sd, strict=False, assign=True)
|
||||
|
||||
|
||||
def _block_state(block: nn.Module) -> dict[str, torch.Tensor]:
|
||||
"""Streamed-eligible tensors of a block: parameters then buffers.
|
||||
Block streaming swaps both params and checkpoint-backed buffers (e.g. Gemma4's
|
||||
per-layer ``layer_scalar``), so the layout, pinned packing, and meta-ordering
|
||||
all consult parameters and buffers together. Non-checkpoint (computed) buffers
|
||||
are harmless here -- only keys present in ``block_key_map`` are ever streamed.
|
||||
"""
|
||||
return {**dict(block.named_parameters()), **dict(block.named_buffers())}
|
||||
|
||||
|
||||
def _block_layouts(
|
||||
blocks: nn.ModuleList,
|
||||
block_key_map: dict[int, list[tuple[str, str]]],
|
||||
dtype: torch.dtype,
|
||||
) -> dict[int, TensorLayout]:
|
||||
"""Per-block layout of the streamed tensors, taken from the meta model.
|
||||
Blocks may differ in shape and even in which tensors they have (e.g. Gemma4's
|
||||
full-attention layers drop ``v_proj``), so each block gets its own layout in
|
||||
``block_key_map`` order. The pinned packing, the disk reader, and the GPU carve
|
||||
all key off this same per-block layout, so the provider's contiguous H2D copy
|
||||
is valid for any entry order (no cross-block ordering required).
|
||||
"""
|
||||
layouts: dict[int, TensorLayout] = {}
|
||||
for idx, entries in block_key_map.items():
|
||||
state = _block_state(blocks[idx])
|
||||
layouts[idx] = derive_layout({param_name: state[param_name] for _sft_key, param_name in entries}, dtype)
|
||||
return layouts
|
||||
|
||||
|
||||
def _scan_checkpoint_keys(
|
||||
|
||||
@@ -9,6 +9,7 @@ import torch
|
||||
|
||||
from ltx_core.block_streaming.utils import allocate_layout_views, make_block_key
|
||||
from ltx_core.loader.fuse_loras import LoraProduct
|
||||
from ltx_core.loader.primitives import LoraStateDictWithStrength, StateDict
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
|
||||
_SAFETENSORS_DTYPE_TO_TORCH: dict[str, torch.dtype] = {
|
||||
@@ -125,6 +126,22 @@ class LoraSource:
|
||||
b_view.copy_(handle.get_tensor(b_keys[prefix]))
|
||||
self._pinned_ab[prefix] = (a_view, b_view)
|
||||
|
||||
def as_state_dict_with_strength(self) -> LoraStateDictWithStrength:
|
||||
"""Return a :class:`LoraStateDictWithStrength` view of the pinned A/B factors.
|
||||
Lets non-block fusion consume the already-loaded disk-streaming LoRA
|
||||
without re-reading the safetensors file or re-applying ``sd_ops``.
|
||||
"""
|
||||
sd: dict[str, torch.Tensor] = {}
|
||||
for prefix, (a, b) in self._pinned_ab.items():
|
||||
sd[f"{prefix}.lora_A.weight"] = a
|
||||
sd[f"{prefix}.lora_B.weight"] = b
|
||||
size = sum(t.numel() * t.element_size() for t in sd.values())
|
||||
dtypes = {t.dtype for t in sd.values()}
|
||||
return LoraStateDictWithStrength(
|
||||
StateDict(sd=sd, device=torch.device("cpu"), size=size, dtype=dtypes),
|
||||
self.strength,
|
||||
)
|
||||
|
||||
def get_ab(
|
||||
self,
|
||||
param_prefix: str,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""Weight buffer pool for block streaming."""
|
||||
"""Raw buffer pool for block streaming."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -7,69 +7,64 @@ from typing import Callable
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.block_streaming.utils import allocate_layout_views
|
||||
from ltx_core.loader.primitives import TensorLayout
|
||||
from ltx_core.block_streaming import utils
|
||||
|
||||
|
||||
class WeightPool:
|
||||
"""Fixed pool of pre-allocated weight buffers with event-based reuse.
|
||||
All slots share a single buffer (CPU or GPU); each slot is a
|
||||
contiguous slice carved out of it via :func:`allocate_layout_views`.
|
||||
class BufferPool:
|
||||
"""Fixed pool of pre-allocated raw buffer slots with event-based reuse.
|
||||
Slots are carved from a single contiguous ``uint8`` buffer; each is
|
||||
``slot_nbytes`` long and handed out as a raw 1-D ``uint8`` tensor.
|
||||
Args:
|
||||
buffer_layout: ``{name: (shape, dtype)}`` for each buffer.
|
||||
capacity: Number of buffers to pre-allocate.
|
||||
slot_nbytes: Byte size of each slot.
|
||||
capacity: Number of slots to pre-allocate.
|
||||
device: Device for allocation.
|
||||
reuse_barrier: Called with the pending event before a buffer is reused.
|
||||
reuse_barrier: Called with the pending event before a slot is reused.
|
||||
pin_memory: Pin buffers (for async H2D copies from CPU).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
buffer_layout: TensorLayout,
|
||||
slot_nbytes: int,
|
||||
capacity: int,
|
||||
device: torch.device,
|
||||
reuse_barrier: Callable[[torch.cuda.Event], None],
|
||||
pin_memory: bool = False,
|
||||
) -> None:
|
||||
self._buffer_layout = buffer_layout
|
||||
self._slot_nbytes = slot_nbytes
|
||||
self._capacity = capacity
|
||||
self._free: deque[dict[str, torch.Tensor]] = deque()
|
||||
self._free: deque[torch.Tensor] = deque()
|
||||
self._events: dict[int, torch.cuda.Event] = {}
|
||||
self._reuse_barrier = reuse_barrier
|
||||
memory_layout = {
|
||||
_make_key(slot, name): (shape, dtype)
|
||||
for slot in range(capacity)
|
||||
for name, (shape, dtype) in buffer_layout.items()
|
||||
}
|
||||
all_views = allocate_layout_views(memory_layout, device=device, pin_memory=pin_memory)
|
||||
buffer = utils.alloc_buffer(max(slot_nbytes * capacity, 1), device, pin_memory)
|
||||
for slot in range(capacity):
|
||||
self._free.append({name: all_views[_make_key(slot, name)] for name in buffer_layout})
|
||||
self._free.append(buffer[slot * slot_nbytes : (slot + 1) * slot_nbytes])
|
||||
|
||||
@property
|
||||
def capacity(self) -> int:
|
||||
return self._capacity
|
||||
|
||||
@property
|
||||
def buffer_layout(self) -> TensorLayout:
|
||||
return self._buffer_layout
|
||||
def slot_nbytes(self) -> int:
|
||||
return self._slot_nbytes
|
||||
|
||||
def acquire(self) -> dict[str, torch.Tensor]:
|
||||
"""Take a free buffer, waiting any pending event before returning."""
|
||||
weights = self._free.popleft()
|
||||
event = self._events.pop(id(weights), None)
|
||||
def acquire(self) -> torch.Tensor:
|
||||
"""Take a free raw slot, waiting any pending event before returning.
|
||||
Raises :class:`RuntimeError` if every slot is currently in use.
|
||||
"""
|
||||
if not self._free:
|
||||
raise RuntimeError(f"BufferPool exhausted: all {self._capacity} buffers are in use")
|
||||
buffer = self._free.popleft()
|
||||
event = self._events.pop(id(buffer), None)
|
||||
if event is not None:
|
||||
self._reuse_barrier(event)
|
||||
return weights
|
||||
return buffer
|
||||
|
||||
def release(self, weights: dict[str, torch.Tensor], event: torch.cuda.Event | None = None) -> None:
|
||||
"""Return a buffer to the free list.
|
||||
If *event* is given it is waited on the next :meth:`acquire`
|
||||
of this buffer, ensuring the prior operation has completed.
|
||||
def release(self, buffer: torch.Tensor, event: torch.cuda.Event | None = None) -> None:
|
||||
"""Return a raw slot to the free list.
|
||||
The *buffer* must be the exact tensor object returned by :meth:`acquire`
|
||||
(reuse is keyed on its identity). If *event* is given it is waited on the
|
||||
next :meth:`acquire` of this slot, ensuring the prior operation finished.
|
||||
"""
|
||||
if event is not None:
|
||||
self._events[id(weights)] = event
|
||||
self._free.append(weights)
|
||||
|
||||
|
||||
def _make_key(slot: int, name: str) -> str:
|
||||
return f"{slot}/{name}"
|
||||
self._events[id(buffer)] = event
|
||||
self._free.append(buffer)
|
||||
|
||||
@@ -3,35 +3,28 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import OrderedDict
|
||||
from typing import NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.block_streaming.disk import LoraSource
|
||||
from ltx_core.block_streaming.pool import WeightPool
|
||||
from ltx_core.block_streaming.pool import BufferPool
|
||||
from ltx_core.block_streaming.source import WeightSource
|
||||
from ltx_core.block_streaming.utils import FP8_DTYPES
|
||||
from ltx_core.loader.fuse_loras import aggregate_lora_products, fuse_cast_fp8_weight
|
||||
from ltx_core.block_streaming.utils import carve_buffer, layout_nbytes
|
||||
from ltx_core.loader.fuse_loras import FuseRule, aggregate_lora_products, bf16_fuse_rule
|
||||
from ltx_core.loader.primitives import StateDict
|
||||
|
||||
_EMPTY_STATE_DICT = StateDict(sd={}, device=torch.device("cpu"), size=0, dtype=set())
|
||||
|
||||
|
||||
def _contiguous_byte_view(weights: dict[str, torch.Tensor]) -> torch.Tensor | None:
|
||||
"""Return a ``uint8`` view spanning every tensor in *weights*, or ``None`` if
|
||||
they don't share one contiguous storage region."""
|
||||
tensors = list(weights.values())
|
||||
if not tensors:
|
||||
return None
|
||||
storage = tensors[0].untyped_storage()
|
||||
storage_ptr = storage.data_ptr()
|
||||
start = end = tensors[0].storage_offset() * tensors[0].element_size()
|
||||
for t in tensors:
|
||||
if t.untyped_storage().data_ptr() != storage_ptr or not t.is_contiguous():
|
||||
return None
|
||||
offset = t.storage_offset() * t.element_size()
|
||||
nbytes = t.numel() * t.element_size()
|
||||
start = min(start, offset)
|
||||
end = max(end, offset + nbytes)
|
||||
view = torch.empty(0, dtype=torch.uint8, device=tensors[0].device)
|
||||
view.set_(storage, start, (end - start,), (1,))
|
||||
return view
|
||||
class CachedBlock(NamedTuple):
|
||||
"""A cached GPU block: the raw pool slot plus the carved per-key views.
|
||||
The raw slot is what is returned to the pool on eviction; the views are
|
||||
what callers consume.
|
||||
"""
|
||||
|
||||
raw: torch.Tensor
|
||||
views: dict[str, torch.Tensor]
|
||||
|
||||
|
||||
class WeightsProvider:
|
||||
@@ -43,63 +36,72 @@ class WeightsProvider:
|
||||
source: Pinned CPU weight source.
|
||||
lora_sources: LoRA adapters fused on H2D copy.
|
||||
blocks_prefix: State-dict prefix for LoRA key matching.
|
||||
fuse_rule: Per-policy LoRA merge rule (must be streaming-compatible:
|
||||
no companion-key emission). Defaults to ``bf16_fuse_rule``.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pool: WeightPool,
|
||||
pool: BufferPool,
|
||||
copy_stream: torch.cuda.Stream,
|
||||
target_device: torch.device,
|
||||
source: WeightSource,
|
||||
lora_sources: list[LoraSource] | None = None,
|
||||
blocks_prefix: str = "",
|
||||
fuse_rule: FuseRule = bf16_fuse_rule,
|
||||
) -> None:
|
||||
self._copy_stream = copy_stream
|
||||
self._pool = pool
|
||||
self._cache: OrderedDict[int, dict[str, torch.Tensor]] = OrderedDict()
|
||||
self._cache: OrderedDict[int, CachedBlock] = OrderedDict()
|
||||
self._events: dict[int, torch.cuda.Event] = {}
|
||||
self._target_device = target_device
|
||||
self._source = source
|
||||
self._lora_sources = lora_sources or []
|
||||
self._blocks_prefix = blocks_prefix
|
||||
self._fuse_rule = fuse_rule
|
||||
|
||||
def get(self, idx: int) -> dict[str, torch.Tensor]:
|
||||
"""Return GPU weights for block *idx*. Does H2D copy on miss."""
|
||||
if idx in self._cache:
|
||||
return self._cache[idx]
|
||||
return self._cache[idx].views
|
||||
|
||||
# Evict oldest GPU buffer if at capacity.
|
||||
if len(self._cache) >= self._pool.capacity:
|
||||
evicted_idx, evicted_weights = self._cache.popitem(last=False)
|
||||
self._pool.release(evicted_weights, event=self._events.pop(evicted_idx, None))
|
||||
evicted_idx, evicted = self._cache.popitem(last=False)
|
||||
self._pool.release(evicted.raw, event=self._events.pop(evicted_idx, None))
|
||||
|
||||
gpu_weights = self._pool.acquire()
|
||||
cpu_weights = self._source.get(idx)
|
||||
layout = self._source.block_layout(idx)
|
||||
raw = self._pool.acquire()
|
||||
gpu_weights = carve_buffer(raw, layout)
|
||||
cpu_buffer = self._source.get(idx)
|
||||
|
||||
h2d_event = self._copy_to_gpu(idx, gpu_weights, cpu_weights)
|
||||
h2d_event = self._copy_to_gpu(idx, raw, gpu_weights, cpu_buffer, layout_nbytes(layout))
|
||||
self._source.release(idx, event=h2d_event)
|
||||
|
||||
self._cache[idx] = gpu_weights
|
||||
self._cache[idx] = CachedBlock(raw, gpu_weights)
|
||||
return gpu_weights
|
||||
|
||||
def _copy_to_gpu(
|
||||
self,
|
||||
idx: int,
|
||||
raw: torch.Tensor,
|
||||
gpu_weights: dict[str, torch.Tensor],
|
||||
cpu_weights: dict[str, torch.Tensor],
|
||||
cpu_buffer: torch.Tensor,
|
||||
nbytes: int,
|
||||
) -> torch.cuda.Event:
|
||||
"""Enqueue H2D copy + LoRA fusion on the copy stream and wait on compute.
|
||||
The wait is intentionally inside this method so callers -- and
|
||||
instrumentation regions wrapping it -- observe the full transfer time.
|
||||
*cpu_buffer* is one contiguous source buffer carved by the same layout as
|
||||
*raw*, so a single byte copy of its leading *nbytes* reproduces every view
|
||||
in *gpu_weights*. The wait is intentionally inside this method so callers --
|
||||
and instrumentation regions wrapping it -- observe the full transfer time.
|
||||
"""
|
||||
if not cpu_buffer.is_contiguous() or cpu_buffer.dtype != torch.uint8 or cpu_buffer.numel() < nbytes:
|
||||
raise ValueError(
|
||||
f"source buffer for block {idx} must be a contiguous uint8 buffer of >= {nbytes} bytes, "
|
||||
f"got {cpu_buffer.dim()}-D {cpu_buffer.dtype} with {cpu_buffer.numel()} elements"
|
||||
)
|
||||
with torch.cuda.stream(self._copy_stream):
|
||||
gpu_view = _contiguous_byte_view(gpu_weights)
|
||||
cpu_view = _contiguous_byte_view(cpu_weights)
|
||||
if gpu_view is not None and cpu_view is not None and gpu_view.numel() == cpu_view.numel():
|
||||
gpu_view.copy_(cpu_view, non_blocking=True)
|
||||
else:
|
||||
for name, gpu_tensor in gpu_weights.items():
|
||||
gpu_tensor.copy_(cpu_weights[name], non_blocking=True)
|
||||
raw[:nbytes].copy_(cpu_buffer[:nbytes], non_blocking=True)
|
||||
if self._lora_sources:
|
||||
self._fuse_block_loras(idx, gpu_weights)
|
||||
h2d_event = torch.cuda.Event()
|
||||
@@ -126,22 +128,19 @@ class WeightsProvider:
|
||||
return len(self._cache)
|
||||
|
||||
def _fuse_block_loras(self, idx: int, weights: dict[str, torch.Tensor]) -> None:
|
||||
"""Fuse LoRA deltas directly into GPU block weights."""
|
||||
"""Fuse LoRA deltas directly into GPU block weights via ``fuse_rule``."""
|
||||
agg_dtype = self._fuse_rule.aggregation_dtype
|
||||
for name, tensor in weights.items():
|
||||
if not name.endswith(".weight"):
|
||||
continue
|
||||
prefix = f"{self._blocks_prefix}.{idx}.{name}".removesuffix(".weight")
|
||||
is_fp8 = tensor.dtype in FP8_DTYPES
|
||||
agg_dtype = torch.bfloat16 if is_fp8 else tensor.dtype
|
||||
products = (
|
||||
ab
|
||||
for ab in (s.get_ab(prefix, device=self._target_device, dtype=agg_dtype) for s in self._lora_sources)
|
||||
if ab is not None
|
||||
)
|
||||
aggregated = aggregate_lora_products(products, agg_dtype)
|
||||
if aggregated is None:
|
||||
deltas = aggregate_lora_products(products, agg_dtype)
|
||||
if deltas is None:
|
||||
continue
|
||||
if is_fp8:
|
||||
tensor.copy_(fuse_cast_fp8_weight(aggregated, tensor, tensor.dtype))
|
||||
else:
|
||||
tensor.add_(aggregated)
|
||||
fused = self._fuse_rule(name, tensor, deltas, _EMPTY_STATE_DICT)
|
||||
tensor.copy_(fused[name])
|
||||
|
||||
@@ -2,31 +2,38 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import OrderedDict
|
||||
from typing import Protocol
|
||||
from typing import NamedTuple, Protocol
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.block_streaming.disk import DiskBlockReader
|
||||
from ltx_core.block_streaming.pool import WeightPool
|
||||
from ltx_core.block_streaming.block_fetcher import BlockFetcher, FetchHandle
|
||||
from ltx_core.block_streaming.pool import BufferPool
|
||||
from ltx_core.block_streaming.utils import carve_buffer, layout_nbytes
|
||||
from ltx_core.loader.primitives import TensorLayout
|
||||
|
||||
|
||||
class WeightSource(Protocol):
|
||||
"""Provides pinned CPU weights for a given block index.
|
||||
Assumes all buffers share an identical layout across all block indices.
|
||||
Blocks may be heterogeneous: each has its own layout, so the source exposes a
|
||||
per-block layout and the byte size of the largest block (which sizes the
|
||||
pool slots -- a smaller block is carved into the front of a max-sized slot).
|
||||
The source is the single source of truth for each block's layout.
|
||||
"""
|
||||
|
||||
def block_layout(self, idx: int) -> TensorLayout:
|
||||
"""Per-block buffer layout (shape + dtype for each param)."""
|
||||
...
|
||||
|
||||
@property
|
||||
def block_layout(self) -> TensorLayout:
|
||||
"""Shared per-block buffer layout (shape + dtype for each param)."""
|
||||
def slot_nbytes(self) -> int:
|
||||
"""Byte size of the largest block; sizes a pool slot (16-byte aligned)."""
|
||||
...
|
||||
|
||||
def get(self, idx: int) -> dict[str, torch.Tensor]:
|
||||
"""Return CPU weights for block *idx*."""
|
||||
def get(self, idx: int) -> torch.Tensor:
|
||||
"""Return one contiguous CPU buffer for block *idx*."""
|
||||
...
|
||||
|
||||
def release(self, idx: int, event: torch.cuda.Event) -> None:
|
||||
def release(self, idx: int, event: torch.cuda.Event | None) -> None:
|
||||
"""Signal that an async operation using these weights is guarded by *event*."""
|
||||
...
|
||||
|
||||
@@ -35,68 +42,133 @@ class WeightSource(Protocol):
|
||||
...
|
||||
|
||||
|
||||
class DiskWeightSource(WeightSource):
|
||||
"""Reads block weights from disk into pinned CPU buffers on demand."""
|
||||
class _Scheduled(NamedTuple):
|
||||
"""A scheduled (possibly in-flight) read: the raw pool slot and its fetch handle."""
|
||||
|
||||
raw: torch.Tensor
|
||||
status: FetchHandle
|
||||
|
||||
|
||||
class DiskWeightSource(WeightSource):
|
||||
"""WeightSource that streams blocks from disk via a :class:`BlockFetcher`.
|
||||
``get(idx)`` must be paired with a ``release(idx)`` before *idx* is fetched
|
||||
again; getting a block that is already in flight raises. Each read acquires a
|
||||
raw pool slot, carves it to block *idx*'s layout, and hands the carved views to
|
||||
the fetcher to fill, so one max-sized slot can serve blocks of differing shapes.
|
||||
``get`` returns that same contiguous slot.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pool: BufferPool,
|
||||
fetcher: BlockFetcher,
|
||||
block_layouts: dict[int, TensorLayout],
|
||||
blocks_number: int,
|
||||
prefetch_depth: int = 0,
|
||||
) -> None:
|
||||
if blocks_number <= 0:
|
||||
raise ValueError(f"blocks_number must be > 0, got {blocks_number}")
|
||||
if prefetch_depth < 0:
|
||||
raise ValueError(f"prefetch_depth must be >= 0, got {prefetch_depth}")
|
||||
max_layout_nbytes = max((layout_nbytes(layout) for layout in block_layouts.values()), default=0)
|
||||
if pool.slot_nbytes < max_layout_nbytes:
|
||||
raise ValueError(
|
||||
f"pool slot is too small for the largest block: slot {pool.slot_nbytes} bytes < {max_layout_nbytes}"
|
||||
)
|
||||
|
||||
def __init__(self, pool: WeightPool, reader: DiskBlockReader) -> None:
|
||||
self._pool = pool
|
||||
self._cache: OrderedDict[int, dict[str, torch.Tensor]] = OrderedDict()
|
||||
self._events: dict[int, torch.cuda.Event] = {}
|
||||
self._reader = reader
|
||||
self._blocks_number = blocks_number
|
||||
self._prefetch_depth = prefetch_depth
|
||||
self._fetcher = fetcher
|
||||
self._block_layouts = block_layouts
|
||||
self._scheduled: dict[int, _Scheduled] = {}
|
||||
self._in_flight: dict[int, torch.Tensor] = {}
|
||||
|
||||
def block_layout(self, idx: int) -> TensorLayout:
|
||||
return self._block_layouts[idx]
|
||||
|
||||
@property
|
||||
def block_layout(self) -> TensorLayout:
|
||||
return self._pool.buffer_layout
|
||||
def slot_nbytes(self) -> int:
|
||||
return self._pool.slot_nbytes
|
||||
|
||||
def get(self, idx: int) -> dict[str, torch.Tensor]:
|
||||
"""Return CPU weights for block *idx*. Reads from disk on miss."""
|
||||
if idx in self._cache:
|
||||
return self._cache[idx]
|
||||
def get(self, idx: int) -> torch.Tensor:
|
||||
if idx in self._in_flight:
|
||||
raise RuntimeError(f"Block {idx} is already in flight; release it before getting it again")
|
||||
|
||||
if len(self._cache) >= self._pool.capacity:
|
||||
evicted_idx, evicted_weights = self._cache.popitem(last=False)
|
||||
self._pool.release(evicted_weights, event=self._events.pop(evicted_idx, None))
|
||||
scheduled = self._scheduled.pop(idx, None)
|
||||
if scheduled is None:
|
||||
scheduled = self._schedule(idx)
|
||||
error = scheduled.status.wait()
|
||||
if error is not None:
|
||||
self._pool.release(scheduled.raw)
|
||||
raise error
|
||||
|
||||
weights = self._pool.acquire()
|
||||
self._reader.read_into(weights, idx)
|
||||
self._cache[idx] = weights
|
||||
return weights
|
||||
self._in_flight[idx] = scheduled.raw
|
||||
for k in range(1, self._prefetch_depth + 1):
|
||||
self._ensure_scheduled((idx + k) % self._blocks_number)
|
||||
return scheduled.raw
|
||||
|
||||
def release(self, idx: int, event: torch.cuda.Event) -> None:
|
||||
"""Attach an H2D event -- waited before this buffer is recycled."""
|
||||
self._events[idx] = event
|
||||
def release(self, idx: int, event: torch.cuda.Event | None) -> None:
|
||||
raw_buffer = self._in_flight.pop(idx)
|
||||
self._pool.release(raw_buffer, event=event)
|
||||
|
||||
def cleanup(self) -> None:
|
||||
"""Clear cache and close the disk reader."""
|
||||
self._cache.clear()
|
||||
self._events.clear()
|
||||
self._reader.cleanup()
|
||||
self._fetcher.cleanup()
|
||||
while self._in_flight:
|
||||
_, raw_buffer = self._in_flight.popitem()
|
||||
self._pool.release(raw_buffer)
|
||||
while self._scheduled:
|
||||
_, scheduled = self._scheduled.popitem()
|
||||
scheduled.status.wait()
|
||||
self._pool.release(scheduled.raw)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._cache)
|
||||
def _ensure_scheduled(self, idx: int) -> None:
|
||||
"""Schedule a read for *idx* if one is not already pending."""
|
||||
if idx not in self._scheduled:
|
||||
self._scheduled[idx] = self._schedule(idx)
|
||||
|
||||
def _schedule(self, idx: int) -> _Scheduled:
|
||||
"""Acquire a raw slot, carve it to block *idx*, enqueue a read, return the handle.
|
||||
The raw slot and its fetch status are returned together so the caller can
|
||||
track both as one unit; the fetcher only receives the carved views to fill.
|
||||
"""
|
||||
raw_buffer = self._pool.acquire()
|
||||
carved = carve_buffer(raw_buffer, self._block_layouts[idx])
|
||||
status = self._fetcher.submit(idx, carved)
|
||||
return _Scheduled(raw_buffer, status)
|
||||
|
||||
|
||||
class PinnedBlock(NamedTuple):
|
||||
"""A pre-loaded pinned block: its single contiguous buffer and the layout to carve it with."""
|
||||
|
||||
buffer: torch.Tensor
|
||||
layout: TensorLayout
|
||||
|
||||
|
||||
class PinnedWeightSource(WeightSource):
|
||||
"""Pre-loaded pinned CPU weights."""
|
||||
"""Pre-loaded pinned CPU weights, one contiguous (possibly heterogeneous) buffer per block."""
|
||||
|
||||
def __init__(self, weights: dict[int, dict[str, torch.Tensor]]) -> None:
|
||||
if not weights:
|
||||
def __init__(self, blocks: dict[int, PinnedBlock]) -> None:
|
||||
if not blocks:
|
||||
raise ValueError("PinnedWeightSource requires at least one block")
|
||||
self._weights = weights
|
||||
self._blocks = blocks
|
||||
self._slot_nbytes = max(layout_nbytes(block.layout) for block in blocks.values())
|
||||
|
||||
def block_layout(self, idx: int) -> TensorLayout:
|
||||
return self._blocks[idx].layout
|
||||
|
||||
@property
|
||||
def block_layout(self) -> TensorLayout:
|
||||
first_block = self._weights[min(self._weights)]
|
||||
return {name: (t.shape, t.dtype) for name, t in first_block.items()}
|
||||
def slot_nbytes(self) -> int:
|
||||
return self._slot_nbytes
|
||||
|
||||
def get(self, idx: int) -> dict[str, torch.Tensor]:
|
||||
return self._weights[idx]
|
||||
def get(self, idx: int) -> torch.Tensor:
|
||||
return self._blocks[idx].buffer
|
||||
|
||||
def release(self, idx: int, event: torch.cuda.Event) -> None:
|
||||
def release(self, idx: int, event: torch.cuda.Event | None) -> None:
|
||||
pass
|
||||
|
||||
def cleanup(self) -> None:
|
||||
self._weights.clear()
|
||||
self._blocks.clear()
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._weights)
|
||||
return len(self._blocks)
|
||||
|
||||
@@ -5,7 +5,7 @@ from __future__ import annotations
|
||||
import math
|
||||
import weakref
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from typing import Any, NamedTuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -84,7 +84,7 @@ def _alloc_pinned_exact(nbytes: int) -> torch.Tensor | None:
|
||||
return buf
|
||||
|
||||
|
||||
def _alloc_buffer(nbytes: int, device: torch.device | None, pin_memory: bool) -> torch.Tensor:
|
||||
def alloc_buffer(nbytes: int, device: torch.device | None, pin_memory: bool) -> torch.Tensor:
|
||||
"""Allocate one ``uint8`` buffer for :func:`allocate_layout_views`.
|
||||
For pinned host buffers, prefer ``cudaHostRegister`` to dodge the caching
|
||||
allocator's power-of-2 rounding. Falls back to the caching allocator if
|
||||
@@ -112,6 +112,47 @@ class _TensorSlice:
|
||||
return math.prod(self.shape) * self.dtype.itemsize
|
||||
|
||||
|
||||
class LayoutSlices(NamedTuple):
|
||||
"""Per-key tensor slices of a layout plus the total aligned buffer size."""
|
||||
|
||||
slices: dict[str, _TensorSlice]
|
||||
nbytes: int
|
||||
|
||||
|
||||
def _layout_slices(layout: TensorLayout) -> LayoutSlices:
|
||||
"""Compute the byte offset of each key in *layout* and the total aligned size.
|
||||
The size is at least one byte so empty layouts still produce a valid buffer.
|
||||
"""
|
||||
slices: dict[str, _TensorSlice] = {}
|
||||
cursor = 0
|
||||
for key, (shape, dtype) in layout.items():
|
||||
cursor = _align_up(cursor, _BUFFER_ALIGN)
|
||||
slices[key] = _TensorSlice(offset=cursor, shape=shape, dtype=dtype)
|
||||
cursor += slices[key].size()
|
||||
return LayoutSlices(slices, max(_align_up(cursor, _BUFFER_ALIGN), 1))
|
||||
|
||||
|
||||
def layout_nbytes(layout: TensorLayout) -> int:
|
||||
"""Byte size of one contiguous, 16-byte-aligned buffer holding *layout* (>= 1)."""
|
||||
return _layout_slices(layout).nbytes
|
||||
|
||||
|
||||
def carve_buffer(buffer: torch.Tensor, layout: TensorLayout) -> dict[str, torch.Tensor]:
|
||||
"""Carve per-key tensor views for *layout* into the front of *buffer*.
|
||||
*buffer* is a 1-D ``uint8`` tensor at least :func:`layout_nbytes` long. Each
|
||||
returned tensor is a non-overlapping slice of its leading bytes reinterpreted
|
||||
at the requested shape and dtype; any trailing bytes are left unused. That
|
||||
slack is what lets one max-sized pool slot hold a smaller (heterogeneous)
|
||||
block. The views keep *buffer*'s storage alive via PyTorch refcounting.
|
||||
"""
|
||||
if buffer.dtype != torch.uint8 or buffer.dim() != 1:
|
||||
raise ValueError(f"carve_buffer expects a 1-D uint8 buffer, got {buffer.dim()}-D {buffer.dtype}")
|
||||
slices, nbytes = _layout_slices(layout)
|
||||
if buffer.numel() < nbytes:
|
||||
raise ValueError(f"buffer too small to carve layout: need {nbytes} bytes, got {buffer.numel()}")
|
||||
return {key: buffer[s.offset : s.offset + s.size()].view(s.dtype).view(s.shape) for key, s in slices.items()}
|
||||
|
||||
|
||||
def allocate_layout_views(
|
||||
layout: TensorLayout,
|
||||
device: torch.device | None = None,
|
||||
@@ -123,12 +164,5 @@ def allocate_layout_views(
|
||||
requested shape and dtype. The views keep the underlying storage alive
|
||||
via PyTorch refcounting — drop them all to release the memory.
|
||||
"""
|
||||
slices: dict[str, _TensorSlice] = {}
|
||||
cursor = 0
|
||||
for key, (shape, dtype) in layout.items():
|
||||
cursor = _align_up(cursor, _BUFFER_ALIGN)
|
||||
slices[key] = _TensorSlice(offset=cursor, shape=shape, dtype=dtype)
|
||||
cursor += slices[key].size()
|
||||
# Allocate at least one byte so empty layouts still produce a valid buffer.
|
||||
buffer = _alloc_buffer(max(_align_up(cursor, _BUFFER_ALIGN), 1), device, pin_memory)
|
||||
return {key: buffer[s.offset : s.offset + s.size()].view(s.dtype).view(s.shape) for key, s in slices.items()}
|
||||
buffer = alloc_buffer(layout_nbytes(layout), device, pin_memory)
|
||||
return carve_buffer(buffer, layout)
|
||||
|
||||
@@ -253,6 +253,11 @@ class MultiModalGuider:
|
||||
and as scale * (cond - uncond) for stg, steering the denoising process away from the unconditioned
|
||||
prediction.
|
||||
"""
|
||||
dtype = cond.dtype
|
||||
cond = cond.float()
|
||||
uncond_text = uncond_text.float() if isinstance(uncond_text, torch.Tensor) else uncond_text
|
||||
uncond_perturbed = uncond_perturbed.float() if isinstance(uncond_perturbed, torch.Tensor) else uncond_perturbed
|
||||
uncond_modality = uncond_modality.float() if isinstance(uncond_modality, torch.Tensor) else uncond_modality
|
||||
pred = (
|
||||
cond
|
||||
+ (self.params.cfg_scale - 1) * (cond - uncond_text)
|
||||
@@ -265,7 +270,7 @@ class MultiModalGuider:
|
||||
factor = self.params.rescale_scale * factor + (1 - self.params.rescale_scale)
|
||||
pred = pred * factor
|
||||
|
||||
return pred
|
||||
return pred.to(dtype)
|
||||
|
||||
def do_unconditional_generation(self) -> bool:
|
||||
"""Returns True if the guider is doing unconditional generation."""
|
||||
|
||||
@@ -17,18 +17,20 @@ class GaussianNoiser(Noiser):
|
||||
|
||||
def __init__(self, generator: torch.Generator):
|
||||
super().__init__()
|
||||
|
||||
self.generator = generator
|
||||
|
||||
def __call__(self, latent_state: LatentState, noise_scale: float = 1.0) -> LatentState:
|
||||
noise = torch.randn(
|
||||
def _sample_noise(self, latent_state: LatentState) -> torch.Tensor:
|
||||
return torch.randn(
|
||||
*latent_state.latent.shape,
|
||||
device=latent_state.latent.device,
|
||||
dtype=latent_state.latent.dtype,
|
||||
generator=self.generator,
|
||||
)
|
||||
scaled_mask = latent_state.denoise_mask * noise_scale
|
||||
latent = noise * scaled_mask + latent_state.latent * (1 - scaled_mask)
|
||||
|
||||
def __call__(self, latent_state: LatentState, noise_scale: float = 1.0) -> LatentState:
|
||||
noise = self._sample_noise(latent_state)
|
||||
latent = torch.lerp(latent_state.latent.float(), noise.float(), noise_scale)
|
||||
latent = torch.lerp(latent_state.clean_latent.float(), latent, latent_state.denoise_mask)
|
||||
return replace(
|
||||
latent_state,
|
||||
latent=latent.to(latent_state.latent.dtype),
|
||||
|
||||
@@ -7,6 +7,7 @@ from ltx_core.conditioning.types import (
|
||||
ConditioningItemAttentionStrengthWrapper,
|
||||
VideoConditionByKeyframeIndex,
|
||||
VideoConditionByLatentIndex,
|
||||
VideoConditionByMask,
|
||||
VideoConditionByReferenceLatent,
|
||||
)
|
||||
|
||||
@@ -17,5 +18,6 @@ __all__ = [
|
||||
"ConditioningItemAttentionStrengthWrapper",
|
||||
"VideoConditionByKeyframeIndex",
|
||||
"VideoConditionByLatentIndex",
|
||||
"VideoConditionByMask",
|
||||
"VideoConditionByReferenceLatent",
|
||||
]
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
from ltx_core.conditioning.types.attention_strength_wrapper import ConditioningItemAttentionStrengthWrapper
|
||||
from ltx_core.conditioning.types.keyframe_cond import VideoConditionByKeyframeIndex
|
||||
from ltx_core.conditioning.types.latent_cond import VideoConditionByLatentIndex
|
||||
from ltx_core.conditioning.types.mask_cond import VideoConditionByMask
|
||||
from ltx_core.conditioning.types.reference_audio_cond import AudioConditionByReferenceLatent
|
||||
from ltx_core.conditioning.types.reference_video_cond import VideoConditionByReferenceLatent
|
||||
|
||||
@@ -11,5 +12,6 @@ __all__ = [
|
||||
"ConditioningItemAttentionStrengthWrapper",
|
||||
"VideoConditionByKeyframeIndex",
|
||||
"VideoConditionByLatentIndex",
|
||||
"VideoConditionByMask",
|
||||
"VideoConditionByReferenceLatent",
|
||||
]
|
||||
|
||||
@@ -10,8 +10,9 @@ from ltx_core.types import LatentState, VideoLatentShape
|
||||
class VideoConditionByKeyframeIndex(ConditioningItem):
|
||||
"""
|
||||
Conditions video generation on keyframe latents at a specific frame index.
|
||||
Appends keyframe tokens to the latent state with positions offset by frame_idx,
|
||||
and sets denoise strength according to the strength parameter.
|
||||
Appends keyframe tokens to the sequence with positions offset by frame_idx: the keyframe
|
||||
latents become clean-latent tokens (placeholder zeros in the noisy latent) and the denoise
|
||||
mask is set from the strength parameter.
|
||||
To add attention masking, wrap with :class:`ConditioningItemAttentionStrengthWrapper`.
|
||||
Args:
|
||||
keyframes: Keyframe latents [B, C, F, H, W].
|
||||
@@ -75,7 +76,7 @@ class VideoConditionByKeyframeIndex(ConditioningItem):
|
||||
)
|
||||
|
||||
return LatentState(
|
||||
latent=torch.cat([latent_state.latent, tokens], dim=1),
|
||||
latent=torch.cat([latent_state.latent, torch.zeros_like(tokens)], dim=1),
|
||||
denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1),
|
||||
positions=torch.cat([latent_state.positions, positions], dim=2),
|
||||
clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1),
|
||||
|
||||
@@ -9,8 +9,8 @@ from ltx_core.types import LatentState
|
||||
class VideoConditionByLatentIndex(ConditioningItem):
|
||||
"""
|
||||
Conditions video generation by injecting latents at a specific latent frame index.
|
||||
Replaces tokens in the latent state at positions corresponding to latent_idx,
|
||||
and sets denoise strength according to the strength parameter.
|
||||
Sets the clean latents at positions corresponding to latent_idx to the injected latents,
|
||||
sets denoise strength according to the strength parameter.
|
||||
"""
|
||||
|
||||
def __init__(self, latent: torch.Tensor, strength: float, latent_idx: int):
|
||||
@@ -37,7 +37,6 @@ class VideoConditionByLatentIndex(ConditioningItem):
|
||||
|
||||
latent_state = latent_state.clone()
|
||||
|
||||
latent_state.latent[:, start_token:stop_token] = tokens
|
||||
latent_state.clean_latent[:, start_token:stop_token] = tokens
|
||||
latent_state.denoise_mask[:, start_token:stop_token] = 1.0 - self.strength
|
||||
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
"""Mask-based conditioning for inpainting and spatial conditioning."""
|
||||
|
||||
from dataclasses import replace
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.conditioning.item import ConditioningItem
|
||||
from ltx_core.tools import LatentTools
|
||||
from ltx_core.types import LatentState
|
||||
|
||||
|
||||
class VideoConditionByMask(ConditioningItem):
|
||||
"""Condition video generation using a binary mask over latent frames.
|
||||
Masked positions (mask=1) receive the provided clean latent values and are
|
||||
excluded from denoising (denoise_mask set to ``1 - strength``). Unmasked
|
||||
positions (mask=0) are left unchanged and denoised normally.
|
||||
The mask operates in **unpatchified latent** space — it should have shape
|
||||
``[B, F, H, W]`` matching the latent dimensions (after VAE encoding,
|
||||
before patchification). This is consistent with the latent input format
|
||||
used by all other conditioning items.
|
||||
Args:
|
||||
latent: Clean conditioning latents in unpatchified format [B, C, F, H, W].
|
||||
Must match the target shape of the latent tools.
|
||||
mask: Binary mask [B, F, H, W] in unpatchified latent space.
|
||||
1 = conditioning position (clean, excluded from denoising),
|
||||
0 = generated position (noised, denoised normally).
|
||||
strength: Conditioning strength for masked positions. 1.0 = fully clean
|
||||
(no denoising), 0.0 = no conditioning effect. Default 1.0.
|
||||
"""
|
||||
|
||||
def __init__(self, latent: torch.Tensor, mask: torch.Tensor, strength: float = 1.0):
|
||||
self.latent = latent
|
||||
self.mask = mask
|
||||
self.strength = strength
|
||||
|
||||
def apply_to(self, latent_state: LatentState, latent_tools: LatentTools) -> LatentState:
|
||||
"""Apply mask-based conditioning to the latent state."""
|
||||
tokens = latent_tools.patchifier.patchify(self.latent)
|
||||
|
||||
mask = latent_tools.patchifier.patchify(self.mask.unsqueeze(1))
|
||||
|
||||
m = mask.to(dtype=latent_state.latent.dtype)
|
||||
inv = 1 - m
|
||||
|
||||
return replace(
|
||||
latent_state,
|
||||
clean_latent=latent_state.clean_latent * inv + tokens * m,
|
||||
denoise_mask=latent_state.denoise_mask * inv + (1.0 - self.strength) * m,
|
||||
)
|
||||
@@ -51,7 +51,7 @@ class AudioConditionByReferenceLatent:
|
||||
)
|
||||
|
||||
return LatentState(
|
||||
latent=torch.cat([latent_state.latent, tokens], dim=1),
|
||||
latent=torch.cat([latent_state.latent, torch.zeros_like(tokens)], dim=1),
|
||||
denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1),
|
||||
positions=torch.cat([latent_state.positions, self.positions], dim=2),
|
||||
clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1),
|
||||
|
||||
@@ -14,7 +14,8 @@ class VideoConditionByReferenceLatent(ConditioningItem):
|
||||
Conditions video generation on a reference video latent for IC-LoRA inference.
|
||||
IC-LoRAs are trained by concatenating reference (control signal) and target tokens,
|
||||
learning to attend across both. This class replicates that setup at inference by
|
||||
appending reference tokens to the latent sequence.
|
||||
appending the reference tokens to the sequence as clean latents (with placeholder zeros
|
||||
in the noisy latent).
|
||||
IC-LoRAs can be trained with lower-resolution references than the target (e.g., 384px
|
||||
reference for 768px output) for efficiency and better generalization. The
|
||||
`downscale_factor` scales reference positions to match target coordinates, preserving
|
||||
@@ -22,9 +23,9 @@ class VideoConditionByReferenceLatent(ConditioningItem):
|
||||
(stored in LoRA metadata).
|
||||
To add attention masking, wrap with :class:`ConditioningItemAttentionStrengthWrapper`.
|
||||
Args:
|
||||
latent: Reference video latents [B, C, F, H, W]
|
||||
downscale_factor: Target/reference resolution ratio (e.g., 2 = half-resolution
|
||||
reference). Spatial positions are scaled by this factor.
|
||||
latent: Reference video latents [B, C, F, H, W].
|
||||
downscale_factor: Target/reference spatial ratio (e.g. 2 = half-res ref).
|
||||
temporal_scale_factor: Target/reference temporal ratio S (e.g. 4 = ref at 1/4 fps).
|
||||
strength: Conditioning strength. 1.0 = full (reference kept clean),
|
||||
0.0 = none (reference denoised). Default 1.0.
|
||||
"""
|
||||
@@ -33,10 +34,12 @@ class VideoConditionByReferenceLatent(ConditioningItem):
|
||||
self,
|
||||
latent: torch.Tensor,
|
||||
downscale_factor: int = 1,
|
||||
temporal_scale_factor: int = 1,
|
||||
strength: float = 1.0,
|
||||
):
|
||||
self.latent = latent
|
||||
self.downscale_factor = downscale_factor
|
||||
self.temporal_scale_factor = temporal_scale_factor
|
||||
self.strength = strength
|
||||
|
||||
def apply_to(
|
||||
@@ -44,10 +47,9 @@ class VideoConditionByReferenceLatent(ConditioningItem):
|
||||
latent_state: LatentState,
|
||||
latent_tools: VideoLatentTools,
|
||||
) -> LatentState:
|
||||
"""Append reference video tokens with scaled positions."""
|
||||
"""Append reference video tokens with positions translated into the target frame."""
|
||||
tokens = latent_tools.patchifier.patchify(self.latent)
|
||||
|
||||
# Compute positions for the reference video's actual dimensions
|
||||
latent_coords = latent_tools.patchifier.get_patch_grid_bounds(
|
||||
output_shape=VideoLatentShape.from_torch_shape(self.latent.shape),
|
||||
device=self.latent.device,
|
||||
@@ -58,9 +60,18 @@ class VideoConditionByReferenceLatent(ConditioningItem):
|
||||
causal_fix=latent_tools.causal_fix,
|
||||
)
|
||||
positions = positions.to(dtype=torch.float32)
|
||||
positions[:, 0, ...] /= latent_tools.fps
|
||||
|
||||
# Scale spatial positions to match target coordinate space
|
||||
# Place ref tokens on their own time spacing (= target_fps / S).
|
||||
positions[:, 0, ...] /= latent_tools.fps / self.temporal_scale_factor
|
||||
|
||||
# Translate into the target's frame so ref's last patch ends with target's last
|
||||
# patch; clamp the causal patch's negative start back to [0, 1/target_fps).
|
||||
if self.temporal_scale_factor != 1:
|
||||
t_target = latent_state.positions[:, 0, 0:1, 1:2].to(dtype=torch.float32) # = 1/target_fps
|
||||
positions[:, 0, ...] = torch.clamp(
|
||||
positions[:, 0, ...] - (self.temporal_scale_factor - 1) * t_target,
|
||||
min=0,
|
||||
)
|
||||
if self.downscale_factor != 1:
|
||||
positions[:, 1, ...] *= self.downscale_factor # height axis
|
||||
positions[:, 2, ...] *= self.downscale_factor # width axis
|
||||
@@ -83,7 +94,7 @@ class VideoConditionByReferenceLatent(ConditioningItem):
|
||||
)
|
||||
|
||||
return LatentState(
|
||||
latent=torch.cat([latent_state.latent, tokens], dim=1),
|
||||
latent=torch.cat([latent_state.latent, torch.zeros_like(tokens)], dim=1),
|
||||
denoise_mask=torch.cat([latent_state.denoise_mask, denoise_mask], dim=1),
|
||||
positions=torch.cat([latent_state.positions, positions], dim=2),
|
||||
clean_latent=torch.cat([latent_state.clean_latent, tokens], dim=1),
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
"""Builder ops for swapping attention backends on a meta model before load."""
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.model.transformer.attention import (
|
||||
Attention,
|
||||
AttentionCallable,
|
||||
AttentionFunction,
|
||||
MaskedAttentionCallable,
|
||||
MaskedAttentionFunction,
|
||||
)
|
||||
|
||||
|
||||
def set_attention_module_op(
|
||||
attention: AttentionFunction | AttentionCallable | None = None,
|
||||
masked_attention: MaskedAttentionFunction | MaskedAttentionCallable | None = None,
|
||||
) -> ModuleOps:
|
||||
"""Build a ``ModuleOps`` that overrides the attention callables on every
|
||||
``Attention`` submodule of a model. Applied via ``create_meta_model`` so
|
||||
the meta model is mutated before weight loading. Matcher returns False
|
||||
for models with no ``Attention`` submodules, so the op is a no-op there.
|
||||
Either or both slots may be supplied; *None* leaves that slot untouched.
|
||||
"""
|
||||
fn = attention.to_callable() if isinstance(attention, AttentionFunction) else attention
|
||||
masked_fn = (
|
||||
masked_attention.to_callable() if isinstance(masked_attention, MaskedAttentionFunction) else masked_attention
|
||||
)
|
||||
|
||||
def matcher(model: torch.nn.Module) -> bool:
|
||||
return any(isinstance(m, Attention) for m in model.modules())
|
||||
|
||||
def mutator(model: torch.nn.Module) -> torch.nn.Module:
|
||||
for module in model.modules():
|
||||
if isinstance(module, Attention):
|
||||
if fn is not None:
|
||||
module.attention_function = fn
|
||||
if masked_fn is not None:
|
||||
module.masked_attention_function = masked_fn
|
||||
return model
|
||||
|
||||
return ModuleOps(name="set_attention_backend", matcher=matcher, mutator=mutator)
|
||||
@@ -1,12 +1,10 @@
|
||||
from collections.abc import Iterable, Iterator
|
||||
from collections.abc import Callable, Iterable, Iterator
|
||||
from dataclasses import dataclass
|
||||
from typing import NamedTuple
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.kernels import TRITON_AVAILABLE
|
||||
from ltx_core.loader.primitives import LoraStateDictWithStrength, StateDict
|
||||
from ltx_core.quantization.fp8_cast import fused_add_round_launch
|
||||
from ltx_core.quantization.fp8_scaled_mm import quantize_weight_to_fp8_per_tensor
|
||||
|
||||
|
||||
class LoraProduct(NamedTuple):
|
||||
@@ -17,6 +15,62 @@ class LoraProduct(NamedTuple):
|
||||
strength: float
|
||||
|
||||
|
||||
#: Signature for a fuse callable used by :class:`FuseRule`.
|
||||
#:
|
||||
#: Args:
|
||||
#: key: The state-dict key being fused (e.g. ``"...layers.0.attn.q.weight"``).
|
||||
#: weight: The current value at ``key`` from ``model_sd``, on the fusion device.
|
||||
#: deltas: The pre-aggregated LoRA delta for ``key``, in ``aggregation_dtype``.
|
||||
#: model_sd: The full state dict, for rules that need companion keys
|
||||
#: (e.g. an existing ``.weight_scale``).
|
||||
#:
|
||||
#: Returns a dict of state-dict keys to overwrite -- at minimum ``{key: new_weight}``,
|
||||
#: plus any companion keys (e.g. an updated ``.weight_scale``) the policy needs to
|
||||
#: keep in sync.
|
||||
FuseFn = Callable[[str, torch.Tensor, torch.Tensor, StateDict], dict[str, torch.Tensor]]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FuseRule:
|
||||
"""Fuse an aggregated LoRA delta into one weight key.
|
||||
Each policy supplies its own rule (see ``QuantizationPolicy.fuse_rule``);
|
||||
``fuse_lora_weights`` is policy-agnostic boilerplate around it.
|
||||
Attributes:
|
||||
aggregation_dtype: Dtype callers must pre-aggregate LoRA deltas in
|
||||
before invoking the rule.
|
||||
fuse_fn: Callable that applies the pre-aggregated deltas to the weight
|
||||
(and any companion keys) and returns a dict of keys to overwrite —
|
||||
at minimum ``{key: new_weight}``, plus any companion keys (e.g. an
|
||||
updated ``.weight_scale`` for scaled-FP8 layouts) the policy needs
|
||||
to keep in sync.
|
||||
"""
|
||||
|
||||
aggregation_dtype: torch.dtype
|
||||
fuse_fn: FuseFn
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
key: str,
|
||||
weight: torch.Tensor,
|
||||
deltas: torch.Tensor,
|
||||
model_sd: StateDict,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
return self.fuse_fn(key, weight, deltas, model_sd)
|
||||
|
||||
|
||||
def _bf16_fuse(
|
||||
key: str,
|
||||
weight: torch.Tensor,
|
||||
deltas: torch.Tensor,
|
||||
model_sd: StateDict, # noqa: ARG001
|
||||
) -> dict[str, torch.Tensor]:
|
||||
deltas.add_(weight)
|
||||
return {key: deltas.to(dtype=weight.dtype)}
|
||||
|
||||
|
||||
bf16_fuse_rule = FuseRule(aggregation_dtype=torch.bfloat16, fuse_fn=_bf16_fuse)
|
||||
|
||||
|
||||
def _get_device() -> torch.device:
|
||||
if torch.cuda.is_available():
|
||||
return torch.device("cuda", torch.cuda.current_device())
|
||||
@@ -25,19 +79,16 @@ def _get_device() -> torch.device:
|
||||
|
||||
def aggregate_lora_products(
|
||||
products: Iterable[LoraProduct],
|
||||
dtype: torch.dtype | None = None,
|
||||
*,
|
||||
out: torch.Tensor | None = None,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor | None:
|
||||
"""Accumulate ``sum((B * strength) @ A)`` across :class:`LoraProduct` items.
|
||||
If ``out`` is provided, ``addmm_`` accumulates directly into it — caller
|
||||
ensures A/B dtypes and devices match ``out``. Otherwise the first product
|
||||
materializes the ``(out, in)``-shape aggregator at ``dtype``; subsequent
|
||||
products use ``addmm_`` to avoid allocating the full intermediate delta.
|
||||
Returns ``out`` (or the new aggregator), or ``None`` if ``products`` was empty
|
||||
and ``out`` was not given.
|
||||
The first product materializes a freshly-allocated aggregator via
|
||||
``torch.matmul(B * strength, A).to(dtype)`` -- preserving the
|
||||
``(B * strength) @ A`` rounding pattern. Subsequent products use
|
||||
``addmm_`` to avoid allocating the full intermediate delta.
|
||||
Returns the aggregator, or ``None`` if ``products`` was empty.
|
||||
"""
|
||||
aggregated = out
|
||||
aggregated: torch.Tensor | None = None
|
||||
for product in products:
|
||||
if aggregated is None:
|
||||
aggregated = torch.matmul(product.b * product.strength, product.a).to(dtype=dtype)
|
||||
@@ -46,62 +97,34 @@ def aggregate_lora_products(
|
||||
return aggregated
|
||||
|
||||
|
||||
def fuse_cast_fp8_weight(
|
||||
delta_bf16: torch.Tensor,
|
||||
weight_fp8: torch.Tensor,
|
||||
target_dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
"""Return ``(delta_bf16 + dequantize(weight_fp8)).to(target_dtype)``.
|
||||
CUDA with Triton uses stochastic rounding; otherwise uses a deterministic bf16 add.
|
||||
``delta_bf16`` is the bf16 accumulator and is mutated in place.
|
||||
"""
|
||||
if delta_bf16.dtype != torch.bfloat16:
|
||||
raise ValueError(f"delta_bf16 must be bfloat16, got {delta_bf16.dtype}")
|
||||
if str(weight_fp8.device).startswith("cuda") and TRITON_AVAILABLE:
|
||||
fused_add_round_launch(delta_bf16, weight_fp8, seed=0)
|
||||
else:
|
||||
delta_bf16.add_(weight_fp8.to(dtype=torch.bfloat16))
|
||||
return delta_bf16.to(dtype=target_dtype)
|
||||
|
||||
|
||||
def fuse_lora_weights(
|
||||
model_sd: StateDict,
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
dtype: torch.dtype | None = None,
|
||||
fuse_rule: FuseRule = bf16_fuse_rule,
|
||||
preserve_input_device: bool = True,
|
||||
) -> Iterator[tuple[str, torch.Tensor]]:
|
||||
"""Yield ``(key, fused_tensor)`` for each weight modified by at least one LoRA.
|
||||
For scaled-FP8 weights, this includes both the updated ``.weight`` tensor
|
||||
and its corresponding ``.weight_scale`` tensor.
|
||||
The fusion math is delegated to ``fuse_rule``.
|
||||
Output dtypes are the rule's responsibility.
|
||||
When ``preserve_input_device`` is False, fused tensors are yielded on the device
|
||||
used for fusion; caller is responsible for moving them to their final
|
||||
destination.
|
||||
"""
|
||||
for key, original_weight in model_sd.sd.items():
|
||||
if original_weight is None or key.endswith(".weight_scale"):
|
||||
fusion_device = _get_device()
|
||||
for key in _affected_weight_keys(lora_sd_and_strengths):
|
||||
original_weight = model_sd.sd.get(key)
|
||||
if original_weight is None:
|
||||
continue
|
||||
original_device = original_weight.device
|
||||
weight = original_weight.to(device=_get_device())
|
||||
target_dtype = dtype if dtype is not None else weight.dtype
|
||||
deltas_dtype = target_dtype if target_dtype not in [torch.float8_e4m3fn, torch.float8_e5m2] else torch.bfloat16
|
||||
|
||||
deltas = _aggregate_deltas(lora_sd_and_strengths, key, deltas_dtype, weight.device)
|
||||
products = _products_for_sd_key(lora_sd_and_strengths, key, fuse_rule.aggregation_dtype, fusion_device)
|
||||
deltas = aggregate_lora_products(products, fuse_rule.aggregation_dtype)
|
||||
if deltas is None:
|
||||
continue
|
||||
|
||||
scale_key = key.replace(".weight", ".weight_scale") if key.endswith(".weight") else None
|
||||
is_scaled_fp8 = scale_key is not None and scale_key in model_sd.sd
|
||||
original_device = original_weight.device
|
||||
weight = original_weight.to(device=fusion_device)
|
||||
|
||||
if weight.dtype == torch.float8_e4m3fn:
|
||||
if is_scaled_fp8:
|
||||
fused = _fuse_delta_with_scaled_fp8(deltas, weight, key, scale_key, model_sd)
|
||||
else:
|
||||
fused = {key: fuse_cast_fp8_weight(deltas, weight, target_dtype)}
|
||||
elif weight.dtype == torch.bfloat16:
|
||||
deltas.add_(weight)
|
||||
fused = {key: deltas.to(dtype=target_dtype)}
|
||||
else:
|
||||
raise ValueError(f"Unsupported dtype: {weight.dtype}")
|
||||
fused = fuse_rule(key, weight, deltas, model_sd)
|
||||
|
||||
for k, v in fused.items():
|
||||
yield k, v.to(device=original_device) if preserve_input_device else v
|
||||
@@ -110,53 +133,46 @@ def fuse_lora_weights(
|
||||
def apply_loras(
|
||||
model_sd: StateDict,
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
dtype: torch.dtype | None = None,
|
||||
fuse_rule: FuseRule = bf16_fuse_rule,
|
||||
destination_sd: StateDict | None = None,
|
||||
) -> StateDict:
|
||||
"""Fuse LoRAs into ``model_sd`` and place the results in ``destination_sd``.
|
||||
When ``destination_sd`` is provided, the fused tensors are placed directly into it.
|
||||
"""
|
||||
fused_iter = fuse_lora_weights(
|
||||
model_sd,
|
||||
lora_sd_and_strengths,
|
||||
fuse_rule=fuse_rule,
|
||||
)
|
||||
if destination_sd is not None:
|
||||
for key, fused in fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype):
|
||||
for key, fused in fused_iter:
|
||||
destination_sd.sd[key] = fused
|
||||
return destination_sd
|
||||
|
||||
fused = dict(fuse_lora_weights(model_sd, lora_sd_and_strengths, dtype))
|
||||
fused = dict(fused_iter)
|
||||
sd = {k: (fused[k] if k in fused else v.clone()) for k, v in model_sd.sd.items()}
|
||||
return StateDict(sd, model_sd.device, model_sd.size, model_sd.dtype)
|
||||
|
||||
|
||||
def _aggregate_deltas(
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength], key: str, dtype: torch.dtype, device: torch.device
|
||||
) -> torch.Tensor | None:
|
||||
def _affected_weight_keys(lora_sd_and_strengths: list[LoraStateDictWithStrength]) -> set[str]:
|
||||
"""Return the set of ``.weight`` keys touched by at least one LoRA in the list."""
|
||||
suffix = ".lora_A.weight"
|
||||
return {k[: -len(suffix)] + ".weight" for lsd, _ in lora_sd_and_strengths for k in lsd.sd if k.endswith(suffix)}
|
||||
|
||||
|
||||
def _products_for_sd_key(
|
||||
lora_sd_and_strengths: list[LoraStateDictWithStrength],
|
||||
key: str,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> Iterator[LoraProduct]:
|
||||
"""Yield :class:`LoraProduct` items matching *key* across state-dict-backed LoRAs."""
|
||||
prefix = key[: -len(".weight")]
|
||||
key_a = f"{prefix}.lora_A.weight"
|
||||
key_b = f"{prefix}.lora_B.weight"
|
||||
|
||||
def _ab_products() -> Iterator[LoraProduct]:
|
||||
for lsd, coef in lora_sd_and_strengths:
|
||||
if key_a not in lsd.sd or key_b not in lsd.sd:
|
||||
continue
|
||||
a = lsd.sd[key_a].to(device=device, dtype=dtype, non_blocking=True)
|
||||
b = lsd.sd[key_b].to(device=device, dtype=dtype, non_blocking=True)
|
||||
yield LoraProduct(a, b, coef)
|
||||
|
||||
return aggregate_lora_products(_ab_products(), dtype)
|
||||
|
||||
|
||||
def _fuse_delta_with_scaled_fp8(
|
||||
deltas: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
key: str,
|
||||
scale_key: str,
|
||||
model_sd: StateDict,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Dequantize scaled FP8 weight, add LoRA delta, and re-quantize."""
|
||||
weight_scale = model_sd.sd[scale_key]
|
||||
|
||||
original_weight = weight.to(torch.float32) * weight_scale
|
||||
|
||||
new_weight = original_weight + deltas.to(torch.float32)
|
||||
|
||||
new_fp8_weight, new_weight_scale = quantize_weight_to_fp8_per_tensor(new_weight)
|
||||
return {key: new_fp8_weight, scale_key: new_weight_scale}
|
||||
for lsd, coef in lora_sd_and_strengths:
|
||||
if key_a not in lsd.sd or key_b not in lsd.sd:
|
||||
continue
|
||||
a = lsd.sd[key_a].to(device=device, dtype=dtype, non_blocking=True)
|
||||
b = lsd.sd[key_b].to(device=device, dtype=dtype, non_blocking=True)
|
||||
yield LoraProduct(a, b, coef)
|
||||
|
||||
@@ -1,17 +1,21 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, NamedTuple, Protocol
|
||||
from typing import TYPE_CHECKING, Any, NamedTuple, Protocol, TypeVar
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.model.model_protocol import ModelType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from typing_extensions import Self
|
||||
|
||||
from ltx_core.loader.fuse_loras import FuseRule
|
||||
from ltx_core.loader.registry import Registry
|
||||
|
||||
BuiltType = TypeVar("BuiltType", covariant=True) # noqa: PLC0105
|
||||
|
||||
|
||||
# Per-key shape and dtype description for a flat collection of tensors.
|
||||
TensorLayout = dict[str, tuple[torch.Size, torch.dtype]]
|
||||
@@ -49,70 +53,75 @@ class StateDictLoader(Protocol):
|
||||
"""
|
||||
Load metadata from path
|
||||
"""
|
||||
...
|
||||
|
||||
def load(self, path: str | list[str], sd_ops: SDOps | None = None, device: torch.device | None = None) -> StateDict:
|
||||
"""
|
||||
Load state dict from path or paths (for sharded model storage) and apply sd_ops
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
class BuilderProtocol(Protocol[ModelType]):
|
||||
class BuilderProtocol(Protocol[BuiltType]):
|
||||
"""Protocol for model builders that produce a model via ``build()``."""
|
||||
|
||||
def build(
|
||||
self, device: torch.device | None = None, dtype: torch.dtype | None = None, **kwargs: object
|
||||
) -> ModelType: ...
|
||||
self,
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype | None = None,
|
||||
**kwargs: Any, # noqa: ANN401
|
||||
) -> BuiltType: ...
|
||||
|
||||
@property
|
||||
def registry(self) -> "Registry": ...
|
||||
|
||||
class ModelBuilderProtocol(BuilderProtocol[ModelType], Protocol[ModelType]):
|
||||
"""
|
||||
Protocol for building PyTorch models from configuration dictionaries.
|
||||
Implementations must provide:
|
||||
- meta_model: Create a model from configuration dictionary and apply module operations
|
||||
- build: Create and initialize a model from state dictionary and apply dtype transformations
|
||||
"""
|
||||
|
||||
model_sd_ops: SDOps | None
|
||||
module_ops: tuple[ModuleOps, ...]
|
||||
loras: tuple["LoraPathStrengthAndSDOps", ...]
|
||||
registry: "Registry"
|
||||
|
||||
def meta_model(self, config: dict, module_ops: list[ModuleOps] | None = None) -> ModelType:
|
||||
"""
|
||||
Create a model on the meta device from a configuration dictionary.
|
||||
This decouples model creation from weight loading, allowing the model
|
||||
architecture to be instantiated without allocating memory for parameters.
|
||||
Args:
|
||||
config: Model configuration dictionary.
|
||||
module_ops: Optional list of module operations to apply (e.g., quantization).
|
||||
Returns:
|
||||
Model instance on meta device (no actual memory allocated for parameters).
|
||||
"""
|
||||
...
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given state-dict key remapping ops."""
|
||||
...
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given module operations (e.g. quantization)."""
|
||||
...
|
||||
|
||||
def with_loras(self, loras: tuple["LoraPathStrengthAndSDOps", ...]) -> "ModelBuilderProtocol[ModelType]":
|
||||
"""Return a copy of this builder with the given LoRAs to fuse at build time."""
|
||||
...
|
||||
|
||||
def with_registry(self, registry: "Registry") -> "ModelBuilderProtocol[ModelType]":
|
||||
def with_registry(self, registry: "Registry") -> "Self":
|
||||
"""Return a copy of this builder using the given weight registry for allocation."""
|
||||
...
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> "ModelBuilderProtocol[ModelType]":
|
||||
|
||||
class ModelBuilderProtocol(BuilderProtocol[BuiltType], Protocol[BuiltType]):
|
||||
"""
|
||||
Protocol for building PyTorch models from configuration dictionaries.
|
||||
Implementations must provide:
|
||||
- build: Create and initialize a model from state dictionary and apply dtype transformations
|
||||
"""
|
||||
|
||||
@property
|
||||
def model_sd_ops(self) -> SDOps | None: ...
|
||||
|
||||
@property
|
||||
def module_ops(self) -> tuple[ModuleOps, ...]: ...
|
||||
|
||||
@property
|
||||
def loras(self) -> tuple["LoraPathStrengthAndSDOps", ...]: ...
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> "Self":
|
||||
"""Return a copy of this builder with the given state-dict key remapping ops."""
|
||||
...
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "Self":
|
||||
"""Return a copy of this builder with the given module operations (e.g. quantization)."""
|
||||
...
|
||||
|
||||
def with_loras(self, loras: tuple["LoraPathStrengthAndSDOps", ...]) -> "Self":
|
||||
"""Return a copy of this builder with the given LoRAs to fuse at build time."""
|
||||
...
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> "Self":
|
||||
"""Return a copy of this builder that loads LoRA weights onto the given device."""
|
||||
...
|
||||
|
||||
def with_fuse_rule(self, fuse_rule: "FuseRule") -> "Self":
|
||||
"""Return a copy of this builder with the given LoRA fuse rule (e.g. from a quantization policy)."""
|
||||
...
|
||||
|
||||
def build(
|
||||
self, device: torch.device | None = None, dtype: torch.dtype | None = None, **kwargs: object
|
||||
) -> ModelType:
|
||||
self,
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype | None = None,
|
||||
**kwargs: Any, # noqa: ANN401
|
||||
) -> BuiltType:
|
||||
"""
|
||||
Build the model
|
||||
Args:
|
||||
@@ -135,8 +144,7 @@ class LoRAAdaptableProtocol(Protocol):
|
||||
- lora: Add a LoRA to the model
|
||||
"""
|
||||
|
||||
def lora(self, lora_path: str, strength: float) -> "LoRAAdaptableProtocol":
|
||||
pass
|
||||
def lora(self, lora_path: str, strength: float, sd_ops: SDOps) -> "LoRAAdaptableProtocol": ...
|
||||
|
||||
|
||||
class LoraPathStrengthAndSDOps(NamedTuple):
|
||||
|
||||
@@ -24,6 +24,7 @@ class ContentMatching:
|
||||
|
||||
prefix: str = ""
|
||||
suffix: str = ""
|
||||
contains: str = ""
|
||||
|
||||
|
||||
class KeyValueOperationResult(NamedTuple):
|
||||
@@ -72,10 +73,10 @@ class SDOps:
|
||||
new_mapping = (*self.mapping, ContentReplacement(content, replacement))
|
||||
return replace(self, mapping=new_mapping)
|
||||
|
||||
def with_matching(self, prefix: str = "", suffix: str = "") -> "SDOps":
|
||||
"""Create a new SDOps instance with the specified prefix and suffix matching added to the mapping."""
|
||||
def with_matching(self, prefix: str = "", suffix: str = "", contains: str = "") -> "SDOps":
|
||||
"""Create a new SDOps instance with the specified prefix, suffix and contains matching added to the mapping."""
|
||||
|
||||
new_mapping = (*self.mapping, ContentMatching(prefix, suffix))
|
||||
new_mapping = (*self.mapping, ContentMatching(prefix, suffix, contains))
|
||||
return replace(self, mapping=new_mapping)
|
||||
|
||||
def with_additional_allowed_keys(self, keys: frozenset[str]) -> "SDOps":
|
||||
@@ -100,7 +101,12 @@ class SDOps:
|
||||
def apply_to_key(self, key: str) -> str | None:
|
||||
"""Apply the mapping to the given name."""
|
||||
matchers = [content for content in self.mapping if isinstance(content, ContentMatching)]
|
||||
valid = any(key.startswith(f.prefix) and key.endswith(f.suffix) for f in matchers)
|
||||
valid = any(
|
||||
key.startswith(matcher.prefix)
|
||||
and key.endswith(matcher.suffix)
|
||||
and (not matcher.contains or matcher.contains in key)
|
||||
for matcher in matchers
|
||||
)
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
|
||||
@@ -1,11 +1,13 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import logging
|
||||
from dataclasses import dataclass, field, replace
|
||||
from typing import Generic
|
||||
from typing import TYPE_CHECKING, Final, Generic
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ltx_core.loader.fuse_loras import apply_loras
|
||||
from ltx_core.loader.fuse_loras import FuseRule, apply_loras, bf16_fuse_rule
|
||||
from ltx_core.loader.helpers import create_meta_model, load_state_dict, read_model_config
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import (
|
||||
@@ -21,6 +23,9 @@ from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.loader.sft_loader import SafetensorsModelStateDictLoader
|
||||
from ltx_core.model.model_protocol import ModelConfigurator, ModelType
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from typing_extensions import Self
|
||||
|
||||
logger: logging.Logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -46,6 +51,7 @@ def _load_model_weights(
|
||||
dtype: torch.dtype | None,
|
||||
model_sd_ops: SDOps | None = None,
|
||||
lora_load_device: torch.device | None = None,
|
||||
fuse_rule: FuseRule = bf16_fuse_rule,
|
||||
) -> None:
|
||||
"""Load base weights and fuse LoRAs into *meta_model* in-place."""
|
||||
if lora_load_device is None:
|
||||
@@ -57,7 +63,7 @@ def _load_model_weights(
|
||||
if not lora_strengths or (min(lora_strengths) == 0 and max(lora_strengths) == 0):
|
||||
sd = model_sd.sd
|
||||
if dtype is not None:
|
||||
sd = {key: value.to(dtype=dtype) for key, value in model_sd.sd.items()}
|
||||
sd = {key: value.to(dtype=dtype) for key, value in sd.items()}
|
||||
meta_model.load_state_dict(sd, strict=False, assign=True)
|
||||
return
|
||||
|
||||
@@ -68,16 +74,21 @@ def _load_model_weights(
|
||||
final_sd = apply_loras(
|
||||
model_sd=model_sd,
|
||||
lora_sd_and_strengths=lora_sd_and_strengths,
|
||||
dtype=dtype,
|
||||
fuse_rule=fuse_rule,
|
||||
destination_sd=model_sd if isinstance(registry, DummyRegistry) else None,
|
||||
)
|
||||
meta_model.load_state_dict(final_sd.sd, strict=False, assign=True)
|
||||
fused_sd = final_sd.sd
|
||||
if dtype is not None:
|
||||
fused_sd = {key: value.to(dtype=dtype) for key, value in fused_sd.items()}
|
||||
meta_model.load_state_dict(fused_sd, strict=False, assign=True)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType], LoRAAdaptableProtocol):
|
||||
"""
|
||||
Builder for PyTorch models residing on a single GPU.
|
||||
The builder is immutable: ``with_*``/``lora`` return modified copies. The
|
||||
``ModelBuilderProtocol`` surface is exposed via read-only properties backed
|
||||
by private attributes.
|
||||
Attributes:
|
||||
model_class_configurator: Class responsible for constructing the model from a config dict.
|
||||
model_path: Path (or tuple of shard paths) to the model's `.safetensors` checkpoint(s).
|
||||
@@ -91,52 +102,109 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType],
|
||||
``torch.device("cpu")``, which keeps LoRA weights in CPU memory and transfers them to
|
||||
the target GPU sequentially during fusion, reducing peak GPU memory usage compared to
|
||||
loading all LoRA weights directly onto the GPU at once.
|
||||
fuse_rule: Per-policy LoRA merge rule. Defaults to ``bf16_fuse_rule``;
|
||||
"""
|
||||
|
||||
model_class_configurator: type[ModelConfigurator[ModelType]]
|
||||
model_path: str | tuple[str, ...]
|
||||
model_sd_ops: SDOps | None = None
|
||||
module_ops: tuple[ModuleOps, ...] = field(default_factory=tuple)
|
||||
loras: tuple[LoraPathStrengthAndSDOps, ...] = field(default_factory=tuple)
|
||||
model_loader: StateDictLoader = field(default_factory=SafetensorsModelStateDictLoader)
|
||||
registry: Registry = field(default_factory=DummyRegistry)
|
||||
lora_load_device: torch.device = field(default_factory=lambda: torch.device("cpu"))
|
||||
def __init__(
|
||||
self,
|
||||
model_class_configurator: type[ModelConfigurator[ModelType]],
|
||||
model_path: str | tuple[str, ...],
|
||||
model_sd_ops: SDOps | None = None,
|
||||
module_ops: tuple[ModuleOps, ...] = (),
|
||||
loras: tuple[LoraPathStrengthAndSDOps, ...] = (),
|
||||
model_loader: StateDictLoader | None = None,
|
||||
registry: Registry | None = None,
|
||||
lora_load_device: torch.device | None = None,
|
||||
fuse_rule: FuseRule = bf16_fuse_rule,
|
||||
) -> None:
|
||||
# Read-only: typed with the covariant ModelType, so it must not be a mutable attribute.
|
||||
self._model_class_configurator: Final = model_class_configurator
|
||||
self._model_path = model_path
|
||||
self._model_sd_ops = model_sd_ops
|
||||
self._module_ops = module_ops
|
||||
self._loras = loras
|
||||
self._model_loader = model_loader if model_loader is not None else SafetensorsModelStateDictLoader()
|
||||
self._registry = registry if registry is not None else DummyRegistry()
|
||||
self._lora_load_device = lora_load_device if lora_load_device is not None else torch.device("cpu")
|
||||
self._fuse_rule = fuse_rule
|
||||
|
||||
def lora(self, lora_path: str, strength: float, sd_ops: SDOps) -> "SingleGPUModelBuilder":
|
||||
return replace(self, loras=(*self.loras, LoraPathStrengthAndSDOps(lora_path, strength, sd_ops)))
|
||||
@property
|
||||
def model_sd_ops(self) -> SDOps | None:
|
||||
return self._model_sd_ops
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> "SingleGPUModelBuilder":
|
||||
return replace(self, model_sd_ops=sd_ops)
|
||||
@property
|
||||
def module_ops(self) -> tuple[ModuleOps, ...]:
|
||||
return self._module_ops
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> "SingleGPUModelBuilder":
|
||||
return replace(self, module_ops=module_ops)
|
||||
@property
|
||||
def loras(self) -> tuple[LoraPathStrengthAndSDOps, ...]:
|
||||
return self._loras
|
||||
|
||||
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> "SingleGPUModelBuilder":
|
||||
return replace(self, loras=loras)
|
||||
@property
|
||||
def registry(self) -> Registry:
|
||||
return self._registry
|
||||
|
||||
def with_registry(self, registry: Registry) -> "SingleGPUModelBuilder":
|
||||
return replace(self, registry=registry)
|
||||
@property
|
||||
def model_path(self) -> str | tuple[str, ...]:
|
||||
return self._model_path
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> "SingleGPUModelBuilder":
|
||||
return replace(self, lora_load_device=device)
|
||||
@property
|
||||
def model_loader(self) -> StateDictLoader:
|
||||
return self._model_loader
|
||||
|
||||
@property
|
||||
def lora_load_device(self) -> torch.device:
|
||||
return self._lora_load_device
|
||||
|
||||
@property
|
||||
def fuse_rule(self) -> FuseRule:
|
||||
return self._fuse_rule
|
||||
|
||||
def lora(self, lora_path: str, strength: float, sd_ops: SDOps) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._loras = (*self._loras, LoraPathStrengthAndSDOps(lora_path, strength, sd_ops))
|
||||
return clone
|
||||
|
||||
def with_sd_ops(self, sd_ops: SDOps | None) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._model_sd_ops = sd_ops
|
||||
return clone
|
||||
|
||||
def with_module_ops(self, module_ops: tuple[ModuleOps, ...]) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._module_ops = module_ops
|
||||
return clone
|
||||
|
||||
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._loras = loras
|
||||
return clone
|
||||
|
||||
def with_registry(self, registry: Registry) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._registry = registry
|
||||
return clone
|
||||
|
||||
def with_lora_load_device(self, device: torch.device) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._lora_load_device = device
|
||||
return clone
|
||||
|
||||
def with_fuse_rule(self, fuse_rule: FuseRule) -> Self:
|
||||
clone = copy.copy(self)
|
||||
clone._fuse_rule = fuse_rule
|
||||
return clone
|
||||
|
||||
def model_config(self) -> dict:
|
||||
return read_model_config(self.model_path, self.model_loader)
|
||||
return read_model_config(self._model_path, self._model_loader)
|
||||
|
||||
def meta_model(self, config: dict, module_ops: tuple[ModuleOps, ...]) -> ModelType:
|
||||
return create_meta_model(self.model_class_configurator, config, module_ops)
|
||||
return create_meta_model(self._model_class_configurator, config, module_ops)
|
||||
|
||||
def load_sd(
|
||||
self, paths: list[str], registry: Registry, device: torch.device | None, sd_ops: SDOps | None = None
|
||||
) -> StateDict:
|
||||
return load_state_dict(paths, self.model_loader, registry, device, sd_ops)
|
||||
|
||||
def _return_model(self, meta_model: ModelType, device: torch.device) -> ModelType:
|
||||
uninitialized = _check_uninitialized(meta_model)
|
||||
if uninitialized:
|
||||
logger.warning(f"Uninitialized parameters or buffers: {uninitialized}")
|
||||
return meta_model
|
||||
return meta_model.to(device)
|
||||
return load_state_dict(paths, self._model_loader, registry, device, sd_ops)
|
||||
|
||||
def build(
|
||||
self,
|
||||
@@ -146,17 +214,23 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType],
|
||||
) -> ModelType:
|
||||
device = torch.device("cuda") if device is None else device
|
||||
config = self.model_config()
|
||||
meta_model = self.meta_model(config, self.module_ops)
|
||||
meta_model = self.meta_model(config, self._module_ops)
|
||||
|
||||
_load_model_weights(
|
||||
meta_model=meta_model,
|
||||
model_path=self.model_path,
|
||||
loras=self.loras,
|
||||
loader=self.model_loader,
|
||||
registry=self.registry,
|
||||
model_path=self._model_path,
|
||||
loras=self._loras,
|
||||
loader=self._model_loader,
|
||||
registry=self._registry,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
model_sd_ops=self.model_sd_ops,
|
||||
lora_load_device=self.lora_load_device,
|
||||
model_sd_ops=self._model_sd_ops,
|
||||
lora_load_device=self._lora_load_device,
|
||||
fuse_rule=self._fuse_rule,
|
||||
)
|
||||
return self._return_model(meta_model, device)
|
||||
|
||||
uninitialized = _check_uninitialized(meta_model)
|
||||
if uninitialized:
|
||||
logger.warning(f"Uninitialized parameters or buffers: {uninitialized}")
|
||||
return meta_model
|
||||
return meta_model.to(device)
|
||||
|
||||
@@ -21,11 +21,14 @@ from ltx_core.types import VideoLatentShape
|
||||
@dataclass(frozen=True)
|
||||
class TilingContext:
|
||||
"""Opaque context produced by :meth:`VideoModalityTilingHelper.tile_modality`.
|
||||
Carries the token-level keep mask and per-conditioning-token blend
|
||||
Carries the token-level keep indices and per-conditioning-token blend
|
||||
weights needed by :meth:`~VideoModalityTilingHelper.blend`.
|
||||
"""
|
||||
|
||||
keep_mask: torch.Tensor
|
||||
keep_indices: torch.Tensor
|
||||
"""``(num_kept,)`` int64 — sorted indices of tokens the tile processes."""
|
||||
num_total_tokens: int
|
||||
"""Total number of tokens in the full (untiled) sequence."""
|
||||
cond_blend_weights: torch.Tensor | None
|
||||
"""``(num_kept_cond,)`` — weight for each kept conditioning token,
|
||||
equal to ``1 / num_tiles_that_keep_this_token``. ``None`` when
|
||||
@@ -81,14 +84,32 @@ class VideoModalityTilingHelper:
|
||||
A ``(tiled_modality, context)`` tuple. Pass *context* to
|
||||
:meth:`blend` together with the model output.
|
||||
"""
|
||||
keep_mask = self._keep_mask(modality, tile)
|
||||
device = modality.positions.device
|
||||
gen_indices = self._generated_token_indices(tile, device=device)
|
||||
num_total = modality.latent.shape[1]
|
||||
|
||||
cond_blend_weights: torch.Tensor | None = None
|
||||
if num_total > self._num_generated_tokens:
|
||||
keep_per_tile_cond = self._all_tiles_cond_keep(modality) # (num_tiles, num_cond) bool
|
||||
tile_idx = next((i for i, t in enumerate(self._tiles) if t.in_coords == tile.in_coords), None)
|
||||
if tile_idx is None:
|
||||
raise ValueError(
|
||||
f"Tile with in_coords={tile.in_coords} is not in this helper's tile set; "
|
||||
f"pass a tile obtained from `helper.tiles`."
|
||||
)
|
||||
my_cond_keep = keep_per_tile_cond[tile_idx]
|
||||
cond_indices = self._num_generated_tokens + my_cond_keep.nonzero(as_tuple=False).squeeze(1)
|
||||
keep_indices = torch.cat([gen_indices, cond_indices])
|
||||
total_keepers = keep_per_tile_cond.sum(dim=0).float() # (num_cond,)
|
||||
cond_blend_weights = 1.0 / total_keepers[my_cond_keep]
|
||||
else:
|
||||
keep_indices = gen_indices
|
||||
|
||||
tile_attention_mask = None
|
||||
if modality.attention_mask is not None:
|
||||
keep_indices = keep_mask.nonzero(as_tuple=False).squeeze(1)
|
||||
tile_attention_mask = modality.attention_mask[:, keep_indices, :][:, :, keep_indices]
|
||||
|
||||
positions = modality.positions[:, :, keep_mask, :]
|
||||
positions = modality.positions[:, :, keep_indices, :]
|
||||
if normalize_positions:
|
||||
num_tile_gen = self._tile_generated_token_count(tile)
|
||||
gen_pos = positions[:, :, :num_tile_gen, :] # (B, 3, num_tile_gen, 2)
|
||||
@@ -97,26 +118,15 @@ class VideoModalityTilingHelper:
|
||||
|
||||
tiled = replace(
|
||||
modality,
|
||||
latent=modality.latent[:, keep_mask, :],
|
||||
timesteps=modality.timesteps[:, keep_mask],
|
||||
latent=modality.latent[:, keep_indices, :],
|
||||
timesteps=modality.timesteps[:, keep_indices],
|
||||
positions=positions,
|
||||
attention_mask=tile_attention_mask,
|
||||
)
|
||||
|
||||
cond_blend_weights = None
|
||||
num_total = modality.latent.shape[1]
|
||||
if num_total > self._num_generated_tokens:
|
||||
cond_keep = keep_mask[self._num_generated_tokens :]
|
||||
# Count how many tiles keep each conditioning token.
|
||||
cond_counts = torch.zeros(cond_keep.sum(), dtype=torch.float32)
|
||||
for t in self._tiles:
|
||||
other_mask = self._keep_mask(modality, t)
|
||||
other_cond = other_mask[self._num_generated_tokens :]
|
||||
# Map other tile's kept cond tokens into this tile's kept subset.
|
||||
cond_counts += other_cond[cond_keep].float()
|
||||
cond_blend_weights = 1.0 / cond_counts
|
||||
|
||||
return tiled, TilingContext(keep_mask=keep_mask, cond_blend_weights=cond_blend_weights)
|
||||
return tiled, TilingContext(
|
||||
keep_indices=keep_indices, num_total_tokens=num_total, cond_blend_weights=cond_blend_weights
|
||||
)
|
||||
|
||||
# -- blend -------------------------------------------------------------
|
||||
|
||||
@@ -147,9 +157,9 @@ class VideoModalityTilingHelper:
|
||||
"""
|
||||
batch, _, dim = tile_to_blend.shape
|
||||
num_tile_gen = self._tile_generated_token_count(tile)
|
||||
gen_indices = self._generated_token_indices(tile)
|
||||
gen_indices = self._generated_token_indices(tile, device=tile_to_blend.device)
|
||||
|
||||
num_total_tokens = context.keep_mask.shape[0]
|
||||
num_total_tokens = context.num_total_tokens
|
||||
expected_shape = (batch, num_total_tokens, dim)
|
||||
|
||||
if output is not None:
|
||||
@@ -168,8 +178,7 @@ class VideoModalityTilingHelper:
|
||||
# Scatter kept conditioning tokens, weighted by 1/N where N is
|
||||
# the number of tiles that keep each token (so they sum to 1).
|
||||
if num_total_tokens > self._num_generated_tokens and context.cond_blend_weights is not None:
|
||||
cond_keep = context.keep_mask[self._num_generated_tokens :]
|
||||
cond_indices = self._num_generated_tokens + cond_keep.nonzero(as_tuple=False).squeeze(1)
|
||||
cond_indices = context.keep_indices[context.keep_indices >= self._num_generated_tokens]
|
||||
weights = context.cond_blend_weights.to(device=tile_to_blend.device, dtype=tile_to_blend.dtype)
|
||||
result[:, cond_indices, :] += tile_to_blend[:, num_tile_gen:, :] * weights[None, :, None]
|
||||
|
||||
@@ -189,46 +198,42 @@ class VideoModalityTilingHelper:
|
||||
)
|
||||
return self._patchifier.get_token_count(tile_shape)
|
||||
|
||||
def _generated_token_indices(self, tile: Tile) -> torch.Tensor:
|
||||
def _generated_token_indices(self, tile: Tile, device: torch.device | None = None) -> torch.Tensor:
|
||||
"""Flat token indices of *tile*'s generated tokens in the full sequence."""
|
||||
frame_slice, height_slice, width_slice = tile.in_coords
|
||||
f = torch.arange(frame_slice.start, frame_slice.stop)
|
||||
h = torch.arange(height_slice.start, height_slice.stop)
|
||||
w = torch.arange(width_slice.start, width_slice.stop)
|
||||
f = torch.arange(frame_slice.start, frame_slice.stop, device=device)
|
||||
h = torch.arange(height_slice.start, height_slice.stop, device=device)
|
||||
w = torch.arange(width_slice.start, width_slice.stop, device=device)
|
||||
return (
|
||||
f[:, None, None] * self._latent_shape.height * self._latent_shape.width
|
||||
+ h[None, :, None] * self._latent_shape.width
|
||||
+ w[None, None, :]
|
||||
).reshape(-1)
|
||||
|
||||
def _keep_mask(self, modality: Modality, tile: Tile) -> torch.Tensor:
|
||||
"""Boolean mask ``(num_total_tokens,)`` — True for tokens the tile processes.
|
||||
Generated tokens are selected by grid position. Conditioning
|
||||
tokens are kept when their ``[start, end)`` intervals overlap
|
||||
the tile in all three dimensions, or when they have a negative
|
||||
time coordinate (reference tokens).
|
||||
def _all_tiles_cond_keep(self, modality: Modality) -> torch.Tensor:
|
||||
"""Vectorized (num_tiles, num_cond) bool: which tiles keep each conditioning token.
|
||||
A conditioning token is kept by a tile when its ``[start, end)`` interval
|
||||
overlaps the tile in all three dimensions, or when it has a negative time
|
||||
coordinate (reference token).
|
||||
"""
|
||||
num_total = modality.latent.shape[1]
|
||||
mask = torch.zeros(num_total, dtype=torch.bool)
|
||||
cond_positions = modality.positions[:, :, self._num_generated_tokens :, :] # (B, 3, num_cond, 2)
|
||||
device = cond_positions.device
|
||||
|
||||
gen_indices = self._generated_token_indices(tile)
|
||||
mask[gen_indices] = True
|
||||
# Per-tile (start, end) bounds along each axis; small Python loop (num_tiles <= ~16).
|
||||
starts_list: list[torch.Tensor] = []
|
||||
ends_list: list[torch.Tensor] = []
|
||||
for t in self._tiles:
|
||||
gen_idx = self._generated_token_indices(t, device=device)
|
||||
gen_positions = modality.positions[:, :, gen_idx, :] # (B, 3, num_tile_gen, 2)
|
||||
starts_list.append(gen_positions[..., 0].amin(dim=2)) # (B, 3)
|
||||
ends_list.append(gen_positions[..., 1].amax(dim=2)) # (B, 3)
|
||||
tile_starts = torch.stack(starts_list, dim=0) # (num_tiles, B, 3)
|
||||
tile_ends = torch.stack(ends_list, dim=0) # (num_tiles, B, 3)
|
||||
|
||||
if num_total > self._num_generated_tokens:
|
||||
gen_positions = modality.positions[:, :, gen_indices, :] # (B, 3, num_tile_gen, 2)
|
||||
tile_start = gen_positions[..., 0].amin(dim=2) # (B, 3)
|
||||
tile_end = gen_positions[..., 1].amax(dim=2) # (B, 3)
|
||||
|
||||
cond_positions = modality.positions[:, :, self._num_generated_tokens :, :] # (B, 3, num_cond, 2)
|
||||
|
||||
overlaps = (cond_positions[..., 0] < tile_end.unsqueeze(2)) & (
|
||||
cond_positions[..., 1] > tile_start.unsqueeze(2)
|
||||
) # (B, 3, num_cond)
|
||||
overlaps_all_dims = overlaps.all(dim=1) # (B, num_cond)
|
||||
|
||||
has_negative_time = cond_positions[:, 0, :, 0] < 0 # (B, num_cond)
|
||||
|
||||
keep_cond = (overlaps_all_dims | has_negative_time).any(dim=0) # (num_cond,)
|
||||
mask[self._num_generated_tokens :] = keep_cond
|
||||
|
||||
return mask
|
||||
cond_starts = cond_positions[..., 0] # (B, 3, num_cond)
|
||||
cond_ends = cond_positions[..., 1] # (B, 3, num_cond)
|
||||
# Broadcast: (1, B, 3, num_cond) vs (num_tiles, B, 3, 1) -> (num_tiles, B, 3, num_cond).
|
||||
overlaps = (cond_starts[None] < tile_ends[..., None]) & (cond_ends[None] > tile_starts[..., None])
|
||||
overlaps_all_dims = overlaps.all(dim=2) # (num_tiles, B, num_cond)
|
||||
has_negative_time = (cond_positions[:, 0, :, 0] < 0)[None] # (1, B, num_cond)
|
||||
return (overlaps_all_dims | has_negative_time).any(dim=1) # (num_tiles, num_cond)
|
||||
|
||||
@@ -1,6 +1,14 @@
|
||||
from typing import Protocol, TypeVar
|
||||
from __future__ import annotations
|
||||
|
||||
ModelType = TypeVar("ModelType")
|
||||
from typing import TYPE_CHECKING, Protocol, TypeVar
|
||||
|
||||
import torch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
|
||||
ModelType = TypeVar("ModelType", covariant=True, bound=torch.nn.Module) # noqa: PLC0105
|
||||
|
||||
|
||||
class ModelConfigurator(Protocol[ModelType]):
|
||||
@@ -8,3 +16,24 @@ class ModelConfigurator(Protocol[ModelType]):
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict) -> ModelType: ...
|
||||
|
||||
|
||||
class LTXModelProtocol(Protocol):
|
||||
"""Velocity-model forward interface shared by ``LTXModel`` and its multi-GPU wrappers.
|
||||
``forward`` pins the real signature (enforced structurally); ``__call__`` mirrors it
|
||||
so protocol-typed values stay callable via ``model(...)``.
|
||||
"""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
video: Modality | None,
|
||||
audio: Modality | None,
|
||||
perturbations: BatchedPerturbationConfig,
|
||||
) -> tuple[torch.Tensor | None, torch.Tensor | None]: ...
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
video: Modality | None,
|
||||
audio: Modality | None,
|
||||
perturbations: BatchedPerturbationConfig,
|
||||
) -> tuple[torch.Tensor | None, torch.Tensor | None]: ...
|
||||
|
||||
@@ -3,13 +3,17 @@
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.model.transformer.model import LTXModel, X0Model
|
||||
from ltx_core.model.transformer.model_configurator import (
|
||||
LTXV_AUDIO_ONLY_MODEL_COMFY_RENAMING_MAP,
|
||||
LTXV_MODEL_COMFY_RENAMING_MAP,
|
||||
LTXAudioOnlyModelConfigurator,
|
||||
LTXModelConfigurator,
|
||||
LTXVideoOnlyModelConfigurator,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LTXV_AUDIO_ONLY_MODEL_COMFY_RENAMING_MAP",
|
||||
"LTXV_MODEL_COMFY_RENAMING_MAP",
|
||||
"LTXAudioOnlyModelConfigurator",
|
||||
"LTXModel",
|
||||
"LTXModelConfigurator",
|
||||
"LTXVideoOnlyModelConfigurator",
|
||||
|
||||
@@ -1,12 +1,35 @@
|
||||
import functools
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Protocol
|
||||
|
||||
import torch
|
||||
from torch.nn.attention import SDPBackend, sdpa_kernel
|
||||
|
||||
from ltx_core.model.transformer.ops import (
|
||||
GatedAttentionCallable,
|
||||
PreAttentionCallable,
|
||||
PytorchGatedAttention,
|
||||
PytorchPreAttention,
|
||||
)
|
||||
from ltx_core.model.transformer.rope import LTXRopeType
|
||||
|
||||
|
||||
def _torch_default_sdpa_priority() -> list[SDPBackend]:
|
||||
"""Fetch torch's current default SDPA priority order at runtime.
|
||||
Used as the default for ``PytorchAttention`` so the wrapper-always
|
||||
code path matches torch's native dispatch order without hard-coding it
|
||||
(which would drift if torch updates the default).
|
||||
``torch._C._get_sdp_priority_order`` is a private API; we accept that
|
||||
risk because the project pins ``torch`` in the lockfile, so any
|
||||
rename/removal surfaces on a controlled torch bump rather than silently.
|
||||
"""
|
||||
return [SDPBackend(p) for p in torch._C._get_sdp_priority_order()]
|
||||
|
||||
from ltx_core.model.transformer.rope import LTXRopeType, apply_rotary_emb
|
||||
|
||||
memory_efficient_attention = None
|
||||
flash_attn_interface = None
|
||||
flash_attn_4_func = None
|
||||
try:
|
||||
from xformers.ops import memory_efficient_attention
|
||||
except ImportError:
|
||||
@@ -17,15 +40,44 @@ try:
|
||||
import flash_attn_interface
|
||||
except ImportError:
|
||||
flash_attn_interface = None
|
||||
try:
|
||||
from flash_attn.cute import flash_attn_func as flash_attn_4_func
|
||||
except ImportError:
|
||||
flash_attn_4_func = None
|
||||
|
||||
|
||||
class AttentionCallable(Protocol):
|
||||
"""Unmasked attention. Backends without a mask kernel (FA3/FA4) implement only
|
||||
this protocol; backends that support masks too (Pytorch/SDPA, xFormers) are
|
||||
structurally usable here and as :class:`MaskedAttentionCallable`."""
|
||||
|
||||
def __call__(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int) -> torch.Tensor: ...
|
||||
|
||||
|
||||
class MaskedAttentionCallable(Protocol):
|
||||
"""Masked attention. Mask is required (not optional) -- the caller has already
|
||||
decided this is the masked path and chosen a backend that can serve it. Used
|
||||
by :class:`Attention` when its forward receives a non-None ``mask``."""
|
||||
|
||||
def __call__(
|
||||
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int, mask: torch.Tensor | None = None
|
||||
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int, mask: torch.Tensor
|
||||
) -> torch.Tensor: ...
|
||||
|
||||
|
||||
class PytorchAttention(AttentionCallable):
|
||||
def __init__(self, priority: list[SDPBackend] | None = None) -> None:
|
||||
# priority=None -> snapshot torch's default SDPA priority at construction.
|
||||
# Always passed through ``sdpa_kernel(..., set_priority=True)`` so the
|
||||
# call site is uniform regardless of how the priority was chosen.
|
||||
self._priority = priority if priority is not None else _torch_default_sdpa_priority()
|
||||
|
||||
@property
|
||||
def label(self) -> str:
|
||||
"""Human-readable identifier for this backend. Encodes the SDPA priority
|
||||
list so a single-backend pin reads differently from the full-priority
|
||||
dispatcher walk."""
|
||||
return f"SDPA[{'>'.join(b.name for b in self._priority)}]"
|
||||
|
||||
def __call__(
|
||||
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int, mask: torch.Tensor | None = None
|
||||
) -> torch.Tensor:
|
||||
@@ -41,12 +93,17 @@ class PytorchAttention(AttentionCallable):
|
||||
if mask.ndim == 3:
|
||||
mask = mask.unsqueeze(1)
|
||||
|
||||
out = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False)
|
||||
with sdpa_kernel(self._priority, set_priority=True):
|
||||
out = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v, attn_mask=mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
|
||||
return out
|
||||
|
||||
|
||||
class XFormersAttention(AttentionCallable):
|
||||
label = "xFormers"
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
@@ -92,13 +149,14 @@ class XFormersAttention(AttentionCallable):
|
||||
|
||||
|
||||
class FlashAttention3(AttentionCallable):
|
||||
label = "FlashAttention3"
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
heads: int,
|
||||
mask: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
if flash_attn_interface is None:
|
||||
raise RuntimeError("FlashAttention3 was selected but `FlashAttention3` is not installed.")
|
||||
@@ -108,32 +166,276 @@ class FlashAttention3(AttentionCallable):
|
||||
|
||||
q, k, v = (t.view(b, -1, heads, dim_head) for t in (q, k, v))
|
||||
|
||||
if mask is not None:
|
||||
raise NotImplementedError("Mask is not supported for FlashAttention3")
|
||||
|
||||
out = flash_attn_interface.flash_attn_func(q.to(v.dtype), k.to(v.dtype), v)
|
||||
out = out.reshape(b, -1, heads * dim_head)
|
||||
return out
|
||||
|
||||
|
||||
class FlashAttention4(AttentionCallable):
|
||||
label = "FlashAttention4"
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
heads: int,
|
||||
) -> torch.Tensor:
|
||||
if flash_attn_4_func is None:
|
||||
raise RuntimeError("FlashAttention4 was selected but `flash-attn-4` is not installed.")
|
||||
|
||||
b, _, dim_head = q.shape
|
||||
dim_head //= heads
|
||||
|
||||
q, k, v = (t.view(b, -1, heads, dim_head) for t in (q, k, v))
|
||||
|
||||
out, _ = flash_attn_4_func(q.to(v.dtype), k.to(v.dtype), v)
|
||||
out = out.reshape(b, -1, heads * dim_head)
|
||||
return out
|
||||
|
||||
|
||||
# --- Automatic selection -----------------------------------------------------
|
||||
# AUTOMATIC inspects installed extras and the GPU arch and returns the fastest
|
||||
# usable callable for each path. The selection runs once per process (cached).
|
||||
# The unmasked and masked picks are independent: each calls its own helper and
|
||||
# may end up on different backends (e.g. FA3 unmasked + xFormers masked on H100).
|
||||
|
||||
|
||||
def _sdpa_can_use(backend: SDPBackend, *, with_mask: bool) -> bool:
|
||||
"""Ask torch whether *backend* can run with the given mask shape.
|
||||
``MATH`` is the universal SDPA fallback (pure PyTorch ops, no kernel
|
||||
requirements) so it returns True everywhere, CPU included. The other
|
||||
backends use ``torch.backends.cuda.can_use_*`` capability checks (no GPU
|
||||
compute, no synchronization) and are False without CUDA. The probe shapes
|
||||
are small but realistic enough to surface constraints (head dim, dtype)
|
||||
that the per-backend rules care about.
|
||||
"""
|
||||
if backend is SDPBackend.MATH:
|
||||
return True
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
q = torch.empty(1, 4, 128, 64, device="cuda", dtype=torch.bfloat16)
|
||||
k = torch.empty(1, 4, 128, 64, device="cuda", dtype=torch.bfloat16)
|
||||
v = torch.empty(1, 4, 128, 64, device="cuda", dtype=torch.bfloat16)
|
||||
mask = torch.zeros(1, 4, 128, 128, device="cuda", dtype=torch.bfloat16) if with_mask else None
|
||||
params = torch.backends.cuda.SDPAParams(q, k, v, mask, 0.0, False, False)
|
||||
if backend is SDPBackend.CUDNN_ATTENTION:
|
||||
return torch.backends.cuda.can_use_cudnn_attention(params, debug=False)
|
||||
if backend is SDPBackend.FLASH_ATTENTION:
|
||||
return torch.backends.cuda.can_use_flash_attention(params, debug=False)
|
||||
if backend is SDPBackend.EFFICIENT_ATTENTION:
|
||||
return torch.backends.cuda.can_use_efficient_attention(params, debug=False)
|
||||
return False
|
||||
|
||||
|
||||
_SDPA_FULL_PRIORITY: tuple[SDPBackend, ...] = (
|
||||
SDPBackend.CUDNN_ATTENTION,
|
||||
SDPBackend.FLASH_ATTENTION,
|
||||
SDPBackend.EFFICIENT_ATTENTION,
|
||||
SDPBackend.MATH,
|
||||
)
|
||||
|
||||
|
||||
def _sdpa_full_priority() -> PytorchAttention:
|
||||
"""Hand SDPA the full backend priority order; let torch's dispatcher pick at call time.
|
||||
``sdpa_kernel(_SDPA_FULL_PRIORITY, set_priority=True)`` enables all four
|
||||
backends and orders them; torch then walks the order at call time and picks
|
||||
the first backend whose ``can_use_*`` check passes for the actual
|
||||
shapes/dtype/mask. FLASH is rejected automatically when a mask is present;
|
||||
CUDNN may be rejected under deterministic mode; MATH is the universal
|
||||
fallback. Probing per-backend usability up front from generic probe shapes
|
||||
cannot anticipate the variety of real call sites (e.g. broadcast key-only
|
||||
masks, large head dim), so we defer the choice to the dispatcher.
|
||||
"""
|
||||
return PytorchAttention(priority=list(_SDPA_FULL_PRIORITY))
|
||||
|
||||
|
||||
def _select_primary_attention() -> AttentionCallable:
|
||||
"""Pick the fastest unmasked attention based on installed extras and GPU arch.
|
||||
Priority by arch:
|
||||
- Hopper (sm_90, H100): FA3 / xFormers (mutually exclusive at import) > FA4 > SDPA.
|
||||
- Datacenter Blackwell (sm_100, B200): FA4 > SDPA. FA4 is intentionally *not*
|
||||
picked on consumer Blackwell (sm_120) -- known regressions in newer
|
||||
FA4 betas; users who want it on sm_120 must opt in explicitly.
|
||||
- Everywhere else (Ada, Ampere, CPU): SDPA with the full backend priority
|
||||
list -- torch's runtime dispatcher picks the best fit at call time.
|
||||
"""
|
||||
if torch.cuda.is_available():
|
||||
major, _ = torch.cuda.get_device_capability(0)
|
||||
if major == 9:
|
||||
if flash_attn_interface is not None:
|
||||
return FlashAttention3()
|
||||
if memory_efficient_attention is not None:
|
||||
return XFormersAttention()
|
||||
if flash_attn_4_func is not None:
|
||||
return FlashAttention4()
|
||||
if major == 10 and flash_attn_4_func is not None:
|
||||
return FlashAttention4()
|
||||
return _sdpa_full_priority()
|
||||
|
||||
|
||||
def _select_masked_attention() -> MaskedAttentionCallable:
|
||||
"""Pick a mask-aware attention. Prefers xFormers when installed; else SDPA with
|
||||
the full priority list (the dispatcher rejects FLASH automatically when a
|
||||
mask is present and walks past it)."""
|
||||
if memory_efficient_attention is not None:
|
||||
return XFormersAttention()
|
||||
return _sdpa_full_priority()
|
||||
|
||||
|
||||
@functools.cache
|
||||
def automatic_attention() -> AttentionCallable:
|
||||
"""Cached AUTOMATIC pick for the unmasked path.
|
||||
Cached so every ``AttentionOps`` in the process shares one instance."""
|
||||
return _select_primary_attention()
|
||||
|
||||
|
||||
@functools.cache
|
||||
def automatic_masked_attention() -> MaskedAttentionCallable:
|
||||
"""Cached AUTOMATIC pick for the masked path. See :func:`automatic_attention`."""
|
||||
return _select_masked_attention()
|
||||
|
||||
|
||||
def attention_label(fn: AttentionCallable | MaskedAttentionCallable) -> str:
|
||||
"""Best-effort human-readable backend name.
|
||||
Built-in callables expose ``.label`` (encoding the SDPA priority list for the
|
||||
Pytorch backends); fall back to the class name for custom or wrapped callables
|
||||
(e.g. the multi-GPU All2All wrappers) that don't define one."""
|
||||
return getattr(fn, "label", type(fn).__name__)
|
||||
|
||||
|
||||
def _resolve_sdpa_variant(backend: SDPBackend, name: str, *, with_mask: bool) -> PytorchAttention:
|
||||
"""Build a single-backend ``PytorchAttention`` pin, raising if the backend
|
||||
can't actually serve the call on this machine. Used by both
|
||||
:meth:`AttentionFunction.to_callable` and :meth:`MaskedAttentionFunction.to_callable`;
|
||||
``with_mask`` differs between the two so the capability check considers
|
||||
the protocol the caller intends to use. Not used for ``MATH`` -- MATH is
|
||||
the universal fallback and would falsely fail the CUDA-only probe on CPU.
|
||||
"""
|
||||
if not _sdpa_can_use(backend, with_mask=with_mask):
|
||||
raise RuntimeError(
|
||||
f"{name} selected but the SDPA {backend.name} backend is not usable on this machine "
|
||||
"(either no CUDA, the backend rejected the probe shapes, or "
|
||||
"torch.use_deterministic_algorithms(True) excluded it)."
|
||||
)
|
||||
return PytorchAttention(priority=[backend])
|
||||
|
||||
|
||||
class AttentionFunction(Enum):
|
||||
PYTORCH = "pytorch"
|
||||
XFORMERS = "xformers"
|
||||
FLASH_ATTENTION_3 = "flash_attention_3"
|
||||
DEFAULT = "default"
|
||||
FLASH_ATTENTION_4 = "flash_attention_4"
|
||||
SDPA_CUDNN = "sdpa_cudnn"
|
||||
SDPA_FLASH = "sdpa_flash"
|
||||
SDPA_EFFICIENT = "sdpa_efficient"
|
||||
SDPA_MATH = "sdpa_math"
|
||||
# Pick the fastest unmasked backend for the current GPU/extras combo; see
|
||||
# :func:`automatic_attention`. Default for :class:`AttentionOps`.
|
||||
AUTOMATIC = "automatic"
|
||||
|
||||
def to_callable(self) -> AttentionCallable:
|
||||
def to_callable(self) -> AttentionCallable: # noqa: PLR0911
|
||||
"""Resolve to a concrete callable. Use this at module init time so that
|
||||
torch.compile can trace through the attention call without graph breaks."""
|
||||
if self is AttentionFunction.PYTORCH:
|
||||
return PytorchAttention()
|
||||
elif self is AttentionFunction.XFORMERS:
|
||||
return XFormersAttention()
|
||||
elif self is AttentionFunction.FLASH_ATTENTION_3:
|
||||
return FlashAttention3()
|
||||
else:
|
||||
# Default behavior: XFormers if installed else - PyTorch
|
||||
return XFormersAttention() if memory_efficient_attention is not None else PytorchAttention()
|
||||
torch.compile can trace through the attention call without graph breaks.
|
||||
Every non-AUTOMATIC variant raises :class:`RuntimeError` when the backend
|
||||
isn't usable on this machine -- missing package or SDPA backend rejected
|
||||
on this hardware (e.g. cuDNN under ``torch.use_deterministic_algorithms``).
|
||||
Opting in means "this kernel or fail loudly". ``AUTOMATIC`` returns the
|
||||
cached :func:`automatic_attention` instance so every build shares one callable.
|
||||
"""
|
||||
match self:
|
||||
case AttentionFunction.AUTOMATIC:
|
||||
return automatic_attention()
|
||||
case AttentionFunction.PYTORCH:
|
||||
return PytorchAttention()
|
||||
case AttentionFunction.XFORMERS:
|
||||
if memory_efficient_attention is None:
|
||||
raise RuntimeError("AttentionFunction.XFORMERS selected but `xformers` is not installed.")
|
||||
return XFormersAttention()
|
||||
case AttentionFunction.FLASH_ATTENTION_3:
|
||||
if flash_attn_interface is None:
|
||||
raise RuntimeError(
|
||||
"AttentionFunction.FLASH_ATTENTION_3 selected but `flash-attn-3` is not installed."
|
||||
)
|
||||
return FlashAttention3()
|
||||
case AttentionFunction.FLASH_ATTENTION_4:
|
||||
if flash_attn_4_func is None:
|
||||
raise RuntimeError(
|
||||
"AttentionFunction.FLASH_ATTENTION_4 selected but `flash-attn-4` is not installed."
|
||||
)
|
||||
return FlashAttention4()
|
||||
case AttentionFunction.SDPA_MATH:
|
||||
return PytorchAttention(priority=[SDPBackend.MATH])
|
||||
case AttentionFunction.SDPA_CUDNN:
|
||||
return _resolve_sdpa_variant(
|
||||
SDPBackend.CUDNN_ATTENTION, "AttentionFunction.SDPA_CUDNN", with_mask=False
|
||||
)
|
||||
case AttentionFunction.SDPA_FLASH:
|
||||
return _resolve_sdpa_variant(
|
||||
SDPBackend.FLASH_ATTENTION, "AttentionFunction.SDPA_FLASH", with_mask=False
|
||||
)
|
||||
case AttentionFunction.SDPA_EFFICIENT:
|
||||
return _resolve_sdpa_variant(
|
||||
SDPBackend.EFFICIENT_ATTENTION, "AttentionFunction.SDPA_EFFICIENT", with_mask=False
|
||||
)
|
||||
|
||||
|
||||
class MaskedAttentionFunction(Enum):
|
||||
"""Backends usable on the masked path. Mirrors :class:`AttentionFunction` minus
|
||||
the variants the torch SDPA dispatcher (or the wrapped kernel) rejects with a
|
||||
mask: ``SDPA_FLASH`` -- FLASH kernel cannot serve an additive ``attn_mask``;
|
||||
``FLASH_ATTENTION_3``/``FLASH_ATTENTION_4`` -- neither has a mask kernel at all.
|
||||
Keeping them out makes "this backend cannot mask" a type error, not a runtime one."""
|
||||
|
||||
PYTORCH = "pytorch"
|
||||
XFORMERS = "xformers"
|
||||
SDPA_CUDNN = "sdpa_cudnn"
|
||||
SDPA_EFFICIENT = "sdpa_efficient"
|
||||
SDPA_MATH = "sdpa_math"
|
||||
# Pick the fastest mask-capable backend for the current extras combo; see
|
||||
# :func:`automatic_masked_attention`. Default for the masked slot of
|
||||
# :class:`AttentionOps`.
|
||||
AUTOMATIC = "automatic"
|
||||
|
||||
def to_callable(self) -> MaskedAttentionCallable:
|
||||
"""Resolve to a concrete masked callable. Same backend classes as
|
||||
:meth:`AttentionFunction.to_callable`; the protocol returned just exposes
|
||||
the masked call signature.
|
||||
Non-AUTOMATIC variants raise :class:`RuntimeError` when the backend isn't
|
||||
usable for the masked path on this machine. SDPA probes run with
|
||||
``with_mask=True`` so the capability check considers the protocol the
|
||||
caller will actually use."""
|
||||
match self:
|
||||
case MaskedAttentionFunction.AUTOMATIC:
|
||||
return automatic_masked_attention()
|
||||
case MaskedAttentionFunction.PYTORCH:
|
||||
return PytorchAttention()
|
||||
case MaskedAttentionFunction.XFORMERS:
|
||||
if memory_efficient_attention is None:
|
||||
raise RuntimeError("MaskedAttentionFunction.XFORMERS selected but `xformers` is not installed.")
|
||||
return XFormersAttention()
|
||||
case MaskedAttentionFunction.SDPA_MATH:
|
||||
return PytorchAttention(priority=[SDPBackend.MATH])
|
||||
case MaskedAttentionFunction.SDPA_CUDNN:
|
||||
return _resolve_sdpa_variant(
|
||||
SDPBackend.CUDNN_ATTENTION, "MaskedAttentionFunction.SDPA_CUDNN", with_mask=True
|
||||
)
|
||||
case MaskedAttentionFunction.SDPA_EFFICIENT:
|
||||
return _resolve_sdpa_variant(
|
||||
SDPBackend.EFFICIENT_ATTENTION, "MaskedAttentionFunction.SDPA_EFFICIENT", with_mask=True
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AttentionOps:
|
||||
"""Pluggable callables consumed by :class:`Attention`."""
|
||||
|
||||
attention_function: AttentionCallable = field(default_factory=lambda: AttentionFunction.AUTOMATIC.to_callable())
|
||||
masked_attention_function: MaskedAttentionCallable = field(
|
||||
default_factory=lambda: MaskedAttentionFunction.AUTOMATIC.to_callable()
|
||||
)
|
||||
preattention_function: PreAttentionCallable = field(default_factory=PytorchPreAttention)
|
||||
gated_attention_function: GatedAttentionCallable = field(default_factory=PytorchGatedAttention)
|
||||
|
||||
|
||||
class Attention(torch.nn.Module):
|
||||
@@ -145,16 +447,17 @@ class Attention(torch.nn.Module):
|
||||
dim_head: int = 64,
|
||||
norm_eps: float = 1e-6,
|
||||
rope_type: LTXRopeType = LTXRopeType.SPLIT,
|
||||
attention_function: AttentionCallable | AttentionFunction = AttentionFunction.DEFAULT,
|
||||
ops: AttentionOps | None = None,
|
||||
apply_gated_attention: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if ops is None:
|
||||
ops = AttentionOps()
|
||||
self.rope_type = rope_type
|
||||
self.attention_function = (
|
||||
attention_function.to_callable()
|
||||
if isinstance(attention_function, AttentionFunction)
|
||||
else attention_function
|
||||
)
|
||||
self.attention_function = ops.attention_function
|
||||
self.masked_attention_function = ops.masked_attention_function
|
||||
self.preattention_function = ops.preattention_function
|
||||
self.gated_attention_function = ops.gated_attention_function
|
||||
|
||||
inner_dim = dim_head * heads
|
||||
context_dim = query_dim if context_dim is None else context_dim
|
||||
@@ -196,7 +499,9 @@ class Attention(torch.nn.Module):
|
||||
context: Key/value context tensor of shape ``(B, S, context_dim)``.
|
||||
Falls back to ``x`` (self-attention) when *None*.
|
||||
mask: Optional attention mask. Interpretation depends on the attention
|
||||
backend (additive bias for xformers/PyTorch SDPA).
|
||||
backend (additive bias for xformers/PyTorch SDPA). A non-None
|
||||
``mask`` routes to ``masked_attention_function``; ``None`` keeps
|
||||
the unmasked path.
|
||||
pe: Rotary positional embeddings applied to both ``q`` and ``k``.
|
||||
k_pe: Separate rotary positional embeddings for ``k`` only. When
|
||||
*None*, ``pe`` is reused for keys.
|
||||
@@ -221,29 +526,17 @@ class Attention(torch.nn.Module):
|
||||
else:
|
||||
q = self.to_q(x)
|
||||
k = self.to_k(context)
|
||||
|
||||
q = self.q_norm(q)
|
||||
k = self.k_norm(k)
|
||||
|
||||
if pe is not None:
|
||||
q = apply_rotary_emb(q, pe, self.rope_type)
|
||||
k = apply_rotary_emb(k, pe if k_pe is None else k_pe, self.rope_type)
|
||||
|
||||
out = self.attention_function(q, k, v, self.heads, mask) # (B, T, H*D)
|
||||
q, k = self.preattention_function(q, k, self, mask, pe, k_pe)
|
||||
if mask is None:
|
||||
out = self.attention_function(q, k, v, self.heads) # (B, T, H*D)
|
||||
else:
|
||||
out = self.masked_attention_function(q, k, v, self.heads, mask)
|
||||
|
||||
if perturbation_mask is not None:
|
||||
out = out * perturbation_mask + v * (1 - perturbation_mask)
|
||||
|
||||
# Apply per-head gating if enabled
|
||||
if self.to_gate_logits is not None:
|
||||
gate_logits = self.to_gate_logits(x) # (B, T, H)
|
||||
b, t, _ = out.shape
|
||||
# Reshape to (B, T, H, D) for per-head gating
|
||||
out = out.view(b, t, self.heads, self.dim_head)
|
||||
# Apply gating: 2 * sigmoid(x) so that zero-init gives identity (2 * 0.5 = 1.0)
|
||||
gates = 2.0 * torch.sigmoid(gate_logits) # (B, T, H)
|
||||
out = out * gates.unsqueeze(-1) # (B, T, H, D) * (B, T, H, 1)
|
||||
# Reshape back to (B, T, H*D)
|
||||
out = out.view(b, t, self.heads * self.dim_head)
|
||||
out = self.gated_attention_function(x, out, self)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
@@ -1,19 +1,112 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig, PerturbationType
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.model.transformer.model import LTXModel
|
||||
from ltx_core.model.transformer.transformer_args import BlockPerturbationsProcessor, TransformerArgs
|
||||
|
||||
# Defaults applied inside the patched forward. Overriding via CompilationConfig
|
||||
# replaces these wholesale; it does not merge.
|
||||
_DEFAULT_INDUCTOR_CONFIG: dict[str, Any] = {}
|
||||
_DEFAULT_DYNAMO_CONFIG: dict[str, Any] = {"inline_inbuilt_nn_modules": True, "cache_size_limit": 256}
|
||||
|
||||
|
||||
def compile_transformer(model: LTXModel) -> LTXModel:
|
||||
model.transformer_blocks = torch.nn.ModuleList(torch.compile(m) for m in model.transformer_blocks)
|
||||
@dataclass(frozen=True)
|
||||
class CompilationConfig:
|
||||
"""``torch.compile`` configuration for transformer blocks. ``None`` keeps eager."""
|
||||
|
||||
mode: str | None = None
|
||||
backend: str = "inductor"
|
||||
fullgraph: bool = False
|
||||
dynamic: bool | None = None
|
||||
inductor_config: dict[str, Any] = field(default_factory=lambda: dict(_DEFAULT_INDUCTOR_CONFIG))
|
||||
dynamo_config: dict[str, Any] = field(default_factory=lambda: dict(_DEFAULT_DYNAMO_CONFIG))
|
||||
|
||||
|
||||
class _SeqDynamicMarkingProcessor:
|
||||
"""Marks the per-block seq dim dynamic, then delegates to an inner processor.
|
||||
Installed by ``compile_transformer`` so the per-block compile artifact stays
|
||||
shape-polymorphic. Wraps whatever ``block_input_processor`` was already on
|
||||
the model -- callers that customised the processor keep their customisation;
|
||||
only the seq-dim marking is layered on top. Lives outside the compiled
|
||||
region, so ``mark_dynamic`` runs in eager mode on the tensors that are
|
||||
about to cross into the trace.
|
||||
"""
|
||||
|
||||
def __init__(self, inner: BlockPerturbationsProcessor) -> None:
|
||||
self.inner = inner
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
args: TransformerArgs,
|
||||
perturbations: BatchedPerturbationConfig,
|
||||
block_idx: int,
|
||||
self_attn_type: PerturbationType,
|
||||
cross_attn_type: PerturbationType,
|
||||
) -> TransformerArgs:
|
||||
# Positional embeddings are second-from-last regardless of rope type:
|
||||
# split rope is (B, H, T, D//2) -- dim -2 == 2; interleaved rope is (B, T, D)
|
||||
# -- dim -2 == 1. Both work via the negative index.
|
||||
torch._dynamo.mark_dynamic(args.x, 1)
|
||||
cos, sin = args.positional_embeddings
|
||||
torch._dynamo.mark_dynamic(cos, cos.ndim - 2)
|
||||
torch._dynamo.mark_dynamic(sin, sin.ndim - 2)
|
||||
if args.cross_positional_embeddings is not None:
|
||||
cross_cos, cross_sin = args.cross_positional_embeddings
|
||||
torch._dynamo.mark_dynamic(cross_cos, cross_cos.ndim - 2)
|
||||
torch._dynamo.mark_dynamic(cross_sin, cross_sin.ndim - 2)
|
||||
if args.self_attention_mask is not None:
|
||||
# Dense form is (B, 1, T, T); key-padding form (from the SP wrapper)
|
||||
# is (B, 1, 1, T) -- leave the size-1 query dim static so Dynamo
|
||||
# keeps the broadcast.
|
||||
if args.self_attention_mask.shape[2] > 1:
|
||||
torch._dynamo.mark_dynamic(args.self_attention_mask, 2)
|
||||
torch._dynamo.mark_dynamic(args.self_attention_mask, 3)
|
||||
if args.context_mask is not None:
|
||||
torch._dynamo.mark_dynamic(args.context_mask, 2)
|
||||
# `timesteps` / `embedded_timestep` are per-token when conditioning sets a
|
||||
# per-position denoise mask, in which case their dim 1 equals the seq length
|
||||
# and must vary with it. When they're a single timestep broadcast across the
|
||||
# sequence (dim 1 == 1), leaving them static lets Dynamo keep the size-1
|
||||
# broadcast.
|
||||
if args.timesteps.shape[1] > 1:
|
||||
torch._dynamo.mark_dynamic(args.timesteps, 1)
|
||||
if args.embedded_timestep.shape[1] > 1:
|
||||
torch._dynamo.mark_dynamic(args.embedded_timestep, 1)
|
||||
# `cross_scale_shift_timestep` is the cross-attn AdaLN scale/shift input
|
||||
# derived from the own-modality per-token timesteps (denoise_mask * sigma),
|
||||
# so its dim 1 equals the seq length when conditioning is per-token.
|
||||
# `cross_gate_timestep` is the cross-modality sigma scalar -- dim 1 is 1
|
||||
# and broadcasts, leave it static. Same guard pattern as `timesteps`.
|
||||
if args.cross_scale_shift_timestep is not None and args.cross_scale_shift_timestep.shape[1] > 1:
|
||||
torch._dynamo.mark_dynamic(args.cross_scale_shift_timestep, 1)
|
||||
return self.inner(args, perturbations, block_idx, self_attn_type, cross_attn_type)
|
||||
|
||||
|
||||
def compile_transformer(model: LTXModel, config: CompilationConfig) -> LTXModel:
|
||||
"""Compile each transformer block via ``torch.compile`` with the given settings.
|
||||
The patched forward emits ``torch.compiler.cudagraph_mark_step_begin()`` once
|
||||
per step. Under CUDA-graph-enabling modes (``"reduce-overhead"`` /
|
||||
``"max-autotune"``) this overrides Dynamo's per-invocation auto-mark
|
||||
heuristic, which would otherwise fire once per compiled block call (48 per
|
||||
forward) and treat each block call as a fresh iteration. Under other modes
|
||||
the mark is a no-op (decrements an unread counter).
|
||||
"""
|
||||
model.transformer_blocks = torch.nn.ModuleList(
|
||||
torch.compile(m, mode=config.mode, backend=config.backend, fullgraph=config.fullgraph, dynamic=config.dynamic)
|
||||
for m in model.transformer_blocks
|
||||
)
|
||||
model.block_input_processor = _SeqDynamicMarkingProcessor(inner=model.block_input_processor)
|
||||
|
||||
def patched_dynamo_forward(*args, **kwargs) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
torch.compiler.cudagraph_mark_step_begin()
|
||||
with (
|
||||
torch._inductor.config.patch(unsafe_skip_cache_dynamic_shape_guards=True),
|
||||
torch._dynamo.config.patch( # type: ignore[attr-defined]
|
||||
inline_inbuilt_nn_modules=True, cache_size_limit=256, allow_unspec_int_on_nn_module=True
|
||||
),
|
||||
torch._inductor.config.patch(**config.inductor_config),
|
||||
torch._dynamo.config.patch(**config.dynamo_config), # type: ignore[attr-defined]
|
||||
):
|
||||
return model.forward_without_compilation(*args, **kwargs)
|
||||
|
||||
@@ -22,11 +115,13 @@ def compile_transformer(model: LTXModel) -> LTXModel:
|
||||
return model
|
||||
|
||||
|
||||
COMPILE_TRANSFORMER = ModuleOps(
|
||||
name="compile_transformer",
|
||||
matcher=lambda model: isinstance(model, LTXModel),
|
||||
mutator=lambda model: compile_transformer(model),
|
||||
)
|
||||
def build_compile_transformer_op(config: CompilationConfig) -> ModuleOps:
|
||||
"""Build a ``ModuleOps`` that compiles transformer blocks with the given settings."""
|
||||
return ModuleOps(
|
||||
name="compile_transformer",
|
||||
matcher=lambda model: isinstance(model, LTXModel),
|
||||
mutator=lambda model: compile_transformer(model, config),
|
||||
)
|
||||
|
||||
|
||||
def modify_sd_ops_for_compilation(original_sd_ops: SDOps, number_of_blocks: int = 48) -> SDOps:
|
||||
|
||||
@@ -17,8 +17,19 @@ class Modality:
|
||||
the batch size, *T* is the total number of tokens (noisy +
|
||||
conditioning), and *D* is the input dimension.
|
||||
timesteps: Per-token timestep embeddings, shape ``(B, T)``.
|
||||
positions: Positional coordinates, shape ``(B, 3, T)`` for video
|
||||
(time, height, width) or ``(B, 1, T)`` for audio.
|
||||
positions: Per-token patch coordinates used to build the RoPE
|
||||
frequencies. With the default ``use_middle_indices_grid=True``,
|
||||
shape is ``(B, n_pos_dims, T, 2)`` where ``n_pos_dims=3`` for
|
||||
video (time, height, width) and ``n_pos_dims=1`` for audio
|
||||
(time); the last dim of size 2 holds the ``[start, end)``
|
||||
index bounds of each patch, and RoPE is evaluated at the
|
||||
*middle* of that range -- hence the flag name. Taking the
|
||||
patch midpoint produces a smoother and more accurate
|
||||
positional signal than indexing by the patch's start when
|
||||
patches span more than one spatial / temporal unit.
|
||||
When ``use_middle_indices_grid=False``, the legacy 3-D form
|
||||
``(B, n_pos_dims, T)`` of integer positional indices is
|
||||
accepted instead and used as-is (no midpoint derivation).
|
||||
context: Text conditioning embeddings from the prompt encoder.
|
||||
enabled: Whether this modality is active in the current forward pass.
|
||||
context_mask: Optional mask for the text context tokens.
|
||||
@@ -34,9 +45,10 @@ class Modality:
|
||||
) # Shape: (B, T, D) where B is the batch size, T is the number of tokens, and D is input dimension
|
||||
sigma: torch.Tensor # Shape: (B,). Current sigma value, used for cross-attention timestep calculation.
|
||||
timesteps: torch.Tensor # Shape: (B, T) where T is the number of timesteps
|
||||
positions: (
|
||||
torch.Tensor
|
||||
) # Shape: (B, 3, T) for video, where 3 is the number of dimensions and T is the number of tokens
|
||||
# Shape: (B, n_pos_dims, T, 2) by default (use_middle_indices_grid=True);
|
||||
# n_pos_dims=3 for video, 1 for audio; last dim holds [start, end) patch bounds.
|
||||
# Legacy form (B, n_pos_dims, T) when use_middle_indices_grid=False.
|
||||
positions: torch.Tensor
|
||||
context: torch.Tensor
|
||||
enabled: bool = True
|
||||
context_mask: torch.Tensor | None = None
|
||||
|
||||
@@ -1,20 +1,30 @@
|
||||
import logging
|
||||
from enum import Enum
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig, PerturbationType
|
||||
from ltx_core.model.model_protocol import LTXModelProtocol
|
||||
from ltx_core.model.transformer.adaln import AdaLayerNormSingle, adaln_embedding_coefficient
|
||||
from ltx_core.model.transformer.attention import AttentionCallable, AttentionFunction
|
||||
from ltx_core.model.transformer.attention import attention_label
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.model.transformer.rope import LTXRopeType
|
||||
from ltx_core.model.transformer.transformer import BasicAVTransformerBlock, TransformerConfig
|
||||
from ltx_core.model.transformer.transformer import (
|
||||
DEFAULT_TRANSFORMER_OPS,
|
||||
BasicAVTransformerBlock,
|
||||
TransformerConfig,
|
||||
TransformerOpsConfig,
|
||||
)
|
||||
from ltx_core.model.transformer.transformer_args import (
|
||||
BlockPerturbationsProcessor,
|
||||
MultiModalTransformerArgsPreprocessor,
|
||||
TransformerArgs,
|
||||
TransformerArgsPreprocessor,
|
||||
)
|
||||
from ltx_core.utils import to_denoised
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LTXModelType(Enum):
|
||||
AudioVideo = "ltx av model"
|
||||
@@ -45,7 +55,7 @@ class LTXModel(torch.nn.Module):
|
||||
num_layers: int = 48,
|
||||
cross_attention_dim: int = 4096,
|
||||
norm_eps: float = 1e-06,
|
||||
attention_type: AttentionFunction | AttentionCallable = AttentionFunction.DEFAULT,
|
||||
ops: TransformerOpsConfig = DEFAULT_TRANSFORMER_OPS,
|
||||
positional_embedding_theta: float = 10000.0,
|
||||
positional_embedding_max_pos: list[int] | None = None,
|
||||
timestep_scale_multiplier: int = 1000,
|
||||
@@ -65,6 +75,15 @@ class LTXModel(torch.nn.Module):
|
||||
cross_attention_adaln: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
# Log the attention backends this transformer is built with. Reading the resolved
|
||||
# ``label`` off the ops reports whatever was selected -- AUTOMATIC, an explicit pin
|
||||
# (PYTORCH/XFORMERS/FA3/FA4/SDPA_*), or a directly supplied callable -- so this is the
|
||||
# single source of truth for which kernel a build uses. Fires once per build.
|
||||
logger.info(
|
||||
"Building transformer with attention backends -- self: %s, masked: %s",
|
||||
attention_label(ops.attention_ops.attention_function),
|
||||
attention_label(ops.attention_ops.masked_attention_function),
|
||||
)
|
||||
self._enable_gradient_checkpointing = False
|
||||
self.cross_attention_adaln = cross_attention_adaln
|
||||
self.use_middle_indices_grid = use_middle_indices_grid
|
||||
@@ -115,9 +134,13 @@ class LTXModel(torch.nn.Module):
|
||||
audio_attention_head_dim=audio_attention_head_dim if model_type.is_audio_enabled() else 0,
|
||||
audio_cross_attention_dim=audio_cross_attention_dim,
|
||||
norm_eps=norm_eps,
|
||||
attention_type=attention_type,
|
||||
ops=ops,
|
||||
apply_gated_attention=apply_gated_attention,
|
||||
)
|
||||
# Hook for per-block input prep. Compile transforms in `compiling.py`
|
||||
# wrap (not replace) this with a processor that also marks the seq dim
|
||||
# dynamic, so any caller customisation here is preserved as the inner.
|
||||
self.block_input_processor = BlockPerturbationsProcessor()
|
||||
|
||||
@property
|
||||
def _adaln_embedding_coefficient(self) -> int:
|
||||
@@ -284,7 +307,7 @@ class LTXModel(torch.nn.Module):
|
||||
audio_attention_head_dim: int,
|
||||
audio_cross_attention_dim: int,
|
||||
norm_eps: float,
|
||||
attention_type: AttentionFunction | AttentionCallable,
|
||||
ops: TransformerOpsConfig,
|
||||
apply_gated_attention: bool,
|
||||
) -> None:
|
||||
"""Initialize transformer blocks for LTX."""
|
||||
@@ -315,14 +338,13 @@ class LTXModel(torch.nn.Module):
|
||||
self.transformer_blocks = torch.nn.ModuleList(
|
||||
[
|
||||
BasicAVTransformerBlock(
|
||||
idx=idx,
|
||||
video=video_config,
|
||||
audio=audio_config,
|
||||
rope_type=self.rope_type,
|
||||
norm_eps=norm_eps,
|
||||
attention_function=attention_type,
|
||||
ops=ops,
|
||||
)
|
||||
for idx in range(num_layers)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
@@ -340,29 +362,44 @@ class LTXModel(torch.nn.Module):
|
||||
self,
|
||||
video: TransformerArgs | None,
|
||||
audio: TransformerArgs | None,
|
||||
perturbations: BatchedPerturbationConfig,
|
||||
) -> tuple[TransformerArgs, TransformerArgs]:
|
||||
"""Process transformer blocks for LTXAV."""
|
||||
perturbations: BatchedPerturbationConfig | None,
|
||||
) -> tuple[TransformerArgs | None, TransformerArgs | None]:
|
||||
"""Process transformer blocks for LTXAV.
|
||||
Per-block perturbation masks are precomputed here and attached to each
|
||||
modality's ``TransformerArgs`` so the block forward has no per-block
|
||||
identity to specialise on — all blocks share a single Dynamo cache slot.
|
||||
"""
|
||||
if perturbations is None:
|
||||
batch_size = (video or audio).x.shape[0]
|
||||
perturbations = BatchedPerturbationConfig.empty(batch_size)
|
||||
|
||||
for block_idx, block in enumerate(self.transformer_blocks):
|
||||
if video is not None:
|
||||
video = self.block_input_processor(
|
||||
video,
|
||||
perturbations,
|
||||
block_idx,
|
||||
self_attn_type=PerturbationType.SKIP_VIDEO_SELF_ATTN,
|
||||
cross_attn_type=PerturbationType.SKIP_A2V_CROSS_ATTN,
|
||||
)
|
||||
if audio is not None:
|
||||
audio = self.block_input_processor(
|
||||
audio,
|
||||
perturbations,
|
||||
block_idx,
|
||||
self_attn_type=PerturbationType.SKIP_AUDIO_SELF_ATTN,
|
||||
cross_attn_type=PerturbationType.SKIP_V2A_CROSS_ATTN,
|
||||
)
|
||||
|
||||
# Process transformer blocks
|
||||
for block in self.transformer_blocks:
|
||||
if self._enable_gradient_checkpointing and self.training:
|
||||
# Use gradient checkpointing to save memory during training.
|
||||
# With use_reentrant=False, we can pass dataclasses directly -
|
||||
# PyTorch will track all tensor leaves in the computation graph.
|
||||
video, audio = torch.utils.checkpoint.checkpoint(
|
||||
block,
|
||||
video,
|
||||
audio,
|
||||
perturbations,
|
||||
use_reentrant=False,
|
||||
)
|
||||
else:
|
||||
video, audio = block(
|
||||
video=video,
|
||||
audio=audio,
|
||||
perturbations=perturbations,
|
||||
)
|
||||
video, audio = block(video=video, audio=audio)
|
||||
|
||||
return video, audio
|
||||
|
||||
@@ -388,7 +425,7 @@ class LTXModel(torch.nn.Module):
|
||||
|
||||
def forward(
|
||||
self, video: Modality | None, audio: Modality | None, perturbations: BatchedPerturbationConfig
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
|
||||
"""
|
||||
Forward pass for LTX models.
|
||||
Returns:
|
||||
@@ -436,7 +473,7 @@ class LegacyX0Model(torch.nn.Module):
|
||||
Returns fully denoised output based on the velocities produced by the base model.
|
||||
"""
|
||||
|
||||
def __init__(self, velocity_model: LTXModel):
|
||||
def __init__(self, velocity_model: LTXModelProtocol):
|
||||
super().__init__()
|
||||
self.velocity_model = velocity_model
|
||||
|
||||
@@ -465,7 +502,7 @@ class X0Model(torch.nn.Module):
|
||||
Applies scaled denoising to the video and audio according to the timesteps = sigma * denoising_mask.
|
||||
"""
|
||||
|
||||
def __init__(self, velocity_model: LTXModel):
|
||||
def __init__(self, velocity_model: LTXModelProtocol):
|
||||
super().__init__()
|
||||
self.velocity_model = velocity_model
|
||||
|
||||
|
||||
@@ -2,10 +2,10 @@ import torch
|
||||
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.model.model_protocol import ModelConfigurator
|
||||
from ltx_core.model.transformer.attention import AttentionFunction
|
||||
from ltx_core.model.transformer.model import LTXModel, LTXModelType
|
||||
from ltx_core.model.transformer.rope import LTXRopeType
|
||||
from ltx_core.model.transformer.text_projection import create_caption_projection
|
||||
from ltx_core.model.transformer.transformer import DEFAULT_TRANSFORMER_OPS, TransformerOpsConfig
|
||||
from ltx_core.utils import check_config_value
|
||||
|
||||
|
||||
@@ -16,7 +16,7 @@ class LTXModelConfigurator(ModelConfigurator[LTXModel]):
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def from_config(cls: type[LTXModel], config: dict) -> LTXModel:
|
||||
def from_config(cls, config: dict, ops: TransformerOpsConfig = DEFAULT_TRANSFORMER_OPS) -> LTXModel:
|
||||
# Build caption projections for 19B models (projection handled in transformer).
|
||||
caption_projection, audio_caption_projection = _build_caption_projections(config, is_av=True)
|
||||
|
||||
@@ -40,6 +40,7 @@ class LTXModelConfigurator(ModelConfigurator[LTXModel]):
|
||||
check_config_value(config, "share_ff", False)
|
||||
check_config_value(config, "av_cross_ada_norm", True)
|
||||
check_config_value(config, "use_middle_indices_grid", True)
|
||||
check_config_value(config, "num_attention_heads", config.get("audio_num_attention_heads", float("nan")))
|
||||
|
||||
return LTXModel(
|
||||
model_type=LTXModelType.AudioVideo,
|
||||
@@ -50,7 +51,7 @@ class LTXModelConfigurator(ModelConfigurator[LTXModel]):
|
||||
num_layers=config.get("num_layers", 48),
|
||||
cross_attention_dim=config.get("cross_attention_dim", 4096),
|
||||
norm_eps=config.get("norm_eps", 1e-06),
|
||||
attention_type=AttentionFunction(config.get("attention_type", "default")),
|
||||
ops=ops,
|
||||
positional_embedding_theta=config.get("positional_embedding_theta", 10000.0),
|
||||
positional_embedding_max_pos=config.get("positional_embedding_max_pos", [20, 2048, 2048]),
|
||||
timestep_scale_multiplier=config.get("timestep_scale_multiplier", 1000),
|
||||
@@ -78,7 +79,7 @@ class LTXVideoOnlyModelConfigurator(ModelConfigurator[LTXModel]):
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def from_config(cls: type[LTXModel], config: dict) -> LTXModel:
|
||||
def from_config(cls, config: dict, ops: TransformerOpsConfig = DEFAULT_TRANSFORMER_OPS) -> LTXModel:
|
||||
# Build caption projection for 19B model (projection handled in transformer).
|
||||
caption_projection, _ = _build_caption_projections(config, is_av=False)
|
||||
|
||||
@@ -109,7 +110,7 @@ class LTXVideoOnlyModelConfigurator(ModelConfigurator[LTXModel]):
|
||||
num_layers=config.get("num_layers", 48),
|
||||
cross_attention_dim=config.get("cross_attention_dim", 4096),
|
||||
norm_eps=config.get("norm_eps", 1e-06),
|
||||
attention_type=AttentionFunction(config.get("attention_type", "default")),
|
||||
ops=ops,
|
||||
positional_embedding_theta=config.get("positional_embedding_theta", 10000.0),
|
||||
positional_embedding_max_pos=config.get("positional_embedding_max_pos", [20, 2048, 2048]),
|
||||
timestep_scale_multiplier=config.get("timestep_scale_multiplier", 1000),
|
||||
@@ -122,6 +123,58 @@ class LTXVideoOnlyModelConfigurator(ModelConfigurator[LTXModel]):
|
||||
)
|
||||
|
||||
|
||||
class LTXAudioOnlyModelConfigurator(ModelConfigurator[LTXModel]):
|
||||
"""
|
||||
Configurator for LTX audio only model.
|
||||
Builds an audio-only LTX model (``model_type=AudioOnly``) so the video
|
||||
transformer weights are never instantiated or loaded. Useful for
|
||||
text-to-audio inference where the video branch is unused.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict, ops: TransformerOpsConfig = DEFAULT_TRANSFORMER_OPS) -> LTXModel:
|
||||
# Build audio caption projection for 19B models (projection handled in transformer).
|
||||
_, audio_caption_projection = _build_caption_projections(config, is_av=True)
|
||||
|
||||
config = config.get("transformer", {})
|
||||
|
||||
check_config_value(config, "dropout", 0.0)
|
||||
check_config_value(config, "attention_bias", True)
|
||||
check_config_value(config, "num_vector_embeds", None)
|
||||
check_config_value(config, "activation_fn", "gelu-approximate")
|
||||
check_config_value(config, "num_embeds_ada_norm", 1000)
|
||||
check_config_value(config, "use_linear_projection", False)
|
||||
check_config_value(config, "only_cross_attention", False)
|
||||
check_config_value(config, "cross_attention_norm", True)
|
||||
check_config_value(config, "double_self_attention", False)
|
||||
check_config_value(config, "upcast_attention", False)
|
||||
check_config_value(config, "standardization_norm", "rms_norm")
|
||||
check_config_value(config, "norm_elementwise_affine", False)
|
||||
check_config_value(config, "qk_norm", "rms_norm")
|
||||
check_config_value(config, "positional_embedding_type", "rope")
|
||||
check_config_value(config, "use_middle_indices_grid", True)
|
||||
|
||||
return LTXModel(
|
||||
model_type=LTXModelType.AudioOnly,
|
||||
num_layers=config.get("num_layers", 48),
|
||||
norm_eps=config.get("norm_eps", 1e-06),
|
||||
ops=ops,
|
||||
timestep_scale_multiplier=config.get("timestep_scale_multiplier", 1000),
|
||||
use_middle_indices_grid=config.get("use_middle_indices_grid", True),
|
||||
audio_num_attention_heads=config.get("audio_num_attention_heads", 32),
|
||||
audio_attention_head_dim=config.get("audio_attention_head_dim", 64),
|
||||
audio_in_channels=config.get("audio_in_channels", 128),
|
||||
audio_out_channels=config.get("audio_out_channels", 128),
|
||||
audio_cross_attention_dim=config.get("audio_cross_attention_dim", 2048),
|
||||
audio_positional_embedding_max_pos=config.get("audio_positional_embedding_max_pos", [20]),
|
||||
rope_type=LTXRopeType(config.get("rope_type", "split")),
|
||||
double_precision_rope=config.get("frequencies_precision", False) == "float64",
|
||||
apply_gated_attention=config.get("apply_gated_attention", False),
|
||||
audio_caption_projection=audio_caption_projection,
|
||||
cross_attention_adaln=config.get("cross_attention_adaln", False),
|
||||
)
|
||||
|
||||
|
||||
def _build_caption_projections(
|
||||
config: dict,
|
||||
is_av: bool,
|
||||
@@ -150,3 +203,16 @@ LTXV_MODEL_COMFY_RENAMING_MAP = (
|
||||
.with_matching(prefix="model.diffusion_model.")
|
||||
.with_replacement("model.diffusion_model.", "")
|
||||
)
|
||||
|
||||
LTXV_AUDIO_ONLY_MODEL_COMFY_RENAMING_MAP = (
|
||||
SDOps("LTXV_AUDIO_ONLY_MODEL_COMFY_MAP")
|
||||
.with_matching(prefix="model.diffusion_model.", contains="audio_attn1")
|
||||
.with_matching(prefix="model.diffusion_model.", contains="audio_attn2")
|
||||
.with_matching(prefix="model.diffusion_model.", contains="audio_ff")
|
||||
.with_matching(prefix="model.diffusion_model.", contains="audio_patchify")
|
||||
.with_matching(prefix="model.diffusion_model.", contains="audio_proj_out")
|
||||
.with_matching(prefix="model.diffusion_model.", contains="audio_adaln_single")
|
||||
.with_matching(prefix="model.diffusion_model.", contains="audio_prompt")
|
||||
.with_matching(prefix="model.diffusion_model.", contains="audio_scale_shift_table")
|
||||
.with_replacement("model.diffusion_model.", "")
|
||||
)
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
from typing import List, Protocol
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ltx_core.model.transformer.rope import apply_rotary_emb
|
||||
from ltx_core.utils import rms_norm
|
||||
|
||||
|
||||
class PreAttentionCallable(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
attn_module: nn.Module,
|
||||
mask: torch.Tensor | None,
|
||||
pe: torch.Tensor | None,
|
||||
k_pe: torch.Tensor | None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]: ...
|
||||
|
||||
|
||||
class PytorchPreAttention(PreAttentionCallable):
|
||||
def __call__(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
attn_module: nn.Module,
|
||||
mask: torch.Tensor | None, # noqa: ARG002
|
||||
pe: torch.Tensor | None,
|
||||
k_pe: torch.Tensor | None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
q = attn_module.q_norm(q)
|
||||
k = attn_module.k_norm(k)
|
||||
if pe is not None:
|
||||
q = apply_rotary_emb(q, pe, attn_module.rope_type)
|
||||
k = apply_rotary_emb(k, pe if k_pe is None else k_pe, attn_module.rope_type)
|
||||
return q, k
|
||||
|
||||
|
||||
class AdaZeroCallable(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
eps: float,
|
||||
scale: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
) -> torch.Tensor: ...
|
||||
|
||||
|
||||
class PytorchAdaZeroFunction(AdaZeroCallable):
|
||||
def __call__(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
eps: float,
|
||||
scale: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return rms_norm(x, eps=eps) * (1 + scale) + shift
|
||||
|
||||
|
||||
class PostSACallable(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
y: torch.Tensor,
|
||||
norm_weights: torch.Tensor | None,
|
||||
eps: float,
|
||||
gate: torch.Tensor,
|
||||
) -> List[torch.Tensor]: ...
|
||||
|
||||
|
||||
class PytorchPostSAFunction(PostSACallable):
|
||||
def __call__(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
y: torch.Tensor,
|
||||
norm_weights: torch.Tensor | None,
|
||||
eps: float,
|
||||
gate: torch.Tensor,
|
||||
) -> List[torch.Tensor]:
|
||||
x_fma = x + y * gate
|
||||
return x_fma, rms_norm(x_fma, norm_weights, eps=eps)
|
||||
|
||||
|
||||
class GatedAttentionCallable(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
attn_out: torch.Tensor,
|
||||
attn_module: nn.Module,
|
||||
) -> torch.Tensor: ...
|
||||
|
||||
|
||||
class PytorchGatedAttention(GatedAttentionCallable):
|
||||
def __call__(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
attn_out: torch.Tensor,
|
||||
attn_module: nn.Module,
|
||||
) -> torch.Tensor:
|
||||
gate_logits = attn_module.to_gate_logits(x) # (B, T, H)
|
||||
b, t, _ = attn_out.shape
|
||||
out = attn_out.view(b, t, attn_module.heads, attn_module.dim_head)
|
||||
gates = 2.0 * torch.sigmoid(gate_logits) # (B, T, H)
|
||||
out = out * gates.unsqueeze(-1) # (B, T, H, D) * (B, T, H, 1)
|
||||
return out.view(b, t, attn_module.heads * attn_module.dim_head)
|
||||
@@ -43,16 +43,26 @@ def apply_interleaved_rotary_emb(
|
||||
def apply_split_rotary_emb(
|
||||
input_tensor: torch.Tensor, cos_freqs: torch.Tensor, sin_freqs: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
needs_reshape = False
|
||||
if input_tensor.ndim != 4 and cos_freqs.ndim == 4:
|
||||
b, h, t, _ = cos_freqs.shape
|
||||
if input_tensor.shape[0] != b:
|
||||
if sin_freqs.shape != cos_freqs.shape:
|
||||
raise ValueError(
|
||||
f"apply_split_rotary_emb: sin_freqs.shape {tuple(sin_freqs.shape)} must equal "
|
||||
f"cos_freqs.shape {tuple(cos_freqs.shape)}."
|
||||
)
|
||||
needs_reshape = input_tensor.ndim != 4 and cos_freqs.ndim == 4
|
||||
if needs_reshape:
|
||||
b_freq = cos_freqs.shape[0]
|
||||
h = cos_freqs.shape[1]
|
||||
b_in = input_tensor.shape[0]
|
||||
if b_freq not in (1, b_in):
|
||||
raise ValueError(
|
||||
f"apply_split_rotary_emb: input_tensor batch ({input_tensor.shape[0]}) "
|
||||
f"must equal cos_freqs batch ({b})."
|
||||
f"apply_split_rotary_emb: cos_freqs batch ({b_freq}) must be 1 "
|
||||
f"(broadcast) or equal input_tensor batch ({b_in})."
|
||||
)
|
||||
input_tensor = input_tensor.reshape(b, t, h, -1).swapaxes(1, 2)
|
||||
needs_reshape = True
|
||||
# `unflatten` only touches the last dim, keeping the batch and seq dims as
|
||||
# the input tensor's own symbolic ints under torch.compile. `reshape(b_in,
|
||||
# t, h, -1)` would have forced Dynamo to specialise those dims because it
|
||||
# cannot prove `b_in == cos_freqs.shape[0]` and `seq == t` across tensors.
|
||||
input_tensor = input_tensor.unflatten(-1, (h, -1)).transpose(1, 2)
|
||||
|
||||
split_input = rearrange(input_tensor, "... (d r) -> ... d r", d=2)
|
||||
first_half_input = split_input[..., :1, :]
|
||||
@@ -67,7 +77,9 @@ def apply_split_rotary_emb(
|
||||
|
||||
output = rearrange(output, "... d r -> ... (d r)")
|
||||
if needs_reshape:
|
||||
output = output.swapaxes(1, 2).reshape(b, t, -1)
|
||||
# `transpose(1, 2).flatten(-2)` keeps the batch and seq dims symbolic; using
|
||||
# `reshape(b_in, t, -1)` would force Dynamo to specialise both axes.
|
||||
output = output.transpose(1, 2).flatten(-2)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
@@ -1,14 +1,29 @@
|
||||
from dataclasses import dataclass, replace
|
||||
from dataclasses import dataclass, field, replace
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig, PerturbationType
|
||||
from ltx_core.model.transformer.adaln import adaln_embedding_coefficient
|
||||
from ltx_core.model.transformer.attention import Attention, AttentionCallable, AttentionFunction
|
||||
from ltx_core.model.transformer.attention import (
|
||||
Attention,
|
||||
AttentionCallable,
|
||||
AttentionFunction,
|
||||
AttentionOps,
|
||||
MaskedAttentionCallable,
|
||||
MaskedAttentionFunction,
|
||||
)
|
||||
from ltx_core.model.transformer.feed_forward import FeedForward
|
||||
from ltx_core.model.transformer.ops import (
|
||||
AdaZeroCallable,
|
||||
GatedAttentionCallable,
|
||||
PostSACallable,
|
||||
PreAttentionCallable,
|
||||
PytorchAdaZeroFunction,
|
||||
PytorchGatedAttention,
|
||||
PytorchPostSAFunction,
|
||||
PytorchPreAttention,
|
||||
)
|
||||
from ltx_core.model.transformer.rope import LTXRopeType
|
||||
from ltx_core.model.transformer.transformer_args import TransformerArgs
|
||||
from ltx_core.utils import rms_norm
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -21,19 +36,68 @@ class TransformerConfig:
|
||||
cross_attention_adaln: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TransformerOpsConfig:
|
||||
"""Pluggable ops for :class:`BasicAVTransformerBlock`.
|
||||
Use :meth:`from_functions` to construct from enum values or partial overrides
|
||||
without spelling out a full :class:`AttentionOps`.
|
||||
"""
|
||||
|
||||
attention_ops: AttentionOps = field(default_factory=AttentionOps)
|
||||
ada_zero_function: AdaZeroCallable = field(default_factory=PytorchAdaZeroFunction)
|
||||
post_sa_function: PostSACallable = field(default_factory=PytorchPostSAFunction)
|
||||
|
||||
@classmethod
|
||||
def from_functions(
|
||||
cls,
|
||||
attention: AttentionFunction | AttentionCallable = AttentionFunction.AUTOMATIC,
|
||||
masked_attention: MaskedAttentionFunction | MaskedAttentionCallable = MaskedAttentionFunction.AUTOMATIC,
|
||||
preattention: PreAttentionCallable | None = None,
|
||||
gated_attention: GatedAttentionCallable | None = None,
|
||||
ada_zero: AdaZeroCallable | None = None,
|
||||
post_sa: PostSACallable | None = None,
|
||||
) -> "TransformerOpsConfig":
|
||||
"""Build a config from individual functions or enums. Each *None* slot
|
||||
falls back to the standard PyTorch implementation."""
|
||||
attention_callable = attention.to_callable() if isinstance(attention, AttentionFunction) else attention
|
||||
masked_callable = (
|
||||
masked_attention.to_callable()
|
||||
if isinstance(masked_attention, MaskedAttentionFunction)
|
||||
else masked_attention
|
||||
)
|
||||
attention_ops = AttentionOps(
|
||||
attention_function=attention_callable,
|
||||
masked_attention_function=masked_callable,
|
||||
preattention_function=preattention if preattention is not None else PytorchPreAttention(),
|
||||
gated_attention_function=(gated_attention if gated_attention is not None else PytorchGatedAttention()),
|
||||
)
|
||||
return cls(
|
||||
attention_ops=attention_ops,
|
||||
ada_zero_function=ada_zero if ada_zero is not None else PytorchAdaZeroFunction(),
|
||||
post_sa_function=post_sa if post_sa is not None else PytorchPostSAFunction(),
|
||||
)
|
||||
|
||||
|
||||
# Frozen, so safe to share as a default argument across callers that want the
|
||||
# stock PyTorch ops without explicit construction.
|
||||
DEFAULT_TRANSFORMER_OPS = TransformerOpsConfig()
|
||||
|
||||
|
||||
class BasicAVTransformerBlock(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
idx: int,
|
||||
video: TransformerConfig | None = None,
|
||||
audio: TransformerConfig | None = None,
|
||||
rope_type: LTXRopeType = LTXRopeType.SPLIT,
|
||||
norm_eps: float = 1e-6,
|
||||
attention_function: AttentionFunction | AttentionCallable = AttentionFunction.DEFAULT,
|
||||
ops: TransformerOpsConfig | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.idx = idx
|
||||
if ops is None:
|
||||
ops = TransformerOpsConfig()
|
||||
self.ada_zero_function = ops.ada_zero_function
|
||||
self.post_sa_function = ops.post_sa_function
|
||||
if video is not None:
|
||||
self.attn1 = Attention(
|
||||
query_dim=video.dim,
|
||||
@@ -42,7 +106,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
context_dim=None,
|
||||
rope_type=rope_type,
|
||||
norm_eps=norm_eps,
|
||||
attention_function=attention_function,
|
||||
ops=ops.attention_ops,
|
||||
apply_gated_attention=video.apply_gated_attention,
|
||||
)
|
||||
self.attn2 = Attention(
|
||||
@@ -52,7 +116,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
dim_head=video.d_head,
|
||||
rope_type=rope_type,
|
||||
norm_eps=norm_eps,
|
||||
attention_function=attention_function,
|
||||
ops=ops.attention_ops,
|
||||
apply_gated_attention=video.apply_gated_attention,
|
||||
)
|
||||
self.ff = FeedForward(video.dim, dim_out=video.dim)
|
||||
@@ -67,7 +131,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
context_dim=None,
|
||||
rope_type=rope_type,
|
||||
norm_eps=norm_eps,
|
||||
attention_function=attention_function,
|
||||
ops=ops.attention_ops,
|
||||
apply_gated_attention=audio.apply_gated_attention,
|
||||
)
|
||||
self.audio_attn2 = Attention(
|
||||
@@ -77,7 +141,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
dim_head=audio.d_head,
|
||||
rope_type=rope_type,
|
||||
norm_eps=norm_eps,
|
||||
attention_function=attention_function,
|
||||
ops=ops.attention_ops,
|
||||
apply_gated_attention=audio.apply_gated_attention,
|
||||
)
|
||||
self.audio_ff = FeedForward(audio.dim, dim_out=audio.dim)
|
||||
@@ -93,7 +157,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
dim_head=audio.d_head,
|
||||
rope_type=rope_type,
|
||||
norm_eps=norm_eps,
|
||||
attention_function=attention_function,
|
||||
ops=ops.attention_ops,
|
||||
apply_gated_attention=video.apply_gated_attention,
|
||||
)
|
||||
|
||||
@@ -105,7 +169,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
dim_head=audio.d_head,
|
||||
rope_type=rope_type,
|
||||
norm_eps=norm_eps,
|
||||
attention_function=attention_function,
|
||||
ops=ops.attention_ops,
|
||||
apply_gated_attention=audio.apply_gated_attention,
|
||||
)
|
||||
|
||||
@@ -157,7 +221,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
|
||||
def _apply_text_cross_attention(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
x_normed: torch.Tensor,
|
||||
context: torch.Tensor,
|
||||
attn: AttentionCallable,
|
||||
scale_shift_table: torch.Tensor,
|
||||
@@ -167,11 +231,14 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
context_mask: torch.Tensor | None,
|
||||
cross_attention_adaln: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Apply text cross-attention, with optional AdaLN modulation."""
|
||||
"""Apply text cross-attention, with optional AdaLN modulation.
|
||||
``x_normed`` is the RMS-normalized self-attention output produced by
|
||||
``post_sa_function`` -- this method does not normalize again.
|
||||
"""
|
||||
if cross_attention_adaln:
|
||||
shift_q, scale_q, gate = self.get_ada_values(scale_shift_table, x.shape[0], timestep, slice(6, 9))
|
||||
shift_q, scale_q, gate = self.get_ada_values(scale_shift_table, x_normed.shape[0], timestep, slice(6, 9))
|
||||
return apply_cross_attention_adaln(
|
||||
x,
|
||||
x_normed,
|
||||
context,
|
||||
attn,
|
||||
shift_q,
|
||||
@@ -180,24 +247,17 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
prompt_scale_shift_table,
|
||||
prompt_timestep,
|
||||
context_mask,
|
||||
self.norm_eps,
|
||||
)
|
||||
return attn(rms_norm(x, eps=self.norm_eps), context=context, mask=context_mask)
|
||||
return attn(x_normed, context=context, mask=context_mask)
|
||||
|
||||
def forward( # noqa: PLR0915
|
||||
self,
|
||||
video: TransformerArgs | None,
|
||||
audio: TransformerArgs | None,
|
||||
perturbations: BatchedPerturbationConfig | None = None,
|
||||
) -> tuple[TransformerArgs | None, TransformerArgs | None]:
|
||||
if video is None and audio is None:
|
||||
raise ValueError("At least one of video or audio must be provided")
|
||||
|
||||
batch_size = (video or audio).x.shape[0]
|
||||
|
||||
if perturbations is None:
|
||||
perturbations = BatchedPerturbationConfig.empty(batch_size)
|
||||
|
||||
vx = video.x if video is not None else None
|
||||
ax = audio.x if audio is not None else None
|
||||
|
||||
@@ -211,30 +271,20 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
vshift_msa, vscale_msa, vgate_msa = self.get_ada_values(
|
||||
self.scale_shift_table, vx.shape[0], video.timesteps, slice(0, 3)
|
||||
)
|
||||
norm_vx = rms_norm(vx, eps=self.norm_eps) * (1 + vscale_msa) + vshift_msa
|
||||
norm_vx = self.ada_zero_function(vx, self.norm_eps, vscale_msa, vshift_msa)
|
||||
del vshift_msa, vscale_msa
|
||||
|
||||
all_perturbed = perturbations.all_in_batch(PerturbationType.SKIP_VIDEO_SELF_ATTN, self.idx)
|
||||
none_perturbed = not perturbations.any_in_batch(PerturbationType.SKIP_VIDEO_SELF_ATTN, self.idx)
|
||||
v_mask = (
|
||||
perturbations.mask_like(PerturbationType.SKIP_VIDEO_SELF_ATTN, self.idx, vx)
|
||||
if not all_perturbed and not none_perturbed
|
||||
else None
|
||||
vx_msa_out = self.attn1(
|
||||
norm_vx,
|
||||
pe=video.positional_embeddings,
|
||||
mask=video.self_attention_mask,
|
||||
perturbation_mask=video.self_attn_perturbation_mask,
|
||||
all_perturbed=video.self_attn_all_perturbed,
|
||||
)
|
||||
vx = (
|
||||
vx
|
||||
+ self.attn1(
|
||||
norm_vx,
|
||||
pe=video.positional_embeddings,
|
||||
mask=video.self_attention_mask,
|
||||
perturbation_mask=v_mask,
|
||||
all_perturbed=all_perturbed,
|
||||
)
|
||||
* vgate_msa
|
||||
)
|
||||
del vgate_msa, norm_vx, v_mask
|
||||
vx, vx_normed = self.post_sa_function(vx, vx_msa_out, None, self.norm_eps, vgate_msa)
|
||||
del vgate_msa, norm_vx, vx_msa_out
|
||||
vx = vx + self._apply_text_cross_attention(
|
||||
vx,
|
||||
vx_normed,
|
||||
video.context,
|
||||
self.attn2,
|
||||
self.scale_shift_table,
|
||||
@@ -244,35 +294,26 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
video.context_mask,
|
||||
cross_attention_adaln=self.cross_attention_adaln,
|
||||
)
|
||||
del vx_normed
|
||||
|
||||
if run_ax:
|
||||
ashift_msa, ascale_msa, agate_msa = self.get_ada_values(
|
||||
self.audio_scale_shift_table, ax.shape[0], audio.timesteps, slice(0, 3)
|
||||
)
|
||||
|
||||
norm_ax = rms_norm(ax, eps=self.norm_eps) * (1 + ascale_msa) + ashift_msa
|
||||
norm_ax = self.ada_zero_function(ax, self.norm_eps, ascale_msa, ashift_msa)
|
||||
del ashift_msa, ascale_msa
|
||||
all_perturbed = perturbations.all_in_batch(PerturbationType.SKIP_AUDIO_SELF_ATTN, self.idx)
|
||||
none_perturbed = not perturbations.any_in_batch(PerturbationType.SKIP_AUDIO_SELF_ATTN, self.idx)
|
||||
a_mask = (
|
||||
perturbations.mask_like(PerturbationType.SKIP_AUDIO_SELF_ATTN, self.idx, ax)
|
||||
if not all_perturbed and not none_perturbed
|
||||
else None
|
||||
ax_msa_out = self.audio_attn1(
|
||||
norm_ax,
|
||||
pe=audio.positional_embeddings,
|
||||
mask=audio.self_attention_mask,
|
||||
perturbation_mask=audio.self_attn_perturbation_mask,
|
||||
all_perturbed=audio.self_attn_all_perturbed,
|
||||
)
|
||||
ax = (
|
||||
ax
|
||||
+ self.audio_attn1(
|
||||
norm_ax,
|
||||
pe=audio.positional_embeddings,
|
||||
mask=audio.self_attention_mask,
|
||||
perturbation_mask=a_mask,
|
||||
all_perturbed=all_perturbed,
|
||||
)
|
||||
* agate_msa
|
||||
)
|
||||
del agate_msa, norm_ax, a_mask
|
||||
ax, ax_normed = self.post_sa_function(ax, ax_msa_out, None, self.norm_eps, agate_msa)
|
||||
del agate_msa, norm_ax, ax_msa_out
|
||||
ax = ax + self._apply_text_cross_attention(
|
||||
ax,
|
||||
ax_normed,
|
||||
audio.context,
|
||||
self.audio_attn2,
|
||||
self.audio_scale_shift_table,
|
||||
@@ -282,13 +323,15 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
audio.context_mask,
|
||||
cross_attention_adaln=self.cross_attention_adaln,
|
||||
)
|
||||
del ax_normed
|
||||
|
||||
# Audio - Video cross attention.
|
||||
if run_a2v or run_v2a:
|
||||
vx_norm3 = rms_norm(vx, eps=self.norm_eps)
|
||||
ax_norm3 = rms_norm(ax, eps=self.norm_eps)
|
||||
|
||||
if run_a2v and not perturbations.all_in_batch(PerturbationType.SKIP_A2V_CROSS_ATTN, self.idx):
|
||||
# Snapshot vx/ax before A2V mutates vx; V2A's video keys/values must
|
||||
# use the pre-A2V state so direction order doesn't bias the result.
|
||||
vx_pre_av = vx
|
||||
ax_pre_av = ax
|
||||
if run_a2v and not video.cross_attn_skip_all:
|
||||
scale_ca_video_a2v, shift_ca_video_a2v, gate_out_a2v = self.get_av_ca_ada_values(
|
||||
self.scale_shift_table_a2v_ca_video,
|
||||
vx.shape[0],
|
||||
@@ -296,7 +339,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
video.cross_gate_timestep,
|
||||
slice(0, 2),
|
||||
)
|
||||
vx_scaled = vx_norm3 * (1 + scale_ca_video_a2v) + shift_ca_video_a2v
|
||||
a2v_vx_scaled = self.ada_zero_function(vx_pre_av, self.norm_eps, scale_ca_video_a2v, shift_ca_video_a2v)
|
||||
del scale_ca_video_a2v, shift_ca_video_a2v
|
||||
|
||||
scale_ca_audio_a2v, shift_ca_audio_a2v, _ = self.get_av_ca_ada_values(
|
||||
@@ -306,22 +349,21 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
audio.cross_gate_timestep,
|
||||
slice(0, 2),
|
||||
)
|
||||
ax_scaled = ax_norm3 * (1 + scale_ca_audio_a2v) + shift_ca_audio_a2v
|
||||
a2v_ax_scaled = self.ada_zero_function(ax_pre_av, self.norm_eps, scale_ca_audio_a2v, shift_ca_audio_a2v)
|
||||
del scale_ca_audio_a2v, shift_ca_audio_a2v
|
||||
a2v_mask = perturbations.mask_like(PerturbationType.SKIP_A2V_CROSS_ATTN, self.idx, vx)
|
||||
vx = vx + (
|
||||
self.audio_to_video_attn(
|
||||
vx_scaled,
|
||||
context=ax_scaled,
|
||||
a2v_vx_scaled,
|
||||
context=a2v_ax_scaled,
|
||||
pe=video.cross_positional_embeddings,
|
||||
k_pe=audio.cross_positional_embeddings,
|
||||
)
|
||||
* gate_out_a2v
|
||||
* a2v_mask
|
||||
* video.cross_attn_perturbation_mask
|
||||
)
|
||||
del gate_out_a2v, a2v_mask, vx_scaled, ax_scaled
|
||||
del gate_out_a2v, a2v_vx_scaled, a2v_ax_scaled
|
||||
|
||||
if run_v2a and not perturbations.all_in_batch(PerturbationType.SKIP_V2A_CROSS_ATTN, self.idx):
|
||||
if run_v2a and not audio.cross_attn_skip_all:
|
||||
scale_ca_audio_v2a, shift_ca_audio_v2a, gate_out_v2a = self.get_av_ca_ada_values(
|
||||
self.scale_shift_table_a2v_ca_audio,
|
||||
ax.shape[0],
|
||||
@@ -329,7 +371,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
audio.cross_gate_timestep,
|
||||
slice(2, 4),
|
||||
)
|
||||
ax_scaled = ax_norm3 * (1 + scale_ca_audio_v2a) + shift_ca_audio_v2a
|
||||
v2a_ax_scaled = self.ada_zero_function(ax_pre_av, self.norm_eps, scale_ca_audio_v2a, shift_ca_audio_v2a)
|
||||
del scale_ca_audio_v2a, shift_ca_audio_v2a
|
||||
scale_ca_video_v2a, shift_ca_video_v2a, _ = self.get_av_ca_ada_values(
|
||||
self.scale_shift_table_a2v_ca_video,
|
||||
@@ -338,28 +380,26 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
video.cross_gate_timestep,
|
||||
slice(2, 4),
|
||||
)
|
||||
vx_scaled = vx_norm3 * (1 + scale_ca_video_v2a) + shift_ca_video_v2a
|
||||
v2a_vx_scaled = self.ada_zero_function(vx_pre_av, self.norm_eps, scale_ca_video_v2a, shift_ca_video_v2a)
|
||||
del scale_ca_video_v2a, shift_ca_video_v2a
|
||||
v2a_mask = perturbations.mask_like(PerturbationType.SKIP_V2A_CROSS_ATTN, self.idx, ax)
|
||||
ax = ax + (
|
||||
self.video_to_audio_attn(
|
||||
ax_scaled,
|
||||
context=vx_scaled,
|
||||
v2a_ax_scaled,
|
||||
context=v2a_vx_scaled,
|
||||
pe=audio.cross_positional_embeddings,
|
||||
k_pe=video.cross_positional_embeddings,
|
||||
)
|
||||
* gate_out_v2a
|
||||
* v2a_mask
|
||||
* audio.cross_attn_perturbation_mask
|
||||
)
|
||||
del gate_out_v2a, v2a_mask, ax_scaled, vx_scaled
|
||||
|
||||
del vx_norm3, ax_norm3
|
||||
del gate_out_v2a, v2a_vx_scaled, v2a_ax_scaled
|
||||
del vx_pre_av, ax_pre_av
|
||||
|
||||
if run_vx:
|
||||
vshift_mlp, vscale_mlp, vgate_mlp = self.get_ada_values(
|
||||
self.scale_shift_table, vx.shape[0], video.timesteps, slice(3, 6)
|
||||
)
|
||||
vx_scaled = rms_norm(vx, eps=self.norm_eps) * (1 + vscale_mlp) + vshift_mlp
|
||||
vx_scaled = self.ada_zero_function(vx, self.norm_eps, vscale_mlp, vshift_mlp)
|
||||
vx = vx + self.ff(vx_scaled) * vgate_mlp
|
||||
|
||||
del vshift_mlp, vscale_mlp, vgate_mlp, vx_scaled
|
||||
@@ -368,7 +408,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
ashift_mlp, ascale_mlp, agate_mlp = self.get_ada_values(
|
||||
self.audio_scale_shift_table, ax.shape[0], audio.timesteps, slice(3, 6)
|
||||
)
|
||||
ax_scaled = rms_norm(ax, eps=self.norm_eps) * (1 + ascale_mlp) + ashift_mlp
|
||||
ax_scaled = self.ada_zero_function(ax, self.norm_eps, ascale_mlp, ashift_mlp)
|
||||
ax = ax + self.audio_ff(ax_scaled) * agate_mlp
|
||||
|
||||
del ashift_mlp, ascale_mlp, agate_mlp, ax_scaled
|
||||
@@ -377,7 +417,7 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
|
||||
|
||||
def apply_cross_attention_adaln(
|
||||
x: torch.Tensor,
|
||||
x_normed: torch.Tensor,
|
||||
context: torch.Tensor,
|
||||
attn: AttentionCallable,
|
||||
q_shift: torch.Tensor,
|
||||
@@ -386,13 +426,17 @@ def apply_cross_attention_adaln(
|
||||
prompt_scale_shift_table: torch.Tensor,
|
||||
prompt_timestep: torch.Tensor,
|
||||
context_mask: torch.Tensor | None = None,
|
||||
norm_eps: float = 1e-6,
|
||||
) -> torch.Tensor:
|
||||
batch_size = x.shape[0]
|
||||
"""Apply query/key AdaLN modulation then cross-attention.
|
||||
``x_normed`` is already RMS-normalized by ``post_sa_function``; this only
|
||||
applies the affine (scale/shift) modulation, so the normalization is not
|
||||
repeated here.
|
||||
"""
|
||||
batch_size = x_normed.shape[0]
|
||||
shift_kv, scale_kv = (
|
||||
prompt_scale_shift_table[None, None].to(device=x.device, dtype=x.dtype)
|
||||
prompt_scale_shift_table[None, None].to(device=x_normed.device, dtype=x_normed.dtype)
|
||||
+ prompt_timestep.reshape(batch_size, prompt_timestep.shape[1], 2, -1)
|
||||
).unbind(dim=2)
|
||||
attn_input = rms_norm(x, eps=norm_eps) * (1 + q_scale) + q_shift
|
||||
attn_input = x_normed * (1 + q_scale) + q_shift
|
||||
encoder_hidden_states = context * (1 + scale_kv) + shift_kv
|
||||
return attn(attn_input, context=encoder_hidden_states, mask=context_mask) * q_gate
|
||||
|
||||
@@ -2,6 +2,7 @@ from dataclasses import dataclass, replace
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.guidance.perturbations import BatchedPerturbationConfig, PerturbationType
|
||||
from ltx_core.model.transformer.adaln import AdaLayerNormSingle
|
||||
from ltx_core.model.transformer.modality import Modality
|
||||
from ltx_core.model.transformer.rope import (
|
||||
@@ -19,8 +20,8 @@ class TransformerArgs:
|
||||
context_mask: torch.Tensor
|
||||
timesteps: torch.Tensor
|
||||
embedded_timestep: torch.Tensor
|
||||
positional_embeddings: torch.Tensor
|
||||
cross_positional_embeddings: torch.Tensor | None
|
||||
positional_embeddings: tuple[torch.Tensor, torch.Tensor]
|
||||
cross_positional_embeddings: tuple[torch.Tensor, torch.Tensor] | None
|
||||
cross_scale_shift_timestep: torch.Tensor | None
|
||||
cross_gate_timestep: torch.Tensor | None
|
||||
enabled: bool
|
||||
@@ -28,6 +29,59 @@ class TransformerArgs:
|
||||
self_attention_mask: torch.Tensor | None = (
|
||||
None # Additive log-space self-attention bias (B, 1, T, T), None = full attention
|
||||
)
|
||||
# Per-block perturbation state, precomputed by `LTXModel._process_transformer_blocks`
|
||||
# so the block forward needs no per-block identity. The bool shortcuts
|
||||
# (`*_all_perturbed`, `cross_attn_skip_all`) are Python bools that Dynamo specialises
|
||||
# on — fine because they're stable across denoising steps for a fixed perturbation
|
||||
# config.
|
||||
self_attn_perturbation_mask: torch.Tensor | None = None
|
||||
self_attn_all_perturbed: bool = False
|
||||
cross_attn_perturbation_mask: torch.Tensor | None = None
|
||||
cross_attn_skip_all: bool = False
|
||||
|
||||
|
||||
class BlockPerturbationsProcessor:
|
||||
"""Per-block preparation of ``TransformerArgs``.
|
||||
The base implementation returns a copy of ``args`` with this block's
|
||||
precomputed perturbation flags and masks attached. Subclasses can layer in
|
||||
operations that must run on each block's inputs but stay outside the
|
||||
compile boundary -- e.g. ``torch._dynamo.mark_dynamic`` for
|
||||
shape-polymorphic block compilation (see ``compiling.py``). Swapping the
|
||||
processor on an ``LTXModel`` instance is how compile transforms opt in to
|
||||
such behaviour without baking it into the model's forward.
|
||||
``self_attn_perturbation_mask`` is None when all or none of the batch is
|
||||
perturbed (the attention call can take the shortcut path). ``cross_attn_*``
|
||||
is None when every sample skips the cross-attention entirely.
|
||||
"""
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
args: "TransformerArgs",
|
||||
perturbations: BatchedPerturbationConfig,
|
||||
block_idx: int,
|
||||
self_attn_type: PerturbationType,
|
||||
cross_attn_type: PerturbationType,
|
||||
) -> "TransformerArgs":
|
||||
device, dtype = args.x.device, args.x.dtype
|
||||
|
||||
all_self = perturbations.all_in_batch(self_attn_type, block_idx)
|
||||
any_self = perturbations.any_in_batch(self_attn_type, block_idx)
|
||||
self_mask: torch.Tensor | None = None
|
||||
if any_self and not all_self:
|
||||
self_mask = perturbations.mask(self_attn_type, block_idx, device, dtype).view(-1, 1, 1)
|
||||
|
||||
all_cross = perturbations.all_in_batch(cross_attn_type, block_idx)
|
||||
cross_mask: torch.Tensor | None = None
|
||||
if not all_cross:
|
||||
cross_mask = perturbations.mask(cross_attn_type, block_idx, device, dtype).view(-1, 1, 1)
|
||||
|
||||
return replace(
|
||||
args,
|
||||
self_attn_perturbation_mask=self_mask,
|
||||
self_attn_all_perturbed=all_self,
|
||||
cross_attn_perturbation_mask=cross_mask,
|
||||
cross_attn_skip_all=all_cross,
|
||||
)
|
||||
|
||||
|
||||
class TransformerArgsPreprocessor:
|
||||
@@ -98,9 +152,12 @@ class TransformerArgsPreprocessor:
|
||||
self, attention_mask: torch.Tensor | None, x_dtype: torch.dtype
|
||||
) -> torch.Tensor | None:
|
||||
"""Prepare self-attention mask by converting [0,1] values to additive log-space bias.
|
||||
Input shape: (B, T, T) with values in [0, 1].
|
||||
Output shape: (B, 1, T, T) with 0.0 for full attention and a large negative value
|
||||
for masked positions.
|
||||
Input shape: 3D ``(B, T_q, T_k)`` with values in [0, 1]. The dense form
|
||||
is ``(B, T, T)``; broadcastable forms like ``(1, 1, T)`` (key-only
|
||||
padding) or ``(B, 1, T)`` are also valid and yield a correspondingly
|
||||
broadcastable output.
|
||||
Output shape: ``(B, 1, T_q, T_k)`` (heads dim inserted) with 0.0 for
|
||||
full attention and a large negative value for masked positions.
|
||||
Positions with attention_mask <= 0 are fully masked (mapped to the dtype's minimum
|
||||
representable value). Strictly positive entries are converted via log-space for
|
||||
smooth attenuation, with small values clamped for numerical stability.
|
||||
@@ -120,7 +177,7 @@ class TransformerArgsPreprocessor:
|
||||
if positive.any():
|
||||
bias[positive] = torch.log(attention_mask[positive].clamp(min=eps)).to(x_dtype)
|
||||
|
||||
return bias.unsqueeze(1) # (B, 1, T, T) for head broadcast
|
||||
return bias.unsqueeze(1) # (B, 1, T_q, T_k) for head broadcast
|
||||
|
||||
def _prepare_positional_embeddings(
|
||||
self,
|
||||
@@ -244,10 +301,6 @@ class MultiModalTransformerArgsPreprocessor:
|
||||
if cross_modality.sigma.ndim != 1:
|
||||
raise ValueError("Cross modality sigma must be a 1D tensor")
|
||||
|
||||
cross_timestep = cross_modality.sigma.view(
|
||||
modality.timesteps.shape[0], 1, *[1] * len(modality.timesteps.shape[2:])
|
||||
)
|
||||
|
||||
cross_pe = self.simple_preprocessor._prepare_positional_embeddings(
|
||||
positions=modality.positions[:, 0:1, :],
|
||||
inner_dim=self.audio_cross_attention_dim,
|
||||
@@ -258,7 +311,8 @@ class MultiModalTransformerArgsPreprocessor:
|
||||
)
|
||||
|
||||
cross_scale_shift_timestep, cross_gate_timestep = self._prepare_cross_attention_timestep(
|
||||
timestep=cross_timestep,
|
||||
modality_timesteps=modality.timesteps,
|
||||
cross_modality_sigma=cross_modality.sigma,
|
||||
timestep_scale_multiplier=self.simple_preprocessor.timestep_scale_multiplier,
|
||||
batch_size=transformer_args.x.shape[0],
|
||||
hidden_dtype=modality.latent.dtype,
|
||||
@@ -273,23 +327,23 @@ class MultiModalTransformerArgsPreprocessor:
|
||||
|
||||
def _prepare_cross_attention_timestep(
|
||||
self,
|
||||
timestep: torch.Tensor | None,
|
||||
modality_timesteps: torch.Tensor,
|
||||
cross_modality_sigma: torch.Tensor,
|
||||
timestep_scale_multiplier: int,
|
||||
batch_size: int,
|
||||
hidden_dtype: torch.dtype,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Prepare cross attention timestep embeddings."""
|
||||
timestep = timestep * timestep_scale_multiplier
|
||||
|
||||
"""Prepare A-V cross-attention AdaLN inputs."""
|
||||
av_ca_factor = self.av_ca_timestep_scale_multiplier / timestep_scale_multiplier
|
||||
|
||||
scale_shift_timestep, _ = self.cross_scale_shift_adaln(
|
||||
timestep.flatten(),
|
||||
(modality_timesteps * timestep_scale_multiplier).flatten(),
|
||||
hidden_dtype=hidden_dtype,
|
||||
)
|
||||
scale_shift_timestep = scale_shift_timestep.view(batch_size, -1, scale_shift_timestep.shape[-1])
|
||||
|
||||
gate_noise_timestep, _ = self.cross_gate_adaln(
|
||||
timestep.flatten() * av_ca_factor,
|
||||
(cross_modality_sigma * timestep_scale_multiplier * av_ca_factor).flatten(),
|
||||
hidden_dtype=hidden_dtype,
|
||||
)
|
||||
gate_noise_timestep = gate_noise_timestep.view(batch_size, -1, gate_noise_timestep.shape[-1])
|
||||
|
||||
@@ -56,6 +56,25 @@ if TYPE_CHECKING:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _memory_format_of(t: torch.Tensor, prefer_channels_last_3d: bool = False) -> torch.memory_format:
|
||||
"""Pick the memory format for a workspace allocation.
|
||||
When ``prefer_channels_last_3d`` is True and ``t`` is 5D, return
|
||||
``channels_last_3d`` regardless of ``t``'s current strides -- the
|
||||
workspace's ``.copy_(t)`` will transcribe the data into the new layout.
|
||||
This is needed because intermediate tensors inside the decoder (after
|
||||
``rearrange`` + slice + residual add in ``_upsample_forward_efficient``)
|
||||
are not NHWC-contiguous, so an auto-detect helper would silently fall
|
||||
back to NCHW for every workspace after the first upsample.
|
||||
Otherwise fall back to inspecting ``t``: ``channels_last_3d`` if ``t``
|
||||
already uses it, else contiguous.
|
||||
"""
|
||||
if prefer_channels_last_3d and t.dim() == 5:
|
||||
return torch.channels_last_3d
|
||||
if t.dim() == 5 and t.is_contiguous(memory_format=torch.channels_last_3d):
|
||||
return torch.channels_last_3d
|
||||
return torch.contiguous_format
|
||||
|
||||
|
||||
def _find_temporal_split_size(num_frames: int) -> int:
|
||||
"""Find chunk size for in-place temporal convolution.
|
||||
The chunk size ensures the last chunk has at least 3 frames
|
||||
@@ -132,6 +151,7 @@ def inplace_conv3d_temporal_chunked(workspace: torch.Tensor, conv: nn.Conv3d) ->
|
||||
workspace.shape[4],
|
||||
device=workspace.device,
|
||||
dtype=workspace.dtype,
|
||||
memory_format=_memory_format_of(workspace),
|
||||
)
|
||||
o_buf = torch.empty_like(x_buf)
|
||||
|
||||
@@ -190,6 +210,7 @@ def _causal_pad(x: torch.Tensor, pad_size: int) -> torch.Tensor:
|
||||
x.shape[4],
|
||||
device=x.device,
|
||||
dtype=x.dtype,
|
||||
memory_format=_memory_format_of(x),
|
||||
)
|
||||
padded[:, :, pad_size:].copy_(x)
|
||||
for i in range(pad_size):
|
||||
@@ -325,6 +346,7 @@ def _midblock_forward_efficient(
|
||||
causal: bool,
|
||||
timestep: torch.Tensor | None,
|
||||
generator: torch.Generator | None,
|
||||
prefer_channels_last_3d: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Memory-efficient ``UNetMidBlock3D`` forward.
|
||||
Allocates a single workspace buffer that is reused across all
|
||||
@@ -351,6 +373,7 @@ def _midblock_forward_efficient(
|
||||
hidden_states.shape[4],
|
||||
device=hidden_states.device,
|
||||
dtype=hidden_states.dtype,
|
||||
memory_format=_memory_format_of(hidden_states, prefer_channels_last_3d),
|
||||
)
|
||||
|
||||
for resnet in block.res_blocks:
|
||||
@@ -366,6 +389,7 @@ def _upsample_forward_efficient(
|
||||
block: DepthToSpaceUpsample,
|
||||
x: torch.Tensor,
|
||||
causal: bool,
|
||||
prefer_channels_last_3d: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Memory-efficient ``DepthToSpaceUpsample`` forward.
|
||||
For non-causal mode the input is copied into a workspace and the
|
||||
@@ -393,6 +417,7 @@ def _upsample_forward_efficient(
|
||||
if causal:
|
||||
x = _causal_pad_free_and_conv(x, block.conv)
|
||||
else:
|
||||
mem_fmt = _memory_format_of(x, prefer_channels_last_3d)
|
||||
workspace = torch.empty(
|
||||
x.shape[0],
|
||||
max(in_channels, out_channels),
|
||||
@@ -401,11 +426,12 @@ def _upsample_forward_efficient(
|
||||
x.shape[4],
|
||||
device=x.device,
|
||||
dtype=x.dtype,
|
||||
memory_format=mem_fmt,
|
||||
)
|
||||
workspace[:, :in_channels, 1:-1].copy_(x)
|
||||
del x
|
||||
inplace_conv3d_temporal_chunked(workspace, conv)
|
||||
x = workspace[:, :out_channels, 1:-1].contiguous()
|
||||
x = workspace[:, :out_channels, 1:-1].contiguous(memory_format=mem_fmt)
|
||||
del workspace
|
||||
|
||||
x = rearrange(
|
||||
@@ -418,7 +444,7 @@ def _upsample_forward_efficient(
|
||||
if block.stride[0] == 2:
|
||||
x = x[:, :, 1:, :, :]
|
||||
if block.residual:
|
||||
x = x + x_in
|
||||
x.add_(x_in)
|
||||
del x_in
|
||||
return x
|
||||
|
||||
@@ -434,12 +460,14 @@ def _final_norm_and_conv_out(
|
||||
causal: bool,
|
||||
scaled_timestep: torch.Tensor | None,
|
||||
batch_size: int,
|
||||
prefer_channels_last_3d: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Workspace-based final norm + [ada] + SiLU + conv_out + unpatchify."""
|
||||
conv_out_mod: CausalConv3d = decoder.conv_out # type: ignore[assignment]
|
||||
conv_out = conv_out_mod.conv
|
||||
feature_channels = sample.shape[1]
|
||||
|
||||
mem_fmt = _memory_format_of(sample, prefer_channels_last_3d)
|
||||
workspace = torch.empty(
|
||||
sample.shape[0],
|
||||
max(feature_channels, conv_out.out_channels),
|
||||
@@ -448,6 +476,7 @@ def _final_norm_and_conv_out(
|
||||
sample.shape[4],
|
||||
device=sample.device,
|
||||
dtype=sample.dtype,
|
||||
memory_format=mem_fmt,
|
||||
)
|
||||
workspace[:, :feature_channels, 1:-1].copy_(sample)
|
||||
del sample
|
||||
@@ -485,7 +514,7 @@ def _final_norm_and_conv_out(
|
||||
del padded
|
||||
else:
|
||||
inplace_conv3d_temporal_chunked(workspace, conv_out)
|
||||
result = workspace[:, : conv_out.out_channels, 1:-1].contiguous()
|
||||
result = workspace[:, : conv_out.out_channels, 1:-1].contiguous(memory_format=mem_fmt)
|
||||
del workspace, interior
|
||||
|
||||
return unpatchify(result, patch_size_hw=decoder.patch_size, patch_size_t=1)
|
||||
@@ -507,6 +536,9 @@ def _memory_efficient_forward(
|
||||
``UNetMidBlock3D`` and ``DepthToSpaceUpsample`` blocks use efficient
|
||||
paths; standalone ``ResnetBlock3D`` blocks fall back to the standard
|
||||
forward. The final norm + ada + SiLU + conv_out is also workspace-based.
|
||||
All workspaces are allocated ``channels_last_3d`` so cuDNN's NHWC 3D
|
||||
conv kernels run end-to-end. The caller (:func:`enable_memory_efficient_decode`)
|
||||
is responsible for converting the input sample and decoder weights to NHWC.
|
||||
"""
|
||||
causal = decoder.causal
|
||||
batch_size = sample.shape[0]
|
||||
@@ -537,6 +569,11 @@ def _memory_efficient_forward(
|
||||
raise ValueError("'timestep' required when timestep_conditioning=True")
|
||||
scaled_timestep = timestep * decoder.timestep_scale_multiplier.to(sample)
|
||||
|
||||
# Workspaces are unconditionally NHWC: rearrange + slice + residual-add
|
||||
# inside _upsample_forward_efficient produces NCHW-default output, so
|
||||
# per-tensor inspection would silently fall back to NCHW for every
|
||||
# workspace after the first upsample.
|
||||
|
||||
# --- Up blocks (dispatch to efficient path per block type) ---
|
||||
for up_block in decoder.up_blocks:
|
||||
if isinstance(up_block, UNetMidBlock3D):
|
||||
@@ -546,15 +583,16 @@ def _memory_efficient_forward(
|
||||
causal=causal,
|
||||
timestep=scaled_timestep if decoder.timestep_conditioning else None,
|
||||
generator=generator,
|
||||
prefer_channels_last_3d=True,
|
||||
)
|
||||
elif isinstance(up_block, DepthToSpaceUpsample):
|
||||
sample = _upsample_forward_efficient(up_block, sample, causal=causal)
|
||||
sample = _upsample_forward_efficient(up_block, sample, causal=causal, prefer_channels_last_3d=True)
|
||||
elif isinstance(up_block, ResnetBlock3D):
|
||||
sample = up_block(sample, causal=causal, generator=generator)
|
||||
else:
|
||||
sample = up_block(sample, causal=causal)
|
||||
|
||||
return _final_norm_and_conv_out(decoder, sample, causal, scaled_timestep, batch_size)
|
||||
return _final_norm_and_conv_out(decoder, sample, causal, scaled_timestep, batch_size, prefer_channels_last_3d=True)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -564,6 +602,10 @@ def _memory_efficient_forward(
|
||||
|
||||
def enable_memory_efficient_decode(decoder: nn.Module) -> nn.Module:
|
||||
"""Patch a ``VideoDecoder`` to use the memory-efficient forward path.
|
||||
The mem-efficient path runs the decoder in ``channels_last_3d`` memory
|
||||
format: weights and inputs are converted on first call so cuDNN's NHWC
|
||||
3D conv kernels are used (~2x faster, avoids the large vol2col scratch
|
||||
buffer of the NCHW path).
|
||||
The original ``forward`` is saved as ``decoder._original_forward`` so
|
||||
that it can be restored later with :func:`disable_memory_efficient_decode`.
|
||||
"""
|
||||
@@ -577,12 +619,20 @@ def enable_memory_efficient_decode(decoder: nn.Module) -> nn.Module:
|
||||
return decoder
|
||||
|
||||
original_forward = decoder.forward
|
||||
weights_converted = False
|
||||
|
||||
def efficient_forward(
|
||||
sample: torch.Tensor,
|
||||
timestep: torch.Tensor | None = None,
|
||||
generator: torch.Generator | None = None,
|
||||
) -> torch.Tensor:
|
||||
nonlocal weights_converted
|
||||
if sample.dim() == 5:
|
||||
if not weights_converted:
|
||||
# Lazy: weights are real by first-call time (meta -> loader -> here).
|
||||
decoder.to(memory_format=torch.channels_last_3d)
|
||||
weights_converted = True
|
||||
sample = sample.to(memory_format=torch.channels_last_3d)
|
||||
return _memory_efficient_forward(decoder, sample, timestep, generator)
|
||||
|
||||
decoder._original_forward = original_forward # type: ignore[attr-defined]
|
||||
|
||||
@@ -2,7 +2,9 @@ from ltx_core.quantization.fp8_cast import (
|
||||
TRANSFORMER_LINEAR_DOWNCAST_MAP,
|
||||
UPCAST_DURING_INFERENCE,
|
||||
UpcastWithStochasticRounding,
|
||||
fp8_cast_fuse_rule,
|
||||
)
|
||||
from ltx_core.quantization.fp8_scaled_mm import fp8_scaled_mm_fuse_rule
|
||||
from ltx_core.quantization.policy import QuantizationPolicy
|
||||
|
||||
__all__ = [
|
||||
@@ -10,4 +12,6 @@ __all__ = [
|
||||
"UPCAST_DURING_INFERENCE",
|
||||
"QuantizationPolicy",
|
||||
"UpcastWithStochasticRounding",
|
||||
"fp8_cast_fuse_rule",
|
||||
"fp8_scaled_mm_fuse_rule",
|
||||
]
|
||||
|
||||
@@ -1,9 +1,15 @@
|
||||
from pathlib import Path
|
||||
|
||||
import safetensors
|
||||
import torch
|
||||
|
||||
from ltx_core.loader.fuse_loras import FuseRule, bf16_fuse_rule
|
||||
from ltx_core.loader.kernels import TRITON_AVAILABLE
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import StateDict
|
||||
from ltx_core.loader.sd_ops import KeyValueOperationResult, SDOps
|
||||
from ltx_core.model.transformer.model import LTXModel
|
||||
from ltx_core.quantization.policy import QuantizationPolicy
|
||||
|
||||
BLOCK_SIZE = 1024
|
||||
|
||||
@@ -113,58 +119,65 @@ def _replace_fwd_with_upcast(layer: torch.nn.Linear, with_stochastic_rounding: b
|
||||
)
|
||||
|
||||
|
||||
# Module-name suffixes for the Linears that participate in fp8 cast. Used by
|
||||
# both the upcast matcher and the sd_ops downcast map so the two cannot drift.
|
||||
# - ``.to_q`` / ``.to_k`` / ``.to_v`` / ``.to_out.0`` have a leading dot so they
|
||||
# only match the attention Linears at ``...attnN.to_q`` etc.
|
||||
# - ``ff.net.0.proj`` / ``ff.net.2`` are intentionally **dotless** so they match
|
||||
# both video FF (``...ff.net.0.proj``) and audio FF (``...audio_ff.net.0.proj``).
|
||||
_FP8_CAST_KEY_PREFIX = "transformer_blocks."
|
||||
_FP8_CAST_LINEAR_SUFFIXES: tuple[str, ...] = (
|
||||
".to_q",
|
||||
".to_k",
|
||||
".to_v",
|
||||
".to_out.0",
|
||||
"ff.net.0.proj",
|
||||
"ff.net.2",
|
||||
)
|
||||
|
||||
|
||||
def _is_fp8_cast_linear(module_name: str) -> bool:
|
||||
"""Return True if *module_name* names a Linear that should be fp8-cast."""
|
||||
if _FP8_CAST_KEY_PREFIX not in module_name:
|
||||
return False
|
||||
return any(module_name.endswith(suffix) for suffix in _FP8_CAST_LINEAR_SUFFIXES)
|
||||
|
||||
|
||||
def _amend_forward_with_upcast(
|
||||
model: torch.nn.Module, with_stochastic_rounding: bool = False, seed: int = 0
|
||||
) -> torch.nn.Module:
|
||||
"""
|
||||
Replace the forward method of the model's Linear layers to forward
|
||||
with upcast and optional stochastic rounding.
|
||||
Replace the forward method of the fp8-cast Linear layers (per
|
||||
:data:`_FP8_CAST_LINEAR_SUFFIXES`) to forward with upcast and optional
|
||||
stochastic rounding.
|
||||
Only the Linears whose weights are downcast by :data:`TRANSFORMER_LINEAR_DOWNCAST_MAP`
|
||||
are retyped. Linears outside that subset (e.g. ``to_gate_logits``) are left as
|
||||
plain ``nn.Linear`` so the meta-model param dtype matches the loaded checkpoint
|
||||
dtype.
|
||||
"""
|
||||
for m in model.modules():
|
||||
if isinstance(m, (torch.nn.Linear)):
|
||||
for name, m in model.named_modules():
|
||||
if isinstance(m, torch.nn.Linear) and _is_fp8_cast_linear(name):
|
||||
_replace_fwd_with_upcast(m, with_stochastic_rounding, seed)
|
||||
return model
|
||||
|
||||
|
||||
TRANSFORMER_LINEAR_DOWNCAST_MAP = (
|
||||
SDOps("TRANSFORMER_LINEAR_DOWNCAST_MAP")
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix=".to_q.weight", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix=".to_q.bias", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix=".to_k.weight", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix=".to_k.bias", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix=".to_v.weight", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix=".to_v.bias", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix=".to_out.0.weight", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix=".to_out.0.bias", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix="ff.net.0.proj.weight", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix="ff.net.0.proj.bias", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix="ff.net.2.weight", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
.with_kv_operation(
|
||||
key_prefix="transformer_blocks.", key_suffix="ff.net.2.bias", operation=_naive_weight_or_bias_downcast
|
||||
)
|
||||
)
|
||||
def _build_transformer_linear_downcast_map() -> SDOps:
|
||||
"""Build the sd_ops downcast map from the same suffix registry as the matcher."""
|
||||
ops = SDOps("TRANSFORMER_LINEAR_DOWNCAST_MAP")
|
||||
for suffix in _FP8_CAST_LINEAR_SUFFIXES:
|
||||
ops = ops.with_kv_operation(
|
||||
key_prefix=_FP8_CAST_KEY_PREFIX,
|
||||
key_suffix=suffix + ".weight",
|
||||
operation=_naive_weight_or_bias_downcast,
|
||||
).with_kv_operation(
|
||||
key_prefix=_FP8_CAST_KEY_PREFIX,
|
||||
key_suffix=suffix + ".bias",
|
||||
operation=_naive_weight_or_bias_downcast,
|
||||
)
|
||||
return ops
|
||||
|
||||
|
||||
TRANSFORMER_LINEAR_DOWNCAST_MAP = _build_transformer_linear_downcast_map()
|
||||
|
||||
UPCAST_DURING_INFERENCE = ModuleOps(
|
||||
name="upcast_fp8_during_linear_forward",
|
||||
@@ -186,3 +199,139 @@ class UpcastWithStochasticRounding(ModuleOps):
|
||||
matcher=lambda model: isinstance(model, LTXModel),
|
||||
mutator=lambda model: _amend_forward_with_upcast(model, True, seed),
|
||||
)
|
||||
|
||||
|
||||
def fuse_cast_fp8_weight(
|
||||
delta_bf16: torch.Tensor,
|
||||
weight_fp8: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Return ``(delta_bf16 + dequantize(weight_fp8)).to(weight_fp8.dtype)``.
|
||||
CUDA with Triton uses stochastic rounding via the fused kernel; otherwise
|
||||
falls back to a deterministic bf16 add. ``delta_bf16`` is the bf16
|
||||
accumulator and is mutated in place.
|
||||
"""
|
||||
if delta_bf16.dtype != torch.bfloat16:
|
||||
raise ValueError(f"delta_bf16 must be bfloat16, got {delta_bf16.dtype}")
|
||||
if str(weight_fp8.device).startswith("cuda") and TRITON_AVAILABLE:
|
||||
fused_add_round_launch(delta_bf16, weight_fp8, seed=0)
|
||||
else:
|
||||
delta_bf16.add_(weight_fp8.to(dtype=torch.bfloat16))
|
||||
return delta_bf16.to(dtype=weight_fp8.dtype)
|
||||
|
||||
|
||||
def _fp8_cast_fuse(
|
||||
key: str,
|
||||
weight: torch.Tensor,
|
||||
deltas: torch.Tensor,
|
||||
model_sd: StateDict,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Cast the dequantized FP8 weight + BF16 deltas back to ``weight.dtype``
|
||||
(FP8) via the fused-add-round kernel on CUDA.
|
||||
Only a subset of linears are FP8-downcast (see ``TRANSFORMER_LINEAR_DOWNCAST_MAP``);
|
||||
LoRAs may also target layers left in BF16 (e.g. audio ``add_q/k/v_proj``, cross-modal
|
||||
projections). For those, fall back to a plain BF16 fuse.
|
||||
"""
|
||||
if weight.dtype not in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
return bf16_fuse_rule(key, weight, deltas, model_sd)
|
||||
return {key: fuse_cast_fp8_weight(deltas, weight)}
|
||||
|
||||
|
||||
fp8_cast_fuse_rule = FuseRule(aggregation_dtype=torch.bfloat16, fuse_fn=_fp8_cast_fuse)
|
||||
|
||||
|
||||
# Raw safetensors storage prefix shared by every diffusion-transformer
|
||||
# parameter (and every prequant `*_scale` sibling). Verified against
|
||||
# ltx-2.3-22b-{dev,distilled}-fp8.safetensors: 2924/2924 and 2992/2992 of
|
||||
# the scale keys start with this exact prefix.
|
||||
_RAW_DIFFUSION_MODEL_PREFIX = "model.diffusion_model."
|
||||
|
||||
|
||||
def _read_scales(checkpoint_path: str | Path) -> dict[str, torch.Tensor]:
|
||||
"""Return ``{post_rename_param_key: scale_tensor}`` for every prequant
|
||||
``*_scale`` sibling in *checkpoint_path*.
|
||||
Keys are returned in the post-rename form the loader will pass to the
|
||||
sd-op (e.g. ``transformer_blocks.0.attn1.to_q.weight``) -- the raw
|
||||
``model.diffusion_model.`` prefix and the ``_scale`` suffix are both
|
||||
stripped. Catches both ``.weight_scale`` and ``.bias_scale``; the
|
||||
latter is absent in the current LTX-2.3 prequant checkpoints but
|
||||
accepted for forward compatibility.
|
||||
"""
|
||||
out: dict[str, torch.Tensor] = {}
|
||||
with safetensors.safe_open(str(checkpoint_path), framework="pt", device="cpu") as h:
|
||||
raw_keys = h.keys()
|
||||
for k in raw_keys:
|
||||
if not k.endswith("_scale"):
|
||||
continue
|
||||
if not k.startswith(_RAW_DIFFUSION_MODEL_PREFIX):
|
||||
raise ValueError(
|
||||
f"Scale key {k!r} does not start with the expected raw prefix {_RAW_DIFFUSION_MODEL_PREFIX!r}"
|
||||
)
|
||||
param_key = k.removeprefix(_RAW_DIFFUSION_MODEL_PREFIX).removesuffix("_scale")
|
||||
out[param_key] = h.get_tensor(k)
|
||||
return out
|
||||
|
||||
|
||||
def _build_prequant_fold_sd_ops(scales: dict[str, torch.Tensor]) -> SDOps:
|
||||
"""Build sd-ops that fold prequant ``*_scale`` siblings into their parent
|
||||
tensor at load time.
|
||||
*scales* is keyed by the **post-rename** param key (e.g.
|
||||
``transformer_blocks.0.attn1.to_q.weight``); see :func:`_read_scales`.
|
||||
Four ``with_kv_operation`` entries (symmetric for ``.weight`` and ``.bias``):
|
||||
* ``.weight`` / ``.bias`` -> if a sibling scale exists in *scales*, fold;
|
||||
then delegate to ``TRANSFORMER_LINEAR_DOWNCAST_MAP`` (downcast covered
|
||||
Linears, pass everything else through). Without a scale, delegate
|
||||
directly.
|
||||
* ``.weight_scale`` / ``.bias_scale`` -> drop (the scale is consumed by
|
||||
the fold). Raises if the scale key doesn't correspond to a known
|
||||
entry in *scales* -- that means the file shipped a scale we didn't
|
||||
pre-register, which would silently desync the fold.
|
||||
"""
|
||||
|
||||
def _on_param(param_key: str, value: torch.Tensor) -> list[KeyValueOperationResult]:
|
||||
scale = scales.get(param_key)
|
||||
if scale is None:
|
||||
return TRANSFORMER_LINEAR_DOWNCAST_MAP.apply_to_key_value(param_key, value)
|
||||
scale = scale.to(device=value.device)
|
||||
if scale.ndim != 0:
|
||||
raise ValueError(f"Unsupported scale shape {tuple(scale.shape)} for {param_key}")
|
||||
bf16 = (value.to(torch.float32) * scale).to(torch.bfloat16)
|
||||
# Delegate the final fp8-vs-bf16 decision to the downcast map: Linears
|
||||
# outside the fp8 subset (e.g. to_gate_logits) stay bf16 to match the
|
||||
# plain nn.Linear that the upcast matcher leaves untouched.
|
||||
return TRANSFORMER_LINEAR_DOWNCAST_MAP.apply_to_key_value(param_key, bf16)
|
||||
|
||||
def _drop_scale(scale_key: str, _value: torch.Tensor) -> list[KeyValueOperationResult]:
|
||||
param_key = scale_key.removesuffix("_scale")
|
||||
if param_key not in scales:
|
||||
raise ValueError(
|
||||
f"Scale key {scale_key!r} has no matching entry in the prequant scales dict; "
|
||||
f"_read_scales and the loader's rename map have drifted"
|
||||
)
|
||||
return []
|
||||
|
||||
# Register the drop ops first so the dict-membership sanity check is the
|
||||
# earliest sd-op that can fire on a scale key -- we crash on a stray scale
|
||||
# before any silently mismatched fold has a chance to land in the state
|
||||
# dict. Registration order is irrelevant for correctness (no overlap
|
||||
# between matchers) but communicates intent.
|
||||
return (
|
||||
SDOps("FP8_CAST_PREQUANT_AWARE")
|
||||
.with_kv_operation(key_suffix=".weight_scale", operation=_drop_scale)
|
||||
.with_kv_operation(key_suffix=".bias_scale", operation=_drop_scale)
|
||||
.with_kv_operation(key_suffix=".weight", operation=_on_param)
|
||||
.with_kv_operation(key_suffix=".bias", operation=_on_param)
|
||||
)
|
||||
|
||||
|
||||
def build_policy(checkpoint_path: str | Path) -> QuantizationPolicy:
|
||||
"""FP8 casting with upcasting during inference.
|
||||
*checkpoint_path* is required (mirroring ``fp8_scaled_mm.build_policy``).
|
||||
For prequantized fp8 checkpoints, sibling ``*_scale`` tensors (weight or
|
||||
bias) are folded into the parent at load time.
|
||||
"""
|
||||
scales = _read_scales(checkpoint_path)
|
||||
return QuantizationPolicy(
|
||||
sd_ops=_build_prequant_fold_sd_ops(scales),
|
||||
module_ops=(UPCAST_DURING_INFERENCE,),
|
||||
fuse_rule=fp8_cast_fuse_rule,
|
||||
)
|
||||
|
||||
@@ -5,8 +5,11 @@ from typing import Callable
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ltx_core.loader.fuse_loras import FuseRule, bf16_fuse_rule
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import StateDict
|
||||
from ltx_core.model.transformer import LTXModel
|
||||
from ltx_core.quantization.policy import QuantizationPolicy
|
||||
from ltx_core.quantization.trtllm_scaled_usable import trtllm_scaled_mm_usable
|
||||
|
||||
|
||||
@@ -173,3 +176,42 @@ def get_fp8_swap_module_ops(checkpoint_path: str) -> tuple[ModuleOps, ...]:
|
||||
mutator=lambda model: _swap_linears_to_fp8(model, _should_swap),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _fp8_scaled_mm_fuse(
|
||||
key: str,
|
||||
weight: torch.Tensor,
|
||||
deltas: torch.Tensor,
|
||||
model_sd: StateDict,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Dequantize via ``weight.float() * weight_scale``, add the BF16 delta,
|
||||
and re-quantize to FP8 with a fresh per-tensor scale.
|
||||
Layers that were not swapped to scaled FP8 (e.g. small embedder linears
|
||||
excluded from the auto-discovered swap set) stay BF16 and have no
|
||||
``.weight_scale`` companion -- for those, fall back to a plain bf16 fuse.
|
||||
"""
|
||||
scale_key = key.replace(".weight", ".weight_scale")
|
||||
if scale_key not in model_sd.sd:
|
||||
return bf16_fuse_rule(key, weight, deltas, model_sd)
|
||||
weight_scale = model_sd.sd[scale_key]
|
||||
original_weight = weight.to(torch.float32) * weight_scale
|
||||
new_weight = original_weight + deltas.to(torch.float32)
|
||||
new_fp8_weight, new_weight_scale = quantize_weight_to_fp8_per_tensor(new_weight)
|
||||
return {key: new_fp8_weight, scale_key: new_weight_scale}
|
||||
|
||||
|
||||
fp8_scaled_mm_fuse_rule = FuseRule(aggregation_dtype=torch.bfloat16, fuse_fn=_fp8_scaled_mm_fuse)
|
||||
|
||||
|
||||
def build_policy(checkpoint_path: str) -> QuantizationPolicy:
|
||||
"""FP8 scaled matmul for checkpoints pre-quantized with per-tensor scales.
|
||||
The set of layers to swap to ``FP8Linear`` is discovered from the
|
||||
checkpoint's ``.weight_scale`` tensors via suffix-matching against the
|
||||
model's named modules. Requires a pre-quantized checkpoint; for BF16
|
||||
checkpoints, use :func:`ltx_core.quantization.fp8_cast.build_policy`.
|
||||
"""
|
||||
return QuantizationPolicy(
|
||||
sd_ops=None,
|
||||
module_ops=get_fp8_swap_module_ops(checkpoint_path),
|
||||
fuse_rule=fp8_scaled_mm_fuse_rule,
|
||||
)
|
||||
|
||||
@@ -1,48 +1,24 @@
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
from ltx_core.loader.fuse_loras import FuseRule, bf16_fuse_rule
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.sd_ops import SDOps
|
||||
from ltx_core.quantization.fp8_cast import TRANSFORMER_LINEAR_DOWNCAST_MAP, UPCAST_DURING_INFERENCE
|
||||
from ltx_core.quantization.fp8_scaled_mm import get_fp8_swap_module_ops
|
||||
from ltx_core.model.model_protocol import ModelConfigurator
|
||||
from ltx_core.model.transformer.model import LTXModel
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QuantizationPolicy:
|
||||
"""Configuration for model quantization during loading.
|
||||
Attributes:
|
||||
kind: Discriminator for the policy variant.
|
||||
sd_ops: State-dict operations applied to each tensor during load.
|
||||
module_ops: Post-load module transformations applied to the meta model.
|
||||
model_configurator: Configurator class to use when constructing the transformer.
|
||||
fuse_rule: How LoRA deltas merge into this policy's weight layout.
|
||||
Default ``bf16_fuse_rule`` is used when no policy is configured.
|
||||
"""
|
||||
|
||||
class Kind(str, Enum):
|
||||
FP8_CAST = "fp8_cast"
|
||||
FP8_SCALED_MM = "fp8_scaled_mm"
|
||||
|
||||
kind: Kind
|
||||
sd_ops: SDOps | None = None
|
||||
module_ops: tuple[ModuleOps, ...] = ()
|
||||
|
||||
@classmethod
|
||||
def fp8_cast(cls) -> "QuantizationPolicy":
|
||||
"""FP8 casting with upcasting during inference."""
|
||||
return cls(
|
||||
kind=cls.Kind.FP8_CAST,
|
||||
sd_ops=TRANSFORMER_LINEAR_DOWNCAST_MAP,
|
||||
module_ops=(UPCAST_DURING_INFERENCE,),
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def fp8_scaled_mm(cls, checkpoint_path: str) -> "QuantizationPolicy":
|
||||
"""FP8 scaled matmul for checkpoints pre-quantized with per-tensor scales.
|
||||
The set of layers to swap to ``FP8Linear`` is discovered from the
|
||||
checkpoint's ``.weight_scale`` tensors via suffix-matching against the
|
||||
model's named modules. Requires a pre-quantized checkpoint; for BF16
|
||||
checkpoints, use :meth:`fp8_cast` instead.
|
||||
"""
|
||||
return cls(
|
||||
kind=cls.Kind.FP8_SCALED_MM,
|
||||
sd_ops=None,
|
||||
module_ops=get_fp8_swap_module_ops(checkpoint_path),
|
||||
)
|
||||
model_configurator: type[ModelConfigurator[LTXModel]] | None = None
|
||||
fuse_rule: FuseRule = bf16_fuse_rule
|
||||
|
||||
@@ -30,21 +30,31 @@ class GemmaTextEncoder(torch.nn.Module):
|
||||
|
||||
def encode(
|
||||
self,
|
||||
text: str,
|
||||
prompts: list[str],
|
||||
padding_side: str = "left", # noqa: ARG002
|
||||
) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
|
||||
"""Run Gemma LLM and return raw hidden states + attention mask.
|
||||
Calls the inner model (self.model.model) to skip lm_head logits computation (~500 MiB saving).
|
||||
Returns:
|
||||
(hidden_states, attention_mask) where hidden_states is a tuple of per-layer tensors.
|
||||
) -> list[tuple[tuple[torch.Tensor, ...], torch.Tensor]]:
|
||||
"""Run a single fused Gemma forward over a batch of prompts.
|
||||
Calls the inner model (self.model.model) to skip lm_head logits computation
|
||||
(~500 MiB saving). The tokenizer pads every prompt to ``max_length`` (1024),
|
||||
so the inputs stack into a single ``[N, 1024]`` batch with no further padding
|
||||
logic; per-prompt outputs are sliced back to ``[1, 1024, D]`` / ``[1, 1024]``
|
||||
in the original order.
|
||||
"""
|
||||
token_pairs = self.tokenizer.tokenize_with_weights(text)["gemma"]
|
||||
input_ids = torch.tensor([[t[0] for t in token_pairs]], device=self.model.device)
|
||||
attention_mask = torch.tensor([[w[1] for w in token_pairs]], device=self.model.device)
|
||||
if not prompts:
|
||||
return []
|
||||
tokenized = [self.tokenizer.tokenize_with_weights(t)["gemma"] for t in prompts]
|
||||
input_ids = torch.tensor(
|
||||
[[tok for tok, _ in pairs] for pairs in tokenized],
|
||||
device=self.model.device,
|
||||
)
|
||||
attention_mask = torch.tensor(
|
||||
[[w for _, w in pairs] for pairs in tokenized],
|
||||
device=self.model.device,
|
||||
)
|
||||
outputs = self.model.model(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True)
|
||||
hidden_states = outputs.hidden_states
|
||||
del outputs
|
||||
return hidden_states, attention_mask
|
||||
return [(tuple(h[i : i + 1] for h in hidden_states), attention_mask[i : i + 1]) for i in range(len(prompts))]
|
||||
|
||||
# --- Prompt enhancement methods ---
|
||||
|
||||
@@ -183,7 +193,7 @@ def module_ops_from_gemma_root(gemma_root: str) -> tuple[ModuleOps, ...]:
|
||||
return module
|
||||
|
||||
def load_processor(module: GemmaTextEncoder) -> GemmaTextEncoder:
|
||||
image_processor = AutoImageProcessor.from_pretrained(processor_root, local_files_only=True)
|
||||
image_processor = AutoImageProcessor.from_pretrained(processor_root, local_files_only=True, use_fast=False)
|
||||
if not module.tokenizer:
|
||||
raise ValueError("Tokenizer model operation must be performed before processor model operation")
|
||||
module.processor = Gemma3Processor(image_processor=image_processor, tokenizer=module.tokenizer.tokenizer)
|
||||
|
||||
+110
-27
@@ -1,4 +1,5 @@
|
||||
import torch
|
||||
import transformers
|
||||
from transformers import Gemma3Config
|
||||
from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS
|
||||
from transformers.models.gemma3 import Gemma3ForConditionalGeneration
|
||||
@@ -19,6 +20,8 @@ from ltx_core.text_encoders.gemma.feature_extractor import (
|
||||
FeatureExtractorV2,
|
||||
)
|
||||
|
||||
_TRANSFORMERS_V5: bool = int(transformers.__version__.split(".", 1)[0]) >= 5
|
||||
|
||||
|
||||
class GemmaTextEncoderConfigurator(ModelConfigurator[GemmaTextEncoder]):
|
||||
@classmethod
|
||||
@@ -100,26 +103,45 @@ def _create_feature_extractor(transformer_config: dict) -> torch.nn.Module:
|
||||
|
||||
# --- Split SDOps: Gemma LLM keys vs Embeddings Processor keys ---
|
||||
|
||||
GEMMA_LLM_KEY_OPS = (
|
||||
SDOps("GEMMA_LLM_KEY_OPS")
|
||||
# 1. Map language model layers (note the double .model prefix)
|
||||
.with_matching(prefix="language_model.model.")
|
||||
.with_replacement("language_model.model.", "model.model.language_model.")
|
||||
# 2. Map the Vision Tower
|
||||
.with_matching(prefix="vision_tower.")
|
||||
.with_replacement("vision_tower.", "model.model.vision_tower.")
|
||||
# 3. Map the Multi-Modal Projector
|
||||
.with_matching(prefix="multi_modal_projector.")
|
||||
.with_replacement("multi_modal_projector.", "model.model.multi_modal_projector.")
|
||||
# 4. Duplicate embed_tokens to lm_head (needed for prompt enhancement via generate())
|
||||
.with_kv_operation(
|
||||
operation=lambda key, value: [
|
||||
KeyValueOperationResult(key, value),
|
||||
KeyValueOperationResult("model.lm_head.weight", value),
|
||||
],
|
||||
key_prefix="model.model.language_model.embed_tokens.weight",
|
||||
|
||||
def _build_gemma_llm_key_ops(*, transformers_v5: bool) -> SDOps:
|
||||
"""Build the checkpoint-key remapping for the Gemma multimodal encoder.
|
||||
The vision-tower mapping differs between transformers <5 and >=5 because
|
||||
upstream PR https://github.com/huggingface/transformers/pull/39847 flattened
|
||||
``Gemma3ForConditionalGeneration.model.vision_tower.vision_model`` into
|
||||
``model.vision_tower``. Checkpoints continue to ship the legacy
|
||||
``vision_tower.vision_model.*`` prefix, so we strip the inner ``vision_model.``
|
||||
when targeting v5 and pass it through unchanged for v4.
|
||||
"""
|
||||
base = (
|
||||
SDOps("GEMMA_LLM_KEY_OPS")
|
||||
# 1. Map language model layers (note the double .model prefix)
|
||||
.with_matching(prefix="language_model.model.")
|
||||
.with_replacement("language_model.model.", "model.model.language_model.")
|
||||
# 2. Map the Vision Tower (version-dependent — see docstring)
|
||||
.with_matching(prefix="vision_tower.")
|
||||
)
|
||||
)
|
||||
if transformers_v5:
|
||||
base = base.with_replacement("vision_tower.vision_model.", "model.model.vision_tower.")
|
||||
else:
|
||||
base = base.with_replacement("vision_tower.", "model.model.vision_tower.")
|
||||
return (
|
||||
base
|
||||
# 3. Map the Multi-Modal Projector
|
||||
.with_matching(prefix="multi_modal_projector.")
|
||||
.with_replacement("multi_modal_projector.", "model.model.multi_modal_projector.")
|
||||
# 4. Duplicate embed_tokens to lm_head (needed for prompt enhancement via generate())
|
||||
.with_kv_operation(
|
||||
operation=lambda key, value: [
|
||||
KeyValueOperationResult(key, value),
|
||||
KeyValueOperationResult("model.lm_head.weight", value),
|
||||
],
|
||||
key_prefix="model.model.language_model.embed_tokens.weight",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
GEMMA_LLM_KEY_OPS = _build_gemma_llm_key_ops(transformers_v5=_TRANSFORMERS_V5)
|
||||
|
||||
EMBEDDINGS_PROCESSOR_KEY_OPS = (
|
||||
SDOps("EMBEDDINGS_PROCESSOR_KEY_OPS")
|
||||
@@ -152,24 +174,85 @@ VIDEO_ONLY_EMBEDDINGS_PROCESSOR_KEY_OPS = (
|
||||
)
|
||||
|
||||
|
||||
def _resolve_local_base_freq(config: object) -> float:
|
||||
rope_parameters = getattr(config, "rope_parameters", None)
|
||||
if isinstance(rope_parameters, dict) and "sliding_attention" in rope_parameters:
|
||||
sliding = rope_parameters["sliding_attention"]
|
||||
if isinstance(sliding, dict) and "rope_theta" in sliding:
|
||||
return float(sliding["rope_theta"])
|
||||
if hasattr(config, "rope_local_base_freq"):
|
||||
return float(config.rope_local_base_freq)
|
||||
raise AttributeError(
|
||||
"Gemma text_config exposes neither rope_local_base_freq nor rope_parameters['sliding_attention']['rope_theta']"
|
||||
)
|
||||
|
||||
|
||||
def _resolve_full_rope_type(config: object) -> str:
|
||||
rope_parameters = getattr(config, "rope_parameters", None)
|
||||
if isinstance(rope_parameters, dict) and "full_attention" in rope_parameters:
|
||||
full = rope_parameters["full_attention"]
|
||||
if isinstance(full, dict) and "rope_type" in full:
|
||||
return str(full["rope_type"])
|
||||
rope_scaling = getattr(config, "rope_scaling", None)
|
||||
if rope_scaling is not None:
|
||||
if isinstance(rope_scaling, dict):
|
||||
if "rope_type" in rope_scaling:
|
||||
return str(rope_scaling["rope_type"])
|
||||
elif hasattr(rope_scaling, "rope_type"):
|
||||
return str(rope_scaling.rope_type)
|
||||
raise AttributeError(
|
||||
"Gemma text_config exposes neither rope_scaling.rope_type nor rope_parameters['full_attention']['rope_type']"
|
||||
)
|
||||
|
||||
|
||||
def _populate_rotary_v4(l_model: torch.nn.Module, config: object) -> None:
|
||||
"""transformers <5 layout: separate ``rotary_emb_local`` + ``rotary_emb``."""
|
||||
dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
|
||||
base = _resolve_local_base_freq(config)
|
||||
local_inv = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(dtype=torch.float) / dim))
|
||||
full_inv, _ = ROPE_INIT_FUNCTIONS[_resolve_full_rope_type(config)](config)
|
||||
l_model.rotary_emb_local.register_buffer("inv_freq", local_inv)
|
||||
l_model.rotary_emb.register_buffer("inv_freq", full_inv)
|
||||
|
||||
|
||||
def _populate_rotary_v5(l_model: torch.nn.Module, config: object) -> None:
|
||||
"""transformers >=5 layout: single ``rotary_emb`` with per-layer-type buffers.
|
||||
Mirrors ``Gemma3PreTrainedModel._init_weights`` for ``Gemma3RotaryEmbedding``
|
||||
so meta-built models reach the same numerical state as a from_pretrained load.
|
||||
"""
|
||||
rope_emb = l_model.rotary_emb
|
||||
for layer_type in dict.fromkeys(config.layer_types):
|
||||
rope_params = config.rope_parameters[layer_type]
|
||||
if rope_params is None:
|
||||
continue
|
||||
rope_type = rope_params["rope_type"]
|
||||
if rope_type == "default":
|
||||
inv_freq, attn_scaling = rope_emb.compute_default_rope_parameters(config, layer_type=layer_type)
|
||||
else:
|
||||
inv_freq, attn_scaling = ROPE_INIT_FUNCTIONS[rope_type](config, layer_type=layer_type)
|
||||
rope_emb.register_buffer(f"{layer_type}_inv_freq", inv_freq, persistent=False)
|
||||
rope_emb.register_buffer(f"{layer_type}_original_inv_freq", inv_freq.clone(), persistent=False)
|
||||
setattr(rope_emb, f"{layer_type}_attention_scaling", attn_scaling)
|
||||
|
||||
|
||||
def create_and_populate(module: GemmaTextEncoder) -> GemmaTextEncoder:
|
||||
model = module.model
|
||||
v_model = model.model.vision_tower.vision_model
|
||||
v_tower = model.model.vision_tower
|
||||
v_model = v_tower.vision_model if hasattr(v_tower, "vision_model") else v_tower
|
||||
l_model = model.model.language_model
|
||||
|
||||
config = model.config.text_config
|
||||
dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
|
||||
base = config.rope_local_base_freq
|
||||
local_rope_freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(dtype=torch.float) / dim))
|
||||
inv_freqs, _ = ROPE_INIT_FUNCTIONS[config.rope_scaling["rope_type"]](config)
|
||||
|
||||
if hasattr(l_model, "rotary_emb_local"):
|
||||
_populate_rotary_v4(l_model, config)
|
||||
else:
|
||||
_populate_rotary_v5(l_model, config)
|
||||
|
||||
positions_length = len(v_model.embeddings.position_ids[0])
|
||||
position_ids = torch.arange(positions_length, dtype=torch.long, device="cpu").unsqueeze(0)
|
||||
v_model.embeddings.register_buffer("position_ids", position_ids)
|
||||
embed_scale = torch.tensor(model.config.text_config.hidden_size**0.5, device="cpu")
|
||||
embed_scale = torch.tensor(config.hidden_size**0.5, device="cpu")
|
||||
l_model.embed_tokens.register_buffer("embed_scale", embed_scale)
|
||||
l_model.rotary_emb_local.register_buffer("inv_freq", local_rope_freqs)
|
||||
l_model.rotary_emb.register_buffer("inv_freq", inv_freqs)
|
||||
|
||||
return module
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ Inference pipelines for LTX-2 audio-video generation. Depends on `ltx-core` for
|
||||
| Pipeline | File | Stages | Model | Sampler | Use case |
|
||||
|----------|------|--------|-------|---------|----------|
|
||||
| `TI2VidOneStagePipeline` | `ti2vid_one_stage.py` | 1 | Full | Euler | Simple text/image-to-video |
|
||||
| `T2AOneStagePipeline` | `t2a_one_stage.py` | 1 | Full | Euler | Text-to-audio (audio-only output, no video branch) |
|
||||
| `TI2VidTwoStagesPipeline` | `ti2vid_two_stages.py` | 2 | Full + distilled LoRA | Euler | Production quality |
|
||||
| `TI2VidTwoStagesHQPipeline` | `ti2vid_two_stages_hq.py` | 2 | Full + distilled LoRA (both stages) | Res2s | Highest quality, fewer steps |
|
||||
| `A2VidPipelineTwoStage` | `a2vid_two_stage.py` | 2 | Full + distilled LoRA | Euler | Audio-conditioned video |
|
||||
|
||||
@@ -58,6 +58,7 @@ Available pipeline modules:
|
||||
- `ltx_pipelines.ti2vid_two_stages` - Two-stage text/image-to-video (recommended).
|
||||
- `ltx_pipelines.ti2vid_two_stages_hq` - Two-stage text/image-to-video (different sampler, better quality).
|
||||
- `ltx_pipelines.ti2vid_one_stage` - Single-stage text/image-to-video.
|
||||
- `ltx_pipelines.t2a_one_stage` - Single-stage text-to-audio (audio-only output).
|
||||
- `ltx_pipelines.distilled` - Fast text/image-to-video pipeline using only the distilled model.
|
||||
- `ltx_pipelines.ic_lora` - Video-to-video with IC-LoRA.
|
||||
- `ltx_pipelines.keyframe_interpolation` - Keyframe interpolation.
|
||||
@@ -254,6 +255,20 @@ Uses IC-LoRA on a **distilled** checkpoint with a **single** lip-dub IC-LoRA app
|
||||
|
||||
---
|
||||
|
||||
### 11. T2AOneStagePipeline
|
||||
|
||||
**Best for:** Text-to-audio — generating speech/audio only (no video) from a text prompt, e.g. driving an audio-style LoRA such as an accent LoRA.
|
||||
|
||||
**Source**: [`src/ltx_pipelines/t2a_one_stage.py`](src/ltx_pipelines/t2a_one_stage.py)
|
||||
|
||||
Single-stage, **audio-only** generation: the video branch is absent (`video=None`), so only the audio modality is denoised and decoded through the audio VAE + vocoder, producing a wave file. Audio duration is derived from `--num-frames` / `--frame-rate` (the same `8k+1` frame convention as video). Audio guidance (CFG/STG) is optional — the `--audio-*` flags default to the model's values; the video→audio cross-modal guidance is disabled since there is no video modality.
|
||||
|
||||
**Extra CLI arguments (all optional, with sensible defaults):** `--num-frames`, `--frame-rate`, `--negative-prompt`, `--audio-cfg-guidance-scale`, `--audio-stg-guidance-scale`, `--audio-stg-blocks`, `--audio-rescale-scale`, `--audio-skip-step`. No `--height/--width/--image` (audio has no spatial dimensions).
|
||||
|
||||
**Use when:** You need speech/audio from text alone, or to evaluate an audio-only LoRA (accent, voice style) without generating video.
|
||||
|
||||
---
|
||||
|
||||
## 🎨 Conditioning Types
|
||||
|
||||
Pipelines use different conditioning methods from [`ltx-core`](../ltx-core/) for controlling generation. See the [ltx-core conditioning documentation](../ltx-core/README.md#conditioning--control) for details.
|
||||
@@ -369,7 +384,9 @@ PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True python -m ltx_pipelines.ti2vid_
|
||||
When authoring custom scripts, pass a `QuantizationPolicy` to pipeline classes:
|
||||
|
||||
```python
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.quantization.fp8_cast import build_policy as build_fp8_cast_policy
|
||||
# Alternative:
|
||||
# from ltx_core.quantization.fp8_scaled_mm import build_policy as build_fp8_scaled_mm_policy
|
||||
|
||||
pipeline = TI2VidTwoStagesPipeline(
|
||||
checkpoint_path=ltx_model_path,
|
||||
@@ -377,7 +394,7 @@ pipeline = TI2VidTwoStagesPipeline(
|
||||
spatial_upsampler_path=upsampler_path,
|
||||
gemma_root=gemma_root_path,
|
||||
loras=[],
|
||||
quantization=QuantizationPolicy.fp8_cast(), # or QuantizationPolicy.fp8_scaled_mm()
|
||||
quantization=build_fp8_cast_policy(ltx_model_path),
|
||||
)
|
||||
pipeline(...)
|
||||
```
|
||||
@@ -398,6 +415,69 @@ By default, pipelines clean GPU memory (especially transformer weights) between
|
||||
# utils.cleanup_memory() # Comment out if you have enough VRAM
|
||||
```
|
||||
|
||||
### Compilation (`torch.compile`)
|
||||
|
||||
Compiling the transformer blocks with `torch.compile` speeds up inference. It is **opt-in and off by default**. The blocks are compiled shape-polymorphically (the sequence dimension is marked dynamic), so one compiled artifact serves any token count without recompiling.
|
||||
|
||||
**CLI** — the `--compile` flag maps directly to `CompilationConfig`:
|
||||
|
||||
| Form | Result |
|
||||
| ---- | ------ |
|
||||
| *(flag absent)* | eager, no compilation |
|
||||
| `--compile` | compile with defaults |
|
||||
| `--compile KEY=VALUE ...` | compile, overriding individual fields |
|
||||
|
||||
```bash
|
||||
# Defaults
|
||||
python -m ltx_pipelines.ti2vid_two_stages --compile --checkpoint-path=...
|
||||
|
||||
# reduce-overhead captures CUDA graphs -- the main latency lever for the denoising loop.
|
||||
# Off by default because graph capture reserves static memory pools (extra VRAM), so it
|
||||
# trades memory for speed; enable it when you have headroom.
|
||||
python -m ltx_pipelines.ti2vid_two_stages --compile mode=reduce-overhead --checkpoint-path=...
|
||||
|
||||
# Several overrides at once
|
||||
python -m ltx_pipelines.ti2vid_two_stages \
|
||||
--compile mode=max-autotune fullgraph=true dynamic=true --checkpoint-path=...
|
||||
```
|
||||
|
||||
| Field | Values | Default | Notes |
|
||||
| ----- | ------ | ------- | ----- |
|
||||
| `mode` | `none`, `reduce-overhead`, `max-autotune`, … | `none` | `reduce-overhead`/`max-autotune` enable CUDA graphs |
|
||||
| `backend` | `inductor`, `eager`, … | `inductor` | |
|
||||
| `fullgraph` | `true`/`false` | `false` | |
|
||||
| `dynamic` | `auto`/`true`/`false` | `auto` | the seq dim is marked dynamic regardless |
|
||||
| `inductor_config` | JSON object or path to a `.json` | `{}` | `torch._inductor.config` overrides |
|
||||
| `dynamo_config` | JSON object or path to a `.json` | `{"inline_inbuilt_nn_modules": true, "cache_size_limit": 256}` | `torch._dynamo.config` overrides |
|
||||
|
||||
**Controlling inductor / dynamo configs.** `inductor_config` and `dynamo_config` take either an inline JSON object or a path to a `.json` file, applied via `torch._inductor.config.patch(...)` / `torch._dynamo.config.patch(...)` around the compiled forward. They **replace the defaults wholesale — they do not merge**, so when overriding `dynamo_config` re-include any defaults you want to keep:
|
||||
|
||||
```bash
|
||||
python -m ltx_pipelines.ti2vid_two_stages \
|
||||
--compile 'inductor_config={"max_autotune": true}' \
|
||||
'dynamo_config={"inline_inbuilt_nn_modules": true, "cache_size_limit": 256, "recompile_limit": 32}' \
|
||||
--checkpoint-path=...
|
||||
```
|
||||
|
||||
**Programmatically**, pass a `CompilationConfig` to the pipeline:
|
||||
|
||||
```python
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
|
||||
pipeline = TI2VidTwoStagesPipeline(
|
||||
...,
|
||||
compilation_config=CompilationConfig(mode="reduce-overhead"),
|
||||
)
|
||||
```
|
||||
|
||||
**Faster cache loads: `unsafe_skip_cache_dynamic_shape_guards` (unsafe, opt-in).** Inductor's FX-graph cache re-checks the dynamic-shape guards stored with each entry on every lookup. Setting this flag skips that re-check (every entry is treated as a guard hit), which speeds up warm and cross-process cache loads. It is **not enabled by default** because it is a correctness hazard: a kernel first compiled at a small sequence length keeps int32 address arithmetic, and reusing it at a larger sequence length (roughly **>58k tokens/rank**) overflows int32 and reads out of bounds — surfacing as a CUDA illegal memory access or silently corrupted output. Only enable it when your token counts stay within the range the cached kernels were compiled for:
|
||||
|
||||
```bash
|
||||
python -m ltx_pipelines.ti2vid_two_stages \
|
||||
--compile 'inductor_config={"unsafe_skip_cache_dynamic_shape_guards": true}' \
|
||||
--checkpoint-path=...
|
||||
```
|
||||
|
||||
### Denoising Loop Optimization
|
||||
|
||||
**Gradient Estimation Denoising Loop:**
|
||||
@@ -414,12 +494,13 @@ def denoising_loop(sigmas, video_state, audio_state, stepper):
|
||||
video_state=video_state,
|
||||
audio_state=audio_state,
|
||||
stepper=stepper,
|
||||
denoise_fn=your_denoise_function,
|
||||
transformer=transformer,
|
||||
denoiser=denoiser,
|
||||
ge_gamma=2.0, # Gradient estimation coefficient
|
||||
)
|
||||
```
|
||||
|
||||
This allows you to use **20-30 steps instead of 40** while maintaining quality. The gradient estimation function is available in [`pipeline_utils.py`](src/ltx_pipelines/utils/helpers.py).
|
||||
This allows you to use **20-30 steps instead of 40** while maintaining quality. The gradient estimation function is defined in [`samplers.py`](src/ltx_pipelines/utils/samplers.py).
|
||||
|
||||
---
|
||||
|
||||
@@ -435,15 +516,18 @@ This allows you to use **20-30 steps instead of 40** while maintaining quality.
|
||||
## 📖 Example: Image-to-Video
|
||||
|
||||
```python
|
||||
from ltx_core.loader import LTXV_LORA_COMFY_RENAMING_MAP, LoraPathStrengthAndSDOps
|
||||
from ltx_pipelines.ti2vid_two_stages import TI2VidTwoStagesPipeline
|
||||
from ltx_core.components.guiders import MultiModalGuiderParams
|
||||
from ltx_core.loader import LTXV_LORA_COMFY_RENAMING_MAP, LoraPathStrengthAndSDOps
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_pipelines.ti2vid_two_stages import TI2VidTwoStagesPipeline
|
||||
from ltx_pipelines.utils.args import ImageConditioningInput
|
||||
from ltx_pipelines.utils.media_io import encode_video
|
||||
|
||||
distilled_lora = [
|
||||
LoraPathStrengthAndSDOps(
|
||||
"/path/to/distilled_lora.safetensors",
|
||||
0.6,
|
||||
LTXV_LORA_COMFY_RENAMING_MAP
|
||||
LTXV_LORA_COMFY_RENAMING_MAP,
|
||||
),
|
||||
]
|
||||
|
||||
@@ -473,19 +557,31 @@ audio_guider_params = MultiModalGuiderParams(
|
||||
stg_blocks=[29],
|
||||
)
|
||||
|
||||
# Generate video from image
|
||||
pipeline(
|
||||
# Generate video from image. The pipeline returns (video_iterator, audio);
|
||||
# the caller is responsible for encoding to file via encode_video().
|
||||
num_frames = 121
|
||||
frame_rate = 25.0
|
||||
tiling_config = TilingConfig.default()
|
||||
video, audio = pipeline(
|
||||
prompt="A serene landscape with mountains in the background",
|
||||
output_path="output.mp4",
|
||||
negative_prompt="worst quality, low quality, blurry, distorted",
|
||||
seed=42,
|
||||
height=512,
|
||||
width=768,
|
||||
num_frames=121,
|
||||
frame_rate=25.0,
|
||||
num_frames=num_frames,
|
||||
frame_rate=frame_rate,
|
||||
num_inference_steps=40,
|
||||
video_guider_params=video_guider_params,
|
||||
audio_guider_params=audio_guider_params,
|
||||
images=[ImageConditioningInput("input_image.jpg", 0, 1.0, 33)], # Image at frame 0, strength 1.0, CRF 33
|
||||
images=[ImageConditioningInput("input_image.jpg", 0, 1.0, 33)], # path, frame_idx=0, strength=1.0, crf=33
|
||||
tiling_config=tiling_config,
|
||||
)
|
||||
encode_video(
|
||||
video=video,
|
||||
fps=frame_rate,
|
||||
audio=audio,
|
||||
output_path="output.mp4",
|
||||
video_chunks_number=get_video_chunks_number(num_frames, tiling_config),
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "ltx-pipelines"
|
||||
version = "1.1.3"
|
||||
version = "v1.1.6"
|
||||
description = "Pipelines implementation for Lightricks' LTX-2 model"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.10"
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
LTX-2 Pipelines: High-level video generation pipelines and utilities.
|
||||
This package provides ready-to-use pipelines for video generation:
|
||||
- TI2VidOneStagePipeline: Text/image-to-video in a single stage
|
||||
- T2AOneStagePipeline: Text-to-audio in a single stage (audio-only output)
|
||||
- TI2VidTwoStagesPipeline: Two-stage generation with upsampling
|
||||
- DistilledPipeline: Fast distilled two-stage generation
|
||||
- ICLoraPipeline: Image/video conditioning with distilled LoRA
|
||||
@@ -18,6 +19,7 @@ from ltx_pipelines.ic_lora import ICLoraPipeline
|
||||
from ltx_pipelines.keyframe_interpolation import KeyframeInterpolationPipeline
|
||||
from ltx_pipelines.lipdub import LipDubPipeline
|
||||
from ltx_pipelines.retake import RetakePipeline
|
||||
from ltx_pipelines.t2a_one_stage import T2AOneStagePipeline
|
||||
from ltx_pipelines.ti2vid_one_stage import TI2VidOneStagePipeline
|
||||
from ltx_pipelines.ti2vid_two_stages import TI2VidTwoStagesPipeline
|
||||
|
||||
@@ -28,6 +30,7 @@ __all__ = [
|
||||
"KeyframeInterpolationPipeline",
|
||||
"LipDubPipeline",
|
||||
"RetakePipeline",
|
||||
"T2AOneStagePipeline",
|
||||
"TI2VidOneStagePipeline",
|
||||
"TI2VidTwoStagesPipeline",
|
||||
]
|
||||
|
||||
@@ -9,6 +9,7 @@ from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio, AudioLatentShape, VideoPixelShape
|
||||
@@ -52,7 +53,7 @@ class A2VidPipelineTwoStage:
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.device = device or get_device()
|
||||
@@ -71,7 +72,7 @@ class A2VidPipelineTwoStage:
|
||||
loras=tuple(loras),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
stage_2_loras = (*tuple(loras), *tuple(distilled_lora))
|
||||
@@ -82,7 +83,7 @@ class A2VidPipelineTwoStage:
|
||||
loras=stage_2_loras,
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.upsampler = VideoUpsampler(
|
||||
@@ -238,7 +239,7 @@ class A2VidPipelineTwoStage:
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
parser = default_2_stage_arg_parser()
|
||||
parser.add_argument(
|
||||
"--audio-path",
|
||||
@@ -266,7 +267,7 @@ def main() -> None:
|
||||
gemma_root=args.gemma_root,
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
|
||||
@@ -6,6 +6,7 @@ import torch
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio
|
||||
@@ -53,7 +54,7 @@ class DistilledPipeline:
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.device = device or get_device()
|
||||
@@ -75,7 +76,7 @@ class DistilledPipeline:
|
||||
loras=tuple(loras),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.upsampler = VideoUpsampler(
|
||||
@@ -180,7 +181,7 @@ class DistilledPipeline:
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
checkpoint_path = detect_checkpoint_path(distilled=True)
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_distilled_arg_parser(params=params)
|
||||
@@ -191,7 +192,7 @@ def main() -> None:
|
||||
gemma_root=args.gemma_root,
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
|
||||
@@ -59,6 +59,7 @@ from ltx_pipelines.utils.constants import DISTILLED_SIGMA_VALUES, STAGE_2_DISTIL
|
||||
from ltx_pipelines.utils.denoisers import SimpleDenoiser
|
||||
from ltx_pipelines.utils.helpers import get_device, modality_from_latent_state
|
||||
from ltx_pipelines.utils.media_io import ResizeMode, align_resolution, load_video_conditioning_hdr
|
||||
from ltx_pipelines.utils.quantization_factory import QuantizationKind
|
||||
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -77,7 +78,7 @@ ALIGNMENT_DIVISOR = 64
|
||||
# to the pipeline constructor.
|
||||
TILED_VAE_ENCODE_PIXEL_THRESHOLD = 512 * 768
|
||||
|
||||
_DEFAULT_QUANTIZATION = QuantizationPolicy.fp8_cast()
|
||||
_DEFAULT_QUANTIZATION = QuantizationKind.FP8_CAST
|
||||
|
||||
# Default stage-2 configuration: one refinement phase with modest 2-way tiling
|
||||
# in every dimension and a short 2-step distilled sigma schedule.
|
||||
@@ -205,7 +206,7 @@ class HDRICLoraPipeline:
|
||||
hdr_lora: str | Path,
|
||||
text_embeddings_path: str | Path,
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy = _DEFAULT_QUANTIZATION,
|
||||
quantization: QuantizationPolicy | QuantizationKind | None = _DEFAULT_QUANTIZATION,
|
||||
registry: Registry | None = None,
|
||||
hdr_lora_config: HdrLoraConfig | None = None,
|
||||
tiled_vae_encode_pixel_threshold: int = TILED_VAE_ENCODE_PIXEL_THRESHOLD,
|
||||
@@ -232,6 +233,8 @@ class HDRICLoraPipeline:
|
||||
"""
|
||||
self.device = device or get_device()
|
||||
self._tiled_vae_encode_threshold = tiled_vae_encode_pixel_threshold
|
||||
if isinstance(quantization, QuantizationKind):
|
||||
quantization = quantization.to_policy(checkpoint_path=distilled_checkpoint_path)
|
||||
if offload_mode != OffloadMode.NONE and quantization is not None:
|
||||
logger.info("Offload mode enabled — disabling quantization (not supported with layer streaming).")
|
||||
quantization = None
|
||||
|
||||
@@ -7,12 +7,14 @@ from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.conditioning import ConditioningItem
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio, VideoPixelShape
|
||||
from ltx_pipelines.iclora_utils import (
|
||||
append_ic_lora_reference_video_conditionings,
|
||||
read_lora_reference_downscale_factor,
|
||||
read_lora_reference_temporal_scale_factor,
|
||||
)
|
||||
from ltx_pipelines.utils.args import (
|
||||
ImageConditioningInput,
|
||||
@@ -60,7 +62,7 @@ class ICLoraPipeline:
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.device = device or get_device()
|
||||
@@ -82,7 +84,7 @@ class ICLoraPipeline:
|
||||
loras=tuple(loras),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.stage_2 = DiffusionStage(
|
||||
@@ -92,7 +94,7 @@ class ICLoraPipeline:
|
||||
loras=(),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.upsampler = VideoUpsampler(
|
||||
@@ -101,10 +103,11 @@ class ICLoraPipeline:
|
||||
self.video_decoder = VideoDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
|
||||
self.audio_decoder = AudioDecoder(distilled_checkpoint_path, self.dtype, self.device, registry=registry)
|
||||
|
||||
# Read reference downscale factor from LoRA metadata.
|
||||
# IC-LoRAs trained with low-resolution reference videos store this factor
|
||||
# so inference can resize reference videos to match training conditions.
|
||||
# Read reference scale factors from LoRA metadata.
|
||||
# IC-LoRAs trained with scaled reference videos store these factors
|
||||
# so inference can resize/subsample reference videos to match training conditions.
|
||||
self.reference_downscale_factor = 1
|
||||
self.reference_temporal_scale_factor = 1
|
||||
for lora in loras:
|
||||
scale = read_lora_reference_downscale_factor(lora.path)
|
||||
if scale != 1:
|
||||
@@ -115,6 +118,15 @@ class ICLoraPipeline:
|
||||
f"specifies {scale}. Cannot combine LoRAs with different reference scales."
|
||||
)
|
||||
self.reference_downscale_factor = scale
|
||||
temporal = read_lora_reference_temporal_scale_factor(lora.path)
|
||||
if temporal != 1:
|
||||
if self.reference_temporal_scale_factor not in (1, temporal):
|
||||
raise ValueError(
|
||||
f"Conflicting reference_temporal_scale_factor values in LoRAs: "
|
||||
f"already have {self.reference_temporal_scale_factor}, but {lora.path} "
|
||||
f"specifies {temporal}. Cannot combine LoRAs with different temporal scales."
|
||||
)
|
||||
self.reference_temporal_scale_factor = temporal
|
||||
|
||||
def __call__( # noqa: PLR0913
|
||||
self,
|
||||
@@ -317,6 +329,7 @@ class ICLoraPipeline:
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
reference_downscale_factor=self.reference_downscale_factor,
|
||||
reference_temporal_scale_factor=self.reference_temporal_scale_factor,
|
||||
conditioning_attention_strength=conditioning_attention_strength,
|
||||
conditioning_attention_mask=conditioning_attention_mask,
|
||||
tiling_config=None,
|
||||
@@ -330,7 +343,7 @@ class ICLoraPipeline:
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
checkpoint_path = detect_checkpoint_path(distilled=True)
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_distilled_arg_parser(params=params)
|
||||
@@ -386,7 +399,7 @@ def main() -> None:
|
||||
gemma_root=args.gemma_root,
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
|
||||
@@ -31,6 +31,17 @@ def read_lora_reference_downscale_factor(lora_path: str) -> int:
|
||||
return 1
|
||||
|
||||
|
||||
def read_lora_reference_temporal_scale_factor(lora_path: str) -> int:
|
||||
"""Read ``reference_temporal_scale_factor`` from LoRA safetensors metadata (default 1)."""
|
||||
try:
|
||||
with safe_open(lora_path, framework="pt") as f:
|
||||
metadata = f.metadata() or {}
|
||||
return int(metadata.get("reference_temporal_scale_factor", 1))
|
||||
except Exception as e:
|
||||
logging.warning("Failed to read metadata from LoRA file '%s': %s", lora_path, e)
|
||||
return 1
|
||||
|
||||
|
||||
def downsample_mask_video_to_latent(
|
||||
mask: torch.Tensor,
|
||||
target_latent_shape: VideoLatentShape,
|
||||
@@ -66,6 +77,12 @@ def downsample_mask_video_to_latent(
|
||||
return rearrange(latent_mask, "b 1 f h w -> b (f h w)")
|
||||
|
||||
|
||||
def temporal_subsample(video: torch.Tensor, temporal_scale_factor: int) -> torch.Tensor:
|
||||
"""VAE-aligned temporal subsampling: keep frame 0, then every Nth frame."""
|
||||
indices = [0, *list(range(1, video.shape[2], temporal_scale_factor))]
|
||||
return video[:, :, indices]
|
||||
|
||||
|
||||
def append_ic_lora_reference_video_conditionings( # noqa: PLR0913
|
||||
conditionings: list[ConditioningItem],
|
||||
video_conditioning: list[tuple[str, float]],
|
||||
@@ -77,6 +94,7 @@ def append_ic_lora_reference_video_conditionings( # noqa: PLR0913
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
reference_downscale_factor: int,
|
||||
reference_temporal_scale_factor: int = 1,
|
||||
conditioning_attention_strength: float,
|
||||
conditioning_attention_mask: torch.Tensor | None,
|
||||
tiling_config: TilingConfig | None = None,
|
||||
@@ -93,6 +111,8 @@ def append_ic_lora_reference_video_conditionings( # noqa: PLR0913
|
||||
for video_path, strength in video_conditioning:
|
||||
frame_gen = decode_video_by_frame(path=video_path, frame_cap=num_frames, device=device)
|
||||
video = video_preprocess(frame_gen, ref_height, ref_width, dtype, device)
|
||||
if reference_temporal_scale_factor > 1:
|
||||
video = temporal_subsample(video, reference_temporal_scale_factor)
|
||||
if tiling_config is not None:
|
||||
encoded_video = video_encoder.tiled_encode(video, tiling_config)
|
||||
else:
|
||||
@@ -113,6 +133,7 @@ def append_ic_lora_reference_video_conditionings( # noqa: PLR0913
|
||||
cond = VideoConditionByReferenceLatent(
|
||||
latent=encoded_video,
|
||||
downscale_factor=scale,
|
||||
temporal_scale_factor=reference_temporal_scale_factor,
|
||||
strength=strength,
|
||||
)
|
||||
if attn_mask is not None:
|
||||
|
||||
@@ -12,6 +12,7 @@ from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio, VideoPixelShape
|
||||
@@ -62,7 +63,7 @@ class KeyframeInterpolationPipeline:
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.device = device or get_device()
|
||||
@@ -80,7 +81,7 @@ class KeyframeInterpolationPipeline:
|
||||
loras=tuple(loras),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
stage_2_loras = (*tuple(loras), *tuple(distilled_lora))
|
||||
@@ -91,7 +92,7 @@ class KeyframeInterpolationPipeline:
|
||||
loras=stage_2_loras,
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.upsampler = VideoUpsampler(
|
||||
@@ -233,7 +234,7 @@ class KeyframeInterpolationPipeline:
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
checkpoint_path = detect_checkpoint_path()
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_arg_parser(params=params)
|
||||
@@ -245,7 +246,7 @@ def main() -> None:
|
||||
gemma_root=args.gemma_root,
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
|
||||
@@ -13,6 +13,7 @@ from ltx_core.conditioning import AudioConditionByReferenceLatent
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.audio_vae import encode_audio as vae_encode_audio
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae import TilingConfig, VideoEncoder, get_video_chunks_number
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio, AudioLatentShape, SpatioTemporalScaleFactors, VideoPixelShape
|
||||
@@ -59,7 +60,7 @@ class LipDubPipeline:
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
) -> None:
|
||||
self.device = device or get_device()
|
||||
@@ -89,7 +90,7 @@ class LipDubPipeline:
|
||||
loras=loras,
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.upsampler = VideoUpsampler(
|
||||
@@ -289,7 +290,7 @@ def patchify_lipdub_audio_reference_latent(
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
checkpoint_path = detect_checkpoint_path(distilled=True)
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = lipdub_arg_parser(params=params)
|
||||
@@ -304,7 +305,7 @@ def main() -> None:
|
||||
gemma_root=args.gemma_root,
|
||||
ic_lora=args.lora[0],
|
||||
quantization=args.quantization,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
|
||||
@@ -11,6 +11,7 @@ from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.conditioning.types.noise_mask_cond import TemporalRegionMask
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import (
|
||||
@@ -73,7 +74,7 @@ class RetakePipeline:
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
distilled: bool = True,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.device = device or get_device()
|
||||
@@ -108,7 +109,7 @@ class RetakePipeline:
|
||||
loras=tuple(loras),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.video_decoder = VideoDecoder(
|
||||
@@ -283,7 +284,7 @@ class RetakePipeline:
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
"""CLI entry point for retake (regenerate a time region)."""
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
parser = video_editing_arg_parser(distilled=True)
|
||||
parser.description = "Retake: regenerate a time region of a video with LTX-2."
|
||||
args = parser.parse_args()
|
||||
@@ -308,7 +309,7 @@ def main() -> None:
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
distilled=True,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
params = detect_params(args.distilled_checkpoint_path)
|
||||
|
||||
@@ -0,0 +1,192 @@
|
||||
import logging
|
||||
|
||||
import torch
|
||||
|
||||
from ltx_core.components.guiders import (
|
||||
MultiModalGuiderFactory,
|
||||
MultiModalGuiderParams,
|
||||
create_multimodal_guider_factory,
|
||||
)
|
||||
from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.transformer import LTXV_AUDIO_ONLY_MODEL_COMFY_RENAMING_MAP, LTXAudioOnlyModelConfigurator
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio
|
||||
from ltx_pipelines.utils import get_device
|
||||
from ltx_pipelines.utils.args import (
|
||||
default_1_stage_t2a_arg_parser,
|
||||
detect_checkpoint_path,
|
||||
)
|
||||
from ltx_pipelines.utils.blocks import (
|
||||
AudioDecoder,
|
||||
DiffusionStage,
|
||||
PromptEncoder,
|
||||
)
|
||||
from ltx_pipelines.utils.constants import detect_params
|
||||
from ltx_pipelines.utils.denoisers import FactoryGuidedDenoiser
|
||||
from ltx_pipelines.utils.media_io import encode_audio
|
||||
from ltx_pipelines.utils.types import ModalitySpec, OffloadMode
|
||||
|
||||
# Placeholder pixel dimensions used for ``VideoPixelShape`` construction.
|
||||
# Audio-only generation reads ``frames`` and ``fps`` from the pixel shape via
|
||||
# ``AudioLatentShape.from_video_pixel_shape`` (height/width are unused).
|
||||
_AUDIO_ONLY_PLACEHOLDER_RES = 512
|
||||
|
||||
|
||||
class T2AOneStagePipeline:
|
||||
"""
|
||||
Single-stage text-to-audio generation pipeline.
|
||||
Generates audio at the target duration in a single diffusion pass with
|
||||
classifier-free guidance (CFG) on the audio modality only. The video
|
||||
modality is fully absent — the transformer runs audio-only by passing
|
||||
``video=None`` to the ``DiffusionStage``.
|
||||
Assumes full non distilled model is provided in the checkpoint_path.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
checkpoint_path: str,
|
||||
gemma_root: str,
|
||||
loras: list[LoraPathStrengthAndSDOps],
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.dtype = torch.bfloat16
|
||||
self.device = device or get_device()
|
||||
self._scheduler = LTX2Scheduler()
|
||||
self.prompt_encoder = PromptEncoder(
|
||||
checkpoint_path=checkpoint_path,
|
||||
gemma_root=gemma_root,
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
registry=registry,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
# Audio-only: build an audio-only transformer (model_configurator) so the video
|
||||
# weights are never instantiated, plus a use-case-specific SDOps that restricts
|
||||
# checkpoint reads to the audio model's keys, so the video weights are never even
|
||||
# read from disk (the loader skips any key the SDOps maps to None).
|
||||
self.stage = DiffusionStage(
|
||||
checkpoint_path=checkpoint_path,
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
loras=tuple(loras),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
model_configurator=LTXAudioOnlyModelConfigurator,
|
||||
model_sd_ops=LTXV_AUDIO_ONLY_MODEL_COMFY_RENAMING_MAP,
|
||||
)
|
||||
self.audio_decoder = AudioDecoder(
|
||||
checkpoint_path=checkpoint_path,
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
registry=registry,
|
||||
)
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
prompt: str,
|
||||
negative_prompt: str,
|
||||
seed: int,
|
||||
num_frames: int,
|
||||
frame_rate: float,
|
||||
num_inference_steps: int,
|
||||
audio_guider_params: MultiModalGuiderParams | MultiModalGuiderFactory,
|
||||
enhance_prompt: bool = False,
|
||||
max_batch_size: int = 1,
|
||||
sigmas: torch.Tensor | None = None,
|
||||
) -> Audio:
|
||||
generator = torch.Generator(device=self.device).manual_seed(seed)
|
||||
noiser = GaussianNoiser(generator=generator)
|
||||
|
||||
ctx_p, ctx_n = self.prompt_encoder(
|
||||
[prompt, negative_prompt],
|
||||
enhance_first_prompt=enhance_prompt,
|
||||
enhance_prompt_image=None,
|
||||
enhance_prompt_seed=seed,
|
||||
)
|
||||
a_context_p = ctx_p.audio_encoding
|
||||
a_context_n = ctx_n.audio_encoding
|
||||
|
||||
sigmas = (sigmas if sigmas is not None else self._scheduler.execute(steps=num_inference_steps)).to(
|
||||
dtype=torch.float32, device=self.device
|
||||
)
|
||||
|
||||
# Normalize to a guider factory. Plain ``MultiModalGuiderParams`` (the default /
|
||||
# CLI case) becomes a simple sigma-independent guider, but callers may also pass
|
||||
# their own factory for sigma-dependent guidance; ``FactoryGuidedDenoiser`` always
|
||||
# consumes a factory.
|
||||
audio_guider_factory = create_multimodal_guider_factory(
|
||||
params=audio_guider_params,
|
||||
negative_context=a_context_n,
|
||||
)
|
||||
|
||||
_, audio_state = self.stage(
|
||||
denoiser=FactoryGuidedDenoiser(
|
||||
v_context=None,
|
||||
a_context=a_context_p,
|
||||
video_guider_factory=None,
|
||||
audio_guider_factory=audio_guider_factory,
|
||||
),
|
||||
sigmas=sigmas,
|
||||
noiser=noiser,
|
||||
width=_AUDIO_ONLY_PLACEHOLDER_RES,
|
||||
height=_AUDIO_ONLY_PLACEHOLDER_RES,
|
||||
frames=num_frames,
|
||||
fps=frame_rate,
|
||||
video=None,
|
||||
audio=ModalitySpec(context=a_context_p),
|
||||
max_batch_size=max_batch_size,
|
||||
)
|
||||
|
||||
return self.audio_decoder(audio_state.latent)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
checkpoint_path = detect_checkpoint_path()
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_1_stage_t2a_arg_parser(params=params)
|
||||
args = parser.parse_args()
|
||||
pipeline = T2AOneStagePipeline(
|
||||
checkpoint_path=args.checkpoint_path,
|
||||
gemma_root=args.gemma_root,
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
audio = pipeline(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
seed=args.seed,
|
||||
num_frames=args.num_frames,
|
||||
frame_rate=args.frame_rate,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
audio_guider_params=MultiModalGuiderParams(
|
||||
cfg_scale=args.audio_cfg_guidance_scale,
|
||||
stg_scale=args.audio_stg_guidance_scale,
|
||||
rescale_scale=args.audio_rescale_scale,
|
||||
# Audio-only generation has no video modality, so the video->audio
|
||||
# (v2a) cross-modal guidance is meaningless here. 1.0 disables it.
|
||||
modality_scale=1.0,
|
||||
skip_step=args.audio_skip_step,
|
||||
stg_blocks=args.audio_stg_blocks,
|
||||
),
|
||||
max_batch_size=args.max_batch_size,
|
||||
)
|
||||
|
||||
encode_audio(audio=audio, output_path=args.output_path)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -12,6 +12,7 @@ from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae.tiling import TilingConfig
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio
|
||||
@@ -55,7 +56,7 @@ class TI2VidOneStagePipeline:
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.dtype = torch.bfloat16
|
||||
@@ -82,7 +83,7 @@ class TI2VidOneStagePipeline:
|
||||
loras=tuple(loras),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.video_decoder = VideoDecoder(
|
||||
@@ -185,7 +186,7 @@ class TI2VidOneStagePipeline:
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
checkpoint_path = detect_checkpoint_path()
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_1_stage_arg_parser(params=params)
|
||||
@@ -195,7 +196,7 @@ def main() -> None:
|
||||
gemma_root=args.gemma_root,
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
video, audio = pipeline(
|
||||
|
||||
@@ -12,6 +12,7 @@ from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio, VideoPixelShape
|
||||
@@ -61,7 +62,7 @@ class TI2VidTwoStagesPipeline:
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.device = device or get_device()
|
||||
@@ -85,7 +86,7 @@ class TI2VidTwoStagesPipeline:
|
||||
loras=tuple(loras),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.stage_2 = DiffusionStage(
|
||||
@@ -95,7 +96,7 @@ class TI2VidTwoStagesPipeline:
|
||||
loras=(*tuple(loras), *distilled_lora),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
|
||||
@@ -223,7 +224,7 @@ class TI2VidTwoStagesPipeline:
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
checkpoint_path = detect_checkpoint_path()
|
||||
params = detect_params(checkpoint_path)
|
||||
parser = default_2_stage_arg_parser(params=params)
|
||||
@@ -235,7 +236,7 @@ def main() -> None:
|
||||
gemma_root=args.gemma_root,
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
|
||||
@@ -9,6 +9,7 @@ from ltx_core.components.noisers import GaussianNoiser
|
||||
from ltx_core.components.schedulers import LTX2Scheduler
|
||||
from ltx_core.loader import LoraPathStrengthAndSDOps
|
||||
from ltx_core.loader.registry import Registry
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.model.video_vae import TilingConfig, get_video_chunks_number
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.types import Audio, VideoLatentShape, VideoPixelShape
|
||||
@@ -60,7 +61,7 @@ class TI2VidTwoStagesHQPipeline:
|
||||
device: torch.device | None = None,
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
):
|
||||
self.device = device or get_device()
|
||||
@@ -95,7 +96,7 @@ class TI2VidTwoStagesHQPipeline:
|
||||
loras=(*loras, distilled_lora_stage_1),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
self.stage_2 = DiffusionStage(
|
||||
@@ -105,7 +106,7 @@ class TI2VidTwoStagesHQPipeline:
|
||||
loras=(*loras, distilled_lora_stage_2),
|
||||
quantization=quantization,
|
||||
registry=registry,
|
||||
torch_compile=torch_compile,
|
||||
compilation_config=compilation_config,
|
||||
offload_mode=offload_mode,
|
||||
)
|
||||
|
||||
@@ -242,7 +243,7 @@ class TI2VidTwoStagesHQPipeline:
|
||||
|
||||
@torch.inference_mode()
|
||||
def main() -> None:
|
||||
logging.getLogger().setLevel(logging.INFO)
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
parser = hq_2_stage_arg_parser(params=LTX_2_3_HQ_PARAMS)
|
||||
args = parser.parse_args()
|
||||
pipeline = TI2VidTwoStagesHQPipeline(
|
||||
@@ -254,7 +255,7 @@ def main() -> None:
|
||||
gemma_root=args.gemma_root,
|
||||
loras=tuple(args.lora) if args.lora else (),
|
||||
quantization=args.quantization,
|
||||
torch_compile=args.compile,
|
||||
compilation_config=args.compile,
|
||||
offload_mode=args.offload_mode,
|
||||
)
|
||||
tiling_config = TilingConfig.default()
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
import argparse
|
||||
import json
|
||||
from collections.abc import Sequence
|
||||
from pathlib import Path
|
||||
from typing import NamedTuple
|
||||
from typing import Any, NamedTuple
|
||||
|
||||
from ltx_core.loader import LTXV_LORA_COMFY_RENAMING_MAP, LoraPathStrengthAndSDOps
|
||||
from ltx_core.model.transformer.compiling import CompilationConfig
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_pipelines.utils.constants import (
|
||||
DEFAULT_IMAGE_CRF,
|
||||
@@ -13,6 +15,7 @@ from ltx_pipelines.utils.constants import (
|
||||
LTX_2_3_PARAMS,
|
||||
PipelineParams,
|
||||
)
|
||||
from ltx_pipelines.utils.quantization_factory import QuantizationKind
|
||||
from ltx_pipelines.utils.types import OffloadMode
|
||||
|
||||
|
||||
@@ -32,7 +35,7 @@ class VideoConditioningAction(argparse.Action):
|
||||
option_string: str | None = None, # noqa: ARG002
|
||||
) -> None:
|
||||
path, strength_str = values
|
||||
resolved_path = resolve_path(path)
|
||||
resolved_path = resolve_existing_path(path)
|
||||
strength = float(strength_str)
|
||||
current = getattr(namespace, self.dest) or []
|
||||
current.append((resolved_path, strength))
|
||||
@@ -58,7 +61,7 @@ class VideoMaskConditioningAction(argparse.Action):
|
||||
msg = f"{option_string} requires exactly 2 arguments (MASK_PATH STRENGTH), got {len(values)}"
|
||||
raise argparse.ArgumentError(self, msg)
|
||||
|
||||
mask_path = resolve_path(values[0])
|
||||
mask_path = resolve_existing_path(values[0])
|
||||
strength = float(values[1])
|
||||
setattr(namespace, self.dest, (mask_path, strength))
|
||||
|
||||
@@ -76,7 +79,7 @@ class ImageAction(argparse.Action):
|
||||
raise argparse.ArgumentError(self, msg)
|
||||
|
||||
conditioning = ImageConditioningInput(
|
||||
path=resolve_path(values[0]),
|
||||
path=resolve_existing_path(values[0]),
|
||||
frame_idx=int(values[1]),
|
||||
strength=float(values[2]),
|
||||
crf=int(values[3]) if len(values) > 3 else DEFAULT_IMAGE_CRF,
|
||||
@@ -101,7 +104,7 @@ class LoraAction(argparse.Action):
|
||||
path = values[0]
|
||||
strength_str = values[1] if len(values) > 1 else str(DEFAULT_LORA_STRENGTH)
|
||||
|
||||
resolved_path = resolve_path(path)
|
||||
resolved_path = resolve_existing_path(path)
|
||||
strength = float(strength_str)
|
||||
|
||||
current = getattr(namespace, self.dest) or []
|
||||
@@ -109,11 +112,118 @@ class LoraAction(argparse.Action):
|
||||
setattr(namespace, self.dest, current)
|
||||
|
||||
|
||||
class CompileAction(argparse.Action):
|
||||
"""Parse ``--compile [KEY=VALUE ...]`` into a :class:`CompilationConfig`.
|
||||
The flag is absent -> ``args.compile`` stays at its default (``None``).
|
||||
The flag is passed alone -> ``CompilationConfig()`` (vanilla torch defaults).
|
||||
The flag is passed with args -> ``CompilationConfig`` with the given fields overridden.
|
||||
Errors (unknown key, malformed value, duplicate key, empty value) raise
|
||||
:class:`argparse.ArgumentError` so argparse formats them as friendly CLI
|
||||
messages rather than uncaught tracebacks.
|
||||
"""
|
||||
|
||||
_ALLOWED_KEYS = frozenset({"mode", "backend", "fullgraph", "dynamic", "inductor_config", "dynamo_config"})
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
parser: argparse.ArgumentParser, # noqa: ARG002
|
||||
namespace: argparse.Namespace,
|
||||
values: list[str],
|
||||
option_string: str | None = None, # noqa: ARG002
|
||||
) -> None:
|
||||
overrides: dict[str, object] = {}
|
||||
for item in values:
|
||||
if "=" not in item:
|
||||
raise argparse.ArgumentError(self, f"expects KEY=VALUE pairs, got: {item!r}")
|
||||
key, _, raw = item.partition("=")
|
||||
key = key.strip()
|
||||
if key not in self._ALLOWED_KEYS:
|
||||
raise argparse.ArgumentError(
|
||||
self,
|
||||
f"{key!r} is not a CompilationConfig field; valid keys: {sorted(self._ALLOWED_KEYS)}",
|
||||
)
|
||||
if key in overrides:
|
||||
raise argparse.ArgumentError(self, f"{key} given more than once")
|
||||
if key == "mode":
|
||||
overrides[key] = self._parse_mode(raw)
|
||||
elif key == "backend":
|
||||
overrides[key] = self._parse_non_empty(key, raw)
|
||||
elif key == "fullgraph":
|
||||
overrides[key] = self._parse_bool(key, raw)
|
||||
elif key == "dynamic":
|
||||
overrides[key] = self._parse_dynamic(raw)
|
||||
elif key in ("inductor_config", "dynamo_config"):
|
||||
overrides[key] = self._parse_json_dict(key, raw)
|
||||
setattr(namespace, self.dest, CompilationConfig(**overrides))
|
||||
|
||||
def _parse_mode(self, raw: str) -> str | None:
|
||||
stripped = raw.strip()
|
||||
if not stripped:
|
||||
raise argparse.ArgumentError(self, "mode=... value cannot be empty (use mode=none to clear)")
|
||||
if stripped.lower() == "none":
|
||||
return None
|
||||
return stripped
|
||||
|
||||
def _parse_non_empty(self, key: str, raw: str) -> str:
|
||||
stripped = raw.strip()
|
||||
if not stripped:
|
||||
raise argparse.ArgumentError(self, f"{key}=... value cannot be empty")
|
||||
return stripped
|
||||
|
||||
def _parse_bool(self, key: str, raw: str) -> bool:
|
||||
normalized = raw.strip().lower()
|
||||
if normalized in ("true", "1"):
|
||||
return True
|
||||
if normalized in ("false", "0"):
|
||||
return False
|
||||
raise argparse.ArgumentError(self, f"{key}=... must be true or false; got {raw!r}")
|
||||
|
||||
def _parse_dynamic(self, raw: str) -> bool | None:
|
||||
normalized = raw.strip().lower()
|
||||
if normalized in ("auto", "none"):
|
||||
return None
|
||||
if normalized in ("true", "1"):
|
||||
return True
|
||||
if normalized in ("false", "0"):
|
||||
return False
|
||||
raise argparse.ArgumentError(self, f"dynamic=... must be auto/true/false; got {raw!r}")
|
||||
|
||||
def _parse_json_dict(self, key: str, raw: str) -> dict[str, Any]:
|
||||
# Inline JSON object starts with '{'; otherwise treat the value as a path to a JSON file.
|
||||
stripped = raw.strip()
|
||||
if not stripped:
|
||||
raise argparse.ArgumentError(self, f"{key}=... value cannot be empty")
|
||||
if stripped.startswith("{"):
|
||||
source = stripped
|
||||
else:
|
||||
path = Path(stripped).expanduser()
|
||||
if not path.is_file():
|
||||
raise argparse.ArgumentError(
|
||||
self, f"{key}=... must be a JSON object or a path to a JSON file; got {raw!r}"
|
||||
)
|
||||
source = path.read_text()
|
||||
try:
|
||||
value = json.loads(source)
|
||||
except json.JSONDecodeError as e:
|
||||
raise argparse.ArgumentError(self, f"{key}=... must be a JSON object; got {raw!r} ({e.msg})") from None
|
||||
if not isinstance(value, dict):
|
||||
raise argparse.ArgumentError(self, f"{key}=... must decode to a JSON object; got {type(value).__name__}")
|
||||
return value
|
||||
|
||||
|
||||
def resolve_path(path: str) -> str:
|
||||
return str(Path(path).expanduser().resolve().as_posix())
|
||||
|
||||
|
||||
QUANTIZATION_POLICIES = ("fp8-cast", "fp8-scaled-mm")
|
||||
def resolve_existing_path(path: str) -> str:
|
||||
"""Resolve *path* and verify it exists."""
|
||||
resolved = resolve_path(path)
|
||||
if not Path(resolved).exists():
|
||||
raise argparse.ArgumentError(None, f"Path not found: {resolved}")
|
||||
return resolved
|
||||
|
||||
|
||||
QUANTIZATION_POLICIES = tuple(k.value for k in QuantizationKind)
|
||||
|
||||
|
||||
def _resolve_quantization(namespace: argparse.Namespace) -> None:
|
||||
@@ -123,16 +233,14 @@ def _resolve_quantization(namespace: argparse.Namespace) -> None:
|
||||
name = getattr(namespace, "quantization", None)
|
||||
if name is None or isinstance(name, QuantizationPolicy):
|
||||
return
|
||||
if name == "fp8-cast":
|
||||
namespace.quantization = QuantizationPolicy.fp8_cast()
|
||||
try:
|
||||
kind = QuantizationKind(name)
|
||||
except ValueError:
|
||||
return
|
||||
if name == "fp8-scaled-mm":
|
||||
ckpt = getattr(namespace, "checkpoint_path", None) or getattr(namespace, "distilled_checkpoint_path", None)
|
||||
if ckpt is None:
|
||||
raise SystemExit(
|
||||
"--quantization fp8-scaled-mm requires --checkpoint-path (or --distilled-checkpoint-path)."
|
||||
)
|
||||
namespace.quantization = QuantizationPolicy.fp8_scaled_mm(ckpt)
|
||||
ckpt = getattr(namespace, "checkpoint_path", None) or getattr(namespace, "distilled_checkpoint_path", None)
|
||||
if ckpt is None:
|
||||
raise SystemExit(f"--quantization {kind.value} requires --checkpoint-path (or --distilled-checkpoint-path).")
|
||||
namespace.quantization = kind.to_policy(checkpoint_path=ckpt)
|
||||
|
||||
|
||||
class _PipelineArgumentParser(argparse.ArgumentParser):
|
||||
@@ -150,7 +258,7 @@ def detect_checkpoint_path(distilled: bool = False) -> str:
|
||||
"""Pre-parse argv to extract the checkpoint path before building the full parser."""
|
||||
pre = argparse.ArgumentParser(add_help=False)
|
||||
flag = "--distilled-checkpoint-path" if distilled else "--checkpoint-path"
|
||||
pre.add_argument(flag, type=resolve_path, required=True)
|
||||
pre.add_argument(flag, type=resolve_existing_path, required=True)
|
||||
known, _ = pre.parse_known_args()
|
||||
return known.distilled_checkpoint_path if distilled else known.checkpoint_path
|
||||
|
||||
@@ -163,14 +271,14 @@ def basic_arg_parser(
|
||||
if distilled:
|
||||
parser.add_argument(
|
||||
"--distilled-checkpoint-path",
|
||||
type=resolve_path,
|
||||
type=resolve_existing_path,
|
||||
required=True,
|
||||
help="Path to LTX-2 distilled model checkpoint (.safetensors file).",
|
||||
)
|
||||
else:
|
||||
parser.add_argument(
|
||||
"--checkpoint-path",
|
||||
type=resolve_path,
|
||||
type=resolve_existing_path,
|
||||
required=True,
|
||||
help="Path to LTX-2 model checkpoint (.safetensors file).",
|
||||
)
|
||||
@@ -185,7 +293,7 @@ def basic_arg_parser(
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gemma-root",
|
||||
type=resolve_path,
|
||||
type=resolve_existing_path,
|
||||
required=True,
|
||||
help="Path to the root directory containing the Gemma text encoder model files.",
|
||||
)
|
||||
@@ -276,8 +384,20 @@ def basic_arg_parser(
|
||||
)
|
||||
parser.add_argument(
|
||||
"--compile",
|
||||
action="store_true",
|
||||
help="Enable torch.compile for transformer blocks to optimize performance.",
|
||||
nargs="*",
|
||||
action=CompileAction,
|
||||
default=None,
|
||||
metavar="KEY=VALUE",
|
||||
help=(
|
||||
"Enable torch.compile for transformer blocks. Pass alone for defaults, "
|
||||
"or with KEY=VALUE overrides for any CompilationConfig field. "
|
||||
"Keys: mode, backend, fullgraph, dynamic, inductor_config, dynamo_config. "
|
||||
"inductor_config/dynamo_config take JSON objects (inline or a path to a .json file) "
|
||||
"that fully replace the defaults. "
|
||||
"Examples: --compile or --compile mode=reduce-overhead or "
|
||||
"--compile mode=reduce-overhead fullgraph=true backend=eager or "
|
||||
"--compile inductor_config='{\"max_autotune\": true}'"
|
||||
),
|
||||
)
|
||||
return parser
|
||||
|
||||
@@ -340,7 +460,7 @@ def video_editing_arg_parser(
|
||||
(no height/width/num-frames; resolution comes from input video). Default is distilled checkpoint only.
|
||||
"""
|
||||
parser = basic_arg_parser(distilled=distilled)
|
||||
parser.add_argument("--video-path", type=resolve_path, required=True, help="Path to the source video.")
|
||||
parser.add_argument("--video-path", type=resolve_existing_path, required=True, help="Path to the source video.")
|
||||
parser.add_argument("--start-time", type=float, required=True, help="Start time of the region to regenerate (s).")
|
||||
parser.add_argument("--end-time", type=float, required=True, help="End time of the region to regenerate (s).")
|
||||
return parser
|
||||
@@ -462,7 +582,7 @@ def default_1_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argpa
|
||||
default=video_guider.skip_step,
|
||||
help=(
|
||||
"Video skip step N controls periodic skipping during the video diffusion process: "
|
||||
"only steps where step_index % (N + 1) == 0 are processed, all others are skipped "
|
||||
"only steps where step_index %% (N + 1) == 0 are processed, all others are skipped "
|
||||
f"(e.g., 0 = no skipping; 1 = skip every other step; 2 = skip 2 of every 3 steps; "
|
||||
f"default: {video_guider.skip_step})."
|
||||
),
|
||||
@@ -522,7 +642,7 @@ def default_1_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argpa
|
||||
default=audio_guider.skip_step,
|
||||
help=(
|
||||
"Audio skip step N controls periodic skipping during the audio diffusion process: "
|
||||
"only steps where step_index % (N + 1) == 0 are processed, all others are skipped "
|
||||
"only steps where step_index %% (N + 1) == 0 are processed, all others are skipped "
|
||||
f"(e.g., 0 = no skipping; 1 = skip every other step; 2 = skip 2 of every 3 steps; "
|
||||
f"default: {audio_guider.skip_step})."
|
||||
),
|
||||
@@ -530,6 +650,62 @@ def default_1_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argpa
|
||||
return parser
|
||||
|
||||
|
||||
def default_1_stage_t2a_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser:
|
||||
"""Argument parser for single-stage text-to-audio pipelines (audio-only)."""
|
||||
audio_guider = params.audio_guider_params
|
||||
parser = basic_arg_parser(params=params)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=params.num_frames,
|
||||
help="Number of frames used to derive audio duration (num-frames / frame-rate).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frame-rate",
|
||||
type=float,
|
||||
default=params.frame_rate,
|
||||
help="Frame rate used with --num-frames to derive the audio duration.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--negative-prompt",
|
||||
type=str,
|
||||
default=DEFAULT_NEGATIVE_PROMPT,
|
||||
help="Negative prompt to steer audio generation away from artifacts.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-cfg-guidance-scale",
|
||||
type=float,
|
||||
default=audio_guider.cfg_scale,
|
||||
help=f"Audio CFG scale (default: {audio_guider.cfg_scale}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-stg-guidance-scale",
|
||||
type=float,
|
||||
default=audio_guider.stg_scale,
|
||||
help=f"Audio STG scale (default: {audio_guider.stg_scale}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-rescale-scale",
|
||||
type=float,
|
||||
default=audio_guider.rescale_scale,
|
||||
help=f"Audio rescale scale (default: {audio_guider.rescale_scale}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-stg-blocks",
|
||||
type=int,
|
||||
nargs="*",
|
||||
default=audio_guider.stg_blocks,
|
||||
help=f"Blocks to perturb for Audio STG (default: {audio_guider.stg_blocks}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-skip-step",
|
||||
type=int,
|
||||
default=audio_guider.skip_step,
|
||||
help=f"Audio skip step (default: {audio_guider.skip_step}).",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def default_2_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argparse.ArgumentParser:
|
||||
parser = default_1_stage_arg_parser(params=params)
|
||||
parser.set_defaults(height=params.stage_2_height, width=params.stage_2_width)
|
||||
@@ -562,7 +738,7 @@ def default_2_stage_arg_parser(params: PipelineParams = LTX_2_3_PARAMS) -> argpa
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spatial-upsampler-path",
|
||||
type=resolve_path,
|
||||
type=resolve_existing_path,
|
||||
required=True,
|
||||
help=(
|
||||
"Path to the spatial upsampler model used to increase the resolution "
|
||||
@@ -605,7 +781,7 @@ def default_2_stage_distilled_arg_parser(params: PipelineParams = LTX_2_3_PARAMS
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spatial-upsampler-path",
|
||||
type=resolve_path,
|
||||
type=resolve_existing_path,
|
||||
required=True,
|
||||
help=(
|
||||
"Path to the spatial upsampler model used to increase the resolution "
|
||||
|
||||
@@ -6,6 +6,7 @@ removes the need for :class:`ModelLedger`.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import logging
|
||||
from collections.abc import Iterator
|
||||
from contextlib import AbstractContextManager, contextmanager
|
||||
@@ -21,6 +22,8 @@ from ltx_core.components.noisers import Noiser
|
||||
from ltx_core.components.patchifiers import AudioPatchifier, VideoLatentPatchifier
|
||||
from ltx_core.components.protocols import DiffusionStepProtocol
|
||||
from ltx_core.loader import SDOps
|
||||
from ltx_core.loader.attention_ops import set_attention_module_op
|
||||
from ltx_core.loader.fuse_loras import bf16_fuse_rule
|
||||
from ltx_core.loader.module_ops import ModuleOps
|
||||
from ltx_core.loader.primitives import BuilderProtocol, LoraPathStrengthAndSDOps, ModelBuilderProtocol
|
||||
from ltx_core.loader.registry import DummyRegistry, Registry
|
||||
@@ -36,13 +39,21 @@ from ltx_core.model.audio_vae import (
|
||||
from ltx_core.model.audio_vae import (
|
||||
decode_audio as vae_decode_audio,
|
||||
)
|
||||
from ltx_core.model.model_protocol import LTXModelProtocol, ModelConfigurator
|
||||
from ltx_core.model.transformer import (
|
||||
LTXV_MODEL_COMFY_RENAMING_MAP,
|
||||
LTXModel,
|
||||
LTXModelConfigurator,
|
||||
X0Model,
|
||||
)
|
||||
from ltx_core.model.transformer.compiling import COMPILE_TRANSFORMER, modify_sd_ops_for_compilation
|
||||
from ltx_core.model.transformer.attention import (
|
||||
AttentionCallable,
|
||||
AttentionFunction,
|
||||
)
|
||||
from ltx_core.model.transformer.compiling import (
|
||||
CompilationConfig,
|
||||
build_compile_transformer_op,
|
||||
modify_sd_ops_for_compilation,
|
||||
)
|
||||
from ltx_core.model.upsampler import LatentUpsamplerConfigurator, upsample_video
|
||||
from ltx_core.model.video_vae import (
|
||||
MEMORY_EFFICIENT_DECODE,
|
||||
@@ -53,7 +64,7 @@ from ltx_core.model.video_vae import (
|
||||
VideoEncoder,
|
||||
VideoEncoderConfigurator,
|
||||
)
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.quantization import QuantizationPolicy, fp8_cast_fuse_rule
|
||||
from ltx_core.text_encoders.gemma import (
|
||||
EMBEDDINGS_PROCESSOR_KEY_OPS,
|
||||
GEMMA_LLM_KEY_OPS,
|
||||
@@ -66,7 +77,6 @@ from ltx_core.text_encoders.gemma.embeddings_processor import EmbeddingsProcesso
|
||||
from ltx_core.tools import AudioLatentTools, LatentTools, VideoLatentTools
|
||||
from ltx_core.types import Audio, AudioLatentShape, LatentState, VideoLatentShape, VideoPixelShape
|
||||
from ltx_core.utils import find_matching_file
|
||||
from ltx_pipelines.multigpu.delegating_builder import DelegatingBuilder
|
||||
from ltx_pipelines.utils.gpu_model import gpu_model
|
||||
from ltx_pipelines.utils.helpers import (
|
||||
cleanup_memory,
|
||||
@@ -101,6 +111,28 @@ def _chain_quantization(
|
||||
return chained_sd_ops, (*module_ops, *quantization.module_ops)
|
||||
|
||||
|
||||
def _apply_compile_ops(
|
||||
sd_ops: SDOps,
|
||||
module_ops: tuple[ModuleOps, ...],
|
||||
loras: tuple[LoraPathStrengthAndSDOps, ...],
|
||||
number_of_layers: int,
|
||||
compilation_config: CompilationConfig,
|
||||
) -> tuple[SDOps, tuple[ModuleOps, ...], tuple[LoraPathStrengthAndSDOps, ...]]:
|
||||
"""Rewrite sd_ops/module_ops/LoRAs for compiled blocks (params land under ``_orig_mod``)."""
|
||||
sd_ops = modify_sd_ops_for_compilation(sd_ops, number_of_layers)
|
||||
compile_op = build_compile_transformer_op(compilation_config)
|
||||
module_ops = (*module_ops, compile_op)
|
||||
loras = tuple(
|
||||
LoraPathStrengthAndSDOps(
|
||||
lora.path,
|
||||
lora.strength,
|
||||
modify_sd_ops_for_compilation(lora.sd_ops, number_of_layers),
|
||||
)
|
||||
for lora in loras
|
||||
)
|
||||
return sd_ops, module_ops, loras
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _streaming_model(
|
||||
builder: StreamingModelBuilder,
|
||||
@@ -111,7 +143,7 @@ def _streaming_model(
|
||||
"""Build a streaming wrapper, yield it, then tear down and free memory."""
|
||||
cpu_slots_count = DISK_CPU_SLOTS if offload_mode == OffloadMode.DISK else None
|
||||
wrapped = builder.build(
|
||||
target_device=target_device,
|
||||
device=target_device,
|
||||
dtype=dtype,
|
||||
cpu_slots_count=cpu_slots_count,
|
||||
)
|
||||
@@ -162,7 +194,7 @@ class DiffusionStage:
|
||||
pattern in every pipeline.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
def __init__( # noqa: PLR0913
|
||||
self,
|
||||
checkpoint_path: str,
|
||||
dtype: torch.dtype,
|
||||
@@ -170,79 +202,118 @@ class DiffusionStage:
|
||||
loras: tuple[LoraPathStrengthAndSDOps, ...] = (),
|
||||
quantization: QuantizationPolicy | None = None,
|
||||
registry: Registry | None = None,
|
||||
torch_compile: bool = False,
|
||||
compilation_config: CompilationConfig | None = None,
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
transformer_builder: ModelBuilderProtocol[LTXModel] | DelegatingBuilder[LTXModel] | None = None,
|
||||
transformer_builder: ModelBuilderProtocol[LTXModelProtocol] | None = None,
|
||||
model_configurator: type[ModelConfigurator] = LTXModelConfigurator,
|
||||
model_sd_ops: SDOps = LTXV_MODEL_COMFY_RENAMING_MAP,
|
||||
) -> None:
|
||||
self._checkpoint_path = checkpoint_path
|
||||
self._dtype = dtype
|
||||
self._device = device
|
||||
self._quantization = quantization
|
||||
self._torch_compile = torch_compile
|
||||
self._compilation_config = compilation_config
|
||||
self._offload_mode = offload_mode
|
||||
# A quantization policy may pin its own configurator; otherwise use the one
|
||||
# provided by the caller (defaults to the audio-video LTXModelConfigurator).
|
||||
configurator = (
|
||||
quantization.model_configurator
|
||||
if quantization is not None and quantization.model_configurator is not None
|
||||
else model_configurator
|
||||
)
|
||||
if transformer_builder is not None:
|
||||
self._transformer_builder = transformer_builder
|
||||
else:
|
||||
self._transformer_builder = Builder(
|
||||
model_path=checkpoint_path,
|
||||
model_class_configurator=LTXModelConfigurator,
|
||||
model_sd_ops=LTXV_MODEL_COMFY_RENAMING_MAP,
|
||||
model_class_configurator=configurator,
|
||||
model_sd_ops=model_sd_ops,
|
||||
loras=tuple(loras),
|
||||
registry=registry or DummyRegistry(),
|
||||
)
|
||||
|
||||
if offload_mode != OffloadMode.NONE:
|
||||
if torch_compile:
|
||||
raise ValueError("torch.compile is not supported with layer streaming")
|
||||
streaming_sd_ops: SDOps = LTXV_MODEL_COMFY_RENAMING_MAP
|
||||
# WeightsProvider currently only supports plain bf16 + fp8_cast LoRA fusion
|
||||
# (no companion-key emission). Quantization policies that emit
|
||||
# companion keys (e.g. ``.weight_scale``) cannot be streamed yet.
|
||||
if quantization is not None and quantization.fuse_rule is not fp8_cast_fuse_rule:
|
||||
raise ValueError(
|
||||
"Block streaming is not supported with this quantization policy "
|
||||
"(only bf16 and fp8_cast are currently supported)."
|
||||
)
|
||||
streaming_sd_ops: SDOps = model_sd_ops
|
||||
streaming_module_ops: tuple[ModuleOps, ...] = ()
|
||||
streaming_loras = tuple(loras)
|
||||
|
||||
if compilation_config:
|
||||
number_of_layers = self._transformer_builder.model_config()["transformer"]["num_layers"]
|
||||
streaming_sd_ops, streaming_module_ops, streaming_loras = _apply_compile_ops(
|
||||
streaming_sd_ops, streaming_module_ops, streaming_loras, number_of_layers
|
||||
)
|
||||
if quantization is not None:
|
||||
if quantization.kind != QuantizationPolicy.Kind.FP8_CAST:
|
||||
raise ValueError(
|
||||
f"Layer streaming supports only QuantizationPolicy.fp8_cast(); "
|
||||
f"got kind={quantization.kind!r} which produces heterogeneous block layouts."
|
||||
)
|
||||
streaming_sd_ops, streaming_module_ops = _chain_quantization(
|
||||
streaming_sd_ops, streaming_module_ops, quantization
|
||||
)
|
||||
self._streaming_builder = StreamingModelBuilder(
|
||||
model_class_configurator=LTXModelConfigurator,
|
||||
model_class_configurator=configurator,
|
||||
model_path=checkpoint_path,
|
||||
model_sd_ops=streaming_sd_ops,
|
||||
module_ops=streaming_module_ops,
|
||||
loras=tuple(loras),
|
||||
loras=streaming_loras,
|
||||
registry=registry or DummyRegistry(),
|
||||
blocks_attr="velocity_model.transformer_blocks",
|
||||
fuse_rule=quantization.fuse_rule if quantization is not None else bf16_fuse_rule,
|
||||
blocks_attr="transformer_blocks",
|
||||
blocks_prefix="transformer_blocks",
|
||||
state_dict_prefix="velocity_model.",
|
||||
model_wrapper=lambda m: X0Model(m).eval(),
|
||||
)
|
||||
|
||||
def with_attention(self, attention: AttentionFunction | AttentionCallable | None) -> "DiffusionStage":
|
||||
"""Return a new ``DiffusionStage`` that pins the transformer build to ``attention``.
|
||||
Functional: never mutates ``self``. The returned stage shares all other
|
||||
configuration with the original; only the underlying builders' ``module_ops``
|
||||
gain a ``set_attention_module_op(attention)`` entry so subsequent transformer
|
||||
builds use that kernel. ``attention=None`` is a no-op (returns ``self``).
|
||||
"""
|
||||
if attention is None:
|
||||
return self
|
||||
op = set_attention_module_op(attention)
|
||||
new = copy.copy(self)
|
||||
new._transformer_builder = self._transformer_builder.with_module_ops(
|
||||
(*self._transformer_builder.module_ops, op),
|
||||
)
|
||||
if self._offload_mode != OffloadMode.NONE:
|
||||
new._streaming_builder = self._streaming_builder.with_module_ops(
|
||||
(*self._streaming_builder.module_ops, op),
|
||||
)
|
||||
return new
|
||||
|
||||
def _build_transformer(self, *, device: torch.device | None = None, **kwargs: object) -> X0Model:
|
||||
target = device or self._device
|
||||
sd_ops = self._transformer_builder.model_sd_ops
|
||||
module_ops = self._transformer_builder.module_ops
|
||||
loras = self._transformer_builder.loras
|
||||
if self._torch_compile:
|
||||
module_ops = (*module_ops, COMPILE_TRANSFORMER)
|
||||
if self._compilation_config is not None:
|
||||
number_of_layers = self._transformer_builder.model_config()["transformer"]["num_layers"]
|
||||
sd_ops = modify_sd_ops_for_compilation(sd_ops, number_of_layers)
|
||||
loras = tuple(
|
||||
LoraPathStrengthAndSDOps(
|
||||
lora.path,
|
||||
lora.strength,
|
||||
modify_sd_ops_for_compilation(lora.sd_ops, number_of_layers),
|
||||
)
|
||||
for lora in loras
|
||||
sd_ops, module_ops, loras = _apply_compile_ops(
|
||||
sd_ops, module_ops, loras, number_of_layers, self._compilation_config
|
||||
)
|
||||
if self._quantization is not None:
|
||||
sd_ops, module_ops = _chain_quantization(sd_ops, module_ops, self._quantization)
|
||||
|
||||
builder = self._transformer_builder.with_module_ops(module_ops).with_sd_ops(sd_ops).with_loras(loras)
|
||||
if self._quantization is not None:
|
||||
builder = builder.with_fuse_rule(self._quantization.fuse_rule)
|
||||
return X0Model(builder.build(device=target, **kwargs)).to(target).eval()
|
||||
|
||||
@contextmanager
|
||||
def _streaming_transformer_ctx(self) -> Iterator[X0Model]:
|
||||
with _streaming_model(
|
||||
self._streaming_builder, self._offload_mode, self._device, self._dtype
|
||||
) as streaming_wrapper:
|
||||
yield X0Model(streaming_wrapper).eval()
|
||||
|
||||
def _transformer_ctx(self, **kwargs: object) -> AbstractContextManager:
|
||||
if self._offload_mode != OffloadMode.NONE:
|
||||
return _streaming_model(self._streaming_builder, self._offload_mode, self._device, self._dtype)
|
||||
return self._streaming_transformer_ctx()
|
||||
return gpu_model(self._build_transformer(**kwargs))
|
||||
|
||||
def model_context(self, **kwargs: object) -> AbstractContextManager:
|
||||
@@ -348,7 +419,17 @@ class DiffusionStage:
|
||||
v_shape = VideoLatentShape.from_pixel_shape(pixel_shape)
|
||||
video_tools = VideoLatentTools(VideoLatentPatchifier(patch_size=1), v_shape, fps)
|
||||
|
||||
mode = "streaming" if self._offload_mode != OffloadMode.NONE else "standard"
|
||||
logger.info("Building transformer (%s) from %s", mode, self._checkpoint_path)
|
||||
with self._transformer_ctx(video_tools=video_tools) as transformer:
|
||||
logger.info(
|
||||
"Running denoising loop (%d steps, %dx%d %d frames @ %.1f fps)",
|
||||
len(sigmas) - 1,
|
||||
width,
|
||||
height,
|
||||
frames,
|
||||
fps,
|
||||
)
|
||||
return self.run(
|
||||
transformer,
|
||||
denoiser,
|
||||
@@ -387,6 +468,8 @@ class PromptEncoder:
|
||||
offload_mode: OffloadMode = OffloadMode.NONE,
|
||||
text_encoder_builder: BuilderProtocol | None = None,
|
||||
) -> None:
|
||||
self._gemma_root = gemma_root
|
||||
self._checkpoint_path = checkpoint_path
|
||||
self._dtype = dtype
|
||||
self._device = device
|
||||
self._offload_mode = offload_mode
|
||||
@@ -432,7 +515,7 @@ class PromptEncoder:
|
||||
|
||||
def _build_embeddings_processor(self) -> EmbeddingsProcessor:
|
||||
"""Build the embeddings processor on the target device."""
|
||||
return self._embeddings_processor_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
return self._embeddings_processor_builder.build(device=self._device, dtype=self._dtype).eval()
|
||||
|
||||
def _text_encoder_ctx(self) -> AbstractContextManager:
|
||||
if self._offload_mode != OffloadMode.NONE:
|
||||
@@ -448,16 +531,20 @@ class PromptEncoder:
|
||||
enhance_prompt_seed: int = 42,
|
||||
) -> list[EmbeddingsProcessorOutput]:
|
||||
"""Encode *prompts* through Gemma -> embeddings processor, freeing each model after use."""
|
||||
logger.info("Building text encoder from %s", self._gemma_root)
|
||||
with self._text_encoder_ctx() as text_encoder:
|
||||
if enhance_first_prompt:
|
||||
prompts = list(prompts)
|
||||
prompts[0] = generate_enhanced_prompt(
|
||||
text_encoder, prompts[0], enhance_prompt_image, seed=enhance_prompt_seed
|
||||
)
|
||||
raw_outputs = [text_encoder.encode(p) for p in prompts]
|
||||
raw_outputs = text_encoder.encode(prompts)
|
||||
logger.info("Text encoder done, building embeddings processor from %s", self._checkpoint_path)
|
||||
|
||||
with gpu_model(self._build_embeddings_processor()) as embeddings_processor:
|
||||
return [embeddings_processor.process_hidden_states(hs, mask) for hs, mask in raw_outputs]
|
||||
result = [embeddings_processor.process_hidden_states(hs, mask) for hs, mask in raw_outputs]
|
||||
logger.info("Prompt encoding complete")
|
||||
return result
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -487,7 +574,7 @@ class ImageConditioner:
|
||||
)
|
||||
|
||||
def _build_encoder(self) -> VideoEncoder:
|
||||
return self._encoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
return self._encoder_builder.build(device=self._device, dtype=self._dtype).eval()
|
||||
|
||||
def __call__(self, fn: Callable[[VideoEncoder], T]) -> T:
|
||||
"""Build video encoder → call *fn(encoder)* → free encoder."""
|
||||
@@ -511,6 +598,7 @@ class VideoUpsampler:
|
||||
device: torch.device,
|
||||
registry: Registry | None = None,
|
||||
) -> None:
|
||||
self._upsampler_path = upsampler_path
|
||||
self._dtype = dtype
|
||||
self._device = device
|
||||
self._encoder_builder = Builder(
|
||||
@@ -527,13 +615,10 @@ class VideoUpsampler:
|
||||
|
||||
def __call__(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
"""Upsample *latent* using video encoder + spatial upsampler, then free both."""
|
||||
logger.info("Building video encoder + spatial upsampler from %s", self._upsampler_path)
|
||||
with (
|
||||
gpu_model(
|
||||
self._encoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
) as encoder,
|
||||
gpu_model(
|
||||
self._upsampler_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
) as upsampler,
|
||||
gpu_model(self._encoder_builder.build(device=self._device, dtype=self._dtype).eval()) as encoder,
|
||||
gpu_model(self._upsampler_builder.build(device=self._device, dtype=self._dtype).eval()) as upsampler,
|
||||
):
|
||||
return upsample_video(latent=latent, video_encoder=encoder, upsampler=upsampler)
|
||||
|
||||
@@ -557,6 +642,7 @@ class VideoDecoder:
|
||||
memory_efficient: bool = True,
|
||||
decoder_builder: BuilderProtocol | None = None,
|
||||
) -> None:
|
||||
self._checkpoint_path = checkpoint_path
|
||||
self._dtype = dtype
|
||||
self._device = device
|
||||
if decoder_builder is not None:
|
||||
@@ -577,7 +663,8 @@ class VideoDecoder:
|
||||
generator: torch.Generator | None = None,
|
||||
) -> Iterator[torch.Tensor]:
|
||||
"""Decode *latent* to pixel-space video chunks. Decoder freed after exhaustion."""
|
||||
decoder = self._decoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
logger.info("Building video decoder from %s", self._checkpoint_path)
|
||||
decoder = self._decoder_builder.build(device=self._device, dtype=self._dtype).eval()
|
||||
return _cleanup_iter(decoder.decode_video(latent, tiling_config, generator), decoder)
|
||||
|
||||
|
||||
@@ -596,6 +683,7 @@ class AudioDecoder:
|
||||
device: torch.device,
|
||||
registry: Registry | None = None,
|
||||
) -> None:
|
||||
self._checkpoint_path = checkpoint_path
|
||||
self._dtype = dtype
|
||||
self._device = device
|
||||
self._decoder_builder = Builder(
|
||||
@@ -613,13 +701,10 @@ class AudioDecoder:
|
||||
|
||||
def __call__(self, latent: torch.Tensor) -> Audio:
|
||||
"""Decode audio *latent* through VAE decoder + vocoder, then free both."""
|
||||
logger.info("Building audio decoder + vocoder from %s", self._checkpoint_path)
|
||||
with (
|
||||
gpu_model(
|
||||
self._decoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
) as decoder,
|
||||
gpu_model(
|
||||
self._vocoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
) as vocoder,
|
||||
gpu_model(self._decoder_builder.build(device=self._device, dtype=self._dtype).eval()) as decoder,
|
||||
gpu_model(self._vocoder_builder.build(device=self._device, dtype=self._dtype).eval()) as vocoder,
|
||||
):
|
||||
return vae_decode_audio(latent, decoder, vocoder)
|
||||
|
||||
@@ -653,7 +738,5 @@ class AudioConditioner:
|
||||
|
||||
def __call__(self, fn: Callable[[torch.nn.Module], T]) -> T:
|
||||
"""Build audio encoder → call *fn(encoder)* → free encoder."""
|
||||
with gpu_model(
|
||||
self._encoder_builder.build(device=self._device, dtype=self._dtype).to(self._device).eval()
|
||||
) as encoder:
|
||||
with gpu_model(self._encoder_builder.build(device=self._device, dtype=self._dtype).eval()) as encoder:
|
||||
return fn(encoder)
|
||||
|
||||
@@ -138,6 +138,8 @@ def _guided_denoise( # noqa: PLR0913,PLR0915
|
||||
ptb_configs = [ptb for _, _, _, ptb in passes]
|
||||
n = len(passes)
|
||||
|
||||
orig_b = (video_state or audio_state).latent.shape[0]
|
||||
|
||||
def _batched_sigma(state: LatentState) -> torch.Tensor:
|
||||
"""Expand scalar sigma to (n * B,) matching the repeated state."""
|
||||
return sigma.expand(state.latent.shape[0] * n)
|
||||
@@ -162,8 +164,16 @@ def _guided_denoise( # noqa: PLR0913,PLR0915
|
||||
enabled=not a_skip,
|
||||
)
|
||||
|
||||
# Replicate each pass's PerturbationConfig to all `orig_b` samples it
|
||||
# carries, so `BatchedPerturbationConfig.mask_like` returns a per-sample
|
||||
# mask (length n*orig_b) instead of a per-pass mask (length n). Without
|
||||
# this expansion the mask is broadcast against a (n*orig_b, T, D) tensor
|
||||
# and the multiplication fails with a batch-dim mismatch whenever
|
||||
# `orig_b > 1` (e.g. multi-prompt benchmark panels).
|
||||
batched_ptb_configs = [ptb for ptb in ptb_configs for _ in range(orig_b)]
|
||||
|
||||
all_v, all_a = transformer(
|
||||
video=batched_video, audio=batched_audio, perturbations=BatchedPerturbationConfig(ptb_configs)
|
||||
video=batched_video, audio=batched_audio, perturbations=BatchedPerturbationConfig(batched_ptb_configs)
|
||||
)
|
||||
|
||||
# Split results back and combine via guiders.
|
||||
@@ -172,6 +182,8 @@ def _guided_denoise( # noqa: PLR0913,PLR0915
|
||||
r = dict(zip(pass_names, zip(splits_v, splits_a, strict=True), strict=True))
|
||||
|
||||
cond_v, cond_a = r["cond"]
|
||||
cond_v = cond_v if isinstance(cond_v, torch.Tensor) else torch.tensor(cond_v)
|
||||
cond_a = cond_a if isinstance(cond_a, torch.Tensor) else torch.tensor(cond_a)
|
||||
uncond_v, uncond_a = r.get("uncond", (0.0, 0.0))
|
||||
ptb_v, ptb_a = r.get("ptb", (0.0, 0.0))
|
||||
mod_v, mod_a = r.get("mod", (0.0, 0.0))
|
||||
|
||||
@@ -391,6 +391,24 @@ def encode_video(
|
||||
logger.info(f"Video saved to {output_path}")
|
||||
|
||||
|
||||
def encode_audio(audio: Audio, output_path: str) -> None:
|
||||
"""Save an audio waveform as a 16-bit PCM ``.wav`` file at the source sampling rate.
|
||||
Reuses :func:`_write_audio` (the same muxing path used by :func:`encode_video`);
|
||||
the only difference is a PCM (``pcm_s16le``) stream in a WAV container instead of
|
||||
the AAC stream used for muxed video.
|
||||
"""
|
||||
container = av.open(output_path, mode="w")
|
||||
audio_stream = container.add_stream("pcm_s16le", rate=audio.sampling_rate)
|
||||
audio_stream.codec_context.sample_rate = audio.sampling_rate
|
||||
audio_stream.codec_context.layout = "stereo"
|
||||
audio_stream.codec_context.time_base = Fraction(1, audio.sampling_rate)
|
||||
try:
|
||||
_write_audio(container, audio_stream, audio)
|
||||
finally:
|
||||
container.close()
|
||||
logger.info(f"Audio saved to {output_path}")
|
||||
|
||||
|
||||
def _encode_chunks_threaded(
|
||||
container: av.container.Container,
|
||||
stream: av.video.stream.VideoStream,
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
"""User-facing quantization-policy dispatch.
|
||||
``ltx-core`` exposes one ``build_policy`` factory per backend. This module
|
||||
provides the user-facing string-keyed dispatch used by CLI args and pipeline
|
||||
defaults — keeping the enum out of ``ltx-core`` so adding/removing backends is
|
||||
a single-file change here.
|
||||
"""
|
||||
|
||||
from enum import Enum
|
||||
|
||||
from typing_extensions import assert_never
|
||||
|
||||
from ltx_core.quantization import QuantizationPolicy
|
||||
from ltx_core.quantization.fp8_cast import build_policy as _build_fp8_cast_policy
|
||||
from ltx_core.quantization.fp8_scaled_mm import build_policy as _build_fp8_scaled_mm_policy
|
||||
|
||||
|
||||
class QuantizationKind(str, Enum):
|
||||
FP8_CAST = "fp8-cast"
|
||||
FP8_SCALED_MM = "fp8-scaled-mm"
|
||||
|
||||
def to_policy(self, checkpoint_path: str | None = None) -> QuantizationPolicy:
|
||||
"""Build the :class:`QuantizationPolicy` for this kind.
|
||||
``checkpoint_path`` is required for both backends: ``FP8_SCALED_MM``
|
||||
uses it to discover the layer set from ``.weight_scale`` tensors,
|
||||
and ``FP8_CAST`` uses it to fold any prequant scales into the fp8
|
||||
weight at load time.
|
||||
"""
|
||||
if checkpoint_path is None:
|
||||
raise ValueError(f"{self.value} quantization requires checkpoint_path.")
|
||||
match self:
|
||||
case QuantizationKind.FP8_CAST:
|
||||
return _build_fp8_cast_policy(checkpoint_path)
|
||||
case QuantizationKind.FP8_SCALED_MM:
|
||||
return _build_fp8_scaled_mm_policy(checkpoint_path)
|
||||
case _:
|
||||
assert_never(self)
|
||||
@@ -436,7 +436,7 @@ def res2s_audio_video_denoising_loop( # noqa: PLR0913,PLR0915,PLR0912
|
||||
return video_state, audio_state
|
||||
|
||||
|
||||
def euler_cfg_pp_denoising_loop(
|
||||
def euler_cfg_pp_denoising_loop( # noqa: PLR0912
|
||||
sigmas: torch.Tensor,
|
||||
video_state: LatentState | None,
|
||||
audio_state: LatentState | None,
|
||||
@@ -514,9 +514,15 @@ def euler_cfg_pp_denoising_loop(
|
||||
)
|
||||
|
||||
if video_state is not None and denoised_video is not None:
|
||||
denoised_video = post_process_latent(denoised_video, video_state.denoise_mask, video_state.clean_latent)
|
||||
denoised_video = post_process_latent(
|
||||
denoised_video.float(), video_state.denoise_mask, video_state.clean_latent
|
||||
)
|
||||
noisy_video = video_state.latent.float()
|
||||
if audio_state is not None and denoised_audio is not None:
|
||||
denoised_audio = post_process_latent(denoised_audio, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
denoised_audio = post_process_latent(
|
||||
denoised_audio.float(), audio_state.denoise_mask, audio_state.clean_latent
|
||||
)
|
||||
noisy_audio = audio_state.latent.float()
|
||||
|
||||
if sigmas[step_idx + 1] == 0:
|
||||
if video_state is not None and denoised_video is not None:
|
||||
@@ -525,30 +531,31 @@ def euler_cfg_pp_denoising_loop(
|
||||
audio_state = replace(audio_state, latent=denoised_audio.to(model_dtype))
|
||||
return video_state, audio_state
|
||||
|
||||
# Draw noise consecutively from the same generator: video first, audio second.
|
||||
noise_video = new_noise_fn(video_state.latent, generator) if (video_state is not None and draw_noise) else None
|
||||
noise_audio = new_noise_fn(audio_state.latent, generator) if (audio_state is not None and draw_noise) else None
|
||||
|
||||
if video_state is not None and denoised_video is not None:
|
||||
video_noise = new_noise_fn(video_state.latent, generator) if draw_noise else None
|
||||
x_next = stepper.step(
|
||||
sample=video_state.latent,
|
||||
sample=noisy_video,
|
||||
denoised_sample=denoised_video,
|
||||
sigmas=sigmas,
|
||||
step_index=step_idx,
|
||||
uncond_denoised=uncond_video,
|
||||
noise=noise_video,
|
||||
noise=video_noise,
|
||||
)
|
||||
if draw_noise:
|
||||
x_next = post_process_latent(x_next, video_state.denoise_mask, video_state.clean_latent)
|
||||
video_state = replace(video_state, latent=x_next.to(model_dtype))
|
||||
|
||||
if audio_state is not None and denoised_audio is not None:
|
||||
audio_noise = new_noise_fn(audio_state.latent, generator) if draw_noise else None
|
||||
x_next = stepper.step(
|
||||
sample=audio_state.latent,
|
||||
sample=noisy_audio,
|
||||
denoised_sample=denoised_audio,
|
||||
sigmas=sigmas,
|
||||
step_index=step_idx,
|
||||
uncond_denoised=uncond_audio,
|
||||
noise=noise_audio,
|
||||
noise=audio_noise,
|
||||
)
|
||||
if draw_noise:
|
||||
x_next = post_process_latent(x_next, audio_state.denoise_mask, audio_state.clean_latent)
|
||||
audio_state = replace(audio_state, latent=x_next.to(model_dtype))
|
||||
|
||||
return video_state, audio_state
|
||||
|
||||
@@ -1,7 +1,4 @@
|
||||
configs/*.yaml
|
||||
!configs/ltx2_av_lora.yaml
|
||||
!configs/ltx2_av_lora_low_vram.yaml
|
||||
!configs/ltx2_v2v_ic_lora.yaml
|
||||
datasets
|
||||
outputs
|
||||
wandb
|
||||
|
||||
+125
-35
@@ -6,10 +6,22 @@ This file provides guidance to AI coding assistants (Claude, Cursor, etc.) when
|
||||
|
||||
**LTX Trainer** is a training toolkit for fine-tuning the Lightricks LTX audio-video generation models. It supports:
|
||||
|
||||
- **Text-to-video (T2V)** - Generate video from text prompts
|
||||
- **Text-to-audio (T2A)** - Generate audio from text prompts
|
||||
- **Image-to-video (I2V)** - Generate video conditioned on a first frame
|
||||
- **Video extension** - Forward (prefix) and backward (suffix) video continuation
|
||||
- **Video inpainting** - Mask-based spatial/temporal inpainting
|
||||
- **Video outpainting** - Spatial crop-based outpainting
|
||||
- **IC-LoRA video-to-video** - In-context control adapters for style/structure transfer
|
||||
- **Audio-to-video (A2V)** and **Video-to-audio (V2A)** - Cross-modal generation with frozen conditioning
|
||||
- **Audio extension** - Forward (prefix) and backward (suffix) audio continuation
|
||||
- **Audio inpainting** - Mask-based audio inpainting
|
||||
- **IC-LoRA audio-to-audio (A2A)** - Audio reference conditioning for style transfer
|
||||
- **AV2AV IC-LoRA** - Combined video and audio reference conditioning
|
||||
- **LoRA training** - Efficient fine-tuning with adapters
|
||||
- **Full fine-tuning** - Complete model training
|
||||
- **Audio-video training** - Joint audio and video generation
|
||||
- **IC-LoRA training** - In-context control adapters for video-to-video transformations
|
||||
|
||||
All conditioning scenarios are expressed through the unified `FlexibleStrategy` configuration.
|
||||
|
||||
**Supported model versions:**
|
||||
|
||||
@@ -39,13 +51,14 @@ packages/ltx-trainer/
|
||||
│ ├── config_display.py # Config pretty-printing
|
||||
│ ├── trainer.py # Main training orchestration with Accelerate
|
||||
│ ├── model_loader.py # Model loading using ltx-core
|
||||
│ ├── validation_sampler.py # Inference for validation samples
|
||||
│ ├── validation_runner.py # ValidationRunner — conditioned validation sampling
|
||||
│ ├── datasets.py # PrecomputedDataset, DummyDataset
|
||||
│ ├── training_strategies/ # Strategy pattern for different training modes
|
||||
│ │ ├── __init__.py # Factory function: get_training_strategy()
|
||||
│ │ ├── base_strategy.py # TrainingStrategy ABC, ModelInputs, TrainingStrategyConfigBase
|
||||
│ │ ├── text_to_video.py # TextToVideoStrategy, TextToVideoConfig
|
||||
│ │ └── video_to_video.py # VideoToVideoStrategy, VideoToVideoConfig
|
||||
│ │ ├── flexible.py # FlexibleStrategy, FlexibleStrategyConfig [RECOMMENDED]
|
||||
│ │ ├── text_to_video.py # TextToVideoStrategy, TextToVideoConfig [DEPRECATED]
|
||||
│ │ └── video_to_video.py # VideoToVideoStrategy, VideoToVideoConfig [DEPRECATED]
|
||||
│ ├── timestep_samplers.py # Flow matching timestep sampling
|
||||
│ ├── gemma_8bit.py # 8-bit Gemma text encoder loading (bitsandbytes)
|
||||
│ ├── quantization.py # Transformer INT8/INT4/FP8 quantization
|
||||
@@ -62,13 +75,25 @@ packages/ltx-trainer/
|
||||
│ ├── process_captions.py # Text embedding computation
|
||||
│ ├── caption_videos.py # Automatic video captioning
|
||||
│ ├── decode_latents.py # Latent decoding for debugging
|
||||
│ ├── inference.py # Inference with trained models
|
||||
│ ├── compute_reference.py # Generate IC-LoRA reference videos
|
||||
│ └── split_scenes.py # Scene detection and splitting
|
||||
├── configs/ # Example training configurations
|
||||
│ ├── ltx2_av_lora.yaml # Audio-video LoRA training
|
||||
│ ├── ltx2_av_lora_low_vram.yaml
|
||||
│ ├── ltx2_v2v_ic_lora.yaml # IC-LoRA video-to-video
|
||||
│ ├── t2v_lora.yaml # Text-to-video LoRA
|
||||
│ ├── t2v_lora_low_vram.yaml # Text-to-video LoRA (low VRAM)
|
||||
│ ├── i2v_lora.yaml # Image-to-video LoRA
|
||||
│ ├── v2v_ic_lora.yaml # IC-LoRA video-to-video
|
||||
│ ├── a2v_lora.yaml # Audio-to-video LoRA
|
||||
│ ├── v2a_lora.yaml # Video-to-audio LoRA
|
||||
│ ├── video_extend_lora.yaml # Video extension (forward)
|
||||
│ ├── video_suffix_lora.yaml # Video extension (backward)
|
||||
│ ├── video_inpainting_lora.yaml # Video inpainting
|
||||
│ ├── video_outpainting_lora.yaml # Video outpainting
|
||||
│ ├── t2a_lora.yaml # Text-to-audio LoRA
|
||||
│ ├── audio_extend_lora.yaml # Audio extension (forward)
|
||||
│ ├── audio_suffix_lora.yaml # Audio extension (backward)
|
||||
│ ├── audio_inpainting_lora.yaml # Audio inpainting
|
||||
│ ├── a2a_ic_lora.yaml # Audio-to-audio IC-LoRA
|
||||
│ ├── av2av_ic_lora.yaml # AV2AV IC-LoRA
|
||||
│ └── accelerate/ # FSDP, DDP configs
|
||||
├── tests/ # Pytest tests
|
||||
└── docs/ # Documentation
|
||||
@@ -83,7 +108,8 @@ packages/ltx-trainer/
|
||||
`load_text_encoder()`, `load_embeddings_processor()`, etc.
|
||||
- Combined loader: `load_model()` returns `LtxModelComponents` dataclass
|
||||
- Uses `SingleGPUModelBuilder` from ltx-core internally
|
||||
- Text encoder and embeddings processor are loaded separately (the text encoder only needs Gemma weights; the embeddings processor only needs the LTX checkpoint)
|
||||
- Text encoder and embeddings processor are loaded separately (the text encoder only needs Gemma weights; the embeddings
|
||||
processor only needs the LTX checkpoint)
|
||||
- 8-bit text encoder loading via `gemma_8bit.py` (bitsandbytes)
|
||||
|
||||
**Training Flow:**
|
||||
@@ -94,7 +120,7 @@ packages/ltx-trainer/
|
||||
kept)
|
||||
4. Each training step: embedding connectors applied → strategy prepares `ModelInputs` → transformer forward pass →
|
||||
strategy computes loss
|
||||
5. Training strategies (`TextToVideoStrategy`, `VideoToVideoStrategy`) handle mode-specific logic
|
||||
5. Training strategies (`FlexibleStrategy`) handle mode-specific logic including conditioning, masking, and loss computation
|
||||
6. Accelerate handles distributed training, mixed precision, and device placement
|
||||
7. Data flows as precomputed latents through `PrecomputedDataset`
|
||||
|
||||
@@ -136,7 +162,12 @@ LTX-2.3) and cross-modality (video↔audio) attention conditioning (both version
|
||||
|
||||
- All config in `src/ltx_trainer/config.py`
|
||||
- Main class: `LtxTrainerConfig`
|
||||
- Training strategy configs: `TextToVideoConfig`, `VideoToVideoConfig`
|
||||
- `TrainingStrategyConfig` - Union of `FlexibleStrategyConfig` | `TextToVideoConfig` (deprecated) | `VideoToVideoConfig` (deprecated)
|
||||
- `FlexibleStrategyConfig` - Unified strategy config with `video`/`audio` `ModalityConfig` blocks
|
||||
- `ModalityConfig` - Per-modality config: `is_generated`, `latents_dir`, `conditions` list
|
||||
- `ConditionConfig` - Discriminated union: `FirstFrameConditionConfig`, `PrefixConditionConfig`, `SuffixConditionConfig`, `SpatialCropConditionConfig`, `MaskConditionConfig`, `ReferenceConditionConfig`
|
||||
- `ValidationSample` - Per-sample validation config with `prompt`, `conditions`, optional `video_dims`/`seed` overrides
|
||||
- `ValidationCondition` - Discriminated union for validation conditions (first_frame, prefix, suffix, spatial_crop, mask, reference, video_to_audio, audio_to_video)
|
||||
- Uses Pydantic field validators and model validators
|
||||
- Config uses `extra="forbid"` — unknown fields cause validation errors
|
||||
- Config files in `configs/` directory
|
||||
@@ -206,8 +237,8 @@ These values are shared across all supported model versions:
|
||||
| Video latent channels | 128 | VAE encoder/decoder, patchifier, `VideoLatentShape` |
|
||||
| Spatial compression | 32× (H and W) | `SpatioTemporalScaleFactors.default()`, config validators |
|
||||
| Temporal compression | 8× | `SpatioTemporalScaleFactors.default()`, config validators |
|
||||
| Frame constraint | `frames % 8 == 1` | Config validators, validation sampler |
|
||||
| Resolution constraint | Width and height divisible by 32 | Config validators, validation sampler |
|
||||
| Frame constraint | `frames % 8 == 1` | Config validators, validation runner |
|
||||
| Resolution constraint | Width and height divisible by 32 | Config validators, validation runner |
|
||||
| Audio latent channels | 8 | `AudioLatentShape`, audio patchifier |
|
||||
| Audio mel bins | 16 | `AudioLatentShape`, audio patchifier |
|
||||
| Patchified token dim (video) | 128 (`128 × 1 × 1 × 1`) | Transformer `in_channels` |
|
||||
@@ -245,12 +276,53 @@ uv run pytest
|
||||
|
||||
```bash
|
||||
# Single GPU
|
||||
uv run python scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run python scripts/train.py configs/t2v_lora.yaml
|
||||
|
||||
# Multi-GPU with Accelerate
|
||||
uv run accelerate launch scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run accelerate launch scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
## Testing Standards
|
||||
|
||||
### Structure
|
||||
|
||||
- **Flat functions only** — use `def test_*()`, never `class Test*` with methods. Pytest collects standalone functions.
|
||||
- **Only test public interfaces** — never call private methods (`_method`) directly. Verify private behavior
|
||||
indirectly through the public API.
|
||||
|
||||
### What to Test
|
||||
|
||||
- **Custom validators and business logic** — cross-field validators, domain constraints, error paths. These catch real
|
||||
bugs.
|
||||
- **Behavioral tests** — call the public method, verify the outputs have the right shape, values, and structure. One
|
||||
behavioral test is worth ten config-only tests.
|
||||
- **Edge cases and error paths** — boundary conditions, composed behaviors, expected exceptions.
|
||||
- **Contract tests** — required fields, rejected invalid inputs, safety mechanisms like `extra="forbid"`.
|
||||
|
||||
### What NOT to Test
|
||||
|
||||
- **Pydantic storing a value** — `Foo(x=1); assert foo.x == 1` tests Pydantic, not your code. If a behavioral test
|
||||
already creates the same config and uses it, the config-only test adds nothing.
|
||||
- **Pydantic Literal defaults** — `assert config.type == "first_frame"` when `type` is `Literal["first_frame"]`.
|
||||
- **Pydantic default factories** — `assert config.conditions == []` when the field has `default_factory=list`.
|
||||
- **Tests already covered by behavioral tests** — if `test_prefix_conditioning` creates a valid `PrefixConditionConfig`
|
||||
and exercises it end-to-end, a separate `test_prefix_valid` that just creates the same config is redundant.
|
||||
- **Trivial instantiation tests** — `strategy = Strategy(config); assert strategy.config is not None` when every other
|
||||
test creates a strategy.
|
||||
|
||||
### Keeping Tests DRY
|
||||
|
||||
- **Use helper functions** for repeated setup patterns (e.g., `_make_strategy(video=_video_modality(...))` instead of
|
||||
6-8 lines of config/strategy creation per test).
|
||||
- **Use named constants** for test dimensions (e.g., `VIDEO_SEQ_LEN`, `TOKENS_PER_FRAME`) instead of magic numbers.
|
||||
- **Merge tests that share identical setup** — when 5+ tests call `prepare_training_inputs` with the exact same
|
||||
config and batch, each checking one assertion, merge them into one test that checks all assertions. Pytest reports
|
||||
the exact failing line anyway.
|
||||
- **Use `@pytest.mark.parametrize`** for the same logic tested with different inputs (e.g., valid/invalid values for
|
||||
a field).
|
||||
- **Use pytest fixtures** for shared batch data and test directories, but prefer explicit helper functions over
|
||||
fixtures for strategy/config creation (makes the test self-documenting).
|
||||
|
||||
## Code Standards
|
||||
|
||||
### Type Hints
|
||||
@@ -286,7 +358,12 @@ Key classes:
|
||||
|
||||
- `LtxTrainerConfig` - Main configuration container
|
||||
- `ModelConfig` - Model paths, training mode (`lora` | `full`), checkpoint loading
|
||||
- `TrainingStrategyConfig` - Union of `TextToVideoConfig` | `VideoToVideoConfig` (discriminated by `name`)
|
||||
- `TrainingStrategyConfig` - Union of `FlexibleStrategyConfig` | `TextToVideoConfig` (deprecated) | `VideoToVideoConfig` (deprecated)
|
||||
- `FlexibleStrategyConfig` - Unified strategy config with `video`/`audio` `ModalityConfig` blocks
|
||||
- `ModalityConfig` - Per-modality config: `is_generated`, `latents_dir`, `conditions` list
|
||||
- `ConditionConfig` - Discriminated union: `FirstFrameConditionConfig`, `PrefixConditionConfig`, `SuffixConditionConfig`, `SpatialCropConditionConfig`, `MaskConditionConfig`, `ReferenceConditionConfig`
|
||||
- `ValidationSample` - Per-sample validation config with `prompt`, `conditions`, optional `video_dims`/`seed` overrides
|
||||
- `ValidationCondition` - Discriminated union for validation conditions (first_frame, prefix, suffix, spatial_crop, mask, reference, video_to_audio, audio_to_video)
|
||||
- `LoraConfig` - Rank, alpha, dropout, target modules
|
||||
- `OptimizationConfig` - Learning rate, batch size, gradient accumulation, scheduler, gradient checkpointing
|
||||
- `AccelerationConfig` - Mixed precision, quantization, 8-bit text encoder
|
||||
@@ -310,21 +387,23 @@ Key classes:
|
||||
- Implements distributed training with Accelerate
|
||||
- Handles mixed precision, gradient accumulation, checkpointing
|
||||
- `_training_step()` applies embedding connectors then delegates to strategy
|
||||
- `_load_text_encoder_and_cache_embeddings()` loads the text encoder + embeddings processor, caches validation embeddings, then unloads the Gemma LLM (keeps only the embeddings processor connectors for training)
|
||||
- `_load_text_encoder_and_cache_embeddings()` loads the text encoder + embeddings processor, caches validation
|
||||
embeddings, then unloads the Gemma LLM (keeps only the embeddings processor connectors for training)
|
||||
- Uses training strategies for mode-specific logic
|
||||
|
||||
**`src/ltx_trainer/training_strategies/`** - Strategy pattern
|
||||
|
||||
- `base_strategy.py`: `TrainingStrategy` ABC, `ModelInputs` dataclass
|
||||
- `text_to_video.py`: Standard text-to-video (with optional audio)
|
||||
- `video_to_video.py`: IC-LoRA video-to-video transformations
|
||||
- `flexible.py`: FlexibleStrategy — unified conditioning framework (recommended)
|
||||
- `text_to_video.py`: TextToVideoStrategy (deprecated — use FlexibleStrategy)
|
||||
- `video_to_video.py`: VideoToVideoStrategy (deprecated — use FlexibleStrategy)
|
||||
|
||||
Key methods each strategy implements:
|
||||
|
||||
- `get_data_sources()` - Required data directories
|
||||
- `prepare_training_inputs()` - Convert batch to `ModelInputs` with `Modality` objects
|
||||
- `compute_loss()` - Calculate training loss (velocity prediction, MSE with masking)
|
||||
- `requires_audio` property - Whether audio components needed
|
||||
|
||||
The strategy's **config** declares its data directories via `get_data_sources()` (single source of truth, used for both dataset wiring and existence validation).
|
||||
|
||||
**`src/ltx_trainer/model_loader.py`** - Model loading
|
||||
|
||||
@@ -339,14 +418,13 @@ Component loaders:
|
||||
- `load_embeddings_processor(checkpoint_path)` → `EmbeddingsProcessor` (feature extractor + connectors)
|
||||
- `load_model()` → `LtxModelComponents` (convenience wrapper)
|
||||
|
||||
**`src/ltx_trainer/validation_sampler.py`** - Inference for validation
|
||||
**`src/ltx_trainer/validation_runner.py`** - Conditioned validation sampling
|
||||
|
||||
Uses ltx-core components for denoising:
|
||||
|
||||
- `LTX2Scheduler` for sigma scheduling
|
||||
- `EulerDiffusionStep` for diffusion steps
|
||||
- `CFGGuider` for classifier-free guidance
|
||||
- `STGGuider` for spatio-temporal guidance
|
||||
- Manages the full validation lifecycle: embedding caching, media encoding, denoising, decoding
|
||||
- Supports all validation condition types: first_frame, prefix, suffix, spatial_crop, mask, reference, video_to_audio, audio_to_video
|
||||
- Handles frozen modality paths (sigma=0 for conditioning modality)
|
||||
- Builds conditioning items using ltx-core's `VideoConditionByLatentIndex`, `VideoConditionByReferenceLatent`, `VideoConditionByMask`
|
||||
- Optional side-by-side reference output for IC-LoRA validation
|
||||
|
||||
**`src/ltx_trainer/timestep_samplers.py`** - Flow matching timestep sampling
|
||||
|
||||
@@ -367,12 +445,18 @@ constructs the `GemmaTextEncoder` with quantized model, feature extractor, and e
|
||||
**`src/ltx_trainer/datasets.py`** - Dataset handling
|
||||
|
||||
- `PrecomputedDataset` loads pre-computed VAE latents and text embeddings
|
||||
- Supports video latents, audio latents, text embeddings, reference latents (for IC-LoRA)
|
||||
- Supports video latents, audio latents, text embeddings, reference video latents, reference audio latents, video masks, and audio masks
|
||||
- Handles legacy patchified format `[seq_len, C]` → automatically unpatchifies to `[C, F, H, W]`
|
||||
- `DummyDataset` for benchmarking and minimal testing
|
||||
|
||||
## Common Development Tasks
|
||||
|
||||
### Agent-Assisted Training
|
||||
|
||||
When a user asks to train, fine-tune, create a LoRA, or produce a custom LTX-2 model, use the repository skill at
|
||||
[`.claude/skills/train-model`](../../.claude/skills/train-model/SKILL.md). The skill is the orchestrator for dataset probing, mode selection, preprocessing,
|
||||
training launch, monitoring, and post-train validation; it treats `packages/ltx-trainer/docs/` as the source of truth.
|
||||
|
||||
### Adding a New Configuration Parameter
|
||||
|
||||
1. Add field to appropriate config class in `src/ltx_trainer/config.py`
|
||||
@@ -382,10 +466,16 @@ constructs the `GemmaTextEncoder` with quantized model, feature extractor, and e
|
||||
|
||||
### Implementing a New Training Strategy
|
||||
|
||||
The `FlexibleStrategy` now covers all use cases (T2V, T2A, I2V, V2V, A2A, AV2AV, inpainting, outpainting, extension, A2V, V2A, IC-LoRA) through
|
||||
configuration alone. A new strategy is only needed for fundamentally different training paradigms that cannot be
|
||||
expressed via `ModalityConfig` + `ConditionConfig` combinations.
|
||||
|
||||
If you do need a new strategy:
|
||||
|
||||
1. Create new file in `src/ltx_trainer/training_strategies/`
|
||||
2. Create config class inheriting `TrainingStrategyConfigBase`
|
||||
2. Create config class inheriting `TrainingStrategyConfigBase` and implement `get_data_sources()`
|
||||
3. Create strategy class inheriting `TrainingStrategy`
|
||||
4. Implement: `get_data_sources()`, `prepare_training_inputs()`, `compute_loss()`
|
||||
4. Implement: `prepare_training_inputs()`, `compute_loss()`
|
||||
5. Add to `__init__.py`: import, add to `TrainingStrategyConfig` union, update factory
|
||||
6. Add discriminator tag to config.py's `TrainingStrategyConfig`
|
||||
7. Create example config file in `configs/`
|
||||
@@ -449,8 +539,8 @@ video_embeds, audio_embeds, binary_mask = text_encoder.embeddings_processor.crea
|
||||
|
||||
- Validation errors: Check validators in `config.py`
|
||||
- Unknown fields: Config uses `extra="forbid"` — all fields must be defined
|
||||
- Strategy validation: IC-LoRA requires `reference_videos` in validation config
|
||||
- Video-to-video strategy requires `training_mode: "lora"`
|
||||
- FlexibleStrategy requires at least one modality with `is_generated: true`
|
||||
- Audio modality cannot use `first_frame` or `spatial_crop` conditions
|
||||
|
||||
**Precomputed Data:**
|
||||
|
||||
@@ -480,7 +570,7 @@ Width and height must be divisible by 32.
|
||||
### Platform Requirements
|
||||
|
||||
- Linux required (uses `triton` which is Linux-only)
|
||||
- CUDA GPU with 24GB+ VRAM recommended (80GB+ for full fine-tuning)
|
||||
- CUDA GPU with 32GB+ VRAM recommended
|
||||
|
||||
## Reference: ltx-core Key Components
|
||||
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
# LTX-2 Trainer
|
||||
|
||||
This package provides tools and scripts for training and fine-tuning
|
||||
Lightricks' **LTX-2** audio-video generation model. It enables LoRA training, full
|
||||
fine-tuning, and training of video-to-video transformations (IC-LoRA) on custom datasets.
|
||||
Lightricks' **LTX-2** audio-video generation model. It supports LoRA training, full
|
||||
fine-tuning, and a flexible conditioning framework covering text-to-video, text-to-audio, image-to-video,
|
||||
video extension, audio extension, video inpainting, audio inpainting, video outpainting, IC-LoRA for video, audio, and joint
|
||||
audio-video references, audio-to-video, and video-to-audio.
|
||||
|
||||
---
|
||||
|
||||
@@ -17,9 +19,16 @@ All detailed guides and technical documentation are in the [docs](./docs/) direc
|
||||
- [🚀 Training Guide](docs/training-guide.md)
|
||||
- [🧪 Inference Guide](../ltx-pipelines/README.md)
|
||||
- [🔧 Utility Scripts](docs/utility-scripts.md)
|
||||
- [🧩 Custom Training Strategies](docs/custom-training-strategies.md)
|
||||
- [📚 LTX-Core Documentation](../ltx-core/README.md)
|
||||
- [🛡️ Troubleshooting Guide](docs/troubleshooting.md)
|
||||
|
||||
### 🤖 Agent-Assisted Training
|
||||
|
||||
Use the [`train-model`](../../.claude/skills/train-model/SKILL.md) repository skill for an end-to-end guided run:
|
||||
it probes your data and hardware, chooses the matching training mode, prepares/preprocesses the dataset, launches
|
||||
training, and monitors the job while using the docs above as the source of truth.
|
||||
|
||||
---
|
||||
|
||||
## 🔧 Requirements
|
||||
@@ -28,7 +37,7 @@ All detailed guides and technical documentation are in the [docs](./docs/) direc
|
||||
- **Gemma Text Encoder** - Local Gemma model directory (required for LTX-2)
|
||||
- **Linux with CUDA** - CUDA 13+ recommended for optimal performance
|
||||
- **Nvidia GPU with 80GB+ VRAM** - Recommended for the standard config. For GPUs with 32GB VRAM (e.g., RTX 5090),
|
||||
use the [low VRAM config](configs/ltx2_av_lora_low_vram.yaml) which enables INT8 quantization and other
|
||||
use the [low VRAM config](configs/t2v_lora_low_vram.yaml) which enables INT8 quantization and other
|
||||
memory optimizations
|
||||
|
||||
---
|
||||
|
||||
@@ -1,315 +0,0 @@
|
||||
# =============================================================================
|
||||
# LTX-2 Audio-Video LoRA Training Configuration
|
||||
# =============================================================================
|
||||
#
|
||||
# This configuration is for training LoRA adapters on the LTX-2 model for
|
||||
# text-to-video generation. It supports both video-only and joint audio-video
|
||||
# training modes.
|
||||
#
|
||||
# Use this configuration when you want to:
|
||||
# - Fine-tune LTX-2 on your own video dataset
|
||||
# - Train with or without audio generation
|
||||
# - Create custom video generation styles or audiovisual concepts
|
||||
#
|
||||
# Dataset structure for text-to-video training:
|
||||
# preprocessed_data_root/
|
||||
# ├── latents/ # Video latents (VAE-encoded videos)
|
||||
# ├── conditions/ # Text embeddings for each video
|
||||
# └── audio_latents/ # Audio latents (only if with_audio: true)
|
||||
#
|
||||
# =============================================================================
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Model Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Specifies the base model to fine-tune and the training mode.
|
||||
model:
|
||||
# Path to the LTX-2 model checkpoint (.safetensors file)
|
||||
# This should be a local path to your downloaded model
|
||||
model_path: "path/to/ltx-2-model.safetensors"
|
||||
|
||||
# Path to the text encoder model directory
|
||||
# For LTX-2, this is typically the Gemma-based text encoder
|
||||
text_encoder_path: "path/to/gemma-text-encoder"
|
||||
|
||||
# Training mode: "lora" for efficient adapter training, "full" for full fine-tuning
|
||||
# LoRA is recommended for most use cases (faster, less memory, prevents overfitting)
|
||||
training_mode: "lora"
|
||||
|
||||
# Optional: Path to resume training from a checkpoint
|
||||
# Can be a checkpoint file (.safetensors) or directory (uses latest checkpoint)
|
||||
load_checkpoint: null
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# LoRA Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls the Low-Rank Adaptation parameters for efficient fine-tuning.
|
||||
lora:
|
||||
# Rank of the LoRA matrices (higher = more capacity but more parameters)
|
||||
# Typical values: 8, 16, 32, 64. Start with 32 for general fine-tuning.
|
||||
rank: 32
|
||||
|
||||
# Alpha scaling factor (usually set equal to rank)
|
||||
# The effective scaling is alpha/rank, so alpha=rank means scaling of 1.0
|
||||
alpha: 32
|
||||
|
||||
# Dropout probability for LoRA layers (0.0 = no dropout)
|
||||
# Can help with regularization if overfitting occurs
|
||||
dropout: 0.0
|
||||
|
||||
# Which transformer modules to apply LoRA to
|
||||
# The LTX-2 transformer has separate attention and FFN blocks for video and audio:
|
||||
#
|
||||
# VIDEO MODULES:
|
||||
# - attn1.to_k, attn1.to_q, attn1.to_v, attn1.to_out.0 (video self-attention)
|
||||
# - attn2.to_k, attn2.to_q, attn2.to_v, attn2.to_out.0 (video cross-attention to text)
|
||||
# - ff.net.0.proj, ff.net.2 (video feed-forward)
|
||||
#
|
||||
# AUDIO MODULES:
|
||||
# - audio_attn1.to_k, audio_attn1.to_q, audio_attn1.to_v, audio_attn1.to_out.0 (audio self-attention)
|
||||
# - audio_attn2.to_k, audio_attn2.to_q, audio_attn2.to_v, audio_attn2.to_out.0 (audio cross-attention to text)
|
||||
# - audio_ff.net.0.proj, audio_ff.net.2 (audio feed-forward)
|
||||
#
|
||||
# AUDIO-VIDEO CROSS-ATTENTION MODULES (for cross-modal interaction):
|
||||
# - audio_to_video_attn.to_k, audio_to_video_attn.to_q, audio_to_video_attn.to_v, audio_to_video_attn.to_out.0
|
||||
# (Q from video, K/V from audio - allows video to attend to audio features)
|
||||
# - video_to_audio_attn.to_k, video_to_audio_attn.to_q, video_to_audio_attn.to_v, video_to_audio_attn.to_out.0
|
||||
# (Q from audio, K/V from video - allows audio to attend to video features)
|
||||
#
|
||||
# Using short patterns like "to_k" matches ALL attention modules (video, audio, and cross-modal).
|
||||
# For audio-video training, this is the recommended approach.
|
||||
target_modules:
|
||||
# Attention layers (matches both video and audio branches)
|
||||
- "to_k"
|
||||
- "to_q"
|
||||
- "to_v"
|
||||
- "to_out.0"
|
||||
# Uncomment below to also train feed-forward layers (can increase the LoRA's capacity):
|
||||
# - "ff.net.0.proj"
|
||||
# - "ff.net.2"
|
||||
# - "audio_ff.net.0.proj"
|
||||
# - "audio_ff.net.2"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Training Strategy Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Defines the text-to-video training approach.
|
||||
training_strategy:
|
||||
# Strategy name: "text_to_video" for standard text-to-video training
|
||||
name: "text_to_video"
|
||||
|
||||
# Probability of conditioning on the first frame during training
|
||||
# Higher values train the model to perform better in image-to-video (I2V) mode,
|
||||
# where a clean first frame is provided and the model generates the rest of the video
|
||||
# Increase this value to train the model to perform better in image-to-video (I2V) mode
|
||||
first_frame_conditioning_p: 0.5
|
||||
|
||||
# Enable joint audio-video training
|
||||
# Set to true if your dataset includes audio and you want to train the audio branch
|
||||
with_audio: true
|
||||
|
||||
# Directory name (within preprocessed_data_root) containing audio latents
|
||||
# Only used when with_audio is true
|
||||
audio_latents_dir: "audio_latents"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Optimization Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls the training optimization parameters.
|
||||
optimization:
|
||||
# Learning rate for the optimizer
|
||||
# Typical range for LoRA: 1e-5 to 1e-4
|
||||
learning_rate: 1e-4
|
||||
|
||||
# Total number of training steps
|
||||
steps: 2000
|
||||
|
||||
# Batch size per GPU
|
||||
# Reduce if running out of memory
|
||||
batch_size: 1
|
||||
|
||||
# Number of gradient accumulation steps
|
||||
# Effective batch size = batch_size * gradient_accumulation_steps * num_gpus
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
# Maximum gradient norm for clipping (helps training stability)
|
||||
max_grad_norm: 1.0
|
||||
|
||||
# Optimizer type: "adamw" (standard) or "adamw8bit" (memory-efficient)
|
||||
optimizer_type: "adamw"
|
||||
|
||||
# Learning rate scheduler type
|
||||
# Options: "constant", "linear", "cosine", "cosine_with_restarts", "polynomial"
|
||||
scheduler_type: "linear"
|
||||
|
||||
# Additional scheduler parameters (depends on scheduler_type)
|
||||
scheduler_params: { }
|
||||
|
||||
# Enable gradient checkpointing to reduce memory usage
|
||||
# Recommended for training with limited GPU memory
|
||||
enable_gradient_checkpointing: true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Acceleration Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Hardware acceleration and memory optimization settings.
|
||||
acceleration:
|
||||
# Mixed precision training mode
|
||||
# Options: "no" (fp32), "fp16" (half precision), "bf16" (bfloat16, recommended)
|
||||
mixed_precision_mode: "bf16"
|
||||
|
||||
# Model quantization for reduced memory usage
|
||||
# Options: null (none), "int8-quanto", "int4-quanto", "int2-quanto", "fp8-quanto", "fp8uz-quanto"
|
||||
quantization: null
|
||||
|
||||
# Load text encoder in 8-bit precision to save memory
|
||||
# Useful when GPU memory is limited
|
||||
load_text_encoder_in_8bit: false
|
||||
|
||||
# Offload optimizer state to CPU during validation video sampling.
|
||||
# Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist
|
||||
# on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP.
|
||||
offload_optimizer_during_validation: false
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Data Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Specifies the training data location and loading parameters.
|
||||
data:
|
||||
# Root directory containing preprocessed training data
|
||||
# Should contain: latents/, conditions/, and optionally audio_latents/
|
||||
preprocessed_data_root: "/path/to/preprocessed/data"
|
||||
|
||||
# Number of worker processes for data loading
|
||||
# Used for parallel data loading to speed up data loading
|
||||
num_dataloader_workers: 2
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Validation Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls validation video generation during training.
|
||||
# NOTE: Validation sampling use simplified inference pipelines and prioritizes speed over
|
||||
# maximum quality. For production-quality inference, use `packages/ltx-pipelines`.
|
||||
validation:
|
||||
# Text prompts for validation video generation
|
||||
# Provide prompts representative of your training data
|
||||
# LTX-2 prefers longer, detailed prompts that describe both visual content and audio
|
||||
prompts:
|
||||
- "A woman with long brown hair sits at a wooden desk in a cozy home office, typing on a laptop while occasionally glancing at notes beside her. Soft natural light streams through a large window, casting warm shadows across the room. She pauses to take a sip from a ceramic mug, then continues working with focused concentration. The audio captures the gentle clicking of keyboard keys, the soft rustle of papers, and ambient room tone with occasional distant bird chirps from outside."
|
||||
- "A chef in a white uniform stands in a professional kitchen, carefully plating a gourmet dish with precise movements. Steam rises from freshly cooked vegetables as he arranges them with tweezers. The stainless steel surfaces gleam under bright overhead lights, and various pots simmer on the stove behind him. The audio features the sizzling of pans, the clinking of utensils against plates, and the ambient hum of kitchen ventilation."
|
||||
|
||||
# Negative prompt to avoid unwanted artifacts
|
||||
negative_prompt: "worst quality, inconsistent motion, blurry, jittery, distorted"
|
||||
|
||||
# Optional: First frame images for image-to-video validation
|
||||
# If provided, must have one image per prompt
|
||||
images: null
|
||||
|
||||
# Output video dimensions [width, height, frames]
|
||||
# Width and height must be divisible by 32
|
||||
# Frames must satisfy: frames % 8 == 1 (e.g., 1, 9, 17, 25, 33, 41, 49, 57, 65, 73, 81, 89, ...)
|
||||
video_dims: [ 576, 576, 89 ]
|
||||
|
||||
# Frame rate for generated videos
|
||||
frame_rate: 25.0
|
||||
|
||||
# Random seed for reproducible validation outputs
|
||||
seed: 42
|
||||
|
||||
# Number of denoising steps for validation inference
|
||||
# Higher values = better quality but slower generation
|
||||
inference_steps: 30
|
||||
|
||||
# Generate validation videos every N training steps
|
||||
# Set to null to disable validation during training
|
||||
interval: 100
|
||||
|
||||
# Classifier-free guidance scale
|
||||
# Higher values = stronger adherence to prompt but may introduce artifacts
|
||||
guidance_scale: 4.0
|
||||
|
||||
# STG (Spatio-Temporal Guidance) parameters for improved video quality
|
||||
# STG is combined with CFG for better temporal coherence
|
||||
stg_scale: 1.0 # Recommended: 1.0 (0.0 disables STG)
|
||||
stg_blocks: [29] # Recommended: single block 29
|
||||
stg_mode: "stg_av" # "stg_av" perturbs both audio and video, "stg_v" video only
|
||||
|
||||
# Whether to generate audio in validation samples
|
||||
# Independent of training_strategy.with_audio - you can generate audio
|
||||
# in validation even when not training the audio branch
|
||||
generate_audio: true
|
||||
|
||||
# Skip validation at the beginning of training (step 0)
|
||||
skip_initial_validation: false
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Checkpoint Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls model checkpoint saving during training.
|
||||
checkpoints:
|
||||
# Save a checkpoint every N steps
|
||||
# Set to null to disable intermediate checkpoints
|
||||
interval: 250
|
||||
|
||||
# Number of most recent checkpoints to keep
|
||||
# Set to -1 to keep all checkpoints
|
||||
keep_last_n: -1
|
||||
|
||||
# Precision to use when saving checkpoint weights
|
||||
# Options: "bfloat16" (default, smaller files) or "float32" (full precision)
|
||||
precision: "bfloat16"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Flow Matching Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Parameters for the flow matching training objective.
|
||||
flow_matching:
|
||||
# Timestep sampling mode
|
||||
# "shifted_logit_normal" is recommended for LTX-2 models
|
||||
timestep_sampling_mode: "shifted_logit_normal"
|
||||
|
||||
# Additional parameters for timestep sampling
|
||||
timestep_sampling_params: { }
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Hugging Face Hub Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Settings for uploading trained models to the Hugging Face Hub.
|
||||
hub:
|
||||
# Whether to push the trained model to the Hub
|
||||
push_to_hub: false
|
||||
|
||||
# Repository ID on Hugging Face Hub (e.g., "username/my-lora-model")
|
||||
# Required if push_to_hub is true
|
||||
hub_model_id: null
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Weights & Biases Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Settings for experiment tracking with W&B.
|
||||
wandb:
|
||||
# Enable W&B logging
|
||||
enabled: false
|
||||
|
||||
# W&B project name
|
||||
project: "ltx-2-trainer"
|
||||
|
||||
# W&B username or team (null uses default account)
|
||||
entity: null
|
||||
|
||||
# Tags to help organize runs
|
||||
tags: [ "ltx2", "lora" ]
|
||||
|
||||
# Log validation videos to W&B
|
||||
log_validation_videos: true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# General Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Global settings for the training run.
|
||||
|
||||
# Random seed for reproducibility
|
||||
seed: 42
|
||||
|
||||
# Directory to save outputs (checkpoints, validation videos, logs)
|
||||
output_dir: "outputs/ltx2_av_lora"
|
||||
@@ -1,327 +0,0 @@
|
||||
# =============================================================================
|
||||
# LTX-2 Audio-Video LoRA Training Configuration (Low VRAM)
|
||||
# =============================================================================
|
||||
#
|
||||
# This is a memory-optimized variant of the standard audio-video LoRA config.
|
||||
# It uses 8-bit optimizer, int8 quantization, and reduced LoRA rank to minimize
|
||||
# GPU memory usage while maintaining good training quality.
|
||||
#
|
||||
# Memory optimizations applied:
|
||||
# - 8-bit AdamW optimizer (reduces optimizer state memory by ~75%)
|
||||
# - INT8 model quantization (reduces model memory by ~50%)
|
||||
# - Lower LoRA rank (16 vs 32, reduces trainable parameters)
|
||||
# - Gradient checkpointing enabled
|
||||
#
|
||||
# Recommended for GPUs with 32GB VRAM (e.g., RTX 5090).
|
||||
#
|
||||
# Use this configuration when you want to:
|
||||
# - Fine-tune LTX-2 on your own video dataset
|
||||
# - Train with or without audio generation
|
||||
# - Create custom video generation styles or audiovisual concepts
|
||||
#
|
||||
# Dataset structure for text-to-video training:
|
||||
# preprocessed_data_root/
|
||||
# ├── latents/ # Video latents (VAE-encoded videos)
|
||||
# ├── conditions/ # Text embeddings for each video
|
||||
# └── audio_latents/ # Audio latents (only if with_audio: true)
|
||||
#
|
||||
# =============================================================================
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Model Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Specifies the base model to fine-tune and the training mode.
|
||||
model:
|
||||
# Path to the LTX-2 model checkpoint (.safetensors file)
|
||||
# This should be a local path to your downloaded model
|
||||
model_path: "path/to/ltx-2-model.safetensors"
|
||||
|
||||
# Path to the text encoder model directory
|
||||
# For LTX-2, this is typically the Gemma-based text encoder
|
||||
text_encoder_path: "path/to/gemma-text-encoder"
|
||||
|
||||
# Training mode: "lora" for efficient adapter training, "full" for full fine-tuning
|
||||
# LoRA is recommended for most use cases (faster, less memory, prevents overfitting)
|
||||
training_mode: "lora"
|
||||
|
||||
# Optional: Path to resume training from a checkpoint
|
||||
# Can be a checkpoint file (.safetensors) or directory (uses latest checkpoint)
|
||||
load_checkpoint: null
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# LoRA Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls the Low-Rank Adaptation parameters for efficient fine-tuning.
|
||||
# Using a lower rank (16) to reduce trainable parameters and memory usage.
|
||||
# This still provides good capacity for many fine-tuning tasks.
|
||||
lora:
|
||||
# Rank of the LoRA matrices (higher = more capacity but more parameters)
|
||||
# Typical values: 8, 16, 32, 64. Using 16 for low VRAM configuration.
|
||||
rank: 16
|
||||
|
||||
# Alpha scaling factor (usually set equal to rank)
|
||||
# The effective scaling is alpha/rank, so alpha=rank means scaling of 1.0
|
||||
alpha: 16
|
||||
|
||||
# Dropout probability for LoRA layers (0.0 = no dropout)
|
||||
# Can help with regularization if overfitting occurs
|
||||
dropout: 0.0
|
||||
|
||||
# Which transformer modules to apply LoRA to
|
||||
# The LTX-2 transformer has separate attention and FFN blocks for video and audio:
|
||||
#
|
||||
# VIDEO MODULES:
|
||||
# - attn1.to_k, attn1.to_q, attn1.to_v, attn1.to_out.0 (video self-attention)
|
||||
# - attn2.to_k, attn2.to_q, attn2.to_v, attn2.to_out.0 (video cross-attention to text)
|
||||
# - ff.net.0.proj, ff.net.2 (video feed-forward)
|
||||
#
|
||||
# AUDIO MODULES:
|
||||
# - audio_attn1.to_k, audio_attn1.to_q, audio_attn1.to_v, audio_attn1.to_out.0 (audio self-attention)
|
||||
# - audio_attn2.to_k, audio_attn2.to_q, audio_attn2.to_v, audio_attn2.to_out.0 (audio cross-attention to text)
|
||||
# - audio_ff.net.0.proj, audio_ff.net.2 (audio feed-forward)
|
||||
#
|
||||
# AUDIO-VIDEO CROSS-ATTENTION MODULES (for cross-modal interaction):
|
||||
# - audio_to_video_attn.to_k, audio_to_video_attn.to_q, audio_to_video_attn.to_v, audio_to_video_attn.to_out.0
|
||||
# (Q from video, K/V from audio - allows video to attend to audio features)
|
||||
# - video_to_audio_attn.to_k, video_to_audio_attn.to_q, video_to_audio_attn.to_v, video_to_audio_attn.to_out.0
|
||||
# (Q from audio, K/V from video - allows audio to attend to video features)
|
||||
#
|
||||
# Using short patterns like "to_k" matches ALL attention modules (video, audio, and cross-modal).
|
||||
# For audio-video training, this is the recommended approach.
|
||||
target_modules:
|
||||
# Attention layers (matches both video and audio branches)
|
||||
- "to_k"
|
||||
- "to_q"
|
||||
- "to_v"
|
||||
- "to_out.0"
|
||||
# Uncomment below to also train feed-forward layers (can increase the LoRA's capacity):
|
||||
# - "ff.net.0.proj"
|
||||
# - "ff.net.2"
|
||||
# - "audio_ff.net.0.proj"
|
||||
# - "audio_ff.net.2"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Training Strategy Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Defines the text-to-video training approach.
|
||||
training_strategy:
|
||||
# Strategy name: "text_to_video" for standard text-to-video training
|
||||
name: "text_to_video"
|
||||
|
||||
# Probability of conditioning on the first frame during training
|
||||
# Higher values train the model to perform better in image-to-video (I2V) mode,
|
||||
# where a clean first frame is provided and the model generates the rest of the video
|
||||
# Increase this value to train the model to perform better in image-to-video (I2V) mode
|
||||
first_frame_conditioning_p: 0.5
|
||||
|
||||
# Enable joint audio-video training
|
||||
# Set to true if your dataset includes audio and you want to train the audio branch
|
||||
with_audio: true
|
||||
|
||||
# Directory name (within preprocessed_data_root) containing audio latents
|
||||
# Only used when with_audio is true
|
||||
audio_latents_dir: "audio_latents"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Optimization Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls the training optimization parameters.
|
||||
optimization:
|
||||
# Learning rate for the optimizer
|
||||
# Typical range for LoRA: 1e-5 to 1e-4
|
||||
learning_rate: 1e-4
|
||||
|
||||
# Total number of training steps
|
||||
steps: 2000
|
||||
|
||||
# Batch size per GPU
|
||||
# Reduce if running out of memory
|
||||
batch_size: 1
|
||||
|
||||
# Number of gradient accumulation steps
|
||||
# Effective batch size = batch_size * gradient_accumulation_steps * num_gpus
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
# Maximum gradient norm for clipping (helps training stability)
|
||||
max_grad_norm: 1.0
|
||||
|
||||
# Optimizer type: "adamw" (standard) or "adamw8bit" (memory-efficient)
|
||||
# Using 8-bit AdamW to reduce optimizer state memory by ~75%
|
||||
optimizer_type: "adamw8bit"
|
||||
|
||||
# Learning rate scheduler type
|
||||
# Options: "constant", "linear", "cosine", "cosine_with_restarts", "polynomial"
|
||||
scheduler_type: "linear"
|
||||
|
||||
# Additional scheduler parameters (depends on scheduler_type)
|
||||
scheduler_params: { }
|
||||
|
||||
# Enable gradient checkpointing to reduce memory usage
|
||||
# Recommended for training with limited GPU memory
|
||||
enable_gradient_checkpointing: true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Acceleration Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Hardware acceleration and memory optimization settings.
|
||||
acceleration:
|
||||
# Mixed precision training mode
|
||||
# Options: "no" (fp32), "fp16" (half precision), "bf16" (bfloat16, recommended)
|
||||
mixed_precision_mode: "bf16"
|
||||
|
||||
# Model quantization for reduced memory usage
|
||||
# Options: null (none), "int8-quanto", "int4-quanto", "int2-quanto", "fp8-quanto", "fp8uz-quanto"
|
||||
# Using INT8 quantization to reduce base model memory consumption by ~50%
|
||||
quantization: "int8-quanto"
|
||||
|
||||
# Load text encoder in 8-bit precision to save memory
|
||||
# Useful when GPU memory is limited
|
||||
load_text_encoder_in_8bit: true
|
||||
|
||||
# Offload optimizer state to CPU during validation video sampling.
|
||||
# Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist
|
||||
# on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP.
|
||||
offload_optimizer_during_validation: true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Data Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Specifies the training data location and loading parameters.
|
||||
data:
|
||||
# Root directory containing preprocessed training data
|
||||
# Should contain: latents/, conditions/, and optionally audio_latents/
|
||||
preprocessed_data_root: "/path/to/preprocessed/data"
|
||||
|
||||
# Number of worker processes for data loading
|
||||
# Used for parallel data loading to speed up data loading
|
||||
num_dataloader_workers: 2
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Validation Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls validation video generation during training.
|
||||
# NOTE: Validation sampling use simplified inference pipelines and prioritizes speed over
|
||||
# maximum quality. For production-quality inference, use `packages/ltx-pipelines`.
|
||||
validation:
|
||||
# Text prompts for validation video generation
|
||||
# Provide prompts representative of your training data
|
||||
# LTX-2 prefers longer, detailed prompts that describe both visual content and audio
|
||||
prompts:
|
||||
- "A woman with long brown hair sits at a wooden desk in a cozy home office, typing on a laptop while occasionally glancing at notes beside her. Soft natural light streams through a large window, casting warm shadows across the room. She pauses to take a sip from a ceramic mug, then continues working with focused concentration. The audio captures the gentle clicking of keyboard keys, the soft rustle of papers, and ambient room tone with occasional distant bird chirps from outside."
|
||||
- "A chef in a white uniform stands in a professional kitchen, carefully plating a gourmet dish with precise movements. Steam rises from freshly cooked vegetables as he arranges them with tweezers. The stainless steel surfaces gleam under bright overhead lights, and various pots simmer on the stove behind him. The audio features the sizzling of pans, the clinking of utensils against plates, and the ambient hum of kitchen ventilation."
|
||||
|
||||
# Negative prompt to avoid unwanted artifacts
|
||||
negative_prompt: "worst quality, inconsistent motion, blurry, jittery, distorted"
|
||||
|
||||
# Optional: First frame images for image-to-video validation
|
||||
# If provided, must have one image per prompt
|
||||
images: null
|
||||
|
||||
# Output video dimensions [width, height, frames]
|
||||
# Width and height must be divisible by 32
|
||||
# Frames must satisfy: frames % 8 == 1 (e.g., 1, 9, 17, 25, 33, 41, 49, 57, 65, 73, 81, 89, ...)
|
||||
video_dims: [ 576, 576, 49 ]
|
||||
|
||||
# Frame rate for generated videos
|
||||
frame_rate: 25.0
|
||||
|
||||
# Random seed for reproducible validation outputs
|
||||
seed: 42
|
||||
|
||||
# Number of denoising steps for validation inference
|
||||
# Higher values = better quality but slower generation
|
||||
inference_steps: 30
|
||||
|
||||
# Generate validation videos every N training steps
|
||||
# Set to null to disable validation during training
|
||||
interval: 100
|
||||
|
||||
# Classifier-free guidance scale
|
||||
# Higher values = stronger adherence to prompt but may introduce artifacts
|
||||
guidance_scale: 4.0
|
||||
|
||||
# STG (Spatio-Temporal Guidance) parameters for improved video quality
|
||||
# STG is combined with CFG for better temporal coherence
|
||||
stg_scale: 1.0 # Recommended: 1.0 (0.0 disables STG)
|
||||
stg_blocks: [ 29 ] # Recommended: single block 29
|
||||
stg_mode: "stg_av" # "stg_av" perturbs both audio and video, "stg_v" video only
|
||||
|
||||
# Whether to generate audio in validation samples
|
||||
# Independent of training_strategy.with_audio - you can generate audio
|
||||
# in validation even when not training the audio branch
|
||||
generate_audio: true
|
||||
|
||||
# Skip validation at the beginning of training (step 0)
|
||||
skip_initial_validation: false
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Checkpoint Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls model checkpoint saving during training.
|
||||
checkpoints:
|
||||
# Save a checkpoint every N steps
|
||||
# Set to null to disable intermediate checkpoints
|
||||
interval: 250
|
||||
|
||||
# Number of most recent checkpoints to keep
|
||||
# Set to -1 to keep all checkpoints
|
||||
keep_last_n: -1
|
||||
|
||||
# Precision to use when saving checkpoint weights
|
||||
# Options: "bfloat16" (default, smaller files) or "float32" (full precision)
|
||||
precision: "bfloat16"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Flow Matching Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Parameters for the flow matching training objective.
|
||||
flow_matching:
|
||||
# Timestep sampling mode
|
||||
# "shifted_logit_normal" is recommended for LTX-2 models
|
||||
timestep_sampling_mode: "shifted_logit_normal"
|
||||
|
||||
# Additional parameters for timestep sampling
|
||||
timestep_sampling_params: { }
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Hugging Face Hub Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Settings for uploading trained models to the Hugging Face Hub.
|
||||
hub:
|
||||
# Whether to push the trained model to the Hub
|
||||
push_to_hub: false
|
||||
|
||||
# Repository ID on Hugging Face Hub (e.g., "username/my-lora-model")
|
||||
# Required if push_to_hub is true
|
||||
hub_model_id: null
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Weights & Biases Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Settings for experiment tracking with W&B.
|
||||
wandb:
|
||||
# Enable W&B logging
|
||||
enabled: false
|
||||
|
||||
# W&B project name
|
||||
project: "ltx-2-trainer"
|
||||
|
||||
# W&B username or team (null uses default account)
|
||||
entity: null
|
||||
|
||||
# Tags to help organize runs
|
||||
tags: [ "ltx2", "lora" ]
|
||||
|
||||
# Log validation videos to W&B
|
||||
log_validation_videos: true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# General Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Global settings for the training run.
|
||||
|
||||
# Random seed for reproducibility
|
||||
seed: 42
|
||||
|
||||
# Directory to save outputs (checkpoints, validation videos, logs)
|
||||
output_dir: "outputs/ltx2_av_lora"
|
||||
@@ -1,331 +0,0 @@
|
||||
# =============================================================================
|
||||
# LTX-2 Video-to-Video (IC-LoRA) Training Configuration
|
||||
# =============================================================================
|
||||
#
|
||||
# This configuration is for training In-Context LoRA (IC-LoRA) adapters that
|
||||
# enable video-to-video transformations. IC-LoRA learns to apply visual
|
||||
# transformations (e.g., depth-to-video, pose control, style transfer, etc.)
|
||||
# by conditioning on reference videos.
|
||||
#
|
||||
# Key differences from text-to-video LoRA:
|
||||
# - Uses reference videos as conditioning input alongside text prompts
|
||||
# - Requires preprocessed reference latents in addition to target latents
|
||||
# - Validation requires reference videos to demonstrate the transformation
|
||||
#
|
||||
# Dataset structure for IC-LoRA training:
|
||||
# preprocessed_data_root/
|
||||
# ├── latents/ # Target video latents (what the model learns to generate)
|
||||
# ├── conditions/ # Text embeddings for each video
|
||||
# └── reference_latents/ # Reference video latents (conditioning input)
|
||||
#
|
||||
# =============================================================================
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Model Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Specifies the base model to fine-tune and the training mode.
|
||||
model:
|
||||
# Path to the LTX-2 model checkpoint (.safetensors file)
|
||||
# This should be a local path to your downloaded model
|
||||
model_path: "path/to/ltx-2-model.safetensors"
|
||||
|
||||
# Path to the text encoder model directory
|
||||
# For LTX-2, this is typically the Gemma-based text encoder
|
||||
text_encoder_path: "path/to/gemma-text-encoder"
|
||||
|
||||
# Training mode: "lora" for efficient adapter training, "full" for full fine-tuning
|
||||
# Note: video_to_video strategy requires "lora" mode
|
||||
training_mode: "lora"
|
||||
|
||||
# Optional: Path to resume training from a checkpoint
|
||||
# Can be a checkpoint file (.safetensors) or directory (uses latest checkpoint)
|
||||
load_checkpoint: null
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# LoRA Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls the Low-Rank Adaptation parameters for efficient fine-tuning.
|
||||
lora:
|
||||
# Rank of the LoRA matrices (higher = more capacity but more parameters)
|
||||
# Typical values: 8, 16, 32, 64. Start with 16-32 for IC-LoRA.
|
||||
rank: 32
|
||||
|
||||
# Alpha scaling factor (usually set equal to rank)
|
||||
# The effective scaling is alpha/rank, so alpha=rank means scaling of 1.0
|
||||
alpha: 32
|
||||
|
||||
# Dropout probability for LoRA layers (0.0 = no dropout)
|
||||
# Can help with regularization if overfitting occurs
|
||||
dropout: 0.0
|
||||
|
||||
# Which transformer modules to apply LoRA to
|
||||
# The LTX-2 transformer has separate attention and FFN blocks for video and audio:
|
||||
#
|
||||
# VIDEO MODULES:
|
||||
# - attn1.to_k, attn1.to_q, attn1.to_v, attn1.to_out.0 (video self-attention)
|
||||
# - attn2.to_k, attn2.to_q, attn2.to_v, attn2.to_out.0 (video cross-attention to text)
|
||||
# - ff.net.0.proj, ff.net.2 (video feed-forward)
|
||||
#
|
||||
# AUDIO MODULES (not used for video-only IC-LoRA):
|
||||
# - audio_attn1.to_k, audio_attn1.to_q, audio_attn1.to_v, audio_attn1.to_out.0 (audio self-attention)
|
||||
# - audio_attn2.to_k, audio_attn2.to_q, audio_attn2.to_v, audio_attn2.to_out.0 (audio cross-attention to text)
|
||||
# - audio_ff.net.0.proj, audio_ff.net.2 (audio feed-forward)
|
||||
#
|
||||
# AUDIO-VIDEO CROSS-ATTENTION MODULES (for cross-modal interaction, not used for video-only IC-LoRA):
|
||||
# - audio_to_video_attn.to_k, audio_to_video_attn.to_q, audio_to_video_attn.to_v, audio_to_video_attn.to_out.0
|
||||
# (Q from video, K/V from audio - allows video to attend to audio features)
|
||||
# - video_to_audio_attn.to_k, video_to_audio_attn.to_q, video_to_audio_attn.to_v, video_to_audio_attn.to_out.0
|
||||
# (Q from audio, K/V from video - allows audio to attend to video features)
|
||||
#
|
||||
# For IC-LoRA (video-only), we explicitly target video modules.
|
||||
# Including FFN layers often improves transformation quality.
|
||||
target_modules:
|
||||
# Video self-attention
|
||||
- "attn1.to_k"
|
||||
- "attn1.to_q"
|
||||
- "attn1.to_v"
|
||||
- "attn1.to_out.0"
|
||||
# Video cross-attention
|
||||
- "attn2.to_k"
|
||||
- "attn2.to_q"
|
||||
- "attn2.to_v"
|
||||
- "attn2.to_out.0"
|
||||
# Video feed-forward (often improves transformation quality)
|
||||
- "ff.net.0.proj"
|
||||
- "ff.net.2"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Training Strategy Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Defines the video-to-video (IC-LoRA) training approach.
|
||||
training_strategy:
|
||||
# Strategy name: "video_to_video" for IC-LoRA training
|
||||
name: "video_to_video"
|
||||
|
||||
# Probability of conditioning on the first frame during training
|
||||
# Higher values train the model to perform better in image-to-video (I2V) mode,
|
||||
# where a clean first frame is provided and the model generates the rest of the video
|
||||
# Increase this value to train the model to perform better in image-to-video (I2V) mode
|
||||
first_frame_conditioning_p: 0.2
|
||||
|
||||
# Directory name (within preprocessed_data_root) containing reference video latents
|
||||
# These are the conditioning inputs that guide the transformation
|
||||
reference_latents_dir: "reference_latents"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Optimization Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls the training optimization parameters.
|
||||
optimization:
|
||||
# Learning rate for the optimizer
|
||||
# Typical range for LoRA: 1e-5 to 1e-4
|
||||
learning_rate: 2e-4
|
||||
|
||||
# Total number of training steps
|
||||
steps: 3000
|
||||
|
||||
# Batch size per GPU
|
||||
# Reduce if running out of memory
|
||||
batch_size: 1
|
||||
|
||||
# Number of gradient accumulation steps
|
||||
# Effective batch size = batch_size * gradient_accumulation_steps * num_gpus
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
# Maximum gradient norm for clipping (helps training stability)
|
||||
max_grad_norm: 1.0
|
||||
|
||||
# Optimizer type: "adamw" (standard) or "adamw8bit" (memory-efficient)
|
||||
optimizer_type: "adamw"
|
||||
|
||||
# Learning rate scheduler type
|
||||
# Options: "constant", "linear", "cosine", "cosine_with_restarts", "polynomial"
|
||||
scheduler_type: "linear"
|
||||
|
||||
# Additional scheduler parameters (depends on scheduler_type)
|
||||
scheduler_params: { }
|
||||
|
||||
# Enable gradient checkpointing to reduce memory usage
|
||||
# Recommended for training with limited GPU memory
|
||||
enable_gradient_checkpointing: true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Acceleration Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Hardware acceleration and memory optimization settings.
|
||||
acceleration:
|
||||
# Mixed precision training mode
|
||||
# Options: "no" (fp32), "fp16" (half precision), "bf16" (bfloat16, recommended)
|
||||
mixed_precision_mode: "bf16"
|
||||
|
||||
# Model quantization for reduced memory usage
|
||||
# Options: null (none), "int8-quanto", "int4-quanto", "int2-quanto", "fp8-quanto", "fp8uz-quanto"
|
||||
quantization: null
|
||||
|
||||
# Load text encoder in 8-bit precision to save memory
|
||||
# Useful when GPU memory is limited
|
||||
load_text_encoder_in_8bit: false
|
||||
|
||||
# Offload optimizer state to CPU during validation video sampling.
|
||||
# Helps avoid OOM when VAE decoder + transformer + optimizer state can't coexist
|
||||
# on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP.
|
||||
offload_optimizer_during_validation: false
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Data Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Specifies the training data location and loading parameters.
|
||||
data:
|
||||
# Root directory containing preprocessed training data
|
||||
# Should contain: latents/, conditions/, and reference_latents/ subdirectories
|
||||
preprocessed_data_root: "/path/to/preprocessed/data"
|
||||
|
||||
# Number of worker processes for data loading
|
||||
# Used for parallel data loading to speed up data loading
|
||||
num_dataloader_workers: 2
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Validation Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls validation video generation during training.
|
||||
# NOTE: Validation sampling use simplified inference pipelines and prioritizes speed over
|
||||
# maximum quality. For production-quality inference, use `packages/ltx-pipelines`.
|
||||
validation:
|
||||
# Text prompts for validation video generation
|
||||
# Provide prompts representative of your training data
|
||||
# LTX-2 prefers longer, detailed prompts that describe both visual content and audio
|
||||
prompts:
|
||||
- "A man in a casual blue jacket walks along a winding path through a lush green park on a bright sunny afternoon. Tall oak trees line the pathway, their leaves rustling gently in the breeze. Dappled sunlight creates shifting patterns on the ground as he strolls at a relaxed pace, occasionally looking up at the scenery around him. The audio captures footsteps on gravel, birds singing in the trees, distant children playing, and the soft whisper of wind through the foliage."
|
||||
- "A fluffy orange tabby cat sits perfectly still on a wooden windowsill, its green eyes intently tracking small birds hopping on a branch just outside the glass. The cat's ears twitch and rotate, following every movement. Warm afternoon light illuminates its fur, creating a soft golden glow. Behind the cat, a cozy living room with a bookshelf and houseplants is visible. The audio features gentle purring, occasional soft meows, muffled bird chirps through the window, and quiet ambient room sounds."
|
||||
|
||||
# Reference videos for validation (REQUIRED for video_to_video strategy)
|
||||
# Must provide one reference video per prompt
|
||||
# These are the conditioning inputs for generating validation outputs
|
||||
reference_videos:
|
||||
- "/path/to/reference_video_1.mp4"
|
||||
- "/path/to/reference_video_2.mp4"
|
||||
|
||||
# Downscale factor for reference videos (for efficient IC-LoRA training)
|
||||
# When > 1, reference videos are processed at 1/n resolution
|
||||
# Must match the --reference-downscale-factor used during dataset preprocessing
|
||||
# Examples: 1 = same resolution, 2 = half resolution (384x384 ref for 768x768 target)
|
||||
reference_downscale_factor: 1
|
||||
|
||||
# Negative prompt to avoid unwanted artifacts
|
||||
negative_prompt: "worst quality, inconsistent motion, blurry, jittery, distorted"
|
||||
|
||||
# Optional: First frame images for additional conditioning
|
||||
# If provided, must have one image per prompt
|
||||
images: null
|
||||
|
||||
# Output video dimensions [width, height, frames]
|
||||
# Width and height must be divisible by 32
|
||||
# Frames must satisfy: frames % 8 == 1 (e.g., 1, 9, 17, 25, 33, 41, 49, 57, 65, 73, 81, 89, ...)
|
||||
video_dims: [ 512, 512, 81 ]
|
||||
|
||||
# Frame rate for generated videos
|
||||
frame_rate: 25.0
|
||||
|
||||
# Random seed for reproducible validation outputs
|
||||
seed: 42
|
||||
|
||||
# Number of denoising steps for validation inference
|
||||
# Higher values = better quality but slower generation
|
||||
inference_steps: 30
|
||||
|
||||
# Generate validation videos every N training steps
|
||||
# Set to null to disable validation during training
|
||||
interval: 100
|
||||
|
||||
# Classifier-free guidance scale
|
||||
# Higher values = stronger adherence to prompt but may introduce artifacts
|
||||
guidance_scale: 4.0
|
||||
|
||||
# STG (Spatio-Temporal Guidance) parameters for improved video quality
|
||||
# STG is combined with CFG for better temporal coherence
|
||||
stg_scale: 1.0 # Recommended: 1.0 (0.0 disables STG)
|
||||
stg_blocks: [29] # Recommended: single block 29
|
||||
stg_mode: "stg_v" # "stg_v" for video-only (no audio training)
|
||||
|
||||
# Whether to generate audio in validation samples
|
||||
# Can be enabled even when not training the audio branch
|
||||
generate_audio: false
|
||||
|
||||
# Skip validation at the beginning of training (step 0)
|
||||
skip_initial_validation: false
|
||||
|
||||
# Concatenate reference video side-by-side with generated output
|
||||
# Useful for visually comparing the transformation quality
|
||||
include_reference_in_output: true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Checkpoint Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Controls model checkpoint saving during training.
|
||||
checkpoints:
|
||||
# Save a checkpoint every N steps
|
||||
# Set to null to disable intermediate checkpoints
|
||||
interval: 250
|
||||
|
||||
# Number of most recent checkpoints to keep
|
||||
# Set to -1 to keep all checkpoints
|
||||
keep_last_n: 3
|
||||
|
||||
# Precision to use when saving checkpoint weights
|
||||
# Options: "bfloat16" (default, smaller files) or "float32" (full precision)
|
||||
precision: "bfloat16"
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Flow Matching Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Parameters for the flow matching training objective.
|
||||
flow_matching:
|
||||
# Timestep sampling mode
|
||||
# "shifted_logit_normal" is recommended for LTX-2 models
|
||||
timestep_sampling_mode: "shifted_logit_normal"
|
||||
|
||||
# Additional parameters for timestep sampling
|
||||
timestep_sampling_params: { }
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Hugging Face Hub Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Settings for uploading trained models to the Hugging Face Hub.
|
||||
hub:
|
||||
# Whether to push the trained model to the Hub
|
||||
push_to_hub: false
|
||||
|
||||
# Repository ID on Hugging Face Hub (e.g., "username/my-ic-lora-model")
|
||||
# Required if push_to_hub is true
|
||||
hub_model_id: null
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# Weights & Biases Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Settings for experiment tracking with W&B.
|
||||
wandb:
|
||||
# Enable W&B logging
|
||||
enabled: false
|
||||
|
||||
# W&B project name
|
||||
project: "ltx-2-trainer"
|
||||
|
||||
# W&B username or team (null uses default account)
|
||||
entity: null
|
||||
|
||||
# Tags to help organize runs
|
||||
tags: [ "ltx2", "ic-lora", "video-to-video" ]
|
||||
|
||||
# Log validation videos to W&B
|
||||
log_validation_videos: true
|
||||
|
||||
# -----------------------------------------------------------------------------
|
||||
# General Configuration
|
||||
# -----------------------------------------------------------------------------
|
||||
# Global settings for the training run.
|
||||
|
||||
# Random seed for reproducibility
|
||||
seed: 42
|
||||
|
||||
# Directory to save outputs (checkpoints, validation videos, logs)
|
||||
output_dir: "outputs/ltx2_v2v_ic_lora"
|
||||
@@ -10,7 +10,7 @@ sub-configurations:
|
||||
|
||||
- **ModelConfig**: Base model and training mode settings
|
||||
- **LoraConfig**: LoRA training parameters
|
||||
- **TrainingStrategyConfig**: Training strategy settings (text-to-video or video-to-video)
|
||||
- **TrainingStrategyConfig**: Training strategy settings (flexible conditioning framework)
|
||||
- **OptimizationConfig**: Learning rate, batch sizes, and scheduler settings
|
||||
- **AccelerationConfig**: Mixed precision and quantization settings
|
||||
- **DataConfig**: Data loading parameters
|
||||
@@ -24,13 +24,29 @@ sub-configurations:
|
||||
|
||||
Check out our example configurations in the `configs` directory:
|
||||
|
||||
- 📄 [Audio-Video LoRA Training](../configs/ltx2_av_lora.yaml) - Joint audio-video generation training
|
||||
- 📄 [Audio-Video LoRA Training (Low VRAM)](../configs/ltx2_av_lora_low_vram.yaml) - Memory-optimized config for 32GB
|
||||
GPUs (uses 8-bit optimizer, INT8 quantization, and reduced LoRA rank)
|
||||
- 📄 [IC-LoRA Training](../configs/ltx2_v2v_ic_lora.yaml) - Video-to-video transformation training
|
||||
- 📄 [Text-to-Video LoRA](../configs/t2v_lora.yaml) - Text-to-video LoRA training
|
||||
- 📄 [Image-to-Video LoRA](../configs/i2v_lora.yaml) - Image-to-video LoRA training
|
||||
- 📄 [IC-LoRA Video-to-Video](../configs/v2v_ic_lora.yaml) - IC-LoRA video-to-video training
|
||||
- 📄 [Audio-to-Video LoRA](../configs/a2v_lora.yaml) - Audio-to-video LoRA training
|
||||
- 📄 [Video-to-Audio LoRA](../configs/v2a_lora.yaml) - Video-to-audio (Foley) LoRA training
|
||||
- 📄 [Video Extension LoRA](../configs/video_extend_lora.yaml) - Video extension (forward) LoRA training
|
||||
- 📄 [Video Suffix LoRA](../configs/video_suffix_lora.yaml) - Video extension (backward) LoRA training
|
||||
- 📄 [Video Inpainting LoRA](../configs/video_inpainting_lora.yaml) - Video inpainting LoRA training
|
||||
- 📄 [Video Outpainting LoRA](../configs/video_outpainting_lora.yaml) - Video outpainting (spatial crop) LoRA training
|
||||
- 📄 [Text-to-Audio LoRA](../configs/t2a_lora.yaml) - Text-to-audio LoRA training
|
||||
- 📄 [Audio Extension LoRA](../configs/audio_extend_lora.yaml) - Audio extension (forward) LoRA training
|
||||
- 📄 [Audio Suffix LoRA](../configs/audio_suffix_lora.yaml) - Audio extension (backward) LoRA training
|
||||
- 📄 [Audio Inpainting LoRA](../configs/audio_inpainting_lora.yaml) - Audio inpainting LoRA training
|
||||
- 📄 [Audio-to-Audio IC-LoRA](../configs/a2a_ic_lora.yaml) - Audio IC-LoRA transformation training
|
||||
- 📄 [AV2AV IC-LoRA](../configs/av2av_ic_lora.yaml) - Audio+video IC-LoRA transformation training
|
||||
- 📄 [T2V LoRA (Low VRAM)](../configs/t2v_lora_low_vram.yaml) - Memory-optimized config for 32GB GPUs
|
||||
|
||||
## ⚙️ Configuration Sections
|
||||
|
||||
> [!NOTE]
|
||||
> The YAML snippets below show **recommended starting values**, not necessarily the code defaults.
|
||||
> Fields you omit from your config file will use the code defaults from [`config.py`](../src/ltx_trainer/config.py).
|
||||
|
||||
### ModelConfig
|
||||
|
||||
Controls the base model and training mode settings.
|
||||
@@ -149,37 +165,60 @@ target_modules:
|
||||
|
||||
### TrainingStrategyConfig
|
||||
|
||||
Configures the training strategy. The trainer includes two built-in strategies described below.
|
||||
For custom use cases, see [Implementing Custom Training Strategies](custom-training-strategies.md).
|
||||
Configures the training strategy. The recommended strategy is `"flexible"`, which supports all conditioning scenarios through configuration.
|
||||
|
||||
#### Text-to-Video Strategy
|
||||
#### Flexible Strategy
|
||||
|
||||
The flexible strategy provides a unified conditioning framework. Each modality (video, audio) is configured
|
||||
independently with its own latents directory, generation flag, and list of conditions.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "text_to_video"
|
||||
first_frame_conditioning_p: 0.1 # Probability of first-frame conditioning
|
||||
with_audio: false # Enable joint audio-video training
|
||||
audio_latents_dir: "audio_latents" # Directory for audio latents (when with_audio: true)
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true # Video is denoised during training
|
||||
latents_dir: "latents" # Directory containing precomputed video latents
|
||||
conditions:
|
||||
- type: first_frame # Use first frame as conditioning
|
||||
probability: 0.5 # Apply this condition 50% of the time
|
||||
audio:
|
||||
is_generated: true # Audio is denoised during training
|
||||
latents_dir: "audio_latents" # Directory containing precomputed audio latents
|
||||
conditions: [] # No additional audio conditions (text-only)
|
||||
```
|
||||
|
||||
#### Video-to-Video Strategy (IC-LoRA)
|
||||
**ModalityConfig parameters:**
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "video_to_video"
|
||||
first_frame_conditioning_p: 0.1
|
||||
reference_latents_dir: "reference_latents" # Directory for reference video latents
|
||||
```
|
||||
| Parameter | Description |
|
||||
|----------------|------------------------------------------------------------------------------------------------------------------|
|
||||
| `is_generated` | `true` = modality is denoised (contributes to loss). `false` = frozen conditioning (sigma=0, no loss). |
|
||||
| `latents_dir` | Directory name within `preprocessed_data_root` containing precomputed latents for this modality. |
|
||||
| `conditions` | List of conditioning configs applied during training (see condition types below). Text conditioning is implicit. |
|
||||
|
||||
**Key parameters:**
|
||||
**Condition types:**
|
||||
|
||||
| Parameter | Description |
|
||||
|------------------------------|------------------------------------------------------------------|
|
||||
| `name` | Strategy type: `"text_to_video"` or `"video_to_video"` |
|
||||
| `first_frame_conditioning_p` | Probability of using first frame as conditioning (0.0-1.0) |
|
||||
| `with_audio` | (text_to_video only) Enable joint audio-video training |
|
||||
| `audio_latents_dir` | (text_to_video only) Directory name for audio latents |
|
||||
| `reference_latents_dir` | (video_to_video only) Directory name for reference video latents |
|
||||
| Type | Parameters | Description |
|
||||
|----------------|-----------------------------------------------------|---------------------------------------------------------------------------------------|
|
||||
| `first_frame` | `probability` | First latent frame is clean, excluded from loss. **Video only.** |
|
||||
| `prefix` | `temporal_boundary`, `probability` | First N latent temporal units are clean. For extension forward. |
|
||||
| `suffix` | `temporal_boundary`, `probability` | Last N latent temporal units are clean. For extension backward. |
|
||||
| `spatial_crop` | `spatial_region` (y1, x1, y2, x2 in px), `probability` | Rectangular region is clean, excluded from loss. For outpainting. **Video only.** |
|
||||
| `mask` | `mask_dir`, `probability` | Per-sample mask directory. Masks are thresholded at `0.5`; `1` means conditioning, `0` means generate. |
|
||||
| `reference` | `latents_dir`, `probability` | IC-LoRA style concatenation. Reference tokens are prepended, clean (timestep=0), no loss. |
|
||||
|
||||
> [!NOTE]
|
||||
> The `prefix`, `suffix`, `mask`, and `reference` condition types work on both video and audio modalities —
|
||||
> place them in the `video.conditions` or `audio.conditions` list as appropriate.
|
||||
> `first_frame` and `spatial_crop` are video-only conditions.
|
||||
|
||||
> [!NOTE]
|
||||
> Training conditions reference **directories** of precomputed data (within `preprocessed_data_root`),
|
||||
> while validation conditions reference **individual files** (images, videos, masks) that are encoded
|
||||
> on-the-fly during validation. The condition `type` names are the same, but the fields differ.
|
||||
|
||||
> [!NOTE]
|
||||
> The legacy `text_to_video` and `video_to_video` strategies are deprecated but remain forward-compatible.
|
||||
> New configs should use `name: "flexible"`.
|
||||
|
||||
### OptimizationConfig
|
||||
|
||||
@@ -206,7 +245,7 @@ optimization:
|
||||
| `steps` | Total number of training steps |
|
||||
| `batch_size` | Batch size per GPU (reduce if running out of memory) |
|
||||
| `gradient_accumulation_steps` | Accumulate gradients over multiple steps |
|
||||
| `scheduler_type` | LR scheduler: `"constant"`, `"linear"`, `"cosine"`, `"cosine_with_restarts"`, `"polynomial"` |
|
||||
| `scheduler_type` | LR scheduler: `"constant"`, `"linear"`, `"cosine"`, `"cosine_with_restarts"`, `"polynomial"`, `"step"` |
|
||||
| `enable_gradient_checkpointing` | Trade training speed for GPU memory savings (recommended for large models) |
|
||||
|
||||
### AccelerationConfig
|
||||
@@ -226,7 +265,7 @@ acceleration:
|
||||
| Parameter | Description |
|
||||
|---------------------------------------|------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| `mixed_precision_mode` | Precision mode - `"bf16"` recommended for modern GPUs |
|
||||
| `quantization` | Model quantization: `null`, `"int8-quanto"`, `"int4-quanto"`, `"fp8-quanto"`, etc. |
|
||||
| `quantization` | Model quantization: `null`, `"int8-quanto"`, `"int4-quanto"`, `"int2-quanto"`, `"fp8-quanto"`, or `"fp8uz-quanto"` |
|
||||
| `load_text_encoder_in_8bit` | Load the Gemma text encoder in 8-bit to save GPU memory |
|
||||
| `offload_optimizer_during_validation` | Move optimizer state to CPU before validation video sampling and back afterwards. Useful when validation OOMs because VAE decoder + transformer + optimizer state can't coexist on the GPU (full fine-tune, high-rank LoRA). No effect for FSDP. |
|
||||
|
||||
@@ -244,50 +283,84 @@ data:
|
||||
|
||||
| Parameter | Description |
|
||||
|--------------------------|--------------------------------------------------------------------------------------------|
|
||||
| `preprocessed_data_root` | Path to your preprocessed dataset (contains `latents/`, `conditions/`, etc.) |
|
||||
| `preprocessed_data_root` | Path to your preprocessed dataset directory produced by `process_dataset.py` (contains `latents/`, `conditions/`, etc.) |
|
||||
| `num_dataloader_workers` | Number of parallel data loading processes (0 = synchronous loading, useful when debugging) |
|
||||
|
||||
### ValidationConfig
|
||||
|
||||
Validation and inference settings for monitoring training progress.
|
||||
Validation and inference settings for monitoring training progress. Validation samples use a self-describing
|
||||
format where each sample specifies its own prompt and conditions.
|
||||
|
||||
```yaml
|
||||
validation:
|
||||
prompts: # Validation prompts
|
||||
- "A cat playing with a ball"
|
||||
- "A dog running in a field"
|
||||
negative_prompt: "worst quality, inconsistent motion, blurry, jittery, distorted"
|
||||
images: null # Optional image paths for image-to-video
|
||||
reference_videos: null # Reference video paths (IC-LoRA only)
|
||||
video_dims: [ 576, 576, 89 ] # Video dimensions [width, height, frames]
|
||||
frame_rate: 25.0 # Frame rate for generated videos
|
||||
seed: 42 # Random seed for reproducibility
|
||||
inference_steps: 30 # Number of inference steps
|
||||
interval: 100 # Steps between validation runs
|
||||
guidance_scale: 4.0 # CFG guidance strength
|
||||
stg_scale: 1.0 # STG guidance strength (0.0 to disable)
|
||||
stg_blocks: [ 29 ] # Transformer blocks to perturb for STG
|
||||
stg_mode: "stg_av" # "stg_av" or "stg_v" (video only)
|
||||
generate_audio: true # Whether to generate audio
|
||||
skip_initial_validation: false # Skip validation at step 0
|
||||
include_reference_in_output: false # Include reference video side-by-side (IC-LoRA)
|
||||
samples:
|
||||
- prompt: "A cat playing with a ball"
|
||||
conditions:
|
||||
- type: first_frame
|
||||
image_or_video: "/path/to/image.png"
|
||||
- prompt: "A dog running in a field"
|
||||
video_dims: [576, 576, 89] # Output dimensions: [width, height, frames]
|
||||
negative_prompt: "worst quality, inconsistent motion, blurry, jittery, distorted" # Negative prompt for all samples
|
||||
frame_rate: 25.0 # Output video frame rate (fps)
|
||||
seed: 42 # Random seed for reproducibility
|
||||
inference_steps: 30 # Number of denoising steps
|
||||
interval: 100 # Run validation every N steps (null to disable)
|
||||
guidance_scale: 4.0 # CFG scale (higher = stronger prompt adherence)
|
||||
stg_scale: 1.0 # STG scale (0.0 to disable)
|
||||
stg_blocks: [29] # Transformer blocks to apply STG perturbation
|
||||
stg_mode: "stg_av" # STG mode: "stg_av" (audio+video) or "stg_v" (video only)
|
||||
generate_audio: true # Whether to generate audio during validation
|
||||
generate_video: true # Whether to generate video during validation
|
||||
skip_initial_validation: false # Skip validation at step 0
|
||||
```
|
||||
|
||||
**Key parameters:**
|
||||
|
||||
| Parameter | Description |
|
||||
|-------------------------------|--------------------------------------------------------------------------------------------------------------------------|
|
||||
| `prompts` | List of text prompts for validation video generation |
|
||||
| `images` | List of image paths for image-to-video validation (must match number of prompts) |
|
||||
| `reference_videos` | List of reference video paths for IC-LoRA validation (must match number of prompts) |
|
||||
| `video_dims` | Output dimensions `[width, height, frames]`. Width/height must be divisible by 32, frames must satisfy `frames % 8 == 1` |
|
||||
| `interval` | Steps between validation runs (set to `null` to disable) |
|
||||
| `guidance_scale` | CFG (Classifier-Free Guidance) scale. Recommended: 4.0 |
|
||||
| `stg_scale` | STG (Spatio-Temporal Guidance) scale. 0.0 disables STG. Recommended: 1.0 |
|
||||
| `stg_blocks` | Transformer blocks to perturb for STG. Recommended: `[29]` (single block) |
|
||||
| `stg_mode` | STG mode: `"stg_av"` perturbs both audio and video, `"stg_v"` perturbs video only |
|
||||
| `generate_audio` | Whether to generate audio in validation samples |
|
||||
| `include_reference_in_output` | For IC-LoRA: concatenate reference video side-by-side with output |
|
||||
| Parameter | Description |
|
||||
|--------------------------|--------------------------------------------------------------------------------------------------------------------------|
|
||||
| `samples` | List of `ValidationSample` objects (see below). Replaces the legacy `prompts`/`images`/`reference_videos` fields. |
|
||||
| `video_dims` | Output dimensions `[width, height, frames]`. Width/height must be divisible by 32, frames must satisfy `frames % 8 == 1` |
|
||||
| `interval` | Steps between validation runs (set to `null` to disable) |
|
||||
| `guidance_scale` | CFG (Classifier-Free Guidance) scale. Recommended: 4.0 |
|
||||
| `stg_scale` | STG (Spatio-Temporal Guidance) scale. 0.0 disables STG. Recommended: 1.0 |
|
||||
| `stg_blocks` | Transformer blocks to perturb for STG. Recommended: `[29]` (single block) |
|
||||
| `stg_mode` | STG mode: `"stg_av"` perturbs both audio and video, `"stg_v"` perturbs video only |
|
||||
| `generate_audio` | Whether to generate audio in validation samples |
|
||||
| `generate_video` | Whether to generate video in validation samples. Set to `false` for V2A (video-to-audio) validation. Default: `true` |
|
||||
| `skip_initial_validation`| Skip validation video sampling at step 0 (beginning of training) |
|
||||
|
||||
#### ValidationSample
|
||||
|
||||
Each sample in the `samples` list has:
|
||||
|
||||
| Field | Description |
|
||||
|--------------|-------------------------------------------------------------------------------------------------|
|
||||
| `prompt` | Text prompt for this validation sample. |
|
||||
| `conditions` | List of validation conditions (see types below). Empty list = text-only generation. |
|
||||
| `video_dims` | Optional per-sample override for `(width, height, frames)`. Inherits from `ValidationConfig` if not set. |
|
||||
| `seed` | Optional per-sample override for random seed. Inherits from `ValidationConfig` if not set. |
|
||||
|
||||
#### Validation Condition Types
|
||||
|
||||
| Type | Parameters | Description |
|
||||
|------------------|------------------------------------------------------------|-------------------------------------------------------------------------|
|
||||
| `first_frame` | `image_or_video` (path) | Use the first frame of the image/video as conditioning. |
|
||||
| `prefix` | `video` or `audio` (path), optional `num_frames`/`duration`| Use a video/audio clip as temporal prefix (for extension forward). |
|
||||
| `suffix` | `video` or `audio` (path), optional `num_frames`/`duration`| Use a video/audio clip as temporal suffix (for extension backward). |
|
||||
| `spatial_crop` | `video` (path), `spatial_region` (y1, x1, y2, x2) | Provide spatial context for outpainting. Video only. |
|
||||
| `mask` | `video` or `audio` (path), `mask` (path) | Mask-based inpainting with a binary mask file. |
|
||||
| `reference` | `video` or `audio` (path), optional video-reference `downscale_factor`, `temporal_scale_factor`, `include_in_output` | IC-LoRA style reference conditioning. |
|
||||
| `video_to_audio` | `video` (path) | Freeze video, generate audio. For Foley/V2A tasks. |
|
||||
| `audio_to_video` | `audio` (path) | Freeze audio, generate video. For audio-driven generation. |
|
||||
|
||||
For video `reference` validation conditions, `downscale_factor` is the spatial reference scale and
|
||||
`temporal_scale_factor` is the temporal reference scale. Set both to match the factors used when
|
||||
preprocessing video reference latents for training; validation media is encoded on the fly and cannot infer
|
||||
those factors from the training dataset.
|
||||
|
||||
> [!NOTE]
|
||||
> The legacy fields `prompts`, `images`, and `reference_videos` are deprecated but auto-converted to `samples`
|
||||
> internally. New configs should use the `samples` format.
|
||||
|
||||
### CheckpointsConfig
|
||||
|
||||
@@ -298,6 +371,8 @@ checkpoints:
|
||||
interval: 250 # Steps between checkpoint saves (null = disabled)
|
||||
keep_last_n: 3 # Number of recent checkpoints to retain
|
||||
precision: bfloat16 # Precision for saved weights (bfloat16 or float32)
|
||||
no_resume: false # Ignore saved state, start from step 0
|
||||
save_training_state: "minimal" # "full", "minimal", or "off"
|
||||
```
|
||||
|
||||
**Key parameters:**
|
||||
@@ -307,6 +382,8 @@ checkpoints:
|
||||
| `interval` | Steps between intermediate checkpoint saves (set to `null` to disable) |
|
||||
| `keep_last_n` | Number of most recent checkpoints to keep (-1 = keep all) |
|
||||
| `precision` | Precision for saved checkpoint weights: `"bfloat16"` (default) or `"float32"` |
|
||||
| `no_resume` | When `true`, ignore saved training state and start from step 0. Model weights from `load_checkpoint` are still loaded. |
|
||||
| `save_training_state` | Save training state for resume: `"full"` (optimizer + scheduler + RNG), `"minimal"` (scheduler + RNG only, sufficient for LoRA), `"off"` (no resume). |
|
||||
|
||||
### HubConfig
|
||||
|
||||
@@ -364,6 +441,20 @@ flow_matching:
|
||||
| `timestep_sampling_mode` | Sampling strategy: `"uniform"` or `"shifted_logit_normal"` |
|
||||
| `timestep_sampling_params` | Additional parameters for the sampling strategy |
|
||||
|
||||
### General Configuration
|
||||
|
||||
Top-level settings for the training run.
|
||||
|
||||
```yaml
|
||||
seed: 42 # Random seed for reproducibility
|
||||
output_dir: "outputs/my_training_run" # Directory for outputs (checkpoints, validation videos, logs)
|
||||
```
|
||||
|
||||
| Parameter | Description |
|
||||
|--------------|----------------------------------------------------------|
|
||||
| `seed` | Random seed for reproducibility (default: `42`) |
|
||||
| `output_dir` | Directory to save outputs (default: `"outputs"`) |
|
||||
|
||||
## 🚀 Next Steps
|
||||
|
||||
Once you've configured your training parameters:
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# Implementing Custom Training Strategies
|
||||
|
||||
This guide explains how to implement your own training strategy for specialized use cases like audio-only training,
|
||||
video inpainting, or other custom training recipes.
|
||||
This guide explains how to implement your own training strategy for specialized recipes that cannot be expressed with
|
||||
the built-in `flexible` strategy.
|
||||
|
||||
## 📋 Overview
|
||||
|
||||
@@ -15,12 +15,20 @@ This architecture lets you implement new training modes without modifying the co
|
||||
|
||||
### When You Need a Custom Strategy
|
||||
|
||||
> [!NOTE]
|
||||
> The built-in `flexible` strategy already supports most conditioning scenarios out of the box:
|
||||
> first-frame conditioning, video extension (prefix/suffix), spatial crop (outpainting),
|
||||
> mask-based inpainting, IC-LoRA reference conditioning, and frozen modality cross-conditioning
|
||||
> (audio-to-video, video-to-audio). Only implement a custom strategy if your use case requires
|
||||
> fundamentally different training logic that cannot be expressed through the flexible strategy's
|
||||
> configuration.
|
||||
|
||||
Consider implementing a custom strategy when you need:
|
||||
|
||||
- **Different input modalities** (e.g., audio-only, audio-to-video conditioning)
|
||||
- **Additional conditioning signals** (e.g., masks for inpainting, depth maps)
|
||||
- **Custom loss computation** (e.g., weighted losses, auxiliary losses)
|
||||
- **Different noise application patterns** (e.g., partial masking)
|
||||
- **Custom loss computation** (e.g., weighted losses, auxiliary losses, perceptual losses)
|
||||
- **Non-standard noise application** (e.g., noise schedules different from flow matching)
|
||||
- **Novel conditioning mechanisms** not covered by the flexible strategy's condition types
|
||||
- **Additional model outputs** beyond the standard video/audio predictions
|
||||
|
||||
## 🏗️ Architecture Overview
|
||||
|
||||
@@ -28,7 +36,7 @@ Consider implementing a custom strategy when you need:
|
||||
|
||||
The trainer delegates all training-mode-specific logic to the strategy:
|
||||
|
||||
1. **Initialization** — The trainer calls `get_data_sources()` to determine which preprocessed data directories to load
|
||||
1. **Initialization** — The trainer calls `config.get_data_sources()` to determine which preprocessed data directories to load
|
||||
2. **Each training step:**
|
||||
- Calls `prepare_training_inputs()` to transform the raw batch into model-ready inputs
|
||||
- Runs the transformer forward pass
|
||||
@@ -52,8 +60,8 @@ The trainer handles everything else: optimization, checkpointing, validation, an
|
||||
Before writing code, answer these questions:
|
||||
|
||||
1. **What additional data does your strategy need?**
|
||||
- Example: Inpainting needs mask latents alongside video latents
|
||||
- Example: Audio-to-video needs reference audio embeddings
|
||||
- Example: A perceptual-loss strategy may need auxiliary feature targets
|
||||
- Example: A novel conditioning mechanism may need an additional precomputed directory
|
||||
|
||||
2. **What does conditioning look like?**
|
||||
- Which tokens should be noised vs. kept clean?
|
||||
@@ -164,6 +172,20 @@ class InpaintingConfig(TrainingStrategyConfigBase):
|
||||
ge=0.0,
|
||||
le=1.0,
|
||||
)
|
||||
|
||||
def get_data_sources(self) -> dict[str, str]:
|
||||
"""Define which data directories to load.
|
||||
|
||||
Returns a mapping of directory names (under preprocessed_data_root) to
|
||||
batch keys. The trainer loads .pt files from each directory and exposes
|
||||
them in the batch under the specified key. The trainer also uses this
|
||||
mapping to validate that all required directories exist.
|
||||
"""
|
||||
return {
|
||||
"latents": "latents", # -> batch["latents"]
|
||||
"conditions": "conditions", # -> batch["conditions"]
|
||||
self.mask_latents_dir: "masks", # -> batch["masks"]
|
||||
}
|
||||
```
|
||||
|
||||
**Key points:**
|
||||
@@ -171,6 +193,7 @@ class InpaintingConfig(TrainingStrategyConfigBase):
|
||||
- Inherit from `TrainingStrategyConfigBase`
|
||||
- Use `Literal["your_strategy_name"]` for the `name` field - this enables automatic strategy selection
|
||||
- Use Pydantic `Field` for validation and documentation
|
||||
- Implement `get_data_sources()` on the config — it's the single source of truth for data directories (used for both dataset wiring and existence validation)
|
||||
|
||||
### Step 4: Implement the Strategy Class
|
||||
|
||||
@@ -187,24 +210,6 @@ class InpaintingStrategy(TrainingStrategy):
|
||||
def __init__(self, config: InpaintingConfig):
|
||||
super().__init__(config)
|
||||
|
||||
@property
|
||||
def requires_audio(self) -> bool:
|
||||
"""Whether this strategy requires audio components."""
|
||||
return False # Set to True if your strategy needs audio
|
||||
|
||||
def get_data_sources(self) -> dict[str, str]:
|
||||
"""Define which data directories to load.
|
||||
|
||||
Returns a mapping of directory names to batch keys.
|
||||
The trainer will load .pt files from each directory and
|
||||
make them available in the batch under the specified key.
|
||||
"""
|
||||
return {
|
||||
"latents": "latents", # -> batch["latents"]
|
||||
"conditions": "conditions", # -> batch["conditions"]
|
||||
self.config.mask_latents_dir: "masks", # -> batch["masks"]
|
||||
}
|
||||
|
||||
def prepare_training_inputs(
|
||||
self,
|
||||
batch: dict[str, Any],
|
||||
@@ -275,7 +280,6 @@ class InpaintingStrategy(TrainingStrategy):
|
||||
batch_size=batch_size,
|
||||
fps=24.0, # Or get from latents_data
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
# Create video Modality
|
||||
@@ -328,7 +332,7 @@ You need to register your strategy in two places:
|
||||
from ltx_trainer.training_strategies.inpainting import InpaintingConfig, InpaintingStrategy
|
||||
|
||||
# Add to the TrainingStrategyConfig type alias
|
||||
TrainingStrategyConfig = TextToVideoConfig | VideoToVideoConfig | InpaintingConfig
|
||||
TrainingStrategyConfig = TextToVideoConfig | VideoToVideoConfig | FlexibleStrategyConfig | InpaintingConfig
|
||||
|
||||
# Add to __all__
|
||||
__all__ = [
|
||||
@@ -356,7 +360,8 @@ from ltx_trainer.training_strategies.inpainting import InpaintingConfig
|
||||
TrainingStrategyConfig = Annotated[
|
||||
Annotated[TextToVideoConfig, Tag("text_to_video")]
|
||||
| Annotated[VideoToVideoConfig, Tag("video_to_video")]
|
||||
| Annotated[InpaintingConfig, Tag("inpainting")], # Add your config
|
||||
| Annotated[FlexibleStrategyConfig, Tag("flexible")]
|
||||
| Annotated[InpaintingConfig, Tag("inpainting")],
|
||||
Discriminator(_get_strategy_discriminator),
|
||||
]
|
||||
```
|
||||
@@ -366,7 +371,7 @@ TrainingStrategyConfig = Annotated[
|
||||
Create an example config in `configs/`:
|
||||
|
||||
```yaml
|
||||
# configs/ltx2_inpainting_lora.yaml
|
||||
# configs/custom_inpainting_lora.yaml
|
||||
|
||||
model:
|
||||
model_path: "/path/to/ltx2.safetensors"
|
||||
@@ -408,8 +413,8 @@ The base `TrainingStrategy` class provides these helper methods:
|
||||
| `_audio_patchifier.patchify(latents)` | Convert `[B, C, T, F]` → `[B, T, C*F]` |
|
||||
| `_get_video_positions(...)` | Generate position embeddings for video |
|
||||
| `_get_audio_positions(...)` | Generate position embeddings for audio |
|
||||
| `_create_per_token_timesteps(mask, sigma)` | Create timesteps with 0 for conditioning tokens |
|
||||
| `_create_first_frame_conditioning_mask(...)` | Create mask for first-frame conditioning |
|
||||
| `_create_per_token_timesteps(conditioning_mask, sampled_sigma)` | Create timesteps with 0 for conditioning tokens |
|
||||
| `_create_first_frame_conditioning_mask(...)` | Create mask for first-frame conditioning |
|
||||
|
||||
## 📊 Understanding ModelInputs
|
||||
|
||||
@@ -418,16 +423,14 @@ The `ModelInputs` dataclass contains everything needed for the forward pass and
|
||||
```python
|
||||
@dataclass
|
||||
class ModelInputs:
|
||||
video: Modality # Video modality data
|
||||
audio: Modality | None # Audio modality (None if video-only)
|
||||
video: Modality | None # Video modality data
|
||||
audio: Modality | None # Audio modality data
|
||||
|
||||
video_targets: Tensor # Target values for loss (velocity)
|
||||
audio_targets: Tensor | None
|
||||
video_targets: Tensor | None # Target values for video loss (velocity)
|
||||
audio_targets: Tensor | None # Target values for audio loss (velocity)
|
||||
|
||||
video_loss_mask: Tensor # Boolean: True = compute loss for this token
|
||||
audio_loss_mask: Tensor | None
|
||||
|
||||
ref_seq_len: int | None = None # For IC-LoRA: reference sequence length
|
||||
video_loss_mask: Tensor | None # Boolean loss mask for video tokens
|
||||
audio_loss_mask: Tensor | None # Boolean loss mask for audio tokens
|
||||
```
|
||||
|
||||
## 📊 Understanding Modality
|
||||
@@ -437,18 +440,20 @@ The `Modality` dataclass (from ltx-core) represents a single modality's data:
|
||||
```python
|
||||
@dataclass(frozen=True)
|
||||
class Modality:
|
||||
enabled: bool # Whether this modality is active
|
||||
latent: Tensor # [B, seq_len, C] - the latent tokens
|
||||
timesteps: Tensor # [B, seq_len] - per-token timesteps (sigmas)
|
||||
positions: Tensor # [B, dims, seq_len, 2] - position bounds
|
||||
context: Tensor # [B, ctx_len, C] - text embeddings
|
||||
context_mask: Tensor # [B, ctx_len] - attention mask for context
|
||||
latent: Tensor # [B, T, D] — patchified latent tokens
|
||||
sigma: Tensor # [B,] — per-batch noise level (for cross-attn conditioning)
|
||||
timesteps: Tensor # [B, T] — per-token timestep embeddings
|
||||
positions: Tensor # [B, 3, T, 2] for video, [B, 1, T, 2] for audio — positional bounds
|
||||
context: Tensor # text conditioning embeddings
|
||||
enabled: bool = True
|
||||
context_mask: Tensor | None = None # attention mask for text context
|
||||
attention_mask: Tensor | None = None # optional 2D self-attention mask [B, T, T]
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> **Per-token timesteps:** Each token in the sequence has its own timestep. Conditioning tokens—those that should remain
|
||||
> un-noised—must have `timestep=0`. This is how the model distinguishes clean reference tokens from tokens to denoise. Use
|
||||
`_create_per_token_timesteps(conditioning_mask, sigma)` to set this up correctly.
|
||||
> `_create_per_token_timesteps(conditioning_mask, sampled_sigma)` to set this up correctly.
|
||||
|
||||
> [!NOTE]
|
||||
> `Modality` is immutable (frozen dataclass). Use `dataclasses.replace()` to create modified copies.
|
||||
@@ -461,7 +466,7 @@ class Modality:
|
||||
from ltx_trainer.config import LtxTrainerConfig
|
||||
import yaml
|
||||
|
||||
with open('configs/ltx2_inpainting_lora.yaml') as f:
|
||||
with open('configs/custom_inpainting_lora.yaml') as f:
|
||||
config = LtxTrainerConfig(**yaml.safe_load(f))
|
||||
print(f'Strategy: {config.training_strategy.name}')
|
||||
"
|
||||
@@ -475,13 +480,13 @@ class Modality:
|
||||
|
||||
config = InpaintingConfig()
|
||||
strategy = get_training_strategy(config)
|
||||
print(f'Data sources: {strategy.get_data_sources()}')
|
||||
print(f'Data sources: {config.get_data_sources()}')
|
||||
"
|
||||
```
|
||||
|
||||
3. **Run a short training test:**
|
||||
```bash
|
||||
uv run python scripts/train.py configs/ltx2_inpainting_lora.yaml
|
||||
uv run python scripts/train.py configs/custom_inpainting_lora.yaml
|
||||
```
|
||||
|
||||
## 💡 Tips and Best Practices
|
||||
@@ -503,7 +508,8 @@ class Modality:
|
||||
|
||||
Study these implementations for guidance:
|
||||
|
||||
| Strategy | Complexity | Key Features |
|
||||
|------------------------------------------------------------------------------------|------------|------------------------------------------------|
|
||||
| [`TextToVideoStrategy`](../src/ltx_trainer/training_strategies/text_to_video.py) | Simple | First-frame conditioning, optional audio |
|
||||
| [`VideoToVideoStrategy`](../src/ltx_trainer/training_strategies/video_to_video.py) | Medium | Reference video concatenation, split loss mask |
|
||||
| Strategy | Complexity | Key Features |
|
||||
|----------|------------|--------------|
|
||||
| [`FlexibleStrategy`](../src/ltx_trainer/training_strategies/flexible.py) | Medium | Unified conditioning framework — supports all built-in modes |
|
||||
| [`TextToVideoStrategy`](../src/ltx_trainer/training_strategies/text_to_video.py) | Simple | First-frame conditioning, optional audio (deprecated) |
|
||||
| [`VideoToVideoStrategy`](../src/ltx_trainer/training_strategies/video_to_video.py) | Medium | Reference video concatenation, split loss mask (deprecated) |
|
||||
|
||||
@@ -33,41 +33,33 @@ uv run python scripts/split_scenes.py --help
|
||||
|
||||
If your dataset doesn't include captions, you can automatically generate them using multimodal models that understand both video and audio.
|
||||
|
||||
The default `qwen_omni` backend talks to a local vLLM server, which you launch once in a separate terminal:
|
||||
|
||||
```bash
|
||||
uv run python scripts/caption_videos.py scenes_output_dir/ \
|
||||
--output scenes_output_dir/dataset.json
|
||||
# Terminal 1: start the captioner server (stays running)
|
||||
uv run python scripts/serve_captioner.py
|
||||
```
|
||||
|
||||
If you're running into VRAM issues, try enabling 8-bit quantization to reduce memory usage:
|
||||
|
||||
```bash
|
||||
# Terminal 2: caption your videos
|
||||
uv run python scripts/caption_videos.py scenes_output_dir/ \
|
||||
--output scenes_output_dir/dataset.json \
|
||||
--use-8bit
|
||||
--output scenes_output_dir/dataset.json
|
||||
```
|
||||
|
||||
This will create a `dataset.json` file containing video paths and their captions.
|
||||
|
||||
**Captioning options:**
|
||||
|
||||
|
||||
| Option | Description |
|
||||
| ------------------ | ---------------------------------------------------------- |
|
||||
| `--captioner-type` | `qwen_omni` (default, local) or `gemini_flash` (API) |
|
||||
| `--use-8bit` | Enable 8-bit quantization for lower VRAM usage |
|
||||
| `--no-audio` | Disable audio processing (video-only captions) |
|
||||
| `--override` | Re-caption files that already have captions |
|
||||
| `--api-key` | API key for Gemini Flash (or set `GOOGLE_API_KEY` env var) |
|
||||
|
||||
| Option | Description |
|
||||
| ------------------ | --------------------------------------------------------------- |
|
||||
| `--captioner-type` | `qwen_omni` (default, local vLLM server) or `gemini_flash` (API) |
|
||||
| `--vllm-url` | Base URL of the vLLM server (default `http://127.0.0.1:8001/v1`) |
|
||||
| `--override` | Re-caption files that already have captions |
|
||||
| `--api-key` | Gemini API key (else `GEMINI_API_KEY`/`GOOGLE_API_KEY`; with no key, uses gcloud/Vertex AI auth) |
|
||||
|
||||
**Caption format:**
|
||||
|
||||
The captioner produces structured captions with sections for:
|
||||
|
||||
- **Visual content**: People, objects, actions, settings, colors, movements
|
||||
- **Speech transcription**: Word-for-word transcription of spoken content
|
||||
- **Sounds**: Music, ambient sounds, sound effects
|
||||
- **On-screen text**: Any visible text overlays
|
||||
Each caption is a single, detailed paragraph describing both the visual content and the audio (speech, music, ambient sounds) of the clip. See the [Utility Scripts Reference](utility-scripts.md#automatic-video-captioning) for backend setup and the full list of options.
|
||||
|
||||
> [!NOTE]
|
||||
> The automatically generated captions may contain inaccuracies or hallucinated content.
|
||||
@@ -80,7 +72,7 @@ This step preprocesses your video dataset by:
|
||||
1. Resizing and cropping videos to fit specified resolution buckets
|
||||
2. Computing and caching video latent representations
|
||||
3. Computing and caching text embeddings for captions
|
||||
4. (Optional) Computing and caching audio latents
|
||||
4. Extracting and caching audio latents from videos (automatic, use `--skip-audio` to disable)
|
||||
|
||||
> [!WARNING]
|
||||
> Very large videos (especially high spatial resolution and/or many frames) can cause GPU out-of-memory (OOM)
|
||||
@@ -97,17 +89,9 @@ uv run python scripts/process_dataset.py dataset.json \
|
||||
--text-encoder-path /path/to/gemma-model
|
||||
```
|
||||
|
||||
### With Audio Processing
|
||||
|
||||
For audio-video training, add the `--with-audio` flag:
|
||||
|
||||
```bash
|
||||
uv run python scripts/process_dataset.py dataset.json \
|
||||
--resolution-buckets "960x544x49" \
|
||||
--model-path /path/to/ltx-2-model.safetensors \
|
||||
--text-encoder-path /path/to/gemma-model \
|
||||
--with-audio
|
||||
```
|
||||
Audio latents are automatically extracted from video files — no extra flag is needed. Use `--skip-audio`
|
||||
to disable this. For standalone audio files (`.wav`), use the `audio` column in your dataset instead
|
||||
(see [Convention-Based Column Detection](#convention-based-column-detection) below).
|
||||
|
||||
### 🚀 Multi-GPU Preprocessing
|
||||
|
||||
@@ -126,7 +110,7 @@ Outputs are written atomically (via a per-process temporary file, then renamed),
|
||||
corrupt files. By default a rerun **resumes** — items whose output `.pt` already exists are skipped.
|
||||
|
||||
> [!IMPORTANT]
|
||||
> Pass `**--overwrite`** when rerunning with changed parameters (different model checkpoint, resolution buckets,
|
||||
> Pass **`--overwrite`** when rerunning with changed parameters (different model checkpoint, resolution buckets,
|
||||
> text encoder, `--lora-trigger`, etc.). Without it the script keeps the stale outputs from the previous run.
|
||||
>
|
||||
> ```bash
|
||||
@@ -152,7 +136,7 @@ The trainer supports videos, single images, or a mix of both in the same dataset
|
||||
> `--resolution-buckets "960x544x1;960x544x49"`. Images are automatically assigned to the `F=1` bucket and
|
||||
> videos to an `F>1` bucket.
|
||||
> - You **must** set `optimization.batch_size: 1` in your training config (see the warning under
|
||||
> [Resolution Buckets](#-resolution-buckets)), since samples with different shapes cannot be collated into a
|
||||
> [Resolution Buckets](#resolution-buckets)), since samples with different shapes cannot be collated into a
|
||||
> single batch. Use `gradient_accumulation_steps` if you need a larger effective batch.
|
||||
> - Per-step cost differs substantially between a single-frame sample and a many-frame sample, which can lead to
|
||||
> uneven gradient magnitudes across steps. Consider weighting the two subsets or tuning the learning rate if
|
||||
@@ -160,7 +144,24 @@ The trainer supports videos, single images, or a mix of both in the same dataset
|
||||
> - If you prefer a fully officially-supported path, train two separate LoRAs (one on stills, one on video) and
|
||||
> stack them at inference.
|
||||
|
||||
The dataset must be a CSV, JSON, or JSONL metadata file with columns for captions and video paths:
|
||||
The dataset must be a CSV, JSON, or JSONL metadata file with columns for captions and media paths.
|
||||
|
||||
#### Convention-Based Column Detection
|
||||
|
||||
The preprocessing script automatically detects and processes columns based on their names. The following columns are recognized:
|
||||
|
||||
| Column | Output Dir | Description |
|
||||
|--------|-----------|-------------|
|
||||
| `video` (or legacy `media_path`) | `latents/` | Target video to encode |
|
||||
| `audio` | `audio_latents/` | Explicit audio file (overrides auto-extraction from video) |
|
||||
| `caption` | `conditions/` | Text caption for the sample |
|
||||
| `reference_video` (or legacy `ref_media_path`) | `reference_latents/` | IC-LoRA reference video |
|
||||
| `reference_audio` | `reference_audio_latents/` | IC-LoRA reference audio |
|
||||
| `video_mask` | `video_masks/` | Binary mask for video inpainting |
|
||||
| `audio_mask` | `audio_masks/` | Binary mask for audio inpainting |
|
||||
|
||||
> [!NOTE]
|
||||
> **Legacy column names:** `media_path` and `ref_media_path` are accepted as aliases for `video` and `reference_video` respectively. Existing datasets using these names will continue to work without modification.
|
||||
|
||||
**JSON format example:**
|
||||
|
||||
@@ -168,11 +169,11 @@ The dataset must be a CSV, JSON, or JSONL metadata file with columns for caption
|
||||
[
|
||||
{
|
||||
"caption": "A cat playing with a ball of yarn",
|
||||
"media_path": "videos/cat_playing.mp4"
|
||||
"video": "videos/cat_playing.mp4"
|
||||
},
|
||||
{
|
||||
"caption": "A dog running in the park",
|
||||
"media_path": "videos/dog_running.mp4"
|
||||
"video": "videos/dog_running.mp4"
|
||||
}
|
||||
]
|
||||
```
|
||||
@@ -180,18 +181,42 @@ The dataset must be a CSV, JSON, or JSONL metadata file with columns for caption
|
||||
**JSONL format example:**
|
||||
|
||||
```jsonl
|
||||
{"caption": "A cat playing with a ball of yarn", "media_path": "videos/cat_playing.mp4"}
|
||||
{"caption": "A dog running in the park", "media_path": "videos/dog_running.mp4"}
|
||||
{"caption": "A cat playing with a ball of yarn", "video": "videos/cat_playing.mp4"}
|
||||
{"caption": "A dog running in the park", "video": "videos/dog_running.mp4"}
|
||||
```
|
||||
|
||||
**CSV format example:**
|
||||
|
||||
```csv
|
||||
caption,media_path
|
||||
caption,video
|
||||
"A cat playing with a ball of yarn","videos/cat_playing.mp4"
|
||||
"A dog running in the park","videos/dog_running.mp4"
|
||||
```
|
||||
|
||||
**Additional dataset format examples:**
|
||||
|
||||
Audio-only dataset:
|
||||
```json
|
||||
{"audio": "song.wav", "caption": "piano melody"}
|
||||
```
|
||||
|
||||
V2V IC-LoRA with reference video:
|
||||
```json
|
||||
{"video": "clip.mp4", "reference_video": "depth.mp4", "caption": "depth to video"}
|
||||
```
|
||||
|
||||
A2A IC-LoRA with reference audio:
|
||||
```json
|
||||
{"video": "clip.mp4", "reference_audio": "ref.wav", "caption": "match this style"}
|
||||
```
|
||||
This form auto-extracts the target audio from `clip.mp4`. For pure audio datasets, use `audio` plus
|
||||
`reference_audio` columns and preprocess with `--audio-durations`.
|
||||
|
||||
Video inpainting with mask:
|
||||
```json
|
||||
{"video": "clip.mp4", "video_mask": "mask.mp4", "caption": "fill the sky"}
|
||||
```
|
||||
|
||||
### 📐 Resolution Buckets
|
||||
|
||||
Videos are organized into "buckets" of specific dimensions (width × height × frames).
|
||||
@@ -268,12 +293,31 @@ The preprocessed data is saved in a `.precomputed` directory:
|
||||
```
|
||||
dataset/
|
||||
└── .precomputed/
|
||||
├── latents/ # Cached video latents
|
||||
├── conditions/ # Cached text embeddings
|
||||
├── audio_latents/ # (only if --with-audio) Cached audio latents
|
||||
└── reference_latents/ # (only for IC-LoRA) Cached reference video latents
|
||||
├── latents/ # Video latents
|
||||
├── conditions/ # Text embeddings
|
||||
├── audio_latents/ # Audio latents (auto-extracted or explicit)
|
||||
├── reference_latents/ # Reference video latents (IC-LoRA)
|
||||
├── reference_audio_latents/ # Reference audio latents (audio IC-LoRA)
|
||||
├── video_masks/ # Video masks (inpainting)
|
||||
└── audio_masks/ # Audio masks (audio inpainting)
|
||||
```
|
||||
|
||||
Set `data.preprocessed_data_root` in your training config to this `.precomputed` directory — the parent directory that
|
||||
contains `latents/`, `conditions/`, and any mode-specific audio/reference/mask directories.
|
||||
|
||||
## 🔊 Audio-Only Dataset Preprocessing
|
||||
|
||||
For datasets containing only audio files (no `video` column), use `--audio-durations` to specify duration buckets:
|
||||
|
||||
```bash
|
||||
uv run python scripts/process_dataset.py dataset.json \
|
||||
--audio-durations "2.0;4.0;8.0" \
|
||||
--model-path /path/to/ltx-2-model.safetensors \
|
||||
--text-encoder-path /path/to/gemma-model
|
||||
```
|
||||
|
||||
The `--audio-durations` flag provides duration buckets (in seconds) for audio-only datasets. Since there is no video column to derive timing from, explicit duration buckets are required.
|
||||
|
||||
## 🪄 IC-LoRA Reference Video Preprocessing
|
||||
|
||||
For IC-LoRA training, you need to preprocess datasets that include reference videos.
|
||||
@@ -281,14 +325,16 @@ Reference videos provide the conditioning input while target videos represent th
|
||||
|
||||
### Dataset Format with Reference Videos
|
||||
|
||||
The `reference_video` column is automatically detected by convention — no extra CLI flags are needed.
|
||||
|
||||
**JSON format:**
|
||||
|
||||
```json
|
||||
[
|
||||
{
|
||||
"caption": "A cat playing with a ball of yarn",
|
||||
"media_path": "videos/cat_playing.mp4",
|
||||
"reference_path": "references/cat_playing_depth.mp4"
|
||||
"video": "videos/cat_playing.mp4",
|
||||
"reference_video": "references/cat_playing_depth.mp4"
|
||||
}
|
||||
]
|
||||
```
|
||||
@@ -296,32 +342,39 @@ Reference videos provide the conditioning input while target videos represent th
|
||||
**JSONL format:**
|
||||
|
||||
```jsonl
|
||||
{"caption": "A cat playing with a ball of yarn", "media_path": "videos/cat_playing.mp4", "reference_path": "references/cat_playing_depth.mp4"}
|
||||
{"caption": "A dog running in the park", "media_path": "videos/dog_running.mp4", "reference_path": "references/dog_running_depth.mp4"}
|
||||
{"caption": "A cat playing with a ball of yarn", "video": "videos/cat_playing.mp4", "reference_video": "references/cat_playing_depth.mp4"}
|
||||
{"caption": "A dog running in the park", "video": "videos/dog_running.mp4", "reference_video": "references/dog_running_depth.mp4"}
|
||||
```
|
||||
|
||||
### Preprocessing with Reference Videos
|
||||
|
||||
To preprocess a dataset with reference videos, add the `--reference-column` argument specifying the name of the field
|
||||
in your dataset JSON/JSONL/CSV that contains the reference video paths:
|
||||
Convention-based detection means you just need the `reference_video` column in your dataset, and `process_dataset.py` will automatically detect and process it. No `--reference-column` flag is needed:
|
||||
|
||||
```bash
|
||||
uv run python scripts/process_dataset.py dataset.json \
|
||||
--resolution-buckets "960x544x49" \
|
||||
--model-path /path/to/ltx-2-model.safetensors \
|
||||
--text-encoder-path /path/to/gemma-model \
|
||||
--reference-column "reference_path"
|
||||
--reference-downscale-factor 2 \
|
||||
--reference-temporal-scale-factor 1
|
||||
```
|
||||
|
||||
This will create an additional `reference_latents/` directory containing the preprocessed reference video latents.
|
||||
Use `--reference-downscale-factor` for spatial subsampling and `--reference-temporal-scale-factor` for temporal
|
||||
subsampling. Validation reference conditions should use matching `downscale_factor` and `temporal_scale_factor` values.
|
||||
|
||||
> [!NOTE]
|
||||
> **Legacy column names:** If your dataset uses `ref_media_path`, it is accepted as an alias for `reference_video`.
|
||||
|
||||
### Generating Reference Videos
|
||||
|
||||
**Dataset Requirements for IC-LoRA:**
|
||||
|
||||
- Your dataset must contain paired videos where each target video has a corresponding reference video
|
||||
- Reference and target videos must have *identical* resolution and length
|
||||
- Both reference and target videos should be preprocessed together using the same resolution buckets
|
||||
- Reference and target videos should cover the same content. Reference videos can optionally be lower spatial
|
||||
resolution or temporally subsampled (see Scaled Reference Conditioning in [Training Modes](training-modes.md)).
|
||||
- Both reference and target videos should be preprocessed together using the same target resolution buckets, plus any
|
||||
reference scale factors you choose.
|
||||
|
||||
We provide an example script, `[scripts/compute_reference.py](../scripts/compute_reference.py)`, to generate reference
|
||||
videos for a given dataset. The default implementation generates Canny edge reference videos.
|
||||
@@ -333,12 +386,63 @@ uv run python scripts/compute_reference.py scenes_output_dir/ \
|
||||
|
||||
The script accepts a JSON file as the dataset configuration and updates it in-place by adding the filenames of the generated reference videos.
|
||||
|
||||
> [!NOTE]
|
||||
> `compute_reference.py` writes generated references to the `reference_video` column, which `process_dataset.py`
|
||||
> detects automatically. The legacy `ref_media_path` column is also accepted.
|
||||
|
||||
If you want to generate a different type of condition (depth maps, pose skeletons, etc.), modify or replace the `compute_reference()` function within this script.
|
||||
|
||||
### Example Dataset
|
||||
|
||||
For reference, see our **[Canny Control Dataset](https://huggingface.co/datasets/Lightricks/Canny-Control-Dataset)** which demonstrates proper IC-LoRA dataset structure with paired videos and Canny edge maps.
|
||||
|
||||
## 🎭 Mask Preprocessing for Inpainting
|
||||
|
||||
For inpainting training with the `mask` condition type, provide `video_mask` or `audio_mask` columns in your dataset
|
||||
metadata. These columns point to mask media files (for example a mask image/video for video inpainting, or a waveform or
|
||||
`.pt` tensor for audio inpainting). `process_dataset.py` downsamples and thresholds them into per-sample `.pt` tensors
|
||||
under `video_masks/` or `audio_masks/`.
|
||||
|
||||
### Processed Video Mask Format
|
||||
|
||||
If you create masks manually instead of using `process_dataset.py`, save them as `.pt` files with the key `"mask"`
|
||||
containing a tensor of shape `[F, H, W]` where:
|
||||
|
||||
- `F` = number of latent frames (temporal dimension)
|
||||
- `H` = latent height (pixel height / 32)
|
||||
- `W` = latent width (pixel width / 32)
|
||||
- Values are thresholded at `0.5`: values `> 0.5` are conditioning tokens (clean, excluded from loss),
|
||||
and values `<= 0.5` are generated tokens (noised, contributes to loss).
|
||||
|
||||
### Audio Mask Format
|
||||
|
||||
Audio masks follow the same thresholding pattern as video masks but with shape `[T]` (temporal dimension only), where `T` is the number of audio latent frames. They are stored in `audio_masks/`.
|
||||
|
||||
### Directory Structure
|
||||
|
||||
Place masks in a directory within your preprocessed data root:
|
||||
|
||||
```
|
||||
preprocessed_data_root/
|
||||
├── latents/ # Video latents
|
||||
├── conditions/ # Text embeddings
|
||||
├── video_masks/ # Video masks (one .pt per sample, matching latent filenames)
|
||||
└── audio_masks/ # Audio masks (one .pt per sample, matching latent filenames)
|
||||
```
|
||||
|
||||
Then reference the mask directory in your training config:
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
conditions:
|
||||
- type: mask
|
||||
mask_dir: "video_masks"
|
||||
```
|
||||
|
||||
## 🎯 LoRA Trigger Words
|
||||
|
||||
When training a LoRA, you can specify a trigger token that will be prepended to all captions:
|
||||
@@ -359,9 +463,9 @@ This acts as a trigger word that activates the LoRA during inference when you in
|
||||
|
||||
## 🔍 Decoding Videos for Verification
|
||||
|
||||
If you add the `--decode` flag, the script will VAE-decode the precomputed latents and save the resulting videos
|
||||
in `.precomputed/decoded_videos`. When audio preprocessing is enabled (`--with-audio`), audio latents will also be
|
||||
decoded and saved to `.precomputed/decoded_audio`. This allows you to visually and audibly inspect the processed data.
|
||||
If you add the `--decode` flag, the script will VAE-decode the precomputed video latents and save the resulting videos
|
||||
in `.precomputed/decoded_videos`. Reference video latents are decoded to `.precomputed/decoded_reference_videos` when
|
||||
present. To inspect audio latents, run `scripts/decode_latents.py` with `--with-audio`.
|
||||
|
||||
```bash
|
||||
uv run python scripts/process_dataset.py dataset.json \
|
||||
@@ -382,6 +486,4 @@ Once your dataset is preprocessed, you can proceed to:
|
||||
- Start training with the [Training Guide](training-guide.md)
|
||||
|
||||
> [!TIP]
|
||||
> If your training recipe requires additional preprocessed data (e.g., masks, conditioning signals), see
|
||||
> [Implementing Custom Training Strategies](custom-training-strategies.md) for guidance on extending the
|
||||
> preprocessing pipeline.
|
||||
> The `flexible` strategy supports masks for inpainting (`mask` condition type) and spatial crop regions for outpainting (`spatial_crop` condition type) out of the box. For other custom preprocessing needs, see [Custom Training Strategies](custom-training-strategies.md).
|
||||
|
||||
@@ -7,12 +7,14 @@ Get up and running with LTX-2 training in just a few steps!
|
||||
Before you begin, ensure you have:
|
||||
|
||||
1. **LTX-2 Model Checkpoint** - A local `.safetensors` file containing the LTX-2 model weights.
|
||||
Download `ltx-2-19b-dev.safetensors` from: [HuggingFace Hub](https://huggingface.co/Lightricks/LTX-2)
|
||||
Download `ltx-2.3-22b-dev.safetensors` from: [HuggingFace Hub](https://huggingface.co/Lightricks/LTX-2.3)
|
||||
The trainer supports LTX-2 and LTX-2.3 checkpoints through the same configuration API; version-specific components
|
||||
are detected from the checkpoint.
|
||||
2. **Gemma Text Encoder** - A local directory containing the Gemma model (required for LTX-2).
|
||||
Download from: [HuggingFace Hub](https://huggingface.co/google/gemma-3-12b-it-qat-q4_0-unquantized/)
|
||||
3. **Linux with CUDA** - The trainer requires `triton` which is Linux-only
|
||||
3. **Linux with CUDA** - The trainer requires `triton` which is Linux-only; CUDA 13+ is recommended
|
||||
4. **GPU with sufficient VRAM** - 80GB recommended for the standard config. For GPUs with 32GB VRAM (e.g., RTX 5090),
|
||||
use the [low VRAM config](../configs/ltx2_av_lora_low_vram.yaml) which enables INT8 quantization and other
|
||||
use the [low VRAM config](../configs/t2v_lora_low_vram.yaml) which enables INT8 quantization and other
|
||||
memory optimizations
|
||||
|
||||
## ⚡ Installation
|
||||
@@ -39,7 +41,18 @@ cd packages/ltx-trainer
|
||||
|
||||
## 🏋 Training Workflow
|
||||
|
||||
### 1. Prepare Your Dataset
|
||||
If you are using an agent-enabled environment with repository skills, you can ask for the
|
||||
[`train-model`](../../../.claude/skills/train-model/SKILL.md) skill to run this workflow with you.
|
||||
It creates a run workspace, confirms the training mode, prepares data, preprocesses latents,
|
||||
launches training, and monitors the run while stopping for approval before expensive steps.
|
||||
|
||||
### 1. Choose a Training Mode
|
||||
|
||||
Start with [`t2v_lora.yaml`](../configs/t2v_lora.yaml) for a first run with videos and captions. For modes such as
|
||||
IC-LoRA, inpainting, or outpainting, check [Training Modes](training-modes.md) first because your metadata needs extra
|
||||
columns such as `reference_video`, `video_mask`, or `audio_mask` before preprocessing.
|
||||
|
||||
### 2. Prepare Your Dataset
|
||||
|
||||
Organize your videos and captions, then preprocess them:
|
||||
|
||||
@@ -57,15 +70,18 @@ uv run python scripts/process_dataset.py dataset.json \
|
||||
--text-encoder-path /path/to/gemma-model
|
||||
```
|
||||
|
||||
By default, preprocessing writes to `.precomputed/`. Use that directory as `data.preprocessed_data_root`
|
||||
in your training config.
|
||||
|
||||
See [Dataset Preparation](dataset-preparation.md) for detailed instructions.
|
||||
|
||||
### 2. Configure Training
|
||||
### 3. Configure Training
|
||||
|
||||
Create or modify a configuration YAML file. Start with one of the example configs:
|
||||
|
||||
- [`configs/ltx2_av_lora.yaml`](../configs/ltx2_av_lora.yaml) - Audio-video LoRA training
|
||||
- [`configs/ltx2_av_lora_low_vram.yaml`](../configs/ltx2_av_lora_low_vram.yaml) - Audio-video LoRA training (optimized for 32GB VRAM)
|
||||
- [`configs/ltx2_v2v_ic_lora.yaml`](../configs/ltx2_v2v_ic_lora.yaml) - IC-LoRA video-to-video
|
||||
- [`configs/t2v_lora.yaml`](../configs/t2v_lora.yaml) - Text-to-video LoRA
|
||||
- [`configs/t2v_lora_low_vram.yaml`](../configs/t2v_lora_low_vram.yaml) - Same as above, tuned for ~32GB VRAM (INT8 quantization and memory optimizations)
|
||||
- [`configs/v2v_ic_lora.yaml`](../configs/v2v_ic_lora.yaml) - IC-LoRA video-to-video
|
||||
|
||||
Key settings to update:
|
||||
|
||||
@@ -82,33 +98,47 @@ output_dir: "outputs/my_training_run"
|
||||
|
||||
See [Configuration Reference](configuration-reference.md) for all available options.
|
||||
|
||||
### 3. Start Training
|
||||
### 4. Start Training
|
||||
|
||||
```bash
|
||||
uv run python scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run python scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
For multi-GPU training:
|
||||
|
||||
```bash
|
||||
uv run accelerate launch scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run accelerate launch scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
See [Training Guide](training-guide.md) for distributed training and advanced options.
|
||||
|
||||
## 🎯 Training Modes
|
||||
|
||||
> [!TIP]
|
||||
> **First time?** Start with [`t2v_lora.yaml`](../configs/t2v_lora.yaml) — it's the simplest mode
|
||||
> and only requires videos with captions. You can explore other modes once you've confirmed your
|
||||
> setup works.
|
||||
|
||||
The trainer supports several training modes:
|
||||
|
||||
| Mode | Description | Config Example |
|
||||
|----------------------|--------------------------------|--------------------------------------------|
|
||||
| **LoRA** | Efficient adapter training | `training_strategy.name: "text_to_video"` |
|
||||
| **Audio-Video LoRA** | Joint audio-video training | `training_strategy.with_audio: true` |
|
||||
| **IC-LoRA** | Video-to-video transformations | `training_strategy.name: "video_to_video"` |
|
||||
| **Full Fine-tuning** | Full model training | `model.training_mode: "full"` |
|
||||
| Mode | Description | Example Config |
|
||||
|-----------------------|--------------------------------------------|-------------------------------------------------------------------|
|
||||
| **Text-to-Video** | Generate video+audio from text prompts | [`t2v_lora.yaml`](../configs/t2v_lora.yaml) |
|
||||
| **Image-to-Video** | Animate from a starting image | [`i2v_lora.yaml`](../configs/i2v_lora.yaml) |
|
||||
| **Video Extension** | Extend videos temporally (forward/backward)| [`video_extend_lora.yaml`](../configs/video_extend_lora.yaml), [`video_suffix_lora.yaml`](../configs/video_suffix_lora.yaml) |
|
||||
| **IC-LoRA (V2V)** | Video-to-video transformations | [`v2v_ic_lora.yaml`](../configs/v2v_ic_lora.yaml) |
|
||||
| **Audio-to-Video** | Generate video conditioned on audio | [`a2v_lora.yaml`](../configs/a2v_lora.yaml) |
|
||||
| **Video-to-Audio** | Generate audio/foley from video | [`v2a_lora.yaml`](../configs/v2a_lora.yaml) |
|
||||
| **Video Inpainting** | Fill in masked regions of video | [`video_inpainting_lora.yaml`](../configs/video_inpainting_lora.yaml) |
|
||||
| **Video Outpainting** | Extend video spatially | [`video_outpainting_lora.yaml`](../configs/video_outpainting_lora.yaml) |
|
||||
| **Text-to-Audio** | Generate audio from text prompts | [`t2a_lora.yaml`](../configs/t2a_lora.yaml) |
|
||||
| **Audio Extension** | Extend audio temporally | [`audio_extend_lora.yaml`](../configs/audio_extend_lora.yaml), [`audio_suffix_lora.yaml`](../configs/audio_suffix_lora.yaml) |
|
||||
| **Audio Inpainting** | Fill in masked regions of audio | [`audio_inpainting_lora.yaml`](../configs/audio_inpainting_lora.yaml) |
|
||||
| **IC-LoRA (A2A)** | Audio-to-audio transformations | [`a2a_ic_lora.yaml`](../configs/a2a_ic_lora.yaml) |
|
||||
| **AV2AV IC-LoRA** | Audio+video IC-LoRA transformations | [`av2av_ic_lora.yaml`](../configs/av2av_ic_lora.yaml) |
|
||||
| **Full Fine-tuning** | Full model training (any mode above) | Set `model.training_mode: "full"` |
|
||||
|
||||
See [Training Modes](training-modes.md) for detailed explanations,
|
||||
or [Custom Training Strategies](custom-training-strategies.md) if you need to implement your own training recipe.
|
||||
See [Training Modes](training-modes.md) for detailed explanations of each mode.
|
||||
|
||||
## Next Steps
|
||||
|
||||
@@ -118,7 +148,7 @@ Once you've completed your first training run, you can:
|
||||
production-ready inference
|
||||
pipelines for various use cases (T2V, I2V, IC-LoRA, etc.). See the package documentation for details.
|
||||
- Learn more about [Dataset Preparation](dataset-preparation.md) for advanced preprocessing
|
||||
- Explore different [Training Modes](training-modes.md) (LoRA, Audio-Video, IC-LoRA)
|
||||
- Explore different [Training Modes](training-modes.md)
|
||||
- Dive deeper into [Training Configuration](configuration-reference.md)
|
||||
- Understand the model architecture in [LTX-Core Documentation](../../ltx-core/README.md)
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ model uploads.
|
||||
After preprocessing your dataset and preparing a configuration file, you can start training using the trainer script:
|
||||
|
||||
```bash
|
||||
uv run python scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run python scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
The trainer will:
|
||||
@@ -19,19 +19,31 @@ The trainer will:
|
||||
4. **Generate validation videos** (if configured)
|
||||
5. **Save the trained weights** in your output directory
|
||||
|
||||
### Agent-Assisted Training
|
||||
|
||||
If your environment supports repository skills, the
|
||||
[`train-model`](../../../.claude/skills/train-model/SKILL.md) skill provides an end-to-end
|
||||
orchestrator for this package. It asks what you want the model to learn, maps that intent to
|
||||
one of the documented [training modes](training-modes.md), probes your filesystem and GPU,
|
||||
prepares/preprocesses the dataset, writes a run-specific config, launches training, and
|
||||
monitors the job. It uses the trainer docs as its source of truth and stops for approval before
|
||||
captioning, preprocessing, or starting expensive training work.
|
||||
|
||||
### Output Files
|
||||
|
||||
**For LoRA training:**
|
||||
|
||||
- `lora_weights.safetensors` - Main LoRA weights file
|
||||
- `checkpoints/lora_weights_step_00000.safetensors` - LoRA checkpoint weights, with the current step in the filename
|
||||
- `training_config.yaml` - Copy of training configuration
|
||||
- `validation_samples/` - Generated validation videos (if enabled)
|
||||
- `samples/` - Generated validation samples (if enabled)
|
||||
- `checkpoints/training_state_step_00000.pt` - Optional resume state, depending on `checkpoints.save_training_state`
|
||||
|
||||
**For full model fine-tuning:**
|
||||
|
||||
- `model_weights.safetensors` - Full model weights
|
||||
- `checkpoints/model_weights_step_00000.safetensors` - Full model checkpoint weights, with the current step in the filename
|
||||
- `training_config.yaml` - Copy of training configuration
|
||||
- `validation_samples/` - Generated validation videos (if enabled)
|
||||
- `samples/` - Generated validation samples (if enabled)
|
||||
- `checkpoints/training_state_step_00000.pt` - Optional resume state, depending on `checkpoints.save_training_state`
|
||||
|
||||
## 🖥️ Distributed / Multi-GPU Training
|
||||
|
||||
@@ -62,22 +74,22 @@ Launch with a specific config using `--config_file`:
|
||||
# DDP (2 GPUs shown as example)
|
||||
CUDA_VISIBLE_DEVICES=0,1 \
|
||||
uv run accelerate launch --config_file configs/accelerate/ddp.yaml \
|
||||
scripts/train.py configs/ltx2_av_lora.yaml
|
||||
scripts/train.py configs/t2v_lora.yaml
|
||||
|
||||
# DDP + torch.compile
|
||||
CUDA_VISIBLE_DEVICES=0,1 \
|
||||
uv run accelerate launch --config_file configs/accelerate/ddp_compile.yaml \
|
||||
scripts/train.py configs/ltx2_av_lora.yaml
|
||||
scripts/train.py configs/t2v_lora.yaml
|
||||
|
||||
# FSDP (4 GPUs shown as example)
|
||||
CUDA_VISIBLE_DEVICES=0,1,2,3 \
|
||||
uv run accelerate launch --config_file configs/accelerate/fsdp.yaml \
|
||||
scripts/train.py configs/ltx2_av_lora.yaml
|
||||
scripts/train.py configs/t2v_lora.yaml
|
||||
|
||||
# FSDP + torch.compile
|
||||
CUDA_VISIBLE_DEVICES=0,1,2,3 \
|
||||
uv run accelerate launch --config_file configs/accelerate/fsdp_compile.yaml \
|
||||
scripts/train.py configs/ltx2_av_lora.yaml
|
||||
scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
**Notes:**
|
||||
@@ -93,13 +105,13 @@ If you prefer to use your default Accelerate profile:
|
||||
|
||||
```bash
|
||||
# Use settings from your default accelerate config
|
||||
uv run accelerate launch scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run accelerate launch scripts/train.py configs/t2v_lora.yaml
|
||||
|
||||
# Override number of processes on the fly (e.g., 2 GPUs)
|
||||
uv run accelerate launch --num_processes 2 scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run accelerate launch --num_processes 2 scripts/train.py configs/t2v_lora.yaml
|
||||
|
||||
# Select specific GPUs
|
||||
CUDA_VISIBLE_DEVICES=0,1 uv run accelerate launch scripts/train.py configs/ltx2_av_lora.yaml
|
||||
CUDA_VISIBLE_DEVICES=0,1 uv run accelerate launch scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
> [!TIP]
|
||||
|
||||
@@ -1,165 +1,167 @@
|
||||
# Training Modes Guide
|
||||
|
||||
The trainer supports several training modes, each suited for different use cases and requirements.
|
||||
The trainer uses the **flexible** training strategy (`name: "flexible"`) — a unified conditioning framework that
|
||||
supports all training modes through configuration. Every scenario is expressed by setting `is_generated` on each
|
||||
modality and adding optional conditions, rather than choosing a separate strategy class.
|
||||
|
||||
## 🎯 Standard LoRA Training (Video-Only)
|
||||
## Key Concepts
|
||||
|
||||
Standard LoRA (Low-Rank Adaptation) training fine-tunes the model by adding small, trainable adapter layers while
|
||||
keeping the base model frozen. This approach:
|
||||
Before diving into individual modes, here are the core ideas behind the flexible strategy:
|
||||
|
||||
- **Requires significantly less memory and compute** than full fine-tuning
|
||||
- **Produces small, portable weight files** (typically a few hundred MB)
|
||||
- **Is ideal for learning specific styles, effects, or concepts**
|
||||
- **Can be easily combined with other LoRAs** during inference
|
||||
- **`is_generated: true`** — the modality is denoised during training and contributes to the loss. This is the
|
||||
modality the model learns to generate.
|
||||
- **`is_generated: false`** — the modality is frozen (sigma=0, no noise, no loss). It passes through the transformer
|
||||
clean and acts as cross-modal conditioning for the generated modality.
|
||||
- **At least one modality must have `is_generated: true`.**
|
||||
- **Conditions** are per-modality and can be composed (e.g., `reference` + `first_frame` together on the video
|
||||
modality).
|
||||
- Audio does **not** support `first_frame` or `spatial_crop` conditions — only `prefix`, `suffix`, `mask`,
|
||||
and `reference`.
|
||||
|
||||
Configure standard LoRA training with:
|
||||
> [!TIP]
|
||||
> If you are using an agent-enabled environment with repository skills and are unsure which mode to choose,
|
||||
> ask for the [`train-model`](../../../.claude/skills/train-model/SKILL.md) skill. It maps your intent to one of
|
||||
> these configs and walks through dataset preparation, preprocessing, launch, and monitoring.
|
||||
|
||||
## 📊 Quick Reference
|
||||
|
||||
| Mode | Video | Audio | Conditions | Config |
|
||||
|-----------------------|-----------|-----------|---------------------|--------|
|
||||
| **T2V** | Generated | Generated | — | [`t2v_lora`](../configs/t2v_lora.yaml) |
|
||||
| **I2V** | Generated | Generated | `first_frame` | [`i2v_lora`](../configs/i2v_lora.yaml) |
|
||||
| **Video Extension** | Generated | Generated | `prefix`/`suffix` | [`video_extend_lora`](../configs/video_extend_lora.yaml) |
|
||||
| **V2V IC-LoRA** | Generated | — | `reference` | [`v2v_ic_lora`](../configs/v2v_ic_lora.yaml) |
|
||||
| **A2V** | Generated | Frozen | — | [`a2v_lora`](../configs/a2v_lora.yaml) |
|
||||
| **V2A (Foley)** | Frozen | Generated | — | [`v2a_lora`](../configs/v2a_lora.yaml) |
|
||||
| **Video Inpainting** | Generated | — | `mask` | [`video_inpainting_lora`](../configs/video_inpainting_lora.yaml) |
|
||||
| **Video Outpainting** | Generated | — | `spatial_crop` | [`video_outpainting_lora`](../configs/video_outpainting_lora.yaml) |
|
||||
| **T2A** | — | Generated | — | [`t2a_lora`](../configs/t2a_lora.yaml) |
|
||||
| **Audio Extension** | — | Generated | `prefix`/`suffix` | [`audio_extend_lora`](../configs/audio_extend_lora.yaml) |
|
||||
| **Audio Inpainting** | — | Generated | `mask` | [`audio_inpainting_lora`](../configs/audio_inpainting_lora.yaml) |
|
||||
| **A2A IC-LoRA** | — | Generated | `reference` | [`a2a_ic_lora`](../configs/a2a_ic_lora.yaml) |
|
||||
| **AV2AV IC-LoRA** | Generated | Generated | `reference` (both) | [`av2av_ic_lora`](../configs/av2av_ic_lora.yaml) |
|
||||
|
||||
---
|
||||
|
||||
## 🎯 Text-to-Video (T2V)
|
||||
|
||||
Generate video and audio from text prompts. Both modalities are denoised with no additional conditions.
|
||||
|
||||
```yaml
|
||||
model:
|
||||
training_mode: "lora"
|
||||
|
||||
training_strategy:
|
||||
name: "text_to_video"
|
||||
first_frame_conditioning_p: 0.1
|
||||
with_audio: false # Video-only training
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
```
|
||||
|
||||
## 🔊 Audio-Video LoRA Training
|
||||
**Example config:** 📄 [t2v_lora.yaml](../configs/t2v_lora.yaml)
|
||||
|
||||
LTX-2 supports joint audio-video generation. You can train LoRA adapters that affect both video and audio output:
|
||||
---
|
||||
|
||||
- **Synchronized audio-video generation** - Audio matches the visual content
|
||||
- **Same efficient LoRA approach** - Just enable audio training
|
||||
- **Requires audio latents** - Dataset must include preprocessed audio
|
||||
## 🖼️ Image-to-Video (I2V)
|
||||
|
||||
Configure audio-video training with:
|
||||
Generate video conditioned on a starting image. The first frame is provided as a clean conditioning signal — no noise,
|
||||
timestep=0, excluded from loss. The `probability` parameter controls how often first-frame conditioning is applied;
|
||||
remaining samples train in pure T2V mode.
|
||||
|
||||
```yaml
|
||||
model:
|
||||
training_mode: "lora"
|
||||
|
||||
training_strategy:
|
||||
name: "text_to_video"
|
||||
first_frame_conditioning_p: 0.1
|
||||
with_audio: true # Enable audio training
|
||||
audio_latents_dir: "audio_latents" # Directory containing audio latents
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
conditions:
|
||||
- type: first_frame
|
||||
probability: 0.5
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
```
|
||||
|
||||
**Example configuration file:**
|
||||
**Example config:** 📄 [i2v_lora.yaml](../configs/i2v_lora.yaml)
|
||||
|
||||
- 📄 [Audio-Video LoRA Training](../configs/ltx2_av_lora.yaml)
|
||||
---
|
||||
|
||||
**Dataset structure for audio-video training:**
|
||||
## ⏩ Video Extension
|
||||
|
||||
Extend a video forward (or backward) in time. Prefix or suffix conditioning provides a span of existing latent frames
|
||||
as clean conditioning. The `temporal_boundary` sets the number of **latent frames** used as context (each latent frame
|
||||
= 8 pixel frames due to temporal compression).
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
conditions:
|
||||
- type: prefix # or "suffix" for backward extension
|
||||
temporal_boundary: 8 # 8 latent frames = 64 pixel frames
|
||||
probability: 1.0
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
```
|
||||
preprocessed_data_root/
|
||||
├── latents/ # Video latents
|
||||
├── conditions/ # Text embeddings
|
||||
└── audio_latents/ # Audio latents (required when with_audio: true)
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> When training audio-video LoRAs, ensure your `target_modules` configuration captures video, audio, and
|
||||
> cross-modal attention branches. Use patterns like `"to_k"` instead of `"attn1.to_k"` to match:
|
||||
> - Video modules: `attn1.to_k`, `attn2.to_k`
|
||||
> - Audio modules: `audio_attn1.to_k`, `audio_attn2.to_k`
|
||||
> - Cross-modal modules: `audio_to_video_attn.to_k`, `video_to_audio_attn.to_k`
|
||||
>
|
||||
> The cross-modal attention modules (`audio_to_video_attn` and `video_to_audio_attn`) enable bidirectional
|
||||
> information flow between audio and video, which is critical for synchronized audiovisual generation.
|
||||
> See [Understanding Target Modules](configuration-reference.md#understanding-target-modules) for detailed guidance.
|
||||
|
||||
> [!NOTE]
|
||||
> You can generate audio during validation even if you're not training the audio branch.
|
||||
> Set `validation.generate_audio: true` independently of `training_strategy.with_audio`.
|
||||
> The `prefix` and `suffix` conditions also work on the audio modality for audio extension.
|
||||
> Set `temporal_boundary` on the audio modality's conditions list to condition on a prefix or suffix
|
||||
> of the audio latents.
|
||||
|
||||
## 🔥 Full Model Fine-tuning
|
||||
**Example configs:** 📄 [video_extend_lora.yaml](../configs/video_extend_lora.yaml) (forward), 📄 [video_suffix_lora.yaml](../configs/video_suffix_lora.yaml) (backward)
|
||||
|
||||
Full model fine-tuning updates all parameters of the base model, providing maximum flexibility but
|
||||
requiring substantial computational resources and larger training datasets:
|
||||
---
|
||||
|
||||
- **Offers the highest potential quality and capability improvements**
|
||||
- **Requires multiple GPUs** and distributed training techniques (e.g., FSDP)
|
||||
- **Produces large checkpoint files** (several GB)
|
||||
- **Best for major model adaptations** or when LoRA limitations are reached
|
||||
## 🔄 IC-LoRA / Video-to-Video (V2V)
|
||||
|
||||
Configure full fine-tuning with:
|
||||
In-Context LoRA learns transformations from paired videos. Pre-encoded reference latents are concatenated to the target
|
||||
sequence — reference tokens participate in bidirectional self-attention but receive no noise and are excluded from loss.
|
||||
This enables control adapters (depth, pose), style transfer, deblurring, colorization, and more.
|
||||
|
||||
```yaml
|
||||
model:
|
||||
training_mode: "full"
|
||||
|
||||
training_strategy:
|
||||
name: "text_to_video"
|
||||
first_frame_conditioning_p: 0.1
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
conditions:
|
||||
- type: reference
|
||||
latents_dir: "reference_latents"
|
||||
probability: 1.0
|
||||
- type: first_frame # optional — composable with reference
|
||||
probability: 0.2
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> Full fine-tuning of LTX-2 requires multiple high-end GPUs (e.g., 4-8× H100 80GB) and distributed
|
||||
> training with FSDP. See [Training Guide](training-guide.md) for multi-GPU setup instructions.
|
||||
> [!NOTE]
|
||||
> IC-LoRA is video-only by default (no audio modality block). Conditions can be composed — the example above also
|
||||
> applies first-frame conditioning with 20% probability alongside the reference.
|
||||
> Use [AV2AV IC-LoRA](#av2av-ic-lora) when both video and audio references should be trained jointly.
|
||||
|
||||
## 🔄 In-Context LoRA (IC-LoRA) Training
|
||||
**Example config:** 📄 [v2v_ic_lora.yaml](../configs/v2v_ic_lora.yaml)
|
||||
|
||||
IC-LoRA is a specialized training mode for video-to-video transformations.
|
||||
Unlike standard training modes that learn from individual videos, IC-LoRA learns transformations from pairs of videos.
|
||||
IC-LoRA enables a wide range of advanced video-to-video applications, such as:
|
||||
### Dataset Requirements
|
||||
|
||||
- **Control adapters** (e.g., Depth, Pose): Learn to map from a control signal (like a depth map or pose skeleton) to a
|
||||
target video
|
||||
- **Video deblurring**: Transform blurry input videos into sharp, high-quality outputs
|
||||
- **Style transfer**: Apply the style of a reference video to a target video sequence
|
||||
- **Colorization**: Convert grayscale reference videos into colorized outputs
|
||||
- **Restoration and enhancement**: Denoise, upscale, or restore old or degraded videos
|
||||
- **Paired videos** — each target video has a corresponding reference video
|
||||
- **Same frame count** between reference and target
|
||||
- Reference videos can optionally be at **lower spatial resolution** (see [Scaled Reference](#scaled-reference-conditioning) below)
|
||||
- Both must be **preprocessed** before training
|
||||
|
||||
By providing paired reference and target videos, IC-LoRA can learn complex transformations that go beyond caption-based
|
||||
conditioning.
|
||||
|
||||
IC-LoRA training fundamentally differs from standard LoRA and full fine-tuning:
|
||||
|
||||
- **Reference videos** provide clean, unnoised conditioning input showing the "before" state
|
||||
- **Target videos** are noised during training and represent the desired "after" state
|
||||
- **The model learns transformations** from reference videos to target videos
|
||||
- **Loss is applied only to the target portion**, not the reference
|
||||
- **Training and inference time increase significantly** due to the doubled sequence length
|
||||
|
||||
To enable IC-LoRA training, configure your YAML file with:
|
||||
|
||||
```yaml
|
||||
model:
|
||||
training_mode: "lora" # Required: IC-LoRA uses LoRA mode
|
||||
|
||||
training_strategy:
|
||||
name: "video_to_video"
|
||||
first_frame_conditioning_p: 0.1
|
||||
reference_latents_dir: "reference_latents" # Directory for reference video latents
|
||||
```
|
||||
|
||||
**Example configuration file:**
|
||||
|
||||
- 📄 [IC-LoRA Training](../configs/ltx2_v2v_ic_lora.yaml) - Video-to-video transformation training
|
||||
|
||||
### Dataset Requirements for IC-LoRA
|
||||
|
||||
- Your dataset must contain **paired videos** where each target video has a corresponding reference video
|
||||
- Reference and target videos must have the **same frame count** (length)
|
||||
- Reference videos can optionally be at **lower spatial resolution** than target videos (
|
||||
see [Scaled Reference Conditioning](#scaled-reference-conditioning) below)
|
||||
- Both reference and target videos should be **preprocessed** before training
|
||||
|
||||
**Dataset structure for IC-LoRA training:**
|
||||
**Dataset structure:**
|
||||
|
||||
```
|
||||
preprocessed_data_root/
|
||||
├── latents/ # Target video latents (what the model learns to generate)
|
||||
├── conditions/ # Text embeddings for each video
|
||||
├── latents/ # Target video latents
|
||||
├── conditions/ # Text embeddings
|
||||
└── reference_latents/ # Reference video latents (conditioning input)
|
||||
```
|
||||
|
||||
### Generating Reference Videos
|
||||
|
||||
We provide an example script to generate reference videos (e.g., Canny edge maps) for a given dataset.
|
||||
The script takes a JSON file as input (e.g., output of `caption_videos.py`) and updates it with the generated reference
|
||||
video paths.
|
||||
Use the `compute_reference.py` script to generate reference videos (e.g., Canny edge maps) for a dataset:
|
||||
|
||||
```bash
|
||||
uv run python scripts/compute_reference.py scenes_output_dir/ \
|
||||
@@ -169,84 +171,392 @@ uv run python scripts/compute_reference.py scenes_output_dir/ \
|
||||
To compute a different condition (depth maps, pose skeletons, etc.), modify the `compute_reference()` function in the
|
||||
script.
|
||||
|
||||
### Configuration Requirements for IC-LoRA
|
||||
|
||||
- You **must** provide `reference_videos` in your validation configuration when using IC-LoRA training
|
||||
- The number of reference videos must match the number of validation prompts
|
||||
|
||||
Example validation configuration for IC-LoRA:
|
||||
|
||||
```yaml
|
||||
validation:
|
||||
prompts:
|
||||
- "First prompt describing the desired output"
|
||||
- "Second prompt describing the desired output"
|
||||
reference_videos:
|
||||
- "/path/to/reference1.mp4"
|
||||
- "/path/to/reference2.mp4"
|
||||
reference_downscale_factor: 1 # Set to match preprocessing (e.g., 2 for half resolution)
|
||||
include_reference_in_output: true # Show reference side-by-side with output
|
||||
```
|
||||
> [!NOTE]
|
||||
> `compute_reference.py` writes generated references to the `reference_video` column, which
|
||||
> `process_dataset.py` detects automatically. The legacy `ref_media_path` column is also accepted.
|
||||
|
||||
### Scaled Reference Conditioning
|
||||
|
||||
For more efficient training and inference, you can use **downscaled reference videos** while keeping target videos at
|
||||
full resolution. This reduces the number of conditioning tokens, leading to:
|
||||
For more efficient training and inference, use **downscaled reference videos** while keeping targets at full
|
||||
resolution. During training, the strategy infers the spatial and temporal scale factors from the preprocessed
|
||||
reference and target latents and adjusts positional encodings accordingly. This reduces conditioning tokens, leading to:
|
||||
|
||||
- **Faster training** due to shorter sequence lengths
|
||||
- **Faster inference** with reduced memory usage
|
||||
- **Faster training** — shorter sequence lengths
|
||||
- **Faster inference** — reduced memory usage
|
||||
- **Same aspect ratio** maintained between reference and target
|
||||
|
||||
#### How It Works
|
||||
|
||||
When the reference video has resolution `H/n × W/n` and the target video has resolution `H × W`, the trainer
|
||||
automatically detects this scale factor `n` and adjusts the positional encodings so that the reference positions
|
||||
map to the correct locations in the target coordinate space.
|
||||
|
||||
#### Preprocessing Datasets with Scaled References
|
||||
|
||||
Use the `--reference-downscale-factor` option when running `process_dataset.py`:
|
||||
Preprocess with the `--reference-downscale-factor` option:
|
||||
|
||||
```bash
|
||||
# Process dataset with scaled reference videos (half resolution)
|
||||
uv run python scripts/process_dataset.py dataset.json \
|
||||
--resolution-buckets 768x768x25 \
|
||||
--model-path /path/to/ltx2.safetensors \
|
||||
--text-encoder-path /path/to/gemma \
|
||||
--reference-column "reference_path" \
|
||||
--reference-downscale-factor 2
|
||||
```
|
||||
|
||||
This will:
|
||||
> [!NOTE]
|
||||
> The `reference_video` column is auto-detected by convention — no `--reference-column` flag needed.
|
||||
|
||||
- Process target videos at 768×768 resolution
|
||||
- Process reference videos at 384×384 resolution (768 / 2)
|
||||
- The trainer will automatically infer the scale factor from the dimension ratio
|
||||
|
||||
**Important**: Set `reference_downscale_factor: 2` in your validation configuration to match the preprocessing:
|
||||
Validation encodes reference media on the fly, so set `downscale_factor` and `temporal_scale_factor`
|
||||
on each `reference` validation condition to match the preprocessing factors:
|
||||
|
||||
```yaml
|
||||
validation:
|
||||
reference_downscale_factor: 2 # Must match the preprocessing factor
|
||||
reference_videos:
|
||||
- "/path/to/reference1.mp4"
|
||||
- "/path/to/reference2.mp4"
|
||||
samples:
|
||||
- prompt: "..."
|
||||
conditions:
|
||||
- type: reference
|
||||
video: "/path/to/reference.mp4"
|
||||
downscale_factor: 2
|
||||
temporal_scale_factor: 1
|
||||
include_in_output: true
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> The scale factor must be a positive integer, and all dimensions must be divisible by 32.
|
||||
> Common scale factors are 1 (no scaling), 2 (half resolution), or 4 (quarter resolution).
|
||||
> Common values are 1 (no scaling), 2 (half resolution), or 4 (quarter resolution).
|
||||
|
||||
## 📊 Training Mode Comparison
|
||||
---
|
||||
|
||||
| Aspect | LoRA | Audio-Video LoRA | Full Fine-tuning | IC-LoRA |
|
||||
|----------------------|--------------------------------|--------------------------------|------------------|--------------------------------|
|
||||
| **Memory Usage** | Low | Low-Medium | High | Medium |
|
||||
| **Training Speed** | Fast | Fast | Slow | Medium |
|
||||
| **Output Size** | 100MB-few GB (depends on rank) | 100MB-few GB (depends on rank) | Tens of GB | 100MB-few GB (depends on rank) |
|
||||
| **Flexibility** | Medium | Medium | High | Specialized |
|
||||
| **Audio Support** | Optional | Yes | Optional | No |
|
||||
| **Reference Videos** | No | No | No | Yes (required) |
|
||||
## 🔊 Audio-to-Video (A2V)
|
||||
|
||||
Generate video conditioned on frozen audio. Audio passes through the transformer clean (sigma=0) and influences video
|
||||
via the built-in cross-modal attention. Only video is denoised.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
audio:
|
||||
is_generated: false
|
||||
latents_dir: "audio_latents"
|
||||
```
|
||||
|
||||
**Example config:** 📄 [a2v_lora.yaml](../configs/a2v_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🎵 Video-to-Audio / Foley (V2A)
|
||||
|
||||
Generate audio (Foley) conditioned on frozen video. Video passes through the transformer clean (sigma=0) and
|
||||
conditions audio via cross-modal attention. Only audio is denoised.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: false
|
||||
latents_dir: "latents"
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
```
|
||||
|
||||
**Example config:** 📄 [v2a_lora.yaml](../configs/v2a_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🎭 Video Inpainting
|
||||
|
||||
Fill in masked regions of a video. Per-sample masks loaded from disk define which tokens are conditioning and which
|
||||
must be generated. Masks are thresholded at `0.5` to match validation/inference: tokens with `mask > 0.5` receive clean
|
||||
latents and timestep=0 and are excluded from loss; tokens with `mask <= 0.5` are denoised normally and contribute to
|
||||
loss.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
conditions:
|
||||
- type: mask
|
||||
mask_dir: "video_masks"
|
||||
probability: 1.0
|
||||
```
|
||||
|
||||
**Dataset structure:**
|
||||
|
||||
```
|
||||
preprocessed_data_root/
|
||||
├── latents/ # Video latents
|
||||
├── conditions/ # Text embeddings
|
||||
└── video_masks/ # Per-sample binary masks (1 → conditioning, 0 → generate)
|
||||
```
|
||||
|
||||
In dataset metadata, provide mask media via the `video_mask` column; preprocessing converts it into `video_masks/`.
|
||||
|
||||
**Example config:** 📄 [video_inpainting_lora.yaml](../configs/video_inpainting_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🌅 Video Outpainting
|
||||
|
||||
Extend a video spatially beyond its original boundaries. A rectangular pixel region is provided as clean conditioning
|
||||
(no noise, timestep=0, excluded from loss) — the model learns to generate the surrounding content. The `spatial_region`
|
||||
is specified in pixel coordinates `[y1, x1, y2, x2]` and automatically converted to latent space.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
conditions:
|
||||
- type: spatial_crop
|
||||
spatial_region: [0, 0, 288, 576] # y1, x1, y2, x2 in pixels
|
||||
probability: 1.0
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> `spatial_crop` is a video-only condition — it is not supported on the audio modality.
|
||||
|
||||
**Example config:** 📄 [video_outpainting_lora.yaml](../configs/video_outpainting_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🔈 Text-to-Audio (T2A)
|
||||
|
||||
Generate audio from text prompts with no video modality. Only the audio branch of the transformer is denoised. Since
|
||||
no video modality is configured, this mode uses **audio-only LoRA targets** — explicitly targeting `audio_attn1`,
|
||||
`audio_attn2`, and `audio_ff` modules.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> With no `video` block in the strategy, the trainer only loads audio latents and text embeddings. LoRA adapters
|
||||
> should explicitly target audio modules (e.g., `audio_attn1.to_k`) rather than short patterns like `to_k` which
|
||||
> would also match video modules. See [LoRA Target Modules Guidance](#lora-target-modules-guidance) below.
|
||||
|
||||
**Example config:** 📄 [t2a_lora.yaml](../configs/t2a_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🔊 Audio Extension
|
||||
|
||||
Extend audio forward (prefix) or backward (suffix) in time — the audio equivalent of Video Extension. A span of
|
||||
existing audio latent frames is provided as clean conditioning, and the model generates the continuation. The
|
||||
`temporal_boundary` sets the number of latent frames used as context. This mode uses **audio-only LoRA targets**.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
conditions:
|
||||
- type: prefix # or "suffix" for backward extension
|
||||
temporal_boundary: 8
|
||||
probability: 1.0
|
||||
```
|
||||
|
||||
**Example configs:** 📄 [audio_extend_lora.yaml](../configs/audio_extend_lora.yaml), 📄 [audio_suffix_lora.yaml](../configs/audio_suffix_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🎭 Audio Inpainting
|
||||
|
||||
Fill in masked regions of audio. Per-sample masks loaded from disk define which audio tokens are conditioning and
|
||||
which must be generated — the audio equivalent of Video Inpainting. Masks are thresholded at `0.5` with the same
|
||||
binary semantics as video inpainting. This mode uses **audio-only LoRA targets**.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
conditions:
|
||||
- type: mask
|
||||
mask_dir: "audio_masks"
|
||||
probability: 1.0
|
||||
```
|
||||
|
||||
**Dataset structure:**
|
||||
|
||||
```
|
||||
preprocessed_data_root/
|
||||
├── conditions/ # Text embeddings
|
||||
├── audio_latents/ # Audio latents
|
||||
└── audio_masks/ # Per-sample binary masks (1 → conditioning, 0 → generate)
|
||||
```
|
||||
|
||||
In dataset metadata, provide mask media via the `audio_mask` column; preprocessing converts it into `audio_masks/`.
|
||||
|
||||
**Example config:** 📄 [audio_inpainting_lora.yaml](../configs/audio_inpainting_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🔄 IC-LoRA / Audio-to-Audio (A2A)
|
||||
|
||||
In-Context LoRA for audio-to-audio transformations. Pre-encoded reference audio latents are concatenated to the target
|
||||
sequence — reference tokens participate in bidirectional self-attention but receive no noise and are excluded from loss.
|
||||
This enables audio style transfer, voice conversion, sound effect transformation, and more. This mode uses
|
||||
**audio-only LoRA targets**.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
conditions:
|
||||
- type: reference
|
||||
latents_dir: "reference_audio_latents"
|
||||
probability: 1.0
|
||||
```
|
||||
|
||||
**Dataset structure:**
|
||||
|
||||
```
|
||||
preprocessed_data_root/
|
||||
├── conditions/ # Text embeddings
|
||||
├── audio_latents/ # Target audio latents
|
||||
└── reference_audio_latents/ # Reference audio latents (conditioning input)
|
||||
```
|
||||
|
||||
**Example config:** 📄 [a2a_ic_lora.yaml](../configs/a2a_ic_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🔄 AV2AV IC-LoRA
|
||||
|
||||
Joint audio-video In-Context LoRA — both modalities have reference conditioning. Pre-encoded reference latents are
|
||||
concatenated to each modality's target sequence independently. This enables joint audiovisual transformations such as
|
||||
synchronized style transfer across both video and audio.
|
||||
|
||||
```yaml
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
conditions:
|
||||
- type: reference
|
||||
latents_dir: "reference_latents"
|
||||
probability: 1.0
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
conditions:
|
||||
- type: reference
|
||||
latents_dir: "reference_audio_latents"
|
||||
probability: 1.0
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> Unlike audio-only IC-LoRA (A2A), AV2AV uses short LoRA target patterns like `"to_k"` to match all branches
|
||||
> (video, audio, and cross-modal attention), since both modalities are trained.
|
||||
|
||||
**Dataset structure:**
|
||||
|
||||
```
|
||||
preprocessed_data_root/
|
||||
├── latents/ # Target video latents
|
||||
├── audio_latents/ # Target audio latents
|
||||
├── conditions/ # Text embeddings
|
||||
├── reference_latents/ # Reference video latents (conditioning input)
|
||||
└── reference_audio_latents/ # Reference audio latents (conditioning input)
|
||||
```
|
||||
|
||||
**Example config:** 📄 [av2av_ic_lora.yaml](../configs/av2av_ic_lora.yaml)
|
||||
|
||||
---
|
||||
|
||||
## 🔥 Full Model Fine-tuning
|
||||
|
||||
All modes above default to `training_mode: "lora"`. For full fine-tuning, set `training_mode: "full"` — this updates
|
||||
all model parameters rather than adding LoRA adapters.
|
||||
|
||||
```yaml
|
||||
model:
|
||||
training_mode: "full"
|
||||
|
||||
training_strategy:
|
||||
name: "flexible"
|
||||
video:
|
||||
is_generated: true
|
||||
latents_dir: "latents"
|
||||
audio:
|
||||
is_generated: true
|
||||
latents_dir: "audio_latents"
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> Full fine-tuning requires multiple high-end GPUs (e.g., 4-8× H100 80GB) and distributed training with FSDP.
|
||||
> See [Training Guide](training-guide.md) for multi-GPU setup instructions.
|
||||
|
||||
---
|
||||
|
||||
## 🎛️ LoRA Target Modules Guidance
|
||||
|
||||
The `target_modules` configuration determines which transformer modules receive LoRA adapters. The right choice depends
|
||||
on whether your training involves cross-modal (audio ↔ video) interaction.
|
||||
|
||||
**For T2V, I2V, A2V, V2A, or any mode involving both modalities** — use short patterns to match all branches
|
||||
(video, audio, and cross-modal attention):
|
||||
|
||||
```yaml
|
||||
target_modules:
|
||||
- "to_k"
|
||||
- "to_q"
|
||||
- "to_v"
|
||||
- "to_out.0"
|
||||
```
|
||||
|
||||
> [!IMPORTANT]
|
||||
> Short patterns like `"to_k"` match video modules (`attn1.to_k`, `attn2.to_k`), audio modules
|
||||
> (`audio_attn1.to_k`, `audio_attn2.to_k`), and cross-modal modules (`audio_to_video_attn.to_k`,
|
||||
> `video_to_audio_attn.to_k`). The cross-modal attention modules enable bidirectional information flow between
|
||||
> audio and video, which is critical for synchronized audiovisual generation.
|
||||
> See [Understanding Target Modules](configuration-reference.md#understanding-target-modules) for detailed guidance.
|
||||
|
||||
**For video-only IC-LoRA** — explicitly target video modules (including FFN layers for better transformation quality):
|
||||
|
||||
```yaml
|
||||
target_modules:
|
||||
- "attn1.to_k"
|
||||
- "attn1.to_q"
|
||||
- "attn1.to_v"
|
||||
- "attn1.to_out.0"
|
||||
- "attn2.to_k"
|
||||
- "attn2.to_q"
|
||||
- "attn2.to_v"
|
||||
- "attn2.to_out.0"
|
||||
- "ff.net.0.proj"
|
||||
- "ff.net.2"
|
||||
```
|
||||
|
||||
**For audio-only modes (T2A, Audio Extension, Audio Inpainting, A2A IC-LoRA)** — explicitly target audio modules:
|
||||
|
||||
```yaml
|
||||
target_modules:
|
||||
- "audio_attn1.to_k"
|
||||
- "audio_attn1.to_q"
|
||||
- "audio_attn1.to_v"
|
||||
- "audio_attn1.to_out.0"
|
||||
- "audio_attn2.to_k"
|
||||
- "audio_attn2.to_q"
|
||||
- "audio_attn2.to_v"
|
||||
- "audio_attn2.to_out.0"
|
||||
- "audio_ff.net.0.proj"
|
||||
- "audio_ff.net.2"
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> Audio-only modes have no `video` block in the strategy, so there is no need to train video or cross-modal
|
||||
> attention modules. Targeting only `audio_*` modules keeps the LoRA small and focused.
|
||||
|
||||
---
|
||||
|
||||
## 🎬 Using Trained Models for Inference
|
||||
|
||||
@@ -255,12 +565,25 @@ LoRAs:
|
||||
|
||||
| Training Mode | Recommended Pipeline |
|
||||
|-------------------------|-------------------------------------------------------|
|
||||
| LoRA / Audio-Video LoRA | `TI2VidOneStagePipeline` or `TI2VidTwoStagesPipeline` |
|
||||
| IC-LoRA | `ICLoraPipeline` |
|
||||
| T2V / I2V / A2V / Extension / Inpainting / Outpainting | `TI2VidOneStagePipeline` or `TI2VidTwoStagesPipeline` |
|
||||
| IC-LoRA (V2V / A2A / AV2AV) | `ICLoraPipeline` |
|
||||
| V2A (Foley) / T2A / Audio Extension / Audio Inpainting | `TI2VidOneStagePipeline` or `TI2VidTwoStagesPipeline` |
|
||||
|
||||
All pipelines support loading custom LoRAs via the `loras` parameter. See the [`ltx-pipelines`](../../ltx-pipelines/)
|
||||
package
|
||||
documentation for detailed usage instructions.
|
||||
package documentation for detailed usage instructions.
|
||||
|
||||
> [!NOTE]
|
||||
> You can generate audio during validation even if you're not training the audio branch.
|
||||
> Set `validation.generate_audio: true` independently of whether audio has `is_generated: true`.
|
||||
|
||||
---
|
||||
|
||||
## 🔄 Migration from Legacy Strategies
|
||||
|
||||
Legacy `text_to_video` and `video_to_video` strategy configs are forward-compatible and will continue to work (with a
|
||||
deprecation warning). We recommend migrating to `flexible` for access to all conditioning modes.
|
||||
|
||||
---
|
||||
|
||||
## 🚀 Next Steps
|
||||
|
||||
@@ -272,6 +595,6 @@ Once you've chosen your training mode:
|
||||
|
||||
> [!TIP]
|
||||
> Need a training mode that's not covered here?
|
||||
> See [Implementing Custom Training Strategies](custom-training-strategies.md)
|
||||
> to learn how to create your own strategy for specialized use cases like video inpainting, audio-only training, or
|
||||
> custom conditioning.
|
||||
> First check whether it can be expressed by composing existing `flexible` conditions. Use
|
||||
> [Implementing Custom Training Strategies](custom-training-strategies.md) only for custom losses,
|
||||
> noising rules, model outputs, or preprocessing that cannot be represented by configuration.
|
||||
|
||||
@@ -8,7 +8,7 @@ Memory management is crucial for successful training with LTX-2.
|
||||
|
||||
> [!TIP]
|
||||
> For GPUs with 32GB VRAM, use the pre-configured low VRAM config:
|
||||
> [`configs/ltx2_av_lora_low_vram.yaml`](../configs/ltx2_av_lora_low_vram.yaml)
|
||||
> [`configs/t2v_lora_low_vram.yaml`](../configs/t2v_lora_low_vram.yaml)
|
||||
> which combines 8-bit optimizer, INT8 quantization, and reduced LoRA rank.
|
||||
|
||||
### Memory Optimization Techniques
|
||||
@@ -111,7 +111,7 @@ Ensure you've installed the dependencies and are using `uv run` to execute scrip
|
||||
# From the repository root
|
||||
uv sync
|
||||
cd packages/ltx-trainer
|
||||
uv run python scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run python scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
> [!TIP]
|
||||
@@ -168,29 +168,29 @@ LTX-2 requires the number of frames to satisfy `frames % 8 == 1`:
|
||||
|
||||
```bash
|
||||
uv run accelerate launch --config_file configs/accelerate/ddp_compile.yaml \
|
||||
scripts/train.py configs/ltx2_av_lora.yaml
|
||||
scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
### Issue: Poor Quality Validation Outputs
|
||||
|
||||
**Solutions:**
|
||||
|
||||
1. **Use Image-to-Video Validation:**
|
||||
For more reliable validation, use image-to-video (first-frame conditioning) rather than pure text-to-video:
|
||||
1. **Use conditioned validation:** For more reliable validation, use image-to-video (first-frame conditioning) rather than pure text-to-video:
|
||||
|
||||
```yaml
|
||||
validation:
|
||||
prompts:
|
||||
- "a professional portrait video of a person"
|
||||
images:
|
||||
- "/path/to/first_frame.png" # One image per prompt
|
||||
samples:
|
||||
- prompt: "a professional portrait video of a person"
|
||||
conditions:
|
||||
- type: first_frame
|
||||
image_or_video: "/path/to/first_frame.png"
|
||||
```
|
||||
|
||||
2. **Increase inference steps:**
|
||||
|
||||
```yaml
|
||||
validation:
|
||||
inference_steps: 50 # Default is 30
|
||||
inference_steps: 30
|
||||
```
|
||||
|
||||
3. **Adjust guidance settings:**
|
||||
|
||||
@@ -35,75 +35,53 @@ uv run python scripts/split_scenes.py video.mp4 scenes/ --max-scenes 50
|
||||
|
||||
### Automatic Video Captioning
|
||||
|
||||
The `scripts/caption_videos.py` script generates captions for videos (with audio) using multimodal models.
|
||||
The `scripts/caption_videos.py` script generates a single, detailed combined audio-visual
|
||||
caption per video as a continuous paragraph of prose. Two backends are available:
|
||||
|
||||
- **`qwen_omni` (default)** — Qwen3-Omni-30B-A3B-Thinking served via a local
|
||||
[vLLM](https://docs.vllm.ai/) HTTP server (~1-3 s/video on H100). Highest quality, runs
|
||||
fully offline once the model is downloaded.
|
||||
- **`gemini_flash`** — Google Gemini (cloud, `gemini-3.5-flash`). No GPU required. Auth is
|
||||
automatic: set `GEMINI_API_KEY` (or `GOOGLE_API_KEY`) for the Developer API, or just have
|
||||
Google Cloud credentials available (`gcloud auth` / an attached service account) and it
|
||||
uses Vertex AI with no extra setup.
|
||||
|
||||
**Step 1 — launch the captioner server** (`qwen_omni` only, one-time).
|
||||
|
||||
`scripts/serve_captioner.py` runs vLLM in an isolated environment via `uvx`, so vLLM's heavy
|
||||
CUDA dependencies never touch the trainer's venv. It defaults to dynamic FP8 quantization
|
||||
(~31 GiB weights, fits on 40 GB GPUs, same speed as BF16 on H100):
|
||||
|
||||
```bash
|
||||
# Generate captions for all videos in a directory (uses Qwen2.5-Omni by default)
|
||||
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json
|
||||
# Terminal 1 - stays running
|
||||
uv run python packages/ltx-trainer/scripts/serve_captioner.py
|
||||
|
||||
# Use 8-bit quantization to reduce VRAM usage
|
||||
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json --use-8bit
|
||||
|
||||
# Use Gemini Flash API instead (requires API key)
|
||||
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json \
|
||||
--captioner-type gemini_flash --api-key YOUR_API_KEY
|
||||
|
||||
# Use Gemini Flash with parallel workers for faster throughput
|
||||
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json \
|
||||
--captioner-type gemini_flash --num-workers 5
|
||||
|
||||
# Caption without audio processing (video-only)
|
||||
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json --no-audio
|
||||
|
||||
# Force re-caption all files
|
||||
uv run python scripts/caption_videos.py videos_dir/ --output dataset.json --override
|
||||
# Useful variants:
|
||||
# --print-cmd show the vLLM command without running it
|
||||
# --quantization bf16 use BF16 instead (needs ~66 GiB free VRAM)
|
||||
# --hf-home /mnt/disk override where the ~65 GB model is downloaded
|
||||
```
|
||||
|
||||
**Key features:**
|
||||
|
||||
- **Audio-visual captioning**: Processes both video and audio content, including speech transcription
|
||||
- **Multiple backends**:
|
||||
- `qwen_omni` (default): Local Qwen2.5-Omni model - processes video + audio locally
|
||||
- `gemini_flash`: Google Gemini Flash API - cloud-based, requires API key
|
||||
- **Parallel captioning** (Gemini Flash only): Use `--num-workers` to run multiple API calls concurrently for faster throughput on large datasets
|
||||
- **Structured output**: Captions include visual description, speech transcription, sounds, and on-screen text
|
||||
- **Memory optimization**: 8-bit quantization option for limited VRAM
|
||||
- **Incremental processing**: Skips already-captioned files by default; progress is saved every 5 videos
|
||||
- **Multiple output formats**: JSON, JSONL, CSV, or TXT
|
||||
|
||||
**Caption format:**
|
||||
|
||||
The captioner produces structured captions with four sections:
|
||||
- `[VISUAL]`: Detailed description of visual content
|
||||
- `[SPEECH]`: Word-for-word transcription of spoken content
|
||||
- `[SOUNDS]`: Description of music, ambient sounds, sound effects
|
||||
- `[TEXT]`: Any on-screen text visible in the video
|
||||
|
||||
**Parallel captioning with Gemini Flash:**
|
||||
|
||||
When using `--captioner-type gemini_flash`, you can speed up large dataset captioning by running multiple API calls at the same time using `--num-workers` (accepts 1–10, default is 1):
|
||||
**Step 2 — caption your videos.**
|
||||
|
||||
```bash
|
||||
export GEMINI_API_KEY="your-key-here"
|
||||
# Terminal 2 - default backend talks to the server above
|
||||
uv run python packages/ltx-trainer/scripts/caption_videos.py videos_dir/ --output dataset.json
|
||||
|
||||
# Caption a large dataset with 5 workers running concurrently
|
||||
uv run python scripts/caption_videos.py videos_dir/ \
|
||||
--output dataset.json \
|
||||
--captioner-type gemini_flash \
|
||||
--num-workers 5
|
||||
# Remote server: --vllm-url http://other-host:8001/v1
|
||||
# Gemini (gemini-3.5-flash): --captioner-type gemini_flash (uses GEMINI_API_KEY, else gcloud/Vertex)
|
||||
# Gemini, parallel calls: --captioner-type gemini_flash --num-workers 5
|
||||
# Re-caption everything: --override
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> `--num-workers` is only supported with `gemini_flash`. Using it with `qwen_omni` or any other local model will raise an error, because local GPU models are not thread-safe.
|
||||
Captioning is incremental (already-captioned files are skipped, progress saves every 5 videos)
|
||||
and writes JSON, JSONL, CSV, or TXT based on the output extension.
|
||||
|
||||
> [!TIP]
|
||||
> Keep `--num-workers` between 3–5 for most use cases. Very high values (8–10) may hit Gemini API rate limits depending on your quota tier.
|
||||
Qwen3-Omni-Thinking can optionally emit a `<think>...</think>` chain-of-thought before the
|
||||
caption (`--enable-thinking`). It is off by default, which is recommended for bulk captioning
|
||||
(thinking is slower as it generates the reasoning trace first).
|
||||
|
||||
**Environment variables (for Gemini Flash):**
|
||||
|
||||
Set one of these to use Gemini Flash without passing `--api-key`:
|
||||
- `GOOGLE_API_KEY`
|
||||
- `GEMINI_API_KEY`
|
||||
For Gemini, keep `--num-workers` at 3-5 (higher values may hit API rate limits).
|
||||
|
||||
### Dataset Preprocessing
|
||||
|
||||
@@ -116,13 +94,6 @@ uv run python scripts/process_dataset.py dataset.json \
|
||||
--model-path /path/to/ltx-2-model.safetensors \
|
||||
--text-encoder-path /path/to/gemma-model
|
||||
|
||||
# With audio processing
|
||||
uv run python scripts/process_dataset.py dataset.json \
|
||||
--resolution-buckets "960x544x49" \
|
||||
--model-path /path/to/ltx-2-model.safetensors \
|
||||
--text-encoder-path /path/to/gemma-model \
|
||||
--with-audio
|
||||
|
||||
# With video decoding for verification
|
||||
uv run python scripts/process_dataset.py dataset.json \
|
||||
--resolution-buckets "960x544x49" \
|
||||
@@ -186,6 +157,10 @@ uv run python scripts/compute_reference.py videos_dir/ --output dataset.json
|
||||
> You can edit this script to generate other types of reference videos for IC-LoRA training,
|
||||
> such as depth maps, segmentation masks, or any custom video transformation.
|
||||
|
||||
> [!NOTE]
|
||||
> `compute_reference.py` writes generated references to the `reference_video` column, which
|
||||
> `process_dataset.py` detects automatically.
|
||||
|
||||
## 🔍 Debugging and Verification Scripts
|
||||
|
||||
### Latents Decoding
|
||||
@@ -224,73 +199,17 @@ uv run python scripts/decode_latents.py /path/to/latents/dir \
|
||||
- **Debug training data**: Visualize what the model actually sees during training
|
||||
- **Quality assessment**: Ensure latent encoding preserves important visual details
|
||||
|
||||
### Inference with Trained Models
|
||||
|
||||
### Inference Script
|
||||
For inference with trained LoRAs, use the [`ltx-pipelines`](../../ltx-pipelines/) package which provides
|
||||
production-ready pipelines:
|
||||
|
||||
The `scripts/inference.py` script runs inference with a trained model.
|
||||
- **Text/Image-to-Video**: `TI2VidOneStagePipeline`, `TI2VidTwoStagesPipeline`
|
||||
- **Distilled (fast) inference**: `DistilledPipeline`
|
||||
- **IC-LoRA video-to-video**: `ICLoraPipeline`
|
||||
- **Keyframe interpolation**: `KeyframeInterpolationPipeline`
|
||||
|
||||
> [!TIP]
|
||||
> For production inference, consider using the [`ltx-pipelines`](../../ltx-pipelines/) package which provides optimized,
|
||||
> feature-rich pipelines for various use cases:
|
||||
> - **Text/Image-to-Video**: `TI2VidOneStagePipeline`, `TI2VidTwoStagesPipeline`
|
||||
> - **Distilled (fast) inference**: `DistilledPipeline`
|
||||
> - **IC-LoRA video-to-video**: `ICLoraPipeline`
|
||||
> - **Keyframe interpolation**: `KeyframeInterpolationPipeline`
|
||||
>
|
||||
> All pipelines support loading custom LoRAs trained with this trainer.
|
||||
|
||||
```bash
|
||||
# Text-to-video inference (with audio by default)
|
||||
# By default, uses CFG scale 4.0 and STG scale 1.0 with block 29
|
||||
uv run python scripts/inference.py \
|
||||
--checkpoint /path/to/model.safetensors \
|
||||
--text-encoder-path /path/to/gemma \
|
||||
--prompt "A cat playing with a ball" \
|
||||
--output output.mp4
|
||||
|
||||
# Video-only (skip audio generation)
|
||||
uv run python scripts/inference.py \
|
||||
--checkpoint /path/to/model.safetensors \
|
||||
--text-encoder-path /path/to/gemma \
|
||||
--prompt "A cat playing with a ball" \
|
||||
--skip-audio \
|
||||
--output output.mp4
|
||||
|
||||
# Image-to-video with conditioning image
|
||||
uv run python scripts/inference.py \
|
||||
--checkpoint /path/to/model.safetensors \
|
||||
--text-encoder-path /path/to/gemma \
|
||||
--prompt "A cat walking" \
|
||||
--condition-image first_frame.png \
|
||||
--output output.mp4
|
||||
|
||||
# Custom guidance settings
|
||||
uv run python scripts/inference.py \
|
||||
--checkpoint /path/to/model.safetensors \
|
||||
--text-encoder-path /path/to/gemma \
|
||||
--prompt "A cat playing with a ball" \
|
||||
--guidance-scale 4.0 \
|
||||
--stg-scale 1.0 \
|
||||
--stg-blocks 29 \
|
||||
--output output.mp4
|
||||
|
||||
# Disable STG (CFG only)
|
||||
uv run python scripts/inference.py \
|
||||
--checkpoint /path/to/model.safetensors \
|
||||
--text-encoder-path /path/to/gemma \
|
||||
--prompt "A cat playing with a ball" \
|
||||
--stg-scale 0.0 \
|
||||
--output output.mp4
|
||||
```
|
||||
|
||||
**Guidance parameters:**
|
||||
|
||||
| Parameter | Default | Description |
|
||||
|-----------|---------|-------------|
|
||||
| `--guidance-scale` | 4.0 | CFG (Classifier-Free Guidance) scale |
|
||||
| `--stg-scale` | 1.0 | STG (Spatio-Temporal Guidance) scale. 0.0 disables STG |
|
||||
| `--stg-blocks` | 29 | Transformer block(s) to perturb for STG |
|
||||
| `--stg-mode` | stg_av | `stg_av` perturbs both audio and video, `stg_v` video only |
|
||||
All pipelines support loading custom LoRAs trained with this trainer.
|
||||
|
||||
## 🚀 Training Scripts
|
||||
|
||||
@@ -300,13 +219,13 @@ Use `scripts/train.py` for both single GPU and multi-GPU runs:
|
||||
|
||||
```bash
|
||||
# Single-GPU training
|
||||
uv run python scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run python scripts/train.py configs/t2v_lora.yaml
|
||||
|
||||
# Multi-GPU (uses your accelerate config)
|
||||
uv run accelerate launch scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run accelerate launch scripts/train.py configs/t2v_lora.yaml
|
||||
|
||||
# Override number of processes
|
||||
uv run accelerate launch --num_processes 4 scripts/train.py configs/ltx2_av_lora.yaml
|
||||
uv run accelerate launch --num_processes 4 scripts/train.py configs/t2v_lora.yaml
|
||||
```
|
||||
|
||||
For detailed usage, see the [Training Guide](training-guide.md).
|
||||
@@ -316,5 +235,5 @@ For detailed usage, see the [Training Guide](training-guide.md).
|
||||
- **Start with `--help`**: Always check available options for each script
|
||||
- **Test on small datasets**: Verify workflows with a few files before processing large datasets
|
||||
- **Use decode verification**: Always decode a few samples to verify preprocessing quality
|
||||
- **Monitor VRAM usage**: Use `--use-8bit` or quantization flags when running into memory issues
|
||||
- **Monitor VRAM usage**: Reach for quantization or lower-memory settings (e.g. FP8 for the captioner server) when running into memory issues
|
||||
- **Keep backups**: Make copies of important dataset files before running conversion scripts
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "ltx-trainer"
|
||||
version = "1.1.3"
|
||||
version = "v1.1.6"
|
||||
description = "LTX-2 training, democratized."
|
||||
readme = "README.md"
|
||||
authors = [
|
||||
@@ -12,9 +12,11 @@ dependencies = [
|
||||
"accelerate>=1.2.1",
|
||||
"av>=14.2.1",
|
||||
"bitsandbytes >=0.45.2; sys_platform == 'linux'",
|
||||
"google-genai>=2.0",
|
||||
"huggingface-hub[hf-xet]>=0.31.4",
|
||||
"imageio>=2.37.0",
|
||||
"imageio-ffmpeg>=0.6.0",
|
||||
"openai>=2.0",
|
||||
"opencv-python>=4.11.0.86",
|
||||
"optimum-quanto>=0.2.6",
|
||||
"pandas>=2.2.3",
|
||||
@@ -25,12 +27,16 @@ dependencies = [
|
||||
"safetensors>=0.5.0",
|
||||
"scenedetect>=0.6.5.2",
|
||||
"sentencepiece>=0.2.0",
|
||||
"soundfile>=0.12.1",
|
||||
"torch>=2.6.0",
|
||||
"torchaudio>=2.7.0",
|
||||
"torchcodec>=0.8.1",
|
||||
# torchcodec must match the torch version (it ships a torch-ABI C++ extension and declares
|
||||
# no torch pin of its own); the 0.9 line matches torch 2.9. torchaudio>=2.9 routes
|
||||
# torchaudio.load() through torchcodec, so audio preprocessing needs it installed.
|
||||
"torchcodec>=0.8.1,<0.10",
|
||||
"torchvision>=0.21.0",
|
||||
"typer>=0.15.1",
|
||||
"wandb>=0.19.11",
|
||||
"wandb>=0.27.0",
|
||||
"setuptools>=79.0.0",
|
||||
]
|
||||
|
||||
@@ -48,8 +54,11 @@ build-backend = "hatchling.build"
|
||||
|
||||
|
||||
[tool.ruff]
|
||||
target-version = "1.1.3"
|
||||
target-version = "v1.1.6"
|
||||
line-length = 120
|
||||
# Restrict isort first-party detection to src/ so stray dirs (e.g. wandb/ run output)
|
||||
# next to pyproject.toml don't get classified as first-party packages. See ruff#10519.
|
||||
src = ["src"]
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = [
|
||||
|
||||
@@ -2,29 +2,28 @@
|
||||
|
||||
"""
|
||||
Auto-caption videos with audio using multimodal models.
|
||||
This script provides a command-line interface for generating captions for videos
|
||||
(including audio) using multimodal models. It supports:
|
||||
- Qwen2.5-Omni: Local model for audio-visual captioning (default)
|
||||
- Gemini Flash: Cloud-based API for audio-visual captioning
|
||||
The paths to videos in the generated dataset/captions file will be RELATIVE to the
|
||||
directory where the output file is stored. This makes the dataset more portable and
|
||||
easier to use in different environments.
|
||||
Backends:
|
||||
- Qwen3-Omni-30B-A3B-Thinking via a local vLLM HTTP server (default,
|
||||
``qwen_omni``). Launch the server once with ``scripts/serve_captioner.py``.
|
||||
- Gemini Flash 3.5 via Google's API (``gemini_flash``).
|
||||
The paths in the output file are RELATIVE to the output file's directory,
|
||||
making the dataset portable.
|
||||
Basic usage:
|
||||
# Caption a single video (includes audio by default)
|
||||
caption_videos.py video.mp4 --output captions.json
|
||||
# Caption all videos in a directory
|
||||
caption_videos.py videos_dir/ --output captions.csv
|
||||
# Caption with custom instruction
|
||||
caption_videos.py video.mp4 --instruction "Describe what happens in this video in detail."
|
||||
# Launch the captioner server once (separate terminal)
|
||||
uv run python scripts/serve_captioner.py
|
||||
# Caption a directory
|
||||
caption_videos.py videos_dir/ --output captions.json
|
||||
# Caption a single video with a custom prompt
|
||||
caption_videos.py video.mp4 --output cap.json --instruction "Describe in detail."
|
||||
Advanced usage:
|
||||
# Use Gemini Flash API (requires GEMINI_API_KEY or GOOGLE_API_KEY env var)
|
||||
# Use Gemini Flash 3.5 (cloud, requires GEMINI_API_KEY)
|
||||
caption_videos.py videos_dir/ --captioner-type gemini_flash
|
||||
# Use Gemini Flash with parallel workers (2-10 workers, cloud API only)
|
||||
# Gemini with parallel workers
|
||||
caption_videos.py videos_dir/ --captioner-type gemini_flash --num-workers 5
|
||||
# Disable audio processing (video-only captions)
|
||||
caption_videos.py videos_dir/ --no-audio
|
||||
# Process videos with specific extensions and save as JSON
|
||||
caption_videos.py videos_dir/ --extensions mp4,mov,avi --output captions.json
|
||||
# Talk to a remote vLLM server
|
||||
caption_videos.py videos_dir/ --vllm-url http://192.168.1.10:8001/v1
|
||||
# Enable Qwen3 chain-of-thought (slower, more detail)
|
||||
caption_videos.py videos_dir/ --enable-thinking
|
||||
"""
|
||||
|
||||
import csv
|
||||
@@ -33,7 +32,6 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from enum import Enum
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import typer
|
||||
from rich.console import Console
|
||||
from rich.progress import (
|
||||
@@ -45,9 +43,14 @@ from rich.progress import (
|
||||
TimeElapsedColumn,
|
||||
TimeRemainingColumn,
|
||||
)
|
||||
from transformers.utils.logging import disable_progress_bar
|
||||
|
||||
from ltx_trainer.captioning import CaptionerType, MediaCaptioningModel, create_captioner
|
||||
from ltx_trainer.captioning import (
|
||||
DEFAULT_QWEN_MODEL,
|
||||
DEFAULT_VLLM_BASE_URL,
|
||||
CaptionerType,
|
||||
MediaCaptioningModel,
|
||||
create_captioner,
|
||||
)
|
||||
|
||||
VIDEO_EXTENSIONS = ["mp4", "avi", "mov", "mkv", "webm"]
|
||||
IMAGE_EXTENSIONS = ["jpg", "jpeg", "png"]
|
||||
@@ -61,8 +64,6 @@ app = typer.Typer(
|
||||
help="Auto-caption videos with audio using multimodal models.",
|
||||
)
|
||||
|
||||
disable_progress_bar()
|
||||
|
||||
|
||||
class OutputFormat(str, Enum):
|
||||
"""Available output formats for captions."""
|
||||
@@ -73,15 +74,13 @@ class OutputFormat(str, Enum):
|
||||
JSONL = "jsonl" # JSON Lines file with one JSON object per line
|
||||
|
||||
|
||||
def caption_media( # noqa: PLR0913
|
||||
def caption_media(
|
||||
input_path: Path,
|
||||
output_path: Path,
|
||||
captioner: MediaCaptioningModel,
|
||||
extensions: list[str],
|
||||
recursive: bool,
|
||||
fps: int,
|
||||
include_audio: bool,
|
||||
clean_caption: bool,
|
||||
output_format: OutputFormat,
|
||||
override: bool,
|
||||
num_workers: int = 1,
|
||||
@@ -94,8 +93,6 @@ def caption_media( # noqa: PLR0913
|
||||
extensions: List of media file extensions to include
|
||||
recursive: Whether to search subdirectories recursively
|
||||
fps: Frames per second to sample from videos (ignored for images)
|
||||
include_audio: Whether to include audio in captioning
|
||||
clean_caption: Whether to clean up captions
|
||||
output_format: Format to save the captions in
|
||||
override: Whether to override existing captions
|
||||
num_workers: Number of parallel workers (only for cloud-based captioners like Gemini)
|
||||
@@ -149,10 +146,10 @@ def caption_media( # noqa: PLR0913
|
||||
caption = captioner.caption(
|
||||
path=media_file,
|
||||
fps=fps,
|
||||
include_audio=include_audio,
|
||||
clean_caption=clean_caption,
|
||||
)
|
||||
rel_path = str(media_file.resolve().relative_to(base_dir))
|
||||
# Don't resolve the file itself, so a symlinked clip keeps its logical path under the
|
||||
# dataset dir instead of jumping to its (possibly external) link target.
|
||||
rel_path = str((media_file.parent.resolve() / media_file.name).relative_to(base_dir))
|
||||
return rel_path, caption
|
||||
|
||||
with progress:
|
||||
@@ -371,16 +368,31 @@ def main( # noqa: PLR0913
|
||||
help="Type of captioner to use. Valid values: 'qwen_omni' (local), 'gemini_flash' (API)",
|
||||
case_sensitive=False,
|
||||
),
|
||||
device: str | None = typer.Option(
|
||||
None,
|
||||
"--device",
|
||||
"-d",
|
||||
help="Device to use for inference (e.g., 'cuda', 'cuda:0', 'cpu'). Only for local models.",
|
||||
vllm_url: str = typer.Option(
|
||||
DEFAULT_VLLM_BASE_URL,
|
||||
"--vllm-url",
|
||||
help=(
|
||||
"Base URL of the vLLM OpenAI-compatible server (qwen_omni only). "
|
||||
"Launch the server with `uv run python scripts/serve_captioner.py`."
|
||||
),
|
||||
),
|
||||
use_8bit: bool = typer.Option(
|
||||
vllm_model: str = typer.Option(
|
||||
DEFAULT_QWEN_MODEL,
|
||||
"--vllm-model",
|
||||
help="Served model identifier on the vLLM server (qwen_omni only).",
|
||||
),
|
||||
enable_thinking: bool = typer.Option(
|
||||
False,
|
||||
"--use-8bit",
|
||||
help="Whether to use 8-bit precision for the captioning model (reduces memory usage)",
|
||||
"--enable-thinking/--no-thinking",
|
||||
help=(
|
||||
"Let Qwen3-Omni produce a <think>...</think> chain-of-thought before the caption. "
|
||||
"Off by default: ~5x slower with marginal quality benefit and occasional hallucinations."
|
||||
),
|
||||
),
|
||||
max_tokens: int = typer.Option(
|
||||
4096,
|
||||
"--max-tokens",
|
||||
help="Maximum new tokens to generate per caption (qwen_omni only).",
|
||||
),
|
||||
instruction: str | None = typer.Option(
|
||||
None,
|
||||
@@ -401,20 +413,14 @@ def main( # noqa: PLR0913
|
||||
help="Search for media files in subdirectories recursively",
|
||||
),
|
||||
fps: int = typer.Option(
|
||||
3,
|
||||
2,
|
||||
"--fps",
|
||||
"-f",
|
||||
help="Frames per second to sample from videos (ignored for images)",
|
||||
),
|
||||
include_audio: bool = typer.Option(
|
||||
True,
|
||||
"--audio/--no-audio",
|
||||
help="Whether to include audio in captioning (for videos with audio tracks)",
|
||||
),
|
||||
clean_caption: bool = typer.Option(
|
||||
True,
|
||||
"--clean-caption/--raw-caption",
|
||||
help="Whether to clean up captions by removing common VLM patterns",
|
||||
help=(
|
||||
"Frames per second to sample from videos. 2 is a typical default; "
|
||||
"lower values use less compute per video. Ignored for images and for the "
|
||||
"Gemini backend (which decides its own sampling rate)."
|
||||
),
|
||||
),
|
||||
override: bool = typer.Option(
|
||||
False,
|
||||
@@ -441,35 +447,36 @@ def main( # noqa: PLR0913
|
||||
),
|
||||
) -> None:
|
||||
"""Auto-caption videos with audio using multimodal models.
|
||||
This script supports audio-visual captioning using:
|
||||
- Qwen2.5-Omni: Local model (default) - processes both video and audio
|
||||
- Gemini Flash: Cloud API - requires GOOGLE_API_KEY environment variable
|
||||
Backends:
|
||||
- ``qwen_omni`` (default): Qwen3-Omni-30B-A3B-Thinking via a local vLLM
|
||||
HTTP server. Launch the server once in a separate terminal with
|
||||
``uv run python scripts/serve_captioner.py``. The server stays loaded
|
||||
across script invocations.
|
||||
- ``gemini_flash``: Google Gemini (``gemini-3.5-flash``) via the google-genai SDK.
|
||||
Auth is automatic -- ``GEMINI_API_KEY``/``GOOGLE_API_KEY`` for the Developer API,
|
||||
or Google Cloud credentials (gcloud / service account) for Vertex AI with no env vars.
|
||||
The paths in the output file will be relative to the output file's directory.
|
||||
Examples:
|
||||
# Caption videos with audio using Qwen2.5-Omni (default)
|
||||
# Caption videos using the local vLLM server (default)
|
||||
caption_videos.py videos_dir/ -o captions.json
|
||||
# Caption using Gemini Flash API
|
||||
# Point at a remote vLLM server
|
||||
caption_videos.py videos_dir/ -o captions.json --vllm-url http://other-host:8001/v1
|
||||
# Caption using Gemini Flash 3.5
|
||||
caption_videos.py videos_dir/ -o captions.json -c gemini_flash
|
||||
# Caption without audio (video-only)
|
||||
caption_videos.py videos_dir/ -o captions.json --no-audio
|
||||
# Caption with custom instruction
|
||||
caption_videos.py video.mp4 -o captions.json -i "Describe this video in detail"
|
||||
"""
|
||||
|
||||
# Parallel workers are only safe for cloud-based (stateless) captioners.
|
||||
# Local models like Qwen-Omni hold GPU state and are not thread-safe.
|
||||
# Parallel workers are only supported for the cloud Gemini backend; qwen_omni
|
||||
# drives a single shared vLLM server and is captioned serially from here.
|
||||
if num_workers > 1 and captioner_type != CaptionerType.GEMINI_FLASH:
|
||||
console.print(
|
||||
"[bold red]Error:[/] --num-workers > 1 is only supported with [bold]--captioner-type gemini_flash[/].\n"
|
||||
"Local models (e.g. qwen_omni) run on GPU and are not thread-safe — "
|
||||
"parallel calls would cause memory corruption or incorrect results.\n"
|
||||
"Either set [bold]--num-workers 1[/] (default) or switch to [bold]--captioner-type gemini_flash[/]."
|
||||
"[bold red]Error:[/] --num-workers > 1 is only supported with "
|
||||
"[bold]--captioner-type gemini_flash[/]. Use --num-workers 1 (default) "
|
||||
"for the qwen_omni backend."
|
||||
)
|
||||
raise typer.Exit(code=1)
|
||||
|
||||
# Determine device for local models
|
||||
device_str = device or ("cuda" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
# Parse extensions
|
||||
ext_list = [ext.strip() for ext in extensions.split(",")]
|
||||
|
||||
@@ -490,14 +497,15 @@ def main( # noqa: PLR0913
|
||||
output = Path(output).resolve()
|
||||
console.print(f"Output will be saved to [bold blue]{output}[/]")
|
||||
|
||||
# Initialize captioning model
|
||||
with console.status("Loading captioning model...", spinner="dots"):
|
||||
with console.status("Initializing captioner...", spinner="dots"):
|
||||
if captioner_type == CaptionerType.QWEN_OMNI:
|
||||
captioner = create_captioner(
|
||||
captioner_type=captioner_type,
|
||||
device=device_str,
|
||||
use_8bit=use_8bit,
|
||||
base_url=vllm_url,
|
||||
model=vllm_model,
|
||||
instruction=instruction,
|
||||
max_tokens=max_tokens,
|
||||
enable_thinking=enable_thinking,
|
||||
)
|
||||
elif captioner_type == CaptionerType.GEMINI_FLASH:
|
||||
captioner = create_captioner(
|
||||
@@ -508,7 +516,7 @@ def main( # noqa: PLR0913
|
||||
else:
|
||||
raise ValueError(f"Unsupported captioner type: {captioner_type}")
|
||||
|
||||
console.print(f"[bold green]✓[/] {captioner_type.value} captioning model loaded successfully")
|
||||
console.print(f"[bold green]✓[/] {captioner_type.value} captioner ready")
|
||||
|
||||
# Caption media files
|
||||
caption_media(
|
||||
@@ -518,8 +526,6 @@ def main( # noqa: PLR0913
|
||||
extensions=ext_list,
|
||||
recursive=recursive,
|
||||
fps=fps,
|
||||
include_audio=include_audio,
|
||||
clean_caption=clean_caption,
|
||||
output_format=output_format,
|
||||
override=override,
|
||||
num_workers=num_workers,
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
Compute reference videos for IC-LoRA training.
|
||||
This script provides a command-line interface for generating reference videos to be used for IC-LoRA training.
|
||||
Note that it reads and writes to the same file (the output of caption_videos.py),
|
||||
where it adds the "reference_path" field to the JSON.
|
||||
where it adds the "reference_video" field to the JSON.
|
||||
Basic usage:
|
||||
# Compute reference videos for all videos in a directory
|
||||
compute_reference.py videos_dir/ --output videos_dir/captions.json
|
||||
@@ -11,7 +11,7 @@ Basic usage:
|
||||
# Standard library imports
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict
|
||||
from typing import Any
|
||||
|
||||
# Third-party imports
|
||||
import cv2
|
||||
@@ -37,6 +37,10 @@ from ltx_trainer.video_utils import read_video, save_video
|
||||
console = Console()
|
||||
disable_progress_bar()
|
||||
|
||||
VIDEO_COLUMNS = ("video", "media_path")
|
||||
REFERENCE_VIDEO_COLUMN = "reference_video"
|
||||
LEGACY_REFERENCE_COLUMN = "reference_path"
|
||||
|
||||
|
||||
def compute_reference(
|
||||
images: torch.Tensor,
|
||||
@@ -79,15 +83,15 @@ def compute_reference(
|
||||
|
||||
def _get_meta_data(
|
||||
output_path: Path,
|
||||
) -> Dict[str, str]:
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Get set of existing reference video paths without loading the actual files.
|
||||
Args:
|
||||
output_path: Path to the reference video paths file
|
||||
Returns:
|
||||
Dictionary mapping media paths to reference video paths
|
||||
Dataset rows with media paths and captions
|
||||
"""
|
||||
if not output_path.exists():
|
||||
return {}
|
||||
return []
|
||||
|
||||
console.print(f"[bold blue]Reading meta data from [cyan]{output_path}[/]...[/]")
|
||||
|
||||
@@ -98,11 +102,18 @@ def _get_meta_data(
|
||||
|
||||
except Exception as e:
|
||||
console.print(f"[bold yellow]Warning: Could not check meta data: {e}[/]")
|
||||
return {}
|
||||
return []
|
||||
|
||||
|
||||
def _get_media_path(item: dict[str, Any]) -> str:
|
||||
for column in VIDEO_COLUMNS:
|
||||
if column in item:
|
||||
return item[column]
|
||||
raise KeyError(f"Dataset row must contain one of {VIDEO_COLUMNS}")
|
||||
|
||||
|
||||
def _save_dataset_json(
|
||||
reference_paths: Dict[str, str],
|
||||
reference_paths: dict[str, str],
|
||||
output_path: Path,
|
||||
) -> None:
|
||||
"""Save dataset json with reference video paths.
|
||||
@@ -115,17 +126,17 @@ def _save_dataset_json(
|
||||
json_data = json.load(f)
|
||||
new_json_data = json_data.copy()
|
||||
for i, item in enumerate(json_data):
|
||||
media_path = item["media_path"]
|
||||
media_path = _get_media_path(item)
|
||||
reference_path = reference_paths[media_path]
|
||||
new_json_data[i]["reference_path"] = reference_path
|
||||
new_json_data[i].pop(LEGACY_REFERENCE_COLUMN, None)
|
||||
new_json_data[i][REFERENCE_VIDEO_COLUMN] = reference_path
|
||||
|
||||
with output_path.open("w", encoding="utf-8") as f:
|
||||
json.dump(new_json_data, f, indent=2, ensure_ascii=False)
|
||||
|
||||
console.print(f"[bold green]✓[/] Reference video paths saved to [cyan]{output_path}[/]")
|
||||
console.print("[bold yellow]Note:[/] Use these files with ImageOrVideoDataset by setting:")
|
||||
console.print(" reference_column='[cyan]reference_path[/]'")
|
||||
console.print(" video_column='[cyan]media_path[/]'")
|
||||
console.print("[bold yellow]Note:[/] Reference videos were written to the '[cyan]reference_video[/]' column.")
|
||||
console.print(" [cyan]process_dataset.py[/] detects this column automatically for IC-LoRA preprocessing.")
|
||||
|
||||
|
||||
def process_media(
|
||||
@@ -158,7 +169,7 @@ def process_media(
|
||||
def media_path_to_reference_path(media_file: Path) -> Path:
|
||||
return media_file.parent / (media_file.stem + "_reference" + media_file.suffix)
|
||||
|
||||
media_files = [base_dir / Path(sample["media_path"]) for sample in meta_data]
|
||||
media_files = [base_dir / Path(_get_media_path(sample)) for sample in meta_data]
|
||||
for media_file in media_files:
|
||||
reference_path = media_path_to_reference_path(media_file)
|
||||
media_to_process.append(media_file)
|
||||
@@ -178,18 +189,23 @@ def process_media(
|
||||
)
|
||||
|
||||
# Process media files
|
||||
media_paths = [item["media_path"] for item in meta_data]
|
||||
media_paths = [_get_media_path(item) for item in meta_data]
|
||||
reference_paths = {rel_path: str(media_path_to_reference_path(Path(rel_path))) for rel_path in media_paths}
|
||||
|
||||
with progress:
|
||||
task = progress.add_task("Computing condition on videos", total=len(media_to_process))
|
||||
|
||||
for media_file in media_to_process:
|
||||
for media_file, rel_path in zip(media_to_process, media_paths, strict=True):
|
||||
progress.update(task, description=f"Processing [bold blue]{media_file.name}[/]")
|
||||
|
||||
rel_path = str(media_file.resolve().relative_to(base_dir))
|
||||
# Key by the original media-path string (matches the dict seeded above). Avoid
|
||||
# resolve()/relative_to here — they crash on symlinked or absolute media paths.
|
||||
reference_path = media_path_to_reference_path(media_file)
|
||||
reference_paths[rel_path] = str(reference_path.relative_to(base_dir))
|
||||
try:
|
||||
ref_stored = str(reference_path.relative_to(base_dir))
|
||||
except ValueError:
|
||||
ref_stored = str(reference_path) # absolute/out-of-tree: keep it next to the source
|
||||
reference_paths[rel_path] = ref_stored
|
||||
|
||||
if not reference_path.resolve().exists() or override:
|
||||
try:
|
||||
|
||||
@@ -310,7 +310,7 @@ def main(
|
||||
help="Device to use for computation",
|
||||
),
|
||||
vae_tiling: bool = typer.Option(
|
||||
default=False,
|
||||
default=True,
|
||||
help="Enable VAE tiling for larger video resolutions",
|
||||
),
|
||||
seed: int | None = typer.Option(
|
||||
|
||||
@@ -1,443 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# ruff: noqa: T201
|
||||
"""
|
||||
CLI script for running LTX video/audio generation inference.
|
||||
Usage:
|
||||
# Text-to-Video + Audio (default behavior)
|
||||
python scripts/inference.py --checkpoint path/to/model.safetensors \
|
||||
--text-encoder-path path/to/gemma \
|
||||
--prompt "A cat playing with a ball" --output output.mp4
|
||||
# Video only (skip audio)
|
||||
python scripts/inference.py --checkpoint path/to/model.safetensors \
|
||||
--text-encoder-path path/to/gemma \
|
||||
--prompt "A cat playing with a ball" --skip-audio --output output.mp4
|
||||
# Image-to-Video
|
||||
python scripts/inference.py --checkpoint path/to/model.safetensors \
|
||||
--text-encoder-path path/to/gemma \
|
||||
--prompt "A cat walking" --condition-image first_frame.png --output output.mp4
|
||||
# Video-to-Video (IC-LoRA style)
|
||||
python scripts/inference.py --checkpoint path/to/model.safetensors \
|
||||
--text-encoder-path path/to/gemma \
|
||||
--prompt "A cat turning into a dog" --reference-video input.mp4 --output output.mp4
|
||||
# With LoRA weights
|
||||
python scripts/inference.py --checkpoint path/to/model.safetensors \
|
||||
--text-encoder-path path/to/gemma \
|
||||
--lora-path path/to/lora.safetensors \
|
||||
--prompt "A cat in my custom style" --output output.mp4
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
from peft import LoraConfig, get_peft_model, set_peft_model_state_dict
|
||||
from safetensors.torch import load_file
|
||||
from torchvision import transforms
|
||||
|
||||
from ltx_trainer.model_loader import load_model
|
||||
from ltx_trainer.progress import StandaloneSamplingProgress
|
||||
from ltx_trainer.utils import open_image_as_srgb
|
||||
from ltx_trainer.validation_sampler import GenerationConfig, ValidationSampler
|
||||
from ltx_trainer.video_utils import read_video, save_video
|
||||
|
||||
|
||||
def load_image(image_path: str) -> torch.Tensor:
|
||||
"""Load an image and convert to tensor [C, H, W] in [0, 1]."""
|
||||
image = open_image_as_srgb(image_path)
|
||||
transform = transforms.ToTensor()
|
||||
return transform(image)
|
||||
|
||||
|
||||
def extract_lora_target_modules(state_dict: dict[str, torch.Tensor]) -> list[str]:
|
||||
"""Extract target module names from LoRA checkpoint keys.
|
||||
LoRA keys follow the pattern (after removing "diffusion_model." prefix):
|
||||
- transformer_blocks.0.attn1.to_k.lora_A.weight
|
||||
- transformer_blocks.0.ff.net.0.proj.lora_B.weight
|
||||
This extracts the full module path like "transformer_blocks.0.attn1.to_k".
|
||||
Using full paths is more robust than partial patterns.
|
||||
"""
|
||||
target_modules = set()
|
||||
# Pattern to extract everything before .lora_A or .lora_B
|
||||
pattern = re.compile(r"(.+)\.lora_[AB]\.")
|
||||
|
||||
for key in state_dict:
|
||||
match = pattern.match(key)
|
||||
if match:
|
||||
module_path = match.group(1)
|
||||
target_modules.add(module_path)
|
||||
|
||||
return sorted(target_modules)
|
||||
|
||||
|
||||
def load_lora_weights(transformer: torch.nn.Module, lora_path: str | Path) -> torch.nn.Module:
|
||||
"""Load LoRA weights into the transformer model.
|
||||
The LoRA rank and target modules are automatically detected from the checkpoint.
|
||||
Alpha is set equal to rank (standard practice for inference).
|
||||
Args:
|
||||
transformer: The base transformer model
|
||||
lora_path: Path to the LoRA weights (.safetensors)
|
||||
Returns:
|
||||
The transformer model with LoRA weights applied
|
||||
"""
|
||||
print(f"Loading LoRA weights from {lora_path}...")
|
||||
|
||||
# Load the LoRA state dict
|
||||
state_dict = load_file(str(lora_path))
|
||||
|
||||
# Remove "diffusion_model." prefix (ComfyUI-compatible format)
|
||||
state_dict = {k.replace("diffusion_model.", "", 1): v for k, v in state_dict.items()}
|
||||
|
||||
# Extract target modules from the checkpoint
|
||||
target_modules = extract_lora_target_modules(state_dict)
|
||||
if not target_modules:
|
||||
raise ValueError(f"Could not extract target modules from LoRA checkpoint: {lora_path}")
|
||||
print(f" Detected {len(target_modules)} target modules")
|
||||
|
||||
# Auto-detect rank from the first lora_A weight shape
|
||||
lora_rank = None
|
||||
for key, value in state_dict.items():
|
||||
if "lora_A" in key and value.ndim == 2:
|
||||
lora_rank = value.shape[0]
|
||||
break
|
||||
if lora_rank is None:
|
||||
raise ValueError("Could not auto-detect LoRA rank from weights")
|
||||
print(f" LoRA rank: {lora_rank}")
|
||||
|
||||
# Create LoRA config and wrap the model
|
||||
# Alpha = rank is standard for inference (maintains the trained scale)
|
||||
lora_config = LoraConfig(
|
||||
r=lora_rank,
|
||||
lora_alpha=lora_rank,
|
||||
target_modules=target_modules,
|
||||
lora_dropout=0.0,
|
||||
init_lora_weights=True,
|
||||
)
|
||||
|
||||
# Wrap the transformer with PEFT to add LoRA layers
|
||||
transformer = get_peft_model(transformer, lora_config)
|
||||
|
||||
# Load the LoRA weights
|
||||
base_model = transformer.get_base_model()
|
||||
set_peft_model_state_dict(base_model, state_dict)
|
||||
|
||||
print("✓ LoRA weights loaded successfully")
|
||||
return transformer
|
||||
|
||||
|
||||
def main() -> None: # noqa: PLR0912, PLR0915
|
||||
parser = argparse.ArgumentParser(
|
||||
description="LTX Video/Audio Generation",
|
||||
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
parser.add_argument(
|
||||
"--checkpoint",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to model checkpoint (.safetensors)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-encoder-path",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Path to Gemma text encoder directory",
|
||||
)
|
||||
|
||||
# LoRA arguments
|
||||
parser.add_argument(
|
||||
"--lora-path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to LoRA weights (.safetensors)",
|
||||
)
|
||||
|
||||
# Generation arguments
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Text prompt for generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--negative-prompt",
|
||||
type=str,
|
||||
default="",
|
||||
help="Negative prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=544,
|
||||
help="Video height (must be divisible by 32)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=960,
|
||||
help="Video width (must be divisible by 32)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=97,
|
||||
help="Number of video frames (must be k*8 + 1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--frame-rate",
|
||||
type=float,
|
||||
default=25.0,
|
||||
help="Video frame rate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-inference-steps",
|
||||
type=int,
|
||||
default=30,
|
||||
help="Number of denoising steps",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-scale",
|
||||
type=float,
|
||||
default=4.0,
|
||||
help="Classifier-free guidance scale (CFG)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stg-scale",
|
||||
type=float,
|
||||
default=1.0,
|
||||
help="STG (Spatio-Temporal Guidance) scale. 0.0 disables STG. Default: 1.0",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stg-blocks",
|
||||
type=int,
|
||||
nargs="*",
|
||||
default=[29],
|
||||
help="Which transformer blocks to perturb for STG. Default: 29 (single block).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stg-mode",
|
||||
type=str,
|
||||
default="stg_av",
|
||||
choices=["stg_av", "stg_v"],
|
||||
help="STG mode: 'stg_av' perturbs both audio and video, 'stg_v' perturbs video only",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=42,
|
||||
help="Random seed for reproducibility",
|
||||
)
|
||||
|
||||
# Conditioning arguments
|
||||
parser.add_argument(
|
||||
"--condition-image",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to conditioning image for image-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--reference-video",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to reference video for video-to-video generation (IC-LoRA style)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--include-reference-in-output",
|
||||
action="store_true",
|
||||
help="Include reference video side-by-side with generated output (only for V2V)",
|
||||
)
|
||||
|
||||
# Audio arguments
|
||||
parser.add_argument(
|
||||
"--skip-audio",
|
||||
action="store_true",
|
||||
help="Skip audio generation (by default, audio is generated alongside video)",
|
||||
)
|
||||
|
||||
# Output arguments
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
type=str,
|
||||
required=True,
|
||||
help="Output video path (.mp4)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--audio-output",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Output audio path (.wav, optional - if not provided, audio will be embedded in video)",
|
||||
)
|
||||
|
||||
# Device arguments
|
||||
parser.add_argument(
|
||||
"--device",
|
||||
type=str,
|
||||
default="cuda",
|
||||
help="Device to run on (cuda/cpu)",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Validate conditioning arguments
|
||||
if args.include_reference_in_output and args.reference_video is None:
|
||||
parser.error("--include-reference-in-output requires --reference-video")
|
||||
|
||||
# Validate arguments
|
||||
generate_audio = not args.skip_audio
|
||||
|
||||
print("=" * 80)
|
||||
print("LTX Video/Audio Generation")
|
||||
print("=" * 80)
|
||||
|
||||
# Determine if we need VAE encoder (for image or video conditioning)
|
||||
need_vae_encoder = args.condition_image is not None or args.reference_video is not None
|
||||
|
||||
components = load_model(
|
||||
checkpoint_path=args.checkpoint,
|
||||
device="cpu", # Load to CPU first, sampler will move to device as needed
|
||||
dtype=torch.bfloat16,
|
||||
with_video_vae_encoder=need_vae_encoder,
|
||||
with_video_vae_decoder=True,
|
||||
with_audio_vae_decoder=generate_audio,
|
||||
with_vocoder=generate_audio,
|
||||
with_text_encoder=True,
|
||||
text_encoder_path=args.text_encoder_path,
|
||||
)
|
||||
|
||||
# Apply LoRA weights if provided
|
||||
transformer = components.transformer
|
||||
if args.lora_path is not None:
|
||||
transformer = load_lora_weights(transformer, args.lora_path)
|
||||
|
||||
# Load conditioning image if provided
|
||||
condition_image = None
|
||||
if args.condition_image:
|
||||
print(f"Loading conditioning image from {args.condition_image}...")
|
||||
condition_image = load_image(args.condition_image)
|
||||
|
||||
# Load reference video if provided
|
||||
reference_video = None
|
||||
if args.reference_video:
|
||||
print(f"Loading reference video from {args.reference_video}...")
|
||||
reference_video, ref_fps = read_video(args.reference_video, max_frames=args.num_frames)
|
||||
print(f" Loaded {reference_video.shape[0]} frames @ {ref_fps:.1f} fps")
|
||||
|
||||
# Determine generation mode
|
||||
if args.reference_video is not None and args.condition_image is not None:
|
||||
mode = "Video-to-Video + Image Conditioning (V2V+I2V)"
|
||||
elif args.reference_video is not None:
|
||||
mode = "Video-to-Video (V2V)"
|
||||
elif args.condition_image is not None:
|
||||
mode = "Image-to-Video (I2V)"
|
||||
else:
|
||||
mode = "Text-to-Video (T2V)"
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("Generation Parameters")
|
||||
print("=" * 80)
|
||||
print(f"Mode: {mode}")
|
||||
print(f"Prompt: {args.prompt}")
|
||||
if args.negative_prompt:
|
||||
print(f"Negative prompt: {args.negative_prompt}")
|
||||
print(f"Resolution: {args.width}x{args.height}")
|
||||
print(f"Frames: {args.num_frames} @ {args.frame_rate} fps")
|
||||
print(f"Inference steps: {args.num_inference_steps}")
|
||||
print(f"CFG scale: {args.guidance_scale}")
|
||||
if args.stg_scale > 0:
|
||||
blocks_str = args.stg_blocks if args.stg_blocks else "all"
|
||||
print(f"STG scale: {args.stg_scale} (mode: {args.stg_mode}, blocks: {blocks_str})")
|
||||
else:
|
||||
print("STG: disabled")
|
||||
print(f"Seed: {args.seed}")
|
||||
if args.lora_path:
|
||||
print(f"LoRA: {args.lora_path}")
|
||||
if condition_image is not None:
|
||||
print(f"Conditioning: Image ({args.condition_image})")
|
||||
if reference_video is not None:
|
||||
print(f"Reference: Video ({args.reference_video})")
|
||||
if args.include_reference_in_output:
|
||||
print(" → Will include reference side-by-side in output")
|
||||
if generate_audio:
|
||||
video_duration = args.num_frames / args.frame_rate
|
||||
print(f"Audio: Enabled (duration will match video: {video_duration:.2f}s)")
|
||||
print("=" * 80)
|
||||
|
||||
print(f"\nGenerating {'video + audio' if generate_audio else 'video'}...")
|
||||
|
||||
# Create generation config
|
||||
gen_config = GenerationConfig(
|
||||
prompt=args.prompt,
|
||||
negative_prompt=args.negative_prompt,
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
frame_rate=args.frame_rate,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
condition_image=condition_image,
|
||||
reference_video=reference_video,
|
||||
generate_audio=generate_audio,
|
||||
include_reference_in_output=args.include_reference_in_output,
|
||||
stg_scale=args.stg_scale,
|
||||
stg_blocks=args.stg_blocks,
|
||||
stg_mode=args.stg_mode,
|
||||
)
|
||||
|
||||
# Generate with progress bar
|
||||
with StandaloneSamplingProgress(num_steps=args.num_inference_steps) as progress:
|
||||
# Create sampler with progress context
|
||||
sampler = ValidationSampler(
|
||||
transformer=transformer,
|
||||
vae_decoder=components.video_vae_decoder,
|
||||
vae_encoder=components.video_vae_encoder,
|
||||
text_encoder=components.text_encoder,
|
||||
audio_decoder=components.audio_vae_decoder if generate_audio else None,
|
||||
vocoder=components.vocoder if generate_audio else None,
|
||||
sampling_context=progress,
|
||||
)
|
||||
video, audio = sampler.generate(
|
||||
config=gen_config,
|
||||
device=args.device,
|
||||
)
|
||||
|
||||
# Save video
|
||||
output_path = Path(args.output)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Get audio sample rate from vocoder if audio was generated
|
||||
audio_sample_rate = None
|
||||
if audio is not None and components.vocoder is not None:
|
||||
audio_sample_rate = components.vocoder.output_sampling_rate
|
||||
|
||||
save_video(
|
||||
video_tensor=video,
|
||||
output_path=output_path,
|
||||
fps=args.frame_rate,
|
||||
audio=audio,
|
||||
audio_sample_rate=audio_sample_rate,
|
||||
)
|
||||
print(f"✓ Video saved to {args.output}")
|
||||
|
||||
# Save separate audio file if requested
|
||||
if audio is not None and args.audio_output is not None:
|
||||
audio_output_path = Path(args.audio_output)
|
||||
audio_output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
torchaudio.save(
|
||||
str(audio_output_path),
|
||||
audio.cpu(),
|
||||
sample_rate=audio_sample_rate,
|
||||
)
|
||||
duration = audio.shape[1] / audio_sample_rate
|
||||
print(f"✓ Audio saved: {duration:.2f}s at {audio_sample_rate}Hz")
|
||||
|
||||
print("\n" + "=" * 80)
|
||||
print("Generation complete!")
|
||||
print("=" * 80)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -154,6 +154,17 @@ class CaptionsDataset(Dataset):
|
||||
else:
|
||||
raise ValueError("Expected `dataset_file` to be a path to a CSV, JSON, or JSONL file.")
|
||||
|
||||
def _embedding_output_path(self, media_path: Path) -> str:
|
||||
"""Output `.pt` path relative to the dataset dir; mirrors `process_videos._output_relative`
|
||||
so caption keys match video/audio latent keys (and absolute paths don't escape output_dir)."""
|
||||
data_root = self.dataset_file.parent
|
||||
resolved = data_root / media_path # pathlib: an absolute media_path overrides data_root
|
||||
try:
|
||||
rel = resolved.relative_to(data_root)
|
||||
except ValueError:
|
||||
rel = Path(*resolved.parts[1:]) if resolved.is_absolute() else resolved
|
||||
return str(rel.with_suffix(".pt"))
|
||||
|
||||
def _load_caption_data_from_csv(self) -> dict[str, str]:
|
||||
"""Load captions from a CSV file and compute output embedding paths."""
|
||||
df = pd.read_csv(self.dataset_file)
|
||||
@@ -166,8 +177,7 @@ class CaptionsDataset(Dataset):
|
||||
caption_data = {}
|
||||
for _, row in df.iterrows():
|
||||
media_path = Path(row[self.media_column].strip())
|
||||
# Convert media path to embedding output path (same structure, .pt extension)
|
||||
output_path = str(media_path.with_suffix(".pt"))
|
||||
output_path = self._embedding_output_path(media_path)
|
||||
caption_data[output_path] = row[self.caption_column]
|
||||
|
||||
return caption_data
|
||||
@@ -188,8 +198,7 @@ class CaptionsDataset(Dataset):
|
||||
raise ValueError(f"Key '{self.media_column}' not found in JSON entry: {entry}")
|
||||
|
||||
media_path = Path(entry[self.media_column].strip())
|
||||
# Convert media path to embedding output path (same structure, .pt extension)
|
||||
output_path = str(media_path.with_suffix(".pt"))
|
||||
output_path = self._embedding_output_path(media_path)
|
||||
caption_data[output_path] = entry[self.caption_column]
|
||||
|
||||
return caption_data
|
||||
@@ -206,8 +215,7 @@ class CaptionsDataset(Dataset):
|
||||
raise ValueError(f"Key '{self.media_column}' not found in JSONL entry: {entry}")
|
||||
|
||||
media_path = Path(entry[self.media_column].strip())
|
||||
# Convert media path to embedding output path (same structure, .pt extension)
|
||||
output_path = str(media_path.with_suffix(".pt"))
|
||||
output_path = self._embedding_output_path(media_path)
|
||||
caption_data[output_path] = entry[self.caption_column]
|
||||
|
||||
return caption_data
|
||||
@@ -326,7 +334,8 @@ def compute_captions_embeddings( # noqa: PLR0913
|
||||
# TODO(batch-tokenization): When tokenizer supports batching, encode all prompts at once.
|
||||
# For now, process one at a time:
|
||||
for i in range(len(batch["prompt"])):
|
||||
hidden_states, prompt_attention_mask = text_encoder.encode(batch["prompt"][i], padding_side="left")
|
||||
encoded = text_encoder.encode([batch["prompt"][i]], padding_side="left")
|
||||
hidden_states, prompt_attention_mask = encoded[0]
|
||||
video_prompt_embeds, audio_prompt_embeds = embeddings_processor.feature_extractor(
|
||||
hidden_states, prompt_attention_mask, "left"
|
||||
)
|
||||
|
||||
@@ -1,14 +1,21 @@
|
||||
#!/usr/bin/env python3
|
||||
|
||||
"""
|
||||
Preprocess a video dataset by computing video clips latents and text captions embeddings.
|
||||
This script provides a command-line interface for preprocessing video datasets by computing
|
||||
latent representations of video clips and text embeddings of their captions. The preprocessed
|
||||
data can be used to accelerate training of video generation models and to save GPU memory.
|
||||
Preprocess a media dataset for LTX-2 training.
|
||||
Automatically detects dataset columns and processes each according to a convention table.
|
||||
Column names determine what gets encoded and where outputs go — no per-role CLI flags needed.
|
||||
Convention table:
|
||||
video → Video VAE → latents/
|
||||
audio → Audio VAE → audio_latents/
|
||||
reference_video → Video VAE → reference_latents/
|
||||
reference_audio → Audio VAE → reference_audio_latents/
|
||||
video_mask → (downsample) → video_masks/
|
||||
audio_mask → (downsample) → audio_masks/
|
||||
caption → Text encoder → conditions/
|
||||
Legacy aliases: media_path → video, ref_media_path → reference_video
|
||||
Basic usage:
|
||||
python scripts/process_dataset.py /path/to/dataset.json --resolution-buckets 768x768x49 \
|
||||
python scripts/process_dataset.py /path/to/dataset.json --resolution-buckets 768x768x49 \\
|
||||
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma
|
||||
The dataset must be a CSV, JSON, or JSONL file with columns for captions and video paths.
|
||||
"""
|
||||
|
||||
from pathlib import Path
|
||||
@@ -16,7 +23,15 @@ from pathlib import Path
|
||||
import typer
|
||||
from decode_latents import LatentsDecoder
|
||||
from process_captions import compute_captions_embeddings
|
||||
from process_videos import compute_latents, compute_scaled_resolution_buckets, parse_resolution_buckets
|
||||
from process_videos import (
|
||||
compute_audio_latents,
|
||||
compute_audio_masks,
|
||||
compute_latents,
|
||||
compute_scaled_resolution_buckets,
|
||||
compute_video_masks,
|
||||
detect_dataset_columns,
|
||||
parse_resolution_buckets,
|
||||
)
|
||||
from rich.console import Console
|
||||
|
||||
from ltx_trainer import logger
|
||||
@@ -27,52 +42,82 @@ console = Console()
|
||||
app = typer.Typer(
|
||||
pretty_exceptions_enable=False,
|
||||
no_args_is_help=True,
|
||||
help="Preprocess a video dataset by computing video clips latents and text captions embeddings. "
|
||||
"The dataset must be a CSV, JSON, or JSONL file with columns for captions and video paths.",
|
||||
help="Preprocess a media dataset for LTX-2 training. "
|
||||
"Automatically detects columns (video, audio, reference_video, reference_audio, caption) "
|
||||
"and processes each with the appropriate encoder.",
|
||||
)
|
||||
|
||||
_KNOWN_ROLES = {"video", "audio", "reference_video", "reference_audio", "video_mask", "audio_mask", "caption"}
|
||||
_LEGACY_ALIASES = {"media_path": "video", "ref_media_path": "reference_video"}
|
||||
|
||||
def preprocess_dataset( # noqa: PLR0913
|
||||
|
||||
def preprocess_dataset( # noqa: PLR0912, PLR0913, PLR0915
|
||||
dataset_file: str,
|
||||
caption_column: str,
|
||||
video_column: str,
|
||||
resolution_buckets: list[tuple[int, int, int]],
|
||||
batch_size: int,
|
||||
output_dir: str | None,
|
||||
lora_trigger: str | None,
|
||||
vae_tiling: bool,
|
||||
decode: bool,
|
||||
resolution_buckets: list[tuple[int, int, int]] | None,
|
||||
model_path: str,
|
||||
text_encoder_path: str,
|
||||
device: str,
|
||||
output_dir: str | None = None,
|
||||
video_column: str | None = None,
|
||||
caption_column: str | None = None,
|
||||
batch_size: int = 1,
|
||||
lora_trigger: str | None = None,
|
||||
vae_tiling: bool = False,
|
||||
decode: bool = False,
|
||||
remove_llm_prefixes: bool = False,
|
||||
reference_column: str | None = None,
|
||||
reference_downscale_factor: int = 1,
|
||||
with_audio: bool = False,
|
||||
reference_temporal_scale_factor: int = 1,
|
||||
skip_audio: bool = False,
|
||||
audio_durations: list[float] | None = None,
|
||||
load_text_encoder_in_8bit: bool = False,
|
||||
overwrite: bool = False,
|
||||
) -> None:
|
||||
"""Run the preprocessing pipeline with the given arguments."""
|
||||
# Validate dataset file
|
||||
"""Run the preprocessing pipeline with convention-based column detection."""
|
||||
_validate_dataset_file(dataset_file)
|
||||
|
||||
# Set up output directories
|
||||
# Detect columns and resolve roles
|
||||
dataset_columns = detect_dataset_columns(dataset_file)
|
||||
roles = _resolve_columns(dataset_columns, video_column, caption_column)
|
||||
|
||||
# Log detected roles
|
||||
for role, col in sorted(roles.items()):
|
||||
alias_note = f" (alias for '{role}')" if col != role else ""
|
||||
logger.info(f"Detected column '{col}'{alias_note} → {role}")
|
||||
|
||||
# Validate: need at least caption
|
||||
if "caption" not in roles:
|
||||
raise ValueError(
|
||||
f"No caption column found. Dataset has columns: {dataset_columns}. "
|
||||
f"Expected 'caption' or use --caption-column to specify."
|
||||
)
|
||||
|
||||
# Validate: need video or audio
|
||||
has_video = "video" in roles
|
||||
has_audio = "audio" in roles
|
||||
if not has_video and not has_audio:
|
||||
raise ValueError(
|
||||
f"No media column found. Dataset has columns: {dataset_columns}. "
|
||||
f"Expected 'video', 'audio', or 'media_path' (legacy)."
|
||||
)
|
||||
|
||||
# Validate: video modes need resolution buckets
|
||||
if has_video and not resolution_buckets:
|
||||
raise ValueError("--resolution-buckets is required when the dataset has a video column.")
|
||||
|
||||
output_base = Path(output_dir) if output_dir else Path(dataset_file).parent / ".precomputed"
|
||||
conditions_dir = output_base / "conditions"
|
||||
latents_dir = output_base / "latents"
|
||||
|
||||
if lora_trigger:
|
||||
logger.info(f'LoRA trigger word "{lora_trigger}" will be prepended to all captions')
|
||||
|
||||
# --- Phase 1: Text encoder ---
|
||||
with free_gpu_memory_context():
|
||||
# Process captions using the dedicated function
|
||||
compute_captions_embeddings(
|
||||
dataset_file=dataset_file,
|
||||
output_dir=str(conditions_dir),
|
||||
output_dir=str(output_base / "conditions"),
|
||||
model_path=model_path,
|
||||
text_encoder_path=text_encoder_path,
|
||||
caption_column=caption_column,
|
||||
media_column=video_column,
|
||||
caption_column=roles["caption"],
|
||||
media_column=roles.get("video") or roles.get("audio") or roles["caption"],
|
||||
lora_trigger=lora_trigger,
|
||||
remove_llm_prefixes=remove_llm_prefixes,
|
||||
batch_size=batch_size,
|
||||
@@ -81,119 +126,177 @@ def preprocess_dataset( # noqa: PLR0913
|
||||
overwrite=overwrite,
|
||||
)
|
||||
|
||||
# Process videos using the dedicated function
|
||||
audio_latents_dir = None
|
||||
if with_audio:
|
||||
logger.info("Audio preprocessing enabled - will extract and encode audio from videos")
|
||||
audio_latents_dir = output_base / "audio_latents"
|
||||
# --- Phase 2: Video VAE (video, reference_video) ---
|
||||
if has_video and resolution_buckets:
|
||||
# Determine if audio should be auto-extracted from video files
|
||||
auto_audio = not skip_audio and "audio" not in roles
|
||||
|
||||
with free_gpu_memory_context():
|
||||
compute_latents(
|
||||
dataset_file=dataset_file,
|
||||
video_column=video_column,
|
||||
resolution_buckets=resolution_buckets,
|
||||
output_dir=str(latents_dir),
|
||||
model_path=model_path,
|
||||
batch_size=batch_size,
|
||||
device=device,
|
||||
vae_tiling=vae_tiling,
|
||||
with_audio=with_audio,
|
||||
audio_output_dir=str(audio_latents_dir) if audio_latents_dir else None,
|
||||
overwrite=overwrite,
|
||||
)
|
||||
|
||||
# Process reference videos if reference_column is provided
|
||||
if reference_column:
|
||||
# Validate: scaled references with multiple buckets can cause ambiguous bucket matching
|
||||
if reference_downscale_factor > 1 and len(resolution_buckets) > 1:
|
||||
raise ValueError(
|
||||
"When using --reference-downscale-factor > 1, only a single resolution bucket is supported. "
|
||||
"Using multiple buckets with scaled references can cause ambiguous bucket matching "
|
||||
"(e.g., a 512x256 reference could match either the scaled-down 1024x512 bucket or the 512x256 "
|
||||
"bucket). Please use a single resolution bucket or set --reference-downscale-factor to 1."
|
||||
)
|
||||
|
||||
# Calculate and validate scaled resolution buckets for reference videos
|
||||
reference_buckets = compute_scaled_resolution_buckets(resolution_buckets, reference_downscale_factor)
|
||||
|
||||
if reference_downscale_factor > 1:
|
||||
logger.info(
|
||||
f"Processing reference videos for IC-LoRA training at 1/{reference_downscale_factor} resolution..."
|
||||
)
|
||||
logger.info(f"Reference resolution buckets: {reference_buckets}")
|
||||
else:
|
||||
logger.info("Processing reference videos for IC-LoRA training...")
|
||||
|
||||
reference_latents_dir = output_base / "reference_latents"
|
||||
audio_latents_dir = str(output_base / "audio_latents") if auto_audio else None
|
||||
if auto_audio:
|
||||
logger.info("Audio will be auto-extracted from video files (use --skip-audio to disable)")
|
||||
|
||||
with free_gpu_memory_context():
|
||||
compute_latents(
|
||||
dataset_file=dataset_file,
|
||||
main_media_column=video_column,
|
||||
video_column=reference_column,
|
||||
resolution_buckets=reference_buckets,
|
||||
output_dir=str(reference_latents_dir),
|
||||
video_column=roles["video"],
|
||||
resolution_buckets=resolution_buckets,
|
||||
output_dir=str(output_base / "latents"),
|
||||
model_path=model_path,
|
||||
batch_size=batch_size,
|
||||
device=device,
|
||||
vae_tiling=vae_tiling,
|
||||
with_audio=auto_audio,
|
||||
audio_output_dir=audio_latents_dir,
|
||||
overwrite=overwrite,
|
||||
)
|
||||
|
||||
# Handle decoding if requested (for verification)
|
||||
# Process reference video if present
|
||||
if "reference_video" in roles:
|
||||
if reference_downscale_factor > 1 and len(resolution_buckets) > 1:
|
||||
raise ValueError(
|
||||
"When using --reference-downscale-factor > 1, only a single resolution bucket is supported."
|
||||
)
|
||||
if reference_temporal_scale_factor > 1 and len(resolution_buckets) > 1:
|
||||
raise ValueError(
|
||||
"When using --reference-temporal-scale-factor > 1, only a single resolution bucket is supported."
|
||||
)
|
||||
|
||||
reference_buckets = compute_scaled_resolution_buckets(resolution_buckets, reference_downscale_factor)
|
||||
if reference_downscale_factor > 1:
|
||||
logger.info(f"Processing reference videos at 1/{reference_downscale_factor} resolution...")
|
||||
if reference_temporal_scale_factor > 1:
|
||||
logger.info(
|
||||
f"Temporally subsampling reference videos by {reference_temporal_scale_factor}x "
|
||||
f"(VAE-aligned pattern)..."
|
||||
)
|
||||
|
||||
with free_gpu_memory_context():
|
||||
compute_latents(
|
||||
dataset_file=dataset_file,
|
||||
main_media_column=roles["video"],
|
||||
video_column=roles["reference_video"],
|
||||
resolution_buckets=reference_buckets,
|
||||
output_dir=str(output_base / "reference_latents"),
|
||||
model_path=model_path,
|
||||
batch_size=batch_size,
|
||||
device=device,
|
||||
vae_tiling=vae_tiling,
|
||||
overwrite=overwrite,
|
||||
temporal_subsample_factor=reference_temporal_scale_factor,
|
||||
)
|
||||
|
||||
# --- Phase 2b: Masks (video_mask, audio_mask) — processed after video latents for alignment ---
|
||||
if "video_mask" in roles and has_video:
|
||||
compute_video_masks(
|
||||
dataset_file=dataset_file,
|
||||
mask_column=roles["video_mask"],
|
||||
latents_dir=str(output_base / "latents"),
|
||||
output_dir=str(output_base / "video_masks"),
|
||||
main_media_column=roles["video"],
|
||||
)
|
||||
|
||||
# --- Phase 3: Audio VAE (audio, reference_audio) ---
|
||||
audio_roles_to_process = [
|
||||
("audio", "audio_latents"),
|
||||
("reference_audio", "reference_audio_latents"),
|
||||
]
|
||||
active_audio_roles = [(role, subdir) for role, subdir in audio_roles_to_process if role in roles]
|
||||
|
||||
if active_audio_roles:
|
||||
# Determine audio duration constraint: video bucket → max_duration, or explicit buckets
|
||||
max_audio_duration = None
|
||||
audio_duration_buckets = None
|
||||
if has_video and resolution_buckets:
|
||||
max_audio_duration = max(f for f, _h, _w in resolution_buckets) / 25.0
|
||||
elif audio_durations:
|
||||
audio_duration_buckets = audio_durations
|
||||
|
||||
for role, output_subdir in active_audio_roles:
|
||||
with free_gpu_memory_context():
|
||||
compute_audio_latents(
|
||||
dataset_file=dataset_file,
|
||||
audio_column=roles[role],
|
||||
output_dir=str(output_base / output_subdir),
|
||||
model_path=model_path,
|
||||
main_media_column=roles.get("video"),
|
||||
max_duration=max_audio_duration,
|
||||
duration_buckets=audio_duration_buckets,
|
||||
device=device,
|
||||
overwrite=overwrite,
|
||||
)
|
||||
|
||||
# --- Phase 4: Audio masks (after audio latents exist for temporal alignment) ---
|
||||
if "audio_mask" in roles:
|
||||
audio_latents_source = output_base / "audio_latents"
|
||||
if audio_latents_source.exists():
|
||||
compute_audio_masks(
|
||||
dataset_file=dataset_file,
|
||||
mask_column=roles["audio_mask"],
|
||||
audio_latents_dir=str(audio_latents_source),
|
||||
output_dir=str(output_base / "audio_masks"),
|
||||
main_media_column=roles.get("video") or roles.get("audio"),
|
||||
)
|
||||
else:
|
||||
logger.warning("audio_mask column found but no audio_latents/ — run with audio first")
|
||||
|
||||
# --- Decode for verification ---
|
||||
if decode:
|
||||
logger.info("Decoding latents for verification...")
|
||||
decoder = LatentsDecoder(model_path=model_path, device=device, vae_tiling=vae_tiling, with_audio=has_audio)
|
||||
if has_video:
|
||||
decoder.decode(output_base / "latents", output_base / "decoded_videos")
|
||||
if "reference_video" in roles and (output_base / "reference_latents").exists():
|
||||
decoder.decode(output_base / "reference_latents", output_base / "decoded_reference_videos")
|
||||
|
||||
decoder = LatentsDecoder(
|
||||
model_path=model_path,
|
||||
device=device,
|
||||
vae_tiling=vae_tiling,
|
||||
with_audio=with_audio,
|
||||
)
|
||||
decoder.decode(latents_dir, output_base / "decoded_videos")
|
||||
|
||||
# Also decode reference videos if they exist
|
||||
if reference_column:
|
||||
reference_latents_dir = output_base / "reference_latents"
|
||||
if reference_latents_dir.exists():
|
||||
logger.info("Decoding reference videos...")
|
||||
decoder.decode(reference_latents_dir, output_base / "decoded_reference_videos")
|
||||
|
||||
# Decode audio latents if they exist
|
||||
if with_audio and audio_latents_dir and audio_latents_dir.exists():
|
||||
logger.info("Decoding audio latents...")
|
||||
decoder.decode_audio(audio_latents_dir, output_base / "decoded_audio")
|
||||
|
||||
# Print summary
|
||||
# --- Summary ---
|
||||
logger.info(f"Dataset preprocessing complete! Results saved to {output_base}")
|
||||
if reference_column:
|
||||
logger.info("Reference videos processed and saved to reference_latents/ directory for IC-LoRA training")
|
||||
if with_audio:
|
||||
logger.info("Audio latents saved to audio_latents/ directory for audio-video training")
|
||||
produced = [d.name for d in output_base.iterdir() if d.is_dir() and not d.name.startswith("decoded")]
|
||||
logger.info(f"Output directories: {', '.join(sorted(produced))}")
|
||||
|
||||
|
||||
def _validate_dataset_file(dataset_path: str) -> None:
|
||||
"""Validate that the dataset file exists and has the correct format."""
|
||||
dataset_file = Path(dataset_path)
|
||||
|
||||
if not dataset_file.exists():
|
||||
raise FileNotFoundError(f"Dataset file does not exist: {dataset_file}")
|
||||
|
||||
if not dataset_file.is_file():
|
||||
raise ValueError(f"Dataset path must be a file, not a directory: {dataset_file}")
|
||||
|
||||
if dataset_file.suffix.lower() not in [".csv", ".json", ".jsonl"]:
|
||||
raise ValueError(f"Dataset file must be CSV, JSON, or JSONL format: {dataset_file}")
|
||||
|
||||
|
||||
def _resolve_columns(
|
||||
dataset_columns: set[str],
|
||||
video_column_override: str | None = None,
|
||||
caption_column_override: str | None = None,
|
||||
) -> dict[str, str]:
|
||||
"""Map canonical role names to actual dataset column names.
|
||||
Returns a dict of role → column_name for recognized roles found in the dataset.
|
||||
"""
|
||||
roles: dict[str, str] = {}
|
||||
for col in dataset_columns:
|
||||
role = _LEGACY_ALIASES.get(col, col)
|
||||
if role in _KNOWN_ROLES:
|
||||
roles[role] = col
|
||||
|
||||
if video_column_override and video_column_override in dataset_columns:
|
||||
roles["video"] = video_column_override
|
||||
if caption_column_override and caption_column_override in dataset_columns:
|
||||
roles["caption"] = caption_column_override
|
||||
|
||||
return roles
|
||||
|
||||
|
||||
@app.command()
|
||||
def main( # noqa: PLR0913
|
||||
dataset_path: str = typer.Argument(
|
||||
...,
|
||||
help="Path to metadata file (CSV/JSON/JSONL) containing captions and video paths",
|
||||
help="Path to metadata file (CSV/JSON/JSONL) with columns matching the convention table",
|
||||
),
|
||||
resolution_buckets: str = typer.Option(
|
||||
...,
|
||||
help='Resolution buckets in format "WxHxF;WxHxF;..." (e.g. "768x768x25;512x512x49")',
|
||||
resolution_buckets: str | None = typer.Option(
|
||||
default=None,
|
||||
help='Resolution buckets in format "WxHxF;WxHxF;..." (e.g. "768x768x25"). '
|
||||
"Required when dataset has a video column.",
|
||||
),
|
||||
model_path: str = typer.Option(
|
||||
...,
|
||||
@@ -203,13 +306,13 @@ def main( # noqa: PLR0913
|
||||
...,
|
||||
help="Path to Gemma text encoder directory",
|
||||
),
|
||||
caption_column: str = typer.Option(
|
||||
default="caption",
|
||||
help="Column name containing captions in the dataset JSON/JSONL/CSV file",
|
||||
caption_column: str | None = typer.Option(
|
||||
default=None,
|
||||
help="Override: treat this column as 'caption' (default: auto-detect 'caption')",
|
||||
),
|
||||
video_column: str = typer.Option(
|
||||
default="media_path",
|
||||
help="Column name containing video paths in the dataset JSON/JSONL/CSV file",
|
||||
video_column: str | None = typer.Option(
|
||||
default=None,
|
||||
help="Override: treat this column as 'video' (default: auto-detect 'video' or 'media_path')",
|
||||
),
|
||||
batch_size: int = typer.Option(
|
||||
default=1,
|
||||
@@ -229,32 +332,43 @@ def main( # noqa: PLR0913
|
||||
),
|
||||
lora_trigger: str | None = typer.Option(
|
||||
default=None,
|
||||
help="Optional trigger word to prepend to each caption (activates the LoRA during inference)",
|
||||
help="Optional trigger word to prepend to each caption",
|
||||
),
|
||||
decode: bool = typer.Option(
|
||||
default=False,
|
||||
help="Decode and save latents after encoding (videos and audio) for verification",
|
||||
help="Decode and save latents after encoding for verification",
|
||||
),
|
||||
remove_llm_prefixes: bool = typer.Option(
|
||||
default=False,
|
||||
help="Remove LLM prefixes from captions",
|
||||
),
|
||||
reference_column: str | None = typer.Option(
|
||||
skip_audio: bool = typer.Option(
|
||||
default=False,
|
||||
help="Don't extract audio from video files (audio extraction is on by default)",
|
||||
),
|
||||
audio_durations: str | None = typer.Option(
|
||||
default=None,
|
||||
help="Column name containing reference video paths (for video-to-video training)",
|
||||
help='Audio duration buckets in seconds for audio-only datasets (e.g. "2.0;4.0;8.0"). '
|
||||
"When set, audio files are trimmed to the best matching duration. "
|
||||
"Not needed when a video column is present (audio duration derived from video bucket).",
|
||||
),
|
||||
with_audio: bool = typer.Option(
|
||||
default=False,
|
||||
help="Extract and encode audio from video files",
|
||||
hidden=True,
|
||||
help="[DEPRECATED: audio is now on by default, use --skip-audio to disable]",
|
||||
),
|
||||
load_text_encoder_in_8bit: bool = typer.Option(
|
||||
default=False,
|
||||
help="Load the Gemma text encoder in 8-bit precision to save GPU memory (requires bitsandbytes)",
|
||||
help="Load the Gemma text encoder in 8-bit precision to save GPU memory",
|
||||
),
|
||||
reference_downscale_factor: int = typer.Option(
|
||||
default=1,
|
||||
help="Downscale factor for reference video resolution. When > 1, reference videos are processed at "
|
||||
"1/n resolution (e.g., 2 means half resolution). Used for efficient IC-LoRA training.",
|
||||
help="Downscale factor for reference video resolution (e.g., 2 = half resolution for IC-LoRA)",
|
||||
),
|
||||
reference_temporal_scale_factor: int = typer.Option(
|
||||
default=1,
|
||||
help="Temporal subsampling factor for reference videos (e.g., 2 = half frame rate, "
|
||||
"VAE-aligned: keeps frame 0, then every Nth frame from frame 1 onwards)",
|
||||
),
|
||||
overwrite: bool = typer.Option(
|
||||
default=False,
|
||||
@@ -262,64 +376,53 @@ def main( # noqa: PLR0913
|
||||
"changed parameters (different model, resolution, etc.) so stale outputs are replaced.",
|
||||
),
|
||||
) -> None:
|
||||
"""Preprocess a video dataset by computing and saving latents and text embeddings.
|
||||
For multi-GPU preprocessing, invoke under ``accelerate launch`` - each process
|
||||
"""Preprocess a media dataset for LTX-2 training.
|
||||
See module docstring for the convention table. Audio is auto-extracted from
|
||||
video files by default — use --skip-audio to disable.
|
||||
For multi-GPU preprocessing, invoke under ``accelerate launch`` -- each process
|
||||
will handle an interleaved shard of the dataset.
|
||||
The dataset must be a CSV, JSON, or JSONL file with columns for captions and video paths.
|
||||
This script is designed for LTX-2 models which use the Gemma text encoder.
|
||||
Examples:
|
||||
# Process a dataset with LTX-2 model
|
||||
python scripts/process_dataset.py dataset.json --resolution-buckets 768x768x25 \\
|
||||
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma
|
||||
# Process dataset with custom column names
|
||||
python scripts/process_dataset.py dataset.json --resolution-buckets 768x768x25 \\
|
||||
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma \\
|
||||
--caption-column "text" --video-column "video_path"
|
||||
# Process dataset with reference videos for IC-LoRA training
|
||||
python scripts/process_dataset.py dataset.json --resolution-buckets 768x768x25 \\
|
||||
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma \\
|
||||
--reference-column "reference_path"
|
||||
# Process dataset with scaled reference videos (half resolution) for efficient IC-LoRA
|
||||
python scripts/process_dataset.py dataset.json --resolution-buckets 768x768x25 \\
|
||||
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma \\
|
||||
--reference-column "reference_path" --reference-downscale-factor 2
|
||||
# Process dataset with audio for audio-video training
|
||||
python scripts/process_dataset.py dataset.json --resolution-buckets 768x512x97 \\
|
||||
--model-path /path/to/ltx2.safetensors --text-encoder-path /path/to/gemma \\
|
||||
--with-audio
|
||||
"""
|
||||
parsed_resolution_buckets = parse_resolution_buckets(resolution_buckets)
|
||||
|
||||
if len(parsed_resolution_buckets) > 1:
|
||||
# Handle deprecated --with-audio flag
|
||||
if with_audio:
|
||||
logger.warning(
|
||||
"Using multiple resolution buckets. "
|
||||
"When training with multiple resolution buckets, you must use a batch size of 1."
|
||||
"--with-audio is deprecated. Audio extraction is now on by default. Use --skip-audio to disable."
|
||||
)
|
||||
|
||||
# Validate reference_downscale_factor
|
||||
parsed_buckets = parse_resolution_buckets(resolution_buckets) if resolution_buckets else None
|
||||
|
||||
if parsed_buckets and len(parsed_buckets) > 1:
|
||||
logger.warning("Using multiple resolution buckets. Training batch size must be 1.")
|
||||
|
||||
if reference_downscale_factor < 1:
|
||||
raise typer.BadParameter("--reference-downscale-factor must be >= 1")
|
||||
|
||||
if reference_downscale_factor > 1 and not reference_column:
|
||||
logger.warning("--reference-downscale-factor specified but no --reference-column provided. Ignoring.")
|
||||
if reference_temporal_scale_factor < 1:
|
||||
raise typer.BadParameter("--reference-temporal-scale-factor must be >= 1")
|
||||
|
||||
parsed_audio_durations = None
|
||||
if audio_durations:
|
||||
parsed_audio_durations = [float(d) for d in audio_durations.split(";")]
|
||||
if any(d <= 0 for d in parsed_audio_durations):
|
||||
raise typer.BadParameter("All audio durations must be positive")
|
||||
|
||||
preprocess_dataset(
|
||||
dataset_file=dataset_path,
|
||||
caption_column=caption_column,
|
||||
video_column=video_column,
|
||||
resolution_buckets=parsed_resolution_buckets,
|
||||
batch_size=batch_size,
|
||||
output_dir=output_dir,
|
||||
lora_trigger=lora_trigger,
|
||||
vae_tiling=vae_tiling,
|
||||
decode=decode,
|
||||
resolution_buckets=parsed_buckets,
|
||||
model_path=model_path,
|
||||
text_encoder_path=text_encoder_path,
|
||||
device=device,
|
||||
output_dir=output_dir,
|
||||
video_column=video_column,
|
||||
caption_column=caption_column,
|
||||
batch_size=batch_size,
|
||||
lora_trigger=lora_trigger,
|
||||
vae_tiling=vae_tiling,
|
||||
decode=decode,
|
||||
remove_llm_prefixes=remove_llm_prefixes,
|
||||
reference_column=reference_column,
|
||||
reference_downscale_factor=reference_downscale_factor,
|
||||
with_audio=with_audio,
|
||||
reference_temporal_scale_factor=reference_temporal_scale_factor,
|
||||
skip_audio=skip_audio,
|
||||
audio_durations=parsed_audio_durations,
|
||||
load_text_encoder_in_8bit=load_text_encoder_in_8bit,
|
||||
overwrite=overwrite,
|
||||
)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user