Compare commits

...
29 Commits
Author SHA1 Message Date
SolitaryThinker 82b6d2dbf0 [misc]: drop .agents/exploration journals added by #1453
AGENTS.md bans experiment journals and branch-state snapshots under
.agents/.
2026-07-05 13:46:48 -07:00
SolitaryThinker c6468e8def [misc]: collect fastvideo/tests/batching in the Modal unit-test lane (#1453)
The new batching unit tests (admission, signature) were not collected by
any CI lane. stages/ is already collected on main, so only batching/ is
added.
2026-07-05 13:46:41 -07:00
SolitaryThinker 2c0d9a0828 [bugfix]: keep completed results when a prompt-file batch group fails (#1453)
The dynamic-batching prompt-file branch let a single failed group raise
out of generate_video, discarding every completed result (and tripping
the strict zip). Restore the legacy per-prompt tolerance: a failed group
now yields an error entry per prompt, keeping results aligned with the
input prompts while completed videos survive.
2026-07-05 13:46:22 -07:00
SolitaryThinker b9876ee06d [bugfix]: fail fast when submitting to a stopped batch scheduler (#1453)
submit() after stop() enqueued jobs the drained run loop would never
resolve, hanging the caller forever on the future. Raise immediately
instead.
2026-07-05 13:45:35 -07:00
SolitaryThinker a568406b8c [bugfix]: drain queued jobs greedily when batching delay is 0 (#1453)
With the default batching_delay_ms=0, _collect_batch called
wait_for(timeout=0) which timed out before draining the queue, so every
dispatched batch had size 1 and dynamic batching silently no-oped. When
delay is 0, drain already-queued jobs non-blocking (until empty,
max_size, or an incompatible job) instead of waiting.
2026-07-05 13:45:21 -07:00
SolitaryThinker 39cc075452 [bugfix]: reserve batch output paths so duplicate prompts do not collide (#1453)
Batch mode resolves every output path via _prepare_output_path before any
file is written, so the os.path.exists dedupe never fired and duplicate
prompts (or prompts equal in their first 100 sanitized chars) overwrote
each other. Track paths reserved within the batch and suffix duplicates,
matching what the legacy sequential path produces for the same inputs.
2026-07-05 13:44:57 -07:00
Mac Lee 88f50fa7d6 [bugfix]: address batching review comments (#1453) 2026-07-05 13:39:16 -07:00
Mac Lee 4b25087bd8 [bugfix]: update batching schema test fixtures (#1453) 2026-07-05 13:39:16 -07:00
Mac Lee a314489b78 [docs]: record checklist closure (#1453) 2026-07-05 13:39:16 -07:00
Mac Lee a1b8a78e9e [misc]: satisfy all-files pre-commit formatting 2026-07-05 13:39:16 -07:00
Mac Lee df55a31824 [docs]: record SSIM validation follow-up 2026-07-05 13:39:16 -07:00
Mac Lee d022cf00e3 [docs]: record final batching validation result 2026-07-05 13:39:16 -07:00
Mac Lee 7dc3166827 [docs]: record multimodal batching validation report 2026-07-05 13:39:16 -07:00
Mac Lee ce4c4f37c7 [fix]: harden dynamic generation batching 2026-07-05 13:39:16 -07:00
Mac Lee 48dff4be69 [misc]: record batching stage 5 state 2026-07-05 13:39:16 -07:00
Mac Lee d9736d5e5d [feat]: add OpenAI video batching scheduler 2026-07-05 13:39:16 -07:00
Mac Lee 7396024f95 [misc]: record batching stage 2 state 2026-07-05 13:39:16 -07:00
Mac Lee 978896f534 [feat]: add generator dynamic batching path 2026-07-05 13:39:16 -07:00
Mac Lee 560810c592 [misc]: record batching stage 1 state 2026-07-05 13:39:16 -07:00
Mac Lee e3a7450a05 [feat]: add generation batching primitives 2026-07-05 13:39:15 -07:00
Mac LeeandSolitaryThinker 6aab7f3832 [ci] cover Hunyuan 1.5 chat-list text preprocessing (#1518)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-05 12:05:25 -07:00
Mac Lee 9cd53fe5f8 [ci] Add metric-specific performance thresholds (#1545) 2026-07-05 12:04:33 -07:00
Mac Lee 6a32cf3a5e [ci]: expose LoRA extraction slash command (#1542) 2026-07-05 06:45:31 -07:00
William Lin 98be9b3da2 [bugfix]: address the three remaining #1447 review findings (#1549) 2026-07-05 06:43:27 -07:00
Mac Lee c53e85b767 [ci] Add v2 performance benchmark config identity fields (#1544) 2026-07-05 06:18:15 -07:00
William Lin 40a8bd2d3b [bugfix]: skip ThunderKittens kernels on aarch64 and document the kernel build matrix (#1548) 2026-07-05 06:13:11 -07:00
Mac LeeandSolitaryThinker 31aa115611 [bugfix]: preserve FSDP hooks for RMSNorm qk norms (#1513)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-05 04:15:49 +00:00
zainnhandSolitaryThinker 98ac10a528 [infra] Auto-rebuild CUDA images when docker/Dockerfile changes on main (#1526)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-04 16:25:46 -07:00
William Lin a5a6d171e5 [attn] Make FA4 explicit opt-in via FASTVIDEO_FA4 and delete the runtime fallback machinery (#1540) 2026-07-04 14:51:00 -07:00
66 changed files with 3765 additions and 364 deletions
@@ -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",
+26
View File
@@ -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"
+2 -1
View File
@@ -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
+30 -5
View File
@@ -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:
+3
View File
@@ -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
+8 -3
View File
@@ -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
View File
@@ -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" && \
+2
View File
@@ -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` |
+119 -23
View File
@@ -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,16 +95,18 @@ 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
@@ -156,9 +161,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 +190,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 +236,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 +257,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 +290,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 +313,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 +335,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 +355,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 +380,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 +409,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,
@@ -41,6 +41,11 @@ surfaces:
disable_autocast: generator.engine.disable_autocast
enable_stage_verification: generator.engine.enable_stage_verification
prompt_txt: request.inputs.prompt_path
batching_mode: generator.engine.batching.mode
batching_max_size: generator.engine.batching.max_size
batching_delay_ms: generator.engine.batching.delay_ms
batching_config: generator.engine.batching.config_path
enable_batching_metrics: generator.engine.batching.enable_metrics
override_text_encoder_safetensors: generator.pipeline.components.text_encoder_weights
override_text_encoder_quant: generator.engine.quantization.text_encoder_quant
transformer_quant: generator.engine.quantization.transformer_quant
+17
View File
@@ -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`**
+31 -5
View File
@@ -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 "============================================================")
+36
View File
@@ -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)
+17
View File
@@ -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
+15
View File
@@ -159,6 +159,16 @@ def legacy_from_pretrained_to_config(
preset_refine["guidance_scale"] = value
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast"}:
engine[key] = value
elif key == "batching_mode":
engine.setdefault("batching", {})["mode"] = value
elif key == "batching_max_size":
engine.setdefault("batching", {})["max_size"] = value
elif key == "batching_delay_ms":
engine.setdefault("batching", {})["delay_ms"] = value
elif key == "batching_config":
engine.setdefault("batching", {})["config_path"] = value
elif key == "enable_batching_metrics":
engine.setdefault("batching", {})["enable_metrics"] = value
elif key == "override_text_encoder_quant":
quantization["text_encoder_quant"] = value
elif key == "workload_type":
@@ -244,6 +254,11 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
"enable_stage_verification": engine.enable_stage_verification,
"use_fsdp_inference": engine.use_fsdp_inference,
"disable_autocast": engine.disable_autocast,
"batching_mode": engine.batching.mode,
"batching_max_size": engine.batching.max_size,
"batching_delay_ms": engine.batching.delay_ms,
"batching_config": engine.batching.config_path,
"enable_batching_metrics": engine.batching.enable_metrics,
}
if normalized.pipeline.workload_type is not None:
kwargs["workload_type"] = normalized.pipeline.workload_type
+11
View File
@@ -71,6 +71,15 @@ class QuantizationConfig:
transformer_quant: str | None = None
@dataclass
class BatchingConfig:
mode: Literal["disabled", "dynamic"] = "disabled"
max_size: int = 1
delay_ms: float = 0.0
config_path: str | None = None
enable_metrics: bool = False
@dataclass
class EngineConfig:
num_gpus: int = 1
@@ -82,6 +91,7 @@ class EngineConfig:
use_fsdp_inference: bool = False
disable_autocast: bool = False
quantization: QuantizationConfig | None = None
batching: BatchingConfig = field(default_factory=BatchingConfig)
@dataclass
@@ -280,6 +290,7 @@ class ServeConfig:
__all__ = [
"BatchingConfig",
"CompileConfig",
"ComponentConfig",
"ContinuationState",
+39 -23
View File
@@ -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)
+66 -73
View File
@@ -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,
+23 -12
View File
@@ -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()
+26
View File
@@ -0,0 +1,26 @@
# SPDX-License-Identifier: Apache-2.0
"""Dynamic generation batching helpers."""
from fastvideo.batching.admission import (
AdmissionLimit,
BatchAdmissionController,
BatchingRule,
load_batching_config,
)
from fastvideo.batching.signature import (
BatchCompatibility,
can_dynamic_batch,
dynamic_batch_signature,
resolution_key,
)
__all__ = [
"AdmissionLimit",
"BatchAdmissionController",
"BatchCompatibility",
"BatchingRule",
"can_dynamic_batch",
"dynamic_batch_signature",
"load_batching_config",
"resolution_key",
]
+297
View File
@@ -0,0 +1,297 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
import os
from dataclasses import dataclass
from difflib import get_close_matches
from typing import Any
from fastvideo.batching.signature import resolution_key
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
_BYTES_PER_GB = 1024**3
_BATCHING_RULE_KEYS = frozenset({
"model",
"model_contains",
"resolution",
"device_memory_gb_min",
"device_memory_gb_max",
"offload",
"max_batch_size",
"max_cost",
"calibration",
})
@dataclass(frozen=True)
class AdmissionLimit:
max_batch_size: int
max_cost: float | None = None
cap_reason: str | None = None
def reject_reason(self, *, batch_size: int, batch_cost: float) -> str | None:
if batch_size > self.max_batch_size:
return self.cap_reason or f"config_cap:{self.max_batch_size}"
if self.max_cost is not None and batch_cost > self.max_cost:
return f"cost_budget:{batch_cost:.0f}>{self.max_cost:.0f}"
return None
def stop_reason_for_next_cost(self, next_batch_cost: float) -> str | None:
if self.max_cost is not None and next_batch_cost > self.max_cost:
return f"cost_budget_next:{next_batch_cost:.0f}>{self.max_cost:.0f}"
return None
@dataclass(frozen=True)
class BatchingRule:
model: str | None = None
model_contains: str | None = None
resolution: str | None = None
device_memory_gb_min: float | None = None
device_memory_gb_max: float | None = None
offload: bool | None = None
max_batch_size: int = 1
max_cost: float | None = None
source: str = "user"
@classmethod
def from_dict(cls, data: dict[str, Any], *, source: str) -> BatchingRule:
if not isinstance(data, dict):
raise ValueError(f"batching config rule from {source} must be an object, got {type(data).__name__}")
_validate_rule_keys(data, source=source)
if "max_batch_size" not in data:
raise ValueError("batching config rule requires max_batch_size")
rule = cls(
model=_optional_str(data.get("model")),
model_contains=_optional_str(data.get("model_contains")),
resolution=_optional_str(data.get("resolution")),
device_memory_gb_min=_optional_float(data.get("device_memory_gb_min")),
device_memory_gb_max=_optional_float(data.get("device_memory_gb_max")),
offload=_optional_bool(data.get("offload")),
max_batch_size=int(data["max_batch_size"]),
max_cost=_optional_float(data.get("max_cost")),
source=source,
)
rule.validate()
return rule
def validate(self) -> None:
if self.model is not None and self.model_contains is not None:
raise ValueError("batching config rule cannot set both model and model_contains")
if self.model is None and self.model_contains is None:
raise ValueError("batching config rule requires model or model_contains")
if self.max_batch_size < 1:
raise ValueError("batching config rule max_batch_size must be >= 1")
if self.max_cost is not None and self.max_cost <= 0.0:
raise ValueError("batching config rule max_cost must be > 0")
if (self.device_memory_gb_min is not None and self.device_memory_gb_max is not None
and self.device_memory_gb_min > self.device_memory_gb_max):
raise ValueError("batching config rule device_memory_gb_min must be <= device_memory_gb_max")
def matches(
self,
*,
model_path: str,
resolution: str | None,
device_memory_gb: float | None,
offload: bool,
) -> bool:
if self.model is not None and self.model != model_path:
return False
if self.model_contains is not None and self.model_contains not in model_path:
return False
if self.resolution not in (None, "*") and self.resolution != resolution:
return False
if self.offload is not None and self.offload != offload:
return False
if device_memory_gb is None:
return True
if self.device_memory_gb_min is not None and device_memory_gb < self.device_memory_gb_min:
return False
return not (self.device_memory_gb_max is not None and device_memory_gb > self.device_memory_gb_max)
class BatchAdmissionController:
def __init__(self, fastvideo_args: FastVideoArgs, *, gpu_id: int = 0):
self._mode = fastvideo_args.batching_mode
self._user_max_batch_size = max(1, int(fastvideo_args.batching_max_size))
self._model_path = fastvideo_args.model_path
self._offload = bool(fastvideo_args.dit_cpu_offload or fastvideo_args.dit_layerwise_offload)
self._device_memory_gb = self._get_device_memory_gb(gpu_id)
self._rules = load_batching_config(fastvideo_args.batching_config)
self._pipeline_config = fastvideo_args.pipeline_config
if self.enabled:
logger.info(
"Batch admission enabled: user_max=%d, device_memory=%.1fGiB, rules=%d",
self._user_max_batch_size,
self._device_memory_gb or 0.0,
len(self._rules),
)
@property
def enabled(self) -> bool:
return self._mode == "dynamic" and self._user_max_batch_size > 1
def reject_reason_for_candidate(self, current_requests: list[Any], candidate_request: Any) -> str | None:
if not self.enabled:
return None
proposed = current_requests + [candidate_request]
limit = self.limit_for(proposed[0])
return limit.reject_reason(
batch_size=len(proposed),
batch_cost=self.estimate_batch_cost(proposed),
)
def batch_is_full(self, requests: list[Any]) -> bool:
if not self.enabled or not requests:
return len(requests) >= self._user_max_batch_size
limit = self.limit_for(requests[0])
if len(requests) >= limit.max_batch_size:
return True
next_cost = self.estimate_batch_cost(requests + [requests[0]])
return limit.max_cost is not None and next_cost > limit.max_cost
def limit_reason_for_batch(self, requests: list[Any]) -> str | None:
if not self.enabled or not requests:
return None
limit = self.limit_for(requests[0])
if len(requests) >= limit.max_batch_size:
return limit.cap_reason or f"config_cap:{limit.max_batch_size}"
next_cost = self.estimate_batch_cost(requests + [requests[0]])
return limit.stop_reason_for_next_cost(next_cost)
def max_admissible_batch_size(self, request: Any) -> int:
return self.limit_for(request).max_batch_size
def limit_for(self, request: Any) -> AdmissionLimit:
rules = self._matching_rules(request)
if not rules:
return AdmissionLimit(max_batch_size=self._user_max_batch_size)
config_cap = min(rule.max_batch_size for rule in rules)
max_batch_size = min(self._user_max_batch_size, config_cap)
cap_reason = f"config_cap:{max_batch_size}" if max_batch_size < self._user_max_batch_size else None
costs = [rule.max_cost for rule in rules if rule.max_cost is not None]
return AdmissionLimit(
max_batch_size=max(1, max_batch_size),
max_cost=min(costs) if costs else None,
cap_reason=cap_reason,
)
def estimate_batch_cost(self, requests: list[Any]) -> float:
return sum(float(self._pipeline_config.estimate_request_cost(request)) for request in requests)
def _matching_rules(self, request: Any) -> list[BatchingRule]:
return [
rule for rule in self._rules if rule.matches(
model_path=self._model_path,
resolution=resolution_key(request),
device_memory_gb=self._device_memory_gb,
offload=self._offload,
)
]
@staticmethod
def _get_device_memory_gb(gpu_id: int) -> float | None:
try:
from fastvideo.platforms import current_platform
return current_platform.get_device_total_memory(gpu_id) / _BYTES_PER_GB
except Exception:
return None
def load_batching_config(path: str | None) -> list[BatchingRule]:
if path is None:
return []
with open(path, encoding="utf-8") as f:
payload = json.load(f)
source = os.path.abspath(path)
entries = _config_entries(payload)
rules = [BatchingRule.from_dict(entry, source=source) for entry in entries]
if not rules:
raise ValueError(f"batching config {source} does not contain any rules")
return rules
def _config_entries(payload: Any) -> list[dict[str, Any]]:
if isinstance(payload, dict) and payload.get("schema_version") not in (None, 1):
raise ValueError("batching config schema_version must be 1")
if isinstance(payload, dict) and isinstance(payload.get("rules"), list):
return payload["rules"]
if isinstance(payload, list):
return payload
if isinstance(payload, dict):
entries: list[dict[str, Any]] = []
for key, value in payload.items():
if key == "schema_version" or not isinstance(value, dict):
continue
model, _sep, resolution = key.partition("|")
entry = dict(value)
if model:
entry.setdefault("model", model)
if resolution:
entry.setdefault("resolution", resolution)
entries.append(entry)
return entries
raise ValueError("batching config must be a {'schema_version': 1, 'rules': [...]} object, "
"a list of rules, or a mapping keyed by model|resolution")
def _validate_rule_keys(data: dict[str, Any], *, source: str) -> None:
unknown = sorted(set(data) - _BATCHING_RULE_KEYS)
if not unknown:
return
hints = []
for key in unknown:
matches = get_close_matches(key, _BATCHING_RULE_KEYS, n=1)
if matches:
hints.append(f"{key!r} (did you mean {matches[0]!r}?)")
else:
hints.append(repr(key))
raise ValueError(f"batching config rule from {source} contains unknown key(s): {', '.join(hints)}")
def _optional_str(value: Any) -> str | None:
if value is None:
return None
return str(value)
def _optional_float(value: Any) -> float | None:
if value is None:
return None
return float(value)
def _optional_bool(value: Any) -> bool | None:
if value is None:
return None
if isinstance(value, bool):
return value
if isinstance(value, int | float):
if value == 1.0:
return True
if value == 0.0:
return False
if isinstance(value, str):
lowered = value.strip().lower()
if lowered in ("1", "true", "yes", "y", "on"):
return True
if lowered in ("0", "false", "no", "n", "off"):
return False
raise ValueError(f"cannot parse boolean batching config value: {value!r}")
+178
View File
@@ -0,0 +1,178 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import dataclasses
from dataclasses import dataclass
from enum import Enum
from typing import Any
from fastvideo.api.sampling_param import SamplingParam
_SIGNATURE_EXCLUDED_FIELDS = frozenset({
"prompt",
"prompt_path",
"output_path",
"output_video_name",
"seed",
"save_video",
"return_frames",
})
_UNSUPPORTED_DYNAMIC_BATCH_FIELDS = frozenset({
"image_path",
"pil_image",
"video_path",
"mouse_cond",
"keyboard_cond",
"grid_sizes",
"pose",
"camera_states",
"camera_trajectory",
"action_list",
"action_speed_list",
"gt_latents",
"conditioning_mask",
"c2ws_plucker_emb",
"refine_from",
"stage1_video",
"trajectory_type",
"movement_distance",
"camera_rotation",
"ltx2_images",
"ltx2_conditioning_latent_stage1",
"ltx2_conditioning_latent_stage2",
"ltx2_video_conditions",
"init_audio",
"inpaint_audio",
"inpaint_mask",
"continuation_state",
})
_UNSUPPORTED_EXTRA_KEYS = frozenset({
"ltx2_audio_latents",
"ltx2_audio_clean_latent",
"ltx2_audio_denoise_mask",
"audio_num_frames",
"video_position_offset_sec",
})
@dataclass(frozen=True)
class BatchCompatibility:
can_batch: bool
reason: str | None = None
def resolution_key(request: Any) -> str:
height = _first_scalar(getattr(request, "height", None))
width = _first_scalar(getattr(request, "width", None))
num_frames = _first_scalar(getattr(request, "num_frames", None))
return f"{height}x{width}x{num_frames}"
def dynamic_batch_signature(
request: SamplingParam,
*,
extra: dict[str, Any] | None = None,
) -> tuple[tuple[str, Any], ...]:
"""Build a hashable compatibility signature for a generation request."""
signature_items: list[tuple[str, Any]] = []
for field in dataclasses.fields(request):
if field.name in _SIGNATURE_EXCLUDED_FIELDS:
continue
signature_items.append((field.name, _freeze_signature_value(getattr(request, field.name, None))))
if extra:
signature_items.append(("extra", _freeze_signature_value(extra)))
return tuple(signature_items)
def can_dynamic_batch(
base: SamplingParam,
candidate: SamplingParam,
*,
base_extra: dict[str, Any] | None = None,
candidate_extra: dict[str, Any] | None = None,
) -> BatchCompatibility:
"""Return whether two FastVideo generation requests can be merged."""
base_ready = _request_is_batchable(base, extra=base_extra)
if not base_ready.can_batch:
return base_ready
candidate_ready = _request_is_batchable(candidate, extra=candidate_extra)
if not candidate_ready.can_batch:
return candidate_ready
base_sig = dynamic_batch_signature(base, extra=base_extra)
candidate_sig = dynamic_batch_signature(candidate, extra=candidate_extra)
if base_sig == candidate_sig:
return BatchCompatibility(can_batch=True)
mismatch = _first_mismatch(base_sig, candidate_sig)
return BatchCompatibility(can_batch=False, reason=mismatch or "signature_mismatch")
def _request_is_batchable(
request: SamplingParam,
*,
extra: dict[str, Any] | None = None,
) -> BatchCompatibility:
if not isinstance(request.prompt, str):
return BatchCompatibility(can_batch=False, reason="prompt_type")
if request.prompt_path is not None:
return BatchCompatibility(can_batch=False, reason="prompt_path")
if request.num_videos_per_prompt != 1:
return BatchCompatibility(can_batch=False, reason="num_videos_per_prompt")
if request.return_continuation_state:
return BatchCompatibility(can_batch=False, reason="return_continuation_state")
for name in _UNSUPPORTED_DYNAMIC_BATCH_FIELDS:
value = getattr(request, name, None)
if _is_present(value):
return BatchCompatibility(can_batch=False, reason=name)
if extra:
unsupported = sorted(set(extra) & _UNSUPPORTED_EXTRA_KEYS)
if unsupported:
return BatchCompatibility(can_batch=False, reason=f"extra.{unsupported[0]}")
return BatchCompatibility(can_batch=True)
def _freeze_signature_value(value: Any) -> Any:
if isinstance(value, str | int | float | bool | type(None)):
return value
if isinstance(value, Enum):
return value.value
if isinstance(value, dict):
return tuple(
(str(key), _freeze_signature_value(item)) for key, item in sorted(value.items(), key=lambda kv: str(kv[0])))
if isinstance(value, list | tuple):
return tuple(_freeze_signature_value(item) for item in value)
return repr(value)
def _is_present(value: Any) -> bool:
if value is None:
return False
if value is False:
return False
return not (isinstance(value, list | tuple | dict | set) and not value)
def _first_scalar(value: Any) -> Any:
if isinstance(value, list | tuple):
return value[0] if value else None
return value
def _first_mismatch(
base_sig: tuple[tuple[str, Any], ...],
candidate_sig: tuple[tuple[str, Any], ...],
) -> str | None:
if len(base_sig) != len(candidate_sig):
return "sampling_params"
for (name, base_value), (candidate_name, candidate_value) in zip(base_sig, candidate_sig, strict=True):
if name != candidate_name:
return "sampling_params"
if base_value != candidate_value:
return f"sampling_params.{name}"
return None
+21
View File
@@ -276,6 +276,27 @@ class PipelineConfig:
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
def estimate_request_cost(self, request: Any) -> float:
"""Estimate relative memory/compute cost for batching admission.
The default is intentionally simple and model-agnostic: pixel count
times frame count. Pipeline subclasses can override this when they have
calibrated costs.
"""
height = getattr(request, "height", None)
width = getattr(request, "width", None)
num_frames = getattr(request, "num_frames", None)
if isinstance(height, list):
height = height[0] if height else None
if isinstance(width, list):
width = width[0] if width else None
if isinstance(num_frames, list):
num_frames = num_frames[0] if num_frames else None
height = int(height or 1)
width = int(width or 1)
num_frames = int(num_frames or 1)
return float(max(1, height) * max(1, width) * max(1, num_frames))
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
del_keys = []
+20 -1
View File
@@ -10,6 +10,7 @@ from fastapi.middleware.cors import CORSMiddleware
from fastvideo.api.presets import validate_preset_selection
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.openai.batching import VideoBatchScheduler
from fastvideo.entrypoints.openai.state import (
DEFAULT_OUTPUT_DIR,
clear_state,
@@ -59,11 +60,29 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
generator = VideoGenerator.from_fastvideo_args(args)
logger.info("Model loaded successfully.")
set_state(generator, args, output_dir, default_request=default_request)
video_batch_scheduler: VideoBatchScheduler | None = None
if args.batching_mode == "dynamic" and args.batching_max_size > 1:
video_batch_scheduler = VideoBatchScheduler(generator, args)
await video_batch_scheduler.start()
logger.info(
"Started dynamic video batch scheduler: max_size=%d delay_ms=%.2f",
args.batching_max_size,
args.batching_delay_ms,
)
set_state(
generator,
args,
output_dir,
default_request=default_request,
video_batch_scheduler=video_batch_scheduler,
)
yield # server is running
logger.info("Shutting down — releasing model resources ...")
if video_batch_scheduler is not None:
await video_batch_scheduler.stop()
generator.shutdown()
clear_state()
logger.info("Shutdown complete.")
+183
View File
@@ -0,0 +1,183 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import asyncio
import time
from collections import deque
from dataclasses import dataclass
from typing import Any
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.batching.signature import can_dynamic_batch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
@dataclass
class _VideoBatchJob:
request_id: str
kwargs: dict[str, Any]
future: asyncio.Future
enqueue_time: float
class VideoBatchScheduler:
"""Async FIFO scheduler for OpenAI-compatible video generation."""
def __init__(self, generator: Any, fastvideo_args: FastVideoArgs) -> None:
self._generator = generator
self._fastvideo_args = fastvideo_args
self._queue: asyncio.Queue[_VideoBatchJob | None] = asyncio.Queue()
self._pending: deque[_VideoBatchJob] = deque()
self._task: asyncio.Task | None = None
self._stopped = False
@property
def enabled(self) -> bool:
return self._fastvideo_args.batching_mode == "dynamic" and self._fastvideo_args.batching_max_size > 1
async def start(self) -> None:
if self._task is not None:
return
self._task = asyncio.create_task(self._run(), name="fastvideo-video-batch-scheduler")
async def stop(self) -> None:
self._stopped = True
await self._queue.put(None)
if self._task is not None:
await self._task
self._task = None
async def submit(self, request_id: str, kwargs: dict[str, Any]) -> Any:
if self._stopped:
raise RuntimeError("Video batch scheduler is stopped; cannot submit new requests")
loop = asyncio.get_running_loop()
future = loop.create_future()
await self._queue.put(
_VideoBatchJob(
request_id=request_id,
kwargs=dict(kwargs),
future=future,
enqueue_time=time.perf_counter(),
))
return await future
async def _run(self) -> None:
while not self._stopped:
job = await self._get_next_job()
if job is None:
break
batch = await self._collect_batch(job)
await self._dispatch(batch)
while self._pending:
pending = self._pending.popleft()
if not pending.future.done():
pending.future.set_exception(RuntimeError("Video batch scheduler stopped before dispatch"))
async def _get_next_job(self) -> _VideoBatchJob | None:
if self._pending:
return self._pending.popleft()
return await self._queue.get()
async def _collect_batch(self, first: _VideoBatchJob) -> list[_VideoBatchJob]:
batch = [first]
max_size = self._fastvideo_args.batching_max_size
delay_s = max(0.0, self._fastvideo_args.batching_delay_ms / 1000.0)
deadline = first.enqueue_time + delay_s
while len(batch) < max_size:
if delay_s > 0:
timeout = deadline - time.perf_counter()
if timeout <= 0:
break
try:
candidate = await asyncio.wait_for(self._get_next_job(), timeout=timeout)
except TimeoutError:
break
else:
# delay=0 means "don't wait": greedily drain whatever is
# already queued so max_size still coalesces.
if self._pending:
candidate = self._pending.popleft()
else:
try:
candidate = self._queue.get_nowait()
except asyncio.QueueEmpty:
break
if candidate is None:
await self._queue.put(None)
break
if self._jobs_are_compatible(batch[0], candidate):
batch.append(candidate)
continue
self._pending.appendleft(candidate)
break
return batch
async def _dispatch(self, batch: list[_VideoBatchJob]) -> None:
loop = asyncio.get_running_loop()
request_ids = [job.request_id for job in batch]
queue_wait_ms = (time.perf_counter() - min(job.enqueue_time for job in batch)) * 1000.0
if self._fastvideo_args.enable_batching_metrics:
logger.info(
"Dispatching video batch: request_ids=%s size=%d queue_wait_ms=%.2f",
request_ids,
len(batch),
queue_wait_ms,
)
try:
results = await loop.run_in_executor(
None,
lambda: self._generator.generate_video_batch([job.kwargs for job in batch]),
)
except Exception as exc:
for job in batch:
if not job.future.done():
job.future.set_exception(exc)
return
if len(results) != len(batch):
error = RuntimeError(f"Video batch returned {len(results)} results for {len(batch)} requests")
for job in batch:
if not job.future.done():
job.future.set_exception(error)
return
for job, result in zip(batch, results, strict=True):
if not job.future.done():
job.future.set_result(result)
def _jobs_are_compatible(self, base: _VideoBatchJob, candidate: _VideoBatchJob) -> bool:
try:
base_sampling, base_extra = self._sampling_param_from_kwargs(base.kwargs)
candidate_sampling, candidate_extra = self._sampling_param_from_kwargs(candidate.kwargs)
except Exception:
return False
return can_dynamic_batch(
base_sampling,
candidate_sampling,
base_extra=base_extra,
candidate_extra=candidate_extra,
).can_batch
def _sampling_param_from_kwargs(self, kwargs: dict[str, Any]) -> tuple[SamplingParam, dict[str, Any]]:
sampling_param = SamplingParam.from_pretrained(self._fastvideo_args.model_path)
updates = dict(kwargs)
prompt = updates.pop("prompt", None)
extra: dict[str, Any] = {}
for key in (
"ltx2_audio_latents",
"ltx2_audio_clean_latent",
"ltx2_audio_denoise_mask",
"audio_num_frames",
"video_position_offset_sec",
):
if key in updates:
extra[key] = updates.pop(key)
sampling_param.update(updates)
sampling_param.prompt = prompt
return sampling_param, extra
+12 -2
View File
@@ -11,6 +11,7 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.openai.batching import VideoBatchScheduler
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.fastvideo_args import FastVideoArgs
@@ -20,6 +21,7 @@ _generator: VideoGenerator | None = None
_fastvideo_args: FastVideoArgs | None = None
_output_dir: str = DEFAULT_OUTPUT_DIR
_default_request: GenerationRequest | None = None
_video_batch_scheduler: VideoBatchScheduler | None = None
def get_generator() -> VideoGenerator:
@@ -44,23 +46,31 @@ def get_default_request() -> GenerationRequest | None:
return _default_request
def get_video_batch_scheduler() -> VideoBatchScheduler | None:
"""Return the video batch scheduler when dynamic batching is enabled."""
return _video_batch_scheduler
def set_state(
generator: VideoGenerator,
fastvideo_args: FastVideoArgs,
output_dir: str,
default_request: GenerationRequest | None = None,
video_batch_scheduler: VideoBatchScheduler | None = None,
) -> None:
"""Set all server state at once (called from lifespan)."""
global _generator, _fastvideo_args, _output_dir, _default_request
global _generator, _fastvideo_args, _output_dir, _default_request, _video_batch_scheduler
_generator = generator
_fastvideo_args = fastvideo_args
_output_dir = output_dir
_default_request = default_request
_video_batch_scheduler = video_batch_scheduler
def clear_state() -> None:
"""Clear server state on shutdown."""
global _generator, _fastvideo_args, _default_request
global _generator, _fastvideo_args, _default_request, _video_batch_scheduler
_generator = None
_fastvideo_args = None
_default_request = None
_video_batch_scheduler = None
+9 -4
View File
@@ -26,6 +26,7 @@ from fastvideo.entrypoints.openai.state import (
get_generator,
get_output_dir,
get_server_args,
get_video_batch_scheduler,
)
from fastvideo.entrypoints.openai.protocol import (
VideoGenerationsRequest,
@@ -151,15 +152,19 @@ async def _run_generation(request_id: str, kwargs: dict[str, Any]) -> None:
is synchronous) and update the store on completion or failure.
"""
generator = get_generator()
scheduler = get_video_batch_scheduler()
loop = asyncio.get_running_loop()
try:
start = time.perf_counter()
result = await loop.run_in_executor(
None,
lambda: generator.generate_video(**kwargs),
)
if scheduler is not None and scheduler.enabled:
result = await scheduler.submit(request_id, kwargs)
else:
result = await loop.run_in_executor(
None,
lambda: generator.generate_video(**kwargs),
)
elapsed = time.perf_counter() - start
update: dict[str, Any] = {
+512 -1
View File
@@ -18,6 +18,7 @@ import warnings
from collections.abc import Mapping
from contextlib import suppress
from copy import deepcopy
from dataclasses import dataclass
from typing import Any
import imageio
@@ -26,6 +27,8 @@ import torch
import torchvision
from einops import rearrange
from fastvideo.batching.admission import BatchAdmissionController
from fastvideo.batching.signature import can_dynamic_batch
from fastvideo.api.compat import (
expand_request_prompt_batch,
generator_config_to_fastvideo_args,
@@ -94,6 +97,11 @@ _FROM_PRETRAINED_CONVENIENCE_KWARGS = frozenset({
"pin_cpu_memory",
"enable_torch_compile",
"torch_compile_kwargs",
"batching_mode",
"batching_max_size",
"batching_delay_ms",
"batching_config",
"enable_batching_metrics",
"output_type",
"nvfp4_fa4",
})
@@ -112,6 +120,17 @@ def _infer_latent_batch_size(batch: ForwardBatch) -> int:
return latent_batch_size
@dataclass
class _GenerationWorkItem:
prompt: str
sampling_param: SamplingParam
fastvideo_args: FastVideoArgs
batch: ForwardBatch
output_path: str
target_height: int
target_width: int
class VideoGenerator:
"""
A unified class for generating videos using diffusion models.
@@ -439,6 +458,70 @@ class VideoGenerator:
if log_queue:
self.executor.clear_log_queue()
def generate_video_batch(self, request_kwargs: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Generate multiple legacy video requests, batching compatible items."""
work_items: list[_GenerationWorkItem] = []
reserved_output_paths: set[str] = set()
fastvideo_args_by_pipeline_override: dict[tuple[tuple[str, str], ...], FastVideoArgs] = {
(): self.fastvideo_args
}
for raw_kwargs in request_kwargs:
kwargs = dict(raw_kwargs)
prompt = kwargs.pop("prompt", None)
if prompt is None:
raise ValueError("Each batched generation request must include prompt")
if not isinstance(prompt, str):
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
sampling_param = kwargs.pop("sampling_param", None)
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(self.fastvideo_args.model_path)
else:
sampling_param = deepcopy(sampling_param)
extra_overrides: dict[str, Any] = {}
for _ek in _BATCH_EXTRA_PASSTHROUGH_KEYS:
if _ek in kwargs:
extra_overrides[_ek] = kwargs.pop(_ek)
request = legacy_generate_call_to_request(
prompt,
sampling_param,
legacy_kwargs=kwargs,
)
if not isinstance(request.prompt, str):
raise TypeError(f"`prompt` must be a string, but got {type(request.prompt)}")
fastvideo_args = self.fastvideo_args
pipeline_overrides = request_to_pipeline_overrides(request)
if pipeline_overrides:
override_key = tuple((key, repr(value)) for key, value in sorted(pipeline_overrides.items()))
fastvideo_args = fastvideo_args_by_pipeline_override.get(override_key)
if fastvideo_args is None:
fastvideo_args = deepcopy(self.fastvideo_args)
for key, value in pipeline_overrides.items():
if not hasattr(fastvideo_args.pipeline_config, key):
raise ValueError(f"Request field {key!r} is not supported by pipeline config overrides")
setattr(fastvideo_args.pipeline_config, key, deepcopy(value))
fastvideo_args_by_pipeline_override[override_key] = fastvideo_args
resolved_sampling_param = request_to_sampling_param(
request,
model_path=self.fastvideo_args.model_path,
)
output_path = self._prepare_output_path(resolved_sampling_param.output_path, request.prompt,
reserved_output_paths)
work_items.append(
self._prepare_generation_work_item(
prompt=request.prompt,
sampling_param=resolved_sampling_param,
fastvideo_args=fastvideo_args,
output_path=output_path,
_extra_overrides=extra_overrides,
))
return self._generate_prepared_work_items(work_items)
def _generate_request_impl(
self,
request: GenerationRequest,
@@ -535,6 +618,29 @@ class VideoGenerator:
logger.info("Found %d prompts in %s", len(prompts), prompt_txt_path)
if self._dynamic_batching_enabled(fastvideo_args):
work_items: list[_GenerationWorkItem] = []
reserved_output_paths: set[str] = set()
for batch_prompt in prompts:
item_kwargs = dict(kwargs)
item_kwargs["output_path"] = self._prepare_output_path(sampling_param.output_path, batch_prompt,
reserved_output_paths)
work_items.append(
self._prepare_generation_work_item(
prompt=batch_prompt,
sampling_param=sampling_param,
fastvideo_args=fastvideo_args,
**item_kwargs,
))
results = self._generate_prepared_work_items(work_items, tolerate_failures=True)
for i, (result, batch_prompt) in enumerate(zip(results, prompts, strict=True)):
result["prompt_index"] = i
result["prompt"] = batch_prompt
logger.info("Completed batch processing. Generated %d videos successfully.",
sum(1 for result in results if "error" not in result))
return results
results = []
for i, batch_prompt in enumerate(prompts):
logger.info("Processing prompt %d/%d: %s...", i + 1, len(prompts), batch_prompt[:100])
@@ -588,6 +694,7 @@ class VideoGenerator:
self,
output_path: str,
prompt: str,
reserved_paths: set[str] | None = None,
) -> str:
"""Build a unique, sanitized output file path.
@@ -602,6 +709,9 @@ class VideoGenerator:
- Invalid filename characters are removed; if the name changes, a
warning is logged.
- If the target path already exists, a numeric suffix is appended.
- ``reserved_paths`` lets batch callers resolve every path before any
file is written: paths in the set are treated as taken, and the
chosen path is added to the set.
"""
target_ext = ".png" if self._is_image_workload() else ".mp4"
@@ -646,15 +756,416 @@ class VideoGenerator:
if output_dir:
os.makedirs(output_dir, exist_ok=True)
def _is_taken(path: str) -> bool:
return os.path.exists(path) or (reserved_paths is not None and path in reserved_paths)
new_output_path = os.path.join(output_dir, out_name)
counter = 1
while os.path.exists(new_output_path):
while _is_taken(new_output_path):
name_part, ext_part = os.path.splitext(out_name)
new_name = f"{name_part}_{counter}{ext_part}"
new_output_path = os.path.join(output_dir, new_name)
counter += 1
if reserved_paths is not None:
reserved_paths.add(new_output_path)
return new_output_path
def _dynamic_batching_enabled(self, fastvideo_args: FastVideoArgs) -> bool:
batching_mode = getattr(fastvideo_args, "batching_mode", "disabled")
batching_max_size = getattr(fastvideo_args, "batching_max_size", 1)
return batching_mode == "dynamic" and batching_max_size > 1
def _prepare_generation_work_item(
self,
prompt: str | list[str],
sampling_param: SamplingParam,
fastvideo_args: FastVideoArgs,
**kwargs,
) -> _GenerationWorkItem:
if isinstance(prompt, str):
prompt_for_output = prompt.strip()
prompt_value: str | list[str] = prompt_for_output
elif isinstance(prompt, list) and all(isinstance(item, str) for item in prompt):
prompt_value = [item.strip() for item in prompt]
prompt_for_output = prompt_value[0] if prompt_value else ""
else:
raise TypeError(f"`prompt` must be a string or list of strings, but got {type(prompt)}")
sampling_param = deepcopy(sampling_param)
output_path = kwargs["output_path"]
sampling_param.prompt = prompt_value
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip()
if sampling_param.height <= 0 or sampling_param.width <= 0 or sampling_param.num_frames <= 0:
raise ValueError(f"Height, width, and num_frames must be positive integers, got "
f"height={sampling_param.height}, width={sampling_param.width}, "
f"num_frames={sampling_param.num_frames}")
target_height = align_to(sampling_param.height, 16)
target_width = align_to(sampling_param.width, 16)
latents_size = [(sampling_param.num_frames - 1) // 4 + 1, sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
debug_str = f"""
height: {target_height}
width: {target_width}
video_length: {sampling_param.num_frames}
prompt: {sampling_param.prompt}
image_path: {sampling_param.image_path}
neg_prompt: {sampling_param.negative_prompt}
seed: {sampling_param.seed}
infer_steps: {sampling_param.num_inference_steps}
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
guidance_scale: {sampling_param.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.pipeline_config.flow_shift}
embedded_guidance_scale: {fastvideo_args.pipeline_config.embedded_cfg_scale}
save_video: {sampling_param.save_video}
output_path: {output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
)
extra_overrides = kwargs.get("_extra_overrides", {})
for _ek, _ev in extra_overrides.items():
batch.extra[_ek] = _ev
return _GenerationWorkItem(
prompt=prompt_for_output,
sampling_param=sampling_param,
fastvideo_args=fastvideo_args,
batch=batch,
output_path=output_path,
target_height=target_height,
target_width=target_width,
)
def _run_forward_batch(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> tuple[ForwardBatch, float, float]:
start_time = time.perf_counter()
result_container = {"output_batch": ForwardBatch(data_type=batch.data_type)}
thread_error: dict[str, BaseException | None] = {"error": None}
thread_error_traceback: dict[str, str] = {"traceback": ""}
def execute_forward_thread():
import traceback
try:
result_container["output_batch"] = self.executor.execute_forward(batch, fastvideo_args)
except BaseException as error: # noqa: BLE001
thread_error["error"] = error
thread_error_traceback["traceback"] = traceback.format_exc()
thread = threading.Thread(target=execute_forward_thread)
thread.start()
thread.join()
if thread_error["error"] is not None:
raise RuntimeError("Forward execution thread failed.\n"
f"{thread_error_traceback['traceback']}") from thread_error["error"]
output_batch = result_container["output_batch"]
if output_batch.output is None:
raise RuntimeError("Forward execution returned no output tensor. "
"This usually means the executor/pipeline failed earlier.")
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
return output_batch, gen_time, start_time
def _samples_from_output(
self,
work_item: _GenerationWorkItem,
output_batch: ForwardBatch,
) -> torch.Tensor:
output = output_batch.output
if output is None:
raise RuntimeError("Forward execution returned no output tensor.")
fastvideo_args = work_item.fastvideo_args
sampling_param = work_item.sampling_param
latent_batch_size = _infer_latent_batch_size(work_item.batch)
skip_pixel_prealloc = fastvideo_args.output_type == "latent"
expected_shape = (
latent_batch_size,
3,
sampling_param.num_frames,
sampling_param.height,
sampling_param.width,
)
if skip_pixel_prealloc:
return output.cpu()
samples = torch.empty(expected_shape, device="cpu", pin_memory=fastvideo_args.pin_cpu_memory)
if output.shape == samples.shape:
samples.copy_(output)
return samples
logger.warning("Output shape %s does not match expected shape %s; use slow path", output.shape, samples.shape)
return output.cpu()
def _postprocess_generation_output(
self,
work_item: _GenerationWorkItem,
output_batch: ForwardBatch,
gen_time: float,
start_time: float,
) -> dict[str, Any]:
batch = work_item.batch
fastvideo_args = work_item.fastvideo_args
output_path = work_item.output_path
samples = self._samples_from_output(work_item, output_batch)
logging_info = output_batch.logging_info
is_latent_output = fastvideo_args.output_type == "latent"
audio_only = bool(output_batch.extra.get("audio_only"))
postprocess_start = time.perf_counter()
frames: list[np.ndarray] | None
if is_latent_output or audio_only:
frames = None if is_latent_output else []
else:
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.permute(1, 2, 0).squeeze(-1)
x = (x * 255).to(torch.uint8)
frames.append(x.contiguous().cpu().numpy())
postprocess_time = time.perf_counter() - postprocess_start
logger.info("PostDecodeFrameProcessStage completed in %.3f s", postprocess_time)
if logging_info is not None:
logging_info.add_stage_execution_time("PostDecodeFrameProcessStage", postprocess_time)
save_to_disk = batch.save_video and not is_latent_output
save_video_time = 0.0
audio_mux_time = 0.0
if save_to_disk:
if audio_only:
output_path = self._rewrite_extension(output_path, ".wav")
save_start = time.perf_counter()
self._write_pcm_wav(
output_path,
output_batch.extra["audio"],
int(output_batch.extra["audio_sample_rate"]),
)
save_video_time = time.perf_counter() - save_start
logger.info("Saved audio to %s", output_path)
elif self._is_image_workload():
assert frames is not None
save_start = time.perf_counter()
imageio.imwrite(output_path, frames[0])
save_video_time = time.perf_counter() - save_start
logger.info("Saved image to %s", output_path)
else:
assert frames is not None
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
if audio is not None and audio_sample_rate is not None:
save_start = time.perf_counter()
save_ok = self._save_video_with_audio_ffmpeg_pipe(
output_path=output_path,
frames=frames,
fps=batch.fps,
audio=audio,
sample_rate=int(audio_sample_rate),
)
if not save_ok:
logger.warning("ffmpeg pipe save failed; trying PyAV single-pass save.")
save_ok = self._save_video_with_audio_single_pass(
output_path=output_path,
frames=frames,
fps=batch.fps,
audio=audio,
sample_rate=int(audio_sample_rate),
)
save_video_time = time.perf_counter() - save_start
if save_ok:
audio_mux_time = 0.0
else:
logger.warning("Single-pass save failed; falling back to two-step save/mux.")
save_start = time.perf_counter()
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
save_video_time = time.perf_counter() - save_start
mux_start = time.perf_counter()
mux_ok = self._mux_audio(output_path, audio, int(audio_sample_rate))
audio_mux_time = time.perf_counter() - mux_start
if not mux_ok:
logger.warning("Audio mux failed; saved video without audio.")
else:
save_start = time.perf_counter()
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
save_video_time = time.perf_counter() - save_start
audio_mux_time = 0.0
logger.info("Saved video to %s", output_path)
logger.info("VideoSaveStage completed in %.3f s", save_video_time)
if logging_info is not None:
logging_info.add_stage_execution_time("VideoSaveStage", save_video_time)
logger.info("AudioMuxStage completed in %.3f s", audio_mux_time)
if logging_info is not None:
logging_info.add_stage_execution_time("AudioMuxStage", audio_mux_time)
e2e_time = time.perf_counter() - start_time
logger.info("End-to-end latency: %.2f seconds", e2e_time)
return {
"prompts": work_item.prompt,
"samples": samples if batch.return_frames else None,
"frames": frames if batch.return_frames else None,
"audio": output_batch.extra.get("audio"),
"audio_sample_rate": output_batch.extra.get("audio_sample_rate"),
"ltx2_audio_latents": output_batch.extra.get("ltx2_audio_latents"),
"size": (work_item.target_height, work_item.target_width, batch.num_frames),
"generation_time": gen_time,
"e2e_latency": e2e_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
"video_path": output_path if save_to_disk else None,
"peak_memory_mb": output_batch.extra.get("peak_memory_mb"),
}
def _split_output_batch(
self,
output_batch: ForwardBatch,
*,
index: int,
batch_size: int,
) -> ForwardBatch:
extra = {}
for key, value in (output_batch.extra or {}).items():
if torch.is_tensor(value) and value.ndim > 0 and value.shape[0] == batch_size:
extra[key] = value[index:index + 1]
elif isinstance(value, list) and len(value) == batch_size:
extra[key] = value[index]
else:
extra[key] = value
result = ForwardBatch(
data_type=output_batch.data_type,
output=(output_batch.output[index:index + 1] if output_batch.output is not None else None),
logging_info=output_batch.logging_info,
extra=extra,
)
if output_batch.trajectory_latents is not None:
result.trajectory_latents = output_batch.trajectory_latents[index:index + 1]
result.trajectory_timesteps = output_batch.trajectory_timesteps
if output_batch.trajectory_decoded is not None:
result.trajectory_decoded = [
decoded[index:index + 1] if torch.is_tensor(decoded) and decoded.shape[0] == batch_size else decoded
for decoded in output_batch.trajectory_decoded
]
return result
def _merge_work_items(self, work_items: list[_GenerationWorkItem]) -> _GenerationWorkItem:
first = work_items[0]
sampling_param = deepcopy(first.sampling_param)
prompts = [item.prompt for item in work_items]
sampling_param.prompt = prompts
sampling_param.seed = work_items[0].sampling_param.seed
merged = self._prepare_generation_work_item(
prompts,
sampling_param,
first.fastvideo_args,
output_path=first.output_path,
_extra_overrides=first.batch.extra,
)
merged.batch.seeds = [int(item.sampling_param.seed) for item in work_items]
merged.batch.extra["dynamic_batch_size"] = len(work_items)
merged.batch.extra["dynamic_batch_output_paths"] = [item.output_path for item in work_items]
return merged
def _can_merge_work_items(
self,
base: _GenerationWorkItem,
candidate: _GenerationWorkItem,
admission: BatchAdmissionController,
current_group: list[_GenerationWorkItem],
) -> bool:
if candidate.fastvideo_args is not base.fastvideo_args:
return False
compatibility = can_dynamic_batch(
base.sampling_param,
candidate.sampling_param,
base_extra=base.batch.extra,
candidate_extra=candidate.batch.extra,
)
if not compatibility.can_batch:
return False
current_requests = [item.sampling_param for item in current_group]
return admission.reject_reason_for_candidate(current_requests, candidate.sampling_param) is None
def _run_work_item_group(self, group: list[_GenerationWorkItem]) -> list[dict[str, Any]]:
if len(group) == 1:
return [self._execute_single_work_item(group[0])]
merged = self._merge_work_items(group)
output_batch, gen_time, start_time = self._run_forward_batch(merged.batch, merged.fastvideo_args)
return [
self._postprocess_generation_output(
item,
self._split_output_batch(output_batch, index=item_index, batch_size=len(group)),
gen_time,
start_time,
) for item_index, item in enumerate(group)
]
def _generate_prepared_work_items(
self,
work_items: list[_GenerationWorkItem],
tolerate_failures: bool = False,
) -> list[dict[str, Any]]:
"""Execute prepared work items, batching compatible neighbors.
With ``tolerate_failures`` (prompt-file semantics), a failed group
yields one ``{"error": ..., "prompt": ...}`` entry per work item so
completed results survive and stay aligned with the inputs; otherwise
the exception propagates.
"""
if not work_items:
return []
def run_group(group: list[_GenerationWorkItem]) -> list[dict[str, Any]]:
if not tolerate_failures:
return self._run_work_item_group(group)
try:
return self._run_work_item_group(group)
except Exception as e:
logger.error("Failed to generate videos for batched prompts %s: %s",
[item.prompt[:100] for item in group], e)
return [{"error": str(e), "prompt": item.prompt} for item in group]
fastvideo_args = work_items[0].fastvideo_args
if not self._dynamic_batching_enabled(fastvideo_args):
return [result for item in work_items for result in run_group([item])]
admission = BatchAdmissionController(fastvideo_args)
results: list[dict[str, Any]] = []
index = 0
while index < len(work_items):
group = [work_items[index]]
index += 1
while index < len(work_items) and len(group) < fastvideo_args.batching_max_size:
candidate = work_items[index]
if not self._can_merge_work_items(group[0], candidate, admission, group):
break
group.append(candidate)
index += 1
results.extend(run_group(group))
return results
def _execute_single_work_item(self, work_item: _GenerationWorkItem) -> dict[str, Any]:
output_batch, gen_time, start_time = self._run_forward_batch(work_item.batch, work_item.fastvideo_args)
return self._postprocess_generation_output(work_item, output_batch, gen_time, start_time)
def _generate_single_video(
self,
prompt: str,
+10
View File
@@ -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"),
+46
View File
@@ -169,6 +169,14 @@ class FastVideoArgs:
# Prompt text file for batch processing
prompt_txt: str | None = None
# Dynamic multimodal generation batching. Defaults preserve the historical
# one-request-at-a-time execution path.
batching_mode: str = "disabled"
batching_max_size: int = 1
batching_delay_ms: float = 0.0
batching_config: str | None = None
enable_batching_metrics: bool = False
# LTX-2 VAE tiling overrides
ltx2_vae_tiling: bool | None = None
ltx2_vae_spatial_tile_size_in_pixels: int | None = None
@@ -446,6 +454,37 @@ class FastVideoArgs:
default=FastVideoArgs.prompt_txt,
help="Path to a text file containing prompts (one per line) for batch processing",
)
parser.add_argument(
"--batching-mode",
type=str,
choices=["disabled", "dynamic"],
default=FastVideoArgs.batching_mode,
help="Request batching mode for inference serving.",
)
parser.add_argument(
"--batching-max-size",
type=int,
default=FastVideoArgs.batching_max_size,
help="Maximum number of compatible generation requests to execute as one batch.",
)
parser.add_argument(
"--batching-delay-ms",
type=float,
default=FastVideoArgs.batching_delay_ms,
help="Maximum queue delay in milliseconds before dispatching a dynamic batch.",
)
parser.add_argument(
"--batching-config",
type=str,
default=FastVideoArgs.batching_config,
help="Optional JSON batching admission rule file.",
)
parser.add_argument(
"--enable-batching-metrics",
action=StoreBoolean,
default=FastVideoArgs.enable_batching_metrics,
help="Log dynamic batching utilization and rejection metrics.",
)
# LTX-2 VAE tiling overrides
parser.add_argument(
@@ -753,6 +792,13 @@ class FastVideoArgs:
WorkloadType), f"Workload type must be a WorkloadType enum, got {type(self.workload_type)}"
assert self.workload_type in WorkloadType.choices(), f"Invalid workload type: {self.workload_type}"
if self.batching_mode not in {"disabled", "dynamic"}:
raise ValueError(f"batching_mode must be 'disabled' or 'dynamic', got {self.batching_mode!r}")
if self.batching_max_size < 1:
raise ValueError("batching_max_size must be >= 1")
if self.batching_delay_ms < 0:
raise ValueError("batching_delay_ms must be >= 0")
if self.mode in [ExecutionMode.DISTILLATION, ExecutionMode.FINETUNING] and self.inference_mode:
logger.warning("Mode is 'training' but inference_mode is True. Setting inference_mode to False.")
self.inference_mode = False
+2 -2
View File
@@ -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))
@@ -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))
+2
View File
@@ -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:
+130
View File
@@ -0,0 +1,130 @@
# 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,
)
+1 -2
View File
@@ -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 {
+2 -20
View File
@@ -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}
+41 -45
View File
@@ -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)
-1
View File
@@ -234,7 +234,6 @@ class DenoisingStage(PipelineStage):
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps if boundary_ratio is not None else None
latent_model_input = latents.to(target_dtype)
assert latent_model_input.shape[0] == 1, "only support batch size 1"
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
# TI2V directly replaces the first frame of the latent with
@@ -35,7 +35,11 @@ class InputValidationStage(PipelineStage):
num_videos_per_prompt = batch.num_videos_per_prompt
assert seed is not None
seeds = [seed + i for i in range(num_videos_per_prompt)]
if batch.seeds is not None:
seeds = batch.seeds
else:
prompt_count = len(batch.prompt) if isinstance(batch.prompt, list) else 1
seeds = [seed + i for i in range(prompt_count * num_videos_per_prompt)]
batch.seeds = seeds
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
+110 -13
View File
@@ -65,12 +65,20 @@ class TextEncodingStage(PipelineStage):
assert batch.prompt is not None
prompt_text: str | list[str] = batch.prompt
all_indices: list[int] = list(range(len(self.text_encoders)))
prompt_embeds_list, prompt_masks_list = self.encode_text(
prompt_text,
fastvideo_args,
encoder_index=all_indices,
return_attention_mask=True,
)
if isinstance(prompt_text, list):
prompt_embeds_list, prompt_masks_list = self._encode_prompt_list_individually(
prompt_text,
fastvideo_args,
encoder_index=all_indices,
return_attention_mask=True,
)
else:
prompt_embeds_list, prompt_masks_list = self.encode_text(
prompt_text,
fastvideo_args,
encoder_index=all_indices,
return_attention_mask=True,
)
if self._last_audio_embeds is not None:
batch.extra["ltx2_audio_prompt_embeds"] = self._last_audio_embeds
@@ -82,13 +90,24 @@ class TextEncodingStage(PipelineStage):
# Encode negative prompt if CFG is enabled
if batch.do_classifier_free_guidance:
assert isinstance(batch.negative_prompt, str)
neg_embeds_list, neg_masks_list = self.encode_text(
batch.negative_prompt,
fastvideo_args,
encoder_index=all_indices,
return_attention_mask=True,
)
assert isinstance(batch.negative_prompt, str | list)
negative_prompt: str | list[str] = batch.negative_prompt
if isinstance(batch.prompt, list) and isinstance(negative_prompt, str):
negative_prompt = [negative_prompt] * len(batch.prompt)
if isinstance(negative_prompt, list):
neg_embeds_list, neg_masks_list = self._encode_prompt_list_individually(
negative_prompt,
fastvideo_args,
encoder_index=all_indices,
return_attention_mask=True,
)
else:
neg_embeds_list, neg_masks_list = self.encode_text(
negative_prompt,
fastvideo_args,
encoder_index=all_indices,
return_attention_mask=True,
)
if self._last_audio_embeds is not None:
batch.extra["ltx2_audio_negative_embeds"] = self._last_audio_embeds
@@ -101,6 +120,81 @@ class TextEncodingStage(PipelineStage):
return batch
def _encode_prompt_list_individually(
self,
texts: list[str],
fastvideo_args: FastVideoArgs,
*,
encoder_index: list[int],
return_attention_mask: bool,
) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
per_prompt_embeds: list[list[torch.Tensor]] = []
per_prompt_masks: list[list[torch.Tensor]] = []
per_prompt_audio_embeds: list[list[torch.Tensor] | None] = []
for text in texts:
embeds, masks = self.encode_text(
text,
fastvideo_args,
encoder_index=encoder_index,
return_attention_mask=return_attention_mask,
)
per_prompt_embeds.append(embeds)
per_prompt_masks.append(masks)
per_prompt_audio_embeds.append(self._last_audio_embeds)
merged_embeds = [
self._cat_tensors([prompt_embeds[encoder_pos] for prompt_embeds in per_prompt_embeds])
for encoder_pos in range(len(per_prompt_embeds[0]))
]
merged_masks = [
self._cat_attention_masks([prompt_masks[encoder_pos] for prompt_masks in per_prompt_masks])
for encoder_pos in range(len(per_prompt_masks[0]))
]
if per_prompt_audio_embeds and all(audio_embeds is not None for audio_embeds in per_prompt_audio_embeds):
audio_embed_lists = [audio_embeds for audio_embeds in per_prompt_audio_embeds if audio_embeds is not None]
self._last_audio_embeds = [
self._cat_tensors([audio_embeds[encoder_pos] for audio_embeds in audio_embed_lists])
for encoder_pos in range(len(audio_embed_lists[0]))
]
else:
self._last_audio_embeds = None
return merged_embeds, merged_masks
@staticmethod
def _cat_tensors(tensors: list[torch.Tensor]) -> torch.Tensor:
base_shape = tensors[0].shape[1:]
if all(tensor.shape[1:] == base_shape for tensor in tensors):
return torch.cat(tensors, dim=0)
if all(tensor.ndim == 3 for tensor in tensors):
base_trailing_shape = tensors[0].shape[2:]
if all(tensor.shape[2:] == base_trailing_shape for tensor in tensors):
max_length = max(tensor.shape[1] for tensor in tensors)
padded_tensors = []
for tensor in tensors:
pad_width = max_length - tensor.shape[1]
if pad_width > 0:
tensor = torch.nn.functional.pad(tensor, (0, 0, 0, pad_width), value=0.0)
padded_tensors.append(tensor)
return torch.cat(padded_tensors, dim=0)
raise ValueError(f"Cannot concatenate tensors with shapes: {[list(tensor.shape) for tensor in tensors]}")
@staticmethod
def _cat_attention_masks(masks: list[torch.Tensor]) -> torch.Tensor:
base_shape = masks[0].shape[1:]
if all(mask.shape[1:] == base_shape for mask in masks):
return torch.cat(masks, dim=0)
if all(mask.ndim == 2 for mask in masks):
max_length = max(mask.shape[1] for mask in masks)
padded_masks = []
for mask in masks:
pad_width = max_length - mask.shape[1]
if pad_width > 0:
mask = torch.nn.functional.pad(mask, (0, pad_width), value=0)
padded_masks.append(mask)
return torch.cat(padded_masks, dim=0)
raise ValueError(f"Cannot concatenate attention masks with shapes: {[list(mask.shape) for mask in masks]}")
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify text encoding stage inputs."""
result = VerificationResult()
@@ -235,6 +329,9 @@ class TextEncodingStage(PipelineStage):
attn_masks_list.append(attention_mask)
return self.return_embeds(embeds_list, attn_masks_list, return_type, return_attention_mask, indices)
if len(processed_texts) > 1 and "padding" not in tok_kwargs:
tok_kwargs["padding"] = True
# If tokenizer is a multimodal processor (e.g. Qwen2_5_VLProcessor),
# use its inner tokenizer for text-only encoding.
tok = getattr(tokenizer, "tokenizer", tokenizer)
+50 -1
View File
@@ -9,7 +9,7 @@ from fastvideo.api.compat import (
generator_config_to_fastvideo_args,
legacy_from_pretrained_to_config,
)
from fastvideo.api.schema import CompileConfig, GeneratorConfig
from fastvideo.api.schema import BatchingConfig, CompileConfig, GeneratorConfig
class TestLegacyTorchCompileKwargsTranslation:
@@ -200,6 +200,48 @@ class TestLegacyTextEncoderCompileTranslation:
assert "enable_torch_compile_text_encoder" not in args.kwargs
class TestBatchingTranslation:
def test_flat_kwargs_promote_to_engine_batching(self) -> None:
config = legacy_from_pretrained_to_config(
"/models/wan",
{
"batching_mode": "dynamic",
"batching_max_size": 4,
"batching_delay_ms": 25.0,
"batching_config": "/tmp/batching.json",
"enable_batching_metrics": True,
},
)
assert config.engine.batching.mode == "dynamic"
assert config.engine.batching.max_size == 4
assert config.engine.batching.delay_ms == 25.0
assert config.engine.batching.config_path == "/tmp/batching.json"
assert config.engine.batching.enable_metrics is True
def test_typed_batching_emits_fastvideo_args_kwargs(self, monkeypatch) -> None:
_stub_fastvideo_args_from_kwargs(monkeypatch)
config = GeneratorConfig(
model_path="/models/wan",
engine=_engine_with_batching(BatchingConfig(
mode="dynamic",
max_size=3,
delay_ms=10.0,
config_path="/tmp/batching.json",
enable_metrics=True,
)),
)
args = generator_config_to_fastvideo_args(config)
assert args.kwargs["batching_mode"] == "dynamic"
assert args.kwargs["batching_max_size"] == 3
assert args.kwargs["batching_delay_ms"] == 10.0
assert args.kwargs["batching_config"] == "/tmp/batching.json"
assert args.kwargs["enable_batching_metrics"] is True
# -------------------------------------------------------------------
# Helpers
# -------------------------------------------------------------------
@@ -213,6 +255,13 @@ def _engine_with_compile(compile_config):
return engine
def _engine_with_batching(batching_config):
from fastvideo.api.schema import EngineConfig
engine = EngineConfig()
engine.batching = batching_config
return engine
def _stub_fastvideo_args_from_kwargs(monkeypatch):
"""Swap ``FastVideoArgs.from_kwargs`` for a capture-only stub so
translation tests don't need to construct a valid FastVideoArgs."""
+7
View File
@@ -130,6 +130,13 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
"use_fsdp_inference": False,
"disable_autocast": False,
"quantization": None,
"batching": {
"mode": "disabled",
"max_size": 1,
"delay_ms": 0.0,
"config_path": None,
"enable_metrics": False,
},
},
"pipeline": {
"workload_type": None,
@@ -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,198 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
import json
import os
import time
from pathlib import Path
from typing import Any
import torch
from fastvideo import VideoGenerator
DEFAULT_PROMPTS = (
"A small robot sketches a city skyline at sunrise, cinematic lighting.",
"A glass teapot steams on a wooden table while rain falls outside.",
)
def _build_init_kwargs(args: argparse.Namespace, *, dynamic: bool) -> dict[str, Any]:
return {
"num_gpus": args.num_gpus,
"sp_size": args.sp_size,
"tp_size": args.tp_size,
"use_fsdp_inference": args.use_fsdp_inference,
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
"flow_shift": args.flow_shift,
"text_encoder_precisions": ("fp32",),
"output_type": "latent",
"batching_mode": "dynamic" if dynamic else "disabled",
"batching_max_size": args.batch_size if dynamic else 1,
"batching_delay_ms": 0.0,
}
def _request_kwargs(args: argparse.Namespace, prompt_index: int) -> dict[str, Any]:
return {
"prompt": args.prompts[prompt_index],
"height": args.height,
"width": args.width,
"num_frames": args.num_frames,
"num_inference_steps": args.num_inference_steps,
"guidance_scale": args.guidance_scale,
"embedded_cfg_scale": args.embedded_cfg_scale,
"seed": args.seed + prompt_index,
"fps": 24,
"save_video": False,
"return_frames": True,
"output_path": str(Path(args.output_dir) / f"request_{prompt_index}.mp4"),
}
def _sync() -> None:
if torch.cuda.is_available():
torch.cuda.synchronize()
def _run_sequential(generator: VideoGenerator, args: argparse.Namespace) -> tuple[list[dict[str, Any]], float]:
_sync()
start = time.perf_counter()
results = []
for index in range(args.batch_size):
kwargs = _request_kwargs(args, index)
prompt = kwargs.pop("prompt")
results.append(generator.generate_video(prompt=prompt, **kwargs))
_sync()
return results, time.perf_counter() - start
def _run_dynamic(generator: VideoGenerator, args: argparse.Namespace) -> tuple[list[dict[str, Any]], float]:
if not hasattr(generator, "generate_video_batch"):
raise RuntimeError("VideoGenerator.generate_video_batch is unavailable in this checkout")
requests = [_request_kwargs(args, index) for index in range(args.batch_size)]
_sync()
start = time.perf_counter()
results = generator.generate_video_batch(requests)
_sync()
return results, time.perf_counter() - start
def _tensor_metrics(sequential: list[dict[str, Any]], dynamic: list[dict[str, Any]]) -> dict[str, Any]:
per_request = []
for index, (seq_result, dyn_result) in enumerate(zip(sequential, dynamic, strict=True)):
seq = seq_result["samples"].detach().cpu().to(torch.float32)
dyn = dyn_result["samples"].detach().cpu().to(torch.float32)
diff = (seq - dyn).abs()
per_request.append({
"index": index,
"shape": list(seq.shape),
"max_abs_diff": float(diff.max().item()),
"mean_abs_diff": float(diff.mean().item()),
"allclose_atol_1e_5": bool(torch.allclose(seq, dyn, atol=1e-5, rtol=1e-5)),
"allclose_atol_1e_4": bool(torch.allclose(seq, dyn, atol=1e-4, rtol=1e-4)),
})
return {
"per_request": per_request,
"max_abs_diff": max(item["max_abs_diff"] for item in per_request),
"mean_abs_diff": sum(item["mean_abs_diff"] for item in per_request) / len(per_request),
"allclose_atol_1e_5": all(item["allclose_atol_1e_5"] for item in per_request),
"allclose_atol_1e_4": all(item["allclose_atol_1e_4"] for item in per_request),
}
def run_parity(args: argparse.Namespace) -> dict[str, Any]:
generator = VideoGenerator.from_pretrained(args.model_path, **_build_init_kwargs(args, dynamic=True))
try:
sequential, sequential_s = _run_sequential(generator, args)
dynamic, dynamic_s = _run_dynamic(generator, args)
metrics = _tensor_metrics(sequential, dynamic)
finally:
generator.shutdown()
return {
"mode": "parity",
"model_path": args.model_path,
"num_gpus": args.num_gpus,
"shape": {
"height": args.height,
"width": args.width,
"num_frames": args.num_frames,
"num_inference_steps": args.num_inference_steps,
},
"batch_size": args.batch_size,
"sequential_time_s": sequential_s,
"dynamic_time_s": dynamic_s,
"speedup": sequential_s / dynamic_s if dynamic_s > 0 else None,
"tensor_metrics": metrics,
}
def run_benchmark(args: argparse.Namespace, *, dynamic: bool) -> dict[str, Any]:
generator = VideoGenerator.from_pretrained(args.model_path, **_build_init_kwargs(args, dynamic=dynamic))
run = _run_dynamic if dynamic else _run_sequential
try:
for _ in range(args.warmup_runs):
run(generator, args)
times = []
for _ in range(args.measurement_runs):
_results, elapsed = run(generator, args)
times.append(elapsed)
finally:
generator.shutdown()
avg = sum(times) / len(times)
return {
"mode": "dynamic" if dynamic else "sequential",
"model_path": args.model_path,
"num_gpus": args.num_gpus,
"batch_size": args.batch_size,
"measurement_runs": args.measurement_runs,
"times_s": times,
"avg_time_s": avg,
"requests_per_second": args.batch_size / avg if avg > 0 else None,
}
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--mode", choices=("parity", "sequential", "dynamic"), required=True)
parser.add_argument("--model-path", default="Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--sp-size", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=1)
parser.add_argument("--use-fsdp-inference", action="store_true")
parser.add_argument("--height", type=int, default=256)
parser.add_argument("--width", type=int, default=256)
parser.add_argument("--num-frames", type=int, default=9)
parser.add_argument("--num-inference-steps", type=int, default=2)
parser.add_argument("--guidance-scale", type=float, default=1.0)
parser.add_argument("--embedded-cfg-scale", type=float, default=6.0)
parser.add_argument("--flow-shift", type=float, default=7.0)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--batch-size", type=int, default=2)
parser.add_argument("--warmup-runs", type=int, default=1)
parser.add_argument("--measurement-runs", type=int, default=3)
parser.add_argument("--output-dir", default="/tmp/fastvideo_dynamic_batching")
parser.add_argument("--output-json", required=True)
parser.add_argument("--prompts", nargs="+", default=list(DEFAULT_PROMPTS))
return parser.parse_args()
def main() -> None:
args = parse_args()
if len(args.prompts) < args.batch_size:
raise ValueError("--prompts must contain at least --batch-size prompts")
os.makedirs(args.output_dir, exist_ok=True)
if args.mode == "parity":
result = run_parity(args)
else:
result = run_benchmark(args, dynamic=args.mode == "dynamic")
output_path = Path(args.output_json)
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(json.dumps(result, indent=2), encoding="utf-8")
print(json.dumps(result, indent=2))
if __name__ == "__main__":
main()
@@ -0,0 +1,89 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from types import SimpleNamespace
import pytest
from fastvideo.batching.admission import (
AdmissionLimit,
BatchAdmissionController,
BatchingRule,
load_batching_config,
)
from fastvideo.configs.pipelines.base import PipelineConfig
def test_admission_limit_rejects_batch_size_and_cost() -> None:
limit = AdmissionLimit(max_batch_size=2, max_cost=10.0)
assert limit.reject_reason(batch_size=3, batch_cost=1.0) == "config_cap:2"
assert limit.reject_reason(batch_size=2, batch_cost=11.0) == "cost_budget:11>10"
assert limit.reject_reason(batch_size=2, batch_cost=10.0) is None
def test_batching_rule_validates_unknown_keys() -> None:
with pytest.raises(ValueError, match="did you mean 'max_batch_size'"):
BatchingRule.from_dict(
{
"model_contains": "wan",
"max_batch_siz": 2,
},
source="unit",
)
@pytest.mark.parametrize(("value", "expected"), [(1, True), (0, False), (1.0, True), (0.0, False)])
def test_batching_rule_parses_numeric_bool_values(value, expected) -> None:
rule = BatchingRule.from_dict(
{
"model_contains": "wan",
"offload": value,
"max_batch_size": 2,
},
source="unit",
)
assert rule.offload is expected
def test_load_batching_config_supports_mapping_form(tmp_path) -> None:
path = tmp_path / "batching.json"
path.write_text(
'{"schema_version": 1, "wan|720x1280x81": {"max_batch_size": 3, "max_cost": 9}}',
encoding="utf-8",
)
rules = load_batching_config(str(path))
assert len(rules) == 1
assert rules[0].model == "wan"
assert rules[0].resolution == "720x1280x81"
assert rules[0].max_batch_size == 3
assert rules[0].max_cost == 9.0
def test_admission_controller_applies_user_and_config_caps(tmp_path, monkeypatch) -> None:
path = tmp_path / "batching.json"
path.write_text(
'{"rules": [{"model_contains": "wan", "resolution": "720x1280x81", "max_batch_size": 3}]}',
encoding="utf-8",
)
monkeypatch.setattr(BatchAdmissionController, "_get_device_memory_gb", staticmethod(lambda gpu_id: 48.0))
args = SimpleNamespace(
batching_mode="dynamic",
batching_max_size=4,
batching_config=str(path),
model_path="/models/wan",
dit_cpu_offload=False,
dit_layerwise_offload=False,
pipeline_config=PipelineConfig(),
)
request = SimpleNamespace(height=720, width=1280, num_frames=81)
controller = BatchAdmissionController(args)
assert controller.enabled is True
assert controller.max_admissible_batch_size(request) == 3
assert controller.batch_is_full([request, request, request]) is True
@@ -0,0 +1,57 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.batching.signature import (
can_dynamic_batch,
dynamic_batch_signature,
resolution_key,
)
def _request(prompt: str = "a prompt", **overrides) -> SamplingParam:
request = SamplingParam(prompt=prompt, height=256, width=384, num_frames=17, num_inference_steps=4)
for key, value in overrides.items():
setattr(request, key, value)
return request
def test_dynamic_batch_signature_excludes_request_local_fields() -> None:
first = _request(seed=1, output_path="/tmp/a.mp4", save_video=True, return_frames=False)
second = _request(seed=2, output_path="/tmp/b.mp4", save_video=False, return_frames=True)
assert dynamic_batch_signature(first) == dynamic_batch_signature(second)
def test_can_dynamic_batch_accepts_matching_text_requests() -> None:
first = _request("first", seed=1)
second = _request("second", seed=2)
result = can_dynamic_batch(first, second)
assert result.can_batch is True
assert result.reason is None
def test_can_dynamic_batch_rejects_sampling_mismatch() -> None:
first = _request(guidance_scale=1.0)
second = _request(guidance_scale=3.0)
result = can_dynamic_batch(first, second)
assert result.can_batch is False
assert result.reason == "sampling_params.guidance_scale"
def test_can_dynamic_batch_rejects_image_conditioning() -> None:
first = _request()
second = _request(image_path="/tmp/image.png")
result = can_dynamic_batch(first, second)
assert result.can_batch is False
assert result.reason == "image_path"
def test_resolution_key_uses_generation_shape() -> None:
assert resolution_key(_request(height=720, width=1280, num_frames=81)) == "720x1280x81"
@@ -1,10 +1,15 @@
"""Unit tests for the OpenAI-compatible API server helpers (no GPU needed)."""
import asyncio
import os
import time
from types import SimpleNamespace
from unittest.mock import patch
import pytest
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.entrypoints.openai.batching import VideoBatchScheduler, _VideoBatchJob
from fastvideo.api.parser import parse_config
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.openai.protocol import (
@@ -21,6 +26,172 @@ from fastvideo.entrypoints.openai.utils import (
parse_size,
)
class _FakeBatchGenerator:
def __init__(self):
self.calls = []
def generate_video_batch(self, request_kwargs):
self.calls.append([dict(item) for item in request_kwargs])
return [{"prompts": item["prompt"], "video_path": item["output_path"]} for item in request_kwargs]
def _make_batch_job(request_id, kwargs):
loop = asyncio.get_running_loop()
return _VideoBatchJob(
request_id=request_id,
kwargs=dict(kwargs),
future=loop.create_future(),
enqueue_time=time.perf_counter(),
)
def _batch_scheduler_args(**overrides):
defaults = dict(
model_path="test-model",
batching_mode="dynamic",
batching_max_size=2,
batching_delay_ms=25.0,
enable_batching_metrics=False,
pipeline_config=PipelineConfig(),
)
defaults.update(overrides)
return SimpleNamespace(**defaults)
def test_video_batch_scheduler_groups_compatible_requests(tmp_path):
async def run():
generator = _FakeBatchGenerator()
scheduler = VideoBatchScheduler(generator, _batch_scheduler_args())
await scheduler.start()
try:
first = {
"prompt": "first",
"height": 256,
"width": 256,
"num_frames": 1,
"num_inference_steps": 2,
"seed": 1,
"output_path": str(tmp_path / "first.mp4"),
"save_video": False,
}
second = {
"prompt": "second",
"height": 256,
"width": 256,
"num_frames": 1,
"num_inference_steps": 2,
"seed": 2,
"output_path": str(tmp_path / "second.mp4"),
"save_video": False,
}
results = await asyncio.gather(
scheduler.submit("req-1", first),
scheduler.submit("req-2", second),
)
finally:
await scheduler.stop()
return generator.calls, results
calls, results = asyncio.run(run())
assert len(calls) == 1
assert [item["prompt"] for item in calls[0]] == ["first", "second"]
assert [result["prompts"] for result in results] == ["first", "second"]
def test_video_batch_scheduler_keeps_incompatible_requests_separate(tmp_path):
async def run():
generator = _FakeBatchGenerator()
scheduler = VideoBatchScheduler(generator, _batch_scheduler_args())
await scheduler.start()
try:
text_only = {
"prompt": "first",
"height": 256,
"width": 256,
"num_frames": 1,
"num_inference_steps": 2,
"seed": 1,
"output_path": str(tmp_path / "first.mp4"),
"save_video": False,
}
image_conditioned = {
"prompt": "second",
"height": 256,
"width": 256,
"num_frames": 1,
"num_inference_steps": 2,
"seed": 2,
"image_path": str(tmp_path / "input.png"),
"output_path": str(tmp_path / "second.mp4"),
"save_video": False,
}
results = await asyncio.gather(
scheduler.submit("req-1", text_only),
scheduler.submit("req-2", image_conditioned),
)
finally:
await scheduler.stop()
return generator.calls, results
calls, results = asyncio.run(run())
assert len(calls) == 2
assert [[item["prompt"] for item in call] for call in calls] == [["first"], ["second"]]
assert [result["prompts"] for result in results] == ["first", "second"]
def test_video_batch_scheduler_requeues_incompatible_pending_job_at_front(tmp_path):
async def run():
generator = _FakeBatchGenerator()
scheduler = VideoBatchScheduler(generator, _batch_scheduler_args())
first = {
"prompt": "first",
"height": 256,
"width": 256,
"num_frames": 1,
"num_inference_steps": 2,
"seed": 1,
"output_path": str(tmp_path / "first.mp4"),
"save_video": False,
}
incompatible = {
"prompt": "second",
"height": 256,
"width": 256,
"num_frames": 1,
"num_inference_steps": 2,
"seed": 2,
"image_path": str(tmp_path / "input.png"),
"output_path": str(tmp_path / "second.mp4"),
"save_video": False,
}
newer = {
"prompt": "third",
"height": 256,
"width": 256,
"num_frames": 1,
"num_inference_steps": 2,
"seed": 3,
"output_path": str(tmp_path / "third.mp4"),
"save_video": False,
}
scheduler._pending.extend([
_make_batch_job("req-2", incompatible),
_make_batch_job("req-3", newer),
])
batch = await scheduler._collect_batch(_make_batch_job("req-1", first))
return [job.request_id for job in batch], [job.request_id for job in scheduler._pending]
batch_ids, pending_ids = asyncio.run(run())
assert batch_ids == ["req-1"]
assert pending_ids == ["req-2", "req-3"]
# ---------------------------------------------------------------------------
# parse_size
# ---------------------------------------------------------------------------
@@ -3,6 +3,7 @@ from types import SimpleNamespace
import warnings
import pytest
import torch
from fastvideo.api import (
GenerationRequest,
@@ -13,8 +14,10 @@ from fastvideo.api import (
load_run_config,
)
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.pipelines import ForwardBatch
def _new_video_generator() -> VideoGenerator:
@@ -37,6 +40,25 @@ def _new_runtime_video_generator() -> VideoGenerator:
return generator
def _batching_fastvideo_args(**overrides):
defaults = dict(
model_path="test-model",
prompt_txt=None,
workload_type=SimpleNamespace(value="t2v"),
batching_mode="dynamic",
batching_max_size=4,
batching_config=None,
dit_cpu_offload=False,
dit_layerwise_offload=False,
output_type="latent",
pin_cpu_memory=False,
VSA_sparsity=0.0,
pipeline_config=PipelineConfig(),
)
defaults.update(overrides)
return SimpleNamespace(**defaults)
def _patch_from_fastvideo_args(monkeypatch):
captured = {}
@@ -151,6 +173,117 @@ def test_prepare_output_path_empty_prompt_fallback(tmp_path):
assert os.path.basename(result) == "output.mp4"
def test_generate_prepared_work_items_merges_compatible_latent_requests(monkeypatch, tmp_path):
vg = _new_video_generator()
vg.fastvideo_args = _batching_fastvideo_args()
calls = []
def fake_device_memory(gpu_id):
return 48.0
def fake_run_forward(batch, fastvideo_args):
calls.append(batch)
batch_size = len(batch.prompt) if isinstance(batch.prompt, list) else 1
output = torch.arange(batch_size * 4, dtype=torch.float32).reshape(batch_size, 4, 1, 1, 1)
return ForwardBatch(data_type=batch.data_type, output=output, extra={"peak_memory_mb": 1.0}), 0.5, 10.0
monkeypatch.setattr(
"fastvideo.batching.admission.BatchAdmissionController._get_device_memory_gb",
staticmethod(fake_device_memory),
)
monkeypatch.setattr(vg, "_run_forward_batch", fake_run_forward)
first = SamplingParam(prompt="one", height=8, width=8, num_frames=1, seed=11, return_frames=True, save_video=False)
second = SamplingParam(prompt="two", height=8, width=8, num_frames=1, seed=22, return_frames=True, save_video=False)
work_items = [
vg._prepare_generation_work_item("one", first, vg.fastvideo_args, output_path=str(tmp_path / "one.mp4")),
vg._prepare_generation_work_item("two", second, vg.fastvideo_args, output_path=str(tmp_path / "two.mp4")),
]
results = vg._generate_prepared_work_items(work_items)
assert len(calls) == 1
assert calls[0].prompt == ["one", "two"]
assert calls[0].seeds == [11, 22]
assert [result["prompts"] for result in results] == ["one", "two"]
assert [result["samples"].shape for result in results] == [(1, 4, 1, 1, 1), (1, 4, 1, 1, 1)]
def test_generate_prepared_work_items_falls_back_for_incompatible_requests(monkeypatch, tmp_path):
vg = _new_video_generator()
vg.fastvideo_args = _batching_fastvideo_args()
calls = []
def fake_run_forward(batch, fastvideo_args):
calls.append(batch)
output = torch.zeros((1, 4, 1, 1, 1), dtype=torch.float32)
return ForwardBatch(data_type=batch.data_type, output=output), 0.5, 10.0
monkeypatch.setattr(vg, "_run_forward_batch", fake_run_forward)
first = SamplingParam(prompt="one", height=8, width=8, num_frames=1, guidance_scale=1.0, save_video=False)
second = SamplingParam(prompt="two", height=8, width=8, num_frames=1, guidance_scale=3.0, save_video=False)
work_items = [
vg._prepare_generation_work_item("one", first, vg.fastvideo_args, output_path=str(tmp_path / "one.mp4")),
vg._prepare_generation_work_item("two", second, vg.fastvideo_args, output_path=str(tmp_path / "two.mp4")),
]
results = vg._generate_prepared_work_items(work_items)
assert len(calls) == 2
assert all(isinstance(call.prompt, str) for call in calls)
assert [result["prompts"] for result in results] == ["one", "two"]
def test_generate_video_batch_routes_compat_kwargs(monkeypatch, tmp_path):
vg = _new_video_generator()
vg.fastvideo_args = _batching_fastvideo_args()
calls = []
def fake_device_memory(gpu_id):
return 48.0
def fake_run_forward(batch, fastvideo_args):
calls.append((batch, fastvideo_args))
output = torch.zeros((len(batch.prompt), 4, 1, 1, 1), dtype=torch.float32)
return ForwardBatch(data_type=batch.data_type, output=output), 0.5, 10.0
monkeypatch.setattr(
"fastvideo.batching.admission.BatchAdmissionController._get_device_memory_gb",
staticmethod(fake_device_memory),
)
monkeypatch.setattr(vg, "_run_forward_batch", fake_run_forward)
results = vg.generate_video_batch([
{
"prompt": "one",
"height": 8,
"width": 8,
"num_frames": 1,
"embedded_cfg_scale": 7.5,
"save_video": False,
"return_frames": True,
"output_path": str(tmp_path / "one.mp4"),
},
{
"prompt": "two",
"height": 8,
"width": 8,
"num_frames": 1,
"embedded_cfg_scale": 7.5,
"save_video": False,
"return_frames": True,
"output_path": str(tmp_path / "two.mp4"),
},
])
assert len(calls) == 1
batch, fastvideo_args = calls[0]
assert batch.prompt == ["one", "two"]
assert fastvideo_args.pipeline_config.embedded_cfg_scale == 7.5
assert [result["prompts"] for result in results] == ["one", "two"]
def test_from_config_normalizes_and_translates(monkeypatch):
captured = _patch_from_fastvideo_args(monkeypatch)
_patch_fastvideo_args_from_kwargs(monkeypatch)
@@ -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))
+9 -7
View File
@@ -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
+11 -8
View File
@@ -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",
}))
@@ -273,7 +276,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/batching/ ./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"
)
+9 -7
View File
@@ -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"),
}
)
)
+92 -75
View File
@@ -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:
+8 -1
View File
@@ -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
@@ -219,6 +296,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 +309,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 +354,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
@@ -0,0 +1,40 @@
# SPDX-License-Identifier: Apache-2.0
from types import SimpleNamespace
from fastvideo.pipelines import ForwardBatch
from fastvideo.pipelines.stages.input_validation import InputValidationStage
def test_input_validation_preserves_explicit_dynamic_batch_seeds() -> None:
batch = ForwardBatch(
data_type="video",
prompt=["one", "two"],
seed=100,
seeds=[17, 23],
height=8,
width=8,
num_frames=1,
num_inference_steps=1,
)
InputValidationStage()._generate_seeds(batch, SimpleNamespace())
assert batch.seeds == [17, 23]
assert [generator.initial_seed() for generator in batch.generator] == [17, 23]
def test_input_validation_generates_one_seed_per_prompt() -> None:
batch = ForwardBatch(
data_type="video",
prompt=["one", "two"],
seed=100,
height=8,
width=8,
num_frames=1,
num_inference_steps=1,
)
InputValidationStage()._generate_seeds(batch, SimpleNamespace())
assert batch.seeds == [100, 101]
assert [generator.initial_seed() for generator in batch.generator] == [100, 101]
@@ -13,7 +13,13 @@ class TensorDict(dict):
return TensorDict({k: v.to(device) for k, v in self.items()})
class FakeTokenizer:
def __init__(self):
self.calls = []
self.texts = []
def __call__(self, texts, **kwargs):
self.calls.append(kwargs)
self.texts.append(list(texts))
B = len(texts)
seq_len = int(kwargs.get("max_length", 4))
return TensorDict({
@@ -21,6 +27,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 +64,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)
@@ -131,6 +165,66 @@ def test_forward_integration_cfg_off_and_on():
assert len(out2.prompt_attention_mask) == 2
assert len(out2.negative_attention_mask) == 2
def test_encode_text_adds_padding_for_prompt_lists():
fastvideo_args, hidden = make_args(num_encoders=1, text_len=4, hidden_size=8)
stage = make_stage(num_encoders=1, hidden_size=hidden)
stage.encode_text(["short", "a longer prompt"], fastvideo_args, encoder_index=[0])
assert stage.tokenizers[0].calls[-1]["padding"] is True
def test_forward_prompt_list_preserves_single_prompt_text_encoding_path():
fastvideo_args, hidden = make_args(num_encoders=1, text_len=4, hidden_size=8)
stage = make_stage(num_encoders=1, hidden_size=hidden)
batch = ForwardBatch(
data_type="video",
prompt=["short", "a longer prompt"],
negative_prompt="",
do_classifier_free_guidance=False,
prompt_embeds=[],
negative_prompt_embeds=None,
prompt_attention_mask=[],
negative_attention_mask=None,
)
out = stage.forward(batch, fastvideo_args)
assert stage.tokenizers[0].texts == [["short"], ["a longer prompt"]]
assert out.prompt_embeds[0].shape == (2, hidden)
assert out.prompt_attention_mask[0].shape == (2, 4)
def test_encode_prompt_list_individually_pads_variable_length_embeds_and_audio():
fastvideo_args, _hidden = make_args(num_encoders=1, text_len=4, hidden_size=8)
stage = TextEncodingStage(text_encoders=[], tokenizers=[])
lengths = {"short": 2, "a longer prompt": 4}
def fake_encode_text(text, *_args, **_kwargs):
length = lengths[text]
embeds = [torch.full((1, length, 3), fill_value=float(length))]
masks = [torch.ones((1, length), dtype=torch.long)]
stage._last_audio_embeds = [torch.full((1, length, 5), fill_value=float(length))]
return embeds, masks
stage.encode_text = fake_encode_text
embeds, masks = stage._encode_prompt_list_individually(
["short", "a longer prompt"],
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
assert embeds[0].shape == (2, 4, 3)
assert masks[0].shape == (2, 4)
assert stage._last_audio_embeds is not None
assert stage._last_audio_embeds[0].shape == (2, 4, 5)
assert torch.equal(embeds[0][0, :2], torch.full((2, 3), 2.0))
assert torch.equal(embeds[0][0, 2:], torch.zeros((2, 3)))
assert torch.equal(stage._last_audio_embeds[0][0, :2], torch.full((2, 5), 2.0))
assert torch.equal(stage._last_audio_embeds[0][0, 2:], torch.zeros((2, 5)))
def test_encode_text_hidden_state_flag_follows_encoder_config():
fastvideo_args, hidden = make_args(num_encoders=1, text_len=4, hidden_size=8)
@@ -156,3 +250,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)