Compare commits
43
Commits
v2
...
maint/pr1471-fixed
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9f28d34d6d | ||
|
|
a7470d2da3 | ||
|
|
821b4bc275 | ||
|
|
09c42cc73a | ||
|
|
d6ed0e54da | ||
|
|
6ce4993ab7 | ||
|
|
b9dfb12502 | ||
|
|
375ba410c7 | ||
|
|
a43284ae1b | ||
|
|
6ab2d6f98d | ||
|
|
9472ed0aab | ||
|
|
4eb4e252eb | ||
|
|
0434e9dcf0 | ||
|
|
36cf64b15f | ||
|
|
7a270e66be | ||
|
|
74de6f27a2 | ||
|
|
6979040225 | ||
|
|
80972d07fa | ||
|
|
260b326d8b | ||
|
|
c57fa1eab0 | ||
|
|
562b951208 | ||
|
|
aa893056e8 | ||
|
|
313c2985f0 | ||
|
|
b3b3858c02 | ||
|
|
d39e2e9fd7 | ||
|
|
a66fb26781 | ||
|
|
f303d94780 | ||
|
|
dae7c0da89 | ||
|
|
0ce6dc928f | ||
|
|
76b0550c15 | ||
|
|
384c1e9493 | ||
|
|
b1dbcc93f6 | ||
|
|
b93833772e | ||
|
|
30b523edd6 | ||
|
|
6aab7f3832 | ||
|
|
9cd53fe5f8 | ||
|
|
6a32cf3a5e | ||
|
|
98be9b3da2 | ||
|
|
c53e85b767 | ||
|
|
40a8bd2d3b | ||
|
|
31aa115611 | ||
|
|
98ac10a528 | ||
|
|
a5a6d171e5 |
@@ -69,7 +69,7 @@ approval, then upload reviewed accepted baseline records.
|
||||
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
|
||||
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
|
||||
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `PERF_MAX_REGRESSION` if set, otherwise `0.05` (5%). |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `0.05` (5%). |
|
||||
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
|
||||
|
||||
Hardcoded defaults:
|
||||
@@ -148,8 +148,7 @@ For each metric with at least two non-null source values:
|
||||
4. Stop if any source record regresses against the batch median by more than
|
||||
`max_intra_batch_regression`.
|
||||
|
||||
Default `max_intra_batch_regression` to `PERF_MAX_REGRESSION` when set,
|
||||
otherwise `0.05`. Print a table with per-source values, batch median, and
|
||||
Default `max_intra_batch_regression` to `0.05`. Print a table with per-source values, batch median, and
|
||||
worst intra-batch regression.
|
||||
|
||||
This check prevents uploading a mixed batch where one JSON is materially
|
||||
@@ -183,7 +182,7 @@ present, that run is not a valid source for baseline reseeding.
|
||||
|
||||
### 2. Sync and back up existing HF records under /tmp
|
||||
|
||||
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
|
||||
Use `fastvideo/performance/hf_store.py` helpers directly. Do **not** use
|
||||
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
|
||||
records, while this step must only fetch and back up existing history.
|
||||
|
||||
@@ -192,7 +191,7 @@ The sync command pattern is:
|
||||
```bash
|
||||
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
|
||||
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
|
||||
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
```
|
||||
|
||||
Then back up only the sanitized model directory under `/tmp`:
|
||||
@@ -200,8 +199,8 @@ Then back up only the sanitized model directory under `/tmp`:
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(PYTHONPATH=fastvideo/tests/performance python - <<'PY'
|
||||
from hf_store import sanitize
|
||||
MODEL_SAFE=$(python - <<'PY'
|
||||
from fastvideo.performance.hf_store import sanitize
|
||||
print(sanitize("<model_id>"))
|
||||
PY
|
||||
)
|
||||
@@ -235,7 +234,7 @@ first baseline seed. Continue, but report that baseline history was empty.
|
||||
Load the last 5 successful records for the target:
|
||||
|
||||
```python
|
||||
from hf_store import load_records_for_model
|
||||
from fastvideo.performance.hf_store import load_records_for_model
|
||||
|
||||
records = load_records_for_model(
|
||||
"/tmp/perf-tracking",
|
||||
@@ -372,7 +371,7 @@ prepared records plus backup on disk.
|
||||
Use the shared storage helper so the path and repo type match CI:
|
||||
|
||||
```python
|
||||
from hf_store import upload_record
|
||||
from fastvideo.performance.hf_store import upload_record
|
||||
|
||||
upload_record("<local_record_path>", record, strict=True)
|
||||
```
|
||||
@@ -460,7 +459,7 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
intentional baseline replacement.
|
||||
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
|
||||
median comparison, and persistence rules.
|
||||
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
|
||||
- `fastvideo/performance/hf_store.py` — HF sync, record loading,
|
||||
`sanitize()`, and `upload_record()`.
|
||||
- `fastvideo/tests/performance/test_inference_performance.py` — source result
|
||||
JSON schema.
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"description": "Wan2.1 T2V 1.3B inference performance",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
|
||||
@@ -114,6 +114,17 @@ steps:
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Extraction Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
@@ -371,6 +382,21 @@ steps:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "scripts/lora_extraction/**"
|
||||
- "fastvideo/tests/lora_extraction/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/training/training_utils.py"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Extraction Tests"
|
||||
env:
|
||||
- TEST_TYPE=lora_extraction
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
# pull_request_target instead of pull_request: the workflow definition and
|
||||
# the hook config are always taken from the BASE branch, so fork /
|
||||
# first-time-contributor PRs run immediately without a maintainer clicking
|
||||
# "Approve and run". The PR head is checked out as data only.
|
||||
pull_request_target:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
inputs:
|
||||
@@ -15,12 +19,25 @@ permissions:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
|
||||
if: github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
# For PR events, lint the PR head — but keep the hook definitions from
|
||||
# the base branch so an untrusted PR cannot alter what gets executed.
|
||||
- name: Save trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
- uses: actions/checkout@v4
|
||||
if: github.event_name == 'pull_request_target'
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha }}
|
||||
persist-credentials: false
|
||||
- name: Restore trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
@@ -125,7 +125,7 @@ jobs:
|
||||
set -euo pipefail
|
||||
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
|
||||
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
|
||||
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
|
||||
exit 1
|
||||
@@ -136,6 +136,7 @@ jobs:
|
||||
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
|
||||
[ssim]=ssim [training]=training
|
||||
[lora-inference]=inference_lora [lora-training]=training_lora
|
||||
[lora-extraction]=lora_extraction
|
||||
[distillation]=distillation_dmd [self-forcing]=self_forcing
|
||||
[vsa]=training_vsa [vmoba]=inference_vmoba
|
||||
[performance]=performance [api]=api_server
|
||||
|
||||
@@ -13,12 +13,33 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
|
||||
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
|
||||
# push trigger is a sufficient change detector on its own -- no separate
|
||||
# detect-changes/paths-filter job is needed now that there is a single
|
||||
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
|
||||
# rocm Dockerfile stay manual-dispatch only.
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'docker/Dockerfile'
|
||||
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
# One static group, no cancellation: every run of this workflow writes the same
|
||||
# mutable registry tags (latest, py3.12-latest, ...), so runs must serialize —
|
||||
# concurrent push/dispatch runs would race on those tags, and cancelling a run
|
||||
# mid-publish can strand the cu126/cu130 tag families at different commits. An
|
||||
# in-flight superseded build wastes its runner time, but its tags are then
|
||||
# overwritten by the newer queued run. GitHub keeps a single pending run per
|
||||
# group: the newest queued run replaces any older queued one.
|
||||
concurrency:
|
||||
group: infra-build-image
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
# CUDA matrix: Python 3.12 x {12.6.3, 13.0.0} x {amd64, arm64}. Each architecture
|
||||
# builds natively and pushes only by digest; publish-cuda-manifests is the sole
|
||||
@@ -28,7 +49,11 @@ jobs:
|
||||
# aliases; 13.0.0/cu130 is published under explicit versioned tags. Flash-attn
|
||||
# 2.8.3 comes from the architecture-specific prebuilt releases.
|
||||
build-cuda-images:
|
||||
if: ${{ github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
|
||||
# on a push that changed docker/Dockerfile (inputs are null on push). The
|
||||
# repository guard keeps fork syncs from auto-building; manual dispatch
|
||||
# still works in forks.
|
||||
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -75,10 +100,10 @@ jobs:
|
||||
secrets: inherit
|
||||
|
||||
publish-cuda-manifests:
|
||||
# !cancelled(): a failed sibling build leg must not skip the manifests for a
|
||||
# CUDA lane whose own digests all exist; the digest-count check below fails
|
||||
# the incomplete lane loudly instead.
|
||||
if: ${{ !cancelled() && github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# !cancelled(): publish lanes whose digests exist even if a sibling build
|
||||
# leg failed (the digest-count check fails incomplete lanes); it also
|
||||
# bypasses skipped-needs propagation, hence the explicit skipped check.
|
||||
if: ${{ !cancelled() && needs.build-cuda-images.result != 'skipped' }}
|
||||
needs: build-cuda-images
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
|
||||
@@ -84,9 +84,12 @@ RUN source /opt/venv/bin/activate \
|
||||
FFMPEG_NATIVE_CXX=/usr/bin/g++ \
|
||||
bash /opt/FastVideo/apps/dreamverse/scripts/install_native_ffmpeg.sh
|
||||
|
||||
# FASTVIDEO_FA4: FA4 (flash_attn.cute) is opt-in; this image installs it via
|
||||
# the dreamverse extra and is validated with it, so enable it here.
|
||||
ENV FASTVIDEO_DREAMVERSE_HOME=/var/lib/dreamverse \
|
||||
STREAM_MODE=av_fmp4 \
|
||||
FASTVIDEO_ENABLE_PROMPT_SAFETY=0 \
|
||||
FASTVIDEO_FA4=1 \
|
||||
HF_HOME=/root/.cache/huggingface
|
||||
|
||||
RUN mkdir -p /var/lib/dreamverse
|
||||
|
||||
@@ -12,13 +12,17 @@ Defaults:
|
||||
|
||||
- `HF_REPO_ID=FastVideo/performance-tracking`
|
||||
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
|
||||
- `PERF_MAX_REGRESSION=0.05`
|
||||
|
||||
Records can include source metadata:
|
||||
Records can include source metadata and rolling-baseline policy context:
|
||||
|
||||
- `run_source`: `pr`, `local`, `scheduled_main`, or `unknown`
|
||||
- `baseline_eligible`: only successful scheduled-main records should be true
|
||||
- Buildkite metadata such as branch, PR number, build URL, build ID, and job ID
|
||||
- `regression_thresholds`: per-metric rolling-baseline percent and absolute
|
||||
floors used for recomputed status context
|
||||
|
||||
Dashboard/API metric payloads expose `threshold_exceeded` for raw threshold
|
||||
crossings; `regressed` remains the gated CI-failure signal.
|
||||
|
||||
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
|
||||
configured dataset repo requires authenticated access:
|
||||
@@ -89,7 +93,8 @@ Trend charts show metric-specific axes and exact point details on hover/focus:
|
||||
- PR number, branch, and Buildkite URL when present
|
||||
|
||||
The latest status table uses the stored JSON `success` value. Recomputed
|
||||
baseline context is shown separately and does not override stored status.
|
||||
baseline context applies each metric's percent and absolute regression floors
|
||||
and does not override stored status.
|
||||
|
||||
## API
|
||||
|
||||
|
||||
@@ -440,6 +440,8 @@ export default function App() {
|
||||
<th>Throughput</th>
|
||||
<th>Memory</th>
|
||||
<th>Worst</th>
|
||||
<th>Exceeded</th>
|
||||
<th>Failing</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
@@ -465,6 +467,12 @@ export default function App() {
|
||||
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
|
||||
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
|
||||
<td>
|
||||
{row.threshold_exceeded_metrics.length
|
||||
? row.threshold_exceeded_metrics.join(", ")
|
||||
: "none"}
|
||||
</td>
|
||||
<td>{row.failing_metrics.length ? row.failing_metrics.join(", ") : "none"}</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
|
||||
@@ -2,6 +2,12 @@ export type MetricValue = {
|
||||
current: number | null;
|
||||
baseline: number | null;
|
||||
regression_pct: number | null;
|
||||
absolute_delta: number | null;
|
||||
threshold_percent: number;
|
||||
threshold_absolute: number;
|
||||
gated: boolean;
|
||||
threshold_exceeded: boolean;
|
||||
regressed: boolean;
|
||||
label: string;
|
||||
lower_is_better: boolean;
|
||||
precision: number;
|
||||
@@ -15,7 +21,8 @@ export type SummaryRow = {
|
||||
success: boolean;
|
||||
baseline_n: number;
|
||||
worst_regression_pct: number | null;
|
||||
regression_threshold_pct: number;
|
||||
threshold_exceeded_metrics: string[];
|
||||
failing_metrics: string[];
|
||||
computed_regression_status: "pass" | "fail";
|
||||
status: "pass" | "fail";
|
||||
run_source: RunSource;
|
||||
|
||||
+4
-2
@@ -170,12 +170,14 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
|
||||
# Then verify both survive so a broken overlay fails the build instead of shipping
|
||||
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay and FA4 falls back to FA2.
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
|
||||
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
|
||||
# arm64 and the image runs FA3/FA2 as usual.
|
||||
RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; FA4 falls back to FA2)"; \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
|
||||
else \
|
||||
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
|
||||
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
|
||||
|
||||
@@ -103,6 +103,7 @@ can merge a PR.
|
||||
|---|---|---|
|
||||
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
|
||||
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
|
||||
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
|
||||
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
|
||||
@@ -144,6 +145,7 @@ Valid direct test names:
|
||||
| `/test training` | `training` |
|
||||
| `/test lora-inference` | `inference_lora` |
|
||||
| `/test lora-training` | `training_lora` |
|
||||
| `/test lora-extraction` | `lora_extraction` |
|
||||
| `/test distillation` | `distillation_dmd` |
|
||||
| `/test self-forcing` | `self_forcing` |
|
||||
| `/test vsa` | `training_vsa` |
|
||||
|
||||
@@ -72,7 +72,10 @@ fastvideo/tests/performance/
|
||||
│ writes Markdown summary + (optionally) uploads new records
|
||||
├── dashboard.py
|
||||
│ └── builds time-series Plotly HTML from HF history
|
||||
└── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
|
||||
fastvideo/performance/
|
||||
├── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
└── metric_policy.py # shared rolling-baseline threshold policy
|
||||
```
|
||||
|
||||
The HF dataset (`FastVideo/performance-tracking` by default) holds one
|
||||
@@ -92,25 +95,28 @@ and recipe changes instead of treating all records for a model as equivalent.
|
||||
|
||||
## Metrics
|
||||
|
||||
Each benchmark records six metrics:
|
||||
Each benchmark records six metrics. The rolling-baseline comparator also has a
|
||||
per-metric policy with direction, percent threshold, absolute threshold, and a
|
||||
`gated` flag.
|
||||
|
||||
| Metric | Raw key | Normalized key | Direction |
|
||||
|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better |
|
||||
| Metric | Raw key | Normalized key | Direction | Default rolling policy |
|
||||
|---|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better | 8% and 0.5 s |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better | 8% and 0.05 FPS |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better | 5% and 256 MB |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better | 5% and 0.25 s |
|
||||
|
||||
`test_inference_performance.py` temporarily sets `FASTVIDEO_STAGE_LOGGING=1`
|
||||
while it runs so pipeline stage execution times are available in
|
||||
`generate_video(...).logging_info`. Stage logs use pipeline-unique keys such as
|
||||
`prompt_encoding_stage` so duplicate stage classes do not collide. For
|
||||
`PipelineStage` entries, the extractor maps the `stage_class` field:
|
||||
`TextEncodingStage` maps to `text_encoder_time_s`, `DenoisingStage` and
|
||||
`DmdDenoisingStage` map to `dit_time_s`, and `DecodingStage` maps to
|
||||
`vae_decode_time_s`, with a fallback for older logs that used the class name as
|
||||
`PipelineStage` entries, shared component stage bases emit a stable
|
||||
`component_metric`: text encoding stages map to `text_encoder_time_s`,
|
||||
denoising stages and subclasses map to `dit_time_s`, and decoding stages map to
|
||||
`vae_decode_time_s`. The extractor falls back to known `stage_class` names for
|
||||
older logs that do not include `component_metric` or that used the class name as
|
||||
the stage key. Generator-side timings such as `PostDecodeFrameProcessStage`,
|
||||
`VideoSaveStage`, and `AudioMuxStage` are intentionally ignored. If a pipeline
|
||||
does not report one of the mapped stages, that component metric is stored as
|
||||
@@ -156,9 +162,22 @@ headroom and almost never need touching.
|
||||
|
||||
`compare_baseline.py` loads the last 5 successful, baseline-eligible records
|
||||
for the same `(model_id, gpu_type)` from the HF dataset, computes the median
|
||||
for each available metric, and fails if the current run regresses by more than
|
||||
`PERF_MAX_REGRESSION` (default 5%). For latency, memory, and component times,
|
||||
higher values are regressions. For throughput, lower values are regressions.
|
||||
for each available metric, and evaluates the current run with the metric's
|
||||
rolling regression policy. For latency, memory, and component times, higher
|
||||
values are regressions. For throughput, lower values are regressions.
|
||||
|
||||
A metric exceeds its rolling threshold when both of these are true:
|
||||
|
||||
```text
|
||||
percent_delta > threshold_percent
|
||||
absolute_delta > threshold_absolute
|
||||
```
|
||||
|
||||
Gated metrics fail CI when that threshold crossing happens. Set `gated: false`
|
||||
for metrics that should remain visible in reports and the dashboard without
|
||||
failing CI. Dashboard/API payloads expose `threshold_exceeded` separately from
|
||||
`regressed`, where `regressed` means a gated CI failure. Missing or `null`
|
||||
metrics are skipped.
|
||||
|
||||
This is the **drift detector** — it catches sub-threshold regressions that
|
||||
slowly add up. Only scheduled-main successful records are baseline eligible.
|
||||
@@ -172,6 +191,45 @@ agent skill to advance the rolling median.
|
||||
|
||||
## Schemas
|
||||
|
||||
### Benchmark config (`.buildkite/performance-benchmarks/tests/*.json`)
|
||||
|
||||
Benchmark configs without `config_schema_version` are treated as legacy v1
|
||||
configs and remain loadable. New or migrated configs should use
|
||||
`config_schema_version: 2` and include explicit comparable identity fields:
|
||||
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1
|
||||
}
|
||||
```
|
||||
|
||||
`benchmark_id` is still required in this phase because raw artifact names,
|
||||
generated-video directories, normalized record paths, and the current rolling
|
||||
baseline comparator still depend on it. The v2 identity fields are config
|
||||
metadata that make the measured workload explicit:
|
||||
|
||||
| Field | Purpose |
|
||||
|---|---|
|
||||
| `workload_id` | Stable benchmark family, such as `wan-t2v-1.3b`. |
|
||||
| `variant_id` | Intentional recipe family, such as `canonical`. |
|
||||
| `benchmark_version` | Version of the measurement protocol and comparison policy. |
|
||||
|
||||
If a config declares `config_schema_version: 2`, loading fails clearly when any
|
||||
required v2 identity field is missing. If v2 identity or metadata fields are
|
||||
added without `config_schema_version: 2`, loading also fails so partial
|
||||
migrations do not silently run as v1 configs. Optional v2 metadata fields
|
||||
reserved for follow-up work, such as `recipe`, `metric_threshold_policy`, and
|
||||
`quality_metadata`, must be JSON objects when present.
|
||||
|
||||
Recipe fingerprinting, hardware/software profile IDs, exact-identity
|
||||
comparison, metric-specific threshold policy behavior, promoted baselines, and
|
||||
dashboard regrouping are separate follow-up changes. Until those land, rolling
|
||||
baseline comparison remains keyed by `(model_id, gpu_type)`.
|
||||
|
||||
### Raw record (`results/perf_*.json`)
|
||||
|
||||
Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
@@ -179,6 +237,10 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
|
||||
"device": "NVIDIA L40S",
|
||||
"num_gpus": 2,
|
||||
@@ -196,6 +258,13 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
"max_dit_time_s": 10.0,
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.10,
|
||||
"threshold_absolute": 1.0,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"commit": "<full sha>",
|
||||
"pr_number": "1234",
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
@@ -222,6 +291,13 @@ result, used as the rolling-baseline source of truth.
|
||||
"text_encoder_time_s": 2.141,
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208,
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.08,
|
||||
"threshold_absolute": 0.5,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
@@ -238,18 +314,17 @@ successful main/full-suite uploads and remain eligible for rolling baselines.
|
||||
|
||||
| Variable | Default | Used by | Purpose |
|
||||
|---|---|---|---|
|
||||
| `PERF_MAX_REGRESSION` | `0.05` | `compare_baseline.py` | Per-metric regression fraction that fails the build. |
|
||||
| `PERFORMANCE_TRACKING_ROOT` | `/tmp/perf-tracking` | `compare_baseline.py`, `dashboard.py` | Local directory the HF dataset is synced to. |
|
||||
| `PERF_REPORTS_DIR` | `/root/data/perf_reports` | `compare_baseline.py`, `dashboard.py` | Where the Markdown summary and Plotly HTML get written for Buildkite to pick up. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `fastvideo/performance/hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `fastvideo/performance/hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `PERF_UPLOAD_POLICY` | `never` | `compare_baseline.py` | Upload policy: `never`, `pass`, or `always`. |
|
||||
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
|
||||
| `TEST_SCOPE` | unset | `compare_baseline.py` | CI context used to infer scheduled-main runs together with `BUILDKITE_BRANCH=main`. |
|
||||
| `BUILDKITE_BRANCH`, `BUILDKITE_COMMIT`, `BUILDKITE_PULL_REQUEST` | unset | `compare_baseline.py`, `test_inference_performance.py` | CI metadata stamped into records. |
|
||||
| `DASHBOARD_DAYS` | `30` | `dashboard.py` | Lookback window for the Plotly trend pages. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `fastvideo/performance/hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `FASTVIDEO_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
|
||||
|
||||
## CI integration
|
||||
@@ -261,9 +336,11 @@ Buildkite artifact upload is in
|
||||
`.buildkite/scripts/pr_test.sh:upload_performance_artifacts`.
|
||||
|
||||
Each performance build runs pytest first. If that fixed-threshold phase fails,
|
||||
`compare_baseline.py` is skipped, so Markdown summaries and normalized JSON
|
||||
artifacts are not emitted. The dashboard still runs best-effort for
|
||||
observability. When pytest passes, the rolling-baseline phase emits:
|
||||
PR/direct runs skip `compare_baseline.py` because they only upload passing
|
||||
records. Scheduled-main runs still execute `compare_baseline.py` with
|
||||
`PERF_PYTEST_RC` set so the failed canonical attempt is visible in normalized
|
||||
JSON and dashboard history. The dashboard runs best-effort for observability.
|
||||
When the rolling-baseline phase runs, it emits:
|
||||
|
||||
* **Markdown summary** — appended to `$GITHUB_STEP_SUMMARY` when that variable
|
||||
is set, and written as `perf_<sha>_<ts>.md` for Buildkite upload. Contains a
|
||||
@@ -279,11 +356,16 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
## Adding a new benchmark
|
||||
|
||||
1. Drop a new JSON config into
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. Required keys:
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. New configs should
|
||||
use v2 identity fields:
|
||||
|
||||
```json
|
||||
{
|
||||
"benchmark_id": "<unique-id>",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "<stable-workload-id>",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"model": { "model_path": "...", "model_short_name": "..." },
|
||||
"init_kwargs": { "num_gpus": 1, ... },
|
||||
"generation_kwargs": { "num_frames": 45, ... },
|
||||
@@ -299,9 +381,17 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"default": { "max_generation_time_s": 120.0, "max_peak_memory_mb": 30000.0 }
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": { "threshold_percent": 0.10, "threshold_absolute": 1.0, "gated": true }
|
||||
}
|
||||
}
|
||||
```
|
||||
}
|
||||
|
||||
```
|
||||
|
||||
Legacy v1 configs without `config_schema_version` still load, but should not
|
||||
gain v2 identity or metadata fields until they are migrated to
|
||||
`config_schema_version: 2`.
|
||||
|
||||
2. The pytest test auto-discovers all configs — no test code needed. CI
|
||||
picks it up on the next `/test performance` run.
|
||||
@@ -320,6 +410,13 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
a useful fixed gate. The rolling baseline will still track component times
|
||||
when static component thresholds are omitted.
|
||||
|
||||
6. Omit `regression_thresholds` to use the default rolling-baseline policy, or
|
||||
include only benchmark-specific deviations. Tune these independently from
|
||||
the fixed thresholds when a metric is noisy or should be informational. The
|
||||
fixed `thresholds` block is an absolute pytest ceiling. The
|
||||
`regression_thresholds` block controls rolling-baseline comparisons against
|
||||
recent scheduled-main records.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
**"No baseline for ... Initializing"** — first run for this `(model_id,
|
||||
@@ -336,5 +433,5 @@ pipelines that did not report a mapped component stage.
|
||||
|
||||
**Component timing is `null`** — the generated result did not include a mapped
|
||||
stage in `logging_info.stages`. Check that the pipeline emits stage logging
|
||||
and that the stage name is listed in `STAGE_METRIC_MAP` in
|
||||
`test_inference_performance.py`.
|
||||
and that the stage emits `component_metric` or is covered by the legacy
|
||||
`STAGE_METRIC_MAP` fallback in `test_inference_performance.py`.
|
||||
|
||||
@@ -74,6 +74,23 @@ uv pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
### Flash Attention 4 (opt-in)
|
||||
|
||||
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
|
||||
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
|
||||
runtime on some GPU/shape combinations. To use FA4, install the pinned
|
||||
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
|
||||
|
||||
```bash
|
||||
export FASTVIDEO_FA4=1
|
||||
```
|
||||
|
||||
On GPUs below sm90 a capability gate routes to FlashAttention-2 the calls FA4
|
||||
cannot serve there: grad-enabled (training) attention (FA4's backward requires
|
||||
sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
|
||||
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
|
||||
FastVideo fails loudly instead of silently falling back.
|
||||
|
||||
### FP4 Flash Attention 4 (Blackwell only)
|
||||
|
||||
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_i2v"
|
||||
|
||||
IMAGE_PATH = "assets/girl.png"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A woman stands up and walks away"
|
||||
)
|
||||
_ = generator.generate_video(
|
||||
prompt,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=1024,
|
||||
width=1024,
|
||||
num_frames=121,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,37 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -12,6 +12,9 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
if(NOT GPU_BACKEND)
|
||||
set(GPU_BACKEND "CUDA")
|
||||
endif()
|
||||
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
@@ -50,7 +53,16 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
|
||||
if(_FASTVIDEO_USER_CUDA_ARCH)
|
||||
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
|
||||
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
|
||||
# Only numeric spellings translate; keywords like "native"/"all" would
|
||||
# otherwise be mangled into nonsense ("nativ.e").
|
||||
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
|
||||
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
|
||||
message(FATAL_ERROR
|
||||
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
|
||||
"supported. Use a numeric arch (e.g. 90a, 121), set "
|
||||
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
|
||||
"auto-detect from the visible GPU.")
|
||||
endif()
|
||||
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
|
||||
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
|
||||
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
|
||||
@@ -173,6 +185,14 @@ else()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
|
||||
# there, and tk's base_types.cuh brace-initializes signed-char vector members
|
||||
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
|
||||
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
|
||||
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
endif()
|
||||
|
||||
if(ENABLE_TK_KERNELS)
|
||||
message(STATUS "ThunderKittens kernels: ENABLED")
|
||||
else()
|
||||
@@ -263,11 +283,6 @@ set(CUDA_FLAGS
|
||||
"--expt-relaxed-constexpr"
|
||||
"-Xcompiler=-fno-strict-aliasing"
|
||||
"-Xcompiler=-fPIC"
|
||||
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
|
||||
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
|
||||
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
|
||||
# is a no-op on x86_64, where char is already signed.
|
||||
"-Xcompiler=-fsigned-char"
|
||||
"-DTORCH_COMPILE"
|
||||
"-Xnvlink=--verbose"
|
||||
"-Xptxas=--verbose"
|
||||
@@ -426,3 +441,14 @@ if(ENABLE_ATTN_QAT_INFER)
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
# One-look answer to "what is this build producing?" — kept last so it is the
|
||||
# final thing configure prints. The per-kernel matrix lives in README.md.
|
||||
message(STATUS "============== fastvideo-kernel build summary ==============")
|
||||
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
|
||||
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
|
||||
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
|
||||
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
|
||||
message(STATUS "============================================================")
|
||||
|
||||
@@ -2,6 +2,42 @@
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Kernel inventory
|
||||
|
||||
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
|
||||
|
||||
| Extension | Kernels | Sources | GPU arch | Build gate |
|
||||
|---|---|---|---|---|
|
||||
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
|
||||
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
|
||||
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
|
||||
|
||||
Runtime-JIT kernels (no build step, ship in every wheel/image):
|
||||
|
||||
| Kernels | Where | Used when |
|
||||
|---|---|---|
|
||||
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
|
||||
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
|
||||
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|
||||
|---|---|---|---|---|---|
|
||||
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
|
||||
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
|
||||
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
|
||||
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
|
||||
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
|
||||
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
|
||||
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
|
||||
|
||||
Notes:
|
||||
|
||||
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
|
||||
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
|
||||
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
|
||||
|
||||
## Installation
|
||||
|
||||
### Standard Installation (Local Development)
|
||||
|
||||
@@ -46,6 +46,23 @@ fi
|
||||
if git rev-parse --git-dir >/dev/null 2>&1; then
|
||||
git submodule update --init --recursive include/cutlass include/tk
|
||||
fi
|
||||
# Fail fast with a clear message if the headers are still missing (e.g. a
|
||||
# Docker context that excluded .git AND the submodule contents) instead of
|
||||
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
|
||||
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
|
||||
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
|
||||
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
|
||||
if [ ! -d include/cutlass/include ]; then
|
||||
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
|
||||
echo " CUTLASS sources must already be present (run" >&2
|
||||
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
|
||||
echo " in the source checkout, or include them in the build context)." >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -d include/tk/include ]; then
|
||||
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
|
||||
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
|
||||
fi
|
||||
|
||||
# Install build dependencies
|
||||
uv pip install scikit-build-core cmake ninja
|
||||
|
||||
@@ -1,15 +1,39 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
from fastvideo import envs
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1, mirroring the
|
||||
# kernel package's FASTVIDEO_VSA_CUTEDSL: its CuTeDSL kernels JIT-compile per
|
||||
# shape family and can fail at runtime on some arch/shape combinations, so it
|
||||
# is never auto-selected just because it is installed. Below sm90 a capability
|
||||
# gate in flash_attn_cute routes to FA2 the calls FA4 cannot serve there:
|
||||
# grad-enabled (its backward asserts sm90+) and GQA (pack_gqa fails CuTeDSL
|
||||
# JIT, observed on sm_89).
|
||||
if envs.FASTVIDEO_FA4:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
fa_version = "4"
|
||||
except ImportError:
|
||||
else:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
@@ -21,6 +45,12 @@ except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
try:
|
||||
if importlib.util.find_spec("flash_attn.cute") is not None:
|
||||
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
|
||||
"set FASTVIDEO_FA4=1 to use it for inference.")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
|
||||
# already a registered torch.library custom op, so dynamo treats it as a
|
||||
@@ -87,9 +117,10 @@ if fa_version in ("2", "3"):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
# FA4 path: `flash_attn_func` (from `flash_attn_cute`) goes through a
|
||||
# registered torch.library custom op (with an FA4 backward on sm90+;
|
||||
# grad-enabled and GQA calls below sm90 route to FA2), so a passthrough
|
||||
# is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
@@ -99,17 +130,6 @@ else:
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_WARNED_NON_FA_DTYPE = False
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
@@ -271,12 +291,8 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
# SP). Cast through bf16 and restore, matching TORCH_SDPA's tolerance.
|
||||
orig_dtype = query.dtype
|
||||
if orig_dtype not in (torch.float16, torch.bfloat16):
|
||||
global _WARNED_NON_FA_DTYPE
|
||||
if not _WARNED_NON_FA_DTYPE:
|
||||
_WARNED_NON_FA_DTYPE = True
|
||||
logger.warning(
|
||||
"FLASH_ATTN received %s inputs; casting to bfloat16 for the "
|
||||
"kernel and restoring on output.", orig_dtype)
|
||||
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
|
||||
f"bfloat16 for the kernel and restoring on output.")
|
||||
query = query.to(torch.bfloat16)
|
||||
key = key.to(torch.bfloat16)
|
||||
value = value.to(torch.bfloat16)
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""NABLA block-sparse flex-attention backend (Kandinsky5 "nabla" checkpoints).
|
||||
|
||||
The block mask is data-dependent: nablaT_v2 mean-pools 64-token blocks of the
|
||||
fractal-ordered sequence, thresholds the softmaxed block map, and ORs it with a
|
||||
precomputed spatio-temporal-window (STA) mask carried on the attention
|
||||
metadata. The mask spans the full sequence, so this backend does not support
|
||||
sequence parallelism — use it via LocalAttention only.
|
||||
"""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
from torch.nn.attention.flex_attention import BlockMask, flex_attention
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
|
||||
CAN_USE_FLEX_ATTN = True
|
||||
except ImportError:
|
||||
CAN_USE_FLEX_ATTN = False
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
|
||||
|
||||
def nablaT_v2(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
sta: torch.Tensor,
|
||||
thr: float = 0.9,
|
||||
) -> "BlockMask":
|
||||
q = q.transpose(1, 2).contiguous()
|
||||
k = k.transpose(1, 2).contiguous()
|
||||
|
||||
# Map estimation
|
||||
B, h, S, D = q.shape
|
||||
s1 = S // 64
|
||||
qa = q.reshape(B, h, s1, 64, D).mean(-2)
|
||||
ka = k.reshape(B, h, s1, 64, D).mean(-2).transpose(-2, -1)
|
||||
map = qa @ ka
|
||||
|
||||
map = torch.softmax(map / math.sqrt(D), dim=-1)
|
||||
# Map binarization
|
||||
vals, inds = map.sort(-1)
|
||||
cvals = vals.cumsum_(-1)
|
||||
mask = (cvals >= 1 - thr).int()
|
||||
mask = mask.gather(-1, inds.argsort(-1))
|
||||
|
||||
mask = torch.logical_or(mask, sta)
|
||||
|
||||
# BlockMask creation
|
||||
kv_nb = mask.sum(-1).to(torch.int32)
|
||||
kv_inds = mask.argsort(dim=-1, descending=True).to(torch.int32)
|
||||
return BlockMask.from_kv_blocks(torch.zeros_like(kv_nb), kv_inds, kv_nb, kv_inds, BLOCK_SIZE=64, mask_mod=None)
|
||||
|
||||
|
||||
class NablaAttentionBackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "NABLA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["NablaAttentionImpl"]:
|
||||
return NablaAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["NablaAttentionMetadata"]:
|
||||
return NablaAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["NablaAttentionMetadataBuilder"]:
|
||||
return NablaAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class NablaAttentionMetadata(AttentionMetadata):
|
||||
# Block-level STA window mask [1, 1, S/64, S/64], precomputed once per run.
|
||||
sta_mask: torch.Tensor = None # type: ignore[assignment]
|
||||
# Cumulative-probability threshold for block-map binarization.
|
||||
P: float = 0.9
|
||||
visual_shape: tuple[int, int, int] = (0, 0, 0)
|
||||
|
||||
|
||||
class NablaAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def prepare(self) -> None:
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
sta_mask: torch.Tensor,
|
||||
P: float,
|
||||
visual_shape: tuple[int, int, int],
|
||||
**kwargs: Any,
|
||||
) -> NablaAttentionMetadata:
|
||||
return NablaAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
sta_mask=sta_mask,
|
||||
P=P,
|
||||
visual_shape=visual_shape,
|
||||
)
|
||||
|
||||
|
||||
class NablaAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
softmax_scale: float,
|
||||
causal: bool = False,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
if not CAN_USE_FLEX_ATTN:
|
||||
raise RuntimeError("NABLA attention requires torch.nn.attention.flex_attention, "
|
||||
"which is unavailable in this PyTorch build.")
|
||||
if causal:
|
||||
raise ValueError("NABLA attention does not support causal masking.")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: NablaAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
# q/k/v: [B, S, heads, head_dim], fractal-ordered by the model; S % 64 == 0.
|
||||
block_mask = nablaT_v2(query, key, attn_metadata.sta_mask, thr=attn_metadata.P)
|
||||
return flex_attention(
|
||||
query=query.transpose(1, 2),
|
||||
key=key.transpose(1, 2),
|
||||
value=value.transpose(1, 2),
|
||||
block_mask=block_mask,
|
||||
).transpose(1, 2)
|
||||
@@ -252,6 +252,7 @@ class LocalAttention(nn.Module):
|
||||
causal: bool = False,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -262,7 +263,10 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
|
||||
attn_backend = get_attn_backend(head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
default_backend=default_backend)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
|
||||
@@ -84,8 +84,9 @@ def get_attn_backend(
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends)
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends, default_backend)
|
||||
|
||||
|
||||
@cache
|
||||
@@ -94,6 +95,7 @@ def _cached_get_attn_backend(
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
@@ -112,6 +114,12 @@ def _cached_get_attn_backend(
|
||||
if backend_by_env_var is not None:
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
|
||||
# Layer-level default (e.g. a checkpoint that requires a specific sparse
|
||||
# backend). Lower precedence than the global force and the env var, so
|
||||
# users can still override it.
|
||||
if selected_backend is None and default_backend is not None:
|
||||
selected_backend = default_backend
|
||||
|
||||
# get device-specific attn_backend
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
|
||||
@@ -4,10 +4,9 @@ import functools
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -16,7 +15,8 @@ if torch.cuda.is_available():
|
||||
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
|
||||
except ImportError:
|
||||
# flash_attn.cute (FA4) is simply not installed -- expected on builds
|
||||
# without it; callers fall back to FA3/FA2 quietly.
|
||||
# without it; callers handle the ImportError (the FASTVIDEO_FA4 gate in
|
||||
# flash_attn.py raises, the FP4 probe treats FA4 as unavailable).
|
||||
raise
|
||||
except Exception as e:
|
||||
# flash_attn.cute IS installed but failed to import -- almost always an
|
||||
@@ -24,23 +24,57 @@ if torch.cuda.is_available():
|
||||
# 'cutlass.cute.core' has no attribute 'ThrMma'" (an AttributeError, not
|
||||
# ImportError). This is fixable by pinning a compatible
|
||||
# nvidia-cutlass-dsl, so warn loudly, then re-raise as ImportError so
|
||||
# callers fall back to FA3/FA2 instead of crashing worker init.
|
||||
# callers can handle it uniformly.
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r); "
|
||||
"falling back to FA3/FA2. This is usually an nvidia-cutlass-dsl "
|
||||
"version mismatch -- pin a compatible nvidia-cutlass-dsl to "
|
||||
"restore FA4.", e)
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r). "
|
||||
"This is usually an nvidia-cutlass-dsl version mismatch -- pin a "
|
||||
"compatible nvidia-cutlass-dsl to restore FA4.", e)
|
||||
raise ImportError(f"flash_attn.cute (FA4) import failed: {e!r}") from e
|
||||
else:
|
||||
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
|
||||
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
|
||||
|
||||
try:
|
||||
# FA2 serves the calls FA4 cute cannot on pre-sm90 GPUs (backward, GQA).
|
||||
# Optional so FA4-only installs can still import this module.
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
except ImportError:
|
||||
_flash_attn_2_func = None
|
||||
_flash_attn_2_varlen_func = None
|
||||
|
||||
|
||||
def _check_dropout(dropout_p: float) -> None:
|
||||
if dropout_p != 0.0:
|
||||
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _sm90_or_newer() -> bool:
|
||||
return current_platform.has_device_capability(90)
|
||||
|
||||
|
||||
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
|
||||
if _sm90_or_newer():
|
||||
return False
|
||||
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
|
||||
# capability gate, not a runtime fallback):
|
||||
# * the backward asserts sm90+ (L40S/sm_89 dies on its arch check);
|
||||
# * GQA fails CuTeDSL JIT in pack_gqa ("ValueError: Operation creation
|
||||
# failed", observed on sm_89 with HunyuanGameCraft/LTX2).
|
||||
if q.shape[-2] != k.shape[-2]:
|
||||
return True
|
||||
return torch.is_grad_enabled() and any(t.requires_grad for t in (q, k, v))
|
||||
|
||||
|
||||
def _fa2_or_raise(fa2_func: Callable | None) -> Callable:
|
||||
if fa2_func is None:
|
||||
raise RuntimeError("this attention call cannot run on FA4 cute below sm90 (its backward and "
|
||||
"GQA support require sm90+) and flash-attn 2, which serves it there, is "
|
||||
"not installed.")
|
||||
return fa2_func
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_cute_forward",
|
||||
mutates_args=(),
|
||||
@@ -243,70 +277,6 @@ torch.library.register_autograd(
|
||||
)
|
||||
|
||||
|
||||
# FA4's CuTeDSL kernels JIT-compile per shape family, and some configurations
|
||||
# fail MLIR op creation at runtime even though the import succeeded (observed:
|
||||
# GQA models on sm_89 dying in pack_gqa with "ValueError: Operation creation
|
||||
# failed"). Degrade to FA2 once, process-wide, instead of crashing inference.
|
||||
class _FA4Policy:
|
||||
"""Per-call gate for the FA4 cute fast path, with FA2 as the fallback.
|
||||
|
||||
FA4 is skipped when:
|
||||
* a previous call failed at runtime -- CuTeDSL JIT compilation is
|
||||
shape-dependent, so the first failure disables FA4 for the rest of
|
||||
the process instead of retrying a broken JIT on every call; or
|
||||
* the call needs autograd -- FA4's backward asserts sm90+ (L40S/sm_89
|
||||
dies on its arch check) and is unvalidated for training in this repo
|
||||
(its lse is not even allocated through our inference-shaped custom
|
||||
op), so training keeps the pre-FA4 behavior: FA2 on every device.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.broken = False
|
||||
|
||||
def use_fa4(self, *tensors: torch.Tensor) -> bool:
|
||||
if self.broken:
|
||||
return False
|
||||
return not (torch.is_grad_enabled() and any(t.requires_grad for t in tensors))
|
||||
|
||||
def mark_broken(self, error: Exception) -> None:
|
||||
if not self.broken:
|
||||
self.broken = True
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) failed at runtime (%r); falling back "
|
||||
"to FA2 for the rest of this process.", error)
|
||||
|
||||
|
||||
_FA4 = _FA4Policy()
|
||||
|
||||
|
||||
def _with_fa2_fallback(fa2_func: Callable) -> Callable:
|
||||
"""Pair an FA4 cute wrapper with its FA2 twin of the same signature.
|
||||
|
||||
The decorated body runs only when ``_FA4`` allows it; otherwise (or after
|
||||
the first FA4 runtime failure) the call is served by ``fa2_func``.
|
||||
``NotImplementedError`` is a contract error (e.g. dropout), not a JIT
|
||||
failure, so it propagates without disabling FA4.
|
||||
"""
|
||||
|
||||
def decorator(fa4_func: Callable) -> Callable:
|
||||
|
||||
@functools.wraps(fa4_func)
|
||||
def wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
if _FA4.use_fa4(q, k, v):
|
||||
try:
|
||||
return fa4_func(q, k, v, *args, **kwargs)
|
||||
except NotImplementedError:
|
||||
raise
|
||||
except Exception as e: # CuTeDSL compile errors surface as ValueError
|
||||
_FA4.mark_broken(e)
|
||||
return fa2_func(q, k, v, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_func)
|
||||
def flash_attn_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -317,6 +287,16 @@ def flash_attn_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
|
||||
return out
|
||||
@@ -392,7 +372,6 @@ def flash_attn_fp4_func(
|
||||
return torch.ops.fastvideo._flash_attn_cute_fp4_forward(q, k, v, sfq, sfk, softmax_scale, causal)
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_varlen_func)
|
||||
def flash_attn_varlen_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -407,6 +386,20 @@ def flash_attn_varlen_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_varlen_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_varlen_forward(
|
||||
q,
|
||||
|
||||
@@ -21,24 +21,35 @@ from einops import rearrange
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
|
||||
from fastvideo import envs
|
||||
|
||||
|
||||
def _resolve_flash_attn_varlen_func() -> Any:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
if envs.FASTVIDEO_FA4:
|
||||
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
|
||||
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
|
||||
# fail loudly here rather than fall through to FA3/FA2. RuntimeError,
|
||||
# not ImportError: importers like bsa_attn.py treat ImportError as
|
||||
# "flash-attn not installed" and silently degrade to reference kernels.
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
|
||||
return flash_attn_varlen_func_cute
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
return flash_attn_varlen_func_flash
|
||||
return flash_attn_varlen_func_flash
|
||||
|
||||
|
||||
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
|
||||
|
||||
@@ -2,14 +2,24 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_kandinsky5_transformer_block(n: str, m) -> bool:
|
||||
return ("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5ArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [
|
||||
lambda n, m:
|
||||
("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_kandinsky5_transformer_block])
|
||||
|
||||
# NABLA block-sparse attention for attention_type="nabla" checkpoints, plus
|
||||
# the dense backends every DiT supports.
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.NABLA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
|
||||
# Native FastVideo implementation uses the same parameter names as diffusers
|
||||
# except FFN internals: Diffusers FFN uses `in_layer/out_layer`, while
|
||||
|
||||
@@ -43,11 +43,16 @@ class TextEncoderArchConfig(EncoderArchConfig):
|
||||
require_processor: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
# update_model_arch re-runs __post_init__ after pipeline configs may
|
||||
# have customized tokenizer_kwargs (e.g. kandinsky5/gen3c/longcat set
|
||||
# "padding"); rebuilding the dict here would silently wipe those
|
||||
# customizations, so only fill in defaults for keys not already set.
|
||||
defaults = {
|
||||
"truncation": True,
|
||||
"max_length": self.text_len,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
self.tokenizer_kwargs = defaults | self.tokenizer_kwargs
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -5,6 +5,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
|
||||
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
|
||||
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
@@ -16,5 +17,6 @@ __all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
|
||||
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
"HYWorldConfig", "Kandinsky5T2VConfig", "Kandinsky5I2VConfig", "MatrixGame2I2V480PConfig",
|
||||
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits import Kandinsky5VideoConfig
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, CLIPTextConfig
|
||||
from fastvideo.configs.models.encoders.reason1 import Reason1Config
|
||||
from fastvideo.configs.models.vaes import HunyuanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
# Byte-exact copy of the upstream Kandinsky5/diffusers template, including the
|
||||
# "promt"/"scren" typos: the checkpoints were trained with this exact system # codespell:ignore promt,scren
|
||||
# prompt, and ENCODE_START_IDX below is the tokenized length of everything
|
||||
# before the user prompt. Fixing the typos shifts user content to index 127
|
||||
# and mis-conditions every generation.
|
||||
KANDINSKY5_PROMPT_TEMPLATE = "\n".join([
|
||||
"<|im_start|>system\nYou are a promt engineer. Describe the video in detail.", # codespell:ignore promt
|
||||
"Describe how the camera moves or shakes, describe the zoom and view angle, whether it follows the objects.",
|
||||
"Describe the location of the video, main characters or objects and their action.",
|
||||
"Describe the dynamism of the video and presented actions.",
|
||||
"Name the visual style of the video: whether it is a professional footage, user generated content, some kind of animation, video game or scren content.", # codespell:ignore scren
|
||||
"Describe the visual effects, postprocessing and transitions if they are presented in the video.",
|
||||
"Pay attention to the order of key actions shown in the scene.<|im_end|>",
|
||||
"<|im_start|>user\n{}<|im_end|>",
|
||||
])
|
||||
KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX = 129
|
||||
|
||||
|
||||
def kandinsky5_qwen_preprocess_text(prompt: str) -> str:
|
||||
if not prompt.strip():
|
||||
prompt = "."
|
||||
return KANDINSKY5_PROMPT_TEMPLATE.format(prompt)
|
||||
|
||||
|
||||
def kandinsky5_qwen_postprocess_text(outputs: BaseEncoderOutput,
|
||||
mask: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if outputs.hidden_states is None:
|
||||
raise RuntimeError("Kandinsky5 Qwen prompt embeddings require hidden_states.")
|
||||
hidden_states = outputs.hidden_states[-1]
|
||||
prompt_embeds = hidden_states[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
|
||||
mask = mask[:, KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX:]
|
||||
if prompt_embeds.shape[1] == 0:
|
||||
prompt_embeds = hidden_states[:, -1:]
|
||||
mask = torch.ones((mask.shape[0], 1), dtype=mask.dtype, device=mask.device)
|
||||
return prompt_embeds, mask
|
||||
|
||||
|
||||
def kandinsky5_clip_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
if outputs.pooler_output is None:
|
||||
raise RuntimeError("Kandinsky5 CLIP pooled output is required.")
|
||||
return outputs.pooler_output
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5T2VConfig(PipelineConfig):
|
||||
"""Kandinsky-5.0 Lite text-to-video pipeline configuration."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=Kandinsky5VideoConfig)
|
||||
vae_config: VAEConfig = field(default_factory=HunyuanVAEConfig)
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (Reason1Config(), CLIPTextConfig()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], Any], ...] = field(
|
||||
default_factory=lambda: (kandinsky5_qwen_preprocess_text, preprocess_text))
|
||||
postprocess_text_funcs: tuple[Callable[..., Any], ...] = field(
|
||||
default_factory=lambda: (kandinsky5_qwen_postprocess_text, kandinsky5_clip_postprocess_text))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", "bf16"))
|
||||
text_encoder_max_lengths: tuple[int, ...] = field(
|
||||
default_factory=lambda: (KANDINSKY5_PROMPT_TEMPLATE_ENCODE_START_IDX + 512, 77))
|
||||
|
||||
flow_shift: float | None = 5.0
|
||||
vae_tiling: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if len(self.text_encoder_configs) != 2:
|
||||
raise ValueError(f"Kandinsky5 pipeline requires exactly 2 text encoders (qwen and clip), "
|
||||
f"but got {len(self.text_encoder_configs)} encoder(s).")
|
||||
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
qwen_cfg = self.text_encoder_configs[0]
|
||||
qwen_cfg.arch_config.output_hidden_states = True
|
||||
qwen_cfg.arch_config.tokenizer_kwargs.update({
|
||||
"padding": True,
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
})
|
||||
|
||||
clip_cfg = self.text_encoder_configs[1]
|
||||
clip_cfg.arch_config.tokenizer_kwargs.update({
|
||||
"padding": "max_length",
|
||||
"max_length": 77,
|
||||
"truncation": True,
|
||||
"add_special_tokens": True,
|
||||
"return_tensors": "pt",
|
||||
})
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5I2VConfig(Kandinsky5T2VConfig):
|
||||
"""Kandinsky-5.0 image-to-video pipeline configuration."""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# I2V needs the VAE encoder to encode the conditioning image.
|
||||
self.vae_config.load_encoder = True
|
||||
@@ -20,6 +20,7 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: str | None = None
|
||||
FASTVIDEO_FA4: bool = False
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: str | None = None
|
||||
@@ -207,9 +208,18 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
|
||||
# - "SAGE_ATTN": use Sage Attention
|
||||
# - "SAGE_ATTN_THREE": use Sage Attention 3
|
||||
# FLASH_ATTN uses FlashAttention-3/2; to run FlashAttention-4 set
|
||||
# FASTVIDEO_FA4=1 as well (see below).
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
# If set (=1), the FLASH_ATTN backend uses FlashAttention-4
|
||||
# (flash_attn.cute). FA4 is opt-in and never auto-selected just because it
|
||||
# is installed. Below sm90, grad-enabled and GQA calls are routed to FA2
|
||||
# (FA4's backward asserts sm90+ and its pack_gqa fails to JIT there).
|
||||
"FASTVIDEO_FA4":
|
||||
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
|
||||
|
||||
@@ -323,9 +323,9 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
|
||||
@@ -10,13 +10,21 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.attention.backends.nabla import CAN_USE_FLEX_ATTN, flex_attention, nablaT_v2
|
||||
from fastvideo.configs.models.dits import Kandinsky5VideoConfig
|
||||
from fastvideo.layers.layernorm import LayerNormScaleShift
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
if not CAN_USE_FLEX_ATTN:
|
||||
logger.warning("torch.nn.attention.flex_attention is unavailable in this PyTorch build; "
|
||||
"Kandinsky5 NABLA sparse attention (Pro checkpoints) cannot be used.")
|
||||
|
||||
FRACTAL_PIXEL_SIZE = 8
|
||||
_ARCH_CONFIG_DEFAULTS = Kandinsky5VideoConfig().arch_config
|
||||
|
||||
@@ -263,10 +271,9 @@ class Kandinsky5Modulation(nn.Module):
|
||||
|
||||
|
||||
def _apply_rotary(x: torch.Tensor, rope: torch.Tensor) -> torch.Tensor:
|
||||
orig_dtype = x.dtype
|
||||
x_ = x.reshape(*x.shape[:-1], -1, 1, 2).to(torch.float32)
|
||||
x_out = (rope * x_).sum(dim=-1)
|
||||
return x_out.reshape(*x.shape).to(orig_dtype)
|
||||
return x_out.reshape(*x.shape).to(x.dtype)
|
||||
|
||||
|
||||
class Kandinsky5Attention(nn.Module):
|
||||
@@ -277,6 +284,7 @@ class Kandinsky5Attention(nn.Module):
|
||||
head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None,
|
||||
prefix: str = "",
|
||||
use_nabla: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
assert num_channels % head_dim == 0
|
||||
@@ -306,6 +314,17 @@ class Kandinsky5Attention(nn.Module):
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
)
|
||||
# NABLA checkpoints get a second attention layer whose backend defaults
|
||||
# to NABLA_ATTN; FASTVIDEO_ATTENTION_BACKEND still overrides it.
|
||||
self.nabla_attention = None
|
||||
if use_nabla:
|
||||
self.nabla_attention = LocalAttention(
|
||||
num_heads=self.num_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
default_backend=AttentionBackendEnum.NABLA_ATTN,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -339,27 +358,54 @@ class Kandinsky5Attention(nn.Module):
|
||||
key = _apply_rotary(key, rotary_emb).type_as(key)
|
||||
|
||||
if sparse_params is not None:
|
||||
raise NotImplementedError(
|
||||
"Sparse attention is not yet supported for Kandinsky5 in FastVideo."
|
||||
)
|
||||
if self.nabla_attention is None:
|
||||
raise RuntimeError("sparse_params passed to an attention layer built without use_nabla; "
|
||||
"this checkpoint/config combination is inconsistent.")
|
||||
try:
|
||||
# Backend impl reads sta_mask/P from the forward-context
|
||||
# attention metadata built by the denoising stage.
|
||||
hidden_states = self.nabla_attention(query, key, value)
|
||||
except AssertionError as exc:
|
||||
# Standalone parity tests call the model without a pipeline
|
||||
# forward context; run the NABLA kernel directly.
|
||||
if "Forward context is not set" not in str(exc):
|
||||
raise
|
||||
attn_mask = nablaT_v2(query, key, sparse_params["sta_mask"], thr=sparse_params["P"])
|
||||
hidden_states = flex_attention(
|
||||
query=query.transpose(1, 2),
|
||||
key=key.transpose(1, 2),
|
||||
value=value.transpose(1, 2),
|
||||
block_mask=attn_mask,
|
||||
).transpose(1, 2)
|
||||
else:
|
||||
try:
|
||||
hidden_states = self.local_attention(query, key, value)
|
||||
|
||||
try:
|
||||
hidden_states = self.local_attention(query, key, value)
|
||||
except AssertionError as exc:
|
||||
# LocalAttention requires pipeline forward context. Standalone
|
||||
# parity tests call the model directly, so fallback to Torch SDPA.
|
||||
if "Forward context is not set" not in str(exc):
|
||||
raise
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
hidden_states = F.scaled_dot_product_attention(query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=None,
|
||||
is_causal=False)
|
||||
hidden_states = hidden_states.transpose(1, 2)
|
||||
hidden_states = hidden_states.flatten(2)
|
||||
except AssertionError as exc:
|
||||
# LocalAttention requires pipeline forward context. Standalone
|
||||
# parity tests call the model directly, so fallback to Torch SDPA.
|
||||
if "Forward context is not set" not in str(exc):
|
||||
raise
|
||||
|
||||
query_shape = query.shape[:-2]
|
||||
key_shape = key.shape[:-2]
|
||||
query = query.reshape(query_shape[0], -1, self.num_heads,
|
||||
query.shape[-1]).transpose(1, 2)
|
||||
key = key.reshape(key_shape[0], -1, self.num_heads,
|
||||
key.shape[-1]).transpose(1, 2)
|
||||
value = value.reshape(key_shape[0], -1, self.num_heads,
|
||||
value.shape[-1]).transpose(1, 2)
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=None,
|
||||
is_causal=False,
|
||||
)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(
|
||||
*query_shape, self.num_heads, -1)
|
||||
|
||||
hidden_states = hidden_states.flatten(-2, -1)
|
||||
|
||||
hidden_states, _ = self.out_layer(hidden_states)
|
||||
return hidden_states
|
||||
@@ -476,7 +522,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
|
||||
head_dim: int,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
prefix: str = ""):
|
||||
prefix: str = "",
|
||||
use_nabla: bool = False):
|
||||
super().__init__()
|
||||
self.visual_modulation = Kandinsky5Modulation(time_dim, model_dim, 9)
|
||||
|
||||
@@ -491,7 +538,8 @@ class Kandinsky5TransformerDecoderBlock(nn.Module):
|
||||
model_dim,
|
||||
head_dim,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.self_attention")
|
||||
prefix=f"{prefix}.self_attention",
|
||||
use_nabla=use_nabla)
|
||||
|
||||
self.cross_attention_norm = LayerNormScaleShift(
|
||||
model_dim,
|
||||
@@ -624,7 +672,8 @@ class Kandinsky5Transformer3DModel(BaseDiT):
|
||||
arch.ff_dim,
|
||||
head_dim,
|
||||
self._supported_attention_backends,
|
||||
prefix=f"{config.prefix}.visual_transformer_blocks.{i}")
|
||||
prefix=f"{config.prefix}.visual_transformer_blocks.{i}",
|
||||
use_nabla=arch.attention_type == "nabla")
|
||||
for i in range(arch.num_visual_blocks)
|
||||
])
|
||||
|
||||
@@ -694,6 +743,7 @@ class Kandinsky5Transformer3DModel(BaseDiT):
|
||||
scale_factor)
|
||||
to_fractal = sparse_params[
|
||||
"to_fractal"] if sparse_params is not None else False
|
||||
|
||||
visual_embed, visual_rope = fractal_flatten(visual_embed, visual_rope,
|
||||
visual_shape,
|
||||
block_mask=to_fractal)
|
||||
@@ -724,6 +774,7 @@ class Kandinsky5Transformer3DModel(BaseDiT):
|
||||
|
||||
if return_dict:
|
||||
return Kandinsky5TransformerOutput(sample=x)
|
||||
|
||||
return x
|
||||
|
||||
def materialize_non_persistent_buffers(self, device: torch.device,
|
||||
|
||||
@@ -537,9 +537,9 @@ class CausalMatrixGame2TransformerBlock(nn.Module):
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
|
||||
@@ -467,6 +467,18 @@ class Qwen2_5_VisionTransformerPretrainedModel(nn.Module):
|
||||
return hidden_states
|
||||
|
||||
|
||||
def _compute_default_rope_parameters(config, device=None, seq_len=None, **kwargs):
|
||||
# transformers>=5 removes the "default" entry from ROPE_INIT_FUNCTIONS and
|
||||
# moves rope_theta inside rope_parameters; replicate the 4.x default init.
|
||||
rope_params = getattr(config, "rope_parameters", None) or getattr(config, "rope_scaling", None) or {}
|
||||
base = rope_params.get("rope_theta", getattr(config, "rope_theta", 10000.0))
|
||||
head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
|
||||
partial_rotary_factor = getattr(config, "partial_rotary_factor", 1.0)
|
||||
dim = int(head_dim * partial_rotary_factor)
|
||||
inv_freq = 1.0 / (base**(torch.arange(0, dim, 2, dtype=torch.int64).float().to(device) / dim))
|
||||
return inv_freq, 1.0
|
||||
|
||||
|
||||
class Qwen2_5_VLRotaryEmbedding(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig, device=None):
|
||||
super().__init__()
|
||||
@@ -479,7 +491,14 @@ class Qwen2_5_VLRotaryEmbedding(nn.Module):
|
||||
self.original_max_seq_len = config.max_position_embeddings
|
||||
|
||||
self.config = config
|
||||
self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
|
||||
if self.rope_type in ROPE_INIT_FUNCTIONS:
|
||||
self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type]
|
||||
elif self.rope_type == "default":
|
||||
# transformers>=5 drops the "default" entry from ROPE_INIT_FUNCTIONS.
|
||||
self.rope_init_fn = _compute_default_rope_parameters
|
||||
else:
|
||||
raise KeyError(f"Unsupported rope_type '{self.rope_type}'; available: "
|
||||
f"{['default', *ROPE_INIT_FUNCTIONS]}")
|
||||
|
||||
inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device)
|
||||
self.register_buffer("inv_freq", inv_freq, persistent=False)
|
||||
@@ -948,6 +967,16 @@ QWEN2_5_VL_ATTENTION_CLASSES = {
|
||||
# If FlashAttention2 is not available, transparently fall back to SDPA.
|
||||
if not is_flash_attn_2_available():
|
||||
QWEN2_5_VL_ATTENTION_CLASSES["flash_attention_2"] = Qwen2_5_VLSdpaAttention
|
||||
else:
|
||||
# transformers>=5 only resolves the flash-attn functions when the model
|
||||
# preloads them via its attention interface; this module bypasses
|
||||
# PreTrainedModel, so _flash_attention_forward(implementation=None) raises
|
||||
# unless we preload here.
|
||||
try:
|
||||
from transformers.modeling_flash_attention_utils import lazy_import_flash_attention
|
||||
lazy_import_flash_attention("flash_attention_2")
|
||||
except (ImportError, ValueError):
|
||||
QWEN2_5_VL_ATTENTION_CLASSES["flash_attention_2"] = Qwen2_5_VLSdpaAttention
|
||||
|
||||
class Qwen2_5_VLDecoderLayer(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int):
|
||||
@@ -1035,7 +1064,7 @@ class Qwen2_5_VLModel(nn.Module):
|
||||
def __init__(self, config: Qwen2_5_VLConfig):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.padding_idx = config.pad_token_id
|
||||
self.padding_idx = getattr(config, "pad_token_id", None)
|
||||
self.vocab_size = config.vocab_size
|
||||
|
||||
self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
|
||||
@@ -1408,6 +1437,18 @@ class Qwen2_5_VLCausalLMOutputWithPast(ModelOutput):
|
||||
rope_deltas: Optional[torch.LongTensor] = None
|
||||
|
||||
|
||||
def _flatten_text_config(config):
|
||||
# transformers>=5 stops forwarding text-model attributes (hidden_size,
|
||||
# vocab_size, rope_scaling, ...) from the composite Qwen2_5_VLConfig to
|
||||
# config.text_config; this module reads them from the top level.
|
||||
text_config = getattr(config, "text_config", None)
|
||||
if text_config is not None:
|
||||
for key, value in text_config.to_dict().items():
|
||||
if not hasattr(config, key):
|
||||
setattr(config, key, value)
|
||||
return config
|
||||
|
||||
|
||||
class Qwen2_5_VLForConditionalGenerationSimple(nn.Module):
|
||||
_tied_weights_keys = ["lm_head.weight"]
|
||||
config_class = Qwen2_5_VLConfig
|
||||
@@ -1415,6 +1456,7 @@ class Qwen2_5_VLForConditionalGenerationSimple(nn.Module):
|
||||
|
||||
def __init__(self, config):
|
||||
super().__init__()
|
||||
config = _flatten_text_config(config)
|
||||
self.config = config
|
||||
self.visual = Qwen2_5_VisionTransformerPretrainedModel(config.vision_config)
|
||||
|
||||
|
||||
@@ -320,6 +320,19 @@ class TextEncoderLoader(ComponentLoader):
|
||||
f"text encoder index {idx} out of range for text_encoder_configs (len={len(encoder_configs)}), model_path={model_path}"
|
||||
)
|
||||
encoder_config = encoder_configs[idx]
|
||||
if (
|
||||
model_config.get("architectures") == ["CLIPModel"]
|
||||
and isinstance(model_config.get("text_config"), dict)
|
||||
):
|
||||
valid_arch_fields = {
|
||||
f.name for f in dataclasses.fields(encoder_config.arch_config)
|
||||
}
|
||||
model_config = {
|
||||
key: value
|
||||
for key, value in deepcopy(model_config["text_config"]).items()
|
||||
if key in valid_arch_fields
|
||||
}
|
||||
model_config["architectures"] = ["CLIPTextModel"]
|
||||
encoder_config.update_model_arch(model_config)
|
||||
if idx < 0 or idx >= len(encoder_precisions):
|
||||
raise IndexError(
|
||||
@@ -404,7 +417,9 @@ class TextEncoderLoader(ComponentLoader):
|
||||
else:
|
||||
loaded_weights: set[str] = model.load_weights(
|
||||
self._get_all_weights(
|
||||
model, model_path, to_cpu=use_cpu_offload
|
||||
model,
|
||||
model_path,
|
||||
to_cpu=fastvideo_args.text_encoder_cpu_offload,
|
||||
)
|
||||
) # type: ignore
|
||||
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Performance benchmark and dashboard utilities."""
|
||||
@@ -57,9 +57,7 @@ def is_baseline_eligible_record(record: dict[str, Any]) -> bool:
|
||||
"""
|
||||
if record.get("baseline_eligible") is True:
|
||||
return True
|
||||
if "baseline_eligible" not in record and "run_source" not in record:
|
||||
return True
|
||||
return False
|
||||
return "baseline_eligible" not in record and "run_source" not in record
|
||||
|
||||
|
||||
def resolve_hf_token() -> str | None:
|
||||
@@ -312,8 +310,12 @@ def load_records_for_model(
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_NUMERIC_COLS = (
|
||||
"latency", "throughput", "memory",
|
||||
"text_encoder_time_s", "dit_time_s", "vae_decode_time_s",
|
||||
"latency",
|
||||
"throughput",
|
||||
"memory",
|
||||
"text_encoder_time_s",
|
||||
"dit_time_s",
|
||||
"vae_decode_time_s",
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Metric policy for rolling performance baseline comparisons."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricPolicy:
|
||||
key: str
|
||||
label: str
|
||||
precision: int
|
||||
lower_is_better: bool
|
||||
threshold_percent: float
|
||||
threshold_absolute: float
|
||||
gated: bool = True
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricDelta:
|
||||
absolute: float
|
||||
percent: float
|
||||
threshold_exceeded: bool
|
||||
regressed: bool
|
||||
|
||||
|
||||
DEFAULT_METRIC_POLICIES: tuple[MetricPolicy, ...] = (
|
||||
MetricPolicy("latency", "Latency", 3, True, 0.08, 0.5),
|
||||
MetricPolicy("throughput", "Throughput", 3, False, 0.08, 0.05),
|
||||
MetricPolicy("memory", "Memory", 1, True, 0.05, 256.0),
|
||||
MetricPolicy("text_encoder_time_s", "Text Enc", 3, True, 0.05, 0.25),
|
||||
MetricPolicy("dit_time_s", "DiT", 3, True, 0.05, 0.25),
|
||||
MetricPolicy("vae_decode_time_s", "VAE Decode", 3, True, 0.05, 0.25),
|
||||
)
|
||||
|
||||
|
||||
def _optional_float(value: Any) -> float | None:
|
||||
if value is None or isinstance(value, bool):
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _optional_bool(value: Any) -> bool | None:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in {"1", "true", "yes", "on"}:
|
||||
return True
|
||||
if normalized in {"0", "false", "no", "off"}:
|
||||
return False
|
||||
return None
|
||||
|
||||
|
||||
def resolve_metric_policies(threshold_overrides: Mapping[str, Any] | None, ) -> tuple[MetricPolicy, ...]:
|
||||
"""Return default metric policies with optional per-metric overrides."""
|
||||
|
||||
if not isinstance(threshold_overrides, Mapping):
|
||||
threshold_overrides = {}
|
||||
policies: list[MetricPolicy] = []
|
||||
for base_policy in DEFAULT_METRIC_POLICIES:
|
||||
raw_override = threshold_overrides.get(base_policy.key, {})
|
||||
if not isinstance(raw_override, Mapping):
|
||||
raw_override = {}
|
||||
|
||||
threshold_percent = _optional_float(raw_override.get("threshold_percent"))
|
||||
threshold_absolute = _optional_float(raw_override.get("threshold_absolute"))
|
||||
gated = _optional_bool(raw_override.get("gated"))
|
||||
|
||||
policies.append(
|
||||
MetricPolicy(
|
||||
key=base_policy.key,
|
||||
label=base_policy.label,
|
||||
precision=base_policy.precision,
|
||||
lower_is_better=base_policy.lower_is_better,
|
||||
threshold_percent=(base_policy.threshold_percent if threshold_percent is None else threshold_percent),
|
||||
threshold_absolute=(base_policy.threshold_absolute
|
||||
if threshold_absolute is None else threshold_absolute),
|
||||
gated=base_policy.gated if gated is None else gated,
|
||||
))
|
||||
return tuple(policies)
|
||||
|
||||
|
||||
def serialize_metric_thresholds(policies: tuple[MetricPolicy, ...], ) -> dict[str, dict[str, float | bool]]:
|
||||
return {
|
||||
policy.key: {
|
||||
"threshold_percent": policy.threshold_percent,
|
||||
"threshold_absolute": policy.threshold_absolute,
|
||||
"gated": policy.gated,
|
||||
}
|
||||
for policy in policies
|
||||
}
|
||||
|
||||
|
||||
def regression_delta(
|
||||
policy: MetricPolicy,
|
||||
current: float,
|
||||
baseline: float,
|
||||
) -> MetricDelta | None:
|
||||
if baseline <= 0:
|
||||
return None
|
||||
absolute_delta = current - baseline if policy.lower_is_better else baseline - current
|
||||
percent_delta = absolute_delta / baseline
|
||||
threshold_exceeded = (percent_delta > policy.threshold_percent and absolute_delta > policy.threshold_absolute)
|
||||
return MetricDelta(
|
||||
absolute=absolute_delta,
|
||||
percent=percent_delta,
|
||||
threshold_exceeded=threshold_exceeded,
|
||||
regressed=policy.gated and threshold_exceeded,
|
||||
)
|
||||
@@ -13,7 +13,7 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
|
||||
from fastvideo.tests.performance import hf_store
|
||||
from fastvideo.performance import hf_store
|
||||
|
||||
from .service import build_latest_summary, build_trends, filter_records
|
||||
|
||||
@@ -150,7 +150,6 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type)
|
||||
rows = build_latest_summary(
|
||||
filtered,
|
||||
max_regression=float(os.environ.get("PERF_MAX_REGRESSION", "0.05")),
|
||||
run_source=run_source,
|
||||
)
|
||||
return {
|
||||
|
||||
@@ -1,26 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Metric definitions shared by the performance dashboard backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
from fastvideo.performance.metric_policy import DEFAULT_METRIC_POLICIES
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricDefinition:
|
||||
key: str
|
||||
label: str
|
||||
precision: int
|
||||
lower_is_better: bool
|
||||
|
||||
|
||||
METRICS: tuple[MetricDefinition, ...] = (
|
||||
MetricDefinition("latency", "Latency", 3, True),
|
||||
MetricDefinition("throughput", "Throughput", 3, False),
|
||||
MetricDefinition("memory", "Memory", 1, True),
|
||||
MetricDefinition("text_encoder_time_s", "Text Encoder", 3, True),
|
||||
MetricDefinition("dit_time_s", "DiT", 3, True),
|
||||
MetricDefinition("vae_decode_time_s", "VAE Decode", 3, True),
|
||||
)
|
||||
METRICS = DEFAULT_METRIC_POLICIES
|
||||
|
||||
METRIC_BY_KEY = {metric.key: metric for metric in METRICS}
|
||||
|
||||
@@ -13,9 +13,8 @@ from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.tests.performance.hf_store import is_baseline_eligible_record, safe_float
|
||||
|
||||
from .metrics import METRICS
|
||||
from fastvideo.performance.hf_store import is_baseline_eligible_record, safe_float
|
||||
from fastvideo.performance.metric_policy import regression_delta, resolve_metric_policies
|
||||
|
||||
Record = dict[str, Any]
|
||||
|
||||
@@ -95,19 +94,9 @@ def baseline_value(records: list[Record], metric_key: str) -> float | None:
|
||||
return float(statistics.median(values))
|
||||
|
||||
|
||||
def regression_percent(metric_key: str, current: float | None, baseline: float | None) -> float | None:
|
||||
if current is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
metric = next(metric for metric in METRICS if metric.key == metric_key)
|
||||
if metric.lower_is_better:
|
||||
return (current - baseline) / baseline * 100.0
|
||||
return (baseline - current) / baseline * 100.0
|
||||
|
||||
|
||||
def build_latest_summary(records: list[Record],
|
||||
*,
|
||||
baseline_window: int = 5,
|
||||
max_regression: float = 0.05,
|
||||
run_source: str | None = None) -> list[Record]:
|
||||
rows: list[Record] = []
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
@@ -123,52 +112,58 @@ def build_latest_summary(records: list[Record],
|
||||
if record is not latest and record.get("success", True) and is_baseline_eligible_record(record)
|
||||
]
|
||||
baseline_records = baseline_pool[-baseline_window:]
|
||||
metric_policies = resolve_metric_policies(latest.get("regression_thresholds"))
|
||||
|
||||
metrics: dict[str, Record] = {}
|
||||
regressions: list[float] = []
|
||||
for metric in METRICS:
|
||||
current = safe_float(latest.get(metric.key))
|
||||
baseline = baseline_value(baseline_records, metric.key)
|
||||
regression = regression_percent(metric.key, current, baseline)
|
||||
metrics[metric.key] = {
|
||||
failing_metrics: list[str] = []
|
||||
threshold_exceeded_metrics: list[str] = []
|
||||
for policy in metric_policies:
|
||||
current = safe_float(latest.get(policy.key))
|
||||
baseline = baseline_value(baseline_records, policy.key)
|
||||
delta = None
|
||||
if current is not None and baseline is not None:
|
||||
delta = regression_delta(policy, current, baseline)
|
||||
regression = None if delta is None else delta.percent * 100.0
|
||||
metrics[policy.key] = {
|
||||
"current": current,
|
||||
"baseline": baseline,
|
||||
"regression_pct": regression,
|
||||
"label": metric.label,
|
||||
"lower_is_better": metric.lower_is_better,
|
||||
"precision": metric.precision,
|
||||
"absolute_delta": None if delta is None else delta.absolute,
|
||||
"threshold_percent": policy.threshold_percent * 100.0,
|
||||
"threshold_absolute": policy.threshold_absolute,
|
||||
"gated": policy.gated,
|
||||
"threshold_exceeded": False if delta is None else delta.threshold_exceeded,
|
||||
"regressed": False if delta is None else delta.regressed,
|
||||
"label": policy.label,
|
||||
"lower_is_better": policy.lower_is_better,
|
||||
"precision": policy.precision,
|
||||
}
|
||||
if regression is not None:
|
||||
regressions.append(regression)
|
||||
if delta is not None and delta.threshold_exceeded:
|
||||
threshold_exceeded_metrics.append(policy.key)
|
||||
if delta is not None and delta.regressed:
|
||||
failing_metrics.append(policy.key)
|
||||
|
||||
worst_regression = max(regressions) if regressions else None
|
||||
success = bool(latest.get("success", True))
|
||||
status = "pass" if success else "fail"
|
||||
|
||||
rows.append({
|
||||
"model_id":
|
||||
model_id,
|
||||
"gpu_type":
|
||||
gpu_type,
|
||||
"timestamp":
|
||||
latest.get("timestamp"),
|
||||
"commit_sha":
|
||||
latest.get("commit_sha"),
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"timestamp": latest.get("timestamp"),
|
||||
"commit_sha": latest.get("commit_sha"),
|
||||
**record_metadata(latest),
|
||||
"success":
|
||||
success,
|
||||
"baseline_n":
|
||||
len(baseline_records),
|
||||
"worst_regression_pct":
|
||||
worst_regression,
|
||||
"regression_threshold_pct":
|
||||
max_regression * 100.0,
|
||||
"computed_regression_status":
|
||||
"fail" if worst_regression is not None and worst_regression > max_regression * 100.0 else "pass",
|
||||
"status":
|
||||
status,
|
||||
"metrics":
|
||||
metrics,
|
||||
"success": success,
|
||||
"baseline_n": len(baseline_records),
|
||||
"worst_regression_pct": worst_regression,
|
||||
"threshold_exceeded_metrics": threshold_exceeded_metrics,
|
||||
"failing_metrics": failing_metrics,
|
||||
"computed_regression_status": "fail" if failing_metrics else "pass",
|
||||
"status": status,
|
||||
"metrics": metrics,
|
||||
})
|
||||
|
||||
return sorted(rows, key=lambda row: (row["status"] != "fail", row["model_id"], row["gpu_type"]))
|
||||
@@ -179,14 +174,15 @@ def build_trends(records: list[Record]) -> list[Record]:
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
points = []
|
||||
for record in group:
|
||||
metric_policies = resolve_metric_policies(record.get("regression_thresholds"))
|
||||
point = {
|
||||
"timestamp": record.get("timestamp"),
|
||||
"commit_sha": record.get("commit_sha"),
|
||||
**record_metadata(record),
|
||||
"success": bool(record.get("success", True)),
|
||||
"metrics": {
|
||||
metric.key: safe_float(record.get(metric.key))
|
||||
for metric in METRICS
|
||||
policy.key: safe_float(record.get(policy.key))
|
||||
for policy in metric_policies
|
||||
},
|
||||
}
|
||||
points.append(point)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,89 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.pipelines.stages.kandinsky5 import (
|
||||
Kandinsky5DecodingStage,
|
||||
Kandinsky5DenoisingStage,
|
||||
Kandinsky5ImageEncodingStage,
|
||||
Kandinsky5LatentPreparationStage,
|
||||
Kandinsky5NormalizationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
from fastvideo.pipelines.stages.timestep_preparation import TimestepPreparationStage
|
||||
|
||||
|
||||
class Kandinsky5I2VPipeline(ComposedPipelineBase):
|
||||
"""Kandinsky-5.0 image-to-video pipeline."""
|
||||
|
||||
_required_config_modules = [
|
||||
"scheduler",
|
||||
"text_encoder",
|
||||
"text_encoder_2",
|
||||
"tokenizer",
|
||||
"tokenizer_2",
|
||||
"transformer",
|
||||
"vae",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="text_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2"),
|
||||
],
|
||||
tokenizers=[
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2"),
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=Kandinsky5LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer"),
|
||||
),
|
||||
)
|
||||
|
||||
# Runs AFTER latent preparation so the initial noise is the seeded
|
||||
# generator's first draw (official kandinsky-5 RNG order); this stage
|
||||
# then samples the image latent and places it into the latents.
|
||||
self.add_stage(
|
||||
stage_name="image_encoding_stage",
|
||||
stage=Kandinsky5ImageEncodingStage(vae=self.get_module("vae")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=Kandinsky5DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="normalization_stage",
|
||||
stage=Kandinsky5NormalizationStage(),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=Kandinsky5DecodingStage(vae=self.get_module("vae"), pipeline=self),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = Kandinsky5I2VPipeline
|
||||
@@ -0,0 +1,74 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.pipelines.stages.kandinsky5 import (
|
||||
Kandinsky5DecodingStage,
|
||||
Kandinsky5DenoisingStage,
|
||||
Kandinsky5LatentPreparationStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
from fastvideo.pipelines.stages.timestep_preparation import TimestepPreparationStage
|
||||
|
||||
|
||||
class Kandinsky5T2VPipeline(ComposedPipelineBase):
|
||||
"""Kandinsky-5.0 Lite text-to-video pipeline."""
|
||||
|
||||
_required_config_modules = [
|
||||
"scheduler",
|
||||
"text_encoder",
|
||||
"text_encoder_2",
|
||||
"tokenizer",
|
||||
"tokenizer_2",
|
||||
"transformer",
|
||||
"vae",
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="text_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2"),
|
||||
],
|
||||
tokenizers=[
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2"),
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=Kandinsky5LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=Kandinsky5DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=Kandinsky5DecodingStage(vae=self.get_module("vae"), pipeline=self),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = Kandinsky5T2VPipeline
|
||||
@@ -0,0 +1,165 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Kandinsky-5 model family pipeline presets."""
|
||||
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_NEGATIVE_PROMPT = ("Static, 2D cartoon, cartoon, 2d animation, paintings, images, worst quality, low quality, ugly, "
|
||||
"deformed, walking backwards")
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Main Kandinsky-5 denoising pass",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
KANDINSKY5_T2V_LITE_5S = InferencePreset(
|
||||
name="kandinsky5_t2v_lite_5s",
|
||||
version=1,
|
||||
model_family="kandinsky5",
|
||||
description="Kandinsky-5.0 Lite T2V 5s",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 50,
|
||||
"negative_prompt": _NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
KANDINSKY5_T2V_LITE_DISTILLED_5S = InferencePreset(
|
||||
name="kandinsky5_t2v_lite_distilled_5s",
|
||||
version=1,
|
||||
model_family="kandinsky5",
|
||||
description="Kandinsky-5.0 Lite T2V Distilled 5s",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 16,
|
||||
"negative_prompt": _NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
KANDINSKY5_T2V_PRO_5S = InferencePreset(
|
||||
name="kandinsky5_t2v_pro_5s",
|
||||
version=1,
|
||||
model_family="kandinsky5",
|
||||
description="Kandinsky-5.0 Pro T2V 5s",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 50,
|
||||
"negative_prompt": _NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
KANDINSKY5_T2V_PRO_DISTILLED_5S = InferencePreset(
|
||||
name="kandinsky5_t2v_pro_distilled_5s",
|
||||
version=1,
|
||||
model_family="kandinsky5",
|
||||
description="Kandinsky-5.0 Pro T2V Distilled 5s",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 16,
|
||||
"negative_prompt": _NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
KANDINSKY5_I2V_LITE_5S = InferencePreset(
|
||||
name="kandinsky5_i2v_lite_5s",
|
||||
version=1,
|
||||
model_family="kandinsky5",
|
||||
description="Kandinsky-5.0 Lite I2V 5s",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 50,
|
||||
"negative_prompt": _NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
KANDINSKY5_I2V_PRO_5S = InferencePreset(
|
||||
name="kandinsky5_i2v_pro_5s",
|
||||
version=1,
|
||||
model_family="kandinsky5",
|
||||
description="Kandinsky-5.0 Pro I2V 5s",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 50,
|
||||
"negative_prompt": _NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
KANDINSKY5_I2V_LITE_DISTILLED_5S = InferencePreset(
|
||||
name="kandinsky5_i2v_lite_distilled_5s",
|
||||
version=1,
|
||||
model_family="kandinsky5",
|
||||
description="Kandinsky-5.0 Lite I2V Distilled 5s",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 16,
|
||||
"negative_prompt": _NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
KANDINSKY5_I2V_PRO_DISTILLED_5S = InferencePreset(
|
||||
name="kandinsky5_i2v_pro_distilled_5s",
|
||||
version=1,
|
||||
model_family="kandinsky5",
|
||||
description="Kandinsky-5.0 Pro I2V Distilled 5s",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 512,
|
||||
"width": 768,
|
||||
"num_frames": 121,
|
||||
"fps": 24,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 16,
|
||||
"negative_prompt": _NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (KANDINSKY5_T2V_LITE_5S, KANDINSKY5_T2V_LITE_DISTILLED_5S, KANDINSKY5_T2V_PRO_5S,
|
||||
KANDINSKY5_T2V_PRO_DISTILLED_5S, KANDINSKY5_I2V_LITE_5S, KANDINSKY5_I2V_LITE_DISTILLED_5S,
|
||||
KANDINSKY5_I2V_PRO_5S, KANDINSKY5_I2V_PRO_DISTILLED_5S)
|
||||
@@ -33,6 +33,8 @@ from fastvideo.pipelines.basic.ltx2.stages import ( # noqa: F401
|
||||
from fastvideo.pipelines.stages.matrixgame2_denoising import MatrixGame2CausalDenoisingStage
|
||||
from fastvideo.pipelines.stages.matrixgame3_denoising import MatrixGame3DenoisingStage
|
||||
from fastvideo.pipelines.stages.hyworld_denoising import HYWorldDenoisingStage
|
||||
from fastvideo.pipelines.stages.kandinsky5 import (Kandinsky5DecodingStage, Kandinsky5DenoisingStage,
|
||||
Kandinsky5LatentPreparationStage)
|
||||
from fastvideo.pipelines.stages.gamecraft_denoising import GameCraftDenoisingStage
|
||||
from fastvideo.pipelines.stages.gen3c_stages import (Gen3CCFGPolicyStage, Gen3CConditioningStage, Gen3CDenoisingStage,
|
||||
Gen3CLatentPreparationStage)
|
||||
@@ -65,6 +67,9 @@ __all__ = [
|
||||
"MatrixGame2CausalDenoisingStage",
|
||||
"MatrixGame3DenoisingStage",
|
||||
"HYWorldDenoisingStage",
|
||||
"Kandinsky5DecodingStage",
|
||||
"Kandinsky5DenoisingStage",
|
||||
"Kandinsky5LatentPreparationStage",
|
||||
"GameCraftDenoisingStage",
|
||||
"Gen3CCFGPolicyStage",
|
||||
"Gen3CConditioningStage",
|
||||
|
||||
@@ -34,6 +34,7 @@ class PipelineStage(ABC):
|
||||
composed with other stages to create a complete pipeline. Each stage is responsible
|
||||
for a specific part of the process, such as prompt encoding, latent preparation, etc.
|
||||
"""
|
||||
performance_component_metric: str | None = None
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""
|
||||
@@ -155,6 +156,9 @@ class PipelineStage(ABC):
|
||||
logger.info("[%s] Execution completed in %s ms", stage_name, execution_time * 1000)
|
||||
batch.logging_info.add_stage_execution_time(stage_key, execution_time)
|
||||
batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name)
|
||||
component_metric = self.performance_component_metric
|
||||
if component_metric is not None:
|
||||
batch.logging_info.add_stage_metric(stage_key, "component_metric", component_metric)
|
||||
except Exception as e:
|
||||
torch.cuda.synchronize()
|
||||
execution_time = time.perf_counter() - start_time
|
||||
|
||||
@@ -28,6 +28,7 @@ class DecodingStage(PipelineStage):
|
||||
This stage handles the decoding of latent representations into the final
|
||||
output format (e.g., pixel values).
|
||||
"""
|
||||
performance_component_metric = "vae_decode_time_s"
|
||||
|
||||
def __init__(self, vae, pipeline=None) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
@@ -51,6 +51,7 @@ class DenoisingStage(PipelineStage):
|
||||
This stage handles the iterative denoising process that transforms
|
||||
the initial noise into the final output.
|
||||
"""
|
||||
performance_component_metric = "dit_time_s"
|
||||
|
||||
def __init__(self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None) -> None:
|
||||
super().__init__()
|
||||
@@ -1187,6 +1188,7 @@ class Cosmos25V2WDenoisingStage(Cosmos25DenoisingStage):
|
||||
|
||||
class Cosmos25AutoDenoisingStage(PipelineStage):
|
||||
"""Route Cosmos 2.5 denoising to T2W vs V2W/I2W."""
|
||||
performance_component_metric = "dit_time_s"
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
|
||||
@@ -0,0 +1,534 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from typing import Any
|
||||
|
||||
import PIL
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.attention.backends.nabla import NablaAttentionMetadataBuilder
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.loader.component_loader import TransformerLoader, VAELoader
|
||||
from fastvideo.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.models.vision_utils import normalize, numpy_to_pt, pil_to_numpy, resize
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class Kandinsky5LatentPreparationStage(PipelineStage):
|
||||
|
||||
def __init__(self, scheduler, transformer) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
self.transformer = transformer
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.height is None or batch.width is None:
|
||||
raise ValueError("height and width must be provided for Kandinsky5.")
|
||||
height = int(batch.height)
|
||||
width = int(batch.width)
|
||||
num_frames = int(batch.num_frames)
|
||||
|
||||
temporal_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
|
||||
|
||||
if num_frames % temporal_ratio != 1:
|
||||
num_frames = num_frames // temporal_ratio * temporal_ratio + 1
|
||||
batch.num_frames = num_frames
|
||||
|
||||
required_divisor = spatial_ratio * patch_size[1]
|
||||
if height % required_divisor != 0 or width % required_divisor != 0:
|
||||
raise ValueError(f"Kandinsky5 height/width must be divisible by {required_divisor}; "
|
||||
f"got height={height}, width={width}.")
|
||||
|
||||
# NABLA sparse attention (Pro checkpoints) reshapes the post-patch grid
|
||||
# into 8x8 blocks; validate here instead of crashing mid-denoise after
|
||||
# all the encoding work is done.
|
||||
arch_cfg = getattr(self.transformer, "config", None) or fastvideo_args.pipeline_config.dit_config.arch_config
|
||||
if getattr(arch_cfg, "attention_type", "regular") == "nabla":
|
||||
nabla_divisor = required_divisor * 8
|
||||
if height % nabla_divisor != 0 or width % nabla_divisor != 0:
|
||||
raise ValueError(f"Kandinsky5 NABLA checkpoints require height/width divisible by {nabla_divisor}; "
|
||||
f"got height={height}, width={width}.")
|
||||
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
batch_size = 1
|
||||
else:
|
||||
batch_size = batch.prompt_embeds[0].shape[0]
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
device = get_local_torch_device()
|
||||
num_latent_frames = (num_frames - 1) // temporal_ratio + 1
|
||||
num_channels = getattr(
|
||||
self.transformer,
|
||||
"in_visual_dim",
|
||||
fastvideo_args.pipeline_config.dit_config.arch_config.in_visual_dim,
|
||||
)
|
||||
shape = (
|
||||
batch_size,
|
||||
num_latent_frames,
|
||||
height // spatial_ratio,
|
||||
width // spatial_ratio,
|
||||
num_channels,
|
||||
)
|
||||
|
||||
if isinstance(batch.generator, list) and len(batch.generator) != batch_size:
|
||||
raise ValueError(f"generator list length {len(batch.generator)} does not match batch size {batch_size}.")
|
||||
|
||||
if batch.latents is None:
|
||||
latents = randn_tensor(shape, generator=batch.generator, device=device, dtype=dtype)
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
else:
|
||||
latents = batch.latents.to(device=device, dtype=dtype)
|
||||
|
||||
visual_cond = getattr(self.transformer, "visual_cond", False)
|
||||
if visual_cond and latents.shape[-1] == num_channels:
|
||||
cond = torch.zeros_like(latents)
|
||||
cond_mask = torch.zeros(
|
||||
(*latents.shape[:-1], 1),
|
||||
device=latents.device,
|
||||
dtype=latents.dtype,
|
||||
)
|
||||
latents = torch.cat([latents, cond, cond_mask], dim=-1)
|
||||
|
||||
# I2V image conditioning is placed by Kandinsky5ImageEncodingStage,
|
||||
# which runs AFTER this stage so the initial noise is the generator's
|
||||
# first draw (matching the official kandinskylab/kandinsky-5 order).
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = (
|
||||
batch_size,
|
||||
num_channels,
|
||||
num_latent_frames,
|
||||
height // spatial_ratio,
|
||||
width // spatial_ratio,
|
||||
)
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("latents", batch.latents, V.none_or_tensor)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
|
||||
class Kandinsky5DenoisingStage(PipelineStage):
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
|
||||
@staticmethod
|
||||
def _scale_factor(height: int, width: int) -> tuple[float, float, float]:
|
||||
if 480 <= height <= 854 and 480 <= width <= 854:
|
||||
return (1.0, 2.0, 2.0)
|
||||
return (1.0, 3.16, 3.16)
|
||||
|
||||
@staticmethod
|
||||
def _text_rope_pos(mask: torch.Tensor, device: torch.device) -> torch.Tensor:
|
||||
seq_len = int(mask.sum(1).max().item())
|
||||
return torch.arange(seq_len, device=device)
|
||||
|
||||
@staticmethod
|
||||
def fast_sta_nabla(
|
||||
T: int,
|
||||
H: int,
|
||||
W: int,
|
||||
wT: int = 3,
|
||||
wH: int = 3,
|
||||
wW: int = 3,
|
||||
device: torch.device | str = "cuda",
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Create a sparse temporal attention (STA) mask for efficient video generation.
|
||||
|
||||
This method generates a mask that limits attention to nearby frames and spatial positions, reducing
|
||||
computational complexity for video generation.
|
||||
|
||||
Args:
|
||||
T (int): Number of temporal frames
|
||||
H (int): Height in latent space
|
||||
W (int): Width in latent space
|
||||
wT (int): Temporal attention window size
|
||||
wH (int): Height attention window size
|
||||
wW (int): Width attention window size
|
||||
device (str): Device to create tensor on
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Sparse attention mask of shape (T*H*W, T*H*W)
|
||||
"""
|
||||
max_extent = int(torch.tensor([T, H, W], device=device).amax().item())
|
||||
r = torch.arange(0, max_extent, 1, dtype=torch.int16, device=device)
|
||||
mat = (r.unsqueeze(1) - r.unsqueeze(0)).abs()
|
||||
sta_t, sta_h, sta_w = (
|
||||
mat[:T, :T].flatten(),
|
||||
mat[:H, :H].flatten(),
|
||||
mat[:W, :W].flatten(),
|
||||
)
|
||||
sta_t = sta_t <= wT // 2
|
||||
sta_h = sta_h <= wH // 2
|
||||
sta_w = sta_w <= wW // 2
|
||||
sta_hw = (sta_h.unsqueeze(1) * sta_w.unsqueeze(0)).reshape(H, H, W, W).transpose(1, 2).flatten()
|
||||
sta = (sta_t.unsqueeze(1) * sta_hw.unsqueeze(0)).reshape(T, T, H * W, H * W).transpose(1, 2)
|
||||
return sta.reshape(T * H * W, T * H * W)
|
||||
|
||||
def get_sparse_params(self, sample: torch.Tensor, device: torch.device) -> dict[str, Any] | None:
|
||||
"""
|
||||
Generate sparse attention parameters for the transformer based on sample dimensions.
|
||||
|
||||
This method computes the sparse attention configuration needed for efficient video processing in the
|
||||
transformer model.
|
||||
|
||||
Args:
|
||||
sample (torch.Tensor): Input sample tensor
|
||||
device (torch.device): Device to place tensors on
|
||||
|
||||
Returns:
|
||||
Dict: Dictionary containing sparse attention parameters
|
||||
"""
|
||||
assert self.transformer.config.patch_size[0] == 1
|
||||
_, T, H, W, _ = sample.shape
|
||||
T, H, W = (
|
||||
T // self.transformer.config.patch_size[0],
|
||||
H // self.transformer.config.patch_size[1],
|
||||
W // self.transformer.config.patch_size[2],
|
||||
)
|
||||
if self.transformer.config.attention_type == "nabla":
|
||||
sta_mask = self.fast_sta_nabla(
|
||||
T,
|
||||
H // 8,
|
||||
W // 8,
|
||||
self.transformer.config.attention_wT,
|
||||
self.transformer.config.attention_wH,
|
||||
self.transformer.config.attention_wW,
|
||||
device=device,
|
||||
)
|
||||
|
||||
sparse_params = {
|
||||
"sta_mask": sta_mask.unsqueeze_(0).unsqueeze_(0),
|
||||
"attention_type": self.transformer.config.attention_type,
|
||||
"to_fractal": True,
|
||||
"P": self.transformer.config.attention_P,
|
||||
"wT": self.transformer.config.attention_wT,
|
||||
"wW": self.transformer.config.attention_wW,
|
||||
"wH": self.transformer.config.attention_wH,
|
||||
"add_sta": self.transformer.config.attention_add_sta,
|
||||
"visual_shape": (T, H, W),
|
||||
"method": self.transformer.config.attention_method,
|
||||
}
|
||||
else:
|
||||
sparse_params = None
|
||||
|
||||
return sparse_params
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.timesteps is None:
|
||||
raise ValueError("timesteps must be prepared before Kandinsky5 denoising.")
|
||||
if batch.latents is None:
|
||||
raise ValueError("latents must be prepared before Kandinsky5 denoising.")
|
||||
if not fastvideo_args.model_loaded["transformer"]:
|
||||
loader = TransformerLoader()
|
||||
self.transformer = loader.load(fastvideo_args.model_paths["transformer"], fastvideo_args)
|
||||
fastvideo_args.model_loaded["transformer"] = True
|
||||
|
||||
device = get_local_torch_device()
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
autocast_enabled = target_dtype != torch.float32 and not fastvideo_args.disable_autocast
|
||||
latents = batch.latents
|
||||
num_channels = getattr(
|
||||
self.transformer,
|
||||
"in_visual_dim",
|
||||
fastvideo_args.pipeline_config.dit_config.arch_config.in_visual_dim,
|
||||
)
|
||||
|
||||
prompt_embeds = batch.prompt_embeds[0].to(device=device, dtype=target_dtype)
|
||||
pooled = batch.prompt_embeds[1].to(device=device, dtype=target_dtype)
|
||||
if batch.prompt_attention_mask is None or not batch.prompt_attention_mask:
|
||||
raise ValueError("Kandinsky5 requires Qwen prompt attention masks.")
|
||||
text_rope_pos = self._text_rope_pos(batch.prompt_attention_mask[0].to(device), device)
|
||||
|
||||
neg_prompt_embeds = None
|
||||
neg_pooled = None
|
||||
negative_text_rope_pos = None
|
||||
if batch.do_classifier_free_guidance and batch.negative_prompt_embeds:
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds[0].to(device=device, dtype=target_dtype)
|
||||
neg_pooled = batch.negative_prompt_embeds[1].to(device=device, dtype=target_dtype)
|
||||
if batch.negative_attention_mask is None or not batch.negative_attention_mask:
|
||||
raise ValueError("Kandinsky5 requires Qwen negative attention masks for CFG.")
|
||||
negative_text_rope_pos = self._text_rope_pos(batch.negative_attention_mask[0].to(device), device)
|
||||
|
||||
height = int(batch.height)
|
||||
width = int(batch.width)
|
||||
temporal_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
|
||||
num_latent_frames = (int(batch.num_frames) - 1) // temporal_ratio + 1
|
||||
visual_rope_pos = [
|
||||
torch.arange(num_latent_frames, device=device),
|
||||
torch.arange(height // spatial_ratio // 2, device=device),
|
||||
torch.arange(width // spatial_ratio // 2, device=device),
|
||||
]
|
||||
scale_factor = self._scale_factor(height, width)
|
||||
|
||||
sparse_params = self.get_sparse_params(latents, device)
|
||||
|
||||
# I2V keeps the first (conditioning) frame fixed during denoising.
|
||||
# Key off the actual image conditioning, not transformer.visual_cond:
|
||||
# official T2V checkpoints also ship visual_cond=True, and skipping
|
||||
# frame 0 for them leaves it as undenoised noise.
|
||||
cond_frames = 1 if batch.image_latent is not None else 0
|
||||
|
||||
trajectory_timesteps: list[torch.Tensor] = []
|
||||
trajectory_latents: list[torch.Tensor] = []
|
||||
|
||||
with tqdm(total=batch.num_inference_steps, desc="Kandinsky5 Denoising") as progress_bar:
|
||||
for i, timestep in enumerate(batch.timesteps):
|
||||
if hasattr(self, "interrupt") and self.interrupt:
|
||||
continue
|
||||
|
||||
t_expand = timestep.unsqueeze(0).repeat(latents.shape[0]).to(device=device, dtype=target_dtype)
|
||||
attn_metadata = None
|
||||
if sparse_params is not None:
|
||||
attn_metadata = NablaAttentionMetadataBuilder().build(
|
||||
current_timestep=i,
|
||||
sta_mask=sparse_params["sta_mask"],
|
||||
P=sparse_params["P"],
|
||||
visual_shape=sparse_params["visual_shape"],
|
||||
)
|
||||
autocast_ctx = (torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled)
|
||||
if device.type == "cuda" else contextlib.nullcontext())
|
||||
with set_forward_context(current_timestep=i, attn_metadata=attn_metadata,
|
||||
forward_batch=batch), autocast_ctx:
|
||||
pred_velocity = self.transformer(
|
||||
hidden_states=latents.to(dtype=target_dtype),
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
pooled_projections=pooled,
|
||||
timestep=t_expand,
|
||||
visual_rope_pos=visual_rope_pos,
|
||||
text_rope_pos=text_rope_pos,
|
||||
scale_factor=scale_factor,
|
||||
sparse_params=sparse_params,
|
||||
return_dict=True,
|
||||
).sample
|
||||
|
||||
if neg_prompt_embeds is not None and neg_pooled is not None:
|
||||
uncond_pred_velocity = self.transformer(
|
||||
hidden_states=latents.to(dtype=target_dtype),
|
||||
encoder_hidden_states=neg_prompt_embeds,
|
||||
pooled_projections=neg_pooled,
|
||||
timestep=t_expand,
|
||||
visual_rope_pos=visual_rope_pos,
|
||||
text_rope_pos=negative_text_rope_pos,
|
||||
scale_factor=scale_factor,
|
||||
sparse_params=sparse_params,
|
||||
return_dict=True,
|
||||
).sample
|
||||
pred_velocity = uncond_pred_velocity + batch.guidance_scale * (pred_velocity -
|
||||
uncond_pred_velocity)
|
||||
|
||||
latents[:, cond_frames:, :, :, :num_channels] = self.scheduler.step(
|
||||
pred_velocity[:, cond_frames:],
|
||||
timestep,
|
||||
latents[:, cond_frames:, :, :, :num_channels],
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if batch.return_trajectory_latents:
|
||||
trajectory_timesteps.append(timestep)
|
||||
# latents is mutated in place, so snapshot a channels-first copy.
|
||||
trajectory_latents.append(latents[..., :num_channels].permute(0, 4, 1, 2, 3).cpu())
|
||||
|
||||
if i == len(batch.timesteps) - 1 or (i + 1) % self.scheduler.order == 0:
|
||||
progress_bar.update()
|
||||
|
||||
if trajectory_latents:
|
||||
batch.trajectory_latents = torch.stack(trajectory_latents, dim=1)
|
||||
batch.trajectory_timesteps = torch.stack(trajectory_timesteps, dim=0).cpu()
|
||||
|
||||
batch.latents = latents[:, :, :, :, :num_channels]
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
return result
|
||||
|
||||
|
||||
class Kandinsky5DecodingStage(DecodingStage):
|
||||
|
||||
def __init__(self, vae: ParallelTiledVAE, pipeline=None) -> None:
|
||||
super().__init__(vae=vae, pipeline=pipeline)
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.latents is None:
|
||||
raise ValueError("latents must be available before Kandinsky5 decoding.")
|
||||
# Kandinsky5 latents are channels-last [B, T, H, W, C]; the base stage
|
||||
# (and the trajectory latents recorded by the denoising stage) work
|
||||
# channels-first.
|
||||
batch.latents = batch.latents.permute(0, 4, 1, 2, 3).contiguous()
|
||||
return super().forward(batch, fastvideo_args)
|
||||
|
||||
|
||||
class Kandinsky5ImageEncodingStage(EncodingStage):
|
||||
"""Encode the conditioning image into a VAE latent for I2V."""
|
||||
|
||||
def __init__(self, vae: ParallelTiledVAE, pipeline=None) -> None:
|
||||
super().__init__(vae=vae)
|
||||
|
||||
@staticmethod
|
||||
def _preprocess(image, height: int, width: int) -> torch.Tensor:
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
image = resize(image, height, width)
|
||||
image = numpy_to_pt(pil_to_numpy(image)) # always lands in [0, 1]
|
||||
return normalize(image) # [0, 1] -> [-1, 1]
|
||||
# Tensor input: no reliable way to tell [0, 1] from an already
|
||||
# normalized [-1, 1] tensor whose values happen to be non-negative,
|
||||
# so mirror diffusers' heuristic and say what we assumed.
|
||||
if image.min() >= 0:
|
||||
logger.warning("Kandinsky5 conditioning image tensor has no negative values; "
|
||||
"assuming range [0, 1] and normalizing to [-1, 1]. "
|
||||
"Pass a [-1, 1] tensor with negative values to skip normalization.")
|
||||
image = normalize(image) # [0, 1] -> [-1, 1]
|
||||
if image.ndim == 3:
|
||||
image = image.unsqueeze(0)
|
||||
if image.shape[-2:] != (height, width):
|
||||
image = torch.nn.functional.interpolate(image.float(),
|
||||
size=(height, width),
|
||||
mode="bilinear",
|
||||
antialias=True)
|
||||
return image
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.pil_image is None:
|
||||
raise ValueError("Kandinsky5 I2V requires an input image.")
|
||||
|
||||
if not fastvideo_args.model_loaded["vae"]:
|
||||
vae = getattr(self, "vae", None)
|
||||
if vae is None:
|
||||
loader = VAELoader()
|
||||
vae = loader.load(fastvideo_args.model_paths["vae"], fastvideo_args)
|
||||
self.vae = vae
|
||||
fastvideo_args.model_loaded["vae"] = True
|
||||
|
||||
device = get_local_torch_device()
|
||||
vae = self.vae.to(device)
|
||||
self.vae = vae
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = vae_dtype != torch.float32 and not fastvideo_args.disable_autocast
|
||||
|
||||
# [B, C, H, W] -> [B, C, 1, H, W]
|
||||
image = self._preprocess(batch.pil_image, int(batch.height), int(batch.width))
|
||||
image = image.to(device=device, dtype=torch.float32).unsqueeze(2)
|
||||
|
||||
# Encode the single conditioning frame without tiling (matches diffusers).
|
||||
# The untested causal-VAE spatial_tiled_encode path corrupts the latent.
|
||||
prev_use_tiling = vae.use_tiling
|
||||
vae.use_tiling = False
|
||||
try:
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
|
||||
if not vae_autocast_enabled:
|
||||
image = image.to(vae_dtype)
|
||||
# Sample with the batch generator (diffusers parity); mode()
|
||||
# would make seed-for-seed reproduction of the reference
|
||||
# pipeline impossible.
|
||||
generator = batch.generator
|
||||
if isinstance(generator, list) and len(generator) != image.shape[0]:
|
||||
generator = generator[0]
|
||||
image_latent = vae.encode(image).sample(generator=generator)
|
||||
finally:
|
||||
vae.use_tiling = prev_use_tiling
|
||||
|
||||
image_latent = image_latent * vae.scaling_factor
|
||||
# [B, C, 1, H, W] -> [B, 1, H, W, C] to match channels-last latents
|
||||
batch.image_latent = image_latent.permute(0, 2, 3, 4, 1).contiguous()
|
||||
|
||||
# Place the conditioning latent into the prepared latents: frame 0 of
|
||||
# the main channels, the visual_cond channel block, and the mask.
|
||||
# NOTE: the official kandinsky-5 repo leaves the visual_cond block
|
||||
# zeros (generation_utils.py generate()), while the diffusers port
|
||||
# copies the image latent into it. A same-seed A/B on
|
||||
# Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers showed the Diffusers
|
||||
# export requires the copy: zeroing the block produces smeared faces
|
||||
# mid-video. Keep the diffusers semantics for Diffusers-format
|
||||
# checkpoints.
|
||||
latents = batch.latents
|
||||
image_latent = batch.image_latent.to(device=latents.device, dtype=latents.dtype)
|
||||
num_channels = image_latent.shape[-1]
|
||||
latents[:, 0:1, :, :, :num_channels] = image_latent
|
||||
if latents.shape[-1] > num_channels:
|
||||
latents[:, 0:1, :, :, num_channels:2 * num_channels] = image_latent
|
||||
latents[:, 0:1, :, :, 2 * num_channels:] = 1.0
|
||||
batch.latents = latents
|
||||
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
vae.to("cpu")
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("pil_image", batch.pil_image, V.not_none)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
# This stage runs after latent preparation and writes into its output.
|
||||
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
result.add_check("image_latent", batch.image_latent, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
|
||||
class Kandinsky5NormalizationStage(PipelineStage):
|
||||
"""Normalize the first latent frames to reduce I2V conditioning artifacts."""
|
||||
|
||||
COND_FRAMES = 4
|
||||
REFERENCE_FRAMES = 5
|
||||
|
||||
@staticmethod
|
||||
def _adaptive_mean_std(source: torch.Tensor, reference: torch.Tensor) -> torch.Tensor:
|
||||
source_mean = source.mean(dim=(1, 2, 3, 4), keepdim=True)
|
||||
source_std = source.std(dim=(1, 2, 3, 4), keepdim=True)
|
||||
# Magic constants limit how far the first frames may drift.
|
||||
ref_mean = torch.clamp(reference.mean(dim=(1, 2, 3, 4), keepdim=True), source_mean - 0.05, source_mean + 0.1)
|
||||
ref_std = torch.clamp(reference.std(dim=(1, 2, 3, 4), keepdim=True), source_std - 0.1, source_std + 0.25)
|
||||
normalized = (source - source_mean) / source_std
|
||||
return normalized * ref_std + ref_mean
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
latents = batch.latents
|
||||
n = self.COND_FRAMES
|
||||
if latents is None or latents.shape[1] <= n:
|
||||
return batch
|
||||
|
||||
reference = latents[:, n:n + min(self.REFERENCE_FRAMES, latents.shape[1] - 1)]
|
||||
latents[:, :n] = self._adaptive_mean_std(latents[:, :n].clone(), reference)
|
||||
batch.latents = latents
|
||||
return batch
|
||||
@@ -8,6 +8,8 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
import torch
|
||||
from typing import Any
|
||||
|
||||
from torch.distributed.tensor import DTensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
@@ -24,6 +26,7 @@ class TextEncodingStage(PipelineStage):
|
||||
This stage handles the encoding of text prompts into the embedding space
|
||||
expected by the diffusion model.
|
||||
"""
|
||||
performance_component_metric = "text_encoder_time_s"
|
||||
|
||||
def __init__(self, text_encoders, tokenizers) -> None:
|
||||
"""
|
||||
@@ -202,6 +205,21 @@ class TextEncodingStage(PipelineStage):
|
||||
encoder_config = encoder_cfgs[i]
|
||||
preprocess_func = preprocess_funcs[i]
|
||||
postprocess_func = postprocess_funcs[i]
|
||||
# cpu_offload semantics: params rest on CPU between calls but the
|
||||
# forward computes on GPU. FSDP2-wrapped encoders (CPUOffloadPolicy;
|
||||
# DTensor params) stream themselves per-layer — leave inputs on the
|
||||
# param device and let FSDP's root pre-forward move them. A plain
|
||||
# module parked on CPU by text_encoder_cpu_offload is swapped to the
|
||||
# target device for the forward and back afterwards, mirroring the
|
||||
# image-encoder/VAE offload pattern.
|
||||
first_param = next(text_encoder.parameters(), None)
|
||||
encoder_device = first_param.device if first_param is not None else torch.device(target_device)
|
||||
moved_for_forward = False
|
||||
if (first_param is not None and not isinstance(first_param, DTensor)
|
||||
and encoder_device.type != torch.device(target_device).type):
|
||||
text_encoder = text_encoder.to(target_device)
|
||||
encoder_device = torch.device(target_device)
|
||||
moved_for_forward = True
|
||||
|
||||
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
|
||||
if max_length is not None:
|
||||
@@ -246,7 +264,7 @@ class TextEncodingStage(PipelineStage):
|
||||
# pre-format prompts into message lists upstream and rely on
|
||||
# the inner tokenizer + full tokenizer_kwargs (which include
|
||||
# add_generation_prompt). Preserve that original path exactly.
|
||||
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(target_device)
|
||||
text_inputs = tok.apply_chat_template(processed_texts, **tok_kwargs).to(encoder_device)
|
||||
else:
|
||||
# Two-step approach matching Diffusers: format with chat
|
||||
# template first, then tokenize the resulting strings.
|
||||
@@ -260,9 +278,9 @@ class TextEncodingStage(PipelineStage):
|
||||
enable_thinking=False,
|
||||
)
|
||||
formatted_texts.append(formatted)
|
||||
text_inputs = tokenizer(formatted_texts, **tok_kwargs).to(target_device)
|
||||
text_inputs = tokenizer(formatted_texts, **tok_kwargs).to(encoder_device)
|
||||
else:
|
||||
text_inputs = tok(processed_texts, **tok_kwargs).to(target_device)
|
||||
text_inputs = tok(processed_texts, **tok_kwargs).to(encoder_device)
|
||||
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
@@ -282,15 +300,18 @@ class TextEncodingStage(PipelineStage):
|
||||
except Exception:
|
||||
prompt_embeds, attention_mask = postprocess_func(outputs, attention_mask)
|
||||
if is_ltx2 and getattr(outputs, "hidden_states", None):
|
||||
audio_embed = outputs.hidden_states[0]
|
||||
audio_embed = outputs.hidden_states[0].to(device=target_device)
|
||||
if dtype is not None:
|
||||
audio_embed = audio_embed.to(dtype=dtype)
|
||||
audio_embeds_list.append(audio_embed)
|
||||
prompt_embeds = prompt_embeds.to(device=target_device)
|
||||
if dtype is not None:
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype)
|
||||
embeds_list.append(prompt_embeds)
|
||||
if return_attention_mask:
|
||||
attn_masks_list.append(attention_mask)
|
||||
attn_masks_list.append(attention_mask.to(device=target_device))
|
||||
if moved_for_forward and fastvideo_args.text_encoder_cpu_offload:
|
||||
text_encoder.to("cpu")
|
||||
self._last_audio_embeds = audio_embeds_list if is_ltx2 else None
|
||||
return self.return_embeds(embeds_list, attn_masks_list, return_type, return_attention_mask, indices)
|
||||
|
||||
@@ -350,6 +371,7 @@ class Cosmos25TextEncodingStage(PipelineStage):
|
||||
Cosmos 2.5 uses Reason1 (Qwen2.5-VL) and relies on the encoder's
|
||||
`compute_text_embeddings_online()`.
|
||||
"""
|
||||
performance_component_metric = "text_encoder_time_s"
|
||||
|
||||
def __init__(self, text_encoder) -> None:
|
||||
super().__init__()
|
||||
|
||||
@@ -157,6 +157,14 @@ class CudaPlatformBase(Platform):
|
||||
"ATTN_QAT_TRAIN selected but fastvideo_kernel.triton_kernels.attn_qat_train is not built. "
|
||||
"Silent fallback would produce a non-QAT training run; refusing to proceed. "
|
||||
"Install the training kernel or pick a different FASTVIDEO_ATTENTION_BACKEND.")
|
||||
elif selected_backend == AttentionBackendEnum.NABLA_ATTN:
|
||||
from fastvideo.attention.backends.nabla import CAN_USE_FLEX_ATTN
|
||||
if CAN_USE_FLEX_ATTN:
|
||||
logger.info("Using NABLA block-sparse flex-attention backend.")
|
||||
return "fastvideo.attention.backends.nabla.NablaAttentionBackend"
|
||||
raise ImportError("NABLA_ATTN selected but torch.nn.attention.flex_attention is unavailable in this "
|
||||
"PyTorch build. Silent fallback to dense attention would be orders of magnitude "
|
||||
"slower and diverge from the reference; upgrade PyTorch or pick a different backend.")
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from fastvideo_kernel import video_sparse_attn # noqa: F401
|
||||
|
||||
@@ -22,6 +22,7 @@ class AttentionBackendEnum(enum.Enum):
|
||||
VMOBA_ATTN = enum.auto()
|
||||
SLA_ATTN = enum.auto()
|
||||
SAGE_SLA_ATTN = enum.auto()
|
||||
NABLA_ATTN = enum.auto()
|
||||
NO_ATTENTION = enum.auto()
|
||||
|
||||
|
||||
|
||||
+179
-4
@@ -27,6 +27,7 @@ from fastvideo.configs.pipelines.hunyuan15 import (Hunyuan15T2V480PConfig, Hunyu
|
||||
Hunyuan15T2V720PConfig, Hunyuan15I2V720PConfig,
|
||||
Hunyuan15SR1080PConfig)
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
|
||||
@@ -117,6 +118,10 @@ class ConfigInfo:
|
||||
workload_types: tuple[WorkloadType, ...]
|
||||
model_family: str | None = None
|
||||
default_preset: str | None = None
|
||||
# When set, overrides the model_index `_class_name` for pipeline resolution.
|
||||
# Lets a model family map to a specific pipeline class by path/detector
|
||||
# (e.g. a T2V and I2V checkpoint that share a `_class_name`).
|
||||
pipeline_cls_name: str | None = None
|
||||
|
||||
|
||||
# The central registry mapping a model name to its configuration information
|
||||
@@ -137,6 +142,7 @@ def register_configs(
|
||||
model_detectors: list[Callable[[str], bool]] | None = None,
|
||||
model_family: str | None = None,
|
||||
default_preset: str | None = None,
|
||||
pipeline_cls_name: str | None = None,
|
||||
) -> None:
|
||||
"""Register config classes for a model family.
|
||||
|
||||
@@ -151,6 +157,7 @@ def register_configs(
|
||||
workload_types=workload_types,
|
||||
model_family=model_family,
|
||||
default_preset=default_preset,
|
||||
pipeline_cls_name=pipeline_cls_name,
|
||||
)
|
||||
|
||||
if hf_model_paths:
|
||||
@@ -490,18 +497,174 @@ def _register_configs() -> None:
|
||||
default_preset="lingbotworld_i2v",
|
||||
)
|
||||
|
||||
def _kandinsky5_detector(require: tuple[str, ...] = (), exclude: tuple[str, ...] = ()) -> Callable[[str], bool]:
|
||||
|
||||
def detect(path: str) -> bool:
|
||||
path_lower = path.lower()
|
||||
if "kandinsky5" not in path_lower and "kandinsky-5" not in path_lower:
|
||||
return False
|
||||
return (all(token in path_lower for token in require) and not any(token in path_lower for token in exclude))
|
||||
|
||||
return detect
|
||||
|
||||
# t2v/i2v exclude each other so a checkpoint stored under a directory
|
||||
# containing the other token (e.g. ~/i2v_experiments/kandinsky5-t2v-ft)
|
||||
# falls through to the model_index _class_name fallback detectors below
|
||||
# instead of being misrouted.
|
||||
_is_kandinsky5_t2v = _kandinsky5_detector(require=("t2v", ), exclude=("i2v", ))
|
||||
_is_kandinsky5_i2v = _kandinsky5_detector(require=("i2v", ), exclude=("t2v", ))
|
||||
_is_kandinsky5_t2v_lite = _kandinsky5_detector(require=("t2v", "lite"), exclude=("i2v", "distilled"))
|
||||
_is_kandinsky5_t2v_pro = _kandinsky5_detector(require=("t2v", "pro"), exclude=("i2v", "distilled"))
|
||||
_is_kandinsky5_t2v_lite_distilled = _kandinsky5_detector(require=("t2v", "lite", "distilled"), exclude=("i2v", ))
|
||||
_is_kandinsky5_t2v_pro_distilled = _kandinsky5_detector(require=("t2v", "pro", "distilled"), exclude=("i2v", ))
|
||||
_is_kandinsky5_i2v_lite = _kandinsky5_detector(require=("i2v", "lite"), exclude=("t2v", "distilled"))
|
||||
_is_kandinsky5_i2v_pro = _kandinsky5_detector(require=("i2v", "pro"), exclude=("t2v", "distilled"))
|
||||
_is_kandinsky5_i2v_lite_distilled = _kandinsky5_detector(require=("i2v", "lite", "distilled"), exclude=("t2v", ))
|
||||
_is_kandinsky5_i2v_pro_distilled = _kandinsky5_detector(require=("i2v", "pro", "distilled"), exclude=("t2v", ))
|
||||
|
||||
# Kandinsky5 Lite T2V
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=PipelineConfig,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: any(token in path.lower() for token in ("kandinsky5", "kandinsky-5")),
|
||||
_is_kandinsky5_t2v_lite,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_lite_5s",
|
||||
pipeline_cls_name="Kandinsky5T2VPipeline",
|
||||
)
|
||||
|
||||
# Kandinsky5 Pro T2V
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_t2v_pro,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_pro_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Lite T2V Distilled
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_t2v_lite_distilled,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_lite_distilled_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Pro T2V Distilled
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_t2v_pro_distilled,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_pro_distilled_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Lite I2V
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v_lite,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_lite_5s",
|
||||
pipeline_cls_name="Kandinsky5I2VPipeline",
|
||||
)
|
||||
|
||||
# Kandinsky5 Pro I2V
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v_pro,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_pro_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Pro I2V Distilled
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=["kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers"],
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v_pro_distilled,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_pro_distilled_5s",
|
||||
)
|
||||
|
||||
# Kandinsky5 Lite I2V Distilled (no official hub repo yet; local
|
||||
# conversions get distilled sampling defaults instead of the sft ones the
|
||||
# fallback would apply).
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(),
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v_lite_distilled,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_lite_distilled_5s",
|
||||
pipeline_cls_name="Kandinsky5I2VPipeline",
|
||||
)
|
||||
|
||||
# Kandinsky5 fallbacks — registered AFTER the variant detectors so those
|
||||
# win first-match. Catch checkpoints the variant detectors cannot resolve:
|
||||
# token-less local paths matched via the model_index _class_name
|
||||
# ("kandinsky5t2vpipeline" carries no lite/pro marker), variant combos
|
||||
# without a dedicated entry (e.g. I2V Lite distilled), and t2v+i2v
|
||||
# ambiguous paths resolved by the checkpoint's _class_name.
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5T2VConfig,
|
||||
workload_types=(),
|
||||
model_detectors=[
|
||||
_is_kandinsky5_t2v,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_t2v_lite_5s",
|
||||
pipeline_cls_name="Kandinsky5T2VPipeline",
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Kandinsky5I2VConfig,
|
||||
workload_types=(),
|
||||
model_detectors=[
|
||||
_is_kandinsky5_i2v,
|
||||
],
|
||||
model_family="kandinsky5",
|
||||
default_preset="kandinsky5_i2v_lite_5s",
|
||||
pipeline_cls_name="Kandinsky5I2VPipeline",
|
||||
)
|
||||
|
||||
# LongCat (T2V, I2V, VC use same config; workload varies by path)
|
||||
@@ -550,7 +713,7 @@ def _register_configs() -> None:
|
||||
"FastVideo/Matrix-Game-2.0-Base-Distilled-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Distilled-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Distilled-Diffusers",
|
||||
# Legacy HF paths (kept for backward compat — pre-rename names):
|
||||
# Legacy HF paths (kept for backward compat - pre-rename names):
|
||||
"FastVideo/Matrix-Game-2.0-Base-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-GTA-Diffusers",
|
||||
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
|
||||
@@ -902,7 +1065,10 @@ def get_model_info(
|
||||
assert config_info is not None, "config_info must be resolved"
|
||||
|
||||
if override_pipeline_cls_name:
|
||||
pipeline_name = override_pipeline_cls_name
|
||||
# Explicit override: skip config resolution entirely so checkpoints
|
||||
# without a diffusers model_index.json keep working (and no download
|
||||
# is triggered just to log the replaced name).
|
||||
pipeline_name: str | None = override_pipeline_cls_name
|
||||
logger.info("Using override pipeline class name %s", pipeline_name)
|
||||
else:
|
||||
if os.path.exists(model_path):
|
||||
@@ -911,6 +1077,12 @@ def get_model_info(
|
||||
config = maybe_download_model_index(model_path)
|
||||
|
||||
pipeline_name = config.get("_class_name")
|
||||
if config_info.pipeline_cls_name is not None:
|
||||
# The resolved (path/detector-based) config pins the pipeline class,
|
||||
# e.g. an I2V checkpoint whose `_class_name` would otherwise resolve
|
||||
# to the T2V pipeline.
|
||||
logger.info("Pinning pipeline class name from %s to %s", pipeline_name, config_info.pipeline_cls_name)
|
||||
pipeline_name = config_info.pipeline_cls_name
|
||||
|
||||
if pipeline_name is None:
|
||||
raise ValueError("Model config does not contain a _class_name attribute. "
|
||||
@@ -961,6 +1133,8 @@ def _register_presets() -> None:
|
||||
ALL_PRESETS as HUNYUAN15_PRESETS, )
|
||||
from fastvideo.pipelines.basic.hyworld.presets import (
|
||||
ALL_PRESETS as HYWORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.kandinsky5.presets import (
|
||||
ALL_PRESETS as KANDINSKY5_PRESETS, )
|
||||
from fastvideo.pipelines.basic.lingbotworld.presets import (
|
||||
ALL_PRESETS as LINGBOTWORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.longcat.presets import (
|
||||
@@ -990,6 +1164,7 @@ def _register_presets() -> None:
|
||||
HUNYUAN_PRESETS,
|
||||
HUNYUAN15_PRESETS,
|
||||
HYWORLD_PRESETS,
|
||||
KANDINSKY5_PRESETS,
|
||||
LINGBOTWORLD_PRESETS,
|
||||
LONGCAT_PRESETS,
|
||||
LTX2_PRESETS,
|
||||
|
||||
@@ -3,15 +3,16 @@
|
||||
|
||||
Landed in PR #1225 slice 5 (Attn-QAT 5/12). The resolver centralises the
|
||||
varlen-flash-attn import-fallback logic that several backends
|
||||
(``bsa_attn.py``, ``video_sparse_attn.py``) used to duplicate. The fallback
|
||||
(``bsa_attn.py``, ``video_sparse_attn.py``) used to duplicate. The resolution
|
||||
order is:
|
||||
|
||||
1. ``fastvideo.attention.utils.flash_attn_cute``
|
||||
1. ``fastvideo.attention.utils.flash_attn_cute`` -- only when
|
||||
``FASTVIDEO_FA4=1`` (explicit opt-in), and then it must import or the
|
||||
resolver raises RuntimeError instead of falling through
|
||||
2. ``flash_attn_interface``
|
||||
3. ``flash_attn``
|
||||
|
||||
These tests verify that the resolver picks the highest-priority impl
|
||||
available and falls through cleanly on ``ImportError``. CPU-only, no
|
||||
These tests verify the opt-in gate and the FA3/FA2 fallthrough. CPU-only, no
|
||||
flash-attn install required.
|
||||
"""
|
||||
|
||||
@@ -35,8 +36,36 @@ def _reload_resolver_module():
|
||||
return importlib.import_module("fastvideo.attention.utils.flash_attn_no_pad")
|
||||
|
||||
|
||||
def test_resolver_falls_back_when_cute_unavailable(monkeypatch) -> None:
|
||||
"""When ``flash_attn_cute`` is unimportable, resolver tries the next impl."""
|
||||
def test_resolver_skips_cute_without_opt_in(monkeypatch) -> None:
|
||||
"""Without ``FASTVIDEO_FA4=1`` the resolver must not even attempt the cute
|
||||
import."""
|
||||
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
|
||||
attempted: list[str] = []
|
||||
real_import = builtins.__import__
|
||||
|
||||
def spying_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
attempted.append(name)
|
||||
return real_import(name, globals, locals, fromlist, level)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", spying_import)
|
||||
|
||||
mod = _reload_resolver_module()
|
||||
resolved = mod._resolve_flash_attn_varlen_func()
|
||||
assert resolved is not None
|
||||
assert resolved.__name__ == "flash_attn_varlen_func"
|
||||
assert "fastvideo.attention.utils.flash_attn_cute" not in attempted
|
||||
|
||||
|
||||
def test_resolver_raises_when_opted_in_but_cute_unavailable(monkeypatch) -> None:
|
||||
"""With ``FASTVIDEO_FA4=1`` an unimportable cute build fails loudly instead
|
||||
of silently falling through to FA3/FA2.
|
||||
|
||||
The resolver runs at module import time, so the reload itself must raise.
|
||||
It raises RuntimeError (not ImportError) so importers that treat
|
||||
ImportError as "flash-attn not installed" (``bsa_attn.py``) cannot swallow
|
||||
the opted-in failure.
|
||||
"""
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
real_import = builtins.__import__
|
||||
|
||||
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
@@ -46,21 +75,17 @@ def test_resolver_falls_back_when_cute_unavailable(monkeypatch) -> None:
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", patched_import)
|
||||
|
||||
mod = _reload_resolver_module()
|
||||
resolved = mod._resolve_flash_attn_varlen_func()
|
||||
assert resolved is not None
|
||||
assert resolved.__name__ == "flash_attn_varlen_func"
|
||||
with pytest.raises(RuntimeError, match="cute disabled for test"):
|
||||
_reload_resolver_module()
|
||||
|
||||
|
||||
def test_resolver_returns_flash_attn_when_cute_and_interface_unavailable(monkeypatch) -> None:
|
||||
def test_resolver_returns_flash_attn_when_interface_unavailable(monkeypatch) -> None:
|
||||
"""The terminal fallback is the plain ``flash_attn`` import."""
|
||||
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
|
||||
real_import = builtins.__import__
|
||||
|
||||
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
if name in {
|
||||
"fastvideo.attention.utils.flash_attn_cute",
|
||||
"flash_attn_interface",
|
||||
}:
|
||||
if name == "flash_attn_interface":
|
||||
raise ImportError(f"{name} disabled for test")
|
||||
return real_import(name, globals, locals, fromlist, level)
|
||||
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Guard: every test directory must be collected by some CI lane or be on
|
||||
the explicit allowlist below.
|
||||
|
||||
Three separate incidents on 2026-07-05 found test files that no CI lane
|
||||
ever collects (fastvideo/tests/stages/, tests/local_tests/ additions in
|
||||
PR #1509, and this sweep found seven dark directories in total): the tests
|
||||
pass review, merge, and then silently never run. This test makes going
|
||||
dark an explicit, reviewed decision instead of an accident: adding a new
|
||||
test directory fails CI until it is either wired into a lane or
|
||||
allowlisted here with a reason.
|
||||
|
||||
Pure text analysis — no fastvideo imports, no GPU, no torch.
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
TESTS_ROOT = REPO_ROOT / "fastvideo" / "tests"
|
||||
|
||||
# Files whose text constitutes "a CI lane references this directory".
|
||||
CI_SOURCES = [
|
||||
TESTS_ROOT / "modal" / "pr_test.py",
|
||||
TESTS_ROOT / "modal" / "ssim_test.py",
|
||||
*sorted((REPO_ROOT / ".buildkite").rglob("*.yml")),
|
||||
*sorted((REPO_ROOT / ".buildkite").rglob("*.sh")),
|
||||
]
|
||||
|
||||
# Directories that intentionally have no CI lane today. Every entry needs a
|
||||
# reason; remove the entry when the directory gets wired into a lane.
|
||||
# State as found on 2026-07-05 — these SHOULD shrink over time, not grow.
|
||||
ALLOWLIST = {
|
||||
"attention": "no lane yet — GPU attention-backend tests, run manually",
|
||||
"audio": "no lane yet — audio encoder tests, run manually",
|
||||
"distributed": "no lane yet — multi-GPU torchrun tests, run manually",
|
||||
"hooks": "no lane yet — run manually",
|
||||
"layers": "no lane yet — torchrun FSDP dispatch tests, run manually",
|
||||
"nightly": "by design: nightly cadence, not per-PR",
|
||||
"ops": "no lane yet — GPU op tests, run manually",
|
||||
"modal": "CI infrastructure itself, not a test suite",
|
||||
}
|
||||
|
||||
|
||||
def _dirs_with_tests() -> list[str]:
|
||||
dirs = []
|
||||
for child in sorted(TESTS_ROOT.iterdir()):
|
||||
if child.is_dir() and any(child.rglob("test_*.py")):
|
||||
dirs.append(child.name)
|
||||
return dirs
|
||||
|
||||
|
||||
def _ci_text() -> str:
|
||||
return "\n".join(
|
||||
src.read_text(errors="replace") for src in CI_SOURCES if src.exists())
|
||||
|
||||
|
||||
def test_every_test_directory_is_collected_or_allowlisted():
|
||||
ci_text = _ci_text()
|
||||
dark = [
|
||||
name for name in _dirs_with_tests()
|
||||
if f"tests/{name}" not in ci_text and name not in ALLOWLIST
|
||||
]
|
||||
assert not dark, (
|
||||
f"Test directories not referenced by any CI lane and not "
|
||||
f"allowlisted: {dark}. Wire them into a lane in "
|
||||
f"fastvideo/tests/modal/pr_test.py (or a Buildkite step), or add an "
|
||||
f"allowlist entry with a reason in {__file__}.")
|
||||
|
||||
|
||||
def test_local_tests_stays_out_of_ci():
|
||||
# tests/local_tests/ (repo root) is developer-local by design (author
|
||||
# decision, 2026-07-05): parity scaffolds and machine-specific checks
|
||||
# that must never gate CI. Fail if any CI source starts collecting it.
|
||||
assert "tests/local_tests" not in _ci_text(), (
|
||||
"tests/local_tests/ is local-only by design; remove the CI "
|
||||
"reference or move the tests into a fastvideo/tests/ lane.")
|
||||
|
||||
|
||||
def test_allowlist_entries_are_still_real_directories():
|
||||
# A stale allowlist hides regressions; entries must track reality.
|
||||
missing = [
|
||||
name for name in ALLOWLIST
|
||||
if name != "modal" and not (TESTS_ROOT / name).is_dir()
|
||||
]
|
||||
assert not missing, (
|
||||
f"Allowlisted directories no longer exist — remove them: {missing}")
|
||||
@@ -0,0 +1,203 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import init_device_mesh
|
||||
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard
|
||||
from torch.distributed.tensor import DTensor
|
||||
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
|
||||
WORLD_SIZE = 2
|
||||
HIDDEN_SIZE = 8
|
||||
SEED = 1379
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
|
||||
|
||||
def _run_torchrun(script_path: Path, mode: str, output_path: Path) -> None:
|
||||
# --standalone binds the rendezvous port atomically, avoiding the
|
||||
# free-port-probe race a hand-picked --master_port would have.
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--standalone",
|
||||
"--nproc_per_node",
|
||||
str(WORLD_SIZE),
|
||||
str(script_path),
|
||||
"--rmsnorm-fsdp-worker",
|
||||
"--mode",
|
||||
mode,
|
||||
"--output",
|
||||
str(output_path),
|
||||
]
|
||||
env = os.environ.copy()
|
||||
env.setdefault("TORCHDYNAMO_DISABLE", "1")
|
||||
try:
|
||||
process = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=env,
|
||||
timeout=120,
|
||||
)
|
||||
except subprocess.TimeoutExpired as error:
|
||||
raise RuntimeError(
|
||||
f"{mode} worker timed out after 120 seconds\n"
|
||||
f"STDOUT:\n{error.stdout}\n"
|
||||
f"STDERR:\n{error.stderr}"
|
||||
) from error
|
||||
if process.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"{mode} worker failed with code {process.returncode}\n"
|
||||
f"STDOUT:\n{process.stdout}\n"
|
||||
f"STDERR:\n{process.stderr}"
|
||||
)
|
||||
|
||||
|
||||
def _summarize_tensor(tensor: torch.Tensor | Any) -> dict[str, Any]:
|
||||
return {
|
||||
"type": type(tensor).__name__,
|
||||
"is_dtensor": isinstance(tensor, DTensor),
|
||||
"shape": list(tensor.shape) if hasattr(tensor, "shape") else None,
|
||||
"device": str(tensor.device) if hasattr(tensor, "device") else None,
|
||||
"dtype": str(tensor.dtype) if hasattr(tensor, "dtype") else None,
|
||||
}
|
||||
|
||||
|
||||
def _run_worker(mode: str, output_path: Path) -> None:
|
||||
if mode not in {
|
||||
"module_no_offload",
|
||||
"direct_no_offload",
|
||||
"module_cpu_offload",
|
||||
"direct_cpu_offload",
|
||||
}:
|
||||
raise ValueError(f"Unsupported mode: {mode}")
|
||||
|
||||
dist.init_process_group("nccl")
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.cuda.set_device(device)
|
||||
torch.manual_seed(SEED + rank)
|
||||
|
||||
try:
|
||||
mesh = init_device_mesh("cuda", (world_size,))
|
||||
norm = RMSNorm(HIDDEN_SIZE, eps=1e-6, has_weight=True).to(device)
|
||||
with torch.no_grad():
|
||||
norm.weight.fill_(1.0)
|
||||
|
||||
fsdp_kwargs: dict[str, Any] = {"mesh": mesh}
|
||||
if mode.endswith("cpu_offload"):
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(pin_memory=False)
|
||||
# fully_shard is applied to the bare RMSNorm to make the hook bypass
|
||||
# observable. Production sharding (fsdp_load.shard_model) only wraps
|
||||
# whole transformer blocks, whose pre-forward all-gather localizes norm
|
||||
# weights before the qk-norm call sites run, so this pins the dispatch
|
||||
# invariant rather than reproducing a production topology.
|
||||
fully_shard(norm, **fsdp_kwargs)
|
||||
|
||||
x = torch.randn(2, 3, HIDDEN_SIZE, device=device, dtype=torch.bfloat16)
|
||||
call_kind = "direct" if mode.startswith("direct") else "module"
|
||||
|
||||
try:
|
||||
if call_kind == "direct":
|
||||
output = norm.forward_native(x)
|
||||
else:
|
||||
output = norm(x)
|
||||
torch.cuda.synchronize(device)
|
||||
result = {
|
||||
"rank": rank,
|
||||
"ok": True,
|
||||
"mode": mode,
|
||||
"weight": _summarize_tensor(norm.weight),
|
||||
"output": _summarize_tensor(output),
|
||||
}
|
||||
except Exception as exc:
|
||||
result = {
|
||||
"rank": rank,
|
||||
"ok": False,
|
||||
"mode": mode,
|
||||
"error_type": type(exc).__name__,
|
||||
"error": str(exc),
|
||||
"weight": _summarize_tensor(norm.weight),
|
||||
}
|
||||
|
||||
gathered = [None for _ in range(world_size)] if rank == 0 else None
|
||||
dist.gather_object(result, object_gather_list=gathered, dst=0)
|
||||
if rank == 0:
|
||||
output_path.write_text(json.dumps(gathered, indent=2), encoding="utf-8")
|
||||
dist.barrier()
|
||||
finally:
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mode", "expect_ok"),
|
||||
[
|
||||
("module_no_offload", True),
|
||||
("direct_no_offload", False),
|
||||
("module_cpu_offload", True),
|
||||
("direct_cpu_offload", False),
|
||||
],
|
||||
)
|
||||
def test_rmsnorm_forward_native_bypasses_fsdp_hooks(mode: str, expect_ok: bool, tmp_path: Path) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("This test requires CUDA.")
|
||||
if torch.cuda.device_count() < WORLD_SIZE:
|
||||
pytest.skip(f"This test requires at least {WORLD_SIZE} CUDA devices.")
|
||||
|
||||
output_path = tmp_path / f"{mode}.json"
|
||||
_run_torchrun(Path(__file__).resolve(), mode, output_path)
|
||||
results = json.loads(output_path.read_text(encoding="utf-8"))
|
||||
print(f"\n{mode} results:\n{json.dumps(results, indent=2)}")
|
||||
|
||||
if expect_ok:
|
||||
failures = [result for result in results if not result["ok"]]
|
||||
assert not failures, json.dumps(results, indent=2)
|
||||
return
|
||||
|
||||
successes = [result for result in results if result["ok"]]
|
||||
assert not successes, json.dumps(results, indent=2)
|
||||
error_text = "\n".join(result.get("error", "") for result in results)
|
||||
# Pin the specific bypassed-hook failure: "got mixed torch.Tensor and
|
||||
# DTensor" ("Tensor" alone is a substring of "DTensor", so it adds nothing).
|
||||
assert "mixed" in error_text and "DTensor" in error_text, json.dumps(results, indent=2)
|
||||
|
||||
|
||||
def test_no_direct_forward_native_calls_in_models() -> None:
|
||||
"""Direct .forward_native(...) calls bypass nn.Module.__call__ and FSDP
|
||||
hooks (issue #1379); model code must use module dispatch instead."""
|
||||
models_dir = REPO_ROOT / "fastvideo" / "models"
|
||||
offenders = [
|
||||
str(path.relative_to(REPO_ROOT))
|
||||
for path in sorted(models_dir.rglob("*.py"))
|
||||
if ".forward_native(" in path.read_text(encoding="utf-8")
|
||||
]
|
||||
assert not offenders, f"Replace .forward_native(...) with module dispatch in: {offenders}"
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--rmsnorm-fsdp-worker", action="store_true")
|
||||
parser.add_argument("--mode", type=str, default=None)
|
||||
parser.add_argument("--output", type=str, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = _parse_args()
|
||||
if not args.rmsnorm_fsdp_worker:
|
||||
raise SystemExit("This module is intended to be run by pytest.")
|
||||
if args.mode is None or args.output is None:
|
||||
raise SystemExit("--mode and --output are required in worker mode.")
|
||||
_run_worker(mode=args.mode, output_path=Path(args.output))
|
||||
@@ -32,7 +32,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -40,6 +41,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App("fastvideo-gpu-job")
|
||||
|
||||
REPO_DIR = "/FastVideo"
|
||||
@@ -72,12 +76,7 @@ local_secrets = modal.Secret.from_dict({
|
||||
# Mutable tags inherit the registry image's baked backend, including custom
|
||||
# FASTVIDEO_MODAL_IMAGE overrides. Explicit CUDA tags also work with older
|
||||
# images that predate the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in IMAGE_TAG.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in IMAGE_TAG.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(IMAGE_TAG)
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(IMAGE_REF, add_python="3.12")
|
||||
@@ -98,6 +97,9 @@ image = (
|
||||
"TOKENIZERS_PARALLELISM": "false",
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); keep CI parity with the seeded
|
||||
# references. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
})
|
||||
)
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ through the ``fastvideo`` package, so they keep working on thin CI hosts that
|
||||
have ``modal`` but not torch.
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import urllib.request
|
||||
|
||||
_REGISTRY = "ghcr.io"
|
||||
@@ -55,3 +56,23 @@ def resolve_image_ref(image_ref: str) -> str:
|
||||
print(f"WARNING: could not resolve {image_ref} to a digest ({error}); "
|
||||
"Modal may reuse a stale cached image for this tag.")
|
||||
return image_ref
|
||||
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
"""UV_TORCH_BACKEND for a launcher image tag.
|
||||
|
||||
A caller-set UV_TORCH_BACKEND always wins. Otherwise sniff the CUDA
|
||||
version from an explicit image tag (cuda13 -> cu130, cuda12.6 -> cu126)
|
||||
so uv resolves torch against the image's toolkit. Mutable tags (e.g.
|
||||
py3.12-latest) return None and inherit the registry image's baked
|
||||
backend, which keeps a latest-tag CUDA transition safe.
|
||||
"""
|
||||
override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if override:
|
||||
return override
|
||||
tag = image_tag.lower()
|
||||
if "cuda13" in tag:
|
||||
return "cu130"
|
||||
if "cuda12.6" in tag:
|
||||
return "cu126"
|
||||
return None
|
||||
|
||||
@@ -5,7 +5,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -13,6 +14,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App()
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
@@ -24,12 +28,7 @@ print(f"Using image: {image_ref}")
|
||||
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
|
||||
# transition safe. Explicit CUDA tags also work with older images that predate
|
||||
# the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
|
||||
|
||||
image = (modal.Image.from_registry(
|
||||
image_ref, add_python="3.12"
|
||||
@@ -61,6 +60,10 @@ image = (modal.Image.from_registry(
|
||||
**({
|
||||
"UV_TORCH_BACKEND": uv_torch_backend_override
|
||||
} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); CI lanes keep it enabled to match the
|
||||
# SSIM/perf baselines. Caller override wins.
|
||||
"FASTVIDEO_FA4":
|
||||
os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
"HF_REPO_ID":
|
||||
"FastVideo/performance-tracking",
|
||||
}))
|
||||
@@ -103,12 +106,30 @@ def run_test_command(test_command: str,
|
||||
if pr_number:
|
||||
print(f"PR number: {pr_number}")
|
||||
|
||||
# For PRs (including forks), use GitHub's PR refs to get the correct commit
|
||||
# The blob-less clone (--filter=blob:none) defers all file-content
|
||||
# downloads to the checkout, so BOTH paths below perform a large lazy blob
|
||||
# fetch from GitHub. Retry transient GitHub/HTTP2 disconnects on each;
|
||||
# otherwise Modal shards can fail before pytest starts.
|
||||
def with_retries(inner_command: str) -> str:
|
||||
return f"""
|
||||
for attempt in 1 2 3; do
|
||||
{inner_command} &&
|
||||
break
|
||||
|
||||
status=$?
|
||||
if [ "$attempt" -eq 3 ]; then
|
||||
exit "$status"
|
||||
fi
|
||||
sleep $((attempt * 5))
|
||||
done"""
|
||||
|
||||
# For PRs (including forks), use GitHub's PR refs to get the correct commit.
|
||||
if pr_number and pr_number != "false":
|
||||
checkout_command = f"git fetch --prune origin refs/pull/{pr_number}/head && git checkout FETCH_HEAD"
|
||||
checkout_command = with_retries(f"git fetch --prune --no-tags --depth=1 origin refs/pull/{pr_number}/head &&"
|
||||
"\n git checkout --detach FETCH_HEAD")
|
||||
print(f"Using PR ref for checkout: {checkout_command}")
|
||||
else:
|
||||
checkout_command = f"git checkout {git_commit}"
|
||||
checkout_command = with_retries(f"git checkout {git_commit}")
|
||||
print(f"Using direct commit checkout: {checkout_command}")
|
||||
|
||||
build_kernel_command = """
|
||||
@@ -122,7 +143,7 @@ def run_test_command(test_command: str,
|
||||
command = f"""
|
||||
source $HOME/.local/bin/env &&
|
||||
source /opt/venv/bin/activate &&
|
||||
git clone {git_repo} /FastVideo &&
|
||||
git clone --filter=blob:none --no-checkout {git_repo} /FastVideo &&
|
||||
cd /FastVideo &&
|
||||
{checkout_command} &&
|
||||
git submodule update --init --recursive &&
|
||||
@@ -273,7 +294,7 @@ def run_self_forcing_tests():
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_unit_test():
|
||||
run_test(
|
||||
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
|
||||
"pytest ./fastvideo/tests/api/ ./fastvideo/tests/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -13,7 +13,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -21,6 +22,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App()
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
@@ -32,12 +36,7 @@ print(f"Using image: {image_ref}")
|
||||
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
|
||||
# transition safe. Explicit CUDA tags also work with older images that predate
|
||||
# the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(image_ref, add_python="3.12")
|
||||
@@ -64,6 +63,9 @@ image = (
|
||||
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"IMAGE_VERSION": image_version,
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); the SSIM references were seeded
|
||||
# with FA4 inference, so keep it enabled in CI. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
@@ -8,8 +8,8 @@ This script:
|
||||
baseline-eligible successful records (filtered by gpu_type),
|
||||
4) writes normalized records back to the HF dataset repo according to
|
||||
PERF_UPLOAD_POLICY,
|
||||
5) exits non-zero if any metric regresses by more than PERF_MAX_REGRESSION
|
||||
(default 5%).
|
||||
5) exits non-zero if any gated metric exceeds both its percent and absolute
|
||||
regression floors.
|
||||
"""
|
||||
|
||||
import glob
|
||||
@@ -21,21 +21,36 @@ from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
try:
|
||||
from .hf_store import (
|
||||
from fastvideo.performance.hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
from fastvideo.performance.metric_policy import (
|
||||
MetricPolicy,
|
||||
regression_delta,
|
||||
resolve_metric_policies,
|
||||
serialize_metric_thresholds,
|
||||
)
|
||||
except ImportError:
|
||||
from hf_store import (
|
||||
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
|
||||
if repo_root not in sys.path:
|
||||
sys.path.insert(0, repo_root)
|
||||
from fastvideo.performance.hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
from fastvideo.performance.metric_policy import (
|
||||
MetricPolicy,
|
||||
regression_delta,
|
||||
resolve_metric_policies,
|
||||
serialize_metric_thresholds,
|
||||
)
|
||||
|
||||
RESULTS_DIR = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)),
|
||||
@@ -46,25 +61,9 @@ TRACKING_ROOT = os.environ.get(
|
||||
"/tmp/perf-tracking",
|
||||
)
|
||||
PERF_REPORTS_DIR = os.environ.get("PERF_REPORTS_DIR", "/root/data/perf_reports")
|
||||
MAX_REGRESSION = float(os.environ.get("PERF_MAX_REGRESSION", "0.05"))
|
||||
UPLOAD_POLICY = os.environ.get("PERF_UPLOAD_POLICY", "never").strip().lower()
|
||||
VALID_UPLOAD_POLICIES = {"never", "pass", "always"}
|
||||
VALID_RUN_SOURCES = {"pr", "local", "scheduled_main", "unknown"}
|
||||
METRICS = (
|
||||
("latency", "Latency", 3),
|
||||
("throughput", "Throughput", 3),
|
||||
("memory", "Memory", 1),
|
||||
("text_encoder_time_s", "Text Enc", 3),
|
||||
("dit_time_s", "DiT", 3),
|
||||
("vae_decode_time_s", "VAE Decode", 3),
|
||||
)
|
||||
LOWER_IS_BETTER_METRICS = {
|
||||
"latency",
|
||||
"memory",
|
||||
"text_encoder_time_s",
|
||||
"dit_time_s",
|
||||
"vae_decode_time_s",
|
||||
}
|
||||
|
||||
|
||||
def _should_persist_tracking() -> bool:
|
||||
@@ -168,6 +167,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
text_encoder_time = safe_float(result.get("text_encoder_time_s"))
|
||||
dit_time = safe_float(result.get("dit_time_s"))
|
||||
vae_decode_time = safe_float(result.get("vae_decode_time_s"))
|
||||
metric_policies = resolve_metric_policies(result.get("regression_thresholds"))
|
||||
|
||||
return {
|
||||
"model_id": model_id,
|
||||
@@ -180,6 +180,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
"text_encoder_time_s": text_encoder_time,
|
||||
"dit_time_s": dit_time,
|
||||
"vae_decode_time_s": vae_decode_time,
|
||||
"regression_thresholds": serialize_metric_thresholds(metric_policies),
|
||||
"success": True,
|
||||
**_record_metadata(_detect_run_source(), result),
|
||||
}
|
||||
@@ -229,55 +230,41 @@ def _baseline_metric(records: list[dict[str, Any]], key: str) -> float | None:
|
||||
return statistics.median(values)
|
||||
|
||||
|
||||
def _metric_policy_summary(policy: MetricPolicy) -> str:
|
||||
gated = "gated" if policy.gated else "info"
|
||||
return (
|
||||
f"{gated}, >{policy.threshold_percent * 100:.1f}% "
|
||||
f"and >{policy.threshold_absolute:.{policy.precision}f}"
|
||||
)
|
||||
|
||||
|
||||
def _check_regressions(
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
) -> list[str]:
|
||||
failures: list[str] = []
|
||||
|
||||
for metric, _label, _precision in METRICS:
|
||||
if metric not in LOWER_IS_BETTER_METRICS:
|
||||
for policy in metric_policies:
|
||||
baseline = _baseline_metric(baseline_records, policy.key)
|
||||
curr = safe_float(current.get(policy.key))
|
||||
if baseline is None or curr is None:
|
||||
continue
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
curr = safe_float(current.get(metric))
|
||||
if baseline is None or curr is None or baseline <= 0:
|
||||
delta = regression_delta(policy, curr, baseline)
|
||||
if delta is None or not delta.regressed:
|
||||
continue
|
||||
regression = (curr - baseline) / baseline
|
||||
if regression > max_regression:
|
||||
failures.append(f"{current['model_id']} {metric} regressed by "
|
||||
f"{regression * 100:.1f}% "
|
||||
f"(current={curr:.3f}, baseline_median={baseline:.3f})")
|
||||
|
||||
baseline_tp = _baseline_metric(baseline_records, "throughput")
|
||||
curr_tp = safe_float(current.get("throughput"))
|
||||
if baseline_tp is not None and curr_tp is not None and baseline_tp > 0:
|
||||
regression = (baseline_tp - curr_tp) / baseline_tp
|
||||
if regression > max_regression:
|
||||
failures.append(f"{current['model_id']} throughput regressed by "
|
||||
f"{regression * 100:.1f}% "
|
||||
f"(current={curr_tp:.3f}, baseline_median={baseline_tp:.3f})")
|
||||
failures.append(
|
||||
f"{current['model_id']} {policy.key} regressed by "
|
||||
f"{delta.percent * 100:.1f}% and "
|
||||
f"{delta.absolute:.{policy.precision}f} "
|
||||
f"(current={curr:.{policy.precision}f}, "
|
||||
f"baseline_median={baseline:.{policy.precision}f}, "
|
||||
f"threshold={_metric_policy_summary(policy)})"
|
||||
)
|
||||
|
||||
return failures
|
||||
|
||||
|
||||
def _metric_delta_percent(
|
||||
metric: str,
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
) -> float | None:
|
||||
curr = safe_float(current.get(metric))
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
if curr is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
|
||||
if metric in LOWER_IS_BETTER_METRICS:
|
||||
return (curr - baseline) / baseline * 100.0
|
||||
if metric == "throughput":
|
||||
return (baseline - curr) / baseline * 100.0
|
||||
return None
|
||||
|
||||
|
||||
def _compact_value(value: float | None, precision: int = 3) -> str:
|
||||
if value is None:
|
||||
return "n/a"
|
||||
@@ -287,23 +274,42 @@ def _compact_value(value: float | None, precision: int = 3) -> str:
|
||||
def _build_summary_row(
|
||||
record: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
has_failed: bool,
|
||||
) -> dict[str, Any]:
|
||||
"""Format a single benchmark result as a row for the Markdown table."""
|
||||
|
||||
metric_values: dict[str, dict[str, float | None]] = {}
|
||||
metric_values: dict[str, dict[str, Any]] = {}
|
||||
regressions: list[float] = []
|
||||
for metric, _label, _precision in METRICS:
|
||||
curr = safe_float(record.get(metric))
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
regression = _metric_delta_percent(metric, record, baseline_records)
|
||||
metric_values[metric] = {
|
||||
failing_metrics: list[str] = []
|
||||
threshold_exceeded_metrics: list[str] = []
|
||||
for policy in metric_policies:
|
||||
curr = safe_float(record.get(policy.key))
|
||||
baseline = _baseline_metric(baseline_records, policy.key)
|
||||
delta = (
|
||||
regression_delta(policy, curr, baseline)
|
||||
if curr is not None and baseline is not None
|
||||
else None
|
||||
)
|
||||
regression = None if delta is None else delta.percent * 100.0
|
||||
absolute_delta = None if delta is None else delta.absolute
|
||||
metric_values[policy.key] = {
|
||||
"curr": curr,
|
||||
"base": baseline,
|
||||
"regression_pct": regression,
|
||||
"absolute_delta": absolute_delta,
|
||||
"threshold_percent": policy.threshold_percent * 100.0,
|
||||
"threshold_absolute": policy.threshold_absolute,
|
||||
"gated": policy.gated,
|
||||
"threshold_exceeded": False if delta is None else delta.threshold_exceeded,
|
||||
"regressed": False if delta is None else delta.regressed,
|
||||
}
|
||||
if regression is not None:
|
||||
regressions.append(regression)
|
||||
if delta is not None and delta.threshold_exceeded:
|
||||
threshold_exceeded_metrics.append(policy.key)
|
||||
if delta is not None and delta.regressed:
|
||||
failing_metrics.append(policy.key)
|
||||
|
||||
worst_regression_pct = max(regressions) if regressions else None
|
||||
|
||||
@@ -313,40 +319,50 @@ def _build_summary_row(
|
||||
"baseline_n": len(baseline_records),
|
||||
"metrics": metric_values,
|
||||
"worst_regression_pct": worst_regression_pct,
|
||||
"threshold_exceeded_metrics": threshold_exceeded_metrics,
|
||||
"failing_metrics": failing_metrics,
|
||||
"failed": has_failed,
|
||||
}
|
||||
|
||||
|
||||
def _build_markdown_summary(
|
||||
summary_rows: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
) -> str:
|
||||
lines = [
|
||||
"## Performance Baseline Comparison",
|
||||
"",
|
||||
f"Threshold: regressions greater than {max_regression * 100:.1f}% fail",
|
||||
"Threshold: gated metrics fail only when both percent and absolute "
|
||||
"regression floors are exceeded.",
|
||||
"",
|
||||
("| Model | GPU | Baseline N | Latency (curr/base) | "
|
||||
"Throughput (curr/base) | Memory (curr/base) | "
|
||||
"Text Enc (curr/base) | DiT (curr/base) | "
|
||||
"VAE Decode (curr/base) | Worst Regression | Status |"),
|
||||
"|---|---|---:|---|---|---|---|---|---|---:|---|",
|
||||
"VAE Decode (curr/base) | Worst Regression | Exceeded Metrics | "
|
||||
"Failing Metrics | Status |"),
|
||||
"|---|---|---:|---|---|---|---|---|---|---:|---|---|---|",
|
||||
]
|
||||
|
||||
for row in summary_rows:
|
||||
metric_cells = []
|
||||
for metric, _label, precision in METRICS:
|
||||
values = row["metrics"][metric]
|
||||
metric_cells.append(f"{_compact_value(values['curr'], precision)} / "
|
||||
f"{_compact_value(values['base'], precision)}")
|
||||
for policy in metric_policies:
|
||||
values = row["metrics"][policy.key]
|
||||
metric_cells.append(f"{_compact_value(values['curr'], policy.precision)} / "
|
||||
f"{_compact_value(values['base'], policy.precision)}")
|
||||
|
||||
worst_reg = ("n/a" if row["worst_regression_pct"] is None else f"{row['worst_regression_pct']:.1f}%")
|
||||
exceeded_metrics = (
|
||||
", ".join(row["threshold_exceeded_metrics"])
|
||||
if row["threshold_exceeded_metrics"]
|
||||
else "none"
|
||||
)
|
||||
failing_metrics = ", ".join(row["failing_metrics"]) if row["failing_metrics"] else "none"
|
||||
status = "FAIL" if row["failed"] else "PASS"
|
||||
|
||||
lines.append(f"| {row['model_id']} | {row['gpu_type']} | "
|
||||
f"{row['baseline_n']} | "
|
||||
f"{' | '.join(metric_cells)} | "
|
||||
f"{worst_reg} | {status} |")
|
||||
f"{worst_reg} | {exceeded_metrics} | {failing_metrics} | {status} |")
|
||||
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
@@ -400,6 +416,7 @@ def main() -> int:
|
||||
|
||||
for raw in current_results:
|
||||
record = _normalize_record(raw)
|
||||
metric_policies = resolve_metric_policies(record.get("regression_thresholds"))
|
||||
|
||||
baseline_records = load_records_for_model(
|
||||
TRACKING_ROOT,
|
||||
@@ -416,7 +433,7 @@ def main() -> int:
|
||||
failures: list[str] = []
|
||||
record["success"] = True
|
||||
else:
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
failures = _check_regressions(record, baseline_records, metric_policies)
|
||||
if static_threshold_failed:
|
||||
failures.append(f"{record['model_id']} fixed-threshold phase failed "
|
||||
f"(PERF_PYTEST_RC={os.environ.get('PERF_PYTEST_RC')})")
|
||||
@@ -434,10 +451,10 @@ def main() -> int:
|
||||
print("Tracking upload skipped for "
|
||||
f"{record['model_id']} ({record['run_source']}, success={record['success']})")
|
||||
|
||||
summary_rows.append(_build_summary_row(record, baseline_records, bool(failures)))
|
||||
summary_rows.append(_build_summary_row(record, baseline_records, metric_policies, bool(failures)))
|
||||
|
||||
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
|
||||
markdown = _build_markdown_summary(summary_rows, MAX_REGRESSION)
|
||||
markdown = _build_markdown_summary(summary_rows, resolve_metric_policies(None))
|
||||
_emit_markdown_summary(markdown, commit_sha)
|
||||
|
||||
if all_failures:
|
||||
|
||||
@@ -1,12 +1,19 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import sys
|
||||
from html import escape
|
||||
from datetime import datetime
|
||||
|
||||
import plotly.express as px
|
||||
import pandas as pd
|
||||
|
||||
from hf_store import sync_from_hf, load_as_dataframe
|
||||
try:
|
||||
from fastvideo.performance.hf_store import load_as_dataframe, sync_from_hf
|
||||
except ImportError:
|
||||
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
|
||||
if repo_root not in sys.path:
|
||||
sys.path.insert(0, repo_root)
|
||||
from fastvideo.performance.hf_store import load_as_dataframe, sync_from_hf
|
||||
|
||||
TRACKING_ROOT = os.environ.get(
|
||||
"PERFORMANCE_TRACKING_ROOT",
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.tests.performance.test_inference_performance import (
|
||||
_benchmark_display_id,
|
||||
_config_identity_metadata,
|
||||
_is_v2_config,
|
||||
_validate_benchmark_config,
|
||||
)
|
||||
|
||||
|
||||
def _v2_config():
|
||||
return {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
}
|
||||
|
||||
|
||||
def test_v1_benchmark_config_without_schema_version_validates():
|
||||
cfg = {
|
||||
"benchmark_id": "legacy-benchmark",
|
||||
}
|
||||
|
||||
_validate_benchmark_config(cfg, "legacy.json")
|
||||
|
||||
assert _is_v2_config(cfg) is False
|
||||
assert _config_identity_metadata(cfg) == {}
|
||||
assert _benchmark_display_id(cfg) == "legacy-benchmark"
|
||||
|
||||
|
||||
def test_v2_benchmark_config_identity_validates_and_is_preserved():
|
||||
cfg = _v2_config()
|
||||
cfg["quality_metadata"] = {"some": "data"}
|
||||
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
assert _is_v2_config(cfg) is True
|
||||
assert _config_identity_metadata(cfg) == {
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"quality_metadata": {"some": "data"},
|
||||
}
|
||||
|
||||
|
||||
def test_v2_benchmark_config_missing_identity_fields_fails_clearly():
|
||||
cfg = _v2_config()
|
||||
del cfg["variant_id"]
|
||||
del cfg["benchmark_version"]
|
||||
|
||||
expected = "wan.json: missing required v2 identity fields: variant_id, benchmark_version"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "value"),
|
||||
[
|
||||
("workload_id", {}),
|
||||
("workload_id", ""),
|
||||
("workload_id", " "),
|
||||
("variant_id", []),
|
||||
("variant_id", ""),
|
||||
("variant_id", " "),
|
||||
],
|
||||
)
|
||||
def test_v2_benchmark_config_rejects_invalid_string_identity_values(field, value):
|
||||
cfg = _v2_config()
|
||||
cfg[field] = value
|
||||
|
||||
expected = f"wan.json: v2 identity field {field!r} must be a non-empty string"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [None, "1", 1.5, True])
|
||||
def test_v2_benchmark_config_rejects_invalid_benchmark_version_values(value):
|
||||
cfg = _v2_config()
|
||||
cfg["benchmark_version"] = value
|
||||
|
||||
expected = "wan.json: v2 identity field 'benchmark_version' must be an integer"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
def test_partial_v2_identity_requires_schema_version():
|
||||
cfg = {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
}
|
||||
|
||||
expected = "wan.json: v2 benchmark identity fields require config_schema_version=2"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
def test_optional_v2_metadata_fields_must_be_objects():
|
||||
cfg = {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"quality_metadata": ["not", "an", "object"],
|
||||
}
|
||||
|
||||
expected = "wan.json: optional v2 metadata field 'quality_metadata' must be an object"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.tests.performance import compare_baseline
|
||||
from fastvideo.performance.metric_policy import resolve_metric_policies
|
||||
|
||||
|
||||
def _raw_result():
|
||||
@@ -72,9 +73,157 @@ def test_normalized_record_includes_source_metadata(monkeypatch):
|
||||
assert record["job_id"] == "job-1"
|
||||
|
||||
|
||||
def test_normalized_record_includes_effective_regression_thresholds(monkeypatch):
|
||||
monkeypatch.setenv("PERF_RUN_SOURCE", "pr")
|
||||
raw = _raw_result()
|
||||
raw["regression_thresholds"] = {
|
||||
"latency": {
|
||||
"threshold_percent": 0.09,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
},
|
||||
"throughput": {
|
||||
"gated": False,
|
||||
},
|
||||
}
|
||||
|
||||
record = compare_baseline.normalize_performance_result(raw)
|
||||
|
||||
assert record["regression_thresholds"]["latency"] == {
|
||||
"threshold_percent": 0.09,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
}
|
||||
assert record["regression_thresholds"]["throughput"]["gated"] is False
|
||||
|
||||
|
||||
def test_invalid_regression_threshold_container_uses_defaults():
|
||||
policies = resolve_metric_policies(["not", "a", "mapping"])
|
||||
|
||||
latency = next(policy for policy in policies if policy.key == "latency")
|
||||
assert latency.threshold_percent == 0.08
|
||||
assert latency.threshold_absolute == 0.5
|
||||
assert latency.gated is True
|
||||
|
||||
|
||||
def test_boolean_regression_threshold_values_are_ignored():
|
||||
policies = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": True,
|
||||
"threshold_absolute": False,
|
||||
"gated": "false",
|
||||
}
|
||||
})
|
||||
|
||||
latency = next(policy for policy in policies if policy.key == "latency")
|
||||
assert latency.threshold_percent == 0.08
|
||||
assert latency.threshold_absolute == 0.5
|
||||
assert latency.gated is False
|
||||
|
||||
|
||||
def test_baseline_eligibility_only_for_successful_scheduled_main():
|
||||
assert compare_baseline._is_baseline_eligible("scheduled_main", True) is True
|
||||
assert compare_baseline._is_baseline_eligible("scheduled_main", False) is False
|
||||
assert compare_baseline._is_baseline_eligible("pr", True) is False
|
||||
assert compare_baseline._is_baseline_eligible("local", True) is False
|
||||
|
||||
|
||||
def test_latency_regression_requires_percent_and_absolute_floors():
|
||||
baseline = [{"latency": 10.0}]
|
||||
current = {"model_id": "wan", "latency": 10.6}
|
||||
|
||||
percent_only = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.75,
|
||||
}
|
||||
})
|
||||
absolute_only = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.10,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
both = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
|
||||
assert compare_baseline._check_regressions(current, baseline, percent_only) == []
|
||||
assert compare_baseline._check_regressions(current, baseline, absolute_only) == []
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, both)
|
||||
assert len(failures) == 1
|
||||
assert "latency regressed by 6.0% and 0.600" in failures[0]
|
||||
|
||||
|
||||
def test_throughput_regression_uses_higher_is_better_direction():
|
||||
baseline = [{"throughput": 10.0}]
|
||||
current = {"model_id": "wan", "throughput": 9.0}
|
||||
policies = resolve_metric_policies({
|
||||
"throughput": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "throughput regressed by 10.0% and 1.000" in failures[0]
|
||||
|
||||
|
||||
def test_memory_regression_uses_metric_specific_absolute_floor():
|
||||
baseline = [{"memory": 10000.0}]
|
||||
current = {"model_id": "wan", "memory": 10600.0}
|
||||
policies = resolve_metric_policies({
|
||||
"memory": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 256.0,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "memory regressed by 6.0% and 600.0" in failures[0]
|
||||
|
||||
|
||||
def test_component_metric_can_gate_independently():
|
||||
baseline = [{"dit_time_s": 8.0}]
|
||||
current = {"model_id": "wan", "dit_time_s": 8.6}
|
||||
policies = resolve_metric_policies({
|
||||
"dit_time_s": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.25,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "dit_time_s regressed by 7.5% and 0.600" in failures[0]
|
||||
|
||||
|
||||
def test_informational_metric_remains_visible_without_failing():
|
||||
baseline = [{"throughput": 10.0}]
|
||||
current = {"model_id": "wan", "gpu_type": "NVIDIA L40S", "throughput": 8.0}
|
||||
policies = resolve_metric_policies({
|
||||
"throughput": {
|
||||
"threshold_percent": 0.01,
|
||||
"threshold_absolute": 0.01,
|
||||
"gated": False,
|
||||
}
|
||||
})
|
||||
|
||||
row = compare_baseline._build_summary_row(current, baseline, policies, False)
|
||||
|
||||
assert compare_baseline._check_regressions(current, baseline, policies) == []
|
||||
assert row["metrics"]["throughput"]["regression_pct"] == 20.0
|
||||
assert row["metrics"]["throughput"]["gated"] is False
|
||||
assert row["metrics"]["throughput"]["threshold_exceeded"] is True
|
||||
assert row["metrics"]["throughput"]["regressed"] is False
|
||||
assert row["threshold_exceeded_metrics"] == ["throughput"]
|
||||
assert row["failing_metrics"] == []
|
||||
|
||||
@@ -68,6 +68,8 @@ def test_summary_endpoint_returns_latest_group_status():
|
||||
assert body["count"] == 1
|
||||
assert body["status_counts"] == {"pass": 1, "fail": 0}
|
||||
assert body["rows"][0]["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert body["rows"][0]["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert body["rows"][0]["threshold_exceeded_metrics"] == ["latency", "throughput"]
|
||||
assert body["rows"][0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.performance import hf_store
|
||||
from fastvideo.performance_dashboard.service import build_latest_summary, build_trends, filter_records
|
||||
from fastvideo.tests.performance import hf_store
|
||||
|
||||
|
||||
def _record(ts, commit, latency, throughput, success=True, **metadata):
|
||||
@@ -28,16 +28,23 @@ def test_build_latest_summary_uses_previous_successful_records_for_baseline():
|
||||
_record("2026-01-03T00:00:00+00:00", "c" * 40, 11.0, 9.0),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05)
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert len(rows) == 1
|
||||
row = rows[0]
|
||||
assert row["baseline_n"] == 1
|
||||
assert row["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert row["metrics"]["latency"]["regression_pct"] == 10.0
|
||||
assert row["metrics"]["latency"]["absolute_delta"] == 1.0
|
||||
assert row["metrics"]["latency"]["threshold_percent"] == 8.0
|
||||
assert row["metrics"]["latency"]["threshold_absolute"] == 0.5
|
||||
assert row["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert row["metrics"]["latency"]["regressed"] is True
|
||||
assert row["metrics"]["throughput"]["regression_pct"] == 10.0
|
||||
assert row["status"] == "pass"
|
||||
assert row["computed_regression_status"] == "fail"
|
||||
assert row["threshold_exceeded_metrics"] == ["latency", "throughput"]
|
||||
assert row["failing_metrics"] == ["latency", "throughput"]
|
||||
|
||||
|
||||
def test_build_latest_summary_status_uses_latest_record_success_field():
|
||||
@@ -73,7 +80,7 @@ def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05, run_source="pr")
|
||||
rows = build_latest_summary(records, run_source="pr")
|
||||
|
||||
assert len(rows) == 1
|
||||
assert rows[0]["run_source"] == "pr"
|
||||
@@ -83,6 +90,60 @@ def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
|
||||
assert rows[0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
def test_build_latest_summary_requires_absolute_floor_for_computed_regression():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record(
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"b" * 40,
|
||||
10.6,
|
||||
10.0,
|
||||
regression_thresholds={
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
}
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert round(rows[0]["metrics"]["latency"]["regression_pct"], 1) == 6.0
|
||||
assert round(rows[0]["metrics"]["latency"]["absolute_delta"], 3) == 0.6
|
||||
assert rows[0]["metrics"]["latency"]["threshold_exceeded"] is False
|
||||
assert rows[0]["metrics"]["latency"]["regressed"] is False
|
||||
assert rows[0]["computed_regression_status"] == "pass"
|
||||
|
||||
|
||||
def test_build_latest_summary_separates_informational_threshold_crossing():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record(
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"b" * 40,
|
||||
10.6,
|
||||
10.0,
|
||||
regression_thresholds={
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
"gated": False,
|
||||
}
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert rows[0]["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert rows[0]["metrics"]["latency"]["regressed"] is False
|
||||
assert rows[0]["threshold_exceeded_metrics"] == ["latency"]
|
||||
assert rows[0]["failing_metrics"] == []
|
||||
assert rows[0]["computed_regression_status"] == "pass"
|
||||
|
||||
|
||||
def test_filter_records_and_trends_preserve_metric_points():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
|
||||
@@ -28,6 +28,17 @@ STAGE_METRIC_MAP: dict[str, str] = {
|
||||
"DmdDenoisingStage": "dit_time_s",
|
||||
"DecodingStage": "vae_decode_time_s",
|
||||
}
|
||||
V2_CONFIG_SCHEMA_VERSION = 2
|
||||
V2_REQUIRED_IDENTITY_FIELDS = (
|
||||
"workload_id",
|
||||
"variant_id",
|
||||
"benchmark_version",
|
||||
)
|
||||
V2_OPTIONAL_METADATA_FIELDS = (
|
||||
"recipe",
|
||||
"metric_threshold_policy",
|
||||
"quality_metadata",
|
||||
)
|
||||
|
||||
# -- Config discovery -------------------------------------------------------
|
||||
|
||||
@@ -42,6 +53,71 @@ _BENCHMARKS_DIR = os.path.join(
|
||||
)
|
||||
|
||||
|
||||
def _has_v2_fields(cfg):
|
||||
v2_fields = V2_REQUIRED_IDENTITY_FIELDS + V2_OPTIONAL_METADATA_FIELDS
|
||||
return any(field in cfg for field in v2_fields)
|
||||
|
||||
|
||||
def _is_v2_config(cfg):
|
||||
return cfg.get("config_schema_version") == V2_CONFIG_SCHEMA_VERSION
|
||||
|
||||
|
||||
def _validate_non_empty_string(value, field, path):
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError(f"{path}: v2 identity field {field!r} must be a non-empty string")
|
||||
|
||||
|
||||
def _validate_integer(value, field, path):
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"{path}: v2 identity field {field!r} must be an integer")
|
||||
|
||||
|
||||
def _validate_benchmark_config(cfg, path="<memory>"):
|
||||
missing_common = [field for field in ("benchmark_id",) if field not in cfg]
|
||||
if missing_common:
|
||||
raise ValueError(f"{path}: missing required benchmark config fields: {', '.join(missing_common)}")
|
||||
|
||||
schema_version = cfg.get("config_schema_version")
|
||||
if schema_version is None:
|
||||
if _has_v2_fields(cfg):
|
||||
raise ValueError(f"{path}: v2 benchmark identity fields require config_schema_version=2")
|
||||
return
|
||||
|
||||
if schema_version != V2_CONFIG_SCHEMA_VERSION:
|
||||
raise ValueError(f"{path}: unsupported benchmark config_schema_version={schema_version!r}")
|
||||
|
||||
missing_v2 = [field for field in V2_REQUIRED_IDENTITY_FIELDS if field not in cfg]
|
||||
if missing_v2:
|
||||
raise ValueError(f"{path}: missing required v2 identity fields: {', '.join(missing_v2)}")
|
||||
|
||||
_validate_non_empty_string(cfg["workload_id"], "workload_id", path)
|
||||
_validate_non_empty_string(cfg["variant_id"], "variant_id", path)
|
||||
_validate_integer(cfg["benchmark_version"], "benchmark_version", path)
|
||||
|
||||
for field in V2_OPTIONAL_METADATA_FIELDS:
|
||||
if field in cfg and not isinstance(cfg[field], Mapping):
|
||||
raise ValueError(f"{path}: optional v2 metadata field {field!r} must be an object")
|
||||
|
||||
|
||||
def _config_identity_metadata(cfg):
|
||||
if not _is_v2_config(cfg):
|
||||
return {}
|
||||
metadata = {
|
||||
"config_schema_version": cfg["config_schema_version"],
|
||||
"workload_id": cfg["workload_id"],
|
||||
"variant_id": cfg["variant_id"],
|
||||
"benchmark_version": cfg["benchmark_version"],
|
||||
}
|
||||
for field in V2_OPTIONAL_METADATA_FIELDS:
|
||||
if field in cfg:
|
||||
metadata[field] = cfg[field]
|
||||
return metadata
|
||||
|
||||
|
||||
def _benchmark_display_id(cfg):
|
||||
return cfg["benchmark_id"]
|
||||
|
||||
|
||||
def _discover_benchmarks():
|
||||
"""Glob benchmark JSON configs and return list of (id, config) tuples."""
|
||||
pattern = os.path.join(_BENCHMARKS_DIR, "*.json")
|
||||
@@ -49,6 +125,7 @@ def _discover_benchmarks():
|
||||
for path in sorted(glob.glob(pattern)):
|
||||
with open(path) as f:
|
||||
cfg = json.load(f)
|
||||
_validate_benchmark_config(cfg, path)
|
||||
configs.append(cfg)
|
||||
return configs
|
||||
|
||||
@@ -102,7 +179,11 @@ def _extract_component_times(result: dict) -> dict[str, float | None]:
|
||||
logger.debug("Skipping malformed stage '%s' data: %r", stage_name, stage_data)
|
||||
continue
|
||||
stage_class = stage_data.get("stage_class", stage_name)
|
||||
metric_key = STAGE_METRIC_MAP.get(stage_class)
|
||||
component_metric = stage_data.get("component_metric")
|
||||
if isinstance(component_metric, str) and component_metric in component_times:
|
||||
metric_key = component_metric
|
||||
else:
|
||||
metric_key = STAGE_METRIC_MAP.get(stage_class)
|
||||
if metric_key is None:
|
||||
logger.debug("Unmapped stage '%s' class '%s' (%.3fs)",
|
||||
stage_name,
|
||||
@@ -219,6 +300,7 @@ def _run_benchmark(cfg):
|
||||
|
||||
results = {
|
||||
"benchmark_id": cfg["benchmark_id"],
|
||||
**_config_identity_metadata(cfg),
|
||||
"model_short_name": model_info.get("model_short_name", ""),
|
||||
"device": device_name,
|
||||
"num_gpus": init_kwargs.get("num_gpus", 1),
|
||||
@@ -231,6 +313,7 @@ def _run_benchmark(cfg):
|
||||
"max_peak_memory_mb": round(max_peak_memory, 1),
|
||||
"individual_peak_memories_mb": [round(m, 1) for m in peak_memories],
|
||||
"thresholds": thresholds,
|
||||
"regression_thresholds": cfg.get("regression_thresholds", {}),
|
||||
"commit": os.environ.get("BUILDKITE_COMMIT", ""),
|
||||
"pr_number": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
@@ -275,7 +358,7 @@ def _run_benchmark(cfg):
|
||||
@pytest.mark.parametrize(
|
||||
"cfg",
|
||||
_BENCHMARK_CONFIGS,
|
||||
ids=[c["benchmark_id"] for c in _BENCHMARK_CONFIGS],
|
||||
ids=[_benchmark_display_id(c) for c in _BENCHMARK_CONFIGS],
|
||||
)
|
||||
def test_inference_performance(cfg):
|
||||
"""Measure generation latency, peak GPU memory, and component-level timings
|
||||
|
||||
@@ -1,13 +1,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import PipelineLoggingInfo
|
||||
from fastvideo.pipelines.stages.denoising import Cosmos25AutoDenoisingStage, DenoisingStage
|
||||
from fastvideo.pipelines.stages.text_encoding import Cosmos25TextEncodingStage
|
||||
from fastvideo.tests.performance.test_inference_performance import _extract_component_times
|
||||
|
||||
|
||||
class SubclassStyleDenoisingStage(DenoisingStage):
|
||||
pass
|
||||
|
||||
|
||||
def test_extract_component_times_handles_pipeline_logging_info_object():
|
||||
logging_info = PipelineLoggingInfo()
|
||||
logging_info.add_stage_execution_time("prompt_encoding_stage", 1.25)
|
||||
logging_info.add_stage_metric("prompt_encoding_stage", "stage_class", "TextEncodingStage")
|
||||
logging_info.add_stage_metric("prompt_encoding_stage", "component_metric", "text_encoder_time_s")
|
||||
|
||||
assert _extract_component_times({"logging_info": logging_info}) == {
|
||||
"text_encoder_time_s": 1.25,
|
||||
@@ -45,6 +52,35 @@ def test_extract_component_times_uses_stage_class_for_pipeline_stage_keys():
|
||||
}
|
||||
|
||||
|
||||
def test_denoising_stage_subclasses_inherit_component_metric():
|
||||
assert SubclassStyleDenoisingStage.performance_component_metric == "dit_time_s"
|
||||
|
||||
|
||||
def test_cosmos25_direct_pipeline_stages_define_component_metrics():
|
||||
assert Cosmos25TextEncodingStage.performance_component_metric == "text_encoder_time_s"
|
||||
assert Cosmos25AutoDenoisingStage.performance_component_metric == "dit_time_s"
|
||||
|
||||
|
||||
def test_extract_component_times_uses_component_metric_for_stage_subclasses():
|
||||
result = {
|
||||
"logging_info": {
|
||||
"stages": {
|
||||
"denoising_stage": {
|
||||
"execution_time": 4.2,
|
||||
"stage_class": "CosmosDenoisingStage",
|
||||
"component_metric": "dit_time_s",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
assert _extract_component_times(result) == {
|
||||
"text_encoder_time_s": None,
|
||||
"dit_time_s": 4.2,
|
||||
"vae_decode_time_s": None,
|
||||
}
|
||||
|
||||
|
||||
def test_extract_component_times_keeps_legacy_class_name_keys():
|
||||
# Backward-compatibility check for logs produced before pipeline-unique
|
||||
# stage keys carried a separate stage_class field.
|
||||
|
||||
@@ -21,6 +21,26 @@ class FakeTokenizer:
|
||||
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
|
||||
})
|
||||
|
||||
|
||||
class FakeChatTokenizer:
|
||||
def __init__(self):
|
||||
self.last_messages = None
|
||||
self.last_kwargs = None
|
||||
|
||||
def apply_chat_template(self, messages, **kwargs):
|
||||
self.last_messages = messages
|
||||
self.last_kwargs = kwargs
|
||||
assert isinstance(messages[0], list)
|
||||
assert messages[0][0]["role"] == "system"
|
||||
assert messages[0][1]["role"] == "user"
|
||||
B = len(messages)
|
||||
seq_len = int(kwargs.get("max_length", 4))
|
||||
return TensorDict({
|
||||
"input_ids": torch.arange(B * seq_len).view(B, seq_len),
|
||||
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
|
||||
})
|
||||
|
||||
|
||||
class FakeTextEncoder(torch.nn.Module):
|
||||
def __init__(self, hidden_size=8):
|
||||
super().__init__()
|
||||
@@ -38,6 +58,14 @@ class FakeTextEncoder(torch.nn.Module):
|
||||
def id_preprocess(x: str) -> str:
|
||||
return x
|
||||
|
||||
|
||||
def chat_list_preprocess(x: str):
|
||||
return [
|
||||
{"role": "system", "content": "Describe the video."},
|
||||
{"role": "user", "content": x if x else " "},
|
||||
]
|
||||
|
||||
|
||||
def take_mean_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
# [B, T, H] -> [B, H]
|
||||
return outputs.last_hidden_state.mean(dim=1)
|
||||
@@ -156,3 +184,32 @@ def test_encode_text_does_not_force_hidden_states_for_ltx2_prefix():
|
||||
stage.encode_text("a", fastvideo_args, encoder_index=[0])
|
||||
|
||||
assert stage.text_encoders[0].last_output_hidden_states is False
|
||||
|
||||
|
||||
def test_chat_list_preprocess_output_is_not_stripped():
|
||||
fastvideo_args, hidden = make_args(num_encoders=1, text_len=5, hidden_size=8)
|
||||
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[0]
|
||||
encoder_config.is_chat_model = True
|
||||
encoder_config.treat_empty_as_dot = True
|
||||
fastvideo_args.pipeline_config.preprocess_text_funcs = (chat_list_preprocess, )
|
||||
|
||||
tokenizer = FakeChatTokenizer()
|
||||
stage = TextEncodingStage(
|
||||
text_encoders=[FakeTextEncoder(hidden_size=hidden)],
|
||||
tokenizers=[tokenizer],
|
||||
)
|
||||
|
||||
embeds, masks = stage.encode_text(
|
||||
"a robotic arm welding a metal structure",
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
|
||||
assert embeds[0].shape == (1, hidden)
|
||||
assert masks[0].shape == (1, 5)
|
||||
assert tokenizer.last_messages == [[
|
||||
{"role": "system", "content": "Describe the video."},
|
||||
{"role": "user", "content": "a robotic arm welding a metal structure"},
|
||||
]]
|
||||
assert tokenizer.last_kwargs["return_tensors"] == "pt"
|
||||
|
||||
@@ -127,4 +127,12 @@ def test_wan_causal_dfsft_single_train_step(
|
||||
|
||||
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
|
||||
# Skips when the current GPU has no seeded reference.
|
||||
check_grad_norm_regression("test_wan_causal_dfsft", model.transformer)
|
||||
# rtol above the harness default: the causal model compiles flex_attention
|
||||
# with max-autotune (required for Wan 1.3B's head config), and the
|
||||
# timing-based kernel selection is bimodal across L40S containers —
|
||||
# observed 3.2562 vs 3.5860 (10.13% apart) with identical code, straddling
|
||||
# the default 10%. 12% covers both winners; real wiring breakage (dead
|
||||
# grads, scale bugs) still lands far outside it.
|
||||
check_grad_norm_regression("test_wan_causal_dfsft",
|
||||
model.transformer,
|
||||
rtol=0.12)
|
||||
|
||||
Reference in New Issue
Block a user