Compare commits
29
Commits
v2
...
maint/pr1453-fixed
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
82b6d2dbf0 | ||
|
|
c6468e8def | ||
|
|
2c0d9a0828 | ||
|
|
b9876ee06d | ||
|
|
a568406b8c | ||
|
|
39cc075452 | ||
|
|
88f50fa7d6 | ||
|
|
4b25087bd8 | ||
|
|
a314489b78 | ||
|
|
a1b8a78e9e | ||
|
|
df55a31824 | ||
|
|
d022cf00e3 | ||
|
|
7dc3166827 | ||
|
|
ce4c4f37c7 | ||
|
|
48dff4be69 | ||
|
|
d9736d5e5d | ||
|
|
7396024f95 | ||
|
|
978896f534 | ||
|
|
560810c592 | ||
|
|
e3a7450a05 | ||
|
|
6aab7f3832 | ||
|
|
9cd53fe5f8 | ||
|
|
6a32cf3a5e | ||
|
|
98be9b3da2 | ||
|
|
c53e85b767 | ||
|
|
40a8bd2d3b | ||
|
|
31aa115611 | ||
|
|
98ac10a528 | ||
|
|
a5a6d171e5 |
@@ -1,5 +1,9 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"description": "Wan2.1 T2V 1.3B inference performance",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
|
||||
@@ -114,6 +114,17 @@ steps:
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Extraction Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
@@ -371,6 +382,21 @@ steps:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "scripts/lora_extraction/**"
|
||||
- "fastvideo/tests/lora_extraction/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/training/training_utils.py"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Extraction Tests"
|
||||
env:
|
||||
- TEST_TYPE=lora_extraction
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
|
||||
@@ -125,7 +125,7 @@ jobs:
|
||||
set -euo pipefail
|
||||
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
|
||||
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
|
||||
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
|
||||
exit 1
|
||||
@@ -136,6 +136,7 @@ jobs:
|
||||
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
|
||||
[ssim]=ssim [training]=training
|
||||
[lora-inference]=inference_lora [lora-training]=training_lora
|
||||
[lora-extraction]=lora_extraction
|
||||
[distillation]=distillation_dmd [self-forcing]=self_forcing
|
||||
[vsa]=training_vsa [vmoba]=inference_vmoba
|
||||
[performance]=performance [api]=api_server
|
||||
|
||||
@@ -13,12 +13,33 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
|
||||
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
|
||||
# push trigger is a sufficient change detector on its own -- no separate
|
||||
# detect-changes/paths-filter job is needed now that there is a single
|
||||
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
|
||||
# rocm Dockerfile stay manual-dispatch only.
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'docker/Dockerfile'
|
||||
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
# One static group, no cancellation: every run of this workflow writes the same
|
||||
# mutable registry tags (latest, py3.12-latest, ...), so runs must serialize —
|
||||
# concurrent push/dispatch runs would race on those tags, and cancelling a run
|
||||
# mid-publish can strand the cu126/cu130 tag families at different commits. An
|
||||
# in-flight superseded build wastes its runner time, but its tags are then
|
||||
# overwritten by the newer queued run. GitHub keeps a single pending run per
|
||||
# group: the newest queued run replaces any older queued one.
|
||||
concurrency:
|
||||
group: infra-build-image
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
# CUDA matrix: Python 3.12 x {12.6.3, 13.0.0} x {amd64, arm64}. Each architecture
|
||||
# builds natively and pushes only by digest; publish-cuda-manifests is the sole
|
||||
@@ -28,7 +49,11 @@ jobs:
|
||||
# aliases; 13.0.0/cu130 is published under explicit versioned tags. Flash-attn
|
||||
# 2.8.3 comes from the architecture-specific prebuilt releases.
|
||||
build-cuda-images:
|
||||
if: ${{ github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
|
||||
# on a push that changed docker/Dockerfile (inputs are null on push). The
|
||||
# repository guard keeps fork syncs from auto-building; manual dispatch
|
||||
# still works in forks.
|
||||
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -75,10 +100,10 @@ jobs:
|
||||
secrets: inherit
|
||||
|
||||
publish-cuda-manifests:
|
||||
# !cancelled(): a failed sibling build leg must not skip the manifests for a
|
||||
# CUDA lane whose own digests all exist; the digest-count check below fails
|
||||
# the incomplete lane loudly instead.
|
||||
if: ${{ !cancelled() && github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# !cancelled(): publish lanes whose digests exist even if a sibling build
|
||||
# leg failed (the digest-count check fails incomplete lanes); it also
|
||||
# bypasses skipped-needs propagation, hence the explicit skipped check.
|
||||
if: ${{ !cancelled() && needs.build-cuda-images.result != 'skipped' }}
|
||||
needs: build-cuda-images
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
|
||||
@@ -84,9 +84,12 @@ RUN source /opt/venv/bin/activate \
|
||||
FFMPEG_NATIVE_CXX=/usr/bin/g++ \
|
||||
bash /opt/FastVideo/apps/dreamverse/scripts/install_native_ffmpeg.sh
|
||||
|
||||
# FASTVIDEO_FA4: FA4 (flash_attn.cute) is opt-in; this image installs it via
|
||||
# the dreamverse extra and is validated with it, so enable it here.
|
||||
ENV FASTVIDEO_DREAMVERSE_HOME=/var/lib/dreamverse \
|
||||
STREAM_MODE=av_fmp4 \
|
||||
FASTVIDEO_ENABLE_PROMPT_SAFETY=0 \
|
||||
FASTVIDEO_FA4=1 \
|
||||
HF_HOME=/root/.cache/huggingface
|
||||
|
||||
RUN mkdir -p /var/lib/dreamverse
|
||||
|
||||
@@ -12,13 +12,17 @@ Defaults:
|
||||
|
||||
- `HF_REPO_ID=FastVideo/performance-tracking`
|
||||
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
|
||||
- `PERF_MAX_REGRESSION=0.05`
|
||||
|
||||
Records can include source metadata:
|
||||
Records can include source metadata and rolling-baseline policy context:
|
||||
|
||||
- `run_source`: `pr`, `local`, `scheduled_main`, or `unknown`
|
||||
- `baseline_eligible`: only successful scheduled-main records should be true
|
||||
- Buildkite metadata such as branch, PR number, build URL, build ID, and job ID
|
||||
- `regression_thresholds`: per-metric rolling-baseline percent and absolute
|
||||
floors used for recomputed status context
|
||||
|
||||
Dashboard/API metric payloads expose `threshold_exceeded` for raw threshold
|
||||
crossings; `regressed` remains the gated CI-failure signal.
|
||||
|
||||
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
|
||||
configured dataset repo requires authenticated access:
|
||||
@@ -89,7 +93,8 @@ Trend charts show metric-specific axes and exact point details on hover/focus:
|
||||
- PR number, branch, and Buildkite URL when present
|
||||
|
||||
The latest status table uses the stored JSON `success` value. Recomputed
|
||||
baseline context is shown separately and does not override stored status.
|
||||
baseline context applies each metric's percent and absolute regression floors
|
||||
and does not override stored status.
|
||||
|
||||
## API
|
||||
|
||||
|
||||
@@ -440,6 +440,8 @@ export default function App() {
|
||||
<th>Throughput</th>
|
||||
<th>Memory</th>
|
||||
<th>Worst</th>
|
||||
<th>Exceeded</th>
|
||||
<th>Failing</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
@@ -465,6 +467,12 @@ export default function App() {
|
||||
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
|
||||
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
|
||||
<td>
|
||||
{row.threshold_exceeded_metrics.length
|
||||
? row.threshold_exceeded_metrics.join(", ")
|
||||
: "none"}
|
||||
</td>
|
||||
<td>{row.failing_metrics.length ? row.failing_metrics.join(", ") : "none"}</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
|
||||
@@ -2,6 +2,12 @@ export type MetricValue = {
|
||||
current: number | null;
|
||||
baseline: number | null;
|
||||
regression_pct: number | null;
|
||||
absolute_delta: number | null;
|
||||
threshold_percent: number;
|
||||
threshold_absolute: number;
|
||||
gated: boolean;
|
||||
threshold_exceeded: boolean;
|
||||
regressed: boolean;
|
||||
label: string;
|
||||
lower_is_better: boolean;
|
||||
precision: number;
|
||||
@@ -15,7 +21,8 @@ export type SummaryRow = {
|
||||
success: boolean;
|
||||
baseline_n: number;
|
||||
worst_regression_pct: number | null;
|
||||
regression_threshold_pct: number;
|
||||
threshold_exceeded_metrics: string[];
|
||||
failing_metrics: string[];
|
||||
computed_regression_status: "pass" | "fail";
|
||||
status: "pass" | "fail";
|
||||
run_source: RunSource;
|
||||
|
||||
+4
-2
@@ -170,12 +170,14 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
|
||||
# Then verify both survive so a broken overlay fails the build instead of shipping
|
||||
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay and FA4 falls back to FA2.
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
|
||||
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
|
||||
# arm64 and the image runs FA3/FA2 as usual.
|
||||
RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; FA4 falls back to FA2)"; \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
|
||||
else \
|
||||
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
|
||||
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
|
||||
|
||||
@@ -103,6 +103,7 @@ can merge a PR.
|
||||
|---|---|---|
|
||||
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
|
||||
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
|
||||
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
|
||||
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
|
||||
@@ -144,6 +145,7 @@ Valid direct test names:
|
||||
| `/test training` | `training` |
|
||||
| `/test lora-inference` | `inference_lora` |
|
||||
| `/test lora-training` | `training_lora` |
|
||||
| `/test lora-extraction` | `lora_extraction` |
|
||||
| `/test distillation` | `distillation_dmd` |
|
||||
| `/test self-forcing` | `self_forcing` |
|
||||
| `/test vsa` | `training_vsa` |
|
||||
|
||||
@@ -72,7 +72,10 @@ fastvideo/tests/performance/
|
||||
│ writes Markdown summary + (optionally) uploads new records
|
||||
├── dashboard.py
|
||||
│ └── builds time-series Plotly HTML from HF history
|
||||
└── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
|
||||
fastvideo/performance/
|
||||
├── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
└── metric_policy.py # shared rolling-baseline threshold policy
|
||||
```
|
||||
|
||||
The HF dataset (`FastVideo/performance-tracking` by default) holds one
|
||||
@@ -92,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
|
||||
|
||||
@@ -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`**
|
||||
|
||||
@@ -12,6 +12,9 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
if(NOT GPU_BACKEND)
|
||||
set(GPU_BACKEND "CUDA")
|
||||
endif()
|
||||
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
@@ -50,7 +53,16 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
|
||||
if(_FASTVIDEO_USER_CUDA_ARCH)
|
||||
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
|
||||
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
|
||||
# Only numeric spellings translate; keywords like "native"/"all" would
|
||||
# otherwise be mangled into nonsense ("nativ.e").
|
||||
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
|
||||
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
|
||||
message(FATAL_ERROR
|
||||
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
|
||||
"supported. Use a numeric arch (e.g. 90a, 121), set "
|
||||
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
|
||||
"auto-detect from the visible GPU.")
|
||||
endif()
|
||||
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
|
||||
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
|
||||
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
|
||||
@@ -173,6 +185,14 @@ else()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
|
||||
# there, and tk's base_types.cuh brace-initializes signed-char vector members
|
||||
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
|
||||
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
|
||||
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
endif()
|
||||
|
||||
if(ENABLE_TK_KERNELS)
|
||||
message(STATUS "ThunderKittens kernels: ENABLED")
|
||||
else()
|
||||
@@ -263,11 +283,6 @@ set(CUDA_FLAGS
|
||||
"--expt-relaxed-constexpr"
|
||||
"-Xcompiler=-fno-strict-aliasing"
|
||||
"-Xcompiler=-fPIC"
|
||||
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
|
||||
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
|
||||
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
|
||||
# is a no-op on x86_64, where char is already signed.
|
||||
"-Xcompiler=-fsigned-char"
|
||||
"-DTORCH_COMPILE"
|
||||
"-Xnvlink=--verbose"
|
||||
"-Xptxas=--verbose"
|
||||
@@ -426,3 +441,14 @@ if(ENABLE_ATTN_QAT_INFER)
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
# One-look answer to "what is this build producing?" — kept last so it is the
|
||||
# final thing configure prints. The per-kernel matrix lives in README.md.
|
||||
message(STATUS "============== fastvideo-kernel build summary ==============")
|
||||
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
|
||||
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
|
||||
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
|
||||
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
|
||||
message(STATUS "============================================================")
|
||||
|
||||
@@ -2,6 +2,42 @@
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Kernel inventory
|
||||
|
||||
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
|
||||
|
||||
| Extension | Kernels | Sources | GPU arch | Build gate |
|
||||
|---|---|---|---|---|
|
||||
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
|
||||
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
|
||||
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
|
||||
|
||||
Runtime-JIT kernels (no build step, ship in every wheel/image):
|
||||
|
||||
| Kernels | Where | Used when |
|
||||
|---|---|---|
|
||||
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
|
||||
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
|
||||
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|
||||
|---|---|---|---|---|---|
|
||||
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
|
||||
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
|
||||
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
|
||||
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
|
||||
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
|
||||
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
|
||||
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
|
||||
|
||||
Notes:
|
||||
|
||||
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
|
||||
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
|
||||
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
|
||||
|
||||
## Installation
|
||||
|
||||
### Standard Installation (Local Development)
|
||||
|
||||
@@ -46,6 +46,23 @@ fi
|
||||
if git rev-parse --git-dir >/dev/null 2>&1; then
|
||||
git submodule update --init --recursive include/cutlass include/tk
|
||||
fi
|
||||
# Fail fast with a clear message if the headers are still missing (e.g. a
|
||||
# Docker context that excluded .git AND the submodule contents) instead of
|
||||
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
|
||||
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
|
||||
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
|
||||
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
|
||||
if [ ! -d include/cutlass/include ]; then
|
||||
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
|
||||
echo " CUTLASS sources must already be present (run" >&2
|
||||
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
|
||||
echo " in the source checkout, or include them in the build context)." >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -d include/tk/include ]; then
|
||||
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
|
||||
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
|
||||
fi
|
||||
|
||||
# Install build dependencies
|
||||
uv pip install scikit-build-core cmake ninja
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -4,10 +4,9 @@ import functools
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -16,7 +15,8 @@ if torch.cuda.is_available():
|
||||
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
|
||||
except ImportError:
|
||||
# flash_attn.cute (FA4) is simply not installed -- expected on builds
|
||||
# without it; callers fall back to FA3/FA2 quietly.
|
||||
# without it; callers handle the ImportError (the FASTVIDEO_FA4 gate in
|
||||
# flash_attn.py raises, the FP4 probe treats FA4 as unavailable).
|
||||
raise
|
||||
except Exception as e:
|
||||
# flash_attn.cute IS installed but failed to import -- almost always an
|
||||
@@ -24,23 +24,57 @@ if torch.cuda.is_available():
|
||||
# 'cutlass.cute.core' has no attribute 'ThrMma'" (an AttributeError, not
|
||||
# ImportError). This is fixable by pinning a compatible
|
||||
# nvidia-cutlass-dsl, so warn loudly, then re-raise as ImportError so
|
||||
# callers fall back to FA3/FA2 instead of crashing worker init.
|
||||
# callers can handle it uniformly.
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r); "
|
||||
"falling back to FA3/FA2. This is usually an nvidia-cutlass-dsl "
|
||||
"version mismatch -- pin a compatible nvidia-cutlass-dsl to "
|
||||
"restore FA4.", e)
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r). "
|
||||
"This is usually an nvidia-cutlass-dsl version mismatch -- pin a "
|
||||
"compatible nvidia-cutlass-dsl to restore FA4.", e)
|
||||
raise ImportError(f"flash_attn.cute (FA4) import failed: {e!r}") from e
|
||||
else:
|
||||
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
|
||||
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
|
||||
|
||||
try:
|
||||
# FA2 serves the calls FA4 cute cannot on pre-sm90 GPUs (backward, GQA).
|
||||
# Optional so FA4-only installs can still import this module.
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
except ImportError:
|
||||
_flash_attn_2_func = None
|
||||
_flash_attn_2_varlen_func = None
|
||||
|
||||
|
||||
def _check_dropout(dropout_p: float) -> None:
|
||||
if dropout_p != 0.0:
|
||||
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _sm90_or_newer() -> bool:
|
||||
return current_platform.has_device_capability(90)
|
||||
|
||||
|
||||
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
|
||||
if _sm90_or_newer():
|
||||
return False
|
||||
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
|
||||
# capability gate, not a runtime fallback):
|
||||
# * the backward asserts sm90+ (L40S/sm_89 dies on its arch check);
|
||||
# * GQA fails CuTeDSL JIT in pack_gqa ("ValueError: Operation creation
|
||||
# failed", observed on sm_89 with HunyuanGameCraft/LTX2).
|
||||
if q.shape[-2] != k.shape[-2]:
|
||||
return True
|
||||
return torch.is_grad_enabled() and any(t.requires_grad for t in (q, k, v))
|
||||
|
||||
|
||||
def _fa2_or_raise(fa2_func: Callable | None) -> Callable:
|
||||
if fa2_func is None:
|
||||
raise RuntimeError("this attention call cannot run on FA4 cute below sm90 (its backward and "
|
||||
"GQA support require sm90+) and flash-attn 2, which serves it there, is "
|
||||
"not installed.")
|
||||
return fa2_func
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_cute_forward",
|
||||
mutates_args=(),
|
||||
@@ -243,70 +277,6 @@ torch.library.register_autograd(
|
||||
)
|
||||
|
||||
|
||||
# FA4's CuTeDSL kernels JIT-compile per shape family, and some configurations
|
||||
# fail MLIR op creation at runtime even though the import succeeded (observed:
|
||||
# GQA models on sm_89 dying in pack_gqa with "ValueError: Operation creation
|
||||
# failed"). Degrade to FA2 once, process-wide, instead of crashing inference.
|
||||
class _FA4Policy:
|
||||
"""Per-call gate for the FA4 cute fast path, with FA2 as the fallback.
|
||||
|
||||
FA4 is skipped when:
|
||||
* a previous call failed at runtime -- CuTeDSL JIT compilation is
|
||||
shape-dependent, so the first failure disables FA4 for the rest of
|
||||
the process instead of retrying a broken JIT on every call; or
|
||||
* the call needs autograd -- FA4's backward asserts sm90+ (L40S/sm_89
|
||||
dies on its arch check) and is unvalidated for training in this repo
|
||||
(its lse is not even allocated through our inference-shaped custom
|
||||
op), so training keeps the pre-FA4 behavior: FA2 on every device.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.broken = False
|
||||
|
||||
def use_fa4(self, *tensors: torch.Tensor) -> bool:
|
||||
if self.broken:
|
||||
return False
|
||||
return not (torch.is_grad_enabled() and any(t.requires_grad for t in tensors))
|
||||
|
||||
def mark_broken(self, error: Exception) -> None:
|
||||
if not self.broken:
|
||||
self.broken = True
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) failed at runtime (%r); falling back "
|
||||
"to FA2 for the rest of this process.", error)
|
||||
|
||||
|
||||
_FA4 = _FA4Policy()
|
||||
|
||||
|
||||
def _with_fa2_fallback(fa2_func: Callable) -> Callable:
|
||||
"""Pair an FA4 cute wrapper with its FA2 twin of the same signature.
|
||||
|
||||
The decorated body runs only when ``_FA4`` allows it; otherwise (or after
|
||||
the first FA4 runtime failure) the call is served by ``fa2_func``.
|
||||
``NotImplementedError`` is a contract error (e.g. dropout), not a JIT
|
||||
failure, so it propagates without disabling FA4.
|
||||
"""
|
||||
|
||||
def decorator(fa4_func: Callable) -> Callable:
|
||||
|
||||
@functools.wraps(fa4_func)
|
||||
def wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
if _FA4.use_fa4(q, k, v):
|
||||
try:
|
||||
return fa4_func(q, k, v, *args, **kwargs)
|
||||
except NotImplementedError:
|
||||
raise
|
||||
except Exception as e: # CuTeDSL compile errors surface as ValueError
|
||||
_FA4.mark_broken(e)
|
||||
return fa2_func(q, k, v, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_func)
|
||||
def flash_attn_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -317,6 +287,16 @@ def flash_attn_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
|
||||
return out
|
||||
@@ -392,7 +372,6 @@ def flash_attn_fp4_func(
|
||||
return torch.ops.fastvideo._flash_attn_cute_fp4_forward(q, k, v, sfq, sfk, softmax_scale, causal)
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_varlen_func)
|
||||
def flash_attn_varlen_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -407,6 +386,20 @@ def flash_attn_varlen_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_varlen_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_varlen_forward(
|
||||
q,
|
||||
|
||||
@@ -21,24 +21,35 @@ from einops import rearrange
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
|
||||
from fastvideo import envs
|
||||
|
||||
|
||||
def _resolve_flash_attn_varlen_func() -> Any:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
if envs.FASTVIDEO_FA4:
|
||||
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
|
||||
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
|
||||
# fail loudly here rather than fall through to FA3/FA2. RuntimeError,
|
||||
# not ImportError: importers like bsa_attn.py treat ImportError as
|
||||
# "flash-attn not installed" and silently degrade to reference kernels.
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
|
||||
return flash_attn_varlen_func_cute
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
return flash_attn_varlen_func_flash
|
||||
return flash_attn_varlen_func_flash
|
||||
|
||||
|
||||
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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}")
|
||||
@@ -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
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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.")
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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] = {
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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:
|
||||
@@ -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,
|
||||
)
|
||||
@@ -13,7 +13,7 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
|
||||
from fastvideo.tests.performance import hf_store
|
||||
from fastvideo.performance import hf_store
|
||||
|
||||
from .service import build_latest_summary, build_trends, filter_records
|
||||
|
||||
@@ -150,7 +150,6 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type)
|
||||
rows = build_latest_summary(
|
||||
filtered,
|
||||
max_regression=float(os.environ.get("PERF_MAX_REGRESSION", "0.05")),
|
||||
run_source=run_source,
|
||||
)
|
||||
return {
|
||||
|
||||
@@ -1,26 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Metric definitions shared by the performance dashboard backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
from fastvideo.performance.metric_policy import DEFAULT_METRIC_POLICIES
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricDefinition:
|
||||
key: str
|
||||
label: str
|
||||
precision: int
|
||||
lower_is_better: bool
|
||||
|
||||
|
||||
METRICS: tuple[MetricDefinition, ...] = (
|
||||
MetricDefinition("latency", "Latency", 3, True),
|
||||
MetricDefinition("throughput", "Throughput", 3, False),
|
||||
MetricDefinition("memory", "Memory", 1, True),
|
||||
MetricDefinition("text_encoder_time_s", "Text Encoder", 3, True),
|
||||
MetricDefinition("dit_time_s", "DiT", 3, True),
|
||||
MetricDefinition("vae_decode_time_s", "VAE Decode", 3, True),
|
||||
)
|
||||
METRICS = DEFAULT_METRIC_POLICIES
|
||||
|
||||
METRIC_BY_KEY = {metric.key: metric for metric in METRICS}
|
||||
|
||||
@@ -13,9 +13,8 @@ from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.tests.performance.hf_store import is_baseline_eligible_record, safe_float
|
||||
|
||||
from .metrics import METRICS
|
||||
from fastvideo.performance.hf_store import is_baseline_eligible_record, safe_float
|
||||
from fastvideo.performance.metric_policy import regression_delta, resolve_metric_policies
|
||||
|
||||
Record = dict[str, Any]
|
||||
|
||||
@@ -95,19 +94,9 @@ def baseline_value(records: list[Record], metric_key: str) -> float | None:
|
||||
return float(statistics.median(values))
|
||||
|
||||
|
||||
def regression_percent(metric_key: str, current: float | None, baseline: float | None) -> float | None:
|
||||
if current is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
metric = next(metric for metric in METRICS if metric.key == metric_key)
|
||||
if metric.lower_is_better:
|
||||
return (current - baseline) / baseline * 100.0
|
||||
return (baseline - current) / baseline * 100.0
|
||||
|
||||
|
||||
def build_latest_summary(records: list[Record],
|
||||
*,
|
||||
baseline_window: int = 5,
|
||||
max_regression: float = 0.05,
|
||||
run_source: str | None = None) -> list[Record]:
|
||||
rows: list[Record] = []
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
@@ -123,52 +112,58 @@ def build_latest_summary(records: list[Record],
|
||||
if record is not latest and record.get("success", True) and is_baseline_eligible_record(record)
|
||||
]
|
||||
baseline_records = baseline_pool[-baseline_window:]
|
||||
metric_policies = resolve_metric_policies(latest.get("regression_thresholds"))
|
||||
|
||||
metrics: dict[str, Record] = {}
|
||||
regressions: list[float] = []
|
||||
for metric in METRICS:
|
||||
current = safe_float(latest.get(metric.key))
|
||||
baseline = baseline_value(baseline_records, metric.key)
|
||||
regression = regression_percent(metric.key, current, baseline)
|
||||
metrics[metric.key] = {
|
||||
failing_metrics: list[str] = []
|
||||
threshold_exceeded_metrics: list[str] = []
|
||||
for policy in metric_policies:
|
||||
current = safe_float(latest.get(policy.key))
|
||||
baseline = baseline_value(baseline_records, policy.key)
|
||||
delta = None
|
||||
if current is not None and baseline is not None:
|
||||
delta = regression_delta(policy, current, baseline)
|
||||
regression = None if delta is None else delta.percent * 100.0
|
||||
metrics[policy.key] = {
|
||||
"current": current,
|
||||
"baseline": baseline,
|
||||
"regression_pct": regression,
|
||||
"label": metric.label,
|
||||
"lower_is_better": metric.lower_is_better,
|
||||
"precision": metric.precision,
|
||||
"absolute_delta": None if delta is None else delta.absolute,
|
||||
"threshold_percent": policy.threshold_percent * 100.0,
|
||||
"threshold_absolute": policy.threshold_absolute,
|
||||
"gated": policy.gated,
|
||||
"threshold_exceeded": False if delta is None else delta.threshold_exceeded,
|
||||
"regressed": False if delta is None else delta.regressed,
|
||||
"label": policy.label,
|
||||
"lower_is_better": policy.lower_is_better,
|
||||
"precision": policy.precision,
|
||||
}
|
||||
if regression is not None:
|
||||
regressions.append(regression)
|
||||
if delta is not None and delta.threshold_exceeded:
|
||||
threshold_exceeded_metrics.append(policy.key)
|
||||
if delta is not None and delta.regressed:
|
||||
failing_metrics.append(policy.key)
|
||||
|
||||
worst_regression = max(regressions) if regressions else None
|
||||
success = bool(latest.get("success", True))
|
||||
status = "pass" if success else "fail"
|
||||
|
||||
rows.append({
|
||||
"model_id":
|
||||
model_id,
|
||||
"gpu_type":
|
||||
gpu_type,
|
||||
"timestamp":
|
||||
latest.get("timestamp"),
|
||||
"commit_sha":
|
||||
latest.get("commit_sha"),
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"timestamp": latest.get("timestamp"),
|
||||
"commit_sha": latest.get("commit_sha"),
|
||||
**record_metadata(latest),
|
||||
"success":
|
||||
success,
|
||||
"baseline_n":
|
||||
len(baseline_records),
|
||||
"worst_regression_pct":
|
||||
worst_regression,
|
||||
"regression_threshold_pct":
|
||||
max_regression * 100.0,
|
||||
"computed_regression_status":
|
||||
"fail" if worst_regression is not None and worst_regression > max_regression * 100.0 else "pass",
|
||||
"status":
|
||||
status,
|
||||
"metrics":
|
||||
metrics,
|
||||
"success": success,
|
||||
"baseline_n": len(baseline_records),
|
||||
"worst_regression_pct": worst_regression,
|
||||
"threshold_exceeded_metrics": threshold_exceeded_metrics,
|
||||
"failing_metrics": failing_metrics,
|
||||
"computed_regression_status": "fail" if failing_metrics else "pass",
|
||||
"status": status,
|
||||
"metrics": metrics,
|
||||
})
|
||||
|
||||
return sorted(rows, key=lambda row: (row["status"] != "fail", row["model_id"], row["gpu_type"]))
|
||||
@@ -179,14 +174,15 @@ def build_trends(records: list[Record]) -> list[Record]:
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
points = []
|
||||
for record in group:
|
||||
metric_policies = resolve_metric_policies(record.get("regression_thresholds"))
|
||||
point = {
|
||||
"timestamp": record.get("timestamp"),
|
||||
"commit_sha": record.get("commit_sha"),
|
||||
**record_metadata(record),
|
||||
"success": bool(record.get("success", True)),
|
||||
"metrics": {
|
||||
metric.key: safe_float(record.get(metric.key))
|
||||
for metric in METRICS
|
||||
policy.key: safe_float(record.get(policy.key))
|
||||
for policy in metric_policies
|
||||
},
|
||||
}
|
||||
points.append(point)
|
||||
|
||||
@@ -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...
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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))
|
||||
@@ -32,7 +32,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -40,6 +41,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App("fastvideo-gpu-job")
|
||||
|
||||
REPO_DIR = "/FastVideo"
|
||||
@@ -72,12 +76,7 @@ local_secrets = modal.Secret.from_dict({
|
||||
# Mutable tags inherit the registry image's baked backend, including custom
|
||||
# FASTVIDEO_MODAL_IMAGE overrides. Explicit CUDA tags also work with older
|
||||
# images that predate the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in IMAGE_TAG.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in IMAGE_TAG.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(IMAGE_TAG)
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(IMAGE_REF, add_python="3.12")
|
||||
@@ -98,6 +97,9 @@ image = (
|
||||
"TOKENIZERS_PARALLELISM": "false",
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); keep CI parity with the seeded
|
||||
# references. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
})
|
||||
)
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ through the ``fastvideo`` package, so they keep working on thin CI hosts that
|
||||
have ``modal`` but not torch.
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import urllib.request
|
||||
|
||||
_REGISTRY = "ghcr.io"
|
||||
@@ -55,3 +56,23 @@ def resolve_image_ref(image_ref: str) -> str:
|
||||
print(f"WARNING: could not resolve {image_ref} to a digest ({error}); "
|
||||
"Modal may reuse a stale cached image for this tag.")
|
||||
return image_ref
|
||||
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
"""UV_TORCH_BACKEND for a launcher image tag.
|
||||
|
||||
A caller-set UV_TORCH_BACKEND always wins. Otherwise sniff the CUDA
|
||||
version from an explicit image tag (cuda13 -> cu130, cuda12.6 -> cu126)
|
||||
so uv resolves torch against the image's toolkit. Mutable tags (e.g.
|
||||
py3.12-latest) return None and inherit the registry image's baked
|
||||
backend, which keeps a latest-tag CUDA transition safe.
|
||||
"""
|
||||
override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if override:
|
||||
return override
|
||||
tag = image_tag.lower()
|
||||
if "cuda13" in tag:
|
||||
return "cu130"
|
||||
if "cuda12.6" in tag:
|
||||
return "cu126"
|
||||
return None
|
||||
|
||||
@@ -5,7 +5,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -13,6 +14,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App()
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
@@ -24,12 +28,7 @@ print(f"Using image: {image_ref}")
|
||||
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
|
||||
# transition safe. Explicit CUDA tags also work with older images that predate
|
||||
# the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
|
||||
|
||||
image = (modal.Image.from_registry(
|
||||
image_ref, add_python="3.12"
|
||||
@@ -61,6 +60,10 @@ image = (modal.Image.from_registry(
|
||||
**({
|
||||
"UV_TORCH_BACKEND": uv_torch_backend_override
|
||||
} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); CI lanes keep it enabled to match the
|
||||
# SSIM/perf baselines. Caller override wins.
|
||||
"FASTVIDEO_FA4":
|
||||
os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
"HF_REPO_ID":
|
||||
"FastVideo/performance-tracking",
|
||||
}))
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -13,7 +13,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -21,6 +22,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App()
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
@@ -32,12 +36,7 @@ print(f"Using image: {image_ref}")
|
||||
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
|
||||
# transition safe. Explicit CUDA tags also work with older images that predate
|
||||
# the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(image_ref, add_python="3.12")
|
||||
@@ -64,6 +63,9 @@ image = (
|
||||
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"IMAGE_VERSION": image_version,
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); the SSIM references were seeded
|
||||
# with FA4 inference, so keep it enabled in CI. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
@@ -8,8 +8,8 @@ This script:
|
||||
baseline-eligible successful records (filtered by gpu_type),
|
||||
4) writes normalized records back to the HF dataset repo according to
|
||||
PERF_UPLOAD_POLICY,
|
||||
5) exits non-zero if any metric regresses by more than PERF_MAX_REGRESSION
|
||||
(default 5%).
|
||||
5) exits non-zero if any gated metric exceeds both its percent and absolute
|
||||
regression floors.
|
||||
"""
|
||||
|
||||
import glob
|
||||
@@ -21,21 +21,36 @@ from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
try:
|
||||
from .hf_store import (
|
||||
from fastvideo.performance.hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
from fastvideo.performance.metric_policy import (
|
||||
MetricPolicy,
|
||||
regression_delta,
|
||||
resolve_metric_policies,
|
||||
serialize_metric_thresholds,
|
||||
)
|
||||
except ImportError:
|
||||
from hf_store import (
|
||||
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
|
||||
if repo_root not in sys.path:
|
||||
sys.path.insert(0, repo_root)
|
||||
from fastvideo.performance.hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
from fastvideo.performance.metric_policy import (
|
||||
MetricPolicy,
|
||||
regression_delta,
|
||||
resolve_metric_policies,
|
||||
serialize_metric_thresholds,
|
||||
)
|
||||
|
||||
RESULTS_DIR = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)),
|
||||
@@ -46,25 +61,9 @@ TRACKING_ROOT = os.environ.get(
|
||||
"/tmp/perf-tracking",
|
||||
)
|
||||
PERF_REPORTS_DIR = os.environ.get("PERF_REPORTS_DIR", "/root/data/perf_reports")
|
||||
MAX_REGRESSION = float(os.environ.get("PERF_MAX_REGRESSION", "0.05"))
|
||||
UPLOAD_POLICY = os.environ.get("PERF_UPLOAD_POLICY", "never").strip().lower()
|
||||
VALID_UPLOAD_POLICIES = {"never", "pass", "always"}
|
||||
VALID_RUN_SOURCES = {"pr", "local", "scheduled_main", "unknown"}
|
||||
METRICS = (
|
||||
("latency", "Latency", 3),
|
||||
("throughput", "Throughput", 3),
|
||||
("memory", "Memory", 1),
|
||||
("text_encoder_time_s", "Text Enc", 3),
|
||||
("dit_time_s", "DiT", 3),
|
||||
("vae_decode_time_s", "VAE Decode", 3),
|
||||
)
|
||||
LOWER_IS_BETTER_METRICS = {
|
||||
"latency",
|
||||
"memory",
|
||||
"text_encoder_time_s",
|
||||
"dit_time_s",
|
||||
"vae_decode_time_s",
|
||||
}
|
||||
|
||||
|
||||
def _should_persist_tracking() -> bool:
|
||||
@@ -168,6 +167,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
text_encoder_time = safe_float(result.get("text_encoder_time_s"))
|
||||
dit_time = safe_float(result.get("dit_time_s"))
|
||||
vae_decode_time = safe_float(result.get("vae_decode_time_s"))
|
||||
metric_policies = resolve_metric_policies(result.get("regression_thresholds"))
|
||||
|
||||
return {
|
||||
"model_id": model_id,
|
||||
@@ -180,6 +180,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
"text_encoder_time_s": text_encoder_time,
|
||||
"dit_time_s": dit_time,
|
||||
"vae_decode_time_s": vae_decode_time,
|
||||
"regression_thresholds": serialize_metric_thresholds(metric_policies),
|
||||
"success": True,
|
||||
**_record_metadata(_detect_run_source(), result),
|
||||
}
|
||||
@@ -229,55 +230,41 @@ def _baseline_metric(records: list[dict[str, Any]], key: str) -> float | None:
|
||||
return statistics.median(values)
|
||||
|
||||
|
||||
def _metric_policy_summary(policy: MetricPolicy) -> str:
|
||||
gated = "gated" if policy.gated else "info"
|
||||
return (
|
||||
f"{gated}, >{policy.threshold_percent * 100:.1f}% "
|
||||
f"and >{policy.threshold_absolute:.{policy.precision}f}"
|
||||
)
|
||||
|
||||
|
||||
def _check_regressions(
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
) -> list[str]:
|
||||
failures: list[str] = []
|
||||
|
||||
for metric, _label, _precision in METRICS:
|
||||
if metric not in LOWER_IS_BETTER_METRICS:
|
||||
for policy in metric_policies:
|
||||
baseline = _baseline_metric(baseline_records, policy.key)
|
||||
curr = safe_float(current.get(policy.key))
|
||||
if baseline is None or curr is None:
|
||||
continue
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
curr = safe_float(current.get(metric))
|
||||
if baseline is None or curr is None or baseline <= 0:
|
||||
delta = regression_delta(policy, curr, baseline)
|
||||
if delta is None or not delta.regressed:
|
||||
continue
|
||||
regression = (curr - baseline) / baseline
|
||||
if regression > max_regression:
|
||||
failures.append(f"{current['model_id']} {metric} regressed by "
|
||||
f"{regression * 100:.1f}% "
|
||||
f"(current={curr:.3f}, baseline_median={baseline:.3f})")
|
||||
|
||||
baseline_tp = _baseline_metric(baseline_records, "throughput")
|
||||
curr_tp = safe_float(current.get("throughput"))
|
||||
if baseline_tp is not None and curr_tp is not None and baseline_tp > 0:
|
||||
regression = (baseline_tp - curr_tp) / baseline_tp
|
||||
if regression > max_regression:
|
||||
failures.append(f"{current['model_id']} throughput regressed by "
|
||||
f"{regression * 100:.1f}% "
|
||||
f"(current={curr_tp:.3f}, baseline_median={baseline_tp:.3f})")
|
||||
failures.append(
|
||||
f"{current['model_id']} {policy.key} regressed by "
|
||||
f"{delta.percent * 100:.1f}% and "
|
||||
f"{delta.absolute:.{policy.precision}f} "
|
||||
f"(current={curr:.{policy.precision}f}, "
|
||||
f"baseline_median={baseline:.{policy.precision}f}, "
|
||||
f"threshold={_metric_policy_summary(policy)})"
|
||||
)
|
||||
|
||||
return failures
|
||||
|
||||
|
||||
def _metric_delta_percent(
|
||||
metric: str,
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
) -> float | None:
|
||||
curr = safe_float(current.get(metric))
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
if curr is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
|
||||
if metric in LOWER_IS_BETTER_METRICS:
|
||||
return (curr - baseline) / baseline * 100.0
|
||||
if metric == "throughput":
|
||||
return (baseline - curr) / baseline * 100.0
|
||||
return None
|
||||
|
||||
|
||||
def _compact_value(value: float | None, precision: int = 3) -> str:
|
||||
if value is None:
|
||||
return "n/a"
|
||||
@@ -287,23 +274,42 @@ def _compact_value(value: float | None, precision: int = 3) -> str:
|
||||
def _build_summary_row(
|
||||
record: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
has_failed: bool,
|
||||
) -> dict[str, Any]:
|
||||
"""Format a single benchmark result as a row for the Markdown table."""
|
||||
|
||||
metric_values: dict[str, dict[str, float | None]] = {}
|
||||
metric_values: dict[str, dict[str, Any]] = {}
|
||||
regressions: list[float] = []
|
||||
for metric, _label, _precision in METRICS:
|
||||
curr = safe_float(record.get(metric))
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
regression = _metric_delta_percent(metric, record, baseline_records)
|
||||
metric_values[metric] = {
|
||||
failing_metrics: list[str] = []
|
||||
threshold_exceeded_metrics: list[str] = []
|
||||
for policy in metric_policies:
|
||||
curr = safe_float(record.get(policy.key))
|
||||
baseline = _baseline_metric(baseline_records, policy.key)
|
||||
delta = (
|
||||
regression_delta(policy, curr, baseline)
|
||||
if curr is not None and baseline is not None
|
||||
else None
|
||||
)
|
||||
regression = None if delta is None else delta.percent * 100.0
|
||||
absolute_delta = None if delta is None else delta.absolute
|
||||
metric_values[policy.key] = {
|
||||
"curr": curr,
|
||||
"base": baseline,
|
||||
"regression_pct": regression,
|
||||
"absolute_delta": absolute_delta,
|
||||
"threshold_percent": policy.threshold_percent * 100.0,
|
||||
"threshold_absolute": policy.threshold_absolute,
|
||||
"gated": policy.gated,
|
||||
"threshold_exceeded": False if delta is None else delta.threshold_exceeded,
|
||||
"regressed": False if delta is None else delta.regressed,
|
||||
}
|
||||
if regression is not None:
|
||||
regressions.append(regression)
|
||||
if delta is not None and delta.threshold_exceeded:
|
||||
threshold_exceeded_metrics.append(policy.key)
|
||||
if delta is not None and delta.regressed:
|
||||
failing_metrics.append(policy.key)
|
||||
|
||||
worst_regression_pct = max(regressions) if regressions else None
|
||||
|
||||
@@ -313,40 +319,50 @@ def _build_summary_row(
|
||||
"baseline_n": len(baseline_records),
|
||||
"metrics": metric_values,
|
||||
"worst_regression_pct": worst_regression_pct,
|
||||
"threshold_exceeded_metrics": threshold_exceeded_metrics,
|
||||
"failing_metrics": failing_metrics,
|
||||
"failed": has_failed,
|
||||
}
|
||||
|
||||
|
||||
def _build_markdown_summary(
|
||||
summary_rows: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
) -> str:
|
||||
lines = [
|
||||
"## Performance Baseline Comparison",
|
||||
"",
|
||||
f"Threshold: regressions greater than {max_regression * 100:.1f}% fail",
|
||||
"Threshold: gated metrics fail only when both percent and absolute "
|
||||
"regression floors are exceeded.",
|
||||
"",
|
||||
("| Model | GPU | Baseline N | Latency (curr/base) | "
|
||||
"Throughput (curr/base) | Memory (curr/base) | "
|
||||
"Text Enc (curr/base) | DiT (curr/base) | "
|
||||
"VAE Decode (curr/base) | Worst Regression | Status |"),
|
||||
"|---|---|---:|---|---|---|---|---|---|---:|---|",
|
||||
"VAE Decode (curr/base) | Worst Regression | Exceeded Metrics | "
|
||||
"Failing Metrics | Status |"),
|
||||
"|---|---|---:|---|---|---|---|---|---|---:|---|---|---|",
|
||||
]
|
||||
|
||||
for row in summary_rows:
|
||||
metric_cells = []
|
||||
for metric, _label, precision in METRICS:
|
||||
values = row["metrics"][metric]
|
||||
metric_cells.append(f"{_compact_value(values['curr'], precision)} / "
|
||||
f"{_compact_value(values['base'], precision)}")
|
||||
for policy in metric_policies:
|
||||
values = row["metrics"][policy.key]
|
||||
metric_cells.append(f"{_compact_value(values['curr'], policy.precision)} / "
|
||||
f"{_compact_value(values['base'], policy.precision)}")
|
||||
|
||||
worst_reg = ("n/a" if row["worst_regression_pct"] is None else f"{row['worst_regression_pct']:.1f}%")
|
||||
exceeded_metrics = (
|
||||
", ".join(row["threshold_exceeded_metrics"])
|
||||
if row["threshold_exceeded_metrics"]
|
||||
else "none"
|
||||
)
|
||||
failing_metrics = ", ".join(row["failing_metrics"]) if row["failing_metrics"] else "none"
|
||||
status = "FAIL" if row["failed"] else "PASS"
|
||||
|
||||
lines.append(f"| {row['model_id']} | {row['gpu_type']} | "
|
||||
f"{row['baseline_n']} | "
|
||||
f"{' | '.join(metric_cells)} | "
|
||||
f"{worst_reg} | {status} |")
|
||||
f"{worst_reg} | {exceeded_metrics} | {failing_metrics} | {status} |")
|
||||
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
@@ -400,6 +416,7 @@ def main() -> int:
|
||||
|
||||
for raw in current_results:
|
||||
record = _normalize_record(raw)
|
||||
metric_policies = resolve_metric_policies(record.get("regression_thresholds"))
|
||||
|
||||
baseline_records = load_records_for_model(
|
||||
TRACKING_ROOT,
|
||||
@@ -416,7 +433,7 @@ def main() -> int:
|
||||
failures: list[str] = []
|
||||
record["success"] = True
|
||||
else:
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
failures = _check_regressions(record, baseline_records, metric_policies)
|
||||
if static_threshold_failed:
|
||||
failures.append(f"{record['model_id']} fixed-threshold phase failed "
|
||||
f"(PERF_PYTEST_RC={os.environ.get('PERF_PYTEST_RC')})")
|
||||
@@ -434,10 +451,10 @@ def main() -> int:
|
||||
print("Tracking upload skipped for "
|
||||
f"{record['model_id']} ({record['run_source']}, success={record['success']})")
|
||||
|
||||
summary_rows.append(_build_summary_row(record, baseline_records, bool(failures)))
|
||||
summary_rows.append(_build_summary_row(record, baseline_records, metric_policies, bool(failures)))
|
||||
|
||||
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
|
||||
markdown = _build_markdown_summary(summary_rows, MAX_REGRESSION)
|
||||
markdown = _build_markdown_summary(summary_rows, resolve_metric_policies(None))
|
||||
_emit_markdown_summary(markdown, commit_sha)
|
||||
|
||||
if all_failures:
|
||||
|
||||
@@ -1,12 +1,19 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import sys
|
||||
from html import escape
|
||||
from datetime import datetime
|
||||
|
||||
import plotly.express as px
|
||||
import pandas as pd
|
||||
|
||||
from hf_store import sync_from_hf, load_as_dataframe
|
||||
try:
|
||||
from fastvideo.performance.hf_store import load_as_dataframe, sync_from_hf
|
||||
except ImportError:
|
||||
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
|
||||
if repo_root not in sys.path:
|
||||
sys.path.insert(0, repo_root)
|
||||
from fastvideo.performance.hf_store import load_as_dataframe, sync_from_hf
|
||||
|
||||
TRACKING_ROOT = os.environ.get(
|
||||
"PERFORMANCE_TRACKING_ROOT",
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.tests.performance.test_inference_performance import (
|
||||
_benchmark_display_id,
|
||||
_config_identity_metadata,
|
||||
_is_v2_config,
|
||||
_validate_benchmark_config,
|
||||
)
|
||||
|
||||
|
||||
def _v2_config():
|
||||
return {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
}
|
||||
|
||||
|
||||
def test_v1_benchmark_config_without_schema_version_validates():
|
||||
cfg = {
|
||||
"benchmark_id": "legacy-benchmark",
|
||||
}
|
||||
|
||||
_validate_benchmark_config(cfg, "legacy.json")
|
||||
|
||||
assert _is_v2_config(cfg) is False
|
||||
assert _config_identity_metadata(cfg) == {}
|
||||
assert _benchmark_display_id(cfg) == "legacy-benchmark"
|
||||
|
||||
|
||||
def test_v2_benchmark_config_identity_validates_and_is_preserved():
|
||||
cfg = _v2_config()
|
||||
cfg["quality_metadata"] = {"some": "data"}
|
||||
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
assert _is_v2_config(cfg) is True
|
||||
assert _config_identity_metadata(cfg) == {
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"quality_metadata": {"some": "data"},
|
||||
}
|
||||
|
||||
|
||||
def test_v2_benchmark_config_missing_identity_fields_fails_clearly():
|
||||
cfg = _v2_config()
|
||||
del cfg["variant_id"]
|
||||
del cfg["benchmark_version"]
|
||||
|
||||
expected = "wan.json: missing required v2 identity fields: variant_id, benchmark_version"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "value"),
|
||||
[
|
||||
("workload_id", {}),
|
||||
("workload_id", ""),
|
||||
("workload_id", " "),
|
||||
("variant_id", []),
|
||||
("variant_id", ""),
|
||||
("variant_id", " "),
|
||||
],
|
||||
)
|
||||
def test_v2_benchmark_config_rejects_invalid_string_identity_values(field, value):
|
||||
cfg = _v2_config()
|
||||
cfg[field] = value
|
||||
|
||||
expected = f"wan.json: v2 identity field {field!r} must be a non-empty string"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [None, "1", 1.5, True])
|
||||
def test_v2_benchmark_config_rejects_invalid_benchmark_version_values(value):
|
||||
cfg = _v2_config()
|
||||
cfg["benchmark_version"] = value
|
||||
|
||||
expected = "wan.json: v2 identity field 'benchmark_version' must be an integer"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
def test_partial_v2_identity_requires_schema_version():
|
||||
cfg = {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
}
|
||||
|
||||
expected = "wan.json: v2 benchmark identity fields require config_schema_version=2"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
def test_optional_v2_metadata_fields_must_be_objects():
|
||||
cfg = {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"quality_metadata": ["not", "an", "object"],
|
||||
}
|
||||
|
||||
expected = "wan.json: optional v2 metadata field 'quality_metadata' must be an object"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.tests.performance import compare_baseline
|
||||
from fastvideo.performance.metric_policy import resolve_metric_policies
|
||||
|
||||
|
||||
def _raw_result():
|
||||
@@ -72,9 +73,157 @@ def test_normalized_record_includes_source_metadata(monkeypatch):
|
||||
assert record["job_id"] == "job-1"
|
||||
|
||||
|
||||
def test_normalized_record_includes_effective_regression_thresholds(monkeypatch):
|
||||
monkeypatch.setenv("PERF_RUN_SOURCE", "pr")
|
||||
raw = _raw_result()
|
||||
raw["regression_thresholds"] = {
|
||||
"latency": {
|
||||
"threshold_percent": 0.09,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
},
|
||||
"throughput": {
|
||||
"gated": False,
|
||||
},
|
||||
}
|
||||
|
||||
record = compare_baseline.normalize_performance_result(raw)
|
||||
|
||||
assert record["regression_thresholds"]["latency"] == {
|
||||
"threshold_percent": 0.09,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
}
|
||||
assert record["regression_thresholds"]["throughput"]["gated"] is False
|
||||
|
||||
|
||||
def test_invalid_regression_threshold_container_uses_defaults():
|
||||
policies = resolve_metric_policies(["not", "a", "mapping"])
|
||||
|
||||
latency = next(policy for policy in policies if policy.key == "latency")
|
||||
assert latency.threshold_percent == 0.08
|
||||
assert latency.threshold_absolute == 0.5
|
||||
assert latency.gated is True
|
||||
|
||||
|
||||
def test_boolean_regression_threshold_values_are_ignored():
|
||||
policies = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": True,
|
||||
"threshold_absolute": False,
|
||||
"gated": "false",
|
||||
}
|
||||
})
|
||||
|
||||
latency = next(policy for policy in policies if policy.key == "latency")
|
||||
assert latency.threshold_percent == 0.08
|
||||
assert latency.threshold_absolute == 0.5
|
||||
assert latency.gated is False
|
||||
|
||||
|
||||
def test_baseline_eligibility_only_for_successful_scheduled_main():
|
||||
assert compare_baseline._is_baseline_eligible("scheduled_main", True) is True
|
||||
assert compare_baseline._is_baseline_eligible("scheduled_main", False) is False
|
||||
assert compare_baseline._is_baseline_eligible("pr", True) is False
|
||||
assert compare_baseline._is_baseline_eligible("local", True) is False
|
||||
|
||||
|
||||
def test_latency_regression_requires_percent_and_absolute_floors():
|
||||
baseline = [{"latency": 10.0}]
|
||||
current = {"model_id": "wan", "latency": 10.6}
|
||||
|
||||
percent_only = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.75,
|
||||
}
|
||||
})
|
||||
absolute_only = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.10,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
both = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
|
||||
assert compare_baseline._check_regressions(current, baseline, percent_only) == []
|
||||
assert compare_baseline._check_regressions(current, baseline, absolute_only) == []
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, both)
|
||||
assert len(failures) == 1
|
||||
assert "latency regressed by 6.0% and 0.600" in failures[0]
|
||||
|
||||
|
||||
def test_throughput_regression_uses_higher_is_better_direction():
|
||||
baseline = [{"throughput": 10.0}]
|
||||
current = {"model_id": "wan", "throughput": 9.0}
|
||||
policies = resolve_metric_policies({
|
||||
"throughput": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "throughput regressed by 10.0% and 1.000" in failures[0]
|
||||
|
||||
|
||||
def test_memory_regression_uses_metric_specific_absolute_floor():
|
||||
baseline = [{"memory": 10000.0}]
|
||||
current = {"model_id": "wan", "memory": 10600.0}
|
||||
policies = resolve_metric_policies({
|
||||
"memory": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 256.0,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "memory regressed by 6.0% and 600.0" in failures[0]
|
||||
|
||||
|
||||
def test_component_metric_can_gate_independently():
|
||||
baseline = [{"dit_time_s": 8.0}]
|
||||
current = {"model_id": "wan", "dit_time_s": 8.6}
|
||||
policies = resolve_metric_policies({
|
||||
"dit_time_s": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.25,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "dit_time_s regressed by 7.5% and 0.600" in failures[0]
|
||||
|
||||
|
||||
def test_informational_metric_remains_visible_without_failing():
|
||||
baseline = [{"throughput": 10.0}]
|
||||
current = {"model_id": "wan", "gpu_type": "NVIDIA L40S", "throughput": 8.0}
|
||||
policies = resolve_metric_policies({
|
||||
"throughput": {
|
||||
"threshold_percent": 0.01,
|
||||
"threshold_absolute": 0.01,
|
||||
"gated": False,
|
||||
}
|
||||
})
|
||||
|
||||
row = compare_baseline._build_summary_row(current, baseline, policies, False)
|
||||
|
||||
assert compare_baseline._check_regressions(current, baseline, policies) == []
|
||||
assert row["metrics"]["throughput"]["regression_pct"] == 20.0
|
||||
assert row["metrics"]["throughput"]["gated"] is False
|
||||
assert row["metrics"]["throughput"]["threshold_exceeded"] is True
|
||||
assert row["metrics"]["throughput"]["regressed"] is False
|
||||
assert row["threshold_exceeded_metrics"] == ["throughput"]
|
||||
assert row["failing_metrics"] == []
|
||||
|
||||
@@ -68,6 +68,8 @@ def test_summary_endpoint_returns_latest_group_status():
|
||||
assert body["count"] == 1
|
||||
assert body["status_counts"] == {"pass": 1, "fail": 0}
|
||||
assert body["rows"][0]["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert body["rows"][0]["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert body["rows"][0]["threshold_exceeded_metrics"] == ["latency", "throughput"]
|
||||
assert body["rows"][0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.performance import hf_store
|
||||
from fastvideo.performance_dashboard.service import build_latest_summary, build_trends, filter_records
|
||||
from fastvideo.tests.performance import hf_store
|
||||
|
||||
|
||||
def _record(ts, commit, latency, throughput, success=True, **metadata):
|
||||
@@ -28,16 +28,23 @@ def test_build_latest_summary_uses_previous_successful_records_for_baseline():
|
||||
_record("2026-01-03T00:00:00+00:00", "c" * 40, 11.0, 9.0),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05)
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert len(rows) == 1
|
||||
row = rows[0]
|
||||
assert row["baseline_n"] == 1
|
||||
assert row["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert row["metrics"]["latency"]["regression_pct"] == 10.0
|
||||
assert row["metrics"]["latency"]["absolute_delta"] == 1.0
|
||||
assert row["metrics"]["latency"]["threshold_percent"] == 8.0
|
||||
assert row["metrics"]["latency"]["threshold_absolute"] == 0.5
|
||||
assert row["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert row["metrics"]["latency"]["regressed"] is True
|
||||
assert row["metrics"]["throughput"]["regression_pct"] == 10.0
|
||||
assert row["status"] == "pass"
|
||||
assert row["computed_regression_status"] == "fail"
|
||||
assert row["threshold_exceeded_metrics"] == ["latency", "throughput"]
|
||||
assert row["failing_metrics"] == ["latency", "throughput"]
|
||||
|
||||
|
||||
def test_build_latest_summary_status_uses_latest_record_success_field():
|
||||
@@ -73,7 +80,7 @@ def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05, run_source="pr")
|
||||
rows = build_latest_summary(records, run_source="pr")
|
||||
|
||||
assert len(rows) == 1
|
||||
assert rows[0]["run_source"] == "pr"
|
||||
@@ -83,6 +90,60 @@ def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
|
||||
assert rows[0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
def test_build_latest_summary_requires_absolute_floor_for_computed_regression():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record(
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"b" * 40,
|
||||
10.6,
|
||||
10.0,
|
||||
regression_thresholds={
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
}
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert round(rows[0]["metrics"]["latency"]["regression_pct"], 1) == 6.0
|
||||
assert round(rows[0]["metrics"]["latency"]["absolute_delta"], 3) == 0.6
|
||||
assert rows[0]["metrics"]["latency"]["threshold_exceeded"] is False
|
||||
assert rows[0]["metrics"]["latency"]["regressed"] is False
|
||||
assert rows[0]["computed_regression_status"] == "pass"
|
||||
|
||||
|
||||
def test_build_latest_summary_separates_informational_threshold_crossing():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record(
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"b" * 40,
|
||||
10.6,
|
||||
10.0,
|
||||
regression_thresholds={
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
"gated": False,
|
||||
}
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert rows[0]["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert rows[0]["metrics"]["latency"]["regressed"] is False
|
||||
assert rows[0]["threshold_exceeded_metrics"] == ["latency"]
|
||||
assert rows[0]["failing_metrics"] == []
|
||||
assert rows[0]["computed_regression_status"] == "pass"
|
||||
|
||||
|
||||
def test_filter_records_and_trends_preserve_metric_points():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
|
||||
@@ -28,6 +28,17 @@ STAGE_METRIC_MAP: dict[str, str] = {
|
||||
"DmdDenoisingStage": "dit_time_s",
|
||||
"DecodingStage": "vae_decode_time_s",
|
||||
}
|
||||
V2_CONFIG_SCHEMA_VERSION = 2
|
||||
V2_REQUIRED_IDENTITY_FIELDS = (
|
||||
"workload_id",
|
||||
"variant_id",
|
||||
"benchmark_version",
|
||||
)
|
||||
V2_OPTIONAL_METADATA_FIELDS = (
|
||||
"recipe",
|
||||
"metric_threshold_policy",
|
||||
"quality_metadata",
|
||||
)
|
||||
|
||||
# -- Config discovery -------------------------------------------------------
|
||||
|
||||
@@ -42,6 +53,71 @@ _BENCHMARKS_DIR = os.path.join(
|
||||
)
|
||||
|
||||
|
||||
def _has_v2_fields(cfg):
|
||||
v2_fields = V2_REQUIRED_IDENTITY_FIELDS + V2_OPTIONAL_METADATA_FIELDS
|
||||
return any(field in cfg for field in v2_fields)
|
||||
|
||||
|
||||
def _is_v2_config(cfg):
|
||||
return cfg.get("config_schema_version") == V2_CONFIG_SCHEMA_VERSION
|
||||
|
||||
|
||||
def _validate_non_empty_string(value, field, path):
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError(f"{path}: v2 identity field {field!r} must be a non-empty string")
|
||||
|
||||
|
||||
def _validate_integer(value, field, path):
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"{path}: v2 identity field {field!r} must be an integer")
|
||||
|
||||
|
||||
def _validate_benchmark_config(cfg, path="<memory>"):
|
||||
missing_common = [field for field in ("benchmark_id",) if field not in cfg]
|
||||
if missing_common:
|
||||
raise ValueError(f"{path}: missing required benchmark config fields: {', '.join(missing_common)}")
|
||||
|
||||
schema_version = cfg.get("config_schema_version")
|
||||
if schema_version is None:
|
||||
if _has_v2_fields(cfg):
|
||||
raise ValueError(f"{path}: v2 benchmark identity fields require config_schema_version=2")
|
||||
return
|
||||
|
||||
if schema_version != V2_CONFIG_SCHEMA_VERSION:
|
||||
raise ValueError(f"{path}: unsupported benchmark config_schema_version={schema_version!r}")
|
||||
|
||||
missing_v2 = [field for field in V2_REQUIRED_IDENTITY_FIELDS if field not in cfg]
|
||||
if missing_v2:
|
||||
raise ValueError(f"{path}: missing required v2 identity fields: {', '.join(missing_v2)}")
|
||||
|
||||
_validate_non_empty_string(cfg["workload_id"], "workload_id", path)
|
||||
_validate_non_empty_string(cfg["variant_id"], "variant_id", path)
|
||||
_validate_integer(cfg["benchmark_version"], "benchmark_version", path)
|
||||
|
||||
for field in V2_OPTIONAL_METADATA_FIELDS:
|
||||
if field in cfg and not isinstance(cfg[field], Mapping):
|
||||
raise ValueError(f"{path}: optional v2 metadata field {field!r} must be an object")
|
||||
|
||||
|
||||
def _config_identity_metadata(cfg):
|
||||
if not _is_v2_config(cfg):
|
||||
return {}
|
||||
metadata = {
|
||||
"config_schema_version": cfg["config_schema_version"],
|
||||
"workload_id": cfg["workload_id"],
|
||||
"variant_id": cfg["variant_id"],
|
||||
"benchmark_version": cfg["benchmark_version"],
|
||||
}
|
||||
for field in V2_OPTIONAL_METADATA_FIELDS:
|
||||
if field in cfg:
|
||||
metadata[field] = cfg[field]
|
||||
return metadata
|
||||
|
||||
|
||||
def _benchmark_display_id(cfg):
|
||||
return cfg["benchmark_id"]
|
||||
|
||||
|
||||
def _discover_benchmarks():
|
||||
"""Glob benchmark JSON configs and return list of (id, config) tuples."""
|
||||
pattern = os.path.join(_BENCHMARKS_DIR, "*.json")
|
||||
@@ -49,6 +125,7 @@ def _discover_benchmarks():
|
||||
for path in sorted(glob.glob(pattern)):
|
||||
with open(path) as f:
|
||||
cfg = json.load(f)
|
||||
_validate_benchmark_config(cfg, path)
|
||||
configs.append(cfg)
|
||||
return configs
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user