4 Commits

Author SHA1 Message Date
Michael Kupchick 9377758131 Merge pull request #248 from Lightricks/pr-2026-07-07-be3b401
Public sync - 2026-07-07
2026-07-08 08:50:18 +03:00
github-actions[bot] 63fd9a4f86 Automated PR - 2026-07-07 2026-07-07 16:57:50 +00:00
Alexey Kravtsov 780984275f Merge pull request #237 from Lightricks/pr-2026-06-17-97c9503
Public sync - 2026-06-17
2026-06-17 17:26:42 +03:00
github-actions[bot] f4b06fb977 Automated PR - 2026-06-17 2026-06-17 14:21:07 +00:00
229 changed files with 28721 additions and 8566 deletions
+277
View File
@@ -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 19 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 / 4060GB / 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 4060GB 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 23 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 12 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.
## 4060GB tier — mid-range (autotune from low-VRAM)
**VRAM range:** 4060 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 **4060GB 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 | 3264 | Default for most concept LoRAs. Start at 32; bump to 64 if validation samples underfit. |
| Multi-character world, dense series, complex multi-concept | 96128 | More capacity for distinguishing several concepts inside one LoRA. |
| Camera move, motion, transition (i.e. behavioural, not visual) | 816 | Motion is a thin signal — high ranks just memorise frame content. |
| IC-LoRA control (V2V depth/pose/Canny/etc., A2A audio reference) | 1632 | Start at 16 for structural control (depth, pose, edges); 2432 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 | 4060GB 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
View File
@@ -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
+29 -3
View File
@@ -17,8 +17,10 @@ checkpoints/
# Other files
.DS_Store
tmp
.wandb
projects/
tmp
wandb/
# Model checkpoints
*.ckpt
@@ -36,14 +38,38 @@ tmp
*.json
*.m4a
*.mov
*.mp3
*.mp4
*.png
*.wav
*.webp
# HDR IC-LoRA e2e test baseline (checked in via Git LFS)
!packages/ltx-pipelines/tests/assets/expected_hdr_ic_lora_exr/frame_*.exr
# Full-params expected results for the --full-e2e quality lane (checked in via Git LFS)
!packages/ltx-pipelines/tests/assets/full_e2e/
!packages/ltx-pipelines/tests/assets/full_e2e/*.mp4
!packages/ltx-pipelines/tests/assets/full_e2e/*.wav
!packages/ltx-pipelines/tests/assets/full_e2e/expected_hdr_ic_lora_exr/frame_*.exr
# HDR IC-LoRA e2e test input clip (checked in via Git LFS)
!packages/ltx-pipelines/tests/assets/hdr_ic_lora_test_input.mp4
# Fast integration-profile goldens (bit-exact decoded video/audio; checked in via Git LFS)
!packages/ltx-pipelines/tests/assets/integration/
!packages/ltx-pipelines/tests/assets/integration/*.safetensors
# 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
+35 -11
View File
@@ -14,19 +14,43 @@
## 🚀 Quick Start
Clone the repo
```bash
# Clone the repository
git clone https://github.com/Lightricks/LTX-2.git
cd LTX-2
# Set up the environment
uv sync --frozen
source .venv/bin/activate
```
### Required Models
Download the relevant [models](https://huggingface.co/Lightricks/LTX-2.3) or use the [Hugging Face CLI](https://huggingface.co/docs/huggingface_hub/guides/cli)
Download the following models from the [LTX-2.3 HuggingFace repository](https://huggingface.co/Lightricks/LTX-2.3):
```bash
hf auth login
hf download Lightricks/LTX-2.3 \
ltx-2.3-22b-distilled-1.1.safetensors ltx-2.3-spatial-upscaler-x2-1.1.safetensors --local-dir models/ltx-2.3
hf download google/gemma-3-12b-it-qat-q4_0-unquantized --local-dir models/gemma-3-12b
```
If you get a 401/403, accept the model terms on Hugging Face and log in with a **Read** token (fine-grained tokens need the "read gated repos" scope enabled).
Generate
```bash
uv run python -m ltx_pipelines.distilled \
--distilled-checkpoint-path models/ltx-2.3/ltx-2.3-22b-distilled-1.1.safetensors \
--spatial-upsampler-path models/ltx-2.3/ltx-2.3-spatial-upscaler-x2-1.1.safetensors \
--gemma-root models/gemma-3-12b \
--seed 42 \
--output-path output.mp4 \
--prompt "A medium close-up shot features a Caucasian man with a beard, wearing a green and white baseball cap without any letters on the front, and a light blue shirt over a white t-shirt. He is positioned in the center of the frame, looking intently directly at the camera, his eyes focused on camera. His facial expression is one of deep concentration, with his brow slightly raised. As he looks straight at the camera, a quick sniff sound is heard, and then he speaks with a deep male voice and a satisfied tone, saying, 'I think it's so good.' The camera remains static throughout, maintaining a shallow depth of field, which keeps the man in sharp focus while the background is softly blurred, showing a beige wall behind him. After a brief pause, another short, audible sniff is heard. The man then continues to speak, his voice maintaining the same quality, as he states, 'So good. So good.' He elaborates further, emphasizing his point with a final statement, 'This got to be, it's got to be the best tool I've ever seen.'"
```
In cases of GPU memory constraints, consider `--quantization fp8-cast --offload {cpu, disk}`. See [additional flags](packages/ltx-pipelines/docs/installation.md#common-cli-flags).
This uses the distilled model and pipeline for fast results. For better quality or other capabilities, see [Models](#full-model-list) and [Pipelines](#available-pipelines).
### Full Model List
For pipelines beyond the quickstart, download the relevant models from the [LTX-2.3 HuggingFace repository](https://huggingface.co/Lightricks/LTX-2.3):
**LTX-2.3 Model Checkpoint** (choose and download one of the following)
* [`ltx-2.3-22b-dev.safetensors`](https://huggingface.co/Lightricks/LTX-2.3/blob/main/ltx-2.3-22b-dev.safetensors) - [Download](https://huggingface.co/Lightricks/LTX-2.3/resolve/main/ltx-2.3-22b-dev.safetensors)
@@ -76,9 +100,9 @@ Download the following models from the [LTX-2.3 HuggingFace repository](https://
### ⚡ Optimization Tips
* **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** - 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))
* **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. On Hopper+ GPUs with native FP8 support, use `--quantization fp8-scaled-mm` for FP8 scaled matrix multiplication. Fp8-scaled-mm should be used with fp8 checkpoints.
* **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 Hopper GPUs, install the FlashAttention 3 wheel. On other CUDA GPUs, PyTorch SDPA is used automatically. An installed backend is selected automatically at runtime; forcing a specific one is a Python-API option (`AttentionFunction.FLASH_ATTENTION_3`/`FLASH_ATTENTION_4`), not a CLI flag.
* **Use gradient estimation** - Reduce inference steps from 40 to 20-30 while maintaining quality (see [pipeline documentation](packages/ltx-pipelines/docs/optimization.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
@@ -94,7 +118,7 @@ When writing prompts, focus on detailed, chronological descriptions of actions a
- Describe lighting and colors
- Note any changes or sudden events
For additional guidance on writing a prompt please refer to <https://ltx.video/blog/how-to-prompt-for-ltx-2>
For additional guidance on writing a prompt please refer to <https://ltx.io/blog/prompting-guide-for-ltx-2>
### Automatic Prompt Enhancement
+38 -5
View File
@@ -8,9 +8,10 @@ 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.
- **`quantization/`**: FP8 quantization backends (FP8 scaled MM, FP8 cast) for reduced memory footprint.
## 🚀 Quick Start
@@ -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
@@ -116,11 +118,9 @@ model = builder.build(device=torch.device("cuda"))
The `quantization/` module provides FP8 quantization support for the LTX-2 transformer, significantly reducing memory usage while maintaining quality. Two backends are available:
#### FP8 Scaled MM (TensorRT-LLM)
#### FP8 Scaled MM
Uses NVIDIA TensorRT-LLM's `cublas_scaled_mm` for efficient FP8 matrix multiplication. Weights are stored in FP8 format with per-tensor scaling, and inputs are quantized dynamically (or statically with calibration data).
**Requirements**: `uv sync --frozen --extra fp8-trtllm`
Uses PyTorch's `torch._scaled_mm` for efficient FP8 matrix multiplication. Weights are stored in FP8 format with per-tensor scaling, and inputs are quantized dynamically.
**Usage with QuantizationPolicy:**
@@ -157,6 +157,39 @@ 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.
---
+7 -31
View File
@@ -1,6 +1,6 @@
[project]
name = "ltx-core"
version = "1.1.5"
version = "1.1.7"
description = "Core implementation of Lightricks' LTX-2 model"
readme = "README.md"
requires-python = ">=3.10"
@@ -13,38 +13,14 @@ dependencies = [
"safetensors",
"accelerate",
"scipy>=1.14",
# Apple Silicon only: Apple's fused MPSGraph SDPA, the AUTOMATIC attention
# backend on MPS. The marker installs it on Apple Silicon and prunes it
# everywhere else (Linux/CUDA), so it is a hard requirement exactly where it
# is the only viable attention kernel. Requires torch>=2.11 (within the
# torch~=2.7 floor); the resolver forks torch to >=2.11 on macOS.
"mps-sdpa>=0.2.0; sys_platform == 'darwin' and platform_machine == 'arm64'",
]
[project.optional-dependencies]
xformers = ["xformers"]
fp8-trtllm = [
"tensorrt-llm==1.0.0",
"onnx>=1.16.0,<1.20.0",
"openmpi",
]
[tool.uv]
conflicts = [
[
{ extra = "xformers" },
{ extra = "fp8-trtllm" },
],
]
[tool.uv.sources]
xformers = { index = "pytorch" }
tensorrt-llm = { index = "nvidia" }
[[tool.uv.index]]
name = "pytorch"
url = "https://download.pytorch.org/whl/cu129"
explicit = true
[[tool.uv.index]]
name = "nvidia"
url = "https://pypi.nvidia.com/"
explicit = true
[build-system]
requires = ["uv_build>=0.9.8,<0.10.0"]
build-backend = "uv_build"
+14 -4
View File
@@ -25,8 +25,12 @@ from ltx_core.model.transformer.modality import Modality
def _split_perturbations(config: BatchedPerturbationConfig, sizes: list[int]) -> list[BatchedPerturbationConfig]:
"""Split a ``BatchedPerturbationConfig`` along the batch dimension."""
it = iter(config.perturbations)
return [BatchedPerturbationConfig([next(it) for _ in range(s)]) for s in sizes]
chunks = []
offset = 0
for size in sizes:
chunks.append(config.batch_slice(offset, offset + size))
offset += size
return chunks
def _merge_tensors(tensors: list[torch.Tensor | None]) -> torch.Tensor | None:
@@ -54,6 +58,10 @@ class BatchSplitAdapter(nn.Module):
self._model = model
self._max_batch_size = max_batch_size
@property
def num_blocks(self) -> int:
return self._model.num_blocks
def _get_chunk_sizes(self, batch_size: int) -> list[int]:
full, remainder = divmod(batch_size, self._max_batch_size)
sizes = [self._max_batch_size] * full
@@ -65,7 +73,7 @@ class BatchSplitAdapter(nn.Module):
self,
video: Modality | None,
audio: Modality | None,
perturbations: BatchedPerturbationConfig,
perturbations: BatchedPerturbationConfig | None,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
batch_size = (video or audio).latent.shape[0]
@@ -77,7 +85,9 @@ class BatchSplitAdapter(nn.Module):
v_chunks = video.split(sizes) if video is not None else [None] * n
a_chunks = audio.split(sizes) if audio is not None else [None] * n
p_chunks = _split_perturbations(perturbations, sizes)
# A None config means "perturb nothing"; forward it per chunk so the inner model
# builds a per-chunk all-keep mask (splitting None has nothing to slice).
p_chunks = _split_perturbations(perturbations, sizes) if perturbations is not None else [None] * n
chunk_results = [
self._model(video=vc, audio=ac, perturbations=pc)
@@ -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,20 +2,31 @@
from __future__ import annotations
import copy
import logging
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.stream_sync import create_stream_sync
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.devices import synchronize_device
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
@@ -25,23 +36,29 @@ from ltx_core.loader.primitives import (
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.
@@ -57,32 +74,116 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
``"transformer_blocks"``).
blocks_prefix: State-dict key prefix for block weights
(e.g. ``"transformer_blocks"``).
cpu_slots_count: Default number of pinned CPU buffer slots used by
:meth:`build` when it is not given an explicit ``cpu_slots_count``.
``None`` = RAM streaming (all blocks pinned); a small value (e.g.
``DISK_CPU_SLOTS``) selects disk streaming. Lets a builder fully
encode its offload behaviour so callers need not re-specify it.
"""
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)
fuse_rule: FuseRule = bf16_fuse_rule
def __init__( # noqa: PLR0913
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 = "",
cpu_slots_count: int | None = None,
) -> 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
self._cpu_slots_count = cpu_slots_count
# Streaming-specific
blocks_attr: str = ""
blocks_prefix: str = ""
@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 checkpoint(self) -> str | tuple[str, ...]:
return self._model_path
def with_loras(self, loras: tuple[LoraPathStrengthAndSDOps, ...]) -> StreamingModelBuilder:
return replace(self, loras=loras)
@property
def model_sd_ops(self) -> SDOps | None:
return self._model_sd_ops
def with_fuse_rule(self, fuse_rule: FuseRule) -> StreamingModelBuilder:
return replace(self, fuse_rule=fuse_rule)
@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
@property
def cpu_slots_count(self) -> int | None:
return self._cpu_slots_count
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."""
@@ -94,23 +195,28 @@ 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``).
cpu_slots_count: Number of pinned CPU buffer slots.
``None`` = RAM streaming (all blocks pre-loaded with LoRA fusion).
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`` falls
back to the builder's configured ``cpu_slots_count``, and if that
is also ``None``, to RAM streaming (all blocks pre-loaded with
LoRA fusion).
gpu_slots_count: Number of GPU buffer slots.
``None`` = ``_DEFAULT_GPU_SLOTS`` (2).
"""
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)
@@ -120,31 +226,40 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
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 self._cpu_slots_count
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,
gpu_slots_count,
target_device,
reuse_barrier=lambda event: copy_stream.wait_event(event),
)
self._load_non_block_weights(meta_model, non_block_keys, device, dtype, non_block_loras)
sync = create_stream_sync(device)
gpu_pool = BufferPool(source.slot_nbytes, gpu_slots_count, device, reuse_barrier=sync.reuse_barrier)
provider = WeightsProvider(
gpu_pool,
copy_stream,
target_device,
sync,
device,
source,
lora_sources,
self.blocks_prefix,
@@ -154,24 +269,12 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
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,
@@ -179,6 +282,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(
@@ -186,87 +308,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, fuse_rule=self.fuse_rule, preserve_input_device=False
block_sd, lora_sd_and_strengths, fuse_rule=self.fuse_rule, 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
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()
synchronize_device()
# 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] = {
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,
lora_sources=lora_sources,
fuse_rule=self.fuse_rule,
)
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,
@@ -277,43 +387,42 @@ 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
@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,
lora_sources: list[LoraSource] | None = None,
fuse_rule: FuseRule = bf16_fuse_rule,
lora_sd_and_strengths: list[LoraStateDictWithStrength],
) -> None:
"""Load non-block weights into *model* on *device*, applying ``sd_ops`` and fusing LoRAs.
Fusion goes through :func:`fuse_lora_weights` under *fuse_rule* so the
bf16 rounding pattern matches the non-streaming and block-streaming
paths, and any quantization-specific fuse rule the builder configures
is honored here as well.
"""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: dict[str, torch.Tensor] = {}
for sft_key, model_key in non_block_keys:
tensor = reader.get_tensor(sft_key).to(device=device, dtype=dtype)
if sd_ops is not None:
for kv in sd_ops.apply_to_key_value(model_key, tensor):
non_block_sd[kv.new_key] = kv.new_value
else:
non_block_sd[model_key] = tensor
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_sources:
lora_sd_and_strengths = [src.as_state_dict_with_strength() for src in lora_sources]
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=fuse_rule,
fuse_rule=self.fuse_rule,
preserve_input_device=True,
):
non_block_sd[key] = fused
@@ -321,6 +430,35 @@ class StreamingModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType])
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(
checkpoint_paths: list[str],
sd_ops: SDOps | None,
@@ -128,8 +128,8 @@ class LoraSource:
def as_state_dict_with_strength(self) -> LoraStateDictWithStrength:
"""Return a :class:`LoraStateDictWithStrength` view of the pinned A/B factors.
Lets :func:`fuse_lora_weights` consume disk-streaming LoRAs without
re-reading the safetensors file or re-applying ``sd_ops``.
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():
@@ -153,9 +153,11 @@ class LoraSource:
if pair is None:
return None
a, b = pair
if device is not None and device.type == "cuda":
a = a.to(device=device, non_blocking=True)
b = b.to(device=device, non_blocking=True)
# Move A/B to a GPU-class target (CUDA/MPS) so the B@A aggregation runs on
# the device; on a CPU target they stay put. non_blocking only helps CUDA.
if device is not None and device.type in ("cuda", "mps"):
a = a.to(device=device, non_blocking=device.type == "cuda")
b = b.to(device=device, non_blocking=device.type == "cuda")
if dtype is not None:
a = a.to(dtype=dtype)
b = b.to(dtype=dtype)
@@ -1,4 +1,4 @@
"""Weight buffer pool for block streaming."""
"""Raw buffer pool for block streaming."""
from __future__ import annotations
@@ -7,69 +7,65 @@ 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
from ltx_core.block_streaming.stream_sync import StreamEvent
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],
reuse_barrier: Callable[[StreamEvent], 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._events: dict[int, torch.cuda.Event] = {}
self._free: deque[torch.Tensor] = deque()
self._events: dict[int, StreamEvent] = {}
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: StreamEvent | 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,45 +3,38 @@
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.loader.fuse_loras import FuseRule, aggregate_lora_products, bf16_fuse_rule
from ltx_core.block_streaming.stream_sync import StreamEvent, StreamSync
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, device_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:
"""Provides GPU-ready block weights via H2D copy from a pinned CPU weight source.
Args:
pool: Pre-allocated GPU weight buffer pool.
copy_stream: Dedicated CUDA stream for async H2D copies.
target_device: GPU device for compute.
sync: Coordinates copy-vs-compute ordering for the backend
(see :class:`StreamSync`).
target_device: device for compute.
source: Pinned CPU weight source.
lora_sources: LoRA adapters fused on H2D copy.
blocks_prefix: State-dict prefix for LoRA key matching.
@@ -51,18 +44,18 @@ class WeightsProvider:
def __init__(
self,
pool: WeightPool,
copy_stream: torch.cuda.Stream,
pool: BufferPool,
sync: StreamSync,
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._sync = sync
self._pool = pool
self._cache: OrderedDict[int, dict[str, torch.Tensor]] = OrderedDict()
self._events: dict[int, torch.cuda.Event] = {}
self._cache: OrderedDict[int, CachedBlock] = OrderedDict()
self._events: dict[int, StreamEvent | None] = {}
self._target_device = target_device
self._source = source
self._lora_sources = lora_sources or []
@@ -72,56 +65,68 @@ class WeightsProvider:
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],
) -> 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: torch.Tensor,
nbytes: int,
) -> StreamEvent | None:
"""Copy block weights to the target device and fuse LoRAs.
*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 copy + fusion run under :meth:`StreamSync.copy_scope`, then
:meth:`StreamSync.commit_copy` orders the copy before compute and returns
a guard event for the source to reuse (the ordering is committed inside
this method so callers -- and instrumentation regions wrapping it --
observe the full transfer time).
"""
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)
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 self._sync.copy_scope():
raw[:nbytes].copy_(cpu_buffer[:nbytes], non_blocking=self._sync.is_async_copy)
if self._lora_sources:
self._fuse_block_loras(idx, gpu_weights)
h2d_event = torch.cuda.Event()
h2d_event.record(self._copy_stream)
torch.cuda.current_stream(self._target_device).wait_event(h2d_event)
return h2d_event
return self._sync.commit_copy()
def release(self, idx: int, event: torch.cuda.Event) -> None:
"""Attach a compute-done event -- waited before this buffer is recycled."""
def release(self, idx: int, event: StreamEvent | None) -> None:
"""Attach a compute-done guard, waited before this buffer is recycled
(``None`` when the backend needs no guard)."""
self._events[idx] = event
def mark_block_done(self, idx: int) -> None:
"""Record a compute-done guard for block *idx* and queue it for slot reuse.
Called once the block's forward pass has been enqueued, so the buffer is
not overwritten by a later copy until this compute completes."""
self.release(idx, self._sync.record_compute_done())
def cleanup(self) -> None:
"""Synchronize streams and release all resources."""
self._copy_stream.synchronize()
torch.cuda.current_stream(self._target_device).synchronize()
"""Drain outstanding copy/compute work and release all resources."""
self._sync.synchronize()
self._cache.clear()
self._events.clear()
self._source.cleanup()
@@ -132,19 +137,29 @@ 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 via ``fuse_rule``."""
agg_dtype = self._fuse_rule.aggregation_dtype
"""Fuse LoRA deltas directly into GPU block weights via ``fuse_rule``.
The fusion device+dtype come from :func:`device_fuse_rule`: on MPS it
aggregates the ``B@A`` on the GPU in fp32 (fast, and fp32 avoids the
bf16-on-MPS unreliability), with the rule casting back to the weight
dtype; CUDA/CPU keep the rule's dtype. ``get_ab`` places A/B on the
target device for CUDA/MPS, so aggregation and the in-place fuse stay
co-located there.
"""
rule = device_fuse_rule(self._target_device, self._fuse_rule)
for name, tensor in weights.items():
if not name.endswith(".weight"):
continue
prefix = f"{self._blocks_prefix}.{idx}.{name}".removesuffix(".weight")
products = (
ab
for ab in (s.get_ab(prefix, device=self._target_device, dtype=agg_dtype) for s in self._lora_sources)
for ab in (
s.get_ab(prefix, device=self._target_device, dtype=rule.aggregation_dtype)
for s in self._lora_sources
)
if ab is not None
)
deltas = aggregate_lora_products(products, agg_dtype)
deltas = aggregate_lora_products(products, rule.aggregation_dtype)
if deltas is None:
continue
fused = self._fuse_rule(name, tensor, deltas, _EMPTY_STATE_DICT)
fused = rule(name, tensor, deltas, _EMPTY_STATE_DICT)
tensor.copy_(fused[name])
@@ -2,31 +2,39 @@
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.stream_sync import StreamEvent
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: StreamEvent | None) -> None:
"""Signal that an async operation using these weights is guarded by *event*."""
...
@@ -35,68 +43,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: StreamEvent | 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: StreamEvent | 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)
@@ -0,0 +1,174 @@
"""Copy/compute synchronization for block streaming, abstracted across backends.
Weight streaming overlaps an H2D weight copy with block compute. The two
operations must be ordered both ways:
* copy -> compute: a block must not be read before its weights have landed.
* compute -> reuse: a GPU buffer slot must not be overwritten by the next
copy until the compute that read it has finished.
On CUDA these are expressed with a dedicated copy stream and cross-stream
events. MPS exposes no user-facing streams (only ``torch.mps.Event`` on a single
implicit queue), and CPU is fully synchronous. :class:`StreamSync` hides those
differences behind one protocol; :func:`create_stream_sync` picks the backend
implementation. The event types (``torch.cuda.Event`` / ``torch.mps.Event``)
share the small :class:`StreamEvent` surface the pool and source rely on.
Kept internal to the streaming module -- nothing else needs stream coordination.
"""
from __future__ import annotations
import contextlib
from typing import Protocol, runtime_checkable
import torch
from ltx_core.devices import is_mps_available
@runtime_checkable
class StreamEvent(Protocol):
"""A device synchronization marker (``torch.cuda.Event`` / ``torch.mps.Event``)."""
def wait(self) -> None:
"""Device-side: make subsequently queued work wait for this event."""
...
def synchronize(self) -> None:
"""Host-side: block the calling thread until this event completes."""
...
class StreamSync(Protocol):
"""Coordinates the H2D copy against block compute for one streaming model."""
@property
def is_async_copy(self) -> bool:
"""Whether H2D copies may be enqueued asynchronously."""
...
def copy_scope(self) -> contextlib.AbstractContextManager[None]:
"""Context to enqueue the H2D copy under (the copy stream on CUDA)."""
...
def commit_copy(self) -> StreamEvent | None:
"""Record a copy-done event and make compute wait on it.
Returns the event so the source can guard reuse of the CPU buffer (the
disk path host-synchronizes on it), or ``None`` when copies are
synchronous and no guard is needed.
"""
...
def record_compute_done(self) -> StreamEvent | None:
"""Record an event marking the end of a block's compute, for slot reuse."""
...
def reuse_barrier(self, event: StreamEvent | None) -> None:
"""Before a slot is overwritten by a new copy, wait for *event* (prior compute)."""
...
def synchronize(self) -> None:
"""Drain all outstanding copy and compute work."""
...
class CudaStreamSync:
"""CUDA: a dedicated copy stream plus cross-stream events.
H2D copies run on ``copy_stream`` so they overlap compute on the default
stream; events order the two directions explicitly.
"""
def __init__(self, device: torch.device) -> None:
self._device = device
self._copy_stream = torch.cuda.Stream(device=device)
@property
def is_async_copy(self) -> bool:
return True
def copy_scope(self) -> contextlib.AbstractContextManager[None]:
return torch.cuda.stream(self._copy_stream)
def commit_copy(self) -> StreamEvent:
event = torch.cuda.Event()
event.record(self._copy_stream)
torch.cuda.current_stream(self._device).wait_event(event)
return event
def record_compute_done(self) -> StreamEvent:
event = torch.cuda.Event()
event.record(torch.cuda.current_stream(self._device))
return event
def reuse_barrier(self, event: StreamEvent | None) -> None:
if event is not None:
self._copy_stream.wait_event(event)
def synchronize(self) -> None:
self._copy_stream.synchronize()
torch.cuda.current_stream(self._device).synchronize()
class MpsStreamSync:
"""MPS: one implicit queue with ``torch.mps.Event`` markers.
There is no user-facing copy stream, so copy and compute already serialize
on the single default queue. The events make that ordering explicit -- and,
crucially, let the buffer pool guard slot reuse on the compute-done event
rather than relying on the implicit single-queue ordering. ``Event.wait``
enqueues a device-side wait on the default queue (it does not block the
host); ``Event.synchronize`` is the host-blocking variant.
"""
@property
def is_async_copy(self) -> bool:
return False
def copy_scope(self) -> contextlib.AbstractContextManager[None]:
return contextlib.nullcontext()
def commit_copy(self) -> StreamEvent:
event = torch.mps.Event()
event.record()
event.wait()
return event
def record_compute_done(self) -> StreamEvent:
event = torch.mps.Event()
event.record()
return event
def reuse_barrier(self, event: StreamEvent | None) -> None:
if event is not None:
event.wait()
def synchronize(self) -> None:
torch.mps.synchronize()
class SynchronousStreamSync:
"""CPU (and any non-accelerator backend): copies are synchronous, no events."""
@property
def is_async_copy(self) -> bool:
return False
def copy_scope(self) -> contextlib.AbstractContextManager[None]:
return contextlib.nullcontext()
def commit_copy(self) -> None:
return None
def record_compute_done(self) -> None:
return None
def reuse_barrier(self, event: StreamEvent | None) -> None: # noqa: ARG002
return None
def synchronize(self) -> None:
return None
def create_stream_sync(device: torch.device) -> StreamSync:
"""Return the :class:`StreamSync` implementation for *device*'s backend."""
if device.type == "cuda":
return CudaStreamSync(device)
if device.type == "mps" and is_mps_available():
return MpsStreamSync()
return SynchronousStreamSync()
@@ -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,16 +84,17 @@ 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
registration fails. Raises if pinning is requested without a CUDA runtime,
since pinning is fundamentally a CUDA driver operation.
registration fails. Pinning is fundamentally a CUDA driver operation; when
requested without a CUDA runtime (e.g. on MPS/CPU, where H2D copies are
synchronous and pinning is meaningless) it degrades to a normal allocation.
"""
if pin_memory and not torch.cuda.is_available():
pin_memory = False # pinning is CUDA-only; degrade gracefully off-CUDA
if pin_memory and (device is None or torch.device(device).type == "cpu"):
if not torch.cuda.is_available():
raise RuntimeError("pin_memory=True requires CUDA, which is not available")
buf = _alloc_pinned_exact(nbytes)
if buf is not None:
return buf
@@ -112,6 +113,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 +165,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)
@@ -42,6 +42,10 @@ class BlockStreamingWrapper(nn.Module):
self._hooks: list[torch.utils.hooks.RemovableHandle] = []
self._register_hooks()
@property
def num_blocks(self) -> int:
return self._model.num_blocks
# ------------------------------------------------------------------
# Hook registration
# ------------------------------------------------------------------
@@ -55,10 +59,10 @@ class BlockStreamingWrapper(nn.Module):
assign_tensor_to_module(block, name, gpu_weights[name])
def _post_hook(self, block_idx: int) -> None:
"""Record a compute-done event and release the block weights."""
compute_done = torch.cuda.Event()
compute_done.record(torch.cuda.current_stream(self._target_device))
self._provider.release(block_idx, event=compute_done)
"""Release the block weights once its forward pass has been enqueued.
The provider guards the buffer against reuse until this block's compute
completes."""
self._provider.mark_block_done(block_idx)
def _register_hooks(self) -> None:
for idx, block in enumerate(self._blocks):
@@ -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),
+92
View File
@@ -0,0 +1,92 @@
"""Device abstraction for CUDA, Apple Silicon (MPS), and CPU backends.
Centralizes backend detection and the handful of APIs that genuinely differ
across accelerators (synchronization, allocator cache, memory queries, RNG
state). Selection order is CUDA -> MPS -> CPU.
CUDA-only optimizations (FlashAttention, Triton blockwise FP8/FP6,
bitsandbytes, NCCL) are gated at their call sites, not here. MPS in particular
has no ``float64`` support and no fp8 dtype support; use
:func:`highest_precision_float` to stay within what the backend can represent.
"""
from __future__ import annotations
import gc
import logging
import torch
logger = logging.getLogger(__name__)
DeviceSpec = torch.device | None
def is_mps_available() -> bool:
"""Return whether PyTorch can use the Apple Metal/MPS backend."""
mps_backend = getattr(torch.backends, "mps", None)
return bool(mps_backend is not None and mps_backend.is_available())
def get_preferred_device(local_rank: int | None = None) -> torch.device:
"""Prefer CUDA, then MPS, then CPU.
``local_rank`` is only meaningful for CUDA multi-process launches. MPS exposes
a single logical device in PyTorch, so rank-based indexing is not used there.
"""
if torch.cuda.is_available():
index = torch.cuda.current_device() if local_rank is None else local_rank
return torch.device("cuda", index)
if is_mps_available():
return torch.device("mps")
return torch.device("cpu")
def resolve_device(device: DeviceSpec = None, *, local_rank: int | None = None) -> torch.device:
"""Return *device*, or the best available accelerator when it is ``None``."""
if device is None:
return get_preferred_device(local_rank=local_rank)
return device
def supports_float64(device: DeviceSpec) -> bool:
"""Return whether *device* can represent ``torch.float64``.
MPS has no double-precision support; CUDA and CPU do.
"""
return resolve_device(device).type != "mps"
def highest_precision_float(device: DeviceSpec) -> torch.dtype:
"""Return the widest float the backend supports: ``float64`` on CUDA/CPU,
``float32`` on MPS.
Use for numerically sensitive accumulators (e.g. sampler ODE math) that
request double precision but must degrade gracefully on MPS.
"""
return torch.float64 if supports_float64(device) else torch.float32
def synchronize_device(device: DeviceSpec = None) -> None:
"""Synchronize CUDA or MPS work if the selected backend supports it."""
resolved = resolve_device(device)
if resolved.type == "cuda" and torch.cuda.is_available():
torch.cuda.synchronize(resolved)
elif resolved.type == "mps" and is_mps_available():
torch.mps.synchronize()
def empty_device_cache(device: DeviceSpec = None) -> None:
"""Release cached allocator memory for CUDA or MPS."""
resolved = resolve_device(device)
if resolved.type == "cuda" and torch.cuda.is_available():
torch.cuda.empty_cache()
elif resolved.type == "mps" and is_mps_available():
torch.mps.empty_cache()
def cleanup_accelerator_memory(device: DeviceSpec = None) -> None:
"""Run Python GC and release CUDA/MPS allocator caches."""
gc.collect()
empty_device_cache(device)
synchronize_device(device)
try:
if hasattr(torch._C, "_host_emptyCache"):
torch._C._host_emptyCache()
except Exception:
logger.warning("Host empty cache cleanup failed; ignoring.", exc_info=True)
@@ -1,17 +1,19 @@
from dataclasses import dataclass
from enum import Enum
from enum import IntEnum
import torch
from torch._prims_common import DeviceLikeType
class PerturbationType(Enum):
"""Types of attention perturbations for STG (Spatio-Temporal Guidance)."""
class PerturbationType(IntEnum):
"""Types of attention perturbations for STG (Spatio-Temporal Guidance).
The integer value is the row index into ``BatchedPerturbationConfig._block_masks`` dim 0.
"""
SKIP_A2V_CROSS_ATTN = "skip_a2v_cross_attn"
SKIP_V2A_CROSS_ATTN = "skip_v2a_cross_attn"
SKIP_VIDEO_SELF_ATTN = "skip_video_self_attn"
SKIP_AUDIO_SELF_ATTN = "skip_audio_self_attn"
SKIP_VIDEO_SELF_ATTN = 0
SKIP_AUDIO_SELF_ATTN = 1
SKIP_A2V_CROSS_ATTN = 2
SKIP_V2A_CROSS_ATTN = 3
@dataclass(frozen=True)
@@ -48,32 +50,84 @@ class PerturbationConfig:
return PerturbationConfig([])
@dataclass(frozen=True)
class BatchedPerturbationConfig:
"""Perturbation configurations for a batch, with utilities for generating attention masks."""
"""Per-block attention keep-masks for a batch, built once from a list of per-sample configs.
Construction materializes ``_block_masks`` -- a ``(len(PerturbationType), num_blocks, B)`` tensor
(1 = keep, 0 = perturbed) whose dim-0 row index is the ``PerturbationType`` value -- from the
perturbation structure. The per-sample config list is NOT retained: every consumer reads the
tensor (``mask`` indexes it; ``any_in_batch`` / ``all_in_batch`` read the host mirror).
The host build (reading the Python structure) happens here, in ``__init__``, so it MUST be run
eagerly OUTSIDE any ``torch.compile`` / CUDA-graph-capture region. The compiled block then reads
perturbation purely as the runtime ``_block_masks`` tensor and never recompiles per config.
"""
perturbations: list[PerturbationConfig]
_block_masks: torch.Tensor # keep-mask on the compute device, indexed [PerturbationType, block, sample]
# Host mirror so any_in_batch / all_in_batch stay sync-free and graph-break-free. Present for
# configs that may hit the eager skip shortcuts; None for compiled-only configs built via
# ``from_masks`` (the compiled processor reads only ``_block_masks``).
_block_masks_cpu: torch.Tensor | None
def mask(
self, perturbation_type: PerturbationType, block: int, device: DeviceLikeType, dtype: torch.dtype
) -> torch.Tensor:
mask = torch.ones((len(self.perturbations),), device=device, dtype=dtype)
for batch_idx, perturbation in enumerate(self.perturbations):
if perturbation.is_perturbed(perturbation_type, block):
mask[batch_idx] = 0
def __init__(
self,
perturbations: list[PerturbationConfig],
num_blocks: int,
device: DeviceLikeType | None = None,
dtype: torch.dtype | None = None,
) -> None:
keep = [
[
[not pc.is_perturbed(PerturbationType(direction), block) for pc in perturbations]
for block in range(num_blocks)
]
for direction in range(len(PerturbationType))
]
self._block_masks_cpu = torch.tensor(keep, dtype=dtype, device="cpu")
self._block_masks = self._block_masks_cpu if device is None else self._block_masks_cpu.to(device)
return mask
@classmethod
def from_masks(
cls, block_masks: torch.Tensor, block_masks_cpu: torch.Tensor | None = None
) -> "BatchedPerturbationConfig":
"""Construct from prebuilt mask tensors (e.g. a batch-dim slice), bypassing the host build.
``block_masks_cpu`` is only consumed by ``any_in_batch`` / ``all_in_batch`` (the eager
processor's skip shortcuts); pass it when the result may take that path. The compiled
processor reads only ``_block_masks``, so callers on that path may omit the mirror.
"""
obj = cls.__new__(cls)
obj._block_masks = block_masks
obj._block_masks_cpu = block_masks_cpu
return obj
def mask_like(self, perturbation_type: PerturbationType, block: int, values: torch.Tensor) -> torch.Tensor:
mask = self.mask(perturbation_type, block, values.device, values.dtype)
return mask.view(mask.numel(), *([1] * len(values.shape[1:])))
def batch_slice(self, start: int, end: int) -> "BatchedPerturbationConfig":
"""A view over samples ``[start:end]`` of the batch, by slicing the mask tensors.
Slicing (never rebuilding) keeps the host mask build outside any compiled / capture region.
"""
cpu_mask = self._block_masks_cpu[:, :, start:end] if self._block_masks_cpu is not None else None
return BatchedPerturbationConfig.from_masks(self._block_masks[:, :, start:end], cpu_mask)
def mask(self, perturbation_type: PerturbationType, block: int) -> torch.Tensor:
"""This block's ``(B, 1, 1)`` keep-mask for one perturbation type, as an OWNED tensor.
A ``clone`` (not a view into ``_block_masks``) so the masks attached to a block
(e.g. self- and cross-attention) don't alias the same storage -- aliased graph inputs are
fragile under ``torch.compile``.
"""
return self._block_masks[perturbation_type, block].reshape(-1, 1, 1).clone()
def any_in_batch(self, perturbation_type: PerturbationType, block: int) -> bool:
return any(perturbation.is_perturbed(perturbation_type, block) for perturbation in self.perturbations)
assert self._block_masks_cpu is not None, "host mirror required by the skip-shortcut processor path"
return bool((self._block_masks_cpu[perturbation_type, block] == 0).any())
def all_in_batch(self, perturbation_type: PerturbationType, block: int) -> bool:
return all(perturbation.is_perturbed(perturbation_type, block) for perturbation in self.perturbations)
assert self._block_masks_cpu is not None, "host mirror required by the skip-shortcut processor path"
return bool((self._block_masks_cpu[perturbation_type, block] == 0).all())
@staticmethod
def empty(batch_size: int) -> "BatchedPerturbationConfig":
return BatchedPerturbationConfig([PerturbationConfig.empty() for _ in range(batch_size)])
def empty(
batch_size: int,
num_blocks: int,
device: DeviceLikeType | None = None,
dtype: torch.dtype | None = None,
) -> "BatchedPerturbationConfig":
return BatchedPerturbationConfig(
[PerturbationConfig.empty() for _ in range(batch_size)], num_blocks, device, dtype
)
@@ -1,5 +1,5 @@
from collections.abc import Callable, Iterable, Iterator
from dataclasses import dataclass
from dataclasses import dataclass, replace
from typing import NamedTuple
import torch
@@ -71,7 +71,26 @@ def _bf16_fuse(
bf16_fuse_rule = FuseRule(aggregation_dtype=torch.bfloat16, fuse_fn=_bf16_fuse)
def _get_device() -> torch.device:
def device_fuse_rule(target_device: torch.device, base_rule: FuseRule) -> FuseRule:
"""Return the fuse rule to use when fusing onto *target_device*.
On MPS, swap the rule's aggregation dtype to fp32: the LoRA ``B@A`` then runs
on the GPU (far faster than fusing on CPU) and fp32 sidesteps the bf16-on-MPS
numerical unreliability that would otherwise force the slow CPU path. The
rule's ``fuse_fn`` still casts the fused result back to the weight dtype.
CUDA/CPU keep *base_rule* unchanged.
"""
if target_device.type == "mps":
return replace(base_rule, aggregation_dtype=torch.float32)
return base_rule
def _fusion_device(target_device: torch.device) -> torch.device:
"""Device to run the fusion on: the target's own accelerator (CUDA/MPS), else
CUDA when present (accelerating a CPU-resident fuse), else CPU. The caller
moves the fused result back to the weight's device afterwards.
"""
if target_device.type in ("cuda", "mps"):
return target_device
if torch.cuda.is_available():
return torch.device("cuda", torch.cuda.current_device())
return torch.device("cpu")
@@ -110,21 +129,22 @@ def fuse_lora_weights(
used for fusion; caller is responsible for moving them to their final
destination.
"""
fusion_device = _get_device()
rule = device_fuse_rule(model_sd.device, fuse_rule)
fusion_device = _fusion_device(model_sd.device)
for key in _affected_weight_keys(lora_sd_and_strengths):
original_weight = model_sd.sd.get(key)
if original_weight is None:
continue
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)
products = _products_for_sd_key(lora_sd_and_strengths, key, rule.aggregation_dtype, fusion_device)
deltas = aggregate_lora_products(products, rule.aggregation_dtype)
if deltas is None:
continue
original_device = original_weight.device
weight = original_weight.to(device=fusion_device)
fused = fuse_rule(key, weight, deltas, model_sd)
fused = 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
@@ -1,18 +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]]
@@ -50,74 +53,80 @@ 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 checkpoint(self) -> str | tuple[str, ...]:
"""Path(s) to the checkpoint this builder loads from (for logging/diagnostics)."""
...
@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") -> "ModelBuilderProtocol[ModelType]":
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:
@@ -140,8 +149,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,6 +1,8 @@
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
@@ -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__)
@@ -78,10 +83,12 @@ def _load_model_weights(
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).
@@ -98,54 +105,110 @@ class SingleGPUModelBuilder(Generic[ModelType], ModelBuilderProtocol[ModelType],
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"))
fuse_rule: FuseRule = bf16_fuse_rule
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 checkpoint(self) -> str | tuple[str, ...]:
return self._model_path
def with_fuse_rule(self, fuse_rule: FuseRule) -> "SingleGPUModelBuilder":
return replace(self, fuse_rule=fuse_rule)
@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,
@@ -155,18 +218,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,
fuse_rule=self.fuse_rule,
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)
@@ -93,7 +93,7 @@ class VideoModalityTilingHelper:
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(
raise RuntimeError(
f"Tile with in_coords={tile.in_coords} is not in this helper's tile set; "
f"pass a tile obtained from `helper.tiles`."
)
@@ -164,7 +164,7 @@ class VideoModalityTilingHelper:
if output is not None:
if output.shape != expected_shape:
raise ValueError(f"Expected output shape {expected_shape}, got {output.shape}")
raise RuntimeError(f"Expected output shape {expected_shape}, got {output.shape}")
result = output
else:
result = torch.zeros(*expected_shape, device=tile_to_blend.device, dtype=tile_to_blend.dtype)
@@ -1,4 +1,6 @@
import contextlib
import math
from collections.abc import Iterator
from typing import List
import einops
@@ -13,6 +15,25 @@ def get_padding(kernel_size: int, dilation: int = 1) -> int:
return int((kernel_size * dilation - dilation) / 2)
@contextlib.contextmanager
def _module_in_fp32(module: nn.Module, *, enabled: bool) -> Iterator[None]:
"""Temporarily cast *module* to float32, restoring its original dtype on exit.
Used for the MPS vocoder path where fp32 autocast is unavailable, so the
weights must be materialized in float32 for the forward pass. Restores to the
module's original weight dtype (captured here), not the input dtype. When
*enabled* is False this is a no-op, so callers can wrap unconditionally.
"""
if not enabled:
yield
return
module_dtype = next(module.parameters()).dtype
module.float()
try:
yield
finally:
module.to(module_dtype)
# ---------------------------------------------------------------------------
# Anti-aliased resampling helpers (kaiser-sinc filters) for BigVGAN v2
# Adopted from https://github.com/NVIDIA/BigVGAN
@@ -564,15 +585,29 @@ class VocoderWithBWE(nn.Module):
# compound through 108 sequential convolutions and degrade spectral
# metrics (mel_l1, MRSTFT) by 40-90% while perceptual quality (CDPAM)
# is unaffected. fp32 eliminates this degradation.
# We use autocast(dtype=float32) rather than self.float() because it
# upcasts bf16 weights per-op at kernel level, avoiding the temporary
# memory spike of self.float() / self.to(original_dtype).
# On CUDA/CPU we use autocast(dtype=float32) rather than self.float()
# because it upcasts bf16 weights per-op at kernel level, avoiding the
# temporary memory spike of self.float() / self.to(original_dtype).
# Benchmarked on H100 (128.5M-param model):
# autocast fp32: +70 MB peak VRAM, 123 ms (vs 482 MB / 95 ms for bf16)
# model.float(): +324 MB peak VRAM, 149 ms
# Tested: both approaches produce bit-identical output.
# MPS autocast does not upcast conv weights to fp32 (it only supports
# lower-precision autocast dtypes), which would leave the float32 input
# running against bf16 conv weights and raise a dtype mismatch. There we
# fall back to materializing the weights in fp32 for the pass (bit-identical
# per the note above; the memory spike is negligible for this small model).
# The vocoder is normally built in fp32 on MPS, so this fallback is then a
# no-op -- it only triggers if a bf16 module is run on MPS directly.
device_type = mel_spec.device.type
module_dtype = next(self.parameters()).dtype
fp32_ctx = (
_module_in_fp32(self, enabled=module_dtype != torch.float32)
if device_type == "mps"
else torch.autocast(device_type=device_type, dtype=torch.float32)
)
with torch.autocast(device_type=mel_spec.device.type, dtype=torch.float32):
with fp32_ctx:
x = self.vocoder(mel_spec.float())
_, _, length_low_rate = x.shape
output_length = length_low_rate * self.output_sampling_rate // self.input_sampling_rate
@@ -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,29 @@ 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(...)``.
"""
@property
def num_blocks(self) -> int:
"""Number of transformer blocks, delegated through any wrappers to the ``LTXModel``."""
...
def forward(
self,
video: Modality | None,
audio: Modality | None,
perturbations: BatchedPerturbationConfig | None,
) -> tuple[torch.Tensor | None, torch.Tensor | None]: ...
def __call__(
self,
video: Modality | None,
audio: Modality | None,
perturbations: BatchedPerturbationConfig | None,
) -> 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,5 +1,5 @@
import functools
import logging
import sys
from dataclasses import dataclass, field
from enum import Enum
from typing import Protocol
@@ -15,8 +15,6 @@ from ltx_core.model.transformer.ops import (
)
from ltx_core.model.transformer.rope import LTXRopeType
logger = logging.getLogger(__name__)
def _torch_default_sdpa_priority() -> list[SDPBackend]:
"""Fetch torch's current default SDPA priority order at runtime.
@@ -30,28 +28,26 @@ def _torch_default_sdpa_priority() -> list[SDPBackend]:
return [SDPBackend(p) for p in torch._C._get_sdp_priority_order()]
memory_efficient_attention = None
flash_attn_interface = None
flash_attn_4_func = None
try:
from xformers.ops import memory_efficient_attention
except ImportError:
memory_efficient_attention = None
try:
# FlashAttention3 and XFormersAttention cannot be used together
if memory_efficient_attention is None:
import flash_attn_interface
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
try:
# macOS only: routes SDPA to Apple's prebuilt MPSGraph attention kernel.
from mps_sdpa import sdpa_opt as _mps_sdpa_opt
except ImportError:
_mps_sdpa_opt = 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
this protocol; backends that support masks too (Pytorch/SDPA) are
structurally usable here and as :class:`MaskedAttentionCallable`."""
def __call__(self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int) -> torch.Tensor: ...
@@ -76,9 +72,8 @@ class PytorchAttention(AttentionCallable):
@property
def label(self) -> str:
"""Human-readable identifier (used in the AUTOMATIC selection log).
Encodes the SDPA priority list so a single-backend pin reads differently
from the full-priority dispatcher walk."""
"""Human-readable identifier. 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__(
@@ -104,50 +99,48 @@ class PytorchAttention(AttentionCallable):
return out
class XFormersAttention(AttentionCallable):
label = "xFormers"
class MPSSdpaAttention(AttentionCallable):
"""Apple-fused scaled-dot-product attention on MPS.
Routes to ``mps_sdpa.sdpa_opt``, which calls Apple's prebuilt
``MPSGraph.scaledDotProductAttention`` kernel (via a zero-copy bridge)
instead of torch's ``sdpa_general_mps`` graph. The Apple kernel does not
materialize the ``[B, H, Nq, Nk]`` score matrix, so it avoids the
long-sequence memory wall that makes torch's materializing MPS SDPA
unusable on video latents (~32x faster at a 14k-token latent on an M4 Pro).
It is a hard dependency on Apple Silicon (the ``mps-sdpa`` platform-marked
requirement), so AUTOMATIC always has it on MPS. Unlike a JIT-compiled Metal
flash kernel it needs no runtime shader compilation, so it is robust
across macOS / Metal revisions.
Accepts an optional additive-float or boolean ``mask`` broadcastable to
``[B, H, Nq, Nk]``, so it serves both the unmasked and masked protocols.
"""
@property
def label(self) -> str:
return "MPS-SDPA"
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 | None = None
) -> torch.Tensor:
if memory_efficient_attention is None:
raise RuntimeError("XFormersAttention was selected but `xformers` is not installed.")
if _mps_sdpa_opt is None:
raise RuntimeError("MPSSdpaAttention was selected but `mps-sdpa` is not installed.")
if q.device.type != "mps":
raise RuntimeError("MPSSdpaAttention requires MPS. Use PyTorch SDPA on CPU or CUDA.")
b, _, dim_head = q.shape
dim_head //= heads
# xformers expects [B, M, H, K]
q, k, v = (t.view(b, -1, heads, dim_head) for t in (q, k, v))
q, k, v = (t.view(b, -1, heads, dim_head).transpose(1, 2) for t in (q, k, v))
if mask is not None:
# add a singleton batch dimension
# add a batch dimension if there isn't already one
if mask.ndim == 2:
mask = mask.unsqueeze(0)
# add a singleton heads dimension
# add a heads dimension if there isn't already one
if mask.ndim == 3:
mask = mask.unsqueeze(1)
# pad to a multiple of 8
pad = 8 - mask.shape[-1] % 8
# the xformers docs says that it's allowed to have a mask of shape (1, Nq, Nk)
# but when using separated heads, the shape has to be (B, H, Nq, Nk)
# in flux, this matrix ends up being over 1GB
# here, we create a mask with the same batch/head size as the input mask (potentially singleton or full)
mask_out = torch.empty(
[mask.shape[0], mask.shape[1], q.shape[1], mask.shape[-1] + pad], dtype=q.dtype, device=q.device
)
mask_out[..., : mask.shape[-1]] = mask
# doesn't this remove the padding again??
mask = mask_out[..., : mask.shape[-1]]
mask = mask.expand(b, heads, -1, -1)
out = memory_efficient_attention(q.to(v.dtype), k.to(v.dtype), v, attn_bias=mask, p=0.0)
out = out.reshape(b, -1, heads * dim_head)
out = _mps_sdpa_opt(q, k, v, attn_mask=mask)
out = out.transpose(1, 2).reshape(b, -1, heads * dim_head)
return out
@@ -163,6 +156,8 @@ class FlashAttention3(AttentionCallable):
) -> torch.Tensor:
if flash_attn_interface is None:
raise RuntimeError("FlashAttention3 was selected but `FlashAttention3` is not installed.")
if q.device.type != "cuda":
raise RuntimeError("FlashAttention3 requires CUDA. Use PyTorch SDPA on CPU or MPS.")
b, _, dim_head = q.shape
dim_head //= heads
@@ -186,6 +181,8 @@ class FlashAttention4(AttentionCallable):
) -> torch.Tensor:
if flash_attn_4_func is None:
raise RuntimeError("FlashAttention4 was selected but `flash-attn-4` is not installed.")
if q.device.type != "cuda":
raise RuntimeError("FlashAttention4 requires CUDA. Use PyTorch SDPA on CPU or MPS.")
b, _, dim_head = q.shape
dim_head //= heads
@@ -199,10 +196,9 @@ class FlashAttention4(AttentionCallable):
# --- 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)
# and logs the resulting label once. 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).
# 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 + SDPA masked on H100).
def _sdpa_can_use(backend: SDPBackend, *, with_mask: bool) -> bool:
@@ -240,6 +236,20 @@ _SDPA_FULL_PRIORITY: tuple[SDPBackend, ...] = (
)
def _on_macos() -> bool:
"""True on macOS, where torch's native SDPA materializes the score matrix and
AUTOMATIC routes to Apple's fused ``mps-sdpa`` kernel instead."""
return sys.platform == "darwin"
def _mps_sdpa_available() -> bool:
"""True when the ``mps-sdpa`` package is importable. It is a platform-marked
hard dependency on Apple Silicon, so this is always True there; it is only
False on non-Apple-Silicon macs (e.g. Intel/CPU), where AUTOMATIC falls back
to torch's SDPA (acceptable on CPU, which has no MPS memory wall)."""
return _mps_sdpa_opt is not None
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
@@ -257,10 +267,14 @@ def _sdpa_full_priority() -> PytorchAttention:
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.
- Hopper (sm_90, H100): FA3 > 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.
- macOS (Apple Silicon / MPS): Apple's fused MPSGraph kernel via ``mps-sdpa``
(a platform-marked hard dependency on Apple Silicon) -- it avoids the
full-score-matrix memory wall on long video sequences. On a non-Apple-Silicon
mac (Intel/CPU) it falls back to torch's SDPA.
- Everywhere else (Ada, Ampere, CPU): SDPA with the full backend priority
list -- torch's runtime dispatcher picks the best fit at call time.
"""
@@ -269,40 +283,45 @@ def _select_primary_attention() -> AttentionCallable:
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()
if _on_macos():
return MPSSdpaAttention() if _mps_sdpa_available() else _sdpa_full_priority()
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()
"""Pick a mask-aware attention. On macOS, Apple's fused MPSGraph kernel via
``mps-sdpa`` (a hard dependency on Apple Silicon, else torch's SDPA on
Intel/CPU macs); else SDPA with the full priority list (the dispatcher
rejects FLASH automatically when a mask is present and walks past it --
torch SDPA handles the additive mask directly)."""
if _on_macos():
return MPSSdpaAttention() if _mps_sdpa_available() else _sdpa_full_priority()
return _sdpa_full_priority()
@functools.cache
def automatic_attention() -> AttentionCallable:
"""Cached AUTOMATIC pick for the unmasked path. Logs the chosen label once
per process."""
fn = _select_primary_attention()
logger.info("Automatic attention selected: %s", fn.label)
return fn
"""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. Logs the chosen label once
per process."""
fn = _select_masked_attention()
logger.info("Automatic masked attention selected: %s", fn.label)
return fn
"""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:
@@ -324,50 +343,60 @@ def _resolve_sdpa_variant(backend: SDPBackend, name: str, *, with_mask: bool) ->
class AttentionFunction(Enum):
PYTORCH = "pytorch"
XFORMERS = "xformers"
FLASH_ATTENTION_3 = "flash_attention_3"
FLASH_ATTENTION_4 = "flash_attention_4"
SDPA_CUDNN = "sdpa_cudnn"
SDPA_FLASH = "sdpa_flash"
SDPA_EFFICIENT = "sdpa_efficient"
SDPA_MATH = "sdpa_math"
# Apple's fused MPSGraph SDPA via the `mps-sdpa` package (macOS/MPS only, a
# platform-marked hard dependency on Apple Silicon). The AUTOMATIC default on
# MPS; never materializes the score matrix.
MPS_SDPA = "mps_sdpa"
# 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: # noqa: PLR0911
def to_callable(self) -> AttentionCallable: # noqa: PLR0911, PLR0912
"""Resolve to a concrete callable. Use this at module init time so that
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 the once-per-process log
fires only on the first resolution.
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."
)
if not torch.cuda.is_available():
raise RuntimeError(
"AttentionFunction.FLASH_ATTENTION_3 requires CUDA. Use PyTorch SDPA on CPU or MPS."
)
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."
)
if not torch.cuda.is_available():
raise RuntimeError(
"AttentionFunction.FLASH_ATTENTION_4 requires CUDA. Use PyTorch SDPA on CPU or MPS."
)
return FlashAttention4()
case AttentionFunction.SDPA_MATH:
return PytorchAttention(priority=[SDPBackend.MATH])
case AttentionFunction.MPS_SDPA:
if _mps_sdpa_opt is None:
raise RuntimeError("AttentionFunction.MPS_SDPA selected but `mps-sdpa` is not installed.")
return MPSSdpaAttention()
case AttentionFunction.SDPA_CUDNN:
return _resolve_sdpa_variant(
SDPBackend.CUDNN_ATTENTION, "AttentionFunction.SDPA_CUDNN", with_mask=False
@@ -390,10 +419,13 @@ class MaskedAttentionFunction(Enum):
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"
# Apple's fused MPSGraph SDPA via the `mps-sdpa` package (macOS/MPS only, a
# platform-marked hard dependency on Apple Silicon); the AUTOMATIC default on
# MPS. Mask-aware.
MPS_SDPA = "mps_sdpa"
# Pick the fastest mask-capable backend for the current extras combo; see
# :func:`automatic_masked_attention`. Default for the masked slot of
# :class:`AttentionOps`.
@@ -412,12 +444,12 @@ class MaskedAttentionFunction(Enum):
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.MPS_SDPA:
if _mps_sdpa_opt is None:
raise RuntimeError("MaskedAttentionFunction.MPS_SDPA selected but `mps-sdpa` is not installed.")
return MPSSdpaAttention()
case MaskedAttentionFunction.SDPA_CUDNN:
return _resolve_sdpa_variant(
SDPBackend.CUDNN_ATTENTION, "MaskedAttentionFunction.SDPA_CUDNN", with_mask=True
@@ -501,7 +533,7 @@ 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). A non-None
backend (additive bias for 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``.
@@ -1,4 +1,4 @@
from dataclasses import dataclass, field
from dataclasses import dataclass, field, replace
from typing import Any
import torch
@@ -11,7 +11,7 @@ from ltx_core.model.transformer.transformer_args import BlockPerturbationsProces
# Defaults applied inside the patched forward. Overriding via CompilationConfig
# replaces these wholesale; it does not merge.
_DEFAULT_INDUCTOR_CONFIG: dict[str, Any] = {"unsafe_skip_cache_dynamic_shape_guards": True}
_DEFAULT_INDUCTOR_CONFIG: dict[str, Any] = {}
_DEFAULT_DYNAMO_CONFIG: dict[str, Any] = {"inline_inbuilt_nn_modules": True, "cache_size_limit": 256}
@@ -27,19 +27,19 @@ class CompilationConfig:
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.
class CompiledBlockPerturbationsProcessor(BlockPerturbationsProcessor):
"""Per-block input prep for compiled blocks: mark the seq dim dynamic, then attach perturbation
as config-independent runtime masks so the block traces ONCE.
The ``mark_dynamic`` calls keep the per-block compile artifact shape-polymorphic; they run in
eager mode (this processor lives outside the compiled region) on the tensors about to cross into
the trace. Both keep-masks are then attached UNCONDITIONALLY and the skip flags pinned False, so
the trace is identical for every pass (cond / uncond / STG): the block never sees a flipped
Python bool (``self_attn_all_perturbed``) or a None-vs-tensor mask that Dynamo would specialise
on, so the STG pass no longer triggers a recompile. An all-keep mask blends to a no-op
(``out*1 + v*0``); an all-zero mask reproduces the skip. Reads only ``mask`` (runtime tensor
indexing), never the host-side ``all_in_batch`` / ``any_in_batch``.
"""
def __init__(self, inner: BlockPerturbationsProcessor) -> None:
self.inner = inner
def __call__(
self,
args: TransformerArgs,
@@ -84,7 +84,14 @@ class _SeqDynamicMarkingProcessor:
# 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)
# Perturbation as config-independent runtime masks (skip flags pinned False -> no recompile).
return replace(
args,
self_attn_perturbation_mask=perturbations.mask(self_attn_type, block_idx),
self_attn_all_perturbed=False,
cross_attn_perturbation_mask=perturbations.mask(cross_attn_type, block_idx),
cross_attn_skip_all=False,
)
def compile_transformer(model: LTXModel, config: CompilationConfig) -> LTXModel:
@@ -100,7 +107,7 @@ def compile_transformer(model: LTXModel, config: CompilationConfig) -> LTXModel:
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)
model.block_input_processor = CompiledBlockPerturbationsProcessor()
def patched_dynamo_forward(*args, **kwargs) -> tuple[torch.Tensor, torch.Tensor]:
torch.compiler.cudagraph_mark_step_begin()
@@ -1,9 +1,12 @@
import logging
from enum import Enum
import torch
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 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 (
@@ -20,6 +23,8 @@ from ltx_core.model.transformer.transformer_args import (
)
from ltx_core.utils import to_denoised
logger = logging.getLogger(__name__)
class LTXModelType(Enum):
AudioVideo = "ltx av model"
@@ -70,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/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
@@ -344,21 +358,18 @@ class LTXModel(torch.nn.Module):
"""
self._enable_gradient_checkpointing = enable
@property
def num_blocks(self) -> int:
"""Number of transformer blocks."""
return len(self.transformer_blocks)
def _process_transformer_blocks(
self,
video: TransformerArgs | None,
audio: TransformerArgs | None,
perturbations: BatchedPerturbationConfig | None,
perturbations: BatchedPerturbationConfig,
) -> 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)
"""Process transformer blocks for LTX."""
for block_idx, block in enumerate(self.transformer_blocks):
if video is not None:
video = self.block_input_processor(
@@ -410,8 +421,8 @@ class LTXModel(torch.nn.Module):
return x
def forward(
self, video: Modality | None, audio: Modality | None, perturbations: BatchedPerturbationConfig
) -> tuple[torch.Tensor, torch.Tensor]:
self, video: Modality | None, audio: Modality | None, perturbations: BatchedPerturbationConfig | None
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""
Forward pass for LTX models.
Returns:
@@ -424,6 +435,11 @@ class LTXModel(torch.nn.Module):
video_args = self.video_args_preprocessor.prepare(video, audio) if video is not None else None
audio_args = self.audio_args_preprocessor.prepare(audio, video) if audio is not None else None
# Materialize the no-perturbation mask here (eager); a None config means "perturb nothing"
# -> all-keep masks. The block loop never builds masks.
if perturbations is None:
ref = (video_args or audio_args).x
perturbations = BatchedPerturbationConfig.empty(ref.shape[0], self.num_blocks, ref.device, ref.dtype)
# Process transformer blocks
video_out, audio_out = self._process_transformer_blocks(
video=video_args,
@@ -459,10 +475,15 @@ 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
@property
def num_blocks(self) -> int:
"""Number of transformer blocks."""
return self.velocity_model.num_blocks
def forward(
self,
video: Modality | None,
@@ -488,15 +509,20 @@ 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
@property
def num_blocks(self) -> int:
"""Number of transformer blocks."""
return self.velocity_model.num_blocks
def forward(
self,
video: Modality | None,
audio: Modality | None,
perturbations: BatchedPerturbationConfig,
perturbations: BatchedPerturbationConfig | None,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""
Denoise the video and audio according to the sigma.
@@ -123,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,
@@ -151,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.", "")
)
@@ -24,7 +24,6 @@ from ltx_core.model.transformer.ops import (
)
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
@@ -222,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,
@@ -232,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,
@@ -245,9 +247,8 @@ 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,
@@ -280,10 +281,10 @@ class BasicAVTransformerBlock(torch.nn.Module):
perturbation_mask=video.self_attn_perturbation_mask,
all_perturbed=video.self_attn_all_perturbed,
)
vx = vx + vx_msa_out * vgate_msa
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,
@@ -293,6 +294,7 @@ 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(
@@ -308,10 +310,10 @@ class BasicAVTransformerBlock(torch.nn.Module):
perturbation_mask=audio.self_attn_perturbation_mask,
all_perturbed=audio.self_attn_all_perturbed,
)
ax = ax + ax_msa_out * agate_msa
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,
@@ -321,6 +323,7 @@ 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:
@@ -414,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,
@@ -423,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
@@ -62,18 +62,16 @@ class BlockPerturbationsProcessor:
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)
self_mask = perturbations.mask(self_attn_type, block_idx)
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)
cross_mask = perturbations.mask(cross_attn_type, block_idx)
return replace(
args,
@@ -0,0 +1,11 @@
"""
Multi-GPU utilities for LTX models.
This package provides utilities for running LTX models across multiple GPUs
using tiled data-parallel techniques and sharded state-dict utilities.
"""
from ltx_core.multigpu import transformer, vae
from ltx_core.multigpu.sharded_sd import ShardedSD
from ltx_core.tiling import DimensionTilingConfig, TileCountConfig
__all__ = ["DimensionTilingConfig", "ShardedSD", "TileCountConfig", "transformer", "vae"]
@@ -0,0 +1,6 @@
"""Multi-GPU utilities for the Gemma text encoder."""
from ltx_core.multigpu.gemma.accelerate_wrapper import AccelerateGemmaWrapper
from ltx_core.multigpu.gemma.loader import load_gemma_with_device_map
__all__ = ["AccelerateGemmaWrapper", "load_gemma_with_device_map"]
@@ -0,0 +1,30 @@
"""Accelerate-based Gemma text encoder wrapper for multi-GPU inference.
One rank (``src_rank``) holds the real ``GemmaTextEncoder`` loaded with
``device_map="auto"``; other ranks hold a lightweight stub. Every public
method runs on the source rank and broadcasts results to all ranks via the
provided NCCL process group.
The ``broadcast_group`` should cover **all ranks that need the
embeddings** (typically the transformer group or world group).
"""
from __future__ import annotations
import torch
from ltx_core.multigpu.gemma.broadcast_wrapper import BroadcastGemmaWrapper
class AccelerateGemmaWrapper(BroadcastGemmaWrapper):
"""Source-rank encode + NCCL broadcast around a sharded ``GemmaTextEncoder``."""
def encode(
self,
prompts: list[str],
padding_side: str = "left",
) -> list[tuple[tuple[torch.Tensor, ...], torch.Tensor]]:
"""Fuse all prompts into one Gemma call on the source rank, broadcast each output."""
if self._rank == self._src_rank:
local_outputs = self._encoder.encode(prompts, padding_side)
else:
local_outputs = [(None, None)] * len(prompts)
return [self._broadcast_encoder_output(hs, mask, self._src_rank) for hs, mask in local_outputs]
@@ -0,0 +1,75 @@
"""Batch-parallel Gemma text encoder wrapper for multi-GPU inference.
Each rank holds a full :class:`GemmaTextEncoder` replica resident on its
own GPU. ``encode`` partitions the prompt list across ranks (each rank
encodes a disjoint slice) and broadcasts every prompt's outputs from its
encoding rank to all other ranks, so all ranks end up with the full list
in the original order.
``enhance_t2v`` / ``enhance_i2v`` (inherited) involve sampling, so they
execute on ``src_rank`` only and the generated string is broadcast.
"""
from __future__ import annotations
import torch
import torch.distributed as dist
from ltx_core.multigpu.gemma.broadcast_wrapper import BroadcastGemmaWrapper
from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder
def _partition(total: int, world_size: int) -> list[int]:
"""Spread ``total`` items across ``world_size`` ranks; remainder lands on the first ranks."""
base, rem = divmod(total, world_size)
return [base + (1 if i < rem else 0) for i in range(world_size)]
class BatchParallelGemmaWrapper(BroadcastGemmaWrapper):
"""Per-rank Gemma replica; ``encode`` parallelises a batch across ranks."""
_encoder: GemmaTextEncoder # always resident on every rank, unlike the base's optional encoder
def __init__(
self,
encoder: GemmaTextEncoder,
broadcast_group: dist.ProcessGroup | None,
src_rank: int,
dtype: torch.dtype = torch.bfloat16,
device: torch.device | None = None,
) -> None:
"""Wrap a per-rank Gemma replica for batch-parallel encoding.
Args:
encoder: Full Gemma replica; required and resident on every rank (unlike
the base, where it is optional and real only on ``src_rank``).
broadcast_group: NCCL group spanning the ranks that share the encode work.
src_rank: Rank within ``broadcast_group`` that runs the inherited sampling
methods (``enhance_t2v`` / ``enhance_i2v``); ``encode`` uses every rank.
dtype: Target dtype for output tensors.
device: Target device for output tensors; defaults to the current CUDA device.
"""
super().__init__(encoder, broadcast_group, src_rank, dtype, device)
self._world_size = dist.get_world_size(broadcast_group)
def encode(
self,
prompts: list[str],
padding_side: str = "left",
) -> list[tuple[tuple[torch.Tensor, ...], torch.Tensor]]:
"""Partition prompts across ranks, encode in parallel, broadcast per-prompt outputs.
With B prompts on W ranks, each rank gets ``ceil(B/W)`` or ``floor(B/W)``
prompts; the typical pos+neg case (B=2, W=2) gives one prompt per rank,
running both Gemma forwards concurrently on different GPUs.
"""
n = len(prompts)
if n == 0:
return []
counts = _partition(n, self._world_size)
start = sum(counts[: self._rank])
local_prompts = prompts[start : start + counts[self._rank]]
local_outputs = self._encoder.encode(local_prompts, padding_side) if local_prompts else []
all_outputs: list[tuple[tuple[torch.Tensor, ...], torch.Tensor]] = []
for owner_rank, owner_count in enumerate(counts):
for slot in range(owner_count):
hs, mask = local_outputs[slot] if owner_rank == self._rank else (None, None)
all_outputs.append(self._broadcast_encoder_output(hs, mask, owner_rank))
return all_outputs
@@ -0,0 +1,108 @@
"""Shared base for the multi-GPU Gemma wrappers: src-rank gating + NCCL broadcast."""
from __future__ import annotations
import torch
import torch.distributed as dist
from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder
class BroadcastGemmaWrapper(torch.nn.Module):
"""Encoder/group plumbing, prompt enhancement, and result broadcast.
Subclasses implement ``encode`` (the stub below raises).
Args:
encoder: The encoder; real on ``src_rank``, may be ``None`` elsewhere.
broadcast_group: NCCL group covering ranks that need the embeddings.
src_rank: Rank *within* ``broadcast_group`` that holds the real encoder and runs
the sampling-based ``enhance_*`` methods; its results are broadcast to every
other rank in the group. Builders derive it from a global driver rank via
``dist.get_group_rank``.
dtype: Target dtype for output tensors.
device: Target device for output tensors.
"""
def __init__(
self,
encoder: GemmaTextEncoder | None,
broadcast_group: dist.ProcessGroup | None,
src_rank: int,
dtype: torch.dtype = torch.bfloat16,
device: torch.device | None = None,
) -> None:
super().__init__()
if device is None and torch.cuda.is_available():
device = torch.device("cuda", torch.cuda.current_device())
self._encoder = encoder
self._group = broadcast_group
self._src_rank = src_rank
self._rank = dist.get_rank(broadcast_group)
self._dtype = dtype
self._device = device
def encode(
self,
prompts: list[str],
padding_side: str = "left",
) -> list[tuple[tuple[torch.Tensor, ...], torch.Tensor]]:
"""Encode a batch of prompts to per-prompt hidden states; implemented by subclasses."""
raise NotImplementedError
def enhance_t2v(
self,
prompt: str,
max_new_tokens: int = 512,
system_prompt: str | None = None,
seed: int = 10,
) -> str:
result = None
if self._rank == self._src_rank:
result = self._encoder.enhance_t2v(prompt, max_new_tokens, system_prompt, seed)
return self._broadcast_str(result)
def enhance_i2v(
self,
prompt: str,
image: torch.Tensor,
max_new_tokens: int = 512,
system_prompt: str | None = None,
seed: int = 10,
) -> str:
result = None
if self._rank == self._src_rank:
result = self._encoder.enhance_i2v(prompt, image, max_new_tokens, system_prompt, seed)
return self._broadcast_str(result)
def _broadcast_str(self, value: str | None) -> str:
obj_list: list[str | None] = [value]
dist.broadcast_object_list(obj_list, src=self._src_rank, group=self._group)
result = obj_list[0]
assert result is not None, "broadcast returned None; check src_rank/broadcast_group"
return result
def _broadcast_encoder_output(
self,
hidden_states: tuple[torch.Tensor, ...] | None,
attention_mask: torch.Tensor | None,
src_rank: int,
) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
"""Broadcast hidden states + attention mask via NCCL from ``src_rank``."""
if self._rank == src_rank:
meta = [{"hs_shapes": [h.shape for h in hidden_states], "mask_shape": attention_mask.shape}]
else:
meta = [None]
dist.broadcast_object_list(meta, src=src_rank, group=self._group)
info = meta[0]
if self._rank != src_rank:
hidden_states = tuple(torch.empty(s, device=self._device, dtype=self._dtype) for s in info["hs_shapes"])
attention_mask = torch.empty(info["mask_shape"], device=self._device, dtype=torch.long)
else:
hidden_states = tuple(h.to(device=self._device, dtype=self._dtype) for h in hidden_states)
attention_mask = attention_mask.to(device=self._device)
for h in hidden_states:
dist.broadcast(h, src=src_rank, group=self._group)
dist.broadcast(attention_mask, src=src_rank, group=self._group)
return hidden_states, attention_mask
@@ -0,0 +1,55 @@
"""Load GemmaTextEncoder with Accelerate ``device_map="auto"``.
The Gemma LLM backbone is spread across available CUDA devices using
HuggingFace Accelerate's automatic device placement.
Mirrors the ``PromptEncoder`` text-encoder loading in
``ltx_pipelines.utils.blocks`` but uses ``device_map="auto"`` instead of
placing the entire model on a single GPU.
"""
from __future__ import annotations
import logging
import torch
from transformers import AutoImageProcessor, Gemma3ForConditionalGeneration, Gemma3Processor
from ltx_core.text_encoders.gemma.encoders.base_encoder import GemmaTextEncoder
from ltx_core.text_encoders.gemma.tokenizer import LTXVGemmaTokenizer
from ltx_core.utils import find_matching_file
logger = logging.getLogger(__name__)
def load_gemma_with_device_map(
gemma_root_path: str,
dtype: torch.dtype = torch.bfloat16,
) -> GemmaTextEncoder:
"""Load GemmaTextEncoder with the LLM backbone spread across GPUs.
Uses ``Gemma3ForConditionalGeneration.from_pretrained(device_map="auto")``
to distribute layers across available CUDA devices.
Args:
gemma_root_path: Path to Gemma model directory.
dtype: Data type for model weights.
"""
model_folder = str(find_matching_file(gemma_root_path, "model*.safetensors").parent)
tokenizer_path = str(find_matching_file(gemma_root_path, "tokenizer.model").parent)
processor_path = str(find_matching_file(gemma_root_path, "preprocessor_config.json").parent)
logger.info("Loading Gemma LLM with device_map='auto'...")
gemma_model = Gemma3ForConditionalGeneration.from_pretrained(
model_folder,
dtype=dtype,
device_map="auto",
local_files_only=True,
)
tokenizer = LTXVGemmaTokenizer(tokenizer_path, 1024)
image_processor = AutoImageProcessor.from_pretrained(processor_path, local_files_only=True, use_fast=False)
processor = Gemma3Processor(image_processor=image_processor, tokenizer=tokenizer.tokenizer)
return GemmaTextEncoder(
model=gemma_model,
tokenizer=tokenizer,
processor=processor,
dtype=dtype,
)
@@ -0,0 +1,191 @@
"""Sharded state dict with distributed weight backup.
Each rank stores ~1/N of the model weights. Bucketed broadcasts
restore weights into a target state dict using only a small,
caller-provided staging buffer.
"""
from __future__ import annotations
import hashlib
from dataclasses import dataclass
import torch
import torch.distributed as dist
def _stable_owner(key: str, world: int) -> int:
"""Deterministic rank assignment (same across all processes)."""
h = hashlib.md5(key.encode("utf-8")).digest()
return int.from_bytes(h[:8], "little") % world
def _nbytes(t: torch.Tensor) -> int:
return t.numel() * t.element_size()
@dataclass
class ShardedSD:
"""Sharded state dict with distributed weight backup.
Distributes model weights across ranks for memory-efficient backup
and restoration. Can be used for any scenario where a full state dict
needs to be restored from sharded storage (e.g. LoRA hot-swap, weight
rollback, checkpoint recovery).
- Deterministic ownership: ``MD5(key) % world_size``
- Local storage only for owned keys (VRAM 1/world_size of model)
- Bucketed broadcast using a single small staging buffer
Usage::
backup = ShardedSD.from_state_dict(model.state_dict(), group)
staging = torch.empty(64 * 1024 * 1024, dtype=torch.uint8, device=device)
backup.broadcast_shards_into(target_sd, staging) # cooperative: all ranks must call
"""
keys: tuple[str, ...]
"""All parameter keys in the original state dict, in insertion order."""
key_sizes: dict[str, int]
"""Byte size of each parameter tensor (numel * element_size)."""
owner_of: dict[str, int]
"""Maps each key to the rank that stores it."""
local_shard: dict[str, torch.Tensor]
"""Tensors owned by this rank (subset of the full state dict)."""
rank: int
"""This process's rank within the group."""
world: int
"""Total number of ranks in the group."""
group: dist.ProcessGroup
"""NCCL process group used for broadcast operations."""
_owner_groups: dict[int, list[str]]
"""Keys grouped by owning rank, sorted by descending tensor size."""
@classmethod
def from_state_dict(
cls,
sd: dict[str, torch.Tensor],
group: dist.ProcessGroup,
clone: bool = True,
) -> ShardedSD:
rank = dist.get_rank(group)
world = dist.get_world_size(group)
keys = tuple(sd.keys())
# Co-locate .weight_scale with its .weight on the same rank.
owner_of: dict[str, int] = {}
for k in keys:
if k.endswith(".weight_scale"):
parent = k.replace(".weight_scale", ".weight")
if parent in sd:
owner_of[k] = _stable_owner(parent, world)
continue
owner_of[k] = _stable_owner(k, world)
key_sizes = {k: _nbytes(v) for k, v in sd.items()}
local_shard: dict[str, torch.Tensor] = {}
for k, v in sd.items():
if owner_of[k] == rank:
local_shard[k] = v.clone() if clone else v
owner_groups: dict[int, list[str]] = {r: [] for r in range(world)}
for k in keys:
owner_groups[owner_of[k]].append(k)
for r in range(world):
owner_groups[r].sort(key=lambda kk: key_sizes[kk], reverse=True)
return cls(
keys=keys,
key_sizes=key_sizes,
owner_of=owner_of,
local_shard=local_shard,
rank=rank,
world=world,
group=group,
_owner_groups=owner_groups,
)
def broadcast_shards_into(
self,
target_sd: dict[str, torch.Tensor],
staging: torch.Tensor,
) -> None:
"""Broadcast stored weights from sharded backup into *target_sd*.
This is a **cooperative operation** all ranks in the process group
must call it simultaneously.
*staging* is a caller-owned ``uint8`` scratch buffer; its size sets the
broadcast granularity (tensors larger than it split across rounds). It
may be shared by instances that never broadcast at the same time. Writes
directly into existing tensors in *target_sd*.
"""
if staging.dtype != torch.uint8 or staging.numel() == 0:
raise ValueError("staging must be a non-empty uint8 buffer")
for owner, klist in self._owner_groups.items():
if klist:
self._broadcast_group(owner, klist, target_sd, staging)
def _broadcast_group(
self,
owner: int,
keys: list[str],
target_sd: dict[str, torch.Tensor],
staging: torch.Tensor,
) -> None:
"""Pack & broadcast params from *owner*, splitting tensors across rounds."""
rounds = self._plan_rounds(keys, staging.numel())
for round_chunks in rounds:
filled = 0
if self.rank == owner:
for k, offset, chunk_size in round_chunks:
src = self.local_shard[k]
if not src.is_contiguous():
raise RuntimeError(f"ShardedSD: local shard tensor '{k}' is not contiguous")
src_bytes = src.view(torch.uint8).view(-1)
staging[filled : filled + chunk_size].copy_(
src_bytes[offset : offset + chunk_size], non_blocking=True
)
filled += chunk_size
else:
filled = sum(chunk_size for (_, _, chunk_size) in round_chunks)
if filled == 0:
continue
view = staging[:filled]
dist.broadcast(view, src=owner, group=self.group)
cursor = 0
for k, offset, chunk_size in round_chunks:
dst = target_sd[k]
if not dst.is_contiguous():
raise RuntimeError(f"ShardedSD: target tensor '{k}' is not contiguous")
dst_bytes = dst.view(torch.uint8).view(-1)
dst_bytes[offset : offset + chunk_size].copy_(staging[cursor : cursor + chunk_size], non_blocking=True)
cursor += chunk_size
def _plan_rounds(self, keys: list[str], capacity: int) -> list[list[tuple[str, int, int]]]:
"""Build rounds that pack a *capacity*-byte buffer, splitting tensors if needed.
Returns a list of rounds, each containing ``(key, byte_offset, chunk_bytes)`` tuples.
"""
rounds: list[list[tuple[str, int, int]]] = []
current: list[tuple[str, int, int]] = []
used = 0
for k in keys:
remaining = self.key_sizes[k]
offset = 0
while remaining > 0:
space = capacity - used
if space == 0:
rounds.append(current)
current = []
used = 0
space = capacity
chunk = min(remaining, space)
current.append((k, offset, chunk))
used += chunk
offset += chunk
remaining -= chunk
if current:
rounds.append(current)
return rounds
@@ -0,0 +1,13 @@
"""
Multi-GPU transformer utilities for LTX models.
This module provides utilities for running LTX transformer models across multiple GPUs
using tiled data parallelism.
"""
from ltx_core.multigpu.transformer.tiled_data_parallel import (
TiledDataParallelModelWrapper,
)
__all__ = [
"TiledDataParallelModelWrapper",
]
@@ -0,0 +1,270 @@
import torch
import torch.distributed as dist
from ltx_core.model.transformer.attention import AttentionCallable, MaskedAttentionCallable
# Mirrors the kernel's DEFAULT_BARRIER_TIMEOUT_SECONDS (configs.cuh), which All2All converts to
# cycles via the device peak SM clock. Stored so the timeout can be read back to reset after a raise.
_DEFAULT_ALL2ALL_TIMEOUT_SECONDS = 10.0
class AttentionManager:
def __init__(
self,
max_tokens: int,
num_heads: int,
head_dim: int,
tensor_dtype: torch.dtype,
group: torch.distributed.ProcessGroup,
copy_out_: bool = False,
) -> None:
# Lazy: ltx_kernels is an optional GPU-only dep, and this constructor already
# requires CUDA -- so importing it here (not at module scope) keeps the multigpu
# modules importable without the kernels installed (e.g. CPU CI test collection).
from ltx_kernels import All2All # noqa: PLC0415
self.rank = dist.get_rank(group)
self.world_size = dist.get_world_size(group)
self.max_tokens = max_tokens
hidden_dim = num_heads * head_dim
num_sms = torch.cuda.get_device_properties(self.rank).multi_processor_count
self.copy_out = copy_out_
buffer_seqlen = (max_tokens + self.world_size - 1) // self.world_size
self.all2all_heads, self.all2all_q = (
All2All(
rank=self.rank,
world_size=self.world_size,
seqlen=buffer_seqlen,
hidden_dim=hidden_dim,
num_sms=num_sms,
tensor_dtype=tensor_dtype,
group=group,
)
for _ in range(2)
)
self.all2all_k, self.all2all_v = (
All2All(
rank=self.rank,
world_size=self.world_size,
seqlen=buffer_seqlen,
hidden_dim=hidden_dim,
num_sms=num_sms,
tensor_dtype=tensor_dtype,
group=group,
)
if not self.copy_out
else self.all2all_q
for _ in range(2)
)
self.group = group
self._all2all_timeout_seconds = _DEFAULT_ALL2ALL_TIMEOUT_SECONDS
def set_seqlen_all2all(self, seqlens: list[int]) -> None:
# Route through the wrappers so the registered custom ops' fake-impl
# shape info gets updated alongside the C++ runtime's rank_tokens.
self.all2all_q.set_rank_tokens(seqlens)
self.all2all_k.set_rank_tokens(seqlens)
self.all2all_v.set_rank_tokens(seqlens)
self.all2all_heads.set_rank_tokens(seqlens)
@property
def all2all_timeout_seconds(self) -> float:
"""The all2all barrier (deadlock-detection) timeout, in seconds, applied to every instance."""
return self._all2all_timeout_seconds
@all2all_timeout_seconds.setter
def all2all_timeout_seconds(self, seconds: float) -> None:
# Raise it for the first ``torch.compile`` forward -- where one rank's recompile can delay its
# all2all kernel launch past the steady-state timeout, tripping the barrier -- then reset to
# the prior value. ``all2all_k``/``all2all_v`` may alias ``all2all_q`` (copy-out path);
# setting twice is idempotent. Fan out first (it validates) so a rejected value leaves the
# stored steady-state value untouched.
for a2a in (self.all2all_q, self.all2all_k, self.all2all_v, self.all2all_heads):
a2a.set_timeout_seconds(seconds)
self._all2all_timeout_seconds = seconds
def send_recv_qkv(
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
t_q = self.all2all_q.send_recv_heads(q, copy_out=self.copy_out)
t_k = self.all2all_k.send_recv_heads(k, copy_out=self.copy_out)
t_v = self.all2all_v.send_recv_heads(v, copy_out=self.copy_out)
return t_q, t_k, t_v
def gather_heads(self, heads_local: torch.Tensor) -> torch.Tensor:
out = self.all2all_heads.gather_heads(heads_local, copy_out=self.copy_out)
return out
class _All2AllRedistribute:
"""Shared redistribute/gather pipeline for self-attention SP wrappers.
Folds the head dim view-and-shuffle so the masked and unmasked variants only
have to choose how to invoke the inner attention (with or without the mask
kwarg) -- the rest of the SP plumbing is identical.
"""
def __init__(self, manager: AttentionManager) -> None:
self.manager = manager
def redistribute(
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
if heads % self.manager.world_size != 0:
raise ValueError(f"heads ({heads}) must be divisible by world_size ({self.manager.world_size})")
head_dim = q.shape[-1] // heads
q = q.view(q.shape[0], q.shape[1], heads, head_dim)
k = k.view(k.shape[0], k.shape[1], heads, head_dim)
v = v.view(v.shape[0], v.shape[1], heads, head_dim)
t_q, t_k, t_v = self.manager.send_recv_qkv(q, k, v)
local_heads = heads // self.manager.world_size
# `flatten` / `unflatten` collapse only the head dims, avoiding a `-1` in the
# seq position -- that would otherwise be ambiguous if the seq is 0 for a
# zero-token modality.
t_q = t_q.flatten(-2)
t_k = t_k.flatten(-2)
t_v = t_v.flatten(-2)
return t_q, t_k, t_v, local_heads, head_dim
def gather(self, hidden_states: torch.Tensor, local_heads: int, head_dim: int) -> torch.Tensor:
hidden_states = hidden_states.unflatten(-1, (local_heads, head_dim))
hidden_states = self.manager.gather_heads(hidden_states)
return hidden_states.flatten(-2)
class All2AllAttention(AttentionCallable):
def __init__(self, manager: AttentionManager, original_attention: AttentionCallable):
self._sp = _All2AllRedistribute(manager)
self.original_attention = original_attention
def __call__(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
heads: int,
) -> torch.Tensor:
t_q, t_k, t_v, local_heads, head_dim = self._sp.redistribute(q, k, v, heads)
hidden_states = self.original_attention(q=t_q, k=t_k, v=t_v, heads=local_heads)
return self._sp.gather(hidden_states, local_heads, head_dim)
class MaskedAll2AllAttention(MaskedAttentionCallable):
def __init__(self, manager: AttentionManager, original_attention: MaskedAttentionCallable):
self._sp = _All2AllRedistribute(manager)
self.original_attention = original_attention
def __call__(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
heads: int,
mask: torch.Tensor,
) -> torch.Tensor:
t_q, t_k, t_v, local_heads, head_dim = self._sp.redistribute(q, k, v, heads)
hidden_states = self.original_attention(q=t_q, k=t_k, v=t_v, heads=local_heads, mask=mask)
return self._sp.gather(hidden_states, local_heads, head_dim)
class _AudioAll2AllRedistribute:
"""Shared redistribute/gather pipeline for audio cross-attention SP wrappers.
Q is sliced locally per rank (no cross-rank shuffle on Q because the audio
sequence length is small enough to replicate); K/V are redistributed across
ranks via ``send_recv_heads``; outputs are gathered via
``all_gather_into_tensor`` along the head dimension. The masked and unmasked
variants share this plumbing and only differ in how they invoke the inner
attention.
"""
def __init__(self, manager: AttentionManager) -> None:
self.manager = manager
def redistribute(
self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, heads: int
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
if heads % self.manager.world_size != 0:
raise ValueError(f"heads ({heads}) must be divisible by world_size ({self.manager.world_size})")
head_dim = q.shape[-1] // heads
heads_per_rank = heads // self.manager.world_size
rank = self.manager.rank
q = q.view(q.shape[0], q.shape[1], heads, head_dim)
k = k.view(k.shape[0], k.shape[1], heads, head_dim)
v = v.view(v.shape[0], v.shape[1], heads, head_dim)
t_q = q[:, :, heads_per_rank * rank : heads_per_rank * (rank + 1), :].clone()
t_k = self.manager.all2all_k.send_recv_heads(k, copy_out=self.manager.copy_out)
t_v = self.manager.all2all_v.send_recv_heads(v, copy_out=self.manager.copy_out)
# `flatten` / `unflatten` collapse only the head dims, avoiding a `-1` in the
# seq position -- that would otherwise be ambiguous if the seq is 0 for a
# zero-token modality.
t_q = t_q.flatten(-2)
t_k = t_k.flatten(-2)
t_v = t_v.flatten(-2)
return t_q, t_k, t_v, heads_per_rank, head_dim
def gather(self, hidden_states: torch.Tensor, heads_per_rank: int, head_dim: int) -> torch.Tensor:
# (B, S, heads_per_rank, head_dim). Move head dim to dim 0 so all_gather_into_tensor
# gathers along it; permute back after the collective.
hidden_states = hidden_states.unflatten(-1, (heads_per_rank, head_dim)).permute(2, 0, 1, 3).contiguous()
gathered = torch.empty(
(heads_per_rank * self.manager.world_size, *hidden_states.shape[1:]),
dtype=hidden_states.dtype,
device=hidden_states.device,
)
dist.all_gather_into_tensor(gathered, hidden_states, group=self.manager.group)
# (heads, B, S, head_dim) -> (B, S, heads, head_dim) -> (B, S, heads * head_dim)
return gathered.permute(1, 2, 0, 3).flatten(-2)
class AudioAll2AllAttention(AttentionCallable):
"""All2All attention for audio cross-attention (video_to_audio).
Q is sliced locally per rank, K/V are redistributed via send_recv_heads,
then outputs are gathered via all_gather across the head dimension.
"""
def __init__(self, manager: AttentionManager, original_attention: AttentionCallable):
self._sp = _AudioAll2AllRedistribute(manager)
self.original_attention = original_attention
def __call__(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
heads: int,
) -> torch.Tensor:
t_q, t_k, t_v, heads_per_rank, head_dim = self._sp.redistribute(q, k, v, heads)
hidden_states = self.original_attention(q=t_q, k=t_k, v=t_v, heads=heads_per_rank)
return self._sp.gather(hidden_states, heads_per_rank, head_dim)
class MaskedAudioAll2AllAttention(MaskedAttentionCallable):
"""Masked counterpart to :class:`AudioAll2AllAttention`.
No current caller invokes A2V / V2A cross-attention with a mask, so the SP
mutator pre-installs an unmasked-only :class:`AudioAll2AllAttention` and the
masked slot stays at the model default. Defined now so adding masked audio
cross-attention later is just an SP-mutator change, not a missing-piece
discovery.
"""
def __init__(self, manager: AttentionManager, original_attention: MaskedAttentionCallable):
self._sp = _AudioAll2AllRedistribute(manager)
self.original_attention = original_attention
def __call__(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
heads: int,
mask: torch.Tensor,
) -> torch.Tensor:
t_q, t_k, t_v, heads_per_rank, head_dim = self._sp.redistribute(q, k, v, heads)
hidden_states = self.original_attention(q=t_q, k=t_k, v=t_v, heads=heads_per_rank, mask=mask)
return self._sp.gather(hidden_states, heads_per_rank, head_dim)
@@ -0,0 +1,300 @@
"""
Multi-GPU inference wrapper for LTX transformer models.
This module provides utilities for running LTX model inference across multiple GPUs
using sequence parallelism. It:
- Tiles the video inputs across GPUs in the sequence (token) dimension
- Patches video self-attention operations with all2all attention
- Runs the model forward pass on each GPU with its local tile
- Gathers all tokens back to all GPUs after the forward pass
"""
from dataclasses import replace
from itertools import accumulate
import torch
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
from ltx_core.loader.module_ops import ModuleOps
from ltx_core.model.transformer.attention import Attention
from ltx_core.model.transformer.modality import Modality
from ltx_core.model.transformer.model import LTXModel
from ltx_core.model.transformer.transformer import BasicAVTransformerBlock
from ltx_core.multigpu.transformer.attention import (
All2AllAttention,
AttentionManager,
AudioAll2AllAttention,
MaskedAll2AllAttention,
MaskedAudioAll2AllAttention,
)
def compute_sequence_partition(
total_tokens: int,
world_size: int,
) -> list[int]:
"""
Compute uniform per-rank token counts.
Requires ``total_tokens % world_size == 0`` callers must pad up-front via
:func:`pad_modality_for_uniform_sharding`. Uniform sharding lets the
All2All custom-op fakes derive output shapes symbolically from input shapes
(``x.shape[1] * world_size`` / ``x.shape[1] // world_size``) instead of from
Python int args.
"""
if total_tokens % world_size != 0:
raise ValueError(
f"compute_sequence_partition expects uniform sharding: total_tokens "
f"({total_tokens}) must be divisible by world_size ({world_size}). "
f"Pad the modality up-front."
)
per_rank = total_tokens // world_size
return [per_rank] * world_size
def pad_modality_for_uniform_sharding(
modality: Modality,
world_size: int,
) -> tuple[Modality, int]:
"""Pad the seq dim up to the next multiple of ``world_size`` and attach a
padding-aware attention bias so the padded keys are ignored.
Returns ``(padded_modality, original_seq_len)``. If no padding is needed
the original modality is returned unchanged.
"""
t_orig = modality.latent.shape[1]
pad = (-t_orig) % world_size
if pad == 0:
return modality, t_orig
t_padded = t_orig + pad
b = modality.latent.shape[0]
device = modality.latent.device
dtype = modality.latent.dtype
latent_pad = torch.zeros(b, pad, modality.latent.shape[2], dtype=dtype, device=device)
latent = torch.cat([modality.latent, latent_pad], dim=1)
timesteps_pad_shape = list(modality.timesteps.shape)
timesteps_pad_shape[1] = pad
timesteps_pad = torch.zeros(timesteps_pad_shape, dtype=modality.timesteps.dtype, device=modality.timesteps.device)
timesteps = torch.cat([modality.timesteps, timesteps_pad], dim=1)
positions_pad_shape = list(modality.positions.shape)
positions_pad_shape[2] = pad
positions_pad = torch.zeros(positions_pad_shape, dtype=modality.positions.dtype, device=modality.positions.device)
positions = torch.cat([modality.positions, positions_pad], dim=2)
if modality.attention_mask is None:
# Key-only padding mask in the canonical [0, 1] form: 1 on valid keys,
# 0 on padded keys. Shape (1, 1, T_padded) broadcasts across batch and
# queries -- O(T) memory instead of materialising a dense (B, T, T)
# matrix just to mask `pad` (< world_size) keys.
# `_prepare_self_attention_mask` does the standard 3D -> 4D log-space
# conversion and produces a (1, 1, 1, T_padded) bias.
attention_mask = torch.ones(1, 1, t_padded, dtype=torch.float32, device=device)
attention_mask[:, :, t_orig:] = 0.0
else:
# User-supplied (B, T, T) [0, 1] mask: extend with padded rows/cols.
# Padded query rows attend to all valid keys so their softmax stays
# well-defined (the outputs are sliced off after the gather, but a
# fully-masked row would produce NaN).
old = modality.attention_mask
attention_mask = torch.zeros(b, t_padded, t_padded, dtype=old.dtype, device=old.device)
attention_mask[:, :t_orig, :t_orig] = old
attention_mask[:, t_orig:, :t_orig] = 1.0
padded = replace(
modality,
latent=latent,
timesteps=timesteps,
positions=positions,
attention_mask=attention_mask,
)
return padded, t_orig
def compute_sequence_offsets(token_counts: list[int]) -> list[int]:
"""
Compute the starting offset for each rank's token partition.
Args:
token_counts: List of token counts per rank.
Returns:
List of starting offsets for each rank.
"""
return [0, *accumulate(token_counts[:-1])]
def tile_modality_for_rank(
modality: Modality,
rank: int,
world_size: int,
) -> tuple[Modality, list[int]]:
"""
Tile a modality's tensors for a specific GPU rank.
Splits the sequence dimension (dim 1 for latent/timesteps, dim 2 for positions)
across GPUs, returning the local tile for the given rank.
Args:
modality: The modality to tile.
rank: Current GPU rank.
world_size: Total number of GPUs.
Returns:
Tuple of (tiled_modality, token_counts_per_rank).
"""
total_tokens = modality.latent.shape[1]
token_counts = compute_sequence_partition(total_tokens, world_size)
offsets = compute_sequence_offsets(token_counts)
start = offsets[rank]
end = start + token_counts[rank]
# Tile latent: (B, T, D) -> (B, T_local, D)
tiled_latent = modality.latent[:, start:end, :]
# Tile timesteps: (B, T) -> (B, T_local)
tiled_timesteps = modality.timesteps[:, start:end]
# Tile positions: (B, 3, T, 2) -> (B, 3, T_local, 2)
tiled_positions = modality.positions[:, :, start:end, :]
tiled_modality = replace(
modality,
latent=tiled_latent,
timesteps=tiled_timesteps,
positions=tiled_positions,
)
return tiled_modality, token_counts
def gather_output_tokens(
local_output: torch.Tensor,
token_counts: list[int],
group: torch.distributed.ProcessGroup | None = None,
) -> torch.Tensor:
"""
Gather output tokens from all GPUs back into a single tensor.
Args:
local_output: Local output tensor of shape (B, T_local, D).
token_counts: Number of tokens on each rank.
group: Process group for communication. If None, uses default group.
Returns:
Gathered tensor of shape (B, T_total, D) on all ranks.
"""
world_size = len(token_counts)
batch_size = local_output.shape[0]
hidden_dim = local_output.shape[2]
# Prepare output tensors for all_gather
max_tokens = max(token_counts)
# Pad local output to max size for uniform all_gather
padded_local = torch.zeros(
batch_size,
max_tokens,
hidden_dim,
dtype=local_output.dtype,
device=local_output.device,
)
padded_local[:, : local_output.shape[1], :] = local_output
# All gather padded outputs
gathered_list = [torch.zeros_like(padded_local) for _ in range(world_size)]
torch.distributed.all_gather(gathered_list, padded_local, group=group)
# Extract actual tokens (remove padding) and concatenate
outputs = []
for i, count in enumerate(token_counts):
outputs.append(gathered_list[i][:, :count, :])
return torch.cat(outputs, dim=1)
def create_video_self_attention_module_ops(
attention_manager: AttentionManager,
) -> ModuleOps:
"""
Create ModuleOps for patching video self-attention with all2all attention.
This patches the `attn1` attribute on BasicAVTransformerBlock instances,
which is the video self-attention module.
Args:
attention_manager: The AttentionManager instance for all2all communication.
Returns:
ModuleOps that can be used to patch the model.
"""
def mutator(module: torch.nn.Module) -> torch.nn.Module:
for block in module.transformer_blocks:
if not isinstance(block, BasicAVTransformerBlock):
continue
# Video self-attention: ``Attention.forward`` may receive a non-None
# ``mask`` (``video.self_attention_mask``), so wrap both slots; the
# branch in ``Attention.forward`` then routes to whichever wrapper
# corresponds to the actual call.
if hasattr(block, "attn1"):
attn1 = block.attn1
if isinstance(attn1, Attention):
attn1.attention_function = All2AllAttention(attention_manager, attn1.attention_function)
attn1.masked_attention_function = MaskedAll2AllAttention(
attention_manager, attn1.masked_attention_function
)
# video_to_audio cross-attention: no current caller passes a mask
# (see ``BasicAVTransformerBlock.forward``), so the masked branch
# is dead code today. Wrap both slots anyway so that if a future
# caller adds a mask, the SP plumbing is already in place rather
# than silently bypassing All2All on that path.
if hasattr(block, "video_to_audio_attn"):
video_to_audio_attn = block.video_to_audio_attn
if isinstance(video_to_audio_attn, Attention):
video_to_audio_attn.attention_function = AudioAll2AllAttention(
attention_manager, video_to_audio_attn.attention_function
)
video_to_audio_attn.masked_attention_function = MaskedAudioAll2AllAttention(
attention_manager, video_to_audio_attn.masked_attention_function
)
return module
return ModuleOps(
name="video_self_attention_all2all",
matcher=lambda module: isinstance(module, LTXModel),
mutator=mutator,
)
class SequenceParallelModelWrapper(torch.nn.Module):
def __init__(self, model: torch.nn.Module, attention_manager: AttentionManager):
super().__init__()
self.model = model
self.attention_manager = attention_manager
@property
def num_blocks(self) -> int:
return self.model.num_blocks
def forward(
self, video: Modality | None, audio: Modality | None, perturbations: BatchedPerturbationConfig | None
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
if video is None:
return self.model(video, audio, perturbations)
# Pad the video seq dim up to a multiple of world_size so all ranks get
# equal shards. The attention mask we attach makes the padded keys
# invisible to attention; padded rows are sliced off after the gather.
video, t_orig = pad_modality_for_uniform_sharding(video, self.attention_manager.world_size)
video_tile, token_counts = tile_modality_for_rank(
video, self.attention_manager.rank, self.attention_manager.world_size
)
total_tokens = sum(token_counts)
if total_tokens > self.attention_manager.max_tokens:
raise ValueError(
f"Total video token count ({total_tokens}) exceeds attention_manager max_tokens "
f"({self.attention_manager.max_tokens}). Use a smaller resolution or fewer frames."
)
self.attention_manager.set_seqlen_all2all(token_counts)
torch.distributed.barrier(self.attention_manager.group)
video, audio = self.model(video_tile, audio, perturbations)
video = gather_output_tokens(video, token_counts, self.attention_manager.group)
# Unpad: drop the rows we added in `pad_modality_for_uniform_sharding` to make
# the seq dim divisible by world_size, restoring the caller's original length.
if video.shape[1] != t_orig:
video = video[:, :t_orig, :]
return video, audio
@@ -0,0 +1,99 @@
"""Tiled Data Parallel model wrapper for the LTX transformer.
Each GPU processes one or more tiles of the patchified
``(frames, height, width)`` latent. Tiles are assigned to ranks via
round-robin, so the number of tiles may exceed the number of GPUs.
Tiles may overlap; overlapping regions are blended with trapezoidal
masks so that seam artefacts are suppressed. Each rank accumulates
its assigned tiles locally, then a single ``all_reduce`` synchronises
the blended output across all ranks.
Conditioning tokens (appended after the generated tokens) are filtered
per tile: only tokens whose positions overlap with the tile's spatial
extent (or that have negative time coordinates) are included.
"""
from __future__ import annotations
from typing import TYPE_CHECKING
import torch
import torch.distributed as dist
from ltx_core.modality_tiling import VideoModalityTilingHelper
from ltx_core.model.transformer.modality import Modality
from ltx_core.tiling import TileCountConfig
from ltx_core.tools import VideoLatentTools
if TYPE_CHECKING:
from ltx_core.guidance.perturbations import BatchedPerturbationConfig
class TiledDataParallelModelWrapper(torch.nn.Module):
"""Wraps an ``X0Model`` for tiled data parallelism.
Tiles are distributed across ranks via round-robin, allowing more
tiles than GPUs (e.g. 16 tiles on 4 GPUs = 4 tiles per rank).
Each rank processes its assigned tiles sequentially, blending each
into a full-size accumulator. A single ``all_reduce(SUM)`` after
all local tiles produces the final result (blend masks sum to 1
globally across all tiles).
Audio is processed untiled on every tile forward; the outputs are
summed via ``all_reduce`` and divided by the total tile count so
that all ranks stay in sync.
"""
def __init__(
self,
model: torch.nn.Module,
*,
video_tools: VideoLatentTools,
tiling: TileCountConfig,
group: dist.ProcessGroup,
normalize_positions: bool = True,
) -> None:
super().__init__()
self.model = model
self.group = group
self.world_size = dist.get_world_size(group)
self._normalize_positions = normalize_positions
self._helper = VideoModalityTilingHelper(tiling, video_tools)
all_tiles = self._helper.tiles
rank = dist.get_rank(group)
self._tiles = [t for i, t in enumerate(all_tiles) if i % self.world_size == rank]
@property
def num_blocks(self) -> int:
return self.model.num_blocks
def forward(
self,
video: Modality | None,
audio: Modality | None,
perturbations: BatchedPerturbationConfig | None,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
if video is None:
return self.model(video, audio, perturbations)
# Each rank processes its assigned tiles and accumulates locally.
denoised_video: torch.Tensor | None = None
denoised_audio: torch.Tensor | None = None
for tile in self._tiles:
tiled_video, ctx = self._helper.tile_modality(video, tile, normalize_positions=self._normalize_positions)
tile_out, audio_out = self.model(tiled_video, audio, perturbations)
blended = self._helper.blend(tile_out, tile, ctx)
denoised_video = blended if denoised_video is None else denoised_video + blended
if audio_out is not None:
denoised_audio = audio_out if denoised_audio is None else denoised_audio + audio_out
assert denoised_video is not None
# All-reduce: sum blended tiles across ranks (masks sum to 1 globally).
denoised_video = denoised_video.contiguous()
dist.all_reduce(denoised_video, op=dist.ReduceOp.SUM, group=self.group)
# Average audio across all tile forwards (each saw different video context).
if denoised_audio is not None:
total_tiles = len(self._helper.tiles)
denoised_audio = denoised_audio.contiguous()
dist.all_reduce(denoised_audio, op=dist.ReduceOp.SUM, group=self.group)
denoised_audio = denoised_audio / total_tiles
return denoised_video, denoised_audio
@@ -0,0 +1,5 @@
"""Multi-GPU utilities for VAE decoding."""
from ltx_core.multigpu.vae.distributed_decoder import DistributedVideoDecoder
__all__ = ["DistributedVideoDecoder"]
@@ -0,0 +1,307 @@
"""Distributed video decoder that partitions the latent across ranks.
Tiles are assigned to ranks via round-robin, so the number of tiles
may exceed the number of GPUs (e.g. 16 tiles on 4 GPUs = 4 tiles per
rank). Each rank decodes its assigned tiles sequentially. Workers
put their list of decoded tiles into a ``mp.Queue`` (CUDA IPC
zero-copy handle sharing). The driver collects all tiles, blends
overlap zones, and returns temporal batches distributed across devices.
The tiling configuration comes from ``MGPUConfig.vae_tiling`` (set at
construction time), NOT from the pipeline's SGPU tiling kwarg. MGPU
tiling controls parallelism; SGPU tiling controls single-GPU VRAM
management they are independent concerns.
"""
from __future__ import annotations
import logging
from collections.abc import Callable, Iterator
from dataclasses import dataclass
from typing import TYPE_CHECKING
import torch
import torch.distributed as dist
from einops import rearrange
from torch.multiprocessing import Queue
from ltx_core.model.video_vae.tiling import TilingConfig
from ltx_core.model.video_vae.video_vae import (
VideoDecoder,
map_spatial_slice,
map_temporal_slice,
to_mapping_operation,
)
from ltx_core.tiling import (
Tile,
create_tiles,
split_by_count,
split_by_count_temporal_causal,
)
from ltx_core.types import SpatioTemporalScaleFactors, VideoLatentShape
if TYPE_CHECKING:
from ltx_core.tiling import TileCountConfig
logger = logging.getLogger(__name__)
# ------------------------------------------------------------------
# Data structures
# ------------------------------------------------------------------
@dataclass(frozen=True)
class DecodedTile:
"""A VAE-decoded tile with pixel-space placement.
Attributes:
pixels: ``[F_tile, H_tile, W_tile, C]`` in the decoder's native dtype.
pixel_tile: Carries ``out_coords`` (f, h, w slices) and ``blend_mask``.
"""
pixels: torch.Tensor
pixel_tile: Tile
# ------------------------------------------------------------------
# Tile construction helpers
# ------------------------------------------------------------------
def _to_decoded_tile(
raw: torch.Tensor,
tile: Tile,
) -> DecodedTile:
"""Convert raw decoder output ``[B, C, F, H, W]`` to a :class:`DecodedTile`.
Rearranges to ``[F, H, W, C]`` and normalises ``[-1, 1] [0, 1]``.
"""
pixels = rearrange(raw[0], "c f h w -> f h w c")
pixels = ((pixels + 1.0) / 2.0).clamp(0.0, 1.0)
return DecodedTile(pixels=pixels, pixel_tile=tile)
# ------------------------------------------------------------------
# Tile assembly
# ------------------------------------------------------------------
def compute_summed_weights(
tiles: list[DecodedTile],
total_frames: int,
output_height: int,
output_width: int,
) -> torch.Tensor:
"""Build the ``[F, H, W]`` denominator for weighted blending."""
weights = torch.zeros(total_frames, output_height, output_width)
for tile in tiles:
f_slice, h_slice, w_slice = tile.pixel_tile.out_coords
weights[f_slice, h_slice, w_slice] += tile.pixel_tile.blend_mask
return weights.clamp(min=1e-8)
def gather_frames(
tiles: list[DecodedTile],
total_frames: int,
output_height: int,
output_width: int,
num_temporal_batches: int,
world_size: int,
weights: torch.Tensor,
device_fn: Callable[[int], str | torch.device] | None = None,
) -> Iterator[torch.Tensor]:
"""Assemble decoded tiles into temporal batches distributed across GPUs.
Each temporal batch is allocated on the device returned by *device_fn(batch_index)*.
By default batches are placed round-robin on ``cuda:0`` ``cuda:<world_size-1>``.
"""
if device_fn is None:
device_fn = lambda b: f"cuda:{b % world_size}" # noqa: E731
batch_size = (total_frames + num_temporal_batches - 1) // num_temporal_batches
for b in range(num_temporal_batches):
batch_range = slice(b * batch_size, min((b + 1) * batch_size, total_frames))
batch_len = batch_range.stop - batch_range.start
if batch_len <= 0:
break
device = device_fn(b)
dtype = tiles[0].pixels.dtype
output = torch.zeros(batch_len, output_height, output_width, 3, device=device, dtype=dtype)
for tile in tiles:
f_slice, h_slice, w_slice = tile.pixel_tile.out_coords
overlap = slice(max(batch_range.start, f_slice.start), min(batch_range.stop, f_slice.stop))
if overlap.start >= overlap.stop:
continue
tile_frames = slice(overlap.start - f_slice.start, overlap.stop - f_slice.start)
out_frames = slice(overlap.start - batch_range.start, overlap.stop - batch_range.start)
blend = tile.pixel_tile.blend_mask[tile_frames].to(device=device)
output[out_frames, h_slice, w_slice, :] += tile.pixels[tile_frames].to(device=device) * blend[:, :, :, None]
batch_weights = weights[batch_range.start : batch_range.stop].to(device=device)
output.div_(batch_weights[:, :, :, None])
yield output
# ------------------------------------------------------------------
# Main class
# ------------------------------------------------------------------
class DistributedVideoDecoder(torch.nn.Module):
"""Distributed VAE decoder with queue-based tile collection.
All ranks decode their latent tile in parallel. Workers send
their :class:`DecodedTile` to the driver rank via the shared
``mp.Queue`` (CUDA IPC zero-copy). The driver collects all
tiles, blends overlapping regions, and returns temporal batches
as an iterator.
Parameters
----------
decoder:
The real (local) ``VideoDecoder`` instance.
queue:
``mp.Queue`` shared across all ranks for CUDA IPC tile transfer.
vae_group:
NCCL process group for the VAE ranks. Used to derive
``rank`` and ``world_size`` within the group.
vae_tiling:
MGPU tiling config that determines how the latent is split.
driver_rank:
Group-local rank of the driver process (the rank that collects
and assembles tiles).
"""
def __init__(
self,
decoder: VideoDecoder,
queue: Queue, # type: ignore[type-arg]
vae_group: dist.ProcessGroup,
vae_tiling: TileCountConfig,
driver_rank: int = 0,
) -> None:
super().__init__()
self.decoder = decoder
self.queue = queue
self.vae_group = vae_group
self.rank = dist.get_rank(vae_group)
self.world_size = dist.get_world_size(vae_group)
self.vae_tiling = vae_tiling
self.driver_rank = driver_rank
def forward(
self,
sample: torch.Tensor,
timestep: torch.Tensor | None = None,
generator: torch.Generator | None = None,
) -> torch.Tensor:
"""Non-tiled path: fall back to local decode."""
return self.decoder(sample, timestep, generator)
def decode_video(
self,
latent: torch.Tensor,
tiling_config: TilingConfig | None = None,
generator: torch.Generator | None = None,
device_fn: Callable[[int], str | torch.device] | None = None,
) -> Iterator[torch.Tensor]:
"""Distributed decode — all ranks decode, driver assembles.
Not a generator so that worker side-effects (decode + queue.put)
execute eagerly regardless of whether the caller iterates.
1. Each rank decodes its latent tile (with optional intra-GPU tiling).
2. Workers send their :class:`DecodedTile` to the driver via the queue.
3. The driver collects all tiles, blends overlaps, and returns
temporal batches distributed across GPUs.
"""
if (
self.vae_tiling.frames.num_tiles > 1
and tiling_config is not None
and tiling_config.temporal_config is not None
):
raise ValueError(
"Cannot combine multi-GPU temporal tiling (vae_tiling.frames.num_tiles > 1) "
"with single-GPU temporal tiling (tiling_config.temporal_config). "
"Use only one to avoid causal decoding conflicts."
)
latent_shape = VideoLatentShape.from_torch_shape(latent.shape)
scale = self.decoder.video_downscale_factors
full_shape = latent_shape.upscale(scale)
# Phase 1: each rank decodes its assigned tiles.
my_tiles = self._decode_tiles(latent, latent_shape, scale, generator, tiling_config)
# Phase 2: workers send tiles to driver.
if self.rank != self.driver_rank:
self.queue.put((self.rank, my_tiles))
return iter([])
# Phase 3: driver collects and assembles.
all_tiles = self._collect_tiles(my_tiles)
weights = compute_summed_weights(all_tiles, full_shape.frames, full_shape.height, full_shape.width)
batches = gather_frames(
all_tiles,
full_shape.frames,
full_shape.height,
full_shape.width,
self.world_size,
self.world_size,
weights,
device_fn=device_fn,
)
return batches
# ------------------------------------------------------------------
# Private helpers
# ------------------------------------------------------------------
def _decode_tiles(
self,
latent: torch.Tensor,
latent_shape: VideoLatentShape,
scale: SpatioTemporalScaleFactors,
generator: torch.Generator | None,
tiling_config: TilingConfig | None = None,
) -> list[DecodedTile]:
"""Decode this rank's assigned latent tiles and convert to :class:`DecodedTile` list."""
all_tiles = create_tiles(
torch.Size([latent_shape.frames, latent_shape.height, latent_shape.width]),
splitters=[
split_by_count_temporal_causal(self.vae_tiling.frames.num_tiles, self.vae_tiling.frames.overlap),
split_by_count(self.vae_tiling.height.num_tiles, self.vae_tiling.height.overlap),
split_by_count(self.vae_tiling.width.num_tiles, self.vae_tiling.width.overlap),
],
mappers=[
to_mapping_operation(map_temporal_slice, scale.time),
to_mapping_operation(map_spatial_slice, scale.height),
to_mapping_operation(map_spatial_slice, scale.width),
],
)
my_tiles = [t for i, t in enumerate(all_tiles) if i % self.world_size == self.rank]
decoded = []
for tile in my_tiles:
latent_slice = latent[:, :, tile.in_coords[0], tile.in_coords[1], tile.in_coords[2]]
if tiling_config is not None:
chunks = list(self.decoder.tiled_decode(latent_slice, tiling_config, generator=generator))
raw = torch.cat(chunks, dim=2)
else:
raw = self.decoder.forward(latent_slice, generator=generator)
decoded.append(_to_decoded_tile(raw, tile))
return decoded
def _collect_tiles(self, driver_tiles: list[DecodedTile]) -> list[DecodedTile]:
"""Collect tiles from all workers via the queue. Returns flat list of all tiles.
Sorted by rank so the downstream reduction in ``gather_frames`` /
``compute_summed_weights`` (in-place ``+=`` over overlapping pixel
regions) processes tiles in a fixed order. Queue-arrival order would
otherwise vary run-to-run and yield 1-ulp bf16 drift from
non-associative floating-point summation.
"""
per_rank: dict[int, list[DecodedTile]] = {self.driver_rank: driver_tiles}
for _ in range(self.world_size - 1):
worker_rank, worker_tiles = self.queue.get()
per_rank[worker_rank] = worker_tiles
result: list[DecodedTile] = []
for rank in sorted(per_rank):
result.extend(per_rank[rank])
return result
@@ -0,0 +1,47 @@
"""Public API for blockwise FP8/FP6 quantization.
The implementation lives in :mod:`._impl`, which imports the compiled
``ltx_kernels.blockwise`` kernels at top level. This module deliberately defers
that import so that ``ltx_core.quantization.blockwise`` remains importable
without those kernels built; the gate fires only when one of the policy builders
is actually called.
"""
from ltx_core.quantization.policy import QuantizationPolicy
__all__ = ["build_fp6_policy", "build_fp8_policy"]
def _import_impl(): # noqa: ANN202 - internal helper
try:
from ltx_core.quantization.blockwise import _impl # noqa: PLC0415
return _impl
except ImportError as e:
raise RuntimeError(
"ltx-kernels not built; blockwise FP8/FP6 quantization requires it. "
"Build it on a CUDA host with `uv sync --group kernels` (or "
"`uv pip install -e packages/ltx-kernels --no-build-isolation`) before "
"calling build_fp8_policy() / build_fp6_policy()."
) from e
def build_fp8_policy() -> QuantizationPolicy:
"""Build a blockwise FP8 quantization policy. Raises ``RuntimeError`` if ``ltx-kernels`` is not built."""
impl = _import_impl()
return QuantizationPolicy(
sd_ops=impl.build_sd_ops_fp8(),
module_ops=(impl.build_module_ops_fp8(),),
model_configurator=impl.BlockwiseFP8LTXModelConfigurator,
fuse_rule=impl.fuse_rule_fp8,
)
def build_fp6_policy() -> QuantizationPolicy:
"""Build a blockwise FP6 quantization policy. Raises ``RuntimeError`` if ``ltx-kernels`` is not built."""
impl = _import_impl()
return QuantizationPolicy(
sd_ops=impl.build_sd_ops_fp6(),
module_ops=(impl.build_module_ops_fp6(),),
model_configurator=impl.BlockwiseFP6LTXModelConfigurator,
fuse_rule=impl.fuse_rule_fp6,
)
@@ -0,0 +1,431 @@
"""Implementation of blockwise FP8/FP6 quantization. Depends on ``ltx_kernels``.
This module imports the compiled ``ltx_kernels.blockwise`` kernels at top level
without them built, simply importing this file raises :class:`ImportError`.
The intended access path is through ``ltx_core.quantization.blockwise.__init__``
which catches that and re-raises as a clean :class:`RuntimeError`. Do not import
this module directly from non-quantization code.
"""
from typing import Callable, ClassVar, List, NamedTuple, Protocol, Type
import torch
from ltx_kernels.blockwise.functional import (
blockwise_dequantize,
blockwise_quantize_adanorm_triton,
blockwise_quantize_rms_fma_triton,
fp6_blockwise_quantize_weights_torch,
fp6_pack_tensor,
fp6_unpack_tensor,
fp8_blockwise_quantize_weights_torch,
gated_attention_triton,
rms_norm_rope,
rms_norm_split_rope,
)
from ltx_kernels.blockwise.linear import BlockwiseFP6Linear, BlockwiseFP8Linear
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.loader.sd_ops import KeyValueOperationResult, SDOps
from ltx_core.model.model_protocol import ModelConfigurator
from ltx_core.model.transformer import LTXModel
from ltx_core.model.transformer.model_configurator import LTXModelConfigurator, LTXVideoOnlyModelConfigurator
from ltx_core.model.transformer.ops import (
AdaZeroCallable,
GatedAttentionCallable,
PostSACallable,
PreAttentionCallable,
)
from ltx_core.model.transformer.rope import LTXRopeType
from ltx_core.model.transformer.transformer import TransformerOpsConfig
class FromLinearProtocol(Protocol):
"""Protocol for nn.Module subclasses that can be constructed from an nn.Linear."""
@classmethod
def from_linear(cls, linear: nn.Linear, transform_weights: bool = True) -> nn.Module: ...
class BlockwiseQuantizedWeight(NamedTuple):
"""Result of blockwise quantization: a quantized weight tensor and its per-block scale.
For FP8: ``weight`` is ``float8_e4m3fn``, ``scale`` is ``float32`` shaped
``[out // 128, in // 128]``.
For FP6: ``weight`` is packed ``uint8`` shaped ``[out, (in // 4) * 3]``,
``scale`` is ``float32`` shaped ``[out // 128, in // 128]``.
"""
weight: torch.Tensor
scale: torch.Tensor
EXCLUDED_LAYER_SUBSTRINGS = (
"patchify_proj",
"adaln_single",
"av_ca_video_scale_shift_adaln_single",
"av_ca_a2v_gate_adaln_single",
"caption_projection",
"proj_out",
"audio_patchify_proj",
"audio_adaln_single",
"av_ca_audio_scale_shift_adaln_single",
"av_ca_v2a_gate_adaln_single",
"audio_caption_projection",
"audio_proj_out",
"to_gate_logits",
"scale_shift_table",
)
_QUANTIZABLE_FLOAT_DTYPES = (torch.bfloat16, torch.float16, torch.float32)
def _is_quantizable_float(x: torch.Tensor | torch.dtype) -> bool:
"""Whether ``x`` is an unquantized high-precision float (bf16 / fp16 / fp32).
FP8 / FP6 weights are floats too but they're already in a quantized layout
and must not be re-quantized.
"""
dtype = x.dtype if isinstance(x, torch.Tensor) else x
return dtype in _QUANTIZABLE_FLOAT_DTYPES
def _should_skip_layer(layer_name: str, excluded_layer_substrings: tuple[str, ...]) -> bool:
return any(substring in layer_name for substring in excluded_layer_substrings)
def _replace_linear_modules(model: torch.nn.Module, linear_cls: Type[FromLinearProtocol]) -> torch.nn.Module:
skip_list = ["to_gate_logits", "scale_shift_table"]
for name, module in model.named_modules():
if "transformer_block" in name and isinstance(module, torch.nn.Linear):
if _should_skip_layer(name, skip_list):
continue
*parent_path, child_name = name.split(".")
parent = model
for part in parent_path:
parent = getattr(parent, part)
setattr(
parent,
child_name,
linear_cls.from_linear(module, False),
)
del module.weight
del module.bias
torch.cuda.empty_cache()
return model
# ---------------------------------------------------------------------------
# Weight quantization helpers
# ---------------------------------------------------------------------------
def _blockwise_quantize_weight_helper(
value: torch.Tensor,
quant_fn: Callable[[torch.Tensor, int], tuple[torch.Tensor, torch.Tensor]],
pack_fn: Callable[[torch.Tensor], torch.Tensor],
) -> BlockwiseQuantizedWeight:
orig_device = value.device
w_quant, w_scales = quant_fn(value.cuda())
return BlockwiseQuantizedWeight(
weight=pack_fn(w_quant).to(device=orig_device),
scale=w_scales.to(device=orig_device),
)
def _fp8_blockwise_quantize_weight(value: torch.Tensor) -> BlockwiseQuantizedWeight:
return _blockwise_quantize_weight_helper(value, fp8_blockwise_quantize_weights_torch, lambda x: x)
def _fp6_blockwise_quantize_weight(value: torch.Tensor) -> BlockwiseQuantizedWeight:
return _blockwise_quantize_weight_helper(value, fp6_blockwise_quantize_weights_torch, fp6_pack_tensor)
def _create_weight_quantize_op(
excluded_layer_substrings: tuple[str, ...],
quantization_func: Callable[[torch.Tensor], BlockwiseQuantizedWeight],
) -> Callable[[str, torch.Tensor], list[KeyValueOperationResult]]:
"""KeyValueOperation that blockwise-quantizes a 2D BF16 ``.weight`` and emits ``.weight_scale``."""
def quantize_weight(key: str, value: torch.Tensor) -> list[KeyValueOperationResult]:
if _should_skip_layer(key, excluded_layer_substrings):
return [KeyValueOperationResult(key, value)]
if value.dim() != 2 or not _is_quantizable_float(value):
return [KeyValueOperationResult(key, value)]
quantized = quantization_func(value)
scale_key = key.replace(".weight", ".weight_scale")
return [
KeyValueOperationResult(key, quantized.weight),
KeyValueOperationResult(scale_key, quantized.scale),
]
return quantize_weight
def _create_bias_to_fp32_op(
excluded_layer_substrings: tuple[str, ...],
) -> Callable[[str, torch.Tensor], list[KeyValueOperationResult]]:
"""KeyValueOperation that casts a ``.bias`` tensor to FP32.
``BlockwiseFP{8,6}Linear`` registers ``.bias`` as float32; the load-time
cast keeps the checkpoint's BF16 bias compatible with that param dtype.
"""
def bias_to_fp32(key: str, value: torch.Tensor) -> list[KeyValueOperationResult]:
if _should_skip_layer(key, excluded_layer_substrings):
return [KeyValueOperationResult(key, value)]
return [KeyValueOperationResult(key, value.float())]
return bias_to_fp32
# ---------------------------------------------------------------------------
# Q8 activation callables (formerly in model.transformer.ops)
# ---------------------------------------------------------------------------
class Q8KernelsPreAttention(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]:
if attn_module.rope_type == LTXRopeType.INTERLEAVED:
rope_func = rms_norm_rope
elif attn_module.rope_type == LTXRopeType.SPLIT:
rope_func = rms_norm_split_rope
else:
raise ValueError(f"Invalid rope type: {attn_module.rope_type}")
if pe is not None:
k_pe = k_pe if k_pe is not None else pe
q = rope_func(q, pe[0], pe[1], attn_module.q_norm.weight, False)
k = rope_func(k, k_pe[0], k_pe[1], attn_module.k_norm.weight, False)
else:
q = attn_module.q_norm(q)
k = attn_module.k_norm(k)
return q, k
class Q8KernelsAdaZeroFunction(AdaZeroCallable):
def __call__(
self,
x: torch.Tensor,
eps: float, # noqa: ARG002
scale: torch.Tensor,
shift: torch.Tensor,
) -> torch.Tensor:
return blockwise_quantize_adanorm_triton(x, None, scale, shift, torch.float8_e4m3fn, 1.0)
class Q8KernelsPostSAFunction(PostSACallable):
def __call__(
self,
x: torch.Tensor,
y: torch.Tensor,
norm_weights: torch.Tensor | None, # noqa: ARG002
eps: float, # noqa: ARG002
gate: torch.Tensor,
) -> List[torch.Tensor]:
# Dequantize the fused result: the cross-attention AdaLN path applies a BF16
# scale/shift, which cannot operate on the (fp8, scales) payload.
normed_fp8 = blockwise_quantize_rms_fma_triton(x, y, gate)
return x, blockwise_dequantize(normed_fp8)
class Q8KernelsGatedAttention(GatedAttentionCallable):
def __call__(
self,
x: torch.Tensor,
attn_out: torch.Tensor,
attn_module: nn.Module,
) -> torch.Tensor:
# Self-attention path: ``x`` arrives as the ``(fp8, scales)`` tuple
# produced by Q8KernelsAdaZeroFunction. Cross-attention path
# (apply_cross_attention_adaln) feeds plain BF16, so dequantize only
# when needed.
if isinstance(x, tuple):
x = blockwise_dequantize(x)
gate_logits = attn_module.to_gate_logits(x)
return gated_attention_triton(attn_out, gate_logits)
# ---------------------------------------------------------------------------
# Fuse rules
# ---------------------------------------------------------------------------
_BLOCK = 128
def _blockwise_dequantize_2d(weight_fp8: torch.Tensor, weight_scale: torch.Tensor) -> torch.Tensor:
"""Dequantize a 2D blockwise-FP8 weight ``[out, in]`` with per-block scale
``[out//128, in//128]`` to BF16.
``ltx_kernels.blockwise.blockwise_dequantize`` is built for 3D activations where
scales are ``[b*s, in//128]`` one row per token. Weights are block-
quantized along the row dim too, so we expand the row axis 128x via
``repeat_interleave`` and reuse the kernel.
"""
out_features, in_features = weight_fp8.shape
scales_per_row = weight_scale.repeat_interleave(_BLOCK, dim=0)
return blockwise_dequantize((weight_fp8.unsqueeze(0), scales_per_row)).view(out_features, in_features)
def _blockwise_fp8_fuse(
key: str,
weight: torch.Tensor,
deltas: torch.Tensor,
model_sd: StateDict,
) -> dict[str, torch.Tensor]:
"""Dequantize the FP8 weight + per-block scale to BF16, add the BF16 delta,
and re-quantize blockwise. Both ``.weight`` and the companion
``.weight_scale`` are emitted so the loaded layer matches what
``BlockwiseFP8Linear`` expects.
Excluded layers (see ``EXCLUDED_LAYER_SUBSTRINGS``) 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]
bf16_weight = _blockwise_dequantize_2d(weight, weight_scale)
merged = bf16_weight + deltas.to(dtype=bf16_weight.dtype)
new_fp8_weight, new_weight_scale = fp8_blockwise_quantize_weights_torch(merged.cuda())
return {
key: new_fp8_weight.to(device=weight.device),
scale_key: new_weight_scale.to(device=weight.device),
}
def _blockwise_fp6_fuse(
key: str,
weight: torch.Tensor,
deltas: torch.Tensor,
model_sd: StateDict,
) -> dict[str, torch.Tensor]:
"""Mirror ``BlockwiseFP6Linear.fp8weight`` for the dequant side: unpack the
packed ``uint8`` weight to ``float8_e4m3fn``, dequantize via the per-block
scale to BF16, add the BF16 delta, re-quantize to FP6, and pack back to
uint8. Both ``.weight`` (packed uint8) and ``.weight_scale`` are emitted.
Note: ``fp6_unpack_tensor`` restores the dropped e_1/e_2 exponent bits as 0,
so the dequant->add->requant round-trip is lossy on those bits even when no
LoRA delta is applied. This matches what ``BlockwiseFP6Linear`` already does
at inference time via its ``fp8weight`` property, so the fused weight is
numerically consistent with the unfused inference path.
Excluded layers (see ``EXCLUDED_LAYER_SUBSTRINGS``) 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]
# Packed shape is [out, (in // 4) * 3]; recover in_features.
original_n = weight.shape[-1] * 4 // 3
fp8_view = fp6_unpack_tensor(weight, original_n).view(torch.float8_e4m3fn)
bf16_weight = _blockwise_dequantize_2d(fp8_view, weight_scale)
merged = bf16_weight + deltas.to(dtype=bf16_weight.dtype)
new_fp8, new_scale = fp6_blockwise_quantize_weights_torch(merged)
new_packed = fp6_pack_tensor(new_fp8.view(torch.uint8))
return {
key: new_packed.to(device=weight.device),
scale_key: new_scale.to(device=weight.device),
}
# ---------------------------------------------------------------------------
# Configurators (TransformerOpsConfig with Q8 activation callables)
# ---------------------------------------------------------------------------
def _build_blockwise_ops_config() -> TransformerOpsConfig:
return TransformerOpsConfig.from_functions(
preattention=Q8KernelsPreAttention(),
gated_attention=Q8KernelsGatedAttention(),
ada_zero=Q8KernelsAdaZeroFunction(),
post_sa=Q8KernelsPostSAFunction(),
)
# FP6 is weight-only; activation ops match FP8.
_BLOCKWISE_OPS = _build_blockwise_ops_config()
class BlockwiseFP8LTXModelConfigurator(ModelConfigurator[LTXModel]):
BASE: ClassVar[type[ModelConfigurator[LTXModel]]] = LTXModelConfigurator
OPS: ClassVar[TransformerOpsConfig] = _BLOCKWISE_OPS
@classmethod
def from_config(cls, config: dict) -> LTXModel:
return cls.BASE.from_config(config, ops=cls.OPS)
class BlockwiseFP8LTXVideoOnlyModelConfigurator(BlockwiseFP8LTXModelConfigurator):
BASE = LTXVideoOnlyModelConfigurator
class BlockwiseFP6LTXModelConfigurator(BlockwiseFP8LTXModelConfigurator):
pass
class BlockwiseFP6LTXVideoOnlyModelConfigurator(BlockwiseFP8LTXVideoOnlyModelConfigurator):
pass
# ---------------------------------------------------------------------------
# SDOps / ModuleOps / FuseRule assembly
# ---------------------------------------------------------------------------
def build_sd_ops_fp8() -> SDOps:
return (
SDOps("blockwise_fp8_weights")
.with_kv_operation(
_create_weight_quantize_op(EXCLUDED_LAYER_SUBSTRINGS, _fp8_blockwise_quantize_weight),
key_prefix="transformer_blocks.",
key_suffix=".weight",
)
.with_kv_operation(
_create_bias_to_fp32_op(EXCLUDED_LAYER_SUBSTRINGS),
key_prefix="transformer_blocks.",
key_suffix=".bias",
)
)
def build_sd_ops_fp6() -> SDOps:
return (
SDOps("blockwise_fp6_weights")
.with_kv_operation(
_create_weight_quantize_op(EXCLUDED_LAYER_SUBSTRINGS, _fp6_blockwise_quantize_weight),
key_prefix="transformer_blocks.",
key_suffix=".weight",
)
.with_kv_operation(
_create_bias_to_fp32_op(EXCLUDED_LAYER_SUBSTRINGS),
key_prefix="transformer_blocks.",
key_suffix=".bias",
)
)
def build_module_ops_fp8() -> ModuleOps:
return ModuleOps(
name="blockwise_fp8_prepare_for_loading",
matcher=lambda model: isinstance(model, LTXModel),
mutator=lambda model: _replace_linear_modules(model, BlockwiseFP8Linear),
)
def build_module_ops_fp6() -> ModuleOps:
return ModuleOps(
name="blockwise_fp6_prepare_for_loading",
matcher=lambda model: isinstance(model, LTXModel),
mutator=lambda model: _replace_linear_modules(model, BlockwiseFP6Linear),
)
fuse_rule_fp8 = FuseRule(aggregation_dtype=torch.bfloat16, fuse_fn=_blockwise_fp8_fuse)
fuse_rule_fp6 = FuseRule(aggregation_dtype=torch.bfloat16, fuse_fn=_blockwise_fp6_fuse)
@@ -10,7 +10,6 @@ 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
def _read_safetensors_dtypes(path: str) -> dict[str, str]:
@@ -50,34 +49,21 @@ class FP8Linear(nn.Module):
def forward(self, x: torch.Tensor) -> torch.Tensor:
origin_shape = x.shape
if trtllm_scaled_mm_usable():
qinput, cur_input_scale = torch.ops.tensorrt_llm.static_quantize_e4m3_per_tensor(x, self.input_scale)
if qinput.dim() == 3:
qinput = qinput.reshape(-1, qinput.shape[-1])
output = torch.ops.trtllm.cublas_scaled_mm(
qinput,
self.weight.t(),
scale_a=cur_input_scale,
scale_b=self.weight_scale,
bias=None,
out_dtype=x.dtype,
)
else:
# Clamp before cast: out-of-range values cast to NaN/saturated FP8, which
# produces black-screen output on some checkpoints (e.g. ltx-2-19b-dev-fp8).
fp8_min = torch.finfo(torch.float8_e4m3fn).min
fp8_max = torch.finfo(torch.float8_e4m3fn).max
qinput = torch.clamp(x * self.input_scale.reciprocal(), fp8_min, fp8_max).to(torch.float8_e4m3fn)
if qinput.dim() == 3:
qinput = qinput.reshape(-1, qinput.shape[-1])
output = torch._scaled_mm(
qinput,
self.weight.t(),
scale_a=self.input_scale,
scale_b=self.weight_scale,
out_dtype=x.dtype,
use_fast_accum=True,
)
# Clamp before cast: out-of-range values cast to NaN/saturated FP8, which
# produces black-screen output on some checkpoints (e.g. ltx-2-19b-dev-fp8).
fp8_min = torch.finfo(torch.float8_e4m3fn).min
fp8_max = torch.finfo(torch.float8_e4m3fn).max
qinput = torch.clamp(x * self.input_scale.reciprocal(), fp8_min, fp8_max).to(torch.float8_e4m3fn)
if qinput.dim() == 3:
qinput = qinput.reshape(-1, qinput.shape[-1])
output = torch._scaled_mm(
qinput,
self.weight.t(),
scale_a=self.input_scale,
scale_b=self.weight_scale,
out_dtype=x.dtype,
use_fast_accum=True,
)
if self.bias is not None:
output = output + self.bias.to(output.dtype)
@@ -1,37 +0,0 @@
"""Runtime detection of TensorRT-LLM FP8 scaled-matmul availability.
When the TRT-LLM ops are usable on the current host (Linux + Hopper-class CUDA
+ tensorrt_llm wheel installed) we use them since they outperform the PyTorch-native
``torch._scaled_mm`` path. Otherwise we fall back to the native implementation,
which is portable across platforms (Windows, macOS, AMD GPUs).
The check runs once and is cached.
"""
from __future__ import annotations
import platform
from functools import cache
import torch
@cache
def trtllm_scaled_mm_usable() -> bool:
if platform.system() != "Linux":
return False
if not torch.cuda.is_available():
return False
major, minor = torch.cuda.get_device_capability()
sm = major * 10 + minor
if sm < 90 or sm >= 120:
return False
# The import is load-bearing — registers the trtllm torch ops as a side effect.
try:
import tensorrt_llm # noqa: F401, PLC0415
except Exception:
return False
return True
@@ -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 ---
@@ -65,7 +75,9 @@ class GemmaTextEncoder(torch.nn.Module):
pad_token_id = self.processor.tokenizer.pad_token_id if self.processor.tokenizer.pad_token_id is not None else 0
model_inputs = _pad_inputs_for_attention_alignment(model_inputs, pad_token_id=pad_token_id)
with torch.inference_mode(), torch.random.fork_rng(devices=[self.model.device]):
# fork_rng device pinning is only supported for CUDA; MPS/CPU fork CPU RNG only.
fork_devices = [self.model.device] if self.model.device.type == "cuda" else []
with torch.inference_mode(), torch.random.fork_rng(devices=fork_devices):
torch.manual_seed(seed)
outputs = self.model.generate(
**model_inputs,
+14
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
import itertools
import math
from dataclasses import dataclass, replace
from typing import Callable, NamedTuple
@@ -462,3 +463,16 @@ class TileCountConfig:
frames: DimensionTilingConfig = DimensionTilingConfig(num_tiles=1, overlap=0)
height: DimensionTilingConfig = DimensionTilingConfig(num_tiles=1, overlap=0)
width: DimensionTilingConfig = DimensionTilingConfig(num_tiles=1, overlap=0)
def balanced_tile_split(num_tiles: int) -> tuple[int, int]:
"""Factor ``num_tiles`` into ``(small, large)`` as square as possible.
``small`` is the largest divisor not exceeding the square root, so
``small * large == num_tiles`` and ``small <= large``. E.g. 2 -> (1, 2),
4 -> (2, 2), 8 -> (2, 4), 16 -> (4, 4). The caller decides which tiled
dimension gets which factor.
"""
if num_tiles < 1:
raise ValueError(f"num_tiles must be >= 1, got {num_tiles}")
small = next(d for d in range(math.isqrt(num_tiles), 0, -1) if num_tiles % d == 0)
return small, num_tiles // small
+1
View File
@@ -0,0 +1 @@
recursive-include csrc *.h *.cuh *.hpp *.cpp *.cu
+83
View File
@@ -0,0 +1,83 @@
# ltx-kernels
Custom CUDA/C++ kernels for `ltx-core`. Three compiled extensions:
- **`all2all_cpp`** -- All2All communication kernels for multi-GPU tensor
parallelism, used by the sequence-parallel inference path.
- **`ops_cpp`** -- Fused element ops for blockwise quantization: `rms_norm_rope`,
`rms_norm_split_rope`, and FP6 pack/unpack.
- **`blockwise_cpp`** -- Blockwise FP8 GEMM. SM89 (GeForce/Ada) kernel always;
the SM90 (Hopper, `deep_gemm`) kernel is added when a `9.0` architecture is
requested.
The Python surface for blockwise quantization lives in
`ltx_kernels.blockwise` (`functional`, `linear`, `triton_ops`).
## Requirements
- CUDA toolkit (nvcc) matching your GPU architecture
- PyTorch with CUDA support
- Linux
## Building
`ltx-kernels` is excluded from the uv workspace, so a plain `uv sync` does not
build it. From the repository root, build it via the opt-in `kernels` group
(editable, no build isolation -- torch must already be installed):
```bash
uv sync --group kernels
```
Equivalently, install it directly:
```bash
uv pip install -e packages/ltx-kernels --no-build-isolation
```
Set `TORCH_CUDA_ARCH_LIST` to target specific architectures (speeds up compilation):
```bash
# H100 only
TORCH_CUDA_ARCH_LIST="9.0" uv pip install -e packages/ltx-kernels --no-build-isolation
# Multiple architectures
TORCH_CUDA_ARCH_LIST="9.0 9.0a 10.0 12.0" uv pip install -e packages/ltx-kernels --no-build-isolation
```
When `TORCH_CUDA_ARCH_LIST` is unset the build targets every supported
architecture (so `uv pip install` "just works" on a dev box); pin it on build
hosts to cut compile time. Any `9.0` entry enables the SM90 GEMM kernel, which
is compiled for `sm_90a` (the deep_gemm kernel uses wgmma/TMA).
### cutlass headers
`blockwise_cpp` includes cute/cutlass headers (header-only; compiled into the
extension, with no runtime dependency). The build fetches them automatically on
first use: a blobless, `include/`-only sparse clone of cutlass pinned to commit
`afa17722` (v3.8.0), cached under `~/.cache/ltx-kernels/` (~25 MB) and reused
across builds.
- Set `CUTLASS_DIR=/path/to/cutlass` to use an existing checkout (uses
`$CUTLASS_DIR/include` and skips the fetch).
- Set `LTX_KERNELS_CACHE_DIR` to override the cache location.
To bump cutlass, change `CUTLASS_REF` in `setup.py`.
## Testing
Tests require a CUDA GPU:
```bash
uv run pytest packages/ltx-kernels/tests/ -v
```
## Operations
`all2all_cpp`:
- **send_recv_heads** -- Redistributes attention heads across GPUs (All2All)
- **gather_heads** -- Inverse of send_recv_heads
- **allgather** -- Gathers sequence tokens from all ranks
All operations support BFloat16 and Float8 (e4m3fn) data types.
@@ -0,0 +1,424 @@
/**
* @file all2all.cpp
* @brief Implementation of All2All communication primitives for multi-GPU tensor parallelism.
*
* This file implements the All2All class which provides efficient inter-GPU communication
* using CUDA IPC (Inter-Process Communication). The implementation supports:
* - Head redistribution for tensor-parallel attention (send_recv_heads, gather_heads)
* - Sequence gathering for cross-rank aggregation (allgather)
*
* All operations use a barrier-based synchronization protocol where each GPU writes
* directly to remote GPU memory via IPC, then signals completion through atomic
* operations on barrier counters.
*/
#include <ATen/cuda/CUDAContext.h>
#include <ATen/cuda/CUDADataType.h>
#include <c10/cuda/CUDAGuard.h>
#include <chrono>
#include <cuda_runtime.h>
#include <memory>
#include <pybind11/functional.h>
#include <torch/python.h>
#include "all2all.hpp"
#include "cuda/api.cuh"
#include "cuda/configs.cuh"
namespace ltx_kernels {
namespace all2all {
/**
* Constructs the All2All communication manager.
*
* Memory Allocation Strategy:
* The constructor allocates a single contiguous GPU memory block that contains:
* 1. Data buffer (tensor_bytes): Space for tensor data exchange
* 2. Barrier signals (MAX_NUM_PEERS * sizeof(int)): Per-rank completion counters
* 3. Buffer pointers (MAX_NUM_PEERS * sizeof(void*)): GPU-accessible pointer array
* 4. Barrier pointer array (MAX_NUM_PEERS * sizeof(int*)): GPU-accessible signal pointers
*
* This layout minimizes memory allocations and allows the entire region to be
* shared via a single IPC handle.
*/
All2All::All2All(int rank, int world_size, int num_tokens, int hidden_dim, int num_sms, at::ScalarType tensor_dtype,
double timeout_seconds)
: rank(rank), world_size(world_size), num_sms(num_sms), max_tokens(num_tokens), num_elems(0), tensor_bytes(0),
tensor_dtype(tensor_dtype) {
num_elems = int64_t(num_tokens) * int64_t(hidden_dim);
tensor_bytes = num_elems * elementSize(tensor_dtype);
// Derive the barrier timeout from the device's peak SM clock so the wall-clock guard is
// correct on any GPU (the kernel counts SM cycles via clock64). Use cudaDeviceGetAttribute,
// not cudaDeviceProp::clockRate, which was removed in CUDA 13. The attribute is in kHz.
int device = 0;
CUDA_CHECK(cudaGetDevice(&device));
int sm_clock_khz = 0;
CUDA_CHECK(cudaDeviceGetAttribute(&sm_clock_khz, cudaDevAttrClockRate, device));
sm_clock_hz_ = static_cast<double>(sm_clock_khz) * 1e3;
set_timeout_seconds(timeout_seconds);
// Calculate sizes for each region of the shared memory block
int64_t ptrs_bytes = MAX_NUM_PEERS * sizeof(void *);
int64_t barrier_signal_bytes = MAX_NUM_PEERS * sizeof(int);
int64_t barrier_signal_ptrs_bytes = MAX_NUM_PEERS * sizeof(int *);
// Allocate GPU memory for token count arrays (used by kernels)
CUDA_CHECK(cudaMalloc(reinterpret_cast<void **>(&rank_tokens_gpu), sizeof(int) * MAX_NUM_PEERS));
CUDA_CHECK(cudaMalloc(reinterpret_cast<void **>(&prefix_rank_tokens_gpu), sizeof(int) * MAX_NUM_PEERS));
// Allocate the main shared memory block and create IPC handle
// Layout: [data_buffer | barrier_signals | buffer_ptrs | barrier_signal_ptrs]
CUDA_CHECK(
cudaMalloc(&buffer_ptrs[rank], tensor_bytes + barrier_signal_bytes + ptrs_bytes + barrier_signal_ptrs_bytes));
CUDA_CHECK(cudaIpcGetMemHandle(&ipc_handlers[rank], buffer_ptrs[rank]));
// Set up pointers to each region within the allocated block
buffer_ptrs_gpu =
reinterpret_cast<void **>(static_cast<uint8_t *>(buffer_ptrs[rank]) + tensor_bytes + barrier_signal_bytes);
barrier_signal_ptrs[rank] = reinterpret_cast<int *>(static_cast<uint8_t *>(buffer_ptrs[rank]) + tensor_bytes);
barrier_signal_ptrs_gpu = reinterpret_cast<int **>(static_cast<uint8_t *>(buffer_ptrs[rank]) + tensor_bytes +
barrier_signal_bytes + ptrs_bytes);
// Initialize barrier signals to zero
CUDA_CHECK(cudaMemset(barrier_signal_ptrs[rank], 0, barrier_signal_bytes));
}
All2All::~All2All() noexcept(false) {
if (!destroyed) {
printf("WARNING: destroy() was not called, which can leak resources.\n");
fflush(stdout);
destroy();
}
}
/**
* Releases all allocated resources.
*
* This must be called explicitly before destruction to ensure proper cleanup of:
* - IPC memory mappings to remote GPUs
* - Local GPU memory allocations
*
* The method synchronizes the device to ensure all pending operations complete
* before releasing resources.
*/
void All2All::destroy() {
if (destroyed) {
return;
}
CUDA_CHECK(cudaDeviceSynchronize());
// Close IPC mappings to remote GPU memory (skip our own rank)
// Only close handles that were actually opened via sync()
for (int i = 0; i < world_size; i++) {
if (i != rank && buffer_ptrs[i] != nullptr) {
CUDA_CHECK(cudaIpcCloseMemHandle(buffer_ptrs[i]));
}
}
// Free local GPU memory allocations
CUDA_CHECK(cudaFree(buffer_ptrs[rank]));
CUDA_CHECK(cudaFree(rank_tokens_gpu));
CUDA_CHECK(cudaFree(prefix_rank_tokens_gpu));
destroyed = true;
}
/**
* Opens IPC memory mappings to all peer GPUs.
*
* This method processes IPC handles gathered from all ranks and opens memory
* mappings to enable direct GPU-to-GPU memory access. After calling this method,
* each GPU can read/write directly to any other GPU's buffer via buffer_ptrs.
*
* The barrier_signal_ptrs are also set up to point to the correct offset within
* each peer's shared memory block.
*/
void All2All::sync(const std::vector<std::optional<pybind11::bytearray>> &all_gathered_handles) {
for (int i = 0; i < world_size; i++) {
auto handle_str = std::string(all_gathered_handles[i].value());
EP_HOST_ASSERT(handle_str.size() == CUDA_IPC_HANDLE_SIZE);
if (i != rank) {
// Open IPC mapping to remote GPU's memory
std::memcpy(ipc_handlers[i].reserved, handle_str.c_str(), CUDA_IPC_HANDLE_SIZE);
CUDA_CHECK(cudaIpcOpenMemHandle(&buffer_ptrs[i], ipc_handlers[i], cudaIpcMemLazyEnablePeerAccess));
// Calculate offset to barrier signals in remote buffer
barrier_signal_ptrs[i] = reinterpret_cast<int *>(static_cast<uint8_t *>(buffer_ptrs[i]) + tensor_bytes);
} else {
// Verify our own handle matches what we sent
EP_HOST_ASSERT(std::memcmp(ipc_handlers[i].reserved, handle_str.c_str(), CUDA_IPC_HANDLE_SIZE) == 0);
}
}
// Copy pointer arrays to GPU for kernel access
CUDA_CHECK(cudaMemcpy(buffer_ptrs_gpu, buffer_ptrs, sizeof(void *) * world_size, cudaMemcpyHostToDevice));
CUDA_CHECK(
cudaMemcpy(barrier_signal_ptrs_gpu, barrier_signal_ptrs, sizeof(int *) * world_size, cudaMemcpyHostToDevice));
CUDA_CHECK(cudaDeviceSynchronize());
}
pybind11::bytearray All2All::get_local_ipc_handle() const {
return {ipc_handlers[rank].reserved, CUDA_IPC_HANDLE_SIZE};
}
/**
* Configures token distribution across ranks for the current batch.
*
* This method computes prefix sums needed by the kernels to calculate source
* and destination offsets. It must be called before any communication operation
* when the token distribution changes between batches.
*
* Example: For rank_num_tokens = {128, 96, 128, 64}
* - rank_tokens = {128, 96, 128, 64}
* - prefix_rank_tokens = {0, 128, 224, 352}
* - total_tokens = 416
*/
void All2All::set_rank_tokens(const std::vector<int> &rank_num_tokens) {
EP_HOST_ASSERT(static_cast<int>(rank_num_tokens.size()) == world_size);
// Initialize prefix sums to zero
for (int i = 0; i < world_size; i++) {
prefix_rank_tokens[i] = 0;
}
// Compute prefix sums (exclusive scan)
for (int i = 0; i < world_size; i++) {
rank_tokens[i] = rank_num_tokens[i];
if (i > 0) {
prefix_rank_tokens[i] = prefix_rank_tokens[i - 1] + rank_tokens[i - 1];
}
}
// Total tokens is the sum of all rank tokens
total_tokens = prefix_rank_tokens[world_size - 1] + rank_tokens[world_size - 1];
// Copy to GPU for kernel access
CUDA_CHECK(cudaMemcpy(rank_tokens_gpu, rank_tokens, sizeof(int) * MAX_NUM_PEERS, cudaMemcpyHostToDevice));
CUDA_CHECK(
cudaMemcpy(prefix_rank_tokens_gpu, prefix_rank_tokens, sizeof(int) * MAX_NUM_PEERS, cudaMemcpyHostToDevice));
CUDA_CHECK(cudaDeviceSynchronize());
}
/**
* Creates a tensor from the local IPC buffer.
*
* This helper method returns either a zero-copy view of the IPC buffer or
* a newly allocated tensor with the data copied. The zero-copy mode is more
* efficient but the tensor lifetime is tied to the All2All instance.
*
* @note The buffer pointer is cast to the template type T for proper interpretation.
*/
at::Tensor All2All::get_local_buffer_tensor(at::Tensor &x, int batch_size, int out_tokens, int out_heads, int head_size,
bool should_copy, cudaStream_t stream) {
auto ptr = buffer_ptrs[rank];
if (should_copy) {
// Allocate new tensor and copy data from IPC buffer
auto out_tensor = torch::empty({batch_size, out_tokens, out_heads, head_size}, x.options());
CUDA_CHECK(cudaMemcpyAsync(out_tensor.data_ptr(), ptr,
int64_t(batch_size) * int64_t(out_tokens) * int64_t(out_heads) * int64_t(head_size) *
int64_t(elementSize(x.scalar_type())),
cudaMemcpyDeviceToDevice, stream));
return out_tensor;
} else {
// Return a view directly into the IPC buffer (zero-copy)
auto out_tensor = torch::from_blob(ptr, {batch_size, out_tokens, out_heads, head_size}, x.options());
return out_tensor;
}
}
/**
* All2All communication to redistribute attention heads across GPUs.
*
* This operation is used in tensor-parallel transformers to exchange attention heads:
* - Before: Each GPU has all tokens but only a subset of heads
* - After: Each GPU has all tokens with heads redistributed
*
* Tensor Layout Transformation:
* Input: [batch, local_tokens, all_heads, head_size] per GPU
* Output: [batch, all_tokens, heads_per_rank, head_size] per GPU
*
* The operation partitions heads evenly: heads_per_rank = all_heads / world_size
* GPU i receives heads [i*heads_per_rank : (i+1)*heads_per_rank] from all GPUs.
*/
at::Tensor All2All::send_recv_heads(at::Tensor &x, bool copy_output) {
// Validate input tensor properties
EP_HOST_ASSERT(x.dim() == 4 and x.is_contiguous());
EP_HOST_ASSERT(x.dtype() == tensor_dtype);
EP_HOST_ASSERT(x.device().is_cuda());
EP_HOST_ASSERT(x.device().index() == rank);
int batch_size = x.size(0);
int num_tokens = x.size(1);
int num_heads = x.size(2);
int head_size = x.size(3);
// Output dimensions after redistribution
int out_tokens = total_tokens; // All tokens from all ranks
int out_heads = num_heads / world_size; // Each rank gets 1/world_size of heads
EP_HOST_ASSERT(int64_t(batch_size) * int64_t(out_tokens) * int64_t(out_heads) * int64_t(head_size) *
int64_t(elementSize(x.scalar_type())) <=
tensor_bytes);
at::cuda::CUDAGuard device_guard{x.device()};
auto stream = at::cuda::getCurrentCUDAStream().stream();
// Launch the All2All kernel
all2all_cuda::all2all_head_launch(buffer_ptrs_gpu, barrier_signal_ptrs_gpu, x.data_ptr(), prefix_rank_tokens_gpu,
rank, world_size, batch_size, total_tokens, num_tokens, num_heads, head_size,
stream, num_sms, tensor_dtype, timeout_cycles_);
return get_local_buffer_tensor(x, batch_size, out_tokens, out_heads, head_size, copy_output, stream);
}
/**
* Inverse All2All to gather heads back to original distribution.
*
* This is the inverse operation of send_recv_heads(). It redistributes data
* so each GPU gets back its original tokens with all attention heads.
*
* Tensor Layout Transformation:
* Input: [batch, all_tokens, heads_per_rank, head_size] per GPU
* Output: [batch, local_tokens, all_heads, head_size] per GPU
*
* Each GPU sends its portion of tokens to the originating rank, reconstructing
* the original head distribution.
*/
at::Tensor All2All::gather_heads(at::Tensor &x, bool copy_output) {
// Validate input tensor properties
EP_HOST_ASSERT(x.dim() == 4 and x.is_contiguous());
EP_HOST_ASSERT(x.dtype() == tensor_dtype);
EP_HOST_ASSERT(x.device().is_cuda());
EP_HOST_ASSERT(x.device().index() == rank);
at::cuda::CUDAGuard device_guard{x.device()};
auto stream = at::cuda::getCurrentCUDAStream().stream();
int batch_size = x.size(0);
int num_heads = x.size(2) * world_size; // Reconstruct total head count
int head_size = x.size(3);
// Output dimensions: this rank's tokens with all heads
int out_tokens = rank_tokens[rank];
int out_heads = num_heads;
EP_HOST_ASSERT(int64_t(batch_size) * int64_t(out_tokens) * int64_t(out_heads) * int64_t(head_size) *
int64_t(elementSize(x.scalar_type())) <=
tensor_bytes);
// Launch the gather kernel
all2all_cuda::all2all_head_gather_launch(buffer_ptrs_gpu, barrier_signal_ptrs_gpu, x.data_ptr(), rank_tokens_gpu,
prefix_rank_tokens_gpu, rank, world_size, batch_size, total_tokens,
num_heads, head_size, stream, num_sms, tensor_dtype, timeout_cycles_);
return get_local_buffer_tensor(x, batch_size, out_tokens, out_heads, head_size, copy_output, stream);
}
/**
* AllGather operation to collect sequence tokens from all ranks.
*
* Each GPU contributes its local sequence tokens, which are gathered into
* a complete sequence replicated on all GPUs. This is typically used after
* tensor-parallel operations to reconstruct the full sequence.
*
* Tensor Layout Transformation:
* Input: [batch, local_seqlen, heads, head_size] per GPU
* Output: [batch, total_seqlen, heads, head_size] per GPU (identical on all GPUs)
*
* Each GPU's tokens are placed at offset prefix_rank_tokens[rank] in the output.
*/
at::Tensor All2All::allgather(at::Tensor &x, bool copy_output) {
// Validate input tensor properties
EP_HOST_ASSERT(x.dim() == 4 and x.is_contiguous());
EP_HOST_ASSERT(x.dtype() == tensor_dtype);
EP_HOST_ASSERT(x.device().is_cuda());
EP_HOST_ASSERT(x.device().index() == rank);
at::cuda::CUDAGuard device_guard{x.device()};
auto stream = at::cuda::getCurrentCUDAStream().stream();
int batch_size = x.size(0);
int seqlen = x.size(1);
int num_heads = x.size(2);
int head_size = x.size(3);
// Output contains all tokens from all ranks
int out_tokens = total_tokens;
int out_heads = num_heads;
int hidden_dim = num_heads * head_size;
EP_HOST_ASSERT(int64_t(batch_size) * int64_t(out_tokens) * int64_t(out_heads) * int64_t(head_size) *
int64_t(elementSize(x.scalar_type())) <=
tensor_bytes);
// Launch the allgather kernel
all2all_cuda::allgather_launch(buffer_ptrs_gpu, barrier_signal_ptrs_gpu, x.data_ptr(), prefix_rank_tokens_gpu, rank,
world_size, batch_size, seqlen, hidden_dim, total_tokens, stream, num_sms,
tensor_dtype, timeout_cycles_);
return get_local_buffer_tensor(x, batch_size, out_tokens, out_heads, head_size, copy_output, stream);
}
} // namespace all2all
} // namespace ltx_kernels
/**
* Python bindings for the All2All communication library.
*
* Usage from Python:
* import all2all_cpp
*
* # Create instance (one per GPU)
* comm = all2all_cpp.All2All(rank, world_size, max_tokens, hidden_dim, num_sms, dtype)
*
* # Exchange IPC handles and synchronize
* handle = comm.get_local_ipc_handle()
* # ... gather handles via NCCL ...
* comm.sync(all_handles)
*
* # Set token distribution
* comm.set_rank_tokens([128, 128, 128, 128])
*
* # Perform operations
* output = comm.send_recv_heads(input_tensor, copy_output=False)
*
* # Cleanup
* comm.destroy()
*/
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.doc() = "High-performance All2All communication library for multi-GPU tensor parallelism.\n\n"
"This library provides IPC-based All2All operations optimized for transformer models.\n"
"Supported operations:\n"
" - send_recv_heads: Redistribute attention heads across GPUs\n"
" - gather_heads: Inverse of send_recv_heads\n"
" - allgather: Gather sequence tokens from all ranks\n";
pybind11::class_<ltx_kernels::all2all::All2All>(
m, "All2All",
"Manages All2All communication state for multi-GPU operations.\n\n"
"Args:\n"
" rank: This GPU's rank (0 to world_size-1)\n"
" world_size: Total number of GPUs\n"
" num_tokens: Maximum tokens per rank\n"
" hidden_dim: Hidden dimension (heads * head_size)\n"
" num_sms: Number of SMs for kernel launches\n"
" tensor_dtype: Tensor data type (torch.bfloat16 or torch.float8_e4m3fn)\n"
" timeout_seconds: Optional initial barrier timeout in seconds (defaults to the kernel default)")
.def(pybind11::init<int, int, int, int, int, at::ScalarType>())
.def(pybind11::init<int, int, int, int, int, at::ScalarType, double>())
.def("get_local_ipc_handle", &ltx_kernels::all2all::All2All::get_local_ipc_handle,
"Returns the IPC handle for this rank's buffer.")
.def("sync", &ltx_kernels::all2all::All2All::sync, "Opens IPC mappings to all peer GPUs using gathered handles.")
.def("destroy", &ltx_kernels::all2all::All2All::destroy,
"Releases all GPU resources. Must be called before destruction.")
.def("send_recv_heads", &ltx_kernels::all2all::All2All::send_recv_heads,
"All2All operation to redistribute attention heads.")
.def("gather_heads", &ltx_kernels::all2all::All2All::gather_heads,
"Inverse All2All to gather heads back to original distribution.")
.def("allgather", &ltx_kernels::all2all::All2All::allgather, "Gathers sequence tokens from all ranks.")
.def("set_rank_tokens", &ltx_kernels::all2all::All2All::set_rank_tokens,
"Sets token counts per rank for the current batch.")
.def("set_timeout_seconds", &ltx_kernels::all2all::All2All::set_timeout_seconds,
"Sets the barrier timeout in seconds (converted to cycles via the device peak SM clock).");
}
@@ -0,0 +1,265 @@
/**
* @file all2all.hpp
* @brief High-performance All2All communication primitives for multi-GPU tensor parallelism.
*
* This library provides efficient All2All communication operations optimized for transformer
* models using tensor parallelism. It uses CUDA IPC (Inter-Process Communication) for
* zero-copy data transfer between GPUs in the same node.
*
* ## Architecture Overview
*
* The All2All class manages shared memory buffers accessible by all GPUs via IPC handles.
* Each GPU allocates a contiguous memory region containing:
* - Data buffer: Stores tensor data for exchange
* - Barrier signals: Synchronization counters for coordination
* - GPU pointer arrays: Device-accessible pointers to all peer buffers
*
* Memory Layout (per GPU):
* ```
* |<---- tensor_bytes ---->|<-- barrier signals -->|<-- buffer_ptrs_gpu -->|<-- barrier_signal_ptrs_gpu -->|
* | Data Buffer | MAX_PEERS * int | MAX_PEERS * void* | MAX_PEERS * int* |
* ```
*
* ## Supported Operations
*
* 1. **send_recv_heads**: Redistributes attention heads across GPUs (All2All)
* - Input: [batch, tokens, heads, head_size] on each GPU
* - Output: [batch, total_tokens, heads/world_size, head_size] on each GPU
*
* 2. **gather_heads**: Inverse of send_recv_heads
* - Gathers distributed heads back to original distribution
*
* 3. **allgather**: Gathers sequence data from all ranks
* - Each GPU contributes its local tokens to form the complete sequence
*
* ## Thread Safety
*
* - The class is NOT thread-safe. Each thread/process should have its own instance.
* - Multiple CUDA streams may use the same instance sequentially.
* - The `destroy()` method MUST be called before destruction to properly release IPC handles.
*
* ## Usage Example
*
* ```cpp
* // Initialize on each GPU
* auto comm = All2All(rank, world_size, max_tokens, hidden_dim, num_sms, dtype);
*
* // Exchange IPC handles (via NCCL or other collective)
* auto my_handle = comm.get_local_ipc_handle();
* // ... gather all handles ...
* comm.sync(all_handles);
*
* // Set token distribution for current batch
* comm.set_rank_tokens({128, 128, 128, 128}); // tokens per rank
*
* // Perform All2All on attention heads
* auto result = comm.send_recv_heads(input_tensor, copy_output=false);
*
* // Clean up
* comm.destroy();
* ```
*/
#pragma once
#include "cuda/configs.cuh"
#include "event.hpp"
#include <cmath>
#include <limits>
#include <pybind11/pybind11.h>
#include <pybind11/pytypes.h>
#include <stdexcept>
#include <torch/types.h>
#include <tuple>
#include <vector>
namespace ltx_kernels {
namespace all2all {
/**
* @class All2All
* @brief Manages All2All communication state and operations for multi-GPU tensor parallelism.
*
* This class encapsulates the IPC-based communication infrastructure needed for
* efficient All2All operations. It maintains shared memory buffers, barrier signals,
* and provides methods for head-parallel tensor redistribution.
*/
struct All2All {
private:
int rank; ///< This GPU's rank (0 to world_size-1)
int world_size; ///< Total number of GPUs in the communication group
int num_sms; ///< Number of SMs to use for kernel launches
int max_tokens; ///< Maximum number of tokens the buffer was allocated for
int64_t num_elems; ///< Number of elements in the data buffer (tokens * hidden_dim)
int64_t tensor_bytes; ///< Size of the data buffer in bytes
/// Host array of pointers to each rank's data buffer (GPU memory)
void *buffer_ptrs[MAX_NUM_PEERS] = {nullptr};
/// Device-accessible array of buffer pointers (copied to GPU)
void **buffer_ptrs_gpu = nullptr;
/// Host array of pointers to each rank's barrier signal buffer
int *barrier_signal_ptrs[MAX_NUM_PEERS] = {nullptr};
/// Device-accessible array of barrier signal pointers
int **barrier_signal_ptrs_gpu = nullptr;
/// IPC handles for sharing memory between processes
cudaIpcMemHandle_t ipc_handlers[MAX_NUM_PEERS];
at::ScalarType tensor_dtype; ///< Data type of tensors (BFloat16 or Float8_e4m3fn)
bool destroyed = false; ///< Flag to track if resources have been released
int total_tokens; ///< Sum of tokens across all ranks for current batch
int rank_tokens[MAX_NUM_PEERS]; ///< Number of tokens on each rank
int prefix_rank_tokens[MAX_NUM_PEERS]; ///< Cumulative sum of tokens (for offset calculation)
int *rank_tokens_gpu = nullptr; ///< Device copy of rank_tokens
int *prefix_rank_tokens_gpu = nullptr; ///< Device copy of prefix_rank_tokens
/// Device peak SM clock in Hz (from cudaDeviceGetAttribute(cudaDevAttrClockRate)), queried
/// once at construction. Used to convert a wall-clock timeout in seconds to barrier cycles.
double sm_clock_hz_ = 0.0;
/// All2All barrier timeout in GPU clock cycles. The constructor sets it from
/// DEFAULT_BARRIER_TIMEOUT_SECONDS and the queried SM clock; raise it (set_timeout_seconds)
/// to tolerate large cross-rank kernel-launch skew during the first torch.compile forward,
/// where one rank's recompile can delay its launch past the steady-state timeout.
uint64_t timeout_cycles_ = 0;
public:
/**
* @brief Constructs an All2All communication manager.
*
* Allocates GPU memory for the local data buffer, barrier signals, and pointer arrays.
* The IPC handle for the local buffer is created and can be retrieved via get_local_ipc_handle().
*
* @param rank This GPU's rank in the communication group (0-indexed)
* @param world_size Total number of GPUs/ranks
* @param num_tokens Maximum number of tokens this rank will handle
* @param hidden_dim Hidden dimension size (heads * head_size)
* @param num_sms Number of CUDA SMs to use for kernel execution
* @param tensor_dtype Data type for tensors (BFloat16 or Float8_e4m3fn)
* @param timeout_seconds Initial barrier timeout in seconds (see set_timeout_seconds); may be
* raised/reset at runtime for the first torch.compile forward
*/
All2All(int rank, int world_size, int num_tokens, int hidden_dim, int num_sms, at::ScalarType tensor_dtype,
double timeout_seconds = DEFAULT_BARRIER_TIMEOUT_SECONDS);
/**
* @brief Destructor - warns if destroy() was not called.
*
* @warning Always call destroy() explicitly before the destructor to properly
* release IPC handles. Failing to do so may leak resources.
*/
~All2All() noexcept(false);
/**
* @brief Synchronizes IPC handles from all ranks and opens remote memory mappings.
*
* This method must be called after all ranks have created their All2All instances
* and exchanged IPC handles via an external collective (e.g., NCCL allgather).
*
* @param all_gathered_handles Vector of IPC handles from all ranks (indexed by rank)
*/
void sync(const std::vector<std::optional<pybind11::bytearray>> &all_gathered_handles);
/**
* @brief Returns the IPC handle for this rank's shared buffer.
*
* The returned handle should be gathered across all ranks and passed to sync().
*
* @return pybind11::bytearray containing the CUDA IPC handle (CUDA_IPC_HANDLE_SIZE bytes)
*/
pybind11::bytearray get_local_ipc_handle() const;
/**
* @brief Creates a tensor view or copy of the local output buffer.
*
* @param x Reference tensor for options (dtype, device)
* @param batch_size Batch dimension size
* @param out_tokens Output token dimension size
* @param out_heads Output heads dimension size
* @param head_size Head dimension size
* @param should_copy If true, copies data to a new tensor; if false, returns a view
* @param stream CUDA stream for async copy
* @return Tensor with shape [batch_size, out_tokens, out_heads, head_size]
*/
at::Tensor get_local_buffer_tensor(at::Tensor &x, int batch_size, int out_tokens, int out_heads, int head_size,
bool should_copy, cudaStream_t stream);
/**
* @brief Releases all GPU resources and closes IPC handles.
*
* This method MUST be called before the object is destroyed. It synchronizes
* the device, closes remote IPC mappings, and frees local GPU memory.
*/
void destroy();
/**
* @brief Performs All2All communication to redistribute attention heads.
*
* Redistributes tensor from [batch, local_tokens, all_heads, head_size] to
* [batch, all_tokens, local_heads, head_size]. Each rank sends its portion
* of heads to the corresponding target rank.
*
* @param x Input tensor with shape [batch, num_tokens, num_heads, head_size]
* @param copy_output If true, returns a copy; if false, returns a view of the IPC buffer
* @return Tensor with shape [batch, total_tokens, num_heads/world_size, head_size]
*/
at::Tensor send_recv_heads(at::Tensor &x, bool copy_output);
/**
* @brief Performs inverse All2All to gather heads back to original distribution.
*
* Inverse of send_recv_heads(). Redistributes from [batch, all_tokens, local_heads, head_size]
* back to [batch, local_tokens, all_heads, head_size].
*
* @param x Input tensor with shape [batch, total_tokens, heads_per_rank, head_size]
* @param copy_output If true, returns a copy; if false, returns a view of the IPC buffer
* @return Tensor with shape [batch, rank_tokens[rank], num_heads, head_size]
*/
at::Tensor gather_heads(at::Tensor &x, bool copy_output);
/**
* @brief Gathers sequence tokens from all ranks.
*
* Each rank contributes its local sequence tokens, which are gathered into
* a complete sequence on all ranks.
*
* @param x Input tensor with shape [batch, seqlen, num_heads, head_size]
* @param copy_output If true, returns a copy; if false, returns a view of the IPC buffer
* @return Tensor with shape [batch, total_tokens, num_heads, head_size]
*/
at::Tensor allgather(at::Tensor &x, bool copy_output);
/**
* @brief Sets the token count for each rank in the current batch.
*
* Must be called before send_recv_heads(), gather_heads(), or allgather()
* to configure the token distribution. This allows variable-length sequences
* across ranks.
*
* @param rank_num_tokens Vector of token counts, one per rank (must have world_size elements)
*/
void set_rank_tokens(const std::vector<int> &rank_num_tokens);
/**
* @brief Sets the all2all barrier timeout in seconds.
*
* Converted to GPU clock cycles using the device's peak SM clock (queried at construction).
* Relaxes deadlock detection during the first torch.compile forward, where asymmetric
* per-rank recompilation can delay a rank's kernel launch beyond the steady-state timeout.
* Reset to the default for steady-state replay.
*/
void set_timeout_seconds(double seconds) {
if (!std::isfinite(seconds) || seconds < 0.0) {
throw std::invalid_argument("All2All timeout (seconds) must be finite and non-negative");
}
// Saturate rather than overflow the float->uint64 cast (out-of-range conversion is UB).
const double cycles = seconds * sm_clock_hz_;
const double max_cycles = static_cast<double>(std::numeric_limits<uint64_t>::max());
timeout_cycles_ = cycles >= max_cycles ? std::numeric_limits<uint64_t>::max() : static_cast<uint64_t>(cycles);
}
};
} // namespace all2all
} // namespace ltx_kernels
@@ -0,0 +1,372 @@
/**
* @file all2all_heads.cu
* @brief CUDA kernels for All2All attention head redistribution.
*
* This file implements the GPU kernels for redistributing attention heads across
* multiple GPUs using IPC-based direct memory access. The kernels are designed
* for tensor-parallel transformer models where attention heads need to be
* exchanged between GPUs.
*
* ## Algorithm Overview
*
* The kernels use a direct-write approach where each GPU writes its data directly
* to the target GPU's memory buffer via IPC. This avoids intermediate copies and
* achieves near-peak memory bandwidth utilization.
*
* ## SM Work Distribution (Round-Robin)
*
* SMs are distributed round-robin among target ranks to handle non-divisible SM counts:
* - SM i writes to rank (i % world_size)
* - With 132 SMs and 8 GPUs: ranks 0-3 get 17 SMs, ranks 4-7 get 16 SMs
* - Each SM group processes all tokens for its assigned target rank
* - Within each group, SMs cooperate to cover all tokens in strided fashion
*
* ## Synchronization Protocol
*
* After data transfer, a barrier synchronization ensures all ranks have completed:
* 1. Each SM atomically increments the target rank's barrier counter for this rank
* 2. SM 0 waits until it has received signals from all ranks
* 3. Barrier counters are reset for the next operation
*/
#include "cuda/configs.cuh"
#include "cuda/exceptions.cuh"
#include "cuda/utils.cuh"
#include <ATen/cuda/CUDADataType.h>
namespace ltx_kernels {
namespace all2all {
namespace all2all_cuda {
/**
* @brief All2All kernel for redistributing attention heads across GPUs.
*
* This kernel performs the "send" phase of All2All: each GPU writes its assigned
* subset of attention heads to all other GPUs. The data layout transformation is:
*
* Source: [batch, num_tokens, num_heads, head_size]
* Dest: [batch, total_tokens, heads_per_rank, head_size]
*
* Each GPU writes heads [target_rank * heads_per_rank : (target_rank+1) * heads_per_rank]
* to target_rank's buffer at token offset prefix_rank_tokens[rank].
*
* ## Memory Layout
*
* Input tensor x (row-major, contiguous):
* - Batch dimension: outermost
* - Token dimension: batch_stride = num_tokens * num_heads * head_size
* - Head dimension: token_stride = num_heads * head_size
* - Head element: head_stride = head_size
*
* Output buffer (per target rank):
* - Similar layout but with heads_per_rank instead of num_heads
* - Tokens from this rank placed at offset prefix_rank_tokens[rank]
*
* ## Thread Block Organization
*
* Each thread block handles multiple tokens cooperatively:
* - Threads are organized in a 2D logical grid (rows=tokens, cols=elements)
* - Each thread copies 16 bytes (int4) per iteration
* - num_threads_per_token = (heads_per_rank * head_size) / elements_per_thread
* - num_tokens_per_copy = num_threads / num_threads_per_token
*
* @tparam ELEM_T Element type (at::BFloat16 or at::Float8_e4m3fn)
* @param buffer_ptrs Device array of pointers to each rank's data buffer
* @param barrier_signal_ptrs Device array of pointers to each rank's barrier signals
* @param x Source tensor data pointer
* @param rank This GPU's rank
* @param world_size Total number of GPUs
* @param batch_size Number of batches
* @param num_tokens Number of tokens on this rank
* @param num_heads Total number of attention heads
* @param head_size Size of each attention head
* @param total_tokens Sum of tokens across all ranks
* @param prefix_rank_tokens Cumulative token counts for offset calculation
*/
template <typename ELEM_T>
__global__ void send_recv_all2all(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int rank, int world_size,
int batch_size, int num_tokens, int num_heads, int head_size, int total_tokens,
int *prefix_rank_tokens, uint64_t timeout_cycles) {
// Grid dimensions
int num_sms = gridDim.x;
int sm_id = blockIdx.x;
int num_threads = blockDim.x;
// === SM Work Distribution (Round-Robin) ===
// Use modular assignment to handle num_sms not divisible by world_size.
// This ensures all SMs are utilized: some ranks get ceil(num_sms/world_size)
// SMs, others get floor(num_sms/world_size) SMs.
int64_t target_rank = get_target_rank(sm_id, world_size);
int64_t rank_local_sm_id = get_rank_local_sm_id(sm_id, world_size);
int64_t num_sms_for_this_rank = get_num_sms_for_rank(target_rank, num_sms, world_size);
// === Head Assignment ===
// Heads are partitioned evenly: rank i gets heads [i*hpr : (i+1)*hpr]
int64_t heads_per_rank = num_heads / world_size;
int64_t head_id = target_rank * heads_per_rank; // Starting head for target rank
// === Thread Mapping ===
// Each thread copies an int4 (16 bytes) per memory operation
// Threads form a 2D grid: (tokens_per_copy, threads_per_token)
int64_t num_elems_per_thread = sizeof(int4) / sizeof(ELEM_T);
int64_t num_threads_per_token = heads_per_rank * head_size / num_elems_per_thread;
int64_t num_tokens_per_copy = num_threads / num_threads_per_token;
// 2D thread coordinates within the logical grid
int64_t copy_thr_col_idx = threadIdx.x % num_threads_per_token; // Element offset
int64_t copy_thr_row_idx = threadIdx.x / num_threads_per_token; // Token offset
// Get target rank's buffer pointer
auto ptr = reinterpret_cast<void *>(static_cast<int8_t *>(buffer_ptrs[target_rank]));
// Use 64-bit arithmetic to avoid overflow for large tensors
int64_t num_tokens_64b = int64_t(num_tokens);
int64_t num_heads_64b = int64_t(num_heads);
int64_t head_size_64b = int64_t(head_size);
// === Main Copy Loop ===
// Iterate over batches and tokens, with SMs in the same group
// working on different token ranges in strided fashion
for (int64_t batch_ind = 0; batch_ind < batch_size; batch_ind++) {
// Strided token iteration: each SM in the group handles different token ranges
for (int64_t token_idx = rank_local_sm_id * num_tokens_per_copy; token_idx < num_tokens;
token_idx += num_tokens_per_copy * num_sms_for_this_rank) {
int64_t copy_token_idx = token_idx + copy_thr_row_idx;
// Destination token index accounts for this rank's offset in the global sequence
int64_t dst_token_idx = prefix_rank_tokens[rank] + copy_token_idx;
if (copy_token_idx >= num_tokens)
break;
// === Pointer Arithmetic ===
// Source: Read from this rank's input tensor at [batch, token, head_id:head_id+hpr, :]
// Note: We read a contiguous chunk of heads starting at head_id
int4 *shuffled_x_ptr =
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(x) +
batch_ind * num_tokens_64b * num_heads_64b * head_size_64b * sizeof(ELEM_T) +
copy_token_idx * num_heads_64b * head_size_64b * sizeof(ELEM_T) +
head_id * head_size_64b * sizeof(ELEM_T)) +
copy_thr_col_idx;
// Destination: Write to target rank's buffer at [batch, dst_token, :, :]
// The buffer has layout [batch, total_tokens, heads_per_rank, head_size]
int4 *shuffled_buffer_ptr =
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(ptr) +
batch_ind * total_tokens * heads_per_rank * head_size_64b * sizeof(ELEM_T) +
dst_token_idx * heads_per_rank * head_size_64b * sizeof(ELEM_T)) +
copy_thr_col_idx;
// Non-allocating store to avoid polluting L1 cache
st_na_global(shuffled_buffer_ptr, __ldg(shuffled_x_ptr));
}
}
// === Barrier Synchronization ===
// Signal completion to target rank and wait for all ranks to finish
barrier_wait_and_reset_roundrobin(barrier_signal_ptrs, target_rank, rank, world_size, num_sms, sm_id, threadIdx.x,
timeout_cycles);
}
/**
* @brief All2All kernel for gathering attention heads back to original distribution.
*
* This kernel performs the inverse of send_recv_all2all: it gathers heads from
* all ranks back to reconstruct the original tensor layout. Each GPU reads from
* its local buffer and writes its portion of heads to all target ranks.
*
* Data layout transformation:
* Source: [batch, total_tokens, heads_per_rank, head_size] (per GPU)
* Dest: [batch, rank_tokens[target], num_heads, head_size] (per target GPU)
*
* ## Memory Layout
*
* Input tensor x (this rank's portion after send_recv_all2all):
* - Contains all tokens but only heads_per_rank heads
* - Layout: [batch, total_tokens, heads_per_rank, head_size]
*
* Output buffer (per target rank):
* - Contains only that rank's tokens but all heads
* - Layout: [batch, rank_tokens[target], num_heads, head_size]
* - This rank writes heads [rank * heads_per_rank : (rank+1) * heads_per_rank]
*
* @tparam ELEM_T Element type (at::BFloat16 or at::Float8_e4m3fn)
* @param buffer_ptrs Device array of pointers to each rank's data buffer
* @param barrier_signal_ptrs Device array of pointers to barrier signals
* @param x Source tensor data (this rank's buffer after send_recv)
* @param rank This GPU's rank
* @param world_size Total number of GPUs
* @param batch_size Number of batches
* @param num_heads Total number of heads (reconstructed)
* @param head_size Size of each attention head
* @param rank_tokens Number of tokens for each rank
* @param total_tokens Sum of tokens across all ranks
* @param prefix_rank_tokens Cumulative token counts for offset calculation
*/
template <typename ELEM_T>
__global__ void gather_heads(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int rank, int world_size,
int batch_size, int num_heads, int head_size, const int *__restrict__ rank_tokens,
int total_tokens, int *prefix_rank_tokens, uint64_t timeout_cycles) {
// Grid dimensions
int num_sms = gridDim.x;
int sm_id = blockIdx.x;
int num_threads = blockDim.x;
// === SM Work Distribution (Round-Robin) ===
// Same partitioning as send_recv_all2all
int64_t target_rank = get_target_rank(sm_id, world_size);
int64_t rank_local_sm_id = get_rank_local_sm_id(sm_id, world_size);
int64_t num_sms_for_this_rank = get_num_sms_for_rank(target_rank, num_sms, world_size);
int64_t heads_per_rank = num_heads / world_size;
// === Thread Mapping ===
int64_t num_elems_per_thread = sizeof(int4) / sizeof(ELEM_T);
int64_t num_threads_per_token = heads_per_rank * head_size / num_elems_per_thread;
int64_t num_tokens_per_copy = num_threads / num_threads_per_token;
int64_t copy_thr_col_idx = threadIdx.x % num_threads_per_token;
int64_t copy_thr_row_idx = threadIdx.x / num_threads_per_token;
// Number of tokens owned by target rank
const int64_t tgt_tokens = int64_t(rank_tokens[target_rank]);
// This rank writes its heads at offset [rank * heads_per_rank] in the output
int64_t head_idx = rank * heads_per_rank;
int64_t num_heads_64b = int64_t(num_heads);
int64_t head_size_64b = int64_t(head_size);
int64_t total_tokens_64b = int64_t(total_tokens);
// Get target rank's buffer pointer
auto ptr = reinterpret_cast<void *>(static_cast<int8_t *>(buffer_ptrs[target_rank]));
// === Main Copy Loop ===
// Process target rank's tokens: read from global position, write to local position
for (int64_t batch_idx = 0; batch_idx < batch_size; batch_idx++) {
for (int64_t token_idx = rank_local_sm_id * num_tokens_per_copy; token_idx < tgt_tokens;
token_idx += num_tokens_per_copy * num_sms_for_this_rank) {
int64_t copy_token = token_idx + copy_thr_row_idx;
if (copy_token >= tgt_tokens)
break;
// Source: Read from global token position (target rank's tokens in our buffer)
int64_t src_token_idx = prefix_rank_tokens[target_rank] + copy_token;
// Destination: Write to local token position in target's buffer
int64_t dst_token_idx = copy_token;
// Source pointer: our input tensor at [batch, src_token, :, :]
int4 *shuffled_x_ptr =
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(x) +
batch_idx * total_tokens_64b * heads_per_rank * head_size_64b * sizeof(ELEM_T) +
src_token_idx * heads_per_rank * head_size_64b * sizeof(ELEM_T)) +
copy_thr_col_idx;
// Destination pointer: target's buffer at [batch, dst_token, head_idx:head_idx+hpr, :]
int4 *shuffled_buffer_ptr =
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(ptr) +
batch_idx * tgt_tokens * num_heads_64b * head_size_64b * sizeof(ELEM_T) +
dst_token_idx * num_heads_64b * head_size_64b * sizeof(ELEM_T) +
head_idx * head_size_64b * sizeof(ELEM_T)) +
copy_thr_col_idx;
st_na_global(shuffled_buffer_ptr, __ldg(shuffled_x_ptr));
}
}
// === Barrier Synchronization ===
barrier_wait_and_reset_roundrobin(barrier_signal_ptrs, target_rank, rank, world_size, num_sms, sm_id, threadIdx.x,
timeout_cycles);
}
/**
* @brief Host function to launch the gather_heads kernel.
*
* Selects the appropriate template instantiation based on tensor data type
* and launches the kernel with the specified number of SMs.
*
* @param buffer_ptrs Device array of buffer pointers
* @param barrier_signal_ptrs Device array of barrier signal pointers
* @param x Input tensor data pointer
* @param rank_tokens Token count per rank (device memory)
* @param prefix_rank_tokens Cumulative token counts (device memory)
* @param rank This GPU's rank
* @param world_size Total number of GPUs
* @param batch_size Number of batches
* @param total_tokens Sum of tokens across all ranks
* @param num_heads Total number of attention heads
* @param head_size Size of each attention head
* @param stream CUDA stream for async execution
* @param num_sms Number of SMs to launch
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
*/
void all2all_head_gather_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, const int *rank_tokens,
int *prefix_rank_tokens, int rank, int world_size, int batch_size, int total_tokens,
int num_heads, int head_size, cudaStream_t stream, int num_sms,
at::ScalarType tensor_dtype, uint64_t timeout_cycles) {
do {
if (tensor_dtype == at::ScalarType::BFloat16) {
gather_heads<at::BFloat16><<<num_sms, DEFAULT_KERNEL_THREADS, 0, stream>>>(
buffer_ptrs, barrier_signal_ptrs, x, rank, world_size, batch_size, num_heads, head_size, rank_tokens,
total_tokens, prefix_rank_tokens, timeout_cycles);
} else if (tensor_dtype == at::ScalarType::Float8_e4m3fn) {
gather_heads<at::Float8_e4m3fn><<<num_sms, DEFAULT_KERNEL_THREADS, 0, stream>>>(
buffer_ptrs, barrier_signal_ptrs, x, rank, world_size, batch_size, num_heads, head_size, rank_tokens,
total_tokens, prefix_rank_tokens, timeout_cycles);
}
// Check for kernel launch errors
cudaError_t e = cudaGetLastError();
if (e != cudaSuccess) {
EPException cuda_exception("CUDA", __FILE__, __LINE__, cudaGetErrorString(e));
fprintf(stderr, "%s\n", cuda_exception.what());
throw cuda_exception;
}
} while (0);
}
/**
* @brief Host function to launch the send_recv_all2all kernel.
*
* Selects the appropriate template instantiation based on tensor data type
* and launches the kernel with the specified number of SMs.
*
* @param buffer_ptrs Device array of buffer pointers
* @param barrier_signal_ptrs Device array of barrier signal pointers
* @param x Input tensor data pointer
* @param prefix_rank_tokens Cumulative token counts (device memory)
* @param rank This GPU's rank
* @param world_size Total number of GPUs
* @param batch_size Number of batches
* @param total_tokens Sum of tokens across all ranks
* @param num_tokens Number of tokens on this rank
* @param num_heads Total number of attention heads
* @param head_size Size of each attention head
* @param stream CUDA stream for async execution
* @param num_sms Number of SMs to launch
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
*/
void all2all_head_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int *prefix_rank_tokens, int rank,
int world_size, int batch_size, int total_tokens, int num_tokens, int num_heads, int head_size,
cudaStream_t stream, int num_sms, at::ScalarType tensor_dtype, uint64_t timeout_cycles) {
do {
if (tensor_dtype == at::ScalarType::BFloat16) {
send_recv_all2all<at::BFloat16><<<num_sms, DEFAULT_KERNEL_THREADS, 0, stream>>>(
buffer_ptrs, barrier_signal_ptrs, x, rank, world_size, batch_size, num_tokens, num_heads, head_size,
total_tokens, prefix_rank_tokens, timeout_cycles);
} else if (tensor_dtype == at::ScalarType::Float8_e4m3fn) {
send_recv_all2all<at::Float8_e4m3fn><<<num_sms, DEFAULT_KERNEL_THREADS, 0, stream>>>(
buffer_ptrs, barrier_signal_ptrs, x, rank, world_size, batch_size, num_tokens, num_heads, head_size,
total_tokens, prefix_rank_tokens, timeout_cycles);
}
// Check for kernel launch errors
cudaError_t e = cudaGetLastError();
if (e != cudaSuccess) {
EPException cuda_exception("CUDA", __FILE__, __LINE__, cudaGetErrorString(e));
fprintf(stderr, "%s\n", cuda_exception.what());
throw cuda_exception;
}
} while (0);
}
} // namespace all2all_cuda
} // namespace all2all
} // namespace ltx_kernels
@@ -0,0 +1,198 @@
/**
* @file allgather.cu
* @brief CUDA kernel for AllGather operation using IPC-based direct memory access.
*
* This file implements the GPU kernel for gathering sequence tokens from all GPUs
* into a complete sequence on each GPU. Unlike the head redistribution kernels,
* this kernel preserves the head dimension and only gathers across the token
* (sequence) dimension.
*
* ## Algorithm Overview
*
* Each GPU broadcasts its local tokens to all other GPUs' buffers:
* - GPU i writes its tokens to position [prefix_rank_tokens[i]] in each buffer
* - After completion, all buffers contain the full sequence [0:total_tokens]
*
* ## Use Case
*
* This is typically used after tensor-parallel computation to reconstruct the
* full sequence for operations that require global context (e.g., output projection).
*/
#include "cuda/configs.cuh"
#include "cuda/exceptions.cuh"
#include "cuda/utils.cuh"
#include <ATen/cuda/CUDADataType.h>
namespace ltx_kernels {
namespace all2all {
namespace all2all_cuda {
/**
* @brief AllGather kernel to collect sequence tokens from all ranks.
*
* Each GPU writes its local sequence tokens to all other GPUs' buffers at the
* appropriate offset. After synchronization, all GPUs have the complete sequence.
*
* Data layout transformation:
* Input per GPU: [batch, seqlen, hidden_dim]
* Output per GPU: [batch, total_tokens, hidden_dim] (identical on all GPUs)
*
* ## Memory Layout
*
* Input tensor x (contiguous):
* - Shape: [batch, seqlen, hidden_dim]
* - hidden_dim = num_heads * head_size (flattened)
*
* Output buffer (per target rank, after gather):
* - Shape: [batch, total_tokens, hidden_dim]
* - This rank's tokens placed at offset rank_tokens_prefix[rank]
*
* ## Thread Mapping
*
* Similar to all2all_heads, threads cooperate to copy tokens:
* - Each thread copies 16 bytes (int4)
* - Threads per token = hidden_dim * sizeof(ELEM_T) / sizeof(int4)
* - Multiple tokens processed per thread block
*
* @tparam ELEM_T Element type (__nv_bfloat16 or at::Float8_e4m3fn)
* @param x Source tensor data pointer (this rank's tokens)
* @param buffer_ptrs Device array of pointers to each rank's data buffer
* @param barrier_signal_ptrs Device array of pointers to barrier signals
* @param batch_size Number of batches
* @param seqlen Number of tokens on this rank
* @param hidden_dim Hidden dimension size (num_heads * head_size)
* @param world_size Total number of GPUs
* @param rank This GPU's rank
* @param total_tokens Sum of tokens across all ranks
* @param rank_tokens_prefix Cumulative token counts (device memory)
*/
template <typename ELEM_T>
__global__ void allgather(void *x, void **buffer_ptrs, int **barrier_signal_ptrs, int batch_size, int seqlen,
int hidden_dim, int world_size, int rank, int total_tokens, int *rank_tokens_prefix,
uint64_t timeout_cycles) {
// Grid dimensions
int num_sms = gridDim.x;
int sm_id = blockIdx.x;
int num_threads = blockDim.x;
// === SM Work Distribution (Round-Robin) ===
// Use modular assignment to handle num_sms not divisible by world_size.
// This ensures all SMs are utilized: some ranks get ceil(num_sms/world_size)
// SMs, others get floor(num_sms/world_size) SMs.
int tgt_rank = get_target_rank(sm_id, world_size);
int rank_local_sm_id = get_rank_local_sm_id(sm_id, world_size);
int num_sms_for_this_rank = get_num_sms_for_rank(tgt_rank, num_sms, world_size);
// Get target rank's buffer pointer
auto ptr = reinterpret_cast<void *>(static_cast<int8_t *>(buffer_ptrs[tgt_rank]));
// === Thread Mapping ===
// Each thread copies one int4 (16 bytes)
int64_t num_elems_per_thread = sizeof(int4) / sizeof(ELEM_T);
int64_t num_threads_per_token = hidden_dim / num_elems_per_thread;
int64_t num_tokens_per_copy = num_threads / num_threads_per_token;
// 2D thread coordinates
int64_t copy_thr_col_idx = threadIdx.x % num_threads_per_token; // Element offset
int64_t copy_thr_row_idx = threadIdx.x / num_threads_per_token; // Token offset
// Use 64-bit arithmetic to avoid overflow
int64_t hidden_dim_64b = int64_t(hidden_dim);
int64_t total_tokens_64b = int64_t(total_tokens);
int64_t seqlen_64b = int64_t(seqlen);
// === Main Copy Loop ===
// Broadcast this rank's tokens to all target ranks' buffers
for (int64_t batch_idx = 0; batch_idx < batch_size; batch_idx++) {
// Strided token iteration within SM group for this target rank
for (int64_t token_idx = rank_local_sm_id * num_tokens_per_copy; token_idx < seqlen;
token_idx += num_tokens_per_copy * num_sms_for_this_rank) {
int64_t copy_token = token_idx + copy_thr_row_idx;
if (copy_token >= seqlen)
break;
// Source: local token index in input tensor
int64_t src_token_idx = copy_token;
// Destination: global token index in output buffer
// This rank's tokens start at prefix_rank_tokens[rank]
int64_t dst_token_idx = copy_token + rank_tokens_prefix[rank];
// Source pointer: input tensor at [batch, src_token, :]
int4 *shuffled_x_ptr = reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(x) +
batch_idx * seqlen_64b * hidden_dim_64b * sizeof(ELEM_T) +
src_token_idx * hidden_dim_64b * sizeof(ELEM_T)) +
copy_thr_col_idx;
// Destination pointer: target buffer at [batch, dst_token, :]
int4 *shuffled_buffer_ptr =
reinterpret_cast<int4 *>(reinterpret_cast<uint8_t *>(ptr) +
batch_idx * total_tokens_64b * hidden_dim_64b * sizeof(ELEM_T) +
dst_token_idx * hidden_dim_64b * sizeof(ELEM_T)) +
copy_thr_col_idx;
// Non-allocating store for better cache behavior
st_na_global(shuffled_buffer_ptr, __ldg(shuffled_x_ptr));
}
}
// === Barrier Synchronization ===
// Signal completion to target rank and wait for all ranks
// Use round-robin variant since SM counts per rank may differ
barrier_wait_and_reset_roundrobin(barrier_signal_ptrs, tgt_rank, rank, world_size, num_sms, sm_id, threadIdx.x,
timeout_cycles);
}
/**
* @brief Host function to launch the allgather kernel.
*
* Launches the AllGather kernel with the specified configuration.
* Uses ALLGATHER_KERNEL_THREADS (1024) threads per block for higher
* occupancy than the All2All kernels.
*
* @param buffer_ptrs Device array of buffer pointers
* @param barrier_signal_ptrs Device array of barrier signal pointers
* @param x Input tensor data pointer
* @param prefix_rank_tokens Cumulative token counts (device memory)
* @param rank This GPU's rank
* @param world_size Total number of GPUs
* @param batch_size Number of batches
* @param seqlen Number of tokens on this rank
* @param hidden_dim Hidden dimension size
* @param total_tokens Sum of tokens across all ranks
* @param stream CUDA stream for async execution
* @param num_sms Number of SMs to launch
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
*/
void allgather_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int *prefix_rank_tokens, int rank,
int world_size, int batch_size, int seqlen, int hidden_dim, int total_tokens, cudaStream_t stream,
int num_sms, at::ScalarType tensor_dtype, uint64_t timeout_cycles) {
do {
if (tensor_dtype == at::ScalarType::BFloat16) {
allgather<at::BFloat16><<<num_sms, ALLGATHER_KERNEL_THREADS, 0, stream>>>(
x, buffer_ptrs, barrier_signal_ptrs, batch_size, seqlen, hidden_dim, world_size, rank, total_tokens,
prefix_rank_tokens, timeout_cycles);
} else if (tensor_dtype == at::ScalarType::Float8_e4m3fn) {
allgather<at::Float8_e4m3fn><<<num_sms, ALLGATHER_KERNEL_THREADS, 0, stream>>>(
x, buffer_ptrs, barrier_signal_ptrs, batch_size, seqlen, hidden_dim, world_size, rank, total_tokens,
prefix_rank_tokens, timeout_cycles);
} else {
EPException dtype_exception("allgather_launch", __FILE__, __LINE__, "Unsupported dtype");
fprintf(stderr, "%s\n", dtype_exception.what());
throw dtype_exception;
}
// Check for kernel launch errors
cudaError_t e = cudaGetLastError();
if (e != cudaSuccess) {
EPException cuda_exception("CUDA", __FILE__, __LINE__, cudaGetErrorString(e));
fprintf(stderr, "%s\n", cuda_exception.what());
throw cuda_exception;
}
} while (0);
}
} // namespace all2all_cuda
} // namespace all2all
} // namespace ltx_kernels
@@ -0,0 +1,99 @@
/**
* @file api.cuh
* @brief CUDA kernel launch function declarations for All2All operations.
*
* This header provides the host-callable interface for launching the All2All
* CUDA kernels. These functions handle template instantiation and kernel
* configuration based on the tensor data type.
*/
#pragma once
#include <ATen/cuda/CUDADataType.h>
#include <vector>
namespace ltx_kernels {
namespace all2all {
namespace all2all_cuda {
/**
* @brief Launches the All2All head redistribution kernel.
*
* Redistributes attention heads across GPUs:
* Input: [batch, num_tokens, num_heads, head_size] per GPU
* Output: [batch, total_tokens, num_heads/world_size, head_size] per GPU
*
* @param buffer_ptrs Device array of pointers to each rank's data buffer
* @param barrier_signal_ptrs Device array of pointers to barrier signals
* @param x Source tensor data pointer
* @param prefix_rank_tokens Cumulative token counts per rank (device memory)
* @param rank This GPU's rank (0 to world_size-1)
* @param world_size Total number of GPUs
* @param batch_size Batch dimension size
* @param total_tokens Sum of tokens across all ranks
* @param num_tokens Number of tokens on this rank
* @param num_heads Total number of attention heads
* @param head_size Size of each attention head
* @param stream CUDA stream for async execution
* @param num_sms Number of SMs to use for the kernel
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
*/
void all2all_head_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int *prefix_rank_tokens, int rank,
int world_size, int batch_size, int total_tokens, int num_tokens, int num_heads, int head_size,
cudaStream_t stream, int num_sms, at::ScalarType tensor_dtype, uint64_t timeout_cycles);
/**
* @brief Launches the gather heads kernel (inverse of all2all_head_launch).
*
* Redistributes tokens back to original head distribution:
* Input: [batch, total_tokens, heads_per_rank, head_size] per GPU
* Output: [batch, rank_tokens[rank], num_heads, head_size] per GPU
*
* @param buffer_ptrs Device array of pointers to each rank's data buffer
* @param barrier_signal_ptrs Device array of pointers to barrier signals
* @param x Source tensor data pointer
* @param rank_tokens Token count for each rank (device memory)
* @param prefix_rank_tokens Cumulative token counts (device memory)
* @param rank This GPU's rank
* @param world_size Total number of GPUs
* @param batch_size Batch dimension size
* @param total_tokens Sum of tokens across all ranks
* @param num_heads Total number of attention heads (reconstructed)
* @param head_size Size of each attention head
* @param stream CUDA stream for async execution
* @param num_sms Number of SMs to use for the kernel
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
*/
void all2all_head_gather_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, const int *rank_tokens,
int *prefix_rank_tokens, int rank, int world_size, int batch_size, int total_tokens,
int num_heads, int head_size, cudaStream_t stream, int num_sms,
at::ScalarType tensor_dtype, uint64_t timeout_cycles);
/**
* @brief Launches the AllGather kernel for sequence tokens.
*
* Gathers sequence tokens from all ranks:
* Input: [batch, seqlen, hidden_dim] per GPU
* Output: [batch, total_tokens, hidden_dim] per GPU (identical on all)
*
* @param buffer_ptrs Device array of pointers to each rank's data buffer
* @param barrier_signal_ptrs Device array of pointers to barrier signals
* @param x Source tensor data pointer
* @param prefix_rank_tokens Cumulative token counts (device memory)
* @param rank This GPU's rank
* @param world_size Total number of GPUs
* @param batch_size Batch dimension size
* @param seqlen Number of tokens on this rank
* @param hidden_dim Hidden dimension size (num_heads * head_size)
* @param total_tokens Sum of tokens across all ranks
* @param stream CUDA stream for async execution
* @param num_sms Number of SMs to use for the kernel
* @param tensor_dtype Data type (BFloat16 or Float8_e4m3fn)
*/
void allgather_launch(void **buffer_ptrs, int **barrier_signal_ptrs, void *x, int *prefix_rank_tokens, int rank,
int world_size, int batch_size, int seqlen, int hidden_dim, int total_tokens, cudaStream_t stream,
int num_sms, at::ScalarType tensor_dtype, uint64_t timeout_cycles);
} // namespace all2all_cuda
} // namespace all2all
} // namespace ltx_kernels
@@ -0,0 +1,89 @@
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <torch/extension.h>
#include <vector>
#include <stdio.h>
#ifdef __SM90__
#include "sm90_fp8_gemm_1d2d_bias.hpp"
#endif
#include "sm89_fp8_gemm_1d2d.hpp"
namespace blockwise{
template <int N>
static auto get_shape(const torch::Tensor& t) {
return [&t] <size_t... Is> (std::index_sequence<Is...>) {
return std::make_tuple(static_cast<int>(t.sizes()[Is])...);
}(std::make_index_sequence<N>());
}
#ifdef __SM90__
static void fp8_gemm_nt_sm90(const std::pair<torch::Tensor, torch::Tensor>& a,
const std::pair<torch::Tensor, torch::Tensor>& b,
const torch::Tensor& d,
const std::optional<torch::Tensor>& bias,
const std::optional<torch::Tensor>& c, const int num_sms) {
// Type and shape checks
const auto& [m , k ] = get_shape<2>(a.first);
const auto& [n , k_] = get_shape<2>(b.first);
const auto& [m_, n_] = get_shape<2>(d);
// The SM90 kernel always adds bias; synthesize a zero bias when the layer is
// bias-less (e.g. the no-bias video FFN of v3 checkpoints), mirroring SM89 below.
torch::Tensor bias_tensor = bias.has_value()
? bias.value()
: torch::zeros({n}, d.options().dtype(torch::kFloat32));
sm90_fp8_gemm_1d2d_bias(a.first, a.second, b.first, b.second, bias_tensor, c, d, m, n, k, num_sms);
}
#endif
static void fp8_gemm_nt_sm89(const std::pair<torch::Tensor, torch::Tensor>& a,
const std::pair<torch::Tensor, torch::Tensor>& b,
const torch::Tensor& d,
const std::optional<torch::Tensor>& bias,
const bool use_fast_accum = true) {
const auto& [m, k] = get_shape<2>(a.first);
const auto& [n, k_] = get_shape<2>(b.first);
const auto& [m_, n_] = get_shape<2>(d);
// The SM89 kernel always adds bias; synthesize a zero bias when the layer is
// bias-less so we add 0 rather than uninitialized memory (mirrors SM90 above).
torch::Tensor bias_tensor = bias.has_value()
? bias.value()
: torch::zeros({n}, d.options().dtype(torch::kFloat32));
blockwise::sm89_fp8_gemm_1d2d_bias(
a.first, a.second, // a data, sfa scales
b.first, b.second, // b data, sfb scales
bias_tensor, // bias (or empty tensor)
d, // output
m, n, k,
use_fast_accum); // pass through accumulation mode
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
// m.def("package_name", &function_name, "function_docstring"")
#ifdef __SM90__
m.def("fp8_gemm_nt_sm90", &fp8_gemm_nt_sm90,
py::arg("a"), py::arg("b"), py::arg("d"),
py::arg("bias") = std::nullopt,
py::arg("c") = std::nullopt,
py::arg("num_sms") = 132
);
#endif
m.def("fp8_gemm_nt_sm89", &fp8_gemm_nt_sm89,
py::arg("a"), py::arg("b"), py::arg("d"),
py::arg("bias") = std::nullopt,
py::arg("use_fast_accum") = true
);
}
};
@@ -0,0 +1,92 @@
#pragma once
#include <torch/python.h>
#include <cute/arch/mma_sm100_umma.hpp>
#include "utils.hpp"
#include "exceptions.hpp"
namespace blockwise{
struct MulticastConfig {
int num_multicast;
bool is_multicast_on_a;
MulticastConfig(const int& num_multicast, const bool& is_multicast_on_a):
num_multicast(num_multicast), is_multicast_on_a(is_multicast_on_a) {
DG_HOST_ASSERT(1 <= num_multicast and num_multicast <= 2);
}
};
struct SharedMemoryConfig {
int smem_size;
int swizzle_a_mode;
int swizzle_b_mode;
int swizzle_cd_mode;
};
struct ThreadConfig {
int num_threads;
// SM90
int num_tma_threads;
int num_math_threads;
// SM100
int num_non_epilogue_threads;
int num_epilogue_threads;
static ThreadConfig sm90(const int& num_tma_threads,
const int& num_math_threads) {
auto config = ThreadConfig();
config.num_threads = num_tma_threads + num_math_threads;
config.num_tma_threads = num_tma_threads;
config.num_math_threads = num_math_threads;
return config;
}
static ThreadConfig sm100(const int& num_non_epilogue_threads,
const int& num_epilogue_threads) {
auto config = ThreadConfig();
config.num_threads = num_non_epilogue_threads + num_epilogue_threads;
config.num_non_epilogue_threads = num_non_epilogue_threads;
config.num_epilogue_threads = num_epilogue_threads;
return config;
}
};
template<int SM>
struct GemmConfig{};
// {
// // Templated configs
// at::ScalarType ab_dtype, cd_dtype;
// bool with_accumulation;
// int block_m, block_n, block_k;
// int num_stages, num_last_stages;
// // Templated device configs
// int num_sms;
// // Structured configs
// MulticastConfig multicast_config;
// SharedMemoryConfig smem_config;
// ThreadConfig thread_config;
// };
template <>
struct GemmConfig<90>
{
at::ScalarType ab_dtype = torch::kFloat8_e4m3fn;
at::ScalarType cd_dtype = torch::kBFloat16;
bool with_accumulation = false;
int block_m = 256;
int block_n = 128;
int block_k = 128;
int num_stages = 3;
int num_last_stages = 2;
int num_sms = 132;
MulticastConfig multicast_config{2, true};
SharedMemoryConfig smem_config{216240, 128, 128, 128};
ThreadConfig thread_config = ThreadConfig::sm90(128, 256);
};
};
@@ -0,0 +1,65 @@
#pragma once
#include <exception>
#include <string>
#include <sstream>
namespace blockwise {
class DGException final : public std::exception {
std::string message = {};
public:
explicit DGException(const char *name, const char* file, const int line, const std::string& error) {
message = std::string(name) + " error (" + file + ":" + std::to_string(line) + "): " + error;
}
const char *what() const noexcept override {
return message.c_str();
}
};
#ifndef DG_STATIC_ASSERT
#define DG_STATIC_ASSERT(cond, ...) static_assert(cond, __VA_ARGS__)
#endif
#ifndef DG_HOST_ASSERT
#define DG_HOST_ASSERT(cond) \
do { \
if (not (cond)) { \
throw DGException("Assertion", __FILE__, __LINE__, #cond); \
} \
} while (0)
#endif
#ifndef DG_HOST_UNREACHABLE
#define DG_HOST_UNREACHABLE(reason) (throw DGException("Assertion", __FILE__, __LINE__, reason))
#endif
// #ifndef DG_CUDA_DRIVER_CHECK
// #define DG_CUDA_DRIVER_CHECK(cmd) \
// do { \
// const auto& e = (cmd); \
// if (e != CUDA_SUCCESS) { \
// std::stringstream ss; \
// const char *name, *info; \
// cuGetErrorName(e, &name), cuGetErrorString(e, &info); \
// ss << static_cast<int>(e) << " (" << name << ", " << info << ")"; \
// throw DGException("CUDA driver", __FILE__, __LINE__, ss.str()); \
// } \
// } while (0)
// #endif
#ifndef DG_CUDA_RUNTIME_CHECK
#define DG_CUDA_RUNTIME_CHECK(cmd) \
do { \
const auto& e = (cmd); \
if (e != cudaSuccess) { \
std::stringstream ss; \
ss << static_cast<int>(e) << " (" << cudaGetErrorName(e) << ", " << cudaGetErrorString(e) << ")"; \
throw DGException("CUDA runtime", __FILE__, __LINE__, ss.str()); \
} \
} while (0)
#endif
} // namespace deep_gemm
@@ -0,0 +1,48 @@
#pragma once
namespace cute {
struct ignore_t {
template <typename T>
constexpr const ignore_t& operator=(T&&) const noexcept {
return *this;
}
};
inline constexpr ignore_t ignore{};
} // namespace cute
#define CUTE_TIE_CONCAT_IMPL(A, B) A##B
#define CUTE_TIE_CONCAT(A, B) CUTE_TIE_CONCAT_IMPL(A, B)
#define CUTE_TIE_GET_NTH_ARG(_1, _2, _3, _4, _5, _6, _7, _8, _9, _10, N, ...) N
#define CUTE_TIE_COUNT_ARGS(...) \
CUTE_TIE_GET_NTH_ARG(__VA_ARGS__, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0)
#define CUTE_TIE_OP_DECL(I, TUPLE, VAR) auto VAR = ::cute::get<I>(TUPLE)
#define CUTE_TIE_OP_ASSIGN(I, TUPLE, VAR) VAR = ::cute::get<I>(TUPLE)
#define CUTE_TIE_APPLY_OP_1(OP, T, V1) OP(0, T, V1);
#define CUTE_TIE_APPLY_OP_2(OP, T, V1, V2) OP(0, T, V1); OP(1, T, V2);
#define CUTE_TIE_APPLY_OP_3(OP, T, V1, V2, V3) OP(0, T, V1); OP(1, T, V2); OP(2, T, V3);
#define CUTE_TIE_APPLY_OP_4(OP, T, V1, V2, V3, V4) OP(0, T, V1); OP(1, T, V2); OP(2, T, V3); OP(3, T, V4);
#define CUTE_TIE_APPLY_OP_5(OP, T, V1, V2, V3, V4, V5) OP(0, T, V1); OP(1, T, V2); OP(2, T, V3); OP(3, T, V4); OP(4, T, V5);
#define CUTE_TIE_DECL(TUPLE_EXPR, ...) \
auto&& CUTE_TIE_CONCAT(cute_tie__temp_tuple_, __LINE__) = (TUPLE_EXPR); \
CUTE_TIE_CONCAT(CUTE_TIE_APPLY_OP_, CUTE_TIE_COUNT_ARGS(__VA_ARGS__)) ( \
CUTE_TIE_OP_DECL, \
CUTE_TIE_CONCAT(cute_tie__temp_tuple_, __LINE__), \
__VA_ARGS__ \
)
#define CUTE_TIE(TUPLE_EXPR, ...) \
do { \
auto&& CUTE_TIE_CONCAT(cute_tie__temp_tuple_, __LINE__) = (TUPLE_EXPR); \
CUTE_TIE_CONCAT(CUTE_TIE_APPLY_OP_, CUTE_TIE_COUNT_ARGS(__VA_ARGS__)) ( \
CUTE_TIE_OP_ASSIGN, \
CUTE_TIE_CONCAT(cute_tie__temp_tuple_, __LINE__), \
__VA_ARGS__ \
); \
} while (0)
@@ -0,0 +1,27 @@
#pragma once
#include <deep_gemm/common/types.hpp>
#include <deep_gemm/common/utils.cuh>
namespace deep_gemm {
struct EpilogueIdentity {
template <uint32_t STORE_BLOCK_N>
__device__ __forceinline__ static uint32_t apply_index_n(const uint32_t &n_idx) {
return n_idx;
}
};
template <uint32_t kLeft, uint32_t kMid, uint32_t kRight>
struct EpilogueHeadSplits: EpilogueIdentity {
template <uint32_t STORE_BLOCK_N>
__device__ __forceinline__ static uint32_t apply_index_n(const uint32_t &n_idx) {
DG_STATIC_ASSERT(kLeft % STORE_BLOCK_N == 0 and kMid % STORE_BLOCK_N == 0
and kRight % STORE_BLOCK_N == 0, "Invalid head splits config");
return n_idx + (n_idx + kRight) / (kLeft + kRight) * kMid;
}
};
#pragma clang diagnostic pop
} // namespace deep_gemm
@@ -0,0 +1,44 @@
#pragma once
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda/std/cstdint>
#include <cuda/std/utility>
#include <deep_gemm/common/utils.cuh>
// Operation functors
template <typename T> struct ReduceSum { __device__ T operator()(T a, T b) const { return a + b; } };
template <typename T> struct ReduceMax { __device__ T operator()(T a, T b) const { return a > b ? a : b; } };
template <typename T> struct ReduceMin { __device__ T operator()(T a, T b) const { return a < b ? a : b; } };
template <typename T> struct ReduceAnd { __device__ T operator()(T a, T b) const { return a & b; } };
template <typename T> struct ReduceOr { __device__ T operator()(T a, T b) const { return a | b; } };
// Unified reduction function
template <int kNumLanesPerGroup, bool kIntergroupReduce, typename T, typename Op>
__forceinline__ __device__ T warp_reduce(T value, Op op) {
DG_STATIC_ASSERT(kNumLanesPerGroup == 32 or kNumLanesPerGroup == 16 or kNumLanesPerGroup == 8 or
kNumLanesPerGroup == 4 or kNumLanesPerGroup == 2 or kNumLanesPerGroup == 1,
"Invalid number of lanes");
constexpr uint32_t mask = 0xffffffff;
if constexpr (kIntergroupReduce) {
if constexpr (kNumLanesPerGroup <= 1) value = op(value, __shfl_xor_sync(mask, value, 1));
if constexpr (kNumLanesPerGroup <= 2) value = op(value, __shfl_xor_sync(mask, value, 2));
if constexpr (kNumLanesPerGroup <= 4) value = op(value, __shfl_xor_sync(mask, value, 4));
if constexpr (kNumLanesPerGroup <= 8) value = op(value, __shfl_xor_sync(mask, value, 8));
if constexpr (kNumLanesPerGroup <= 16) value = op(value, __shfl_xor_sync(mask, value, 16));
} else {
if constexpr (kNumLanesPerGroup >= 32) value = op(value, __shfl_xor_sync(mask, value, 16));
if constexpr (kNumLanesPerGroup >= 16) value = op(value, __shfl_xor_sync(mask, value, 8));
if constexpr (kNumLanesPerGroup >= 8) value = op(value, __shfl_xor_sync(mask, value, 4));
if constexpr (kNumLanesPerGroup >= 4) value = op(value, __shfl_xor_sync(mask, value, 2));
if constexpr (kNumLanesPerGroup >= 2) value = op(value, __shfl_xor_sync(mask, value, 1));
}
return value;
}
// Convenience aliases
template <int kNumLanesPerGroup = 32, bool kIntergroupReduce = false, typename T>
__forceinline__ __device__ T warp_reduce_sum(T value) {
return warp_reduce<kNumLanesPerGroup, kIntergroupReduce, T>(value, ReduceSum<T>{});
}
@@ -0,0 +1,239 @@
#pragma once
#include <deep_gemm/common/types.hpp>
#include <deep_gemm/common/utils.cuh>
namespace deep_gemm {
enum class KGroupedIndexType {
MN,
K,
SF_K,
};
template <GemmType kGemmType, uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t kNumSMs, bool kIsMulticastOnA>
static constexpr uint32_t get_num_1d_blocks_per_group() {
// Select the best from candidates
uint32_t num_best_blocks = 0, min_usage = cute::numeric_limits<uint32_t>::max();
for (const auto& candidate: {8u, 16u}) {
const auto& usage = kIsMulticastOnA ?
candidate * BLOCK_N + constexpr_ceil_div(kNumSMs, candidate) * BLOCK_M: // Grouping on N
candidate * BLOCK_M + constexpr_ceil_div(kNumSMs, candidate) * BLOCK_N; // Grouping on M
if (usage < min_usage)
min_usage = usage, num_best_blocks = candidate;
}
return num_best_blocks;
}
#pragma clang diagnostic push
#pragma ide diagnostic ignored "cppcoreguidelines-pro-type-member-init"
template <GemmType kGemmType,
uint32_t BLOCK_M, uint32_t BLOCK_N,
uint32_t kNumGroups,
uint32_t kNumMulticast, bool kIsMulticastOnA,
uint32_t kNumSMs,
uint32_t SF_K_ALIGNMENT = 512u, // for k-grouped GEMM only: 128 (SM90 float SF) or 512 (SM100 UE8M0 SF)
uint32_t kNum1DBlocksPerGroup = get_num_1d_blocks_per_group<kGemmType, BLOCK_M, BLOCK_N, kNumSMs, kIsMulticastOnA>()>
struct Scheduler {
int current_iter = -1;
// Block configs
uint32_t num_blocks;
uint32_t num_m_blocks;
uint32_t num_n_blocks;
// For SM90 multicast checks
uint32_t num_blocks_in_group;
bool is_peer_cta_alive = true;
// For grouped GEMM
int* grouped_layout;
uint32_t current_group_idx = 0;
// Only used for masked layout
uint32_t current_m_cumsum = 0;
// Only used for k-grouped layout
uint32_t current_shape_k, current_num_valid_groups = 0, current_k_cumsum = 0, current_sf_k_cumsum = 0;
uint32_t next_group_idx, next_shape_k;
// Only used for k-grouped gemm
__device__ __forceinline__ void get_next_k_group(uint32_t &group_idx, uint32_t &shape_k) const {
for (; group_idx < kNumGroups; ++ group_idx) {
shape_k = __ldg(grouped_layout + group_idx);
if (shape_k > 0)
break;
}
}
// ReSharper disable once CppPossiblyUninitializedMember
__device__ __forceinline__ explicit Scheduler(const uint32_t& shape_m, const uint32_t& shape_n, const uint32_t& shape_k,
int* grouped_layout = nullptr) {
num_m_blocks = ceil_div(shape_m, BLOCK_M);
num_n_blocks = ceil_div(shape_n, BLOCK_N);
current_shape_k = shape_k;
if constexpr (kGemmType == GemmType::Normal) {
num_blocks = num_m_blocks * num_n_blocks;
} else if (kGemmType == GemmType::MGroupedContiguous) {
num_blocks = num_m_blocks * num_n_blocks;
this->grouped_layout = grouped_layout;
} else if (kGemmType == GemmType::MGroupedMasked) {
this->grouped_layout = grouped_layout;
} else if (kGemmType == GemmType::KGroupedContiguous) {
this->grouped_layout = grouped_layout;
get_next_k_group(current_group_idx, current_shape_k);
next_group_idx = current_group_idx + 1;
get_next_k_group(next_group_idx, next_shape_k);
}
}
__device__ __forceinline__ void get_swizzled_block_idx(const uint32_t& block_idx, uint32_t& m_block_idx, uint32_t& n_block_idx) {
DG_STATIC_ASSERT(kNum1DBlocksPerGroup % kNumMulticast == 0, "Invalid group size");
// Swizzle for better L2 usages
const auto& primary_num_blocks = kIsMulticastOnA ? num_n_blocks : num_m_blocks;
const auto& secondary_num_blocks = kIsMulticastOnA ? num_m_blocks : num_n_blocks;
const auto& num_blocks_per_group = secondary_num_blocks * kNum1DBlocksPerGroup;
const auto& group_idx = block_idx / num_blocks_per_group;
auto first_block_idx = group_idx * kNum1DBlocksPerGroup;
auto in_group_idx = block_idx % num_blocks_per_group;
num_blocks_in_group = min(kNum1DBlocksPerGroup, primary_num_blocks - first_block_idx);
// Fix unaligned TMA multicast
// NOTES: for SM90 only, as SM90 can dynamically disable TMA multicast
// while SM100 uses 2-CTA, which can not be dynamically disabled
#if __CUDA_ARCH__ < 1000
if (kNumMulticast > 1 and num_blocks_in_group % 2 != 0) {
if (in_group_idx < (num_blocks_in_group ^ 1) * secondary_num_blocks) {
num_blocks_in_group = num_blocks_in_group ^ 1;
} else {
in_group_idx = in_group_idx - (num_blocks_in_group ^ 1) * secondary_num_blocks;
first_block_idx += num_blocks_in_group ^ 1;
num_blocks_in_group = 1;
}
}
#endif
// Convert to final M/N block indices
// `kIsMulticastOnA == true` leads to groups on N
if constexpr (kIsMulticastOnA) {
m_block_idx = in_group_idx / num_blocks_in_group;
n_block_idx = first_block_idx + in_group_idx % num_blocks_in_group;
} else {
m_block_idx = first_block_idx + in_group_idx % num_blocks_in_group;
n_block_idx = in_group_idx / num_blocks_in_group;
}
}
template <bool kWithGroupOffset, KGroupedIndexType kIndexType = KGroupedIndexType::MN>
__device__ __forceinline__ uint32_t get_global_idx(const uint32_t shape_dim, const uint32_t block_size,
const uint32_t& block_idx, const uint32_t& m_block_idx = 0) {
if constexpr (kGemmType == GemmType::Normal) {
return block_idx * block_size;
} else if constexpr (kGemmType == GemmType::MGroupedContiguous) {
const auto offset = kWithGroupOffset ? cute::max(0, __ldg(grouped_layout + m_block_idx * BLOCK_M)) : 0;
return offset * shape_dim + block_idx * block_size;
} else if constexpr (kGemmType == GemmType::MGroupedMasked) {
const auto offset = kWithGroupOffset ? current_group_idx : 0;
return offset * shape_dim + block_idx * block_size;
} else if constexpr (kGemmType == GemmType::KGroupedContiguous) {
auto offset = 0;
if constexpr (kWithGroupOffset) {
if constexpr (kIndexType == KGroupedIndexType::MN)
offset = current_group_idx * shape_dim;
else if constexpr (kIndexType == KGroupedIndexType::K)
offset = current_k_cumsum;
else if constexpr (kIndexType == KGroupedIndexType::SF_K)
offset = current_sf_k_cumsum;
}
return offset + block_idx * block_size;
}
}
__device__ __forceinline__ bool get_next_block(uint32_t& m_block_idx, uint32_t& n_block_idx) {
const auto next_block_idx = (++ current_iter) * kNumSMs + blockIdx.x;
if constexpr (kGemmType == GemmType::MGroupedMasked) {
while (true) {
// End of the task
if (current_group_idx == kNumGroups)
return false;
// Within current group
num_m_blocks = ceil_div(static_cast<uint32_t>(__ldg(grouped_layout + current_group_idx)), BLOCK_M);
const auto current_m_block_cumsum = current_m_cumsum + num_m_blocks;
if (next_block_idx < current_m_block_cumsum * num_n_blocks)
break;
// Move to check the next group
current_group_idx ++, current_m_cumsum = current_m_block_cumsum;
}
get_swizzled_block_idx(next_block_idx - current_m_cumsum * num_n_blocks, m_block_idx, n_block_idx);
} else if (kGemmType == GemmType::KGroupedContiguous) {
while (true) {
// End of the task
if (current_group_idx == kNumGroups)
return false;
// Within current group
if (next_block_idx < (current_num_valid_groups + 1) * num_m_blocks * num_n_blocks)
break;
// Move to check the next group
current_k_cumsum += current_shape_k;
current_sf_k_cumsum += ceil_div(current_shape_k, SF_K_ALIGNMENT);
current_num_valid_groups ++;
current_group_idx = next_group_idx ++;
current_shape_k = next_shape_k;
get_next_k_group(next_group_idx, next_shape_k);
}
get_swizzled_block_idx(next_block_idx - current_num_valid_groups * num_m_blocks * num_n_blocks, m_block_idx, n_block_idx);
} else {
if (next_block_idx >= num_blocks)
return false;
// For SM90 only
// NOTES: we don't have to set `is_peer_cta_alive` for masked grouped GEMM, as it must be aligned
is_peer_cta_alive = num_n_blocks % kNumMulticast == 0 or // Always aligned on N (constant bypass)
num_m_blocks % kNumMulticast == 0 or // Always aligned on M (constant bypass)
(next_block_idx ^ 1) < num_blocks; // Peer CTA in bound
get_swizzled_block_idx(next_block_idx, m_block_idx, n_block_idx);
}
return true;
}
// For SM90 only
__device__ __forceinline__ bool is_tma_multicast_valid(const uint32_t& m_block_idx) const {
if (num_blocks_in_group == 1)
return false;
if constexpr (kGemmType == GemmType::Normal or kGemmType == GemmType::MGroupedMasked or kGemmType == GemmType::KGroupedContiguous) {
return true;
} else {
DG_STATIC_ASSERT(kGemmType == GemmType::MGroupedContiguous, "Invalid Gemm type");
if constexpr (kIsMulticastOnA) {
return true;
} else {
const auto& group_idx = __ldg(grouped_layout + m_block_idx * BLOCK_M);
const auto& peer_group_idx = __ldg(grouped_layout + (m_block_idx ^ 1) * BLOCK_M);
return group_idx == peer_group_idx;
}
}
}
// For SM90 only
// ReSharper disable once CppNotAllPathsReturnValue
__device__ __forceinline__ bool is_computation_valid(const uint32_t& m_block_idx, const uint32_t& m_offset) const {
if constexpr (kGemmType == GemmType::Normal) {
return true;
} else if constexpr (kGemmType == GemmType::MGroupedContiguous) {
return __ldg(grouped_layout + m_offset + m_block_idx * BLOCK_M) >= 0;
} else if constexpr (kGemmType == GemmType::MGroupedMasked) {
return m_offset + m_block_idx * BLOCK_M < __ldg(grouped_layout + current_group_idx);
}
}
};
#pragma clang diagnostic pop
} // namespace deep_gemm
@@ -0,0 +1,260 @@
#pragma once
#include <cute/atom/mma_traits_sm100.hpp>
#include <cute/arch/mma_sm100_umma.hpp>
#include <cute/arch/tmem_allocator_sm100.hpp>
#include <deep_gemm/common/utils.cuh>
namespace deep_gemm::sm100 {
template <uint32_t BLOCK_INNER, uint32_t kSwizzleMode, typename dtype_t>
constexpr uint32_t get_inner_block_atom_size() {
return kSwizzleMode == 0 ? BLOCK_INNER : kSwizzleMode / sizeof(dtype_t);
}
template <uint32_t BLOCK_INNER, uint32_t BLOCK_OUTER,
uint32_t kSwizzleMode, uint32_t kNumMulticast,
typename dtype_t>
__device__ __forceinline__ void
tma_copy(void const* desc_ptr, cutlass::arch::ClusterTransactionBarrier* barrier_ptr,
dtype_t* smem_ptr, const uint32_t& inner_idx, const int32_t& outer_idx) {
DG_STATIC_ASSERT(1 <= kNumMulticast and kNumMulticast <= 2, "Invalid multicast config");
DG_STATIC_ASSERT(static_cast<uint64_t>(cute::TMA::CacheHintSm90::EVICT_NORMAL) ==
static_cast<uint64_t>(cute::TMA::CacheHintSm100::EVICT_NORMAL), "Invalid cache hint");
// 2-CTA function will send signals to the leader CTA only
const auto copy_func = kNumMulticast == 1 ? cute::SM90_TMA_LOAD_2D::copy : cute::SM100_TMA_2SM_LOAD_2D::copy;
// Issue multiple TMAs
constexpr uint32_t BLOCK_INNER_ATOM = get_inner_block_atom_size<BLOCK_INNER, kSwizzleMode, dtype_t>();
#pragma unroll
for (uint32_t i = 0; i < BLOCK_INNER / BLOCK_INNER_ATOM; ++ i) {
copy_func(desc_ptr, reinterpret_cast<uint64_t*>(barrier_ptr),
static_cast<uint64_t>(cute::TMA::CacheHintSm100::EVICT_NORMAL),
smem_ptr + i * BLOCK_OUTER * BLOCK_INNER_ATOM, inner_idx + i * BLOCK_INNER_ATOM, outer_idx);
}
}
__device__ __forceinline__
cute::UMMA::SmemDescriptor make_smem_desc(cute::UMMA::LayoutType layout, void* smem_ptr,
uint32_t stride_byte_offset, uint32_t leading_byte_offset) {
cute::UMMA::SmemDescriptor desc;
// Set the version for SM100
desc.version_ = 1;
// Legacy mode
desc.lbo_mode_ = 0;
// Layout
desc.layout_type_ = static_cast<uint8_t>(layout);
// Start address
const auto uint_ptr = cute::cast_smem_ptr_to_uint(smem_ptr);
desc.start_address_ = static_cast<uint16_t>(uint_ptr >> 4);
// Base offset
desc.base_offset_ = 0;
// SBO and LBO
desc.stride_byte_offset_ = stride_byte_offset >> 4;
desc.leading_byte_offset_ = leading_byte_offset >> 4;
return desc;
}
__device__ __forceinline__
cute::UMMA::SmemDescriptor make_sf_desc(void* smem_ptr) {
// NOTES: the UTCCP layout is K-major by default
// Atom size: 8 x 128 bits
// {SBO, LBO} means the byte stride between atoms on {MN, K}
// Since the UTCCP we used is 128b-wide (only 1 atom on K), so LBO can be zero
return make_smem_desc(cute::UMMA::LayoutType::SWIZZLE_NONE, smem_ptr, 8 * 16, 0);
}
__device__ __forceinline__
void replace_smem_desc_addr(cute::UMMA::SmemDescriptor& desc, const void* smem_ptr) {
const auto uint_ptr = cute::cast_smem_ptr_to_uint(smem_ptr);
desc.start_address_ = static_cast<uint16_t>(uint_ptr >> 4);
}
__device__ __forceinline__
static uint32_t get_atom_base(const cute::UMMA::LayoutType& layout_type) {
return layout_type == cute::UMMA::LayoutType::SWIZZLE_128B_BASE32B ? 32 : 16;
}
// ReSharper disable once CppNotAllPathsReturnValue
template <cute::UMMA::Major kMajorMode, uint32_t kSwizzleMode, bool kUseBase32, typename dtype_t>
constexpr static cute::UMMA::LayoutType to_umma_layout_type() {
DG_STATIC_ASSERT(kSwizzleMode == 0 or kSwizzleMode == 16 or
kSwizzleMode == 32 or kSwizzleMode == 64 or
kSwizzleMode == 128, "Invalid swizzling mode");
// A special case
if constexpr ((cute::is_same_v<dtype_t, float> and kMajorMode == cute::UMMA::Major::MN) or kUseBase32) {
DG_STATIC_ASSERT(kUseBase32, "Invalid swizzling base");
return cute::UMMA::LayoutType::SWIZZLE_128B_BASE32B;
}
// Normal cases
if constexpr (kSwizzleMode == 0) return cute::UMMA::LayoutType::SWIZZLE_NONE;
if constexpr (kSwizzleMode == 16) return cute::UMMA::LayoutType::SWIZZLE_NONE;
if constexpr (kSwizzleMode == 32) return cute::UMMA::LayoutType::SWIZZLE_32B;
if constexpr (kSwizzleMode == 64) return cute::UMMA::LayoutType::SWIZZLE_64B;
if constexpr (kSwizzleMode == 128) return cute::UMMA::LayoutType::SWIZZLE_128B;
}
template <cute::UMMA::Major kMajorMode, uint32_t BLOCK_MN, uint32_t kSwizzleMode, typename dtype_t>
__device__ __forceinline__
constexpr uint32_t get_umma_desc_stride_k() {
return kMajorMode == cute::UMMA::Major::K ? 1 : get_inner_block_atom_size<BLOCK_MN, kSwizzleMode, dtype_t>();
}
template <cute::UMMA::Major kMajorMode, uint32_t BLOCK_MN, uint32_t kSwizzleMode, typename dtype_t>
__device__ __forceinline__
uint32_t advance_umma_desc_lo(const uint32_t& base, const uint32_t& offset, const uint32_t& k_idx) {
return base + (((offset + k_idx * get_umma_desc_stride_k<kMajorMode, BLOCK_MN, kSwizzleMode, dtype_t>()) * static_cast<uint32_t>(sizeof(dtype_t))) >> 4u);
}
template <cute::UMMA::Major kMajorMode, uint32_t BLOCK_MN, uint32_t BLOCK_K, uint32_t kSwizzleMode, bool kUseBase32 = false, typename dtype_t>
__device__ __forceinline__
cute::UMMA::SmemDescriptor make_umma_desc(dtype_t* base_smem_ptr, uint32_t mn_idx, uint32_t k_idx) {
const uint32_t stride_k = get_umma_desc_stride_k<kMajorMode, BLOCK_MN, kSwizzleMode, dtype_t>();
const auto& layout_type = to_umma_layout_type<kMajorMode, kSwizzleMode, kUseBase32, dtype_t>();
const auto& num_non_contiguous = 128 / get_atom_base(layout_type);
if constexpr (kMajorMode == cute::UMMA::Major::K) {
// NOTES: for K-major layout, the swizzle must be 128B (also, atom index must be 0), as `BLOCK_K` is always 128
DG_STATIC_ASSERT(kSwizzleMode == BLOCK_K * sizeof(dtype_t), "Unexpected value");
// Atom size: 8 x `kSwizzleMode` (in bytes, on K)
// {SBO, LBO} means the byte stride between atoms on {MN, K}
// NOTES: on K, there is only 1 atom as asserted previously, so LBO can be 0
const uint32_t stride_byte_offset = num_non_contiguous * BLOCK_K * sizeof(dtype_t);
const uint32_t leading_byte_offset = 0;
return make_smem_desc(layout_type,
base_smem_ptr + mn_idx * BLOCK_K + k_idx * stride_k,
stride_byte_offset, leading_byte_offset);
} else {
constexpr uint32_t BLOCK_MN_ATOM = get_inner_block_atom_size<BLOCK_MN, kSwizzleMode, dtype_t>();
// Must have no in-atom MN-idx
// NOTES: no worries for the runtime assert, the `mn_idx` are constants at compilation time
DG_DEVICE_ASSERT(mn_idx % BLOCK_MN_ATOM == 0);
DG_STATIC_ASSERT(kSwizzleMode > 0, "Invalid swizzling");
// Atom size: `kSwizzleMode` (in bytes, on MN) x 8
// NOTES: `kSwizzleMode == 16` mean non-swizzling but interleaving
// {SBO, LBO} means the byte stride between atoms on {K, MN} for swizzling
// {SBO, LBO} means the byte stride between atoms on {MN, K} for non-swizzling
uint32_t stride_byte_offset = num_non_contiguous * BLOCK_MN_ATOM * sizeof(dtype_t);
uint32_t leading_byte_offset = BLOCK_K * BLOCK_MN_ATOM * sizeof(dtype_t);
if constexpr (kSwizzleMode == 16)
swap(stride_byte_offset, leading_byte_offset);
return make_smem_desc(layout_type,
base_smem_ptr + mn_idx * BLOCK_K + k_idx * stride_k,
stride_byte_offset, leading_byte_offset);
}
}
__device__ __forceinline__
uint64_t make_runtime_instr_desc_with_sf_id(cute::UMMA::InstrDescriptorBlockScaled desc, const uint32_t& sf_id) {
desc.a_sf_id_ = sf_id, desc.b_sf_id_ = sf_id;
return static_cast<uint64_t>(static_cast<uint32_t>(desc)) << 32;
}
template <uint32_t kNumCols>
__device__ constexpr uint32_t get_num_aligned_tmem_cols() {
DG_STATIC_ASSERT(kNumCols <= 512, "Too many tensor memory columns");
if (kNumCols <= 32) return 32;
if (kNumCols <= 64) return 64;
if (kNumCols <= 128) return 128;
if (kNumCols <= 256) return 256;
return 512;
}
__device__ __forceinline__ void tcgen05_before_thread_sync() {
asm volatile("tcgen05.fence::before_thread_sync;");
}
__device__ __forceinline__ void tcgen05_after_thread_sync() {
asm volatile("tcgen05.fence::after_thread_sync;");
}
// UMMA versions with relaxed assertions
struct SM100_MMA_F16BF16_SS {
__device__ static void
fma(uint64_t const& desc_a,
uint64_t const& desc_b,
uint32_t const& tmem_c,
uint32_t const& scale_c,
uint64_t const& desc) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %4, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::f16 [%0], %1, %2, %3, p; \n\t"
"}\n"
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
}
};
struct SM100_MMA_F16BF16_2x1SM_SS {
__device__ static void
fma(uint64_t const& desc_a,
uint64_t const& desc_b,
uint32_t const& tmem_c,
uint32_t const& scale_c,
uint64_t const& desc) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %4, 0;\n\t"
"tcgen05.mma.cta_group::2.kind::f16 [%0], %1, %2, %3, p; \n\t"
"}\n"
:: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c));
}
};
struct SM100_MMA_MXF8F6F4_SS {
__device__ static void
fma(uint64_t const& desc_a,
uint64_t const& desc_b,
uint32_t const& tmem_c,
uint32_t const& scale_c,
uint64_t const& desc,
uint32_t const& tmem_sfa,
uint32_t const& tmem_sfb) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %4, 0;\n\t"
"tcgen05.mma.cta_group::1.kind::mxf8f6f4.block_scale [%0], %1, %2, %3, [%5], [%6], p; \n\t"
"}\n"
:
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c),
"r"(tmem_sfa), "r"(tmem_sfb));
}
};
struct SM100_MMA_MXF8F6F4_2x1SM_SS {
__device__ static void
fma(uint64_t const& desc_a,
uint64_t const& desc_b,
uint32_t const& tmem_c,
uint32_t const& scale_c,
uint64_t const& desc,
uint32_t const& tmem_sfa,
uint32_t const& tmem_sfb) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t"
"setp.ne.b32 p, %4, 0;\n\t"
"tcgen05.mma.cta_group::2.kind::mxf8f6f4.block_scale [%0], %1, %2, %3, [%5], [%6], p; \n\t"
"}\n"
:
: "r"(tmem_c), "l"(desc_a), "l"(desc_b), "r"(static_cast<uint32_t>(desc >> 32)), "r"(scale_c),
"r"(tmem_sfa), "r"(tmem_sfb));
}
};
} // namespace `deep_gemm::sm100`
@@ -0,0 +1,283 @@
#pragma once
#include <cute/arch/copy_sm90_tma.hpp>
#include <cute/arch/cluster_sm90.hpp>
#include <cute/arch/mma_sm90_gmma.hpp>
#include <cute/arch/mma_sm90_gmma_ext.hpp>
#include <deep_gemm/common/utils.cuh>
namespace deep_gemm::sm90 {
template <int N_, typename MMA>
struct FP8MMA {
template <size_t ...Idx>
__forceinline__ __device__ static void call_fma_impl(uint64_t const& desc_a, uint64_t const& desc_b, float* d, bool scale_d, cute::index_sequence<Idx...>) {
using namespace cute::SM90::GMMA;
MMA::fma(desc_a, desc_b, d[Idx]..., (scale_d ? ScaleOut::One : ScaleOut::Zero));
}
__forceinline__ __device__ static void wgmma(uint64_t const& desc_a, uint64_t const& desc_b, float* d, bool scale_d) {
call_fma_impl(desc_a, desc_b, d, scale_d, cute::make_index_sequence<N_/2>{});
}
static constexpr int M = 64;
static constexpr int N = N_;
static constexpr int K = 32;
static constexpr int kNumAccum = M * N / 128;
};
template <int N>
struct FP8MMASelector {
static constexpr auto select_mma() {
using namespace cute::SM90::GMMA;
if constexpr (N == 8) return MMA_64x8x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 16) return MMA_64x16x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 24) return MMA_64x24x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 32) return MMA_64x32x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 40) return MMA_64x40x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 48) return MMA_64x48x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 56) return MMA_64x56x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 64) return MMA_64x64x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 72) return MMA_64x72x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 80) return MMA_64x80x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 88) return MMA_64x88x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 96) return MMA_64x96x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 104) return MMA_64x104x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 112) return MMA_64x112x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 120) return MMA_64x120x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 128) return MMA_64x128x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 136) return MMA_64x136x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 144) return MMA_64x144x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 152) return MMA_64x152x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 160) return MMA_64x160x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 168) return MMA_64x168x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 176) return MMA_64x176x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 184) return MMA_64x184x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 192) return MMA_64x192x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 200) return MMA_64x200x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 208) return MMA_64x208x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 216) return MMA_64x216x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 224) return MMA_64x224x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 232) return MMA_64x232x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 240) return MMA_64x240x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 248) return MMA_64x248x32_F32E4M3E4M3_SS_TN();
if constexpr (N == 256) return MMA_64x256x32_F32E4M3E4M3_SS_TN();
}
static constexpr auto select_type() {
return FP8MMA<N, decltype(select_mma())>();
}
using type = decltype(select_type());
};
template <int N_, typename MMA>
struct BF16MMA {
template <size_t ...Idx>
__forceinline__ __device__ static void call_fma_impl(uint64_t const& desc_a, uint64_t const& desc_b, float* d, bool scale_d, cute::index_sequence<Idx...>) {
using namespace cute::SM90::GMMA;
MMA::fma(desc_a, desc_b, d[Idx]..., (scale_d ? ScaleOut::One : ScaleOut::Zero));
}
__forceinline__ __device__ static void wgmma(uint64_t const& desc_a, uint64_t const& desc_b, float* d, bool scale_d) {
call_fma_impl(desc_a, desc_b, d, scale_d, cute::make_index_sequence<N_/2>{});
}
static constexpr int M = 64;
static constexpr int N = N_;
static constexpr int K = 16;
static constexpr int kNumAccum = M * N / 128;
};
template <int N>
struct BF16MMASelector {
static constexpr auto select_mma() {
using namespace cute::SM90::GMMA;
if constexpr (N == 8) return MMA_64x8x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 16) return MMA_64x16x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 24) return MMA_64x24x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 32) return MMA_64x32x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 40) return MMA_64x40x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 48) return MMA_64x48x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 56) return MMA_64x56x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 64) return MMA_64x64x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 72) return MMA_64x72x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 80) return MMA_64x80x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 88) return MMA_64x88x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 96) return MMA_64x96x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 104) return MMA_64x104x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 112) return MMA_64x112x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 120) return MMA_64x120x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 128) return MMA_64x128x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 136) return MMA_64x136x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 144) return MMA_64x144x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 152) return MMA_64x152x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 160) return MMA_64x160x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 168) return MMA_64x168x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 176) return MMA_64x176x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 184) return MMA_64x184x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 192) return MMA_64x192x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 200) return MMA_64x200x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 208) return MMA_64x208x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 216) return MMA_64x216x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 224) return MMA_64x224x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 232) return MMA_64x232x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 240) return MMA_64x240x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 248) return MMA_64x248x16_F32BF16BF16_SS<Major::K, Major::K>();
if constexpr (N == 256) return MMA_64x256x16_F32BF16BF16_SS<Major::K, Major::K>();
}
static constexpr auto select_type() {
return BF16MMA<N, decltype(select_mma())>();
}
using type = decltype(select_type());
};
template <typename dtype_t>
struct SM90_U32x2_STSM_N {
__device__ __forceinline__ static void
copy(dtype_t src_0, dtype_t src_1, void* smem_dst) {
const uint32_t src[2] = {*reinterpret_cast<uint32_t*>(&src_0), *reinterpret_cast<uint32_t*>(&src_1)};
asm volatile("stmatrix.sync.aligned.x2.m8n8.shared.b16 [%0], {%1, %2};\n"
:: "l"(smem_dst), "r"(src[0]), "r"(src[1]));
}
};
struct SM90_U32x2_LDSM_N {
__device__ __forceinline__ static void
copy(uint32_t& dst_0, uint32_t& dst_1, void* smem_src) {
asm volatile("ldmatrix.sync.aligned.x2.m8n8.shared.b16 {%0, %1}, [%2];\n"
: "=r"(dst_0), "=r"(dst_1)
: "l"(smem_src));
}
};
struct SM90_U32x4_LDSM_N {
__device__ __forceinline__ static void
copy(uint32_t& dst_0, uint32_t& dst_1, uint32_t& dst_2, uint32_t& dst_3, void* smem_src) {
asm volatile("ldmatrix.sync.aligned.x4.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n"
: "=r"(dst_0), "=r"(dst_1), "=r"(dst_2), "=r"(dst_3)
: "l"(smem_src));
}
};
__forceinline__ __device__ void warpgroup_arrive() {
asm volatile("wgmma.fence.sync.aligned;\n" ::: "memory");
}
__forceinline__ __device__ void warpgroup_commit_batch() {
asm volatile("wgmma.commit_group.sync.aligned;\n" ::: "memory");
}
__forceinline__ __device__ void warpgroup_fence_operand(float& reg) {
asm volatile("" : "+f"(reg) :: "memory");
}
template <int N>
__forceinline__ __device__ void warpgroup_wait() {
DG_STATIC_ASSERT(N >= 0 and N <= 7, "WGMMA wait: N must be in range [0, 7]");
asm volatile("wgmma.wait_group.sync.aligned %0;\n" :: "n"(N) : "memory");
}
// TODO: replace with CUTLASS solution
union GmmaDescriptor {
__host__ __device__ constexpr GmmaDescriptor() noexcept: desc_(0) {}
__host__ __device__ constexpr GmmaDescriptor(uint64_t desc) noexcept: desc_(desc) {}
__host__ __device__ constexpr GmmaDescriptor(GmmaDescriptor const &t) noexcept: desc_(t.desc_) {}
__host__ __device__ constexpr GmmaDescriptor(GmmaDescriptor &&t) noexcept: desc_(t.desc_) {}
__host__ __device__ constexpr GmmaDescriptor &operator=(GmmaDescriptor const &t) noexcept {
desc_ = t.desc_;
return *this;
}
__host__ __device__ constexpr GmmaDescriptor &operator=(GmmaDescriptor &&t) noexcept {
desc_ = t.desc_;
return *this;
}
uint64_t desc_;
uint32_t reg32_[2];
uint16_t reg16_[4];
struct {
uint16_t start_address_: 14, : 2;
uint16_t leading_byte_offset_: 14, : 2;
uint16_t stride_byte_offset_: 14, : 2;
uint8_t : 1, base_offset_: 3, : 4;
uint8_t : 6, layout_type_: 2;
} bitfield;
// Decay to an `uint64_t`
__host__ __device__ constexpr operator uint64_t() const noexcept { return desc_; }
};
template <class PointerType>
__device__ GmmaDescriptor make_smem_desc(PointerType smem_ptr, const int& layout_type,
const int& leading_byte_offset = 0,
const int& stride_byte_offset = 1024) {
GmmaDescriptor desc;
const auto& uint_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(smem_ptr));
desc.bitfield.start_address_ = uint_ptr >> 4;
desc.bitfield.layout_type_ = layout_type;
desc.bitfield.leading_byte_offset_ = leading_byte_offset >> 4;
desc.bitfield.stride_byte_offset_ = stride_byte_offset >> 4;
desc.bitfield.base_offset_ = 0;
return desc;
}
__device__ __forceinline__ void
tma_copy(void const* desc_ptr, uint64_t* barrier_ptr, void* smem_ptr,
const uint32_t& crd_0, const uint32_t& crd_1, const uint32_t& num_tma_multicast = 1) {
constexpr auto cache_hint = static_cast<uint64_t>(cute::TMA::CacheHintSm90::EVICT_NORMAL);
if (num_tma_multicast == 1) {
cute::SM90_TMA_LOAD_2D::copy(desc_ptr, barrier_ptr, cache_hint, smem_ptr, crd_0, crd_1);
} else if (cute::block_rank_in_cluster() == 0) {
cute::SM90_TMA_LOAD_MULTICAST_2D::copy(desc_ptr, barrier_ptr, (1 << num_tma_multicast) - 1, cache_hint, smem_ptr, crd_0, crd_1);
}
}
__device__ __forceinline__ void
tma_3d_copy(void const* desc_ptr, uint64_t* barrier_ptr, void* smem_ptr,
const uint32_t& crd_0, const uint32_t& crd_1, const uint32_t& crd_2) {
constexpr auto cache_hint = static_cast<uint64_t>(cute::TMA::CacheHintSm90::EVICT_NORMAL);
cute::SM90_TMA_LOAD_3D::copy(desc_ptr, barrier_ptr, cache_hint, smem_ptr, crd_0, crd_1, crd_2);
}
// Tensormap related
__device__ __forceinline__ void tensor_map_release_cta() {
asm volatile ("fence.proxy.tensormap::generic.release.cta;");
}
__device__ __forceinline__ void tensor_map_acquire_cta(const cute::TmaDescriptor* gmem_desc_ptr) {
auto gmem_int_desc = reinterpret_cast<uint64_t>(gmem_desc_ptr);
asm volatile ("fence.proxy.tensormap::generic.acquire.cta [%0], 128;" :: "l"(gmem_int_desc) : "memory");
}
__device__ __forceinline__ void tensor_map_replace_global_addr_in_smem(cute::TmaDescriptor* smem_desc, const void* new_addr) {
auto smem_int_desc = static_cast<uint32_t>(__cvta_generic_to_shared(smem_desc));
const auto new_int64_addr = reinterpret_cast<uint64_t>(new_addr);
asm volatile ("tensormap.replace.tile.global_address.shared::cta.b1024.b64 [%0], %1;" :: "r"(smem_int_desc), "l"(new_int64_addr));
}
__device__ __forceinline__ void tensor_map_replace_global_inner_dim_stride_in_smem(cute::TmaDescriptor* smem_desc, const uint32_t& new_dim, const uint64_t& new_stride) {
auto smem_int_desc = __cvta_generic_to_shared(smem_desc);
asm volatile ("tensormap.replace.tile.global_dim.shared::cta.b1024.b32 [%0], 0, %1;" :: "l"(smem_int_desc), "r"(new_dim));
#if ((__CUDACC_VER_MAJOR__ > 12) or ((__CUDACC_VER_MAJOR__ == 12) and (__CUDACC_VER_MINOR__ >= 3)))
asm volatile("tensormap.replace.tile.global_stride.shared::cta.b1024.b64 [%0], 0, %1;" :: "l"(smem_int_desc), "l"(new_stride));
#else
DG_STATIC_ASSERT(false, "Invalid CUDA version");
#endif
}
} // namespace `deep_gemm::sm90`
@@ -0,0 +1,18 @@
#pragma once
namespace deep_gemm {
enum class GemmType {
Normal = 0,
MGroupedContiguous = 1,
MGroupedMasked = 2,
KGroupedContiguous = 3,
};
enum class KernelType {
Kernel1D1D = 0,
Kernel1D2D = 1,
KernelNoSF = 2
};
} // namespace deep_gemm
@@ -0,0 +1,179 @@
#pragma once
#include <cuda_bf16.h>
#include <cuda_fp8.h>
#include <cuda/std/cstdint>
#include <cuda/std/utility>
#include <cute/container/tuple.hpp>
#include "cute_tie.cuh"
#ifdef __CLION_IDE__
__host__ __device__ __forceinline__ void host_device_printf(const char* format, ...) {
asm volatile("trap;");
}
#define printf host_device_printf
#endif
#ifndef DG_DEVICE_ASSERT
#define DG_DEVICE_ASSERT(cond) \
do { \
if (not (cond)) { \
printf("Assertion failed: %s:%d, condition: %s\n", __FILE__, __LINE__, #cond); \
asm("trap;"); \
} \
} while (0)
#endif
#ifndef DG_TRAP_ONLY_DEVICE_ASSERT
#define DG_TRAP_ONLY_DEVICE_ASSERT(cond) \
do { \
if (not (cond)) \
asm("trap;"); \
} while (0)
#endif
#ifndef DG_STATIC_ASSERT
#define DG_STATIC_ASSERT(cond, ...) static_assert(cond, __VA_ARGS__)
#endif
namespace deep_gemm {
template <typename FuncT>
struct PatternVisitor {
FuncT func;
__device__ __host__
explicit PatternVisitor(FuncT&& func): func(std::forward<FuncT>(func)) {}
__device__ __host__
auto operator [](const uint32_t& i) {
return func(i);
}
};
template <typename T>
__device__ __host__ T ceil_div(T a, T b) {
return (a + b - 1) / b;
}
template <typename T>
__device__ __host__ constexpr T constexpr_ceil_div(T a, T b) {
return (a + b - 1) / b;
}
template <typename T>
__device__ __host__ T align(T a, T b) {
return ceil_div(a, b) * b;
}
template <typename T>
__device__ __host__ constexpr T constexpr_align(T a, T b) {
return constexpr_ceil_div(a, b) * b;
}
template <typename T>
__device__ __host__ constexpr T constexpr_gcd(T a, T b) {
return b == 0 ? a : constexpr_gcd(b, a % b);
}
template<typename T>
__forceinline__ __device__ void swap(T& a, T& b) {
T temp = a;
a = b;
b = temp;
}
__forceinline__ __device__ uint32_t get_sm_idx() {
uint32_t sm_idx;
asm ("mov.u32 %0, %%smid;" : "=r"(sm_idx));
return sm_idx;
}
__forceinline__ __device__ uint32_t get_lane_idx() {
uint32_t lane_id;
asm ("mov.u32 %0, %laneid;" : "=r"(lane_id));
return lane_id;
}
__device__ __forceinline__ uint32_t ld_shared(const uint32_t* ptr) {
uint32_t ret;
asm volatile("ld.shared.u32 %0, [%1];" : "=r"(ret) : "l"(ptr));
return ret;
}
__device__ __forceinline__ float2 ld_shared(const float2* ptr) {
float2 ret;
asm volatile("ld.shared.v2.f32 {%0, %1}, [%2];" : "=f"(ret.x), "=f"(ret.y) : "l"(ptr));
return ret;
}
__device__ __forceinline__ float4 ld_shared(const float4* ptr) {
float4 ret;
asm volatile("ld.shared.v4.f32 {%0, %1, %2, %3}, [%4];" : "=f"(ret.x), "=f"(ret.y), "=f"(ret.z), "=f"(ret.w) : "l"(ptr));
return ret;
}
__device__ __forceinline__ uint4 ld_shared(const uint4* ptr) {
uint4 ret;
asm volatile("ld.shared.v4.u32 {%0, %1, %2, %3}, [%4];" : "=r"(ret.x), "=r"(ret.y), "=r"(ret.z), "=r"(ret.w) : "l"(ptr));
return ret;
}
__device__ __forceinline__ float ld_shared(const float* ptr) {
float ret;
asm volatile("ld.shared.f32 %0, [%1];" : "=f"(ret) : "l"(ptr));
return ret;
}
__device__ __forceinline__ void st_shared(const float* ptr, float val) {
asm volatile("st.shared.f32 [%0], %1;" :: "l"(ptr), "f"(val));
}
__device__ __forceinline__ void st_shared(const float2* ptr, float2 val) {
asm volatile("st.shared.v2.f32 [%0], {%1, %2};" :: "l"(ptr), "f"(val.x), "f"(val.y));
}
__device__ __forceinline__ void st_shared(const uint32_t* ptr, uint32_t val) {
asm volatile("st.shared.u32 [%0], %1;" :: "l"(ptr), "r"(val));
}
__device__ __forceinline__ void st_shared(const void* ptr, uint32_t x, uint32_t y) {
asm volatile("st.shared.v2.u32 [%0], {%1, %2};" :: "l"(ptr), "r"(x), "r"(y));
}
__device__ __forceinline__ void st_shared(const void* ptr, uint32_t x, uint32_t y, uint32_t z, uint32_t w) {
asm volatile("st.shared.v4.u32 [%0], {%1, %2, %3, %4};" :: "l"(ptr), "r"(x), "r"(y), "r"(z), "r"(w));
}
template <typename old_t>
__device__ __forceinline__ int cast_into_bf16_and_pack(old_t& x, old_t& y) {
auto bf16x2 = __float22bfloat162_rn({*reinterpret_cast<float*>(&x), *reinterpret_cast<float*>(&y)});
return *reinterpret_cast<int*>(&bf16x2);
}
__device__ __forceinline__ void prefetch_l1(void *ptr) {
asm volatile("prefetch.global.L1 [%0];" :: "l"(ptr));
}
template <uint32_t kNumBytes>
struct Vectorized {
static auto zeros() {
// TODO: add `ulonglong4` for SM100 once `__ldg` support this
if constexpr (kNumBytes > 0 and kNumBytes % 16 == 0) {
return make_uint4(0, 0, 0, 0);
} else if constexpr (kNumBytes > 0 and kNumBytes % 8 == 0) {
return make_uint2(0, 0);
} else if constexpr (kNumBytes > 0 and kNumBytes % 4 == 0) {
return 0;
} else {
DG_STATIC_ASSERT(kNumBytes > 0 and kNumBytes % 4 == 0, "Invalid vectorization");
}
}
using vec_t = decltype(zeros());
};
} // namespace `deep_gemm`
@@ -0,0 +1,408 @@
#pragma once
#pragma clang diagnostic push
#pragma clang diagnostic ignored "-Wunknown-attributes"
#include <cutlass/arch/barrier.h>
#include <cutlass/arch/reg_reconfig.h>
#include <cute/arch/cluster_sm90.hpp>
#include <cute/arch/copy_sm90_desc.hpp>
#include <cute/arch/copy_sm90_tma.hpp>
#include <deep_gemm/common/epilogue_utils.cuh>
#include <deep_gemm/common/utils.cuh>
#include <deep_gemm/common/scheduler.cuh>
#include <deep_gemm/common/sm90_utils.cuh>
namespace deep_gemm {
using namespace deep_gemm::sm90;
template <uint32_t kNumFormerIters, uint32_t kGap, uint32_t kEnd, typename func_t>
__device__ void dispatch_num_former_iters(uint32_t num_former_iters, const func_t& func) {
if (num_former_iters == kNumFormerIters) {
func(cute::Int<kNumFormerIters>{});
return;
}
if constexpr (kNumFormerIters + kGap <= kEnd)
dispatch_num_former_iters<kNumFormerIters + kGap, kGap, kEnd>(num_former_iters, func);
}
template <uint32_t SHAPE_M, uint32_t SHAPE_N, uint32_t SHAPE_K,
uint32_t kNumGroups,
uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t BLOCK_K,
uint32_t kSwizzleDMode,
uint32_t kNumStages, uint32_t kNumLastStages,
uint32_t kNumTMAThreads, uint32_t kNumMathThreads,
uint32_t kNumTMAMulticast, bool kIsTMAMulticastOnA,
uint32_t kNumSMs, GemmType kGemmType,
typename epilogue_type_t>
__global__ __launch_bounds__(kNumTMAThreads + kNumMathThreads, 1) void
sm90_fp8_gemm_1d2d_impl(float* sfb, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const __grid_constant__ cute::TmaDescriptor tensor_map_a,
const __grid_constant__ cute::TmaDescriptor tensor_map_b,
const __grid_constant__ cute::TmaDescriptor tensor_map_d,
const __grid_constant__ cute::TmaDescriptor tensor_map_sfa) {
#if (defined(__CUDA_ARCH__) and (__CUDA_ARCH__ >= 900)) or defined(__CLION_IDE__)
// Scaling checks
DG_STATIC_ASSERT(BLOCK_K == 128, "Only support per-128-channel FP8 scaling");
DG_STATIC_ASSERT(constexpr_ceil_div(BLOCK_N, BLOCK_K) == 1 or (constexpr_gcd(BLOCK_N, BLOCK_K) == BLOCK_N - BLOCK_K), "Too much B scales in a single block");
// Types
using WGMMA = typename FP8MMASelector<BLOCK_N>::type;
using Barrier = cutlass::arch::ClusterTransactionBarrier;
DG_STATIC_ASSERT(BLOCK_M % WGMMA::M == 0, "Invalid block size");
// Overwrite shape constants if the compiler gives
shape_m = SHAPE_M != 0 ? SHAPE_M : shape_m;
shape_n = SHAPE_N != 0 ? SHAPE_N : shape_n;
shape_k = SHAPE_K != 0 ? SHAPE_K : shape_k;
// Shared memory
static constexpr bool kMustUseUniformedScaleB = (BLOCK_K % BLOCK_N == 0);
static constexpr uint32_t SMEM_D_SIZE = BLOCK_M * BLOCK_N * sizeof(__nv_bfloat16);
static constexpr uint32_t SMEM_A_SIZE_PER_STAGE = BLOCK_M * BLOCK_K * sizeof(__nv_fp8_e4m3);
static constexpr uint32_t SMEM_B_SIZE_PER_STAGE = BLOCK_N * BLOCK_K * sizeof(__nv_fp8_e4m3);
static constexpr uint32_t SMEM_SFA_SIZE_PER_STAGE = BLOCK_M * sizeof(float);
const uint32_t& shape_k_scales = ceil_div(shape_k, BLOCK_K);
const uint32_t& smem_sfb_size = align<uint32_t>(shape_k_scales * (kMustUseUniformedScaleB ? 1 : 2) * sizeof(float), sizeof(Barrier));
// Configs
const uint32_t num_total_k_blocks = ceil_div(shape_k, BLOCK_K);
const uint32_t warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
const uint32_t lane_idx = get_lane_idx();
// Prefetch TMA descriptors at the very beginning
if (warp_idx == kNumMathThreads / 32 and cute::elect_one_sync()) {
cute::prefetch_tma_descriptor(&tensor_map_a);
cute::prefetch_tma_descriptor(&tensor_map_b);
cute::prefetch_tma_descriptor(&tensor_map_sfa);
cute::prefetch_tma_descriptor(&tensor_map_d);
}
__syncwarp();
// Align to 1024 bytes for swizzle-128B
extern __shared__ __align__(1024) uint8_t smem_buffer[];
DG_STATIC_ASSERT(SMEM_D_SIZE % 1024 == 0, "Shared memory of A/B must be aligned to 1024 bytes");
// Data on shared memory
auto smem_d = reinterpret_cast<__nv_bfloat16*>(smem_buffer);
auto smem_a = PatternVisitor([&](const uint32_t& i) {
return reinterpret_cast<__nv_fp8_e4m3*>(smem_buffer + SMEM_D_SIZE + i * SMEM_A_SIZE_PER_STAGE);
});
auto smem_b = PatternVisitor([&](const uint32_t& i) {
return reinterpret_cast<__nv_fp8_e4m3*>(smem_buffer + SMEM_D_SIZE + kNumStages * SMEM_A_SIZE_PER_STAGE + i * SMEM_B_SIZE_PER_STAGE);
});
constexpr uint32_t SMEM_SF_OFFSET = SMEM_D_SIZE + kNumStages * (SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE);
auto smem_sfa = PatternVisitor([&](const uint32_t& i) {
return reinterpret_cast<float*>(smem_buffer + SMEM_SF_OFFSET + i * SMEM_SFA_SIZE_PER_STAGE);
});
auto smem_sfb = reinterpret_cast<float*>(smem_buffer + SMEM_SF_OFFSET + kNumStages * SMEM_SFA_SIZE_PER_STAGE);
// Fill barriers
auto barrier_start_ptr = reinterpret_cast<Barrier*>(reinterpret_cast<uint8_t*>(smem_sfb) + smem_sfb_size);
auto full_barriers = PatternVisitor([&](const uint32_t& i) { return barrier_start_ptr + i; });
auto empty_barriers = PatternVisitor([&](const uint32_t& i) { return barrier_start_ptr + kNumStages + i; });
// Initialize barriers
DG_STATIC_ASSERT(kNumTMAMulticast <= 32, "Too many TMA multicast");
if (warp_idx == kNumMathThreads / 32 + 1 and cute::elect_one_sync()) {
// NOTES: we always use `lane_idx` to arrive for the `lane_idx`-th CTA in the cluster,
// even with TMA multicast disabled, we want to make the behavior aligned
#pragma unroll
for (uint32_t i = 0; i < kNumStages; ++ i) {
full_barriers[i]->init(1);
empty_barriers[i]->init(kNumTMAMulticast * kNumMathThreads / 32);
}
// Make initialized barrier visible in async proxy
cutlass::arch::fence_barrier_init();
}
// Synchronize all threads to make barrier visible in normal memory model
(kNumTMAMulticast > 1) ? cute::cluster_sync() : __syncthreads();
// Register reconfigurations
constexpr uint32_t kNumTMARegisters = 40;
constexpr uint32_t kNumMathRegisters = 232;
// Block scheduler
uint32_t m_block_idx, n_block_idx;
auto scheduler = Scheduler<kGemmType, BLOCK_M, BLOCK_N, kNumGroups, kNumTMAMulticast, kIsTMAMulticastOnA, kNumSMs>(shape_m, shape_n, shape_k, grouped_layout);
// Pipeline and TMA phases
uint32_t stage_idx = 0, phase = 0;
auto advance_pipeline = [&](uint32_t& k_block_idx) {
++ k_block_idx;
// Flip phases only if reach the next first stage
stage_idx = stage_idx == kNumStages - 1 ? 0 : stage_idx + 1;
phase ^= stage_idx == 0;
};
if (warp_idx >= kNumMathThreads / 32) {
// TMA warp-group for loading data
cutlass::arch::warpgroup_reg_dealloc<kNumTMARegisters>();
// NOTES: only one thread (or warp) will be used
if (warp_idx == kNumMathThreads / 32 and cute::elect_one_sync()) {
// Persistently schedule over blocks
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
// Assign TMA multicast number into A and B
// NOTES: there may be additional odd rows/columns or cases where multicast is not possible.
const bool is_tma_multicast_valid = scheduler.is_tma_multicast_valid(m_block_idx);
const uint32_t num_tma_multicast_a = (kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
const uint32_t num_tma_multicast_b = (not kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
DG_STATIC_ASSERT(kNumTMAMulticast <= 2, "Scheduler does not support > 2 TMA multicast");
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
// Wait consumer release
empty_barriers[stage_idx]->wait(phase ^ 1);
// Issue TMA A
constexpr bool kWithGroupOffsetA = kGemmType == GemmType::MGroupedMasked;
auto& full_barrier = *full_barriers[stage_idx];
const uint32_t k_idx = k_block_idx * BLOCK_K;
tma_copy(&tensor_map_a, reinterpret_cast<uint64_t*>(&full_barrier),
smem_a[stage_idx], k_idx, scheduler.get_global_idx<kWithGroupOffsetA>(shape_m, BLOCK_M, m_block_idx),
num_tma_multicast_a);
tma_copy(&tensor_map_sfa, reinterpret_cast<uint64_t*>(&full_barrier),
smem_sfa[stage_idx], m_block_idx * BLOCK_M, scheduler.get_global_idx<kWithGroupOffsetA>(shape_k_scales, 1, k_block_idx),
num_tma_multicast_a);
// Issue TMA B
tma_copy(&tensor_map_b, reinterpret_cast<uint64_t*>(&full_barrier),
smem_b[stage_idx], k_idx, scheduler.get_global_idx<true>(shape_n, BLOCK_N, n_block_idx, m_block_idx),
num_tma_multicast_b);
full_barrier.arrive_and_expect_tx(SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE + SMEM_SFA_SIZE_PER_STAGE);
}
}
// To safely deconstruct distributed shared barriers, we need another round of empty waits
if constexpr (kNumTMAMulticast > 1) {
for (uint32_t i = 0; i < kNumStages; advance_pipeline(i))
empty_barriers[stage_idx]->wait(phase ^ 1);
}
}
} else {
// Math warp-groups for WGMMA
cutlass::arch::warpgroup_reg_alloc<kNumMathRegisters>();
// NOTES: use `__shfl_sync` to encourage NVCC to use unified registers
const auto math_wg_idx = __shfl_sync(0xffffffff, threadIdx.x / 128, 0);
const auto r_0 = warp_idx * 16 + lane_idx / 4, r_1 = r_0 + 8;
auto a_desc = make_smem_desc(smem_a[0] + math_wg_idx * WGMMA::M * BLOCK_K, 1);
auto b_desc = make_smem_desc(smem_b[0], 1);
const uint32_t a_desc_lo = __shfl_sync(0xffffffff, a_desc.reg32_[0], 0);
const uint32_t b_desc_lo = __shfl_sync(0xffffffff, b_desc.reg32_[0], 0);
// Persistently schedule over blocks
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
// Decide the number of scales B to load
DG_TRAP_ONLY_DEVICE_ASSERT(shape_n % 8 == 0);
uint32_t num_former_iters = BLOCK_N / 8, num_full_iters = num_former_iters;
if constexpr (not kMustUseUniformedScaleB) {
num_former_iters = min(BLOCK_N, BLOCK_K - n_block_idx * BLOCK_N % BLOCK_K) / 8;
num_full_iters = min(shape_n - n_block_idx * BLOCK_N, BLOCK_N) / 8;
}
uint32_t num_sfb = shape_k_scales * (num_former_iters >= num_full_iters ? 1 : 2);
// Load B scales with math warp-groups
// NOTES: except the first warp, we want to overlap loading B scales with TMA stores between tasks
if (threadIdx.x >= 32) {
auto num_previous_lines = scheduler.get_global_idx<true>(ceil_div(shape_n, BLOCK_K), 0, 0, m_block_idx);
auto local_sfb = sfb + (num_previous_lines + ((n_block_idx * BLOCK_N) / BLOCK_K)) * shape_k_scales;
#pragma unroll
for (uint32_t i = threadIdx.x - 32; i < num_sfb; i += kNumMathThreads - 32)
st_shared(smem_sfb + i, __ldg(local_sfb + i));
}
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
// Accumulation for WGMMA or CUDA promotion
constexpr uint32_t WAVE_BLOCK_M = WGMMA::M * (BLOCK_M <= 64 ? 1 : 2);
DG_STATIC_ASSERT(BLOCK_M % WAVE_BLOCK_M == 0, "Invalid block sizes");
float accum[WGMMA::kNumAccum], final_accum[WGMMA::kNumAccum * (BLOCK_M / WAVE_BLOCK_M)] = {0};
// Empty barrier arrival
auto empty_barrier_arrive = [&]() {
if constexpr (kNumTMAMulticast == 1) {
lane_idx == 0 ? empty_barriers[stage_idx]->arrive() : void();
} else {
auto target_cta = scheduler.is_peer_cta_alive ? lane_idx : cute::block_rank_in_cluster();
lane_idx < kNumTMAMulticast ? empty_barriers[stage_idx]->arrive(target_cta) : void();
}
};
// Skip useless computations
if (scheduler.is_computation_valid(m_block_idx, math_wg_idx * WGMMA::M)) {
// The compiler must know the dynamic variable `num_former_iters`'s real value
constexpr bool kShouldOptimize = BLOCK_K / constexpr_gcd(BLOCK_K, BLOCK_N) <= 4 and not kMustUseUniformedScaleB;
constexpr uint32_t kGap = constexpr_gcd(BLOCK_K, BLOCK_N) / 8;
constexpr uint32_t kEnd = kShouldOptimize ? BLOCK_K / 8 : 0;
// Dispatch `num_former_iters` and launch MMAs
dispatch_num_former_iters<0, kGap, kEnd>(kShouldOptimize ? num_former_iters : 0, [&](auto _) {
#pragma unroll 8
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
const auto& a_desc_base_lo = a_desc_lo + stage_idx * (SMEM_A_SIZE_PER_STAGE / 16);
const auto& b_desc_base_lo = b_desc_lo + stage_idx * (SMEM_B_SIZE_PER_STAGE / 16);
// Read B scales
float scale_b_0 = ld_shared(smem_sfb + k_block_idx), scale_b_1;
// NOTES: even some blocks do not need to read the second row, but we still load one to align with other blocks
if constexpr (not kMustUseUniformedScaleB)
scale_b_1 = ld_shared(smem_sfb + k_block_idx + shape_k_scales);
// Wait TMA arrivals
full_barriers[stage_idx]->wait(phase);
// TODO: remove some useless computation for unaligned Ms
#pragma unroll
for (uint32_t local_idx = 0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx) {
auto m_offset = local_idx * WAVE_BLOCK_M;
// Read A scales
// NOTES: all shared memory read must be prior to `warpgroup_arrive` to avoid next scheduled block polluting the results
auto scale_a_0 = ld_shared(smem_sfa[stage_idx] + r_0 + m_offset);
auto scale_a_1 = ld_shared(smem_sfa[stage_idx] + r_1 + m_offset);
// Commit WGMMA instructions
#pragma unroll
for (uint32_t i = 0; i < WGMMA::kNumAccum; ++ i)
warpgroup_fence_operand(accum[i]);
warpgroup_arrive();
#pragma unroll
for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) {
a_desc.reg32_[0] = a_desc_base_lo + (m_offset * BLOCK_K + k * WGMMA::K) / 16;
b_desc.reg32_[0] = b_desc_base_lo + k * WGMMA::K / 16;
WGMMA::wgmma(a_desc, b_desc, accum, k);
}
warpgroup_commit_batch();
#pragma unroll
for (uint32_t i = 0; i < WGMMA::kNumAccum; ++ i)
warpgroup_fence_operand(accum[i]);
warpgroup_wait<0>();
// Notify barrier arrival at the last warpgroup wave
if (local_idx == BLOCK_M / WAVE_BLOCK_M - 1)
empty_barrier_arrive();
// Promote with scales
// NOTES: making it as predicates is very important for performance, comparing to two loops
float scale_0_0 = scale_a_0 * scale_b_0, scale_1_0 = scale_a_1 * scale_b_0;
float scale_0_1, scale_1_1;
if constexpr (not kMustUseUniformedScaleB)
scale_0_1 = scale_a_0 * scale_b_1, scale_1_1 = scale_a_1 * scale_b_1;
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
#pragma unroll
for (uint32_t i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
// NOTES: for unrolled `num_former_iters` cases, we expect the compiler to automatically make it a constant
bool predicate = kMustUseUniformedScaleB or i < num_former_iters;
shifted_accum[i * 4 + 0] += (predicate ? scale_0_0 : scale_0_1) * accum[i * 4 + 0];
shifted_accum[i * 4 + 1] += (predicate ? scale_0_0 : scale_0_1) * accum[i * 4 + 1];
shifted_accum[i * 4 + 2] += (predicate ? scale_1_0 : scale_1_1) * accum[i * 4 + 2];
shifted_accum[i * 4 + 3] += (predicate ? scale_1_0 : scale_1_1) * accum[i * 4 + 3];
}
}
}
});
} else {
#pragma unroll
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
full_barriers[stage_idx]->wait(phase);
empty_barrier_arrive();
}
}
// TMA checks
constexpr uint32_t kNumElemBytes = sizeof(nv_bfloat16);
constexpr uint32_t TMA_D_BLOCK_N = kSwizzleDMode == 0 ? BLOCK_N : (kSwizzleDMode / kNumElemBytes);
constexpr uint32_t WGMMA_M_PER_WARP = WGMMA::M / 4;
DG_STATIC_ASSERT(BLOCK_M % 8 == 0, "Invalid swizzling atom");
DG_STATIC_ASSERT(BLOCK_N % TMA_D_BLOCK_N == 0 and BLOCK_N / TMA_D_BLOCK_N <= 32,
"Unaligned TMA store or too many TMA store instructions");
DG_STATIC_ASSERT(TMA_D_BLOCK_N % 8 == 0, "Invalid TMA block N");
// Wait last TMA store to be finished
if (threadIdx.x < BLOCK_N / TMA_D_BLOCK_N)
cute::tma_store_wait<0>();
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
// Write back to shared memory using STSM and issue TMA stores
DG_STATIC_ASSERT(WGMMA::kNumAccum % 4 == 0, "Invalid STSM x2 vectorization");
#pragma unroll
for (uint32_t local_idx = 0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx) {
auto m_offset = local_idx * WAVE_BLOCK_M;
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
#pragma unroll
for (auto i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
// Swizzle or padding into the correct address
uint8_t* smem_ptr = nullptr;
if constexpr (kSwizzleDMode > 0) {
// Calculate the swizzling atom offset and in-atom offset
constexpr uint32_t kNumBankGroupBytes = 16;
auto atom_offset = i / (TMA_D_BLOCK_N / 8), in_atom_offset = i % (TMA_D_BLOCK_N / 8);
// Calculate the index of the bank group to be written in the atom
auto bank_group_index = in_atom_offset + lane_idx * (kSwizzleDMode / kNumBankGroupBytes);
// Reshape the atom in another view and swizzle
// - original: `(BLOCK_M, kSwizzleDMode / kNumBankGroupBytes)`
// - new: `(BLOCK_M * kSwizzleDMode / kNumBankGroupBytes / 8, 8)`
constexpr bool kHasShortcut = (kSwizzleDMode / kNumBankGroupBytes) == 8;
auto row = kHasShortcut ? (in_atom_offset / 8 + lane_idx) : (bank_group_index / 8);
auto col = kHasShortcut ? (in_atom_offset) : (bank_group_index % 8);
col ^= row % (kSwizzleDMode / 16);
// Add back into the base pointer
// NOTES: think twice before modifying this, as changes may affect the number of instructions
smem_ptr = reinterpret_cast<uint8_t*>(smem_d) + // Base pointer
warp_idx * (WGMMA_M_PER_WARP * kSwizzleDMode) + // Warp offset
m_offset * kSwizzleDMode + // Wave offset
atom_offset * BLOCK_M * kSwizzleDMode + // Swizzle atom offset (constants)
row * (kNumBankGroupBytes * 8) + col * kNumBankGroupBytes; // In-atom offset
} else {
// No swizzling, just padding
smem_ptr = reinterpret_cast<uint8_t*>(smem_d + (m_offset + warp_idx * WGMMA_M_PER_WARP + lane_idx) * BLOCK_N + i * 8);
}
// NOTES: only 16 lanes' addresses are used
SM90_U32x2_STSM_N<nv_bfloat162>::copy(
__float22bfloat162_rn({shifted_accum[i * 4 + 0], shifted_accum[i * 4 + 1]}),
__float22bfloat162_rn({shifted_accum[i * 4 + 2], shifted_accum[i * 4 + 3]}),
smem_ptr
);
}
}
cute::tma_store_fence();
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
// Use TMA store to write back to global memory
// TODO: compatible with FP32 output
constexpr bool kWithGroupOffsetD = kGemmType == GemmType::MGroupedMasked;
DG_STATIC_ASSERT(kNumMathThreads >= BLOCK_N / TMA_D_BLOCK_N, "Too many TMA blocks");
if (threadIdx.x < BLOCK_N / TMA_D_BLOCK_N) {
auto in_block_n_offset = threadIdx.x * TMA_D_BLOCK_N;
auto smem_ptr = smem_d + in_block_n_offset * BLOCK_M;
cute::SM90_TMA_STORE_2D::copy(&tensor_map_d, smem_ptr,
epilogue_type_t::apply_index_n<TMA_D_BLOCK_N>(n_block_idx * BLOCK_N + in_block_n_offset),
scheduler.get_global_idx<kWithGroupOffsetD>(shape_m, BLOCK_M, m_block_idx));
cute::tma_store_arrive();
}
__syncwarp();
}
}
#else
if (blockIdx.x == 0 and threadIdx.x == 0)
DG_DEVICE_ASSERT(false and "This kernel only support sm_90a");
#endif
}
}; // namespace deep_gemm
#pragma clang diagnostic pop
@@ -0,0 +1,590 @@
#include <cutlass/arch/barrier.h>
#include <cutlass/arch/reg_reconfig.h>
#include <cute/arch/cluster_sm90.hpp>
#include <cute/arch/copy_sm90_desc.hpp>
#include <cute/arch/copy_sm90_tma.hpp>
#include <deep_gemm/common/epilogue_utils.cuh>
#include <deep_gemm/common/utils.cuh>
#include <deep_gemm/common/scheduler.cuh>
#include <deep_gemm/common/sm90_utils.cuh>
#include <stdexcept>
#include <string>
// LT-PATCH: upstream hard-#defines `__CUDA_ARCH__ 900` here, which forces the wgmma
// kernel body on every compile pass and makes this source impossible to place in a
// multi-arch fat binary (it emits sm_90-only instructions during e.g. the sm_89 pass
// -> ptxas error). Removed so the existing `#if __CUDA_ARCH__ >= 900 ... #else assert
// #endif` guard takes effect per-arch: the real body is built only into the sm_90a
// cubin, other arches get a host-visible assert stub. The sm_90a pass is unchanged
// (nvcc defines __CUDA_ARCH__=900 there regardless).
namespace deep_gemm {
using namespace deep_gemm::sm90;
template <uint32_t kNumFormerIters, uint32_t kGap, uint32_t kEnd, typename func_t>
__device__ void dispatch_num_former_iters(uint32_t num_former_iters, const func_t& func) {
if (num_former_iters == kNumFormerIters) {
func(cute::Int<kNumFormerIters>{});
return;
}
if constexpr (kNumFormerIters + kGap <= kEnd)
dispatch_num_former_iters<kNumFormerIters + kGap, kGap, kEnd>(num_former_iters, func);
}
template <uint32_t SHAPE_M, uint32_t SHAPE_N, uint32_t SHAPE_K,
uint32_t kNumGroups,
uint32_t BLOCK_M, uint32_t BLOCK_N, uint32_t BLOCK_K,
uint32_t kSwizzleDMode,
uint32_t kNumStages, uint32_t kNumLastStages,
uint32_t kNumTMAThreads, uint32_t kNumMathThreads,
uint32_t kNumTMAMulticast, bool kIsTMAMulticastOnA,
uint32_t kNumSMs, GemmType kGemmType,
typename epilogue_type_t>
__global__ __launch_bounds__(kNumTMAThreads + kNumMathThreads, 1) void
sm90_fp8_gemm_1d2d_bias_impl(float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const __grid_constant__ cute::TmaDescriptor tensor_map_a,
const __grid_constant__ cute::TmaDescriptor tensor_map_b,
const __grid_constant__ cute::TmaDescriptor tensor_map_d,
const __grid_constant__ cute::TmaDescriptor tensor_map_sfa) {
// LT-PATCH: was `__CUDA_ARCH__ >= 900`. Tightened to Hopper-only (< 1000) so that in a
// multi-arch fat binary that also targets Blackwell (sm_100/sm_120), this wgmma body is
// NOT emitted for those passes (wgmma is sm_90a-only) -- they get the `#else` assert stub
// instead. Blackwell dispatches to the SM89 kernel at runtime, so the stub is never run.
#if (defined(__CUDA_ARCH__) and (__CUDA_ARCH__ >= 900) and (__CUDA_ARCH__ < 1000)) or defined(__CLION_IDE__)
// Scaling checks
DG_STATIC_ASSERT(BLOCK_K == 128, "Only support per-128-channel FP8 scaling");
DG_STATIC_ASSERT(constexpr_ceil_div(BLOCK_N, BLOCK_K) == 1 or (constexpr_gcd(BLOCK_N, BLOCK_K) == BLOCK_N - BLOCK_K), "Too much B scales in a single block");
// Types
using WGMMA = typename FP8MMASelector<BLOCK_N>::type;
using Barrier = cutlass::arch::ClusterTransactionBarrier;
DG_STATIC_ASSERT(BLOCK_M % WGMMA::M == 0, "Invalid block size");
// Overwrite shape constants if the compiler gives
shape_m = SHAPE_M != 0 ? SHAPE_M : shape_m;
shape_n = SHAPE_N != 0 ? SHAPE_N : shape_n;
shape_k = SHAPE_K != 0 ? SHAPE_K : shape_k;
// Shared memory
static constexpr bool kMustUseUniformedScaleB = (BLOCK_K % BLOCK_N == 0);
static constexpr uint32_t SMEM_D_SIZE = BLOCK_M * BLOCK_N * sizeof(__nv_bfloat16);
static constexpr uint32_t SMEM_A_SIZE_PER_STAGE = BLOCK_M * BLOCK_K * sizeof(__nv_fp8_e4m3);
static constexpr uint32_t SMEM_B_SIZE_PER_STAGE = BLOCK_N * BLOCK_K * sizeof(__nv_fp8_e4m3);
static constexpr uint32_t SMEM_SFA_SIZE_PER_STAGE = BLOCK_M * sizeof(float);
const uint32_t& shape_k_scales = ceil_div(shape_k, BLOCK_K);
const uint32_t& smem_sfb_size = align<uint32_t>(shape_k_scales * (kMustUseUniformedScaleB ? 1 : 2) * sizeof(float), sizeof(Barrier));
// Configs
const uint32_t num_total_k_blocks = ceil_div(shape_k, BLOCK_K);
const uint32_t warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
const uint32_t lane_idx = get_lane_idx();
// Prefetch TMA descriptors at the very beginning
if (warp_idx == kNumMathThreads / 32 and cute::elect_one_sync()) {
cute::prefetch_tma_descriptor(&tensor_map_a);
cute::prefetch_tma_descriptor(&tensor_map_b);
cute::prefetch_tma_descriptor(&tensor_map_sfa);
cute::prefetch_tma_descriptor(&tensor_map_d);
}
__syncwarp();
// Align to 1024 bytes for swizzle-128B
extern __shared__ __align__(1024) uint8_t smem_buffer[];
DG_STATIC_ASSERT(SMEM_D_SIZE % 1024 == 0, "Shared memory of A/B must be aligned to 1024 bytes");
// Data on shared memory
auto smem_d = reinterpret_cast<__nv_bfloat16*>(smem_buffer);
auto smem_a = PatternVisitor([&](const uint32_t& i) {
return reinterpret_cast<__nv_fp8_e4m3*>(smem_buffer + SMEM_D_SIZE + i * SMEM_A_SIZE_PER_STAGE);
});
auto smem_b = PatternVisitor([&](const uint32_t& i) {
return reinterpret_cast<__nv_fp8_e4m3*>(smem_buffer + SMEM_D_SIZE + kNumStages * SMEM_A_SIZE_PER_STAGE + i * SMEM_B_SIZE_PER_STAGE);
});
constexpr uint32_t SMEM_SF_OFFSET = SMEM_D_SIZE + kNumStages * (SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE);
auto smem_sfa = PatternVisitor([&](const uint32_t& i) {
return reinterpret_cast<float*>(smem_buffer + SMEM_SF_OFFSET + i * SMEM_SFA_SIZE_PER_STAGE);
});
auto smem_sfb = reinterpret_cast<float*>(smem_buffer + SMEM_SF_OFFSET + kNumStages * SMEM_SFA_SIZE_PER_STAGE);
// Fill barriers
auto barrier_start_ptr = reinterpret_cast<Barrier*>(reinterpret_cast<uint8_t*>(smem_sfb) + smem_sfb_size);
auto full_barriers = PatternVisitor([&](const uint32_t& i) { return barrier_start_ptr + i; });
auto empty_barriers = PatternVisitor([&](const uint32_t& i) { return barrier_start_ptr + kNumStages + i; });
// Initialize barriers
DG_STATIC_ASSERT(kNumTMAMulticast <= 32, "Too many TMA multicast");
if (warp_idx == kNumMathThreads / 32 + 1 and cute::elect_one_sync()) {
// NOTES: we always use `lane_idx` to arrive for the `lane_idx`-th CTA in the cluster,
// even with TMA multicast disabled, we want to make the behavior aligned
#pragma unroll
for (uint32_t i = 0; i < kNumStages; ++ i) {
full_barriers[i]->init(1);
empty_barriers[i]->init(kNumTMAMulticast * kNumMathThreads / 32);
}
// Make initialized barrier visible in async proxy
cutlass::arch::fence_barrier_init();
}
// Synchronize all threads to make barrier visible in normal memory model
(kNumTMAMulticast > 1) ? cute::cluster_sync() : __syncthreads();
// Register reconfigurations
constexpr uint32_t kNumTMARegisters = 40;
constexpr uint32_t kNumMathRegisters = 232;
// Block scheduler
uint32_t m_block_idx, n_block_idx;
auto scheduler = Scheduler<kGemmType, BLOCK_M, BLOCK_N, kNumGroups, kNumTMAMulticast, kIsTMAMulticastOnA, kNumSMs>(shape_m, shape_n, shape_k, grouped_layout);
// Pipeline and TMA phases
uint32_t stage_idx = 0, phase = 0;
auto advance_pipeline = [&](uint32_t& k_block_idx) {
++ k_block_idx;
// Flip phases only if reach the next first stage
stage_idx = stage_idx == kNumStages - 1 ? 0 : stage_idx + 1;
phase ^= stage_idx == 0;
};
if (warp_idx >= kNumMathThreads / 32) {
// TMA warp-group for loading data
cutlass::arch::warpgroup_reg_dealloc<kNumTMARegisters>();
// NOTES: only one thread (or warp) will be used
if (warp_idx == kNumMathThreads / 32 and cute::elect_one_sync()) {
// Persistently schedule over blocks
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
// Assign TMA multicast number into A and B
// NOTES: there may be additional odd rows/columns or cases where multicast is not possible.
const bool is_tma_multicast_valid = scheduler.is_tma_multicast_valid(m_block_idx);
const uint32_t num_tma_multicast_a = (kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
const uint32_t num_tma_multicast_b = (not kIsTMAMulticastOnA and is_tma_multicast_valid) ? kNumTMAMulticast : 1;
DG_STATIC_ASSERT(kNumTMAMulticast <= 2, "Scheduler does not support > 2 TMA multicast");
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
// Wait consumer release
empty_barriers[stage_idx]->wait(phase ^ 1);
// Issue TMA A
constexpr bool kWithGroupOffsetA = kGemmType == GemmType::MGroupedMasked;
auto& full_barrier = *full_barriers[stage_idx];
const uint32_t k_idx = k_block_idx * BLOCK_K;
tma_copy(&tensor_map_a, reinterpret_cast<uint64_t*>(&full_barrier),
smem_a[stage_idx], k_idx, scheduler.get_global_idx<kWithGroupOffsetA>(shape_m, BLOCK_M, m_block_idx),
num_tma_multicast_a);
tma_copy(&tensor_map_sfa, reinterpret_cast<uint64_t*>(&full_barrier),
smem_sfa[stage_idx], m_block_idx * BLOCK_M, scheduler.get_global_idx<kWithGroupOffsetA>(shape_k_scales, 1, k_block_idx),
num_tma_multicast_a);
// Issue TMA B
tma_copy(&tensor_map_b, reinterpret_cast<uint64_t*>(&full_barrier),
smem_b[stage_idx], k_idx, scheduler.get_global_idx<true>(shape_n, BLOCK_N, n_block_idx, m_block_idx),
num_tma_multicast_b);
full_barrier.arrive_and_expect_tx(SMEM_A_SIZE_PER_STAGE + SMEM_B_SIZE_PER_STAGE + SMEM_SFA_SIZE_PER_STAGE);
}
}
// To safely deconstruct distributed shared barriers, we need another round of empty waits
if constexpr (kNumTMAMulticast > 1) {
for (uint32_t i = 0; i < kNumStages; advance_pipeline(i))
empty_barriers[stage_idx]->wait(phase ^ 1);
}
}
} else {
// Math warp-groups for WGMMA
cutlass::arch::warpgroup_reg_alloc<kNumMathRegisters>();
// NOTES: use `__shfl_sync` to encourage NVCC to use unified registers
const auto math_wg_idx = __shfl_sync(0xffffffff, threadIdx.x / 128, 0);
const auto r_0 = warp_idx * 16 + lane_idx / 4, r_1 = r_0 + 8;
auto a_desc = make_smem_desc(smem_a[0] + math_wg_idx * WGMMA::M * BLOCK_K, 1);
auto b_desc = make_smem_desc(smem_b[0], 1);
const uint32_t a_desc_lo = __shfl_sync(0xffffffff, a_desc.reg32_[0], 0);
const uint32_t b_desc_lo = __shfl_sync(0xffffffff, b_desc.reg32_[0], 0);
// Persistently schedule over blocks
while (scheduler.get_next_block(m_block_idx, n_block_idx)) {
// Decide the number of scales B to load
DG_TRAP_ONLY_DEVICE_ASSERT(shape_n % 8 == 0);
uint32_t num_former_iters = BLOCK_N / 8, num_full_iters = num_former_iters;
if constexpr (not kMustUseUniformedScaleB) {
num_former_iters = min(BLOCK_N, BLOCK_K - n_block_idx * BLOCK_N % BLOCK_K) / 8;
num_full_iters = min(shape_n - n_block_idx * BLOCK_N, BLOCK_N) / 8;
}
uint32_t num_sfb = shape_k_scales * (num_former_iters >= num_full_iters ? 1 : 2);
// Load B scales with math warp-groups
// NOTES: except the first warp, we want to overlap loading B scales with TMA stores between tasks
if (threadIdx.x >= 32) {
auto num_previous_lines = scheduler.get_global_idx<true>(ceil_div(shape_n, BLOCK_K), 0, 0, m_block_idx);
auto local_sfb = sfb + (num_previous_lines + ((n_block_idx * BLOCK_N) / BLOCK_K)) * shape_k_scales;
#pragma unroll
for (uint32_t i = threadIdx.x - 32; i < num_sfb; i += kNumMathThreads - 32)
st_shared(smem_sfb + i, __ldg(local_sfb + i));
}
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
// Accumulation for WGMMA or CUDA promotion
constexpr uint32_t WAVE_BLOCK_M = WGMMA::M * (BLOCK_M <= 64 ? 1 : 2);
DG_STATIC_ASSERT(BLOCK_M % WAVE_BLOCK_M == 0, "Invalid block sizes");
float accum[WGMMA::kNumAccum], final_accum[WGMMA::kNumAccum * (BLOCK_M / WAVE_BLOCK_M)] = {0};
// Empty barrier arrival
auto empty_barrier_arrive = [&]() {
if constexpr (kNumTMAMulticast == 1) {
lane_idx == 0 ? empty_barriers[stage_idx]->arrive() : void();
} else {
auto target_cta = scheduler.is_peer_cta_alive ? lane_idx : cute::block_rank_in_cluster();
lane_idx < kNumTMAMulticast ? empty_barriers[stage_idx]->arrive(target_cta) : void();
}
};
// Skip useless computations
if (scheduler.is_computation_valid(m_block_idx, math_wg_idx * WGMMA::M)) {
// The compiler must know the dynamic variable `num_former_iters`'s real value
constexpr bool kShouldOptimize = BLOCK_K / constexpr_gcd(BLOCK_K, BLOCK_N) <= 4 and not kMustUseUniformedScaleB;
constexpr uint32_t kGap = constexpr_gcd(BLOCK_K, BLOCK_N) / 8;
constexpr uint32_t kEnd = kShouldOptimize ? BLOCK_K / 8 : 0;
// Dispatch `num_former_iters` and launch MMAs
dispatch_num_former_iters<0, kGap, kEnd>(kShouldOptimize ? num_former_iters : 0, [&](auto _) {
#pragma unroll 8
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
const auto& a_desc_base_lo = a_desc_lo + stage_idx * (SMEM_A_SIZE_PER_STAGE / 16);
const auto& b_desc_base_lo = b_desc_lo + stage_idx * (SMEM_B_SIZE_PER_STAGE / 16);
// Read B scales
float scale_b_0 = ld_shared(smem_sfb + k_block_idx), scale_b_1;
// NOTES: even some blocks do not need to read the second row, but we still load one to align with other blocks
if constexpr (not kMustUseUniformedScaleB)
scale_b_1 = ld_shared(smem_sfb + k_block_idx + shape_k_scales);
// Wait TMA arrivals
full_barriers[stage_idx]->wait(phase);
// TODO: remove some useless computation for unaligned Ms
#pragma unroll
for (uint32_t local_idx = 0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx) {
auto m_offset = local_idx * WAVE_BLOCK_M;
// Read A scales
// NOTES: all shared memory read must be prior to `warpgroup_arrive` to avoid next scheduled block polluting the results
auto scale_a_0 = ld_shared(smem_sfa[stage_idx] + r_0 + m_offset);
auto scale_a_1 = ld_shared(smem_sfa[stage_idx] + r_1 + m_offset);
// Commit WGMMA instructions
#pragma unroll
for (uint32_t i = 0; i < WGMMA::kNumAccum; ++ i)
warpgroup_fence_operand(accum[i]);
warpgroup_arrive();
#pragma unroll
for (uint32_t k = 0; k < BLOCK_K / WGMMA::K; ++ k) {
a_desc.reg32_[0] = a_desc_base_lo + (m_offset * BLOCK_K + k * WGMMA::K) / 16;
b_desc.reg32_[0] = b_desc_base_lo + k * WGMMA::K / 16;
WGMMA::wgmma(a_desc, b_desc, accum, k);
}
warpgroup_commit_batch();
#pragma unroll
for (uint32_t i = 0; i < WGMMA::kNumAccum; ++ i)
warpgroup_fence_operand(accum[i]);
warpgroup_wait<0>();
// Notify barrier arrival at the last warpgroup wave
if (local_idx == BLOCK_M / WAVE_BLOCK_M - 1)
empty_barrier_arrive();
// Promote with scales
// NOTES: making it as predicates is very important for performance, comparing to two loops
float scale_0_0 = scale_a_0 * scale_b_0, scale_1_0 = scale_a_1 * scale_b_0;
float scale_0_1, scale_1_1;
if constexpr (not kMustUseUniformedScaleB)
scale_0_1 = scale_a_0 * scale_b_1, scale_1_1 = scale_a_1 * scale_b_1;
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
#pragma unroll
for (uint32_t i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
// NOTES: for unrolled `num_former_iters` cases, we expect the compiler to automatically make it a constant
bool predicate = kMustUseUniformedScaleB or i < num_former_iters;
shifted_accum[i * 4 + 0] += (predicate ? scale_0_0 : scale_0_1) * accum[i * 4 + 0];
shifted_accum[i * 4 + 1] += (predicate ? scale_0_0 : scale_0_1) * accum[i * 4 + 1];
shifted_accum[i * 4 + 2] += (predicate ? scale_1_0 : scale_1_1) * accum[i * 4 + 2];
shifted_accum[i * 4 + 3] += (predicate ? scale_1_0 : scale_1_1) * accum[i * 4 + 3];
}
}
}
});
} else {
#pragma unroll
for (uint32_t k_block_idx = 0; k_block_idx < num_total_k_blocks; advance_pipeline(k_block_idx)) {
full_barriers[stage_idx]->wait(phase);
empty_barrier_arrive();
}
}
// TMA checks
constexpr uint32_t kNumElemBytes = sizeof(nv_bfloat16);
constexpr uint32_t TMA_D_BLOCK_N = kSwizzleDMode == 0 ? BLOCK_N : (kSwizzleDMode / kNumElemBytes);
constexpr uint32_t WGMMA_M_PER_WARP = WGMMA::M / 4;
DG_STATIC_ASSERT(BLOCK_M % 8 == 0, "Invalid swizzling atom");
DG_STATIC_ASSERT(BLOCK_N % TMA_D_BLOCK_N == 0 and BLOCK_N / TMA_D_BLOCK_N <= 32,
"Unaligned TMA store or too many TMA store instructions");
DG_STATIC_ASSERT(TMA_D_BLOCK_N % 8 == 0, "Invalid TMA block N");
// Wait last TMA store to be finished
float* bias_ptr = bias + n_block_idx*BLOCK_N + (lane_idx % 4) * 2;
#pragma unroll
for(uint32_t local_idx=0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx){
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
#pragma unroll
for (auto i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
shifted_accum[4*i + 0] += bias_ptr[8*i + 0];
shifted_accum[4*i + 1] += bias_ptr[8*i + 1];
shifted_accum[4*i + 2] += bias_ptr[8*i + 0];
shifted_accum[4*i + 3] += bias_ptr[8*i + 1];
}
}
if (threadIdx.x < BLOCK_N / TMA_D_BLOCK_N)
cute::tma_store_wait<0>();
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
// Write back to shared memory using STSM and issue TMA stores
DG_STATIC_ASSERT(WGMMA::kNumAccum % 4 == 0, "Invalid STSM x2 vectorization");
#pragma unroll
for (uint32_t local_idx = 0; local_idx < BLOCK_M / WAVE_BLOCK_M; ++ local_idx) {
auto m_offset = local_idx * WAVE_BLOCK_M;
auto shifted_accum = final_accum + WGMMA::kNumAccum * local_idx;
#pragma unroll
for (auto i = 0; i < WGMMA::kNumAccum / 4; ++ i) {
// Swizzle or padding into the correct address
uint8_t* smem_ptr = nullptr;
if constexpr (kSwizzleDMode > 0) {
// Calculate the swizzling atom offset and in-atom offset
constexpr uint32_t kNumBankGroupBytes = 16;
auto atom_offset = i / (TMA_D_BLOCK_N / 8), in_atom_offset = i % (TMA_D_BLOCK_N / 8);
// Calculate the index of the bank group to be written in the atom
auto bank_group_index = in_atom_offset + lane_idx * (kSwizzleDMode / kNumBankGroupBytes);
// Reshape the atom in another view and swizzle
// - original: `(BLOCK_M, kSwizzleDMode / kNumBankGroupBytes)`
// - new: `(BLOCK_M * kSwizzleDMode / kNumBankGroupBytes / 8, 8)`
constexpr bool kHasShortcut = (kSwizzleDMode / kNumBankGroupBytes) == 8;
auto row = kHasShortcut ? (in_atom_offset / 8 + lane_idx) : (bank_group_index / 8);
auto col = kHasShortcut ? (in_atom_offset) : (bank_group_index % 8);
col ^= row % (kSwizzleDMode / 16);
// Add back into the base pointer
// NOTES: think twice before modifying this, as changes may affect the number of instructions
smem_ptr = reinterpret_cast<uint8_t*>(smem_d) + // Base pointer
warp_idx * (WGMMA_M_PER_WARP * kSwizzleDMode) + // Warp offset
m_offset * kSwizzleDMode + // Wave offset
atom_offset * BLOCK_M * kSwizzleDMode + // Swizzle atom offset (constants)
row * (kNumBankGroupBytes * 8) + col * kNumBankGroupBytes; // In-atom offset
} else {
// No swizzling, just padding
smem_ptr = reinterpret_cast<uint8_t*>(smem_d + (m_offset + warp_idx * WGMMA_M_PER_WARP + lane_idx) * BLOCK_N + i * 8);
}
// NOTES: only 16 lanes' addresses are used
SM90_U32x2_STSM_N<nv_bfloat162>::copy(
__float22bfloat162_rn({shifted_accum[i * 4 + 0], shifted_accum[i * 4 + 1]}),
__float22bfloat162_rn({shifted_accum[i * 4 + 2], shifted_accum[i * 4 + 3]}),
smem_ptr
);
}
}
cute::tma_store_fence();
cutlass::arch::NamedBarrier::sync(kNumMathThreads, 0);
// Use TMA store to write back to global memory
// TODO: compatible with FP32 output
constexpr bool kWithGroupOffsetD = kGemmType == GemmType::MGroupedMasked;
DG_STATIC_ASSERT(kNumMathThreads >= BLOCK_N / TMA_D_BLOCK_N, "Too many TMA blocks");
if (threadIdx.x < BLOCK_N / TMA_D_BLOCK_N) {
auto in_block_n_offset = threadIdx.x * TMA_D_BLOCK_N;
auto smem_ptr = smem_d + in_block_n_offset * BLOCK_M;
cute::SM90_TMA_STORE_2D::copy(&tensor_map_d, smem_ptr,
epilogue_type_t::apply_index_n<TMA_D_BLOCK_N>(n_block_idx * BLOCK_N + in_block_n_offset),
scheduler.get_global_idx<kWithGroupOffsetD>(shape_m, BLOCK_M, m_block_idx));
cute::tma_store_arrive();
}
__syncwarp();
}
}
#else
if (blockIdx.x == 0 and threadIdx.x == 0)
DG_DEVICE_ASSERT(false and "This kernel only support sm_90a");
#endif
}
static cudaLaunchConfig_t construct_launch_config(const cudaStream_t& stream, const int& smem_size,
const dim3& grid_dim, const dim3& block_dim, const int& cluster_dim) {
cudaLaunchConfig_t config;
config.gridDim = grid_dim;
config.blockDim = block_dim;
config.dynamicSmemBytes = smem_size;
config.stream = stream;
config.numAttrs = 0;
config.attrs = nullptr;
// NOTES: must use `static` or the `attr` will be deconstructed
static cudaLaunchAttribute attr;
if (cluster_dim > 1) {
attr.id = cudaLaunchAttributeClusterDimension;
attr.val.clusterDim = {static_cast<unsigned>(cluster_dim), 1, 1};
config.attrs = &attr;
config.numAttrs = 1;
}
return config;
}
// static auto launch_kernel(auto kernel, const cudaLaunchConfig_t& config, float* sfb, float* bias, int* grouped_layout,
// uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
// const CUtensorMap tensor_map_a,
// const CUtensorMap tensor_map_b,
// const CUtensorMap tensor_map_d,
// const CUtensorMap tensor_map_sfa) {
// // void* ptr_args[] = {&sfb, &bias, &grouped_layout, &shape_m, &shape_n, &shape_k, &tensor_map_a, &tensor_map_b, &tensor_map_d, &tensor_map_sfa};
// return
// }
template<int N, int K>
void sm90_fp8_gemm_1d2d_bias_launch(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa){
dim3 grid{num_sms, 1, 1};
dim3 block{num_threads, 1, 1};
const auto config = construct_launch_config(stream, smem_size, grid, block, cluster_dim);
if(num_sms == 132){
auto kernel = &sm90_fp8_gemm_1d2d_bias_impl<0, N, K, 1, 256, 128, 128, 128, 3, (K / 128) % 3, 128, 256, 2, true, 132, GemmType::Normal, EpilogueIdentity>;
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
cudaLaunchKernelEx(&config, kernel, sfb, bias, grouped_layout, shape_m, shape_n, shape_k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);
} else if(num_sms == 116) {
auto kernel = &sm90_fp8_gemm_1d2d_bias_impl<0, N, K, 1, 256, 128, 128, 128, 3, (K / 128) % 3, 128, 256, 2, true, 116, GemmType::Normal, EpilogueIdentity>;
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
cudaLaunchKernelEx(&config, kernel, sfb, bias, grouped_layout, shape_m, shape_n, shape_k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);
} else if (num_sms == 100) {
auto kernel = &sm90_fp8_gemm_1d2d_bias_impl<0, N, K, 1, 256, 128, 128, 128, 3, (K / 128) % 3, 128, 256, 2, true, 100, GemmType::Normal, EpilogueIdentity>;
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
cudaLaunchKernelEx(&config, kernel, sfb, bias, grouped_layout, shape_m, shape_n, shape_k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);
} else {
// The supported SM counts are exactly the branches above (the only kernels
// instantiated). Fail loudly instead of falling through with no launch,
// which would leave the output buffer uninitialized.
throw std::runtime_error("Unsupported num_sms=" + std::to_string(num_sms)
+ " (blockwise SM90 GEMM is built for 132, 116, and 100 SMs)");
}
// launch_kernel(kernel, config, sfb, bias, grouped_layout, shape_m, shape_n, shape_k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);
}
template void sm90_fp8_gemm_1d2d_bias_launch<2048, 2048>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<4096, 2048>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<8192, 2048>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<16384, 2048>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<2048, 4096>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<4096, 4096>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<8192, 4096>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<16384, 4096>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<2048, 8192>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<4096, 8192>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<8192, 8192>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<16384, 8192>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<2048, 16384>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<4096, 16384>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<8192, 16384>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
template void sm90_fp8_gemm_1d2d_bias_launch<16384, 16384>(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
}; // namespace deep_gemm
@@ -0,0 +1,287 @@
#include "cutlass/cutlass.h"
#include "cutlass/layout/layout.h"
#include <cute/tensor.hpp>
#include <c10/cuda/CUDAException.h>
#include <torch/extension.h>
#include <torch/python.h>
#include <cuda_runtime.h>
#include <iostream>
#include "kernel_traits.cuh"
#include "static_switch.h"
namespace sm89{
using namespace cute;
__device__ static void copy_1d(float* gmem_src, float* smem_dst)
{
uint32_t smem_int_ptr = cast_smem_ptr_to_uint((void*)smem_dst);
asm volatile("cp.async.ca.shared.global.L2::128B [%0], [%1], %2;\n"
:: "r"(smem_int_ptr),
"l"(gmem_src),
"n"(sizeof(float)));
}
template <typename KernelTraits=gemm_traits<128, 256, 2, 4096, 2, 4, true, half_t, bfloat16_t>>
__global__ void gemm_fp8_kernel(float_e4m3_t* Aptr, float* sfa, float_e4m3_t* Bptr, float* sfb, float* bias_ptr, void* out, int M, int N, int K, int TMA_ALIGNED_M){
using output_t = typename KernelTraits::out_t;
using SmemLayoutA = typename KernelTraits::SmemLayoutA;
using SmemLayoutB = typename KernelTraits::SmemLayoutB;
using SmemLayoutC = typename KernelTraits::SmemLayoutC;
constexpr int BM = KernelTraits::BM;
constexpr int BN = KernelTraits::BN;
constexpr int BK = KernelTraits::BK;
constexpr int Ksfa = KernelTraits::KSF;
constexpr bool has_bias = KernelTraits::HasBias;
extern __shared__ float smem_[];
float *bias_shm = smem_;
float *sfa_shm = reinterpret_cast<float*>(bias_shm + cosize(typename KernelTraits::SmemLayoutBias{}));
output_t* C_shm = reinterpret_cast<output_t*>(sfa_shm + cosize(typename KernelTraits::SmemLayoutSFA{}));
float_e4m3_t* A_shm = reinterpret_cast<float_e4m3_t*>(sfa_shm + cosize(typename KernelTraits::SmemLayoutSFA{}));
float_e4m3_t* B_shm = reinterpret_cast<float_e4m3_t*>(A_shm + cosize(SmemLayoutA{}));
int idx = threadIdx.x;
int ix = blockIdx.x;
int iy = blockIdx.y;
// sfa += BM * iy;
sfb += KernelTraits::NUM_SFB_PER_STEP * ix * Ksfa;
output_t* Cptr = reinterpret_cast<output_t*>(out);
Tensor A = make_tensor(make_gmem_ptr(Aptr), make_shape(M, K), make_stride(K, Int<1>{}));
Tensor B = make_tensor(make_gmem_ptr(Bptr), make_shape(N, K), make_stride(K, Int<1>{}));
Tensor D = make_tensor(make_gmem_ptr(Cptr), make_shape(M, N), make_stride(N, Int<1>{}));
Tensor SFA = make_tensor(make_gmem_ptr(sfa), make_shape(M, Ksfa), make_stride(Int<1>{}, TMA_ALIGNED_M));
Tensor gA = local_tile(A, make_tile(Int<BM>{}, Int<BK>{}), make_coord(iy, _));
Tensor gB = local_tile(B, make_tile(Int<BN>{}, Int<BK>{}), make_coord(ix, _));
Tensor gD = local_tile(D, make_tile(Int<BM>{}, Int<BN>{}), make_coord(iy, ix));
Tensor gSFA = local_tile(SFA, make_tile(Int<BM>{}, Int<1>{}), make_coord(iy, _));
auto sBias = make_tensor(make_smem_ptr(bias_shm), typename KernelTraits::SmemLayoutBias{});
if constexpr (has_bias){
Tensor Bias = make_tensor(make_gmem_ptr(bias_ptr), make_shape(_1{}, N), make_stride(N, Int<1>{}));
Tensor gBias = local_tile(Bias, make_tile(Int<1>{}, Int<BN>{}), make_coord(_, ix));
typename KernelTraits::G2SBiasCopy g2s_bias_copy;
auto g2s_bias_thr_copy = g2s_bias_copy.get_slice(idx);
auto tCBiasgBias = g2s_bias_thr_copy.partition_S(gBias);
auto tCBiassBias = g2s_bias_thr_copy.partition_D(sBias);
if(idx < BN){
copy_1d((float*)&gBias(0) + idx, (float*)&sBias(0) + idx);
}
}
auto sSFA = make_tensor(make_smem_ptr(sfa_shm), typename KernelTraits::SmemLayoutSFA{});
auto sA = make_tensor(make_smem_ptr(A_shm), SmemLayoutA{});
auto sB = make_tensor(make_smem_ptr(B_shm), SmemLayoutB{});
typename KernelTraits::MMATile tiled_mma;
auto thr_mma = tiled_mma.get_slice(threadIdx.x);
auto tCrA = thr_mma.partition_fragment_A(gA(_, _, 0));
auto tCrB = thr_mma.partition_fragment_B(gB(_, _, 0));
auto tCrD = thr_mma.partition_fragment_C(gD);
clear(tCrD);
auto tCrD_fp32 = make_tensor_like<float>(tCrD);
clear(tCrD_fp32);
typename KernelTraits::G2STiledCopy g2s_tiled_copy;
auto g2s_thr_copy = g2s_tiled_copy.get_slice(idx);
auto tAgA_copy = g2s_thr_copy.partition_S(gA);
auto tAsA_copy = g2s_thr_copy.partition_D(sA);
auto tBgB_copy = g2s_thr_copy.partition_S(gB);
auto tBsB_copy = g2s_thr_copy.partition_D(sB);
auto s2r_tiled_copy_a = make_tiled_copy_A(typename KernelTraits::S2RCopyAtomA{}, tiled_mma);
auto s2r_thr_copy_a = s2r_tiled_copy_a.get_slice(idx);
auto tAsA = s2r_thr_copy_a.partition_S(sA);
auto tCrA_view = s2r_thr_copy_a.retile_D(tCrA);
auto s2r_tiled_copy_b = make_tiled_copy_B(typename KernelTraits::S2RCopyAtomB{}, tiled_mma);
auto s2r_thr_copy_b = s2r_tiled_copy_b.get_slice(idx);
auto tBsB = s2r_thr_copy_b.partition_S(sB);
auto tCrB_view = s2r_thr_copy_b.retile_D(tCrB);
auto cA = make_identity_tensor(make_shape(size<0>(sA), size<1>(sA)));
auto tAcA = g2s_thr_copy.partition_S(cA);
int residual = M - iy*BM;
int itile_to_read = 0;
int ismem_read = 0;
int ismem_write = 0;
int ismem_read_sfa = 0;
constexpr int kStages = KernelTraits::KStages;
#pragma unroll
for(int istage=0; istage<kStages - 1; ++istage){
for (size_t m = 0; m < size<1>(tAsA_copy); m++)
{
for (size_t k = 0; k < size<2>(tAsA_copy); k++)
{
if(get<0>(tAcA(0, m, k)) < residual){
cute::copy(g2s_tiled_copy, tAgA_copy(_, m, k, istage), tAsA_copy(_, m, k, istage));
}
}
}
if(idx < KernelTraits::THREADS_SFA_COPY && (BM * iy + idx * KernelTraits::SFA_ELEMS_PER_COPY < M)) {
copy_1d((float*)&gSFA(0, 0, istage) + idx*KernelTraits::SFA_ELEMS_PER_COPY, (float*)&sSFA(0, istage) + idx*KernelTraits::SFA_ELEMS_PER_COPY);
}
cute::copy(g2s_tiled_copy, tBgB_copy(_, _, _, istage), tBsB_copy(_, _, _, istage));
cp_async_fence();
++itile_to_read;
++ismem_write;
}
cp_async_wait<kStages - 2>();
__syncthreads();
cute::copy(s2r_tiled_copy_a, tAsA(_, _, 0, ismem_read), tCrA_view(_, _, 0));
cute::copy(s2r_tiled_copy_b, tBsB(_, _, 0, ismem_read), tCrB_view(_, _, 0));
static constexpr int nk = size<2>(tCrA);
auto sfa_tv = typename KernelTraits::SFAThreadLayout{};
static constexpr int NTILES = KernelTraits::NTiles;
#pragma unroll
for(int itile = 0; itile < NTILES; itile++){
clear(tCrD);
#pragma unroll
for(int ik = 0; ik < nk; ik++){
int ik_next = (ik + 1) % nk;
if(ik == nk - 1) {
cp_async_wait<kStages - 2>();
__syncthreads();
ismem_read = (ismem_read + 1) % kStages;
}
cute::copy(s2r_tiled_copy_a, tAsA(_, _, ik_next, ismem_read), tCrA_view(_, _, ik_next));
cute::copy(s2r_tiled_copy_b, tBsB(_, _, ik_next, ismem_read), tCrB_view(_, _, ik_next));
if(ik == 0){
if(itile_to_read < NTILES){
for (size_t m = 0; m < size<1>(tAsA_copy); m++)
{
for (size_t k = 0; k < size<2>(tAsA_copy); k++)
{
if(get<0>(tAcA(0, m, k)) < residual){
cute::copy(g2s_tiled_copy, tAgA_copy(_, m, k, itile_to_read), tAsA_copy(_, m, k, ismem_write));
}
}
}
cute::copy(g2s_tiled_copy, tBgB_copy(_, _, _, itile_to_read), tBsB_copy(_, _, _, ismem_write));
if(idx < KernelTraits::THREADS_SFA_COPY && (BM * iy + idx * KernelTraits::SFA_ELEMS_PER_COPY < M)) {
copy_1d((float*)&gSFA(0, 0, itile_to_read) + idx * KernelTraits::SFA_ELEMS_PER_COPY, (float*)&sSFA(0, ismem_write) + idx*KernelTraits::SFA_ELEMS_PER_COPY);
}
++itile_to_read;
ismem_write = (ismem_write + 1) % kStages;
}
cp_async_fence();
}
cute::gemm(tiled_mma, tCrD, tCrA(_, _, ik), tCrB(_, _, ik), tCrD);
}
int sf_ind = itile / KernelTraits::TILES_PER_BLOCK;
float sfb_val = sfb[sf_ind];
#pragma unroll
for(int i = 0; i < size<1>(tCrD); i++){ // (MMA, MMA_M, MMA_N) = (4, 4, 4)
float sfa_val_1 = sSFA(sfa_tv(idx) + i * KernelTraits::MMA_WARP_M, ismem_read_sfa);
float sfa_val_2 = sSFA(sfa_tv(idx) + 8 + i * KernelTraits::MMA_WARP_M, ismem_read_sfa);
#pragma unroll
for(int j = 0; j < size<2>(tCrD); j++){
tCrD_fp32(0, i, j) += sfa_val_1 * sfb_val * float(tCrD(0, i, j));
tCrD_fp32(1, i, j) += sfa_val_1 * sfb_val * float(tCrD(1, i, j));
tCrD_fp32(2, i, j) += sfa_val_2 * sfb_val * float(tCrD(2, i, j));
tCrD_fp32(3, i, j) += sfa_val_2 * sfb_val * float(tCrD(3, i, j));
}
}
ismem_read_sfa = (ismem_read_sfa + 1) % kStages;
__syncthreads();
}
auto tCrBias = make_tensor<float>(Layout<Shape<_2, Int<size<2>(tCrD_fp32)>>>{});
auto bias_threads = typename KernelTraits::BiasThreadLayout{};
if constexpr (has_bias){
#pragma unroll
for(int i = 0; i<size<2>(tCrD_fp32); i++){
tCrBias(0, i) = sBias(bias_threads(idx) + i * KernelTraits::MMA_WARP_N);
tCrBias(1, i) = sBias(1 + bias_threads(idx) + i * KernelTraits::MMA_WARP_N);
}
#pragma unroll
for(int i = 0; i<size<1>(tCrD_fp32); i++){
#pragma unroll
for (int j = 0; j < size<2>(tCrD_fp32) ; j++)
{
tCrD_fp32(0, i, j) += tCrBias(0, j);
tCrD_fp32(1, i, j) += tCrBias(1, j);
tCrD_fp32(2, i, j) += tCrBias(0, j);
tCrD_fp32(3, i, j) += tCrBias(1, j);
}
}
}
auto sC = make_tensor(make_smem_ptr(C_shm), SmemLayoutC{});
auto r2s_tiled_copy_c = make_tiled_copy_C(typename KernelTraits::R2SCopyAtomC{}, tiled_mma);
auto r2s_thr_copy_c = r2s_tiled_copy_c.get_slice(idx);
auto tCrC_r2s = r2s_thr_copy_c.retile_S(tCrD_fp32);
auto tCsC_r2s = r2s_thr_copy_c.partition_D(sC);
typename KernelTraits::S2GCopyC s2g_tiled_copy_c;
auto s2g_thr_copy_c = s2g_tiled_copy_c.get_thread_slice(idx);
auto tCsC_s2g = s2g_thr_copy_c.partition_S(sC);
auto tCgC_s2g = s2g_thr_copy_c.partition_D(gD);
int pipe = size<2>(tCsC_r2s);
auto cC = make_identity_tensor(make_shape(size<0>(gD), size<1>(gD)));
auto tCcC = s2g_thr_copy_c.partition_D(cC);
for(int i = 0; i< size<1>(tCrC_r2s); i++){
for(int j = 0; j < size<2>(tCrC_r2s); j+=pipe){
for(int step = 0; step < pipe; ++step){
auto fragment = make_tensor_like<output_t>(tCrC_r2s(_, i, j + step));
cute::copy(tCrC_r2s(_, i, j + step), fragment);
cute::copy(r2s_tiled_copy_c, fragment, tCsC_r2s(_, 0, step));
}
__syncthreads();
if (get<0>(tCcC(0, i, j / pipe)) < residual){
cute::copy(s2g_tiled_copy_c, tCsC_s2g(_, 0, 0), tCgC_s2g(_, i, j / pipe));
}
__syncthreads();
}
}
}
template <bool has_bias, typename accum_type>
void fp8_kernel_launch(void* Aptr, void* sfa, void* Bptr, void* sfb, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream) {
int TMA_ALIGNED_M = ((M + sizeof(float) - 1) / sizeof(float)) * sizeof(float); // SIZEOF(float) = 4
BLOCK_K_SWITCH(K_, M_SWITCH(
using KernelTraits = gemm_traits<BM, BN, 3, K_, WARP_ROW, WARP_COL, has_bias, accum_type, bfloat16_t>;
auto kernel = &gemm_fp8_kernel<KernelTraits>;
int BX = (N + KernelTraits::BN - 1) / KernelTraits::BN;
int BY = (M + KernelTraits::BM - 1) / KernelTraits::BM;
dim3 block(KernelTraits::NUM_THREADS);
dim3 gridDim(BX, BY);
cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, KernelTraits::SmemSize);
kernel<<<gridDim, KernelTraits::NUM_THREADS, KernelTraits::SmemSize, stream>>>((float_e4m3_t*)Aptr, (float*)sfa, (float_e4m3_t*)Bptr, (float*)sfb, (float*)bias_ptr, out, M, N, K, TMA_ALIGNED_M);
C10_CUDA_KERNEL_LAUNCH_CHECK();))
}
template<bool use_fast_accum>
void fp8_bias_gemm_cuda(void* Aptr, void* SFA, void* Bptr, void* SFB, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream){
using accum_type = std::conditional_t<use_fast_accum, half_t, float>;
fp8_kernel_launch<true, accum_type>(Aptr, SFA, Bptr, SFB, bias_ptr, out, M, N, K, stream);
}
// template<bool use_fast_accum>
// void fp8_gemm_cuda(void* Aptr, void* SFA, void* Bptr, void* SFB, void* out, int M, int N, int K, cudaStream_t stream){
// using accum_type = std::conditional_t<use_fast_accum, half_t, float>;
// BLOCK_K_SWITCH(num_acc_upcast_steps, fp8_kernel_launch<false, num_acc_upcast_steps, accum_type>(Aptr, SFA, Bptr, SFB, nullptr, out, M, N, K, stream);)
// }
// template void fp8_gemm_cuda<true>(void* Aptr, void* SFA, void* Bptr, void* SFB, void* out, int M, int N, int K, cudaStream_t stream);
// template void fp8_gemm_cuda<false>(void* Aptr, void* SFA, void* Bptr, void* SFB, void* out, int M, int N, int K, cudaStream_t stream);
template void fp8_bias_gemm_cuda<true>(void* Aptr, void* SFA, void* Bptr, void* SFB, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream);
template void fp8_bias_gemm_cuda<false>(void* Aptr, void* SFA, void* Bptr, void* SFB, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream);
}; // namespace sm89
@@ -0,0 +1,121 @@
#pragma once
#include <cute/tensor.hpp>
#include <cutlass/cutlass.h>
#include <cutlass/layout/layout.h>
#include <cutlass/numeric_types.h>
#include "mma_sm89_fp16.hpp"
#include "mma_traits_sm89_fp16.hpp"
using namespace cute;
template<int BYTES> struct BytesToType {};
template<> struct BytesToType<4> {
using Type = uint32_t;
static_assert(sizeof(Type) == 4);
};
template<> struct BytesToType<2> {
using Type = uint16_t;
static_assert(sizeof(Type) == 2);
};
template<int BM_, int BN_, int KStages_, int K_, int WARP_ROW_=2, int WARP_COL_=2, bool HasBias_=false, typename accum_t_=cutlass::half_t, typename out_t_=cutlass::bfloat16_t>
struct gemm_traits {
static constexpr int BLOCK_SIZE = 128;
static constexpr int K = K_;
static constexpr int BM = BM_;
static constexpr int BN = BN_;
static constexpr int BK = 128;
static constexpr int TILES_PER_BLOCK = BLOCK_SIZE / BK;
static constexpr int NUM_SFB_PER_STEP = BN / BLOCK_SIZE;
static constexpr int NTiles = K / BK;
static constexpr int KSF = K / BLOCK_SIZE;
static constexpr int KStages = KStages_;
static constexpr int WARP_ROW = WARP_ROW_;
static constexpr int WARP_COL = WARP_COL_;
static constexpr int NUM_WARPS = WARP_ROW * WARP_COL;
static constexpr int NUM_THREADS = NUM_WARPS * 32;
static constexpr int MMA_WARP_M = WARP_ROW * 16;
static constexpr int MMA_WARP_N = WARP_COL * 8;
static constexpr int MMA_WARP_K = 32;
using accum_t = accum_t_;
using out_t = out_t_;
using SwizzleLayoutO = std::conditional_t<
std::is_same_v<out_t_, cutlass::bfloat16_t>,
Swizzle<3, 3, 3>,
Swizzle<2, 4, 3>
>;
using SwizzleLayoutAB = Swizzle<2, 4, 3>;
using MMA_Atom_SM89 = std::conditional_t<
std::is_same_v<accum_t, cutlass::half_t>,
MMA_Atom<SM89_16x8x32_F16E4M3E4M3F16_TN>,
MMA_Atom<SM89_16x8x32_F32E4M3E4M3F32_TN>
>;
static constexpr int INPUT_ELEMS_PER_COPY = sizeof(uint128_t) / sizeof(float_e4m3_t);
static constexpr int OUTPUT_ELEMS_PER_COPY = sizeof(uint128_t) / sizeof(out_t_);
static constexpr int THREADS_PER_ROW = BK / INPUT_ELEMS_PER_COPY;
using GMEMLayout = Layout< Shape <Int<NUM_THREADS / THREADS_PER_ROW>, Int<THREADS_PER_ROW>>, Stride<Int<THREADS_PER_ROW>, _1>>;
using G2SCopyAtom = Copy_Atom<SM80_CP_ASYNC_CACHEGLOBAL<cute::uint128_t>, float_e4m3_t>;
using G2STiledCopy = decltype(
make_tiled_copy(
G2SCopyAtom{},
GMEMLayout{},
Layout<Shape<_1, Int<INPUT_ELEMS_PER_COPY>>>{}
)
);
using S2RCopyAtomA = Copy_Atom<SM75_U32x4_LDSM_N, float_e4m3_t>;
using S2RCopyAtomB = Copy_Atom<SM75_U32x2_LDSM_N, float_e4m3_t>;
using SmemLayoutAtom = decltype(composition(
Swizzle<2, 4, 3>{},
make_layout(make_shape(Int<8>{}, Int<BK>{}),
make_stride(Int<BK>{}, Int<1>{}))));
using SmemLayoutA = decltype(
tile_to_shape(SmemLayoutAtom{}, make_shape(Int<BM>{}, Int<BK>{}, Int<KStages>{}))
);
using SmemLayoutB = decltype(
tile_to_shape(SmemLayoutAtom{}, make_shape(Int<BN>{}, Int<BK>{}, Int<KStages>{}))
);
using MMATile = decltype(
make_tiled_mma(
MMA_Atom_SM89{},
Layout<Shape<Int<WARP_ROW>, Int<WARP_COL>, _1>>{},
Tile<Int<MMA_WARP_M>, Int<MMA_WARP_N>, Int<MMA_WARP_K>>{}
)
);
static constexpr int ELEMS_PER_TILE = MMA_WARP_M * MMA_WARP_N;
static constexpr int NUM_ELEMS_PER_WRITE = NUM_THREADS * sizeof(cute::uint128_t) / sizeof(out_t_);
static constexpr int OUT_PIPE = NUM_ELEMS_PER_WRITE / ELEMS_PER_TILE;
// using SmemLayoutC = Layout<Shape<Int<BM>, Int<BN>>, Stride<Int<BN>, Int<1>>>;
using SmemLayoutC = decltype(
make_layout(
make_shape(Int<MMA_WARP_M>{}, Int<MMA_WARP_N*OUT_PIPE>{}),
make_stride(Int<MMA_WARP_N*OUT_PIPE>{}, Int<1>{})
)
);
static constexpr int THREADS_PER_ROW_WRITE = MMA_WARP_N * OUT_PIPE / OUTPUT_ELEMS_PER_COPY;
using R2SCopyAtomC = Copy_Atom<UniversalCopy<typename BytesToType<2*sizeof(out_t)>::Type>, out_t>;
using S2GCopyAtomC = Copy_Atom<UniversalCopy<cute::uint128_t>, out_t>;
using S2GCopyC = decltype(make_tiled_copy(S2GCopyAtomC{},
make_layout(make_shape(Int<NUM_THREADS / THREADS_PER_ROW_WRITE>{}, Int<THREADS_PER_ROW_WRITE>{}),
make_stride(Int<THREADS_PER_ROW_WRITE>{}, Int<1>{})),
make_layout(make_shape(Int<1>{}, Int<OUTPUT_ELEMS_PER_COPY>{}))));
using G2SBiasCopyAtom = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<float>, float>;
using G2SBiasCopy = decltype(make_tiled_copy(G2SBiasCopyAtom{}, make_layout(
make_shape(Int<1>{},Int<BN>{}), make_stride(Int<BN>{}, Int<1>{})),
make_layout(make_shape(Int<1>{},Int<1>{}), make_stride(Int<1>{}, Int<1>{}))));
using sfa_copy_vtype = float;
static constexpr int SFA_ELEMS_PER_COPY = sizeof(sfa_copy_vtype)/sizeof(float);
static constexpr int THREADS_SFA_COPY = BM * sizeof(float) / sizeof(sfa_copy_vtype);
// using G2SSFACopyAtom = Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS<cute::uint128_t>, float>;
static constexpr bool HasBias = HasBias_;
using SmemLayoutBias = Layout<Shape<Int<1>, Int<BN>>, Stride<Int<BN>, Int<1>>>;
using SmemLayoutSFA = Layout<Shape<Int<BM>, Int<KStages>>, Stride<Int<1>, Int<BM>>>;
using BiasThreadLayout = Layout<Shape<Shape<_4, _8>, Shape<Int<WARP_ROW>, Int<WARP_COL>>>, Stride<Stride<_2, _0>, Stride<_0, _8>>>;
using SFAThreadLayout = Layout<Shape<Shape<_4, _8>, Shape<Int<WARP_ROW>, Int<WARP_COL>>>, Stride<Stride<_0, _1>, Stride<_16, _0>>>;
static constexpr int SmemSize = cute::max(cute::cosize(SmemLayoutA{})+cute::cosize(SmemLayoutB{}), cute::cosize(SmemLayoutC{})*sizeof(out_t)) + cute::cosize(SmemLayoutBias{}) * sizeof(float) + cute::cosize(SmemLayoutSFA{})*sizeof(float);
};
@@ -0,0 +1,84 @@
#pragma once
#include <cute/config.hpp>
#include <cute/arch/mma.hpp>
////////////////////////////////////////////////////////////////////////////////
#if (__CUDACC_VER_MAJOR__ > 12) || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ >= 4)
# define CUTE_ARCH_MMA_F32_SM89_SUPPORTED
#endif
#if (__CUDACC_VER_MAJOR__ > 12) || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ >= 8)
# define CUTE_ARCH_MMA_F16_SM89_SUPPORTED
#endif
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 890)
# if defined(CUTE_ARCH_MMA_F32_SM89_SUPPORTED)
# define CUTE_ARCH_MMA_F32_SM89_ENABLED
# endif
# if defined(CUTE_ARCH_MMA_F16_SM89_SUPPORTED)
# define CUTE_ARCH_MMA_F16_SM89_ENABLED
# endif
#endif
namespace cute {
struct SM89_16x8x32_F32E4M3E4M3F32_TN
{
using DRegisters = float[4];
using ARegisters = uint32_t[4];
using BRegisters = uint32_t[2];
using CRegisters = float[4];
CUTE_HOST_DEVICE static void
fma(float & d0, float & d1, float & d2, float & d3,
uint32_t const& a0, uint32_t const& a1, uint32_t const& a2, uint32_t const& a3,
uint32_t const& b0, uint32_t const& b1,
float const& c0, float const& c1, float const& c2, float const& c3)
{
#if defined(CUTE_ARCH_MMA_F32_SM89_ENABLED)
asm(
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
"{%0,%1,%2,%3}, {%4,%5,%6,%7}, {%8,%9}, {%10,%11,%12,%13};\n"
: "=f"(d0), "=f"(d1), "=f"(d2), "=f"(d3)
:
"r"(a0), "r"(a1), "r"(a2), "r"(a3),
"r"(b0), "r"(b1),
"f"(c0), "f"(c1), "f"(c2), "f"(c3)
);
#else
CUTE_INVALID_CONTROL_PATH("Attempting to use SM89_16x8x32_F32E4M3E4M3F32_TN without CUTE_ARCH_MMA_F32_SM89_ENABLED");
#endif
}
};
// MMA 16x8x32 TN
struct SM89_16x8x32_F16E4M3E4M3F16_TN
{
using DRegisters = uint32_t[2];
using ARegisters = uint32_t[4];
using BRegisters = uint32_t[2];
using CRegisters = uint32_t[2];
CUTE_HOST_DEVICE static void
fma(uint32_t & d0, uint32_t & d1,
uint32_t const& a0, uint32_t const& a1, uint32_t const& a2, uint32_t const& a3,
uint32_t const& b0, uint32_t const& b1,
uint32_t const& c0, uint32_t const& c1)
{
#if defined(CUTE_ARCH_MMA_F16_SM89_ENABLED)
asm(
"mma.sync.aligned.m16n8k32.row.col.f16.e4m3.e4m3.f16 "
"{%0,%1}, {%2,%3,%4,%5}, {%6,%7}, {%8,%9};\n"
: "=r"(d0), "=r"(d1)
:
"r"(a0), "r"(a1), "r"(a2), "r"(a3),
"r"(b0), "r"(b1),
"r"(c0), "r"(c1)
);
#else
CUTE_INVALID_CONTROL_PATH("Attempting to use SM89_16x8x32_F32E4M3E4M3F32_TN without CUTE_ARCH_MMA_F16_SM89_ENABLED");
#endif
}
};
}
@@ -0,0 +1,50 @@
#pragma once
#include <cute/atom/mma_traits.hpp>
#include <cute/layout.hpp>
#include <cute/numeric/numeric_types.hpp>
#include "mma_sm89_fp16.hpp"
namespace cute
{
namespace {
// (T32,V4) -> (M16,N8)
using SM80_16x8_Row = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
}
template <>
struct MMA_Traits<SM89_16x8x32_F32E4M3E4M3F32_TN> {
using ValTypeD = float;
using ValTypeA = float_e4m3_t;
using ValTypeB = float_e4m3_t;
using ValTypeC = float;
using Shape_MNK = Shape<_16,_8,_32>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _4,_2, _2>>,
Stride<Stride<_64,_1>,Stride<_16,_8,_256>>>;
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_4, _2>>,
Stride<Stride<_32,_1>,Stride<_8,_128>>>;
using CLayout = SM80_16x8_Row;
};
template <>
struct MMA_Traits<SM89_16x8x32_F16E4M3E4M3F16_TN> {
using ValTypeD = half_t;
using ValTypeA = float_e4m3_t;
using ValTypeB = float_e4m3_t;
using ValTypeC = half_t;
using Shape_MNK = Shape<_16,_8,_32>;
using ThrID = Layout<_32>;
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _4,_2, _2>>,
Stride<Stride<_64,_1>,Stride<_16,_8,_256>>>;
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_4, _2>>,
Stride<Stride<_32,_1>,Stride<_8,_128>>>;
using CLayout = SM80_16x8_Row;
};
}
@@ -0,0 +1,35 @@
#pragma once
#define BOOL_SWITCH(COND, CONST_NAME, ...) \
if (COND) { \
constexpr static bool CONST_NAME = true; \
__VA_ARGS__ \
} else { \
constexpr static bool CONST_NAME = false; \
__VA_ARGS__ \
}
//K/128
#define BLOCK_K_SWITCH(COSNT_NAME, ...) \
if (K == 2048) { \
constexpr static int COSNT_NAME = 2048; \
__VA_ARGS__ \
} \
else if (K == 4096) { \
constexpr static int COSNT_NAME = 4096; \
__VA_ARGS__ \
} else if (K == 8192) { \
constexpr static int COSNT_NAME = 8192; \
__VA_ARGS__ \
} else if (K == 16384) { \
constexpr static int COSNT_NAME = 16384; \
__VA_ARGS__ \
} else { \
TORCH_CHECK(false, "Unsupported K value: ", K); \
}
#define M_SWITCH(...) \
constexpr static int BM = 64; \
constexpr static int BN = 128; \
constexpr static int WARP_ROW = 2; \
constexpr static int WARP_COL = 4; \
__VA_ARGS__
@@ -0,0 +1,206 @@
#pragma once
#include <cuda.h>
#include <cuda_runtime.h>
#include <torch/python.h>
#include "exceptions.hpp"
namespace blockwise {
template <typename T>
static T ceil_div(const T& a, const T& b) {
return (a + b - 1) / b;
}
template <typename T>
static constexpr T align(const T& a, const T& b) {
return ceil_div(a, b) * b;
}
static int get_tma_aligned_size(const int& x, const int& element_size) {
constexpr int kNumTMAAlignmentBytes = 16;
DG_HOST_ASSERT(kNumTMAAlignmentBytes % element_size == 0);
return align(x, kNumTMAAlignmentBytes / element_size);
}
static std::pair<int, int> get_inner_outer_dims(const cute::UMMA::Major& major, const int& k, const int& mn) {
return major == cute::UMMA::Major::K ? std::make_pair(k, mn) : std::make_pair(mn, k);
}
static int get_non_contiguous_dim(const cute::UMMA::Major& major) {
return major == cute::UMMA::Major::K ? -2 : -1;
}
static int get_compiled_dim(const int& dim, const char& name, const std::string& compiled_dims) {
for (const char& c: compiled_dims) {
if (name == c)
return dim;
}
return 0;
}
static CUtensorMapDataType aten_dtype_to_tensor_map_dtype(const at::ScalarType& dtype,
const bool& allow_tf32) {
if (allow_tf32 and dtype == torch::kFloat)
return CU_TENSOR_MAP_DATA_TYPE_TFLOAT32;
switch (dtype) {
case torch::kInt: return CU_TENSOR_MAP_DATA_TYPE_INT32;
case torch::kFloat: return CU_TENSOR_MAP_DATA_TYPE_FLOAT32;
case torch::kBFloat16: return CU_TENSOR_MAP_DATA_TYPE_BFLOAT16;
case torch::kFloat8_e4m3fn: return CU_TENSOR_MAP_DATA_TYPE_UINT8;
default: DG_HOST_UNREACHABLE("Unsupported dtype");
}
}
static CUtensorMapSwizzle mode_into_tensor_map_swizzle(const int& mode, const int& base) {
#if CUDA_VERSION >= 12080
if (base != 0) {
DG_HOST_ASSERT(base == 32 and mode == 128);
return CU_TENSOR_MAP_SWIZZLE_128B_ATOM_32B;
}
#endif
DG_HOST_ASSERT(base == 0);
switch (mode) {
case 0:
case 16: return CU_TENSOR_MAP_SWIZZLE_NONE;
case 32: return CU_TENSOR_MAP_SWIZZLE_32B;
case 64: return CU_TENSOR_MAP_SWIZZLE_64B;
case 128: return CU_TENSOR_MAP_SWIZZLE_128B;
default: DG_HOST_UNREACHABLE("Unsupported swizzling mode");
}
}
static CUtensorMap make_tma_2d_desc(const torch::Tensor& t,
int gmem_inner_dim, int gmem_outer_dim,
int smem_inner_dim, int smem_outer_dim,
const int& gmem_outer_stride,
const int& swizzle_mode, const int& swizzle_base = 0,
const bool& allow_tf32 = false) {
const auto& elem_size = static_cast<int>(t.element_size());
if (swizzle_mode != 0)
smem_inner_dim = swizzle_mode / elem_size;
CUtensorMap tensor_map;
const cuuint64_t gmem_dims[2] = {static_cast<cuuint64_t>(gmem_inner_dim), static_cast<cuuint64_t>(gmem_outer_dim)};
const cuuint32_t smem_dims[2] = {static_cast<cuuint32_t>(smem_inner_dim), static_cast<cuuint32_t>(smem_outer_dim)};
const cuuint64_t gmem_strides[1] = {static_cast<cuuint64_t>(gmem_outer_stride * elem_size), };
const cuuint32_t elem_strides[2] = {1, 1};
// if (get_env<int>("DG_JIT_DEBUG")) {
// printf("Making TMA desc: global memory: %d %d, shared memory: %d %d, outer stride: %d, swizzle: %d (base: %d), elem size: %d\n",
// gmem_inner_dim, gmem_outer_dim, smem_inner_dim, smem_outer_dim,
// gmem_outer_stride, swizzle_mode, swizzle_base, elem_size);
// }
cuTensorMapEncodeTiled(
&tensor_map, aten_dtype_to_tensor_map_dtype(t.scalar_type(), allow_tf32),
2, t.data_ptr(), gmem_dims, gmem_strides, smem_dims, elem_strides,
CU_TENSOR_MAP_INTERLEAVE_NONE, mode_into_tensor_map_swizzle(swizzle_mode, swizzle_base),
CU_TENSOR_MAP_L2_PROMOTION_L2_256B, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
return tensor_map;
}
static CUtensorMap make_tma_3d_desc(const torch::Tensor& t,
const int& gmem_dim_0, const int& gmem_dim_1, const int& gmem_dim_2,
const int& smem_dim_0, const int& smem_dim_1, const int& smem_dim_2,
const int& gmem_stride_0, const int& gmem_stride_1,
const int& swizzle_mode, const int& swizzle_base = 0,
const bool& allow_tf32 = false) {
const auto& elem_size = static_cast<int>(t.element_size());
if (swizzle_mode != 0)
DG_HOST_ASSERT(smem_dim_0 == swizzle_mode / elem_size);
CUtensorMap tensor_map;
const cuuint64_t gmem_dims[3] = {static_cast<cuuint64_t>(gmem_dim_0), static_cast<cuuint64_t>(gmem_dim_1), static_cast<cuuint64_t>(gmem_dim_2),};
const cuuint32_t smem_dims[3] = {static_cast<cuuint32_t>(smem_dim_0), static_cast<cuuint32_t>(smem_dim_1), static_cast<cuuint32_t>(smem_dim_2)};
const cuuint64_t gmem_strides[2] = {static_cast<cuuint64_t>(gmem_stride_0 * elem_size), static_cast<cuuint64_t>(gmem_stride_1 * elem_size)};
const cuuint32_t elem_strides[3] = {1, 1, 1};
// if (get_env<int>("DG_JIT_DEBUG")) {
// printf("Making 3D TMA desc: global memory: %d %d %d, shared memory: %d %d %d, outer stride: %d %d, swizzle: %d, elem size: %d\n",
// gmem_dim_0, gmem_dim_1, gmem_dim_2, smem_dim_0, smem_dim_1, smem_dim_2,
// gmem_stride_0, gmem_stride_1, swizzle_mode, elem_size);
// }
cuTensorMapEncodeTiled(
&tensor_map, aten_dtype_to_tensor_map_dtype(t.scalar_type(), allow_tf32),
3, t.data_ptr(), gmem_dims, gmem_strides, smem_dims, elem_strides,
CU_TENSOR_MAP_INTERLEAVE_NONE, mode_into_tensor_map_swizzle(swizzle_mode, swizzle_base),
CU_TENSOR_MAP_L2_PROMOTION_L2_256B, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE);
return tensor_map;
}
static CUtensorMap make_tma_a_desc(const cute::UMMA::Major& major,
const torch::Tensor& t,
const int& shape_m, const int& shape_k,
const int& block_m, const int& block_k,
const int& outer_stride,
const int& num_groups,
const int& swizzle_mode, const int& swizzle_base = 0,
const bool& allow_tf32 = false) {
if (num_groups > 1)
DG_HOST_ASSERT(major == cute::UMMA::Major::K);
const auto& [gmem_inner_dim, gmem_outer_dim] = get_inner_outer_dims(major, shape_k, shape_m * num_groups);
const auto& [smem_inner_dim, smem_outer_dim] = get_inner_outer_dims(major, block_k, block_m);
return make_tma_2d_desc(t,
gmem_inner_dim, gmem_outer_dim,
smem_inner_dim, smem_outer_dim,
outer_stride,
swizzle_mode, swizzle_base,
allow_tf32);
}
static CUtensorMap make_tma_b_desc(const cute::UMMA::Major& major,
const torch::Tensor& t,
const int& shape_n, const int& shape_k,
const int& block_n, const int& block_k,
const int& outer_stride,
const int& num_groups,
const int& swizzle_mode, const int& swizzle_base = 0,
const bool& allow_tf32 = false) {
const auto& [gmem_inner_dim, gmem_outer_dim] = get_inner_outer_dims(major, shape_k, shape_n);
const auto& [smem_inner_dim, smem_outer_dim] = get_inner_outer_dims(major, block_k, block_n);
// `num_groups` is always applied into the outer dimensions
return make_tma_2d_desc(t,
gmem_inner_dim, gmem_outer_dim * num_groups,
smem_inner_dim, smem_outer_dim,
outer_stride,
swizzle_mode, swizzle_base,
allow_tf32);
}
static CUtensorMap make_tma_cd_desc(const torch::Tensor& t,
const int& shape_m, const int& shape_n,
const int& block_m, const int& block_n,
const int& outer_stride,
const int& num_groups,
const int& swizzle_mode, const int& swizzle_base = 0,
const bool& allow_tf32 = false) {
// Swizzling requires the inner box dim to be less or equal than `kSwizzleCDMode`
// bytes, so `BLOCK_N * sizeof(T) / kSwizzleCDMode` TMA stores are required
return make_tma_2d_desc(t,
shape_n, shape_m * num_groups,
block_n, block_m,
outer_stride,
swizzle_mode, swizzle_base,
allow_tf32);
}
static CUtensorMap make_tma_sf_desc(const cute::UMMA::Major& major,
const torch::Tensor& t,
int shape_mn, int shape_k,
const int& block_mn, const int& block_k,
const int& num_groups,
const int& swizzle_mode, const int& swizzle_base = 0,
const bool& allow_tf32 = false) {
DG_HOST_ASSERT(major == cute::UMMA::Major::MN);
// TODO: maybe swizzle SF as well
DG_HOST_ASSERT(swizzle_mode == 0);
shape_mn = get_tma_aligned_size(shape_mn, static_cast<int>(t.element_size()));
return make_tma_2d_desc(t,
shape_mn, ceil_div(shape_k, block_k * (t.scalar_type() == torch::kFloat ? 1 : 4)) * num_groups,
block_mn, 1,
shape_mn,
swizzle_mode, swizzle_base,
allow_tf32);
}
} // namespace deep_gemm
@@ -0,0 +1,39 @@
#pragma once
#include <cuda.h>
#include <cuda_runtime.h>
#include <nvrtc.h>
#include <torch/python.h>
#include <ATen/cuda/CUDAContext.h>
#include "kernels/geforce/static_switch.h"
namespace sm89 {
template<bool use_fast_accum>
void fp8_bias_gemm_cuda(void* Aptr, void* SFA, void* Bptr, void* SFB, void* bias_ptr, void* out, int M, int N, int K, cudaStream_t stream);
}
namespace blockwise {
static void sm89_fp8_gemm_1d2d_bias(const torch::Tensor& a, const torch::Tensor& sfa,
const torch::Tensor& b, const torch::Tensor& sfb,
const torch::Tensor& bias,
const torch::Tensor& d,
const int& m, const int& n, const int& k,
const bool use_fast_accum) {
auto stream = at::cuda::getCurrentCUDAStream().stream();
if (use_fast_accum) {
sm89::fp8_bias_gemm_cuda<true>(
a.data_ptr(), sfa.data_ptr(),
b.data_ptr(), sfb.data_ptr(),
bias.data_ptr(), d.data_ptr(),
m, n, k, stream);
} else {
sm89::fp8_bias_gemm_cuda<false>(
a.data_ptr(), sfa.data_ptr(),
b.data_ptr(), sfb.data_ptr(),
bias.data_ptr(), d.data_ptr(),
m, n, k, stream);
}
}
}
@@ -0,0 +1,66 @@
#pragma once
#include <cuda.h>
#include <cuda_runtime.h>
#include <nvrtc.h>
#include <torch/python.h>
#include <ATen/cuda/CUDAContext.h>
#include <cute/arch/mma_sm100_desc.hpp>
#include "runtime_utils.hpp"
#include "config.hpp"
#include "static_switch.hpp"
namespace deep_gemm{
template<int N, int K>
void sm90_fp8_gemm_1d2d_bias_launch(int num_sms, int num_threads, int cluster_dim, int smem_size, cudaStream_t stream, float* sfb, float* bias, int* grouped_layout,
uint32_t shape_m, uint32_t shape_n, uint32_t shape_k,
const CUtensorMap tensor_map_a,
const CUtensorMap tensor_map_b,
const CUtensorMap tensor_map_d,
const CUtensorMap tensor_map_sfa);
};
namespace blockwise{
static void sm90_fp8_gemm_1d2d_bias(const torch::Tensor& a, const torch::Tensor& sfa,
const torch::Tensor& b, const torch::Tensor& sfb,
const torch::Tensor& bias,
const std::optional<torch::Tensor>& c,
const torch::Tensor& d,
const int& m, const int& n, const int& k, const int num_sms) {
// DG_HOST_ASSERT(not c.has_value() and d.scalar_type() == torch::kBFloat16);
const auto& config = GemmConfig<90>();
// Requires no TMA splits
// DG_HOST_ASSERT(config.smem_config.swizzle_a_mode == config.block_k);
// DG_HOST_ASSERT(config.smem_config.swizzle_b_mode == config.block_k);
int smem_size = k == 16384 || k == 8192 ? 216624 : config.smem_config.smem_size;
const auto& tensor_map_a = make_tma_a_desc(cute::UMMA::Major::K, a, m, k,
config.block_m,
config.block_k,
static_cast<int>(a.stride(-2)), 1,
config.smem_config.swizzle_a_mode);
const auto& tensor_map_b = make_tma_b_desc(cute::UMMA::Major::K, b, n, k,
config.block_n,
config.block_k,
static_cast<int>(b.stride(-2)), 1,
config.smem_config.swizzle_b_mode);
const auto& tensor_map_d = make_tma_cd_desc(d, m, static_cast<int>(d.size(-1)),
config.block_m,
config.block_n,
static_cast<int>(d.stride(-2)), 1,
config.smem_config.swizzle_cd_mode);
const auto& tensor_map_sfa = make_tma_sf_desc(cute::UMMA::Major::MN, sfa, m, k,
config.block_m, config.block_k, 1, 0);
auto stream = at::cuda::getCurrentCUDAStream().stream();
// Launch
DIM_SWITCH(k, K,
DIM_SWITCH(n, N,
deep_gemm::sm90_fp8_gemm_1d2d_bias_launch<N, K>(num_sms, config.thread_config.num_threads, config.multicast_config.num_multicast, smem_size, stream, (float*)sfb.data_ptr(), (float*)bias.data_ptr(), nullptr, m, n, k, tensor_map_a, tensor_map_b, tensor_map_d, tensor_map_sfa);)
)
}
};
@@ -0,0 +1,60 @@
#pragma once
#define DIM_SWITCH(VAR_NAME, CONST_NAME, ...) \
if (VAR_NAME == 4096) { \
constexpr static int CONST_NAME = 4096; \
__VA_ARGS__ \
} else if (VAR_NAME == 2048){ \
constexpr static int CONST_NAME = 2048; \
__VA_ARGS__ \
} else if (VAR_NAME == 8192){ \
constexpr static int CONST_NAME = 8192; \
__VA_ARGS__ \
} else if(VAR_NAME == 16384) { \
constexpr static int CONST_NAME = 16384; \
__VA_ARGS__ \
} else { \
TORCH_CHECK(false, "Unsupported DIM_SWITCH value: ", VAR_NAME); \
}
#define BOOL_SWITCH(COND, CONST_NAME, ...) \
if (COND) { \
constexpr static bool CONST_NAME = true; \
__VA_ARGS__ \
} else { \
constexpr static bool CONST_NAME = false; \
__VA_ARGS__ \
} \
//K/128
#define BLOCK_K_SWITCH(COSNT_NAME, ...) \
if (K == 2048) { \
constexpr static int COSNT_NAME = 16; \
__VA_ARGS__ \
} \
else if (K == 4096) { \
constexpr static int COSNT_NAME = 32; \
__VA_ARGS__ \
} else if (K == 8192) { \
constexpr static int COSNT_NAME = 64; \
__VA_ARGS__ \
} else if (K == 16384) { \
constexpr static int COSNT_NAME = 128; \
__VA_ARGS__ \
} else { \
TORCH_CHECK(false, "Unsupported K value: ", K); \
}
#define M_SWITCH(...) \
if (M <= 1024) { \
constexpr static int BM = 128; \
constexpr static int BN = 128; \
constexpr static int WARP_ROW = 2; \
constexpr static int WARP_COL = 2; \
__VA_ARGS__ \
} else { \
constexpr static int BM = 128; \
constexpr static int BN = 256; \
constexpr static int WARP_ROW = 2; \
constexpr static int WARP_COL = 4; \
__VA_ARGS__ \
}

Some files were not shown because too many files have changed in this diff Show More