Compare commits
25
Commits
v2
...
maint/pr1505-fixed
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4d8713f6e4 | ||
|
|
fda43fc610 | ||
|
|
5a8329ddf6 | ||
|
|
fcb5b465c5 | ||
|
|
4f13fb0fee | ||
|
|
d96ec99b4b | ||
|
|
2555f25cce | ||
|
|
fbe56fce8d | ||
|
|
5064bcdc47 | ||
|
|
2aa824b615 | ||
|
|
f934efc58d | ||
|
|
76b0550c15 | ||
|
|
384c1e9493 | ||
|
|
b1dbcc93f6 | ||
|
|
b93833772e | ||
|
|
30b523edd6 | ||
|
|
6aab7f3832 | ||
|
|
9cd53fe5f8 | ||
|
|
6a32cf3a5e | ||
|
|
98be9b3da2 | ||
|
|
c53e85b767 | ||
|
|
40a8bd2d3b | ||
|
|
31aa115611 | ||
|
|
98ac10a528 | ||
|
|
a5a6d171e5 |
@@ -69,7 +69,7 @@ approval, then upload reviewed accepted baseline records.
|
||||
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
|
||||
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
|
||||
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `PERF_MAX_REGRESSION` if set, otherwise `0.05` (5%). |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `0.05` (5%). |
|
||||
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
|
||||
|
||||
Hardcoded defaults:
|
||||
@@ -148,8 +148,7 @@ For each metric with at least two non-null source values:
|
||||
4. Stop if any source record regresses against the batch median by more than
|
||||
`max_intra_batch_regression`.
|
||||
|
||||
Default `max_intra_batch_regression` to `PERF_MAX_REGRESSION` when set,
|
||||
otherwise `0.05`. Print a table with per-source values, batch median, and
|
||||
Default `max_intra_batch_regression` to `0.05`. Print a table with per-source values, batch median, and
|
||||
worst intra-batch regression.
|
||||
|
||||
This check prevents uploading a mixed batch where one JSON is materially
|
||||
@@ -183,7 +182,7 @@ present, that run is not a valid source for baseline reseeding.
|
||||
|
||||
### 2. Sync and back up existing HF records under /tmp
|
||||
|
||||
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
|
||||
Use `fastvideo/performance/hf_store.py` helpers directly. Do **not** use
|
||||
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
|
||||
records, while this step must only fetch and back up existing history.
|
||||
|
||||
@@ -192,7 +191,7 @@ The sync command pattern is:
|
||||
```bash
|
||||
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
|
||||
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
|
||||
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
```
|
||||
|
||||
Then back up only the sanitized model directory under `/tmp`:
|
||||
@@ -200,8 +199,8 @@ Then back up only the sanitized model directory under `/tmp`:
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(PYTHONPATH=fastvideo/tests/performance python - <<'PY'
|
||||
from hf_store import sanitize
|
||||
MODEL_SAFE=$(python - <<'PY'
|
||||
from fastvideo.performance.hf_store import sanitize
|
||||
print(sanitize("<model_id>"))
|
||||
PY
|
||||
)
|
||||
@@ -235,7 +234,7 @@ first baseline seed. Continue, but report that baseline history was empty.
|
||||
Load the last 5 successful records for the target:
|
||||
|
||||
```python
|
||||
from hf_store import load_records_for_model
|
||||
from fastvideo.performance.hf_store import load_records_for_model
|
||||
|
||||
records = load_records_for_model(
|
||||
"/tmp/perf-tracking",
|
||||
@@ -372,7 +371,7 @@ prepared records plus backup on disk.
|
||||
Use the shared storage helper so the path and repo type match CI:
|
||||
|
||||
```python
|
||||
from hf_store import upload_record
|
||||
from fastvideo.performance.hf_store import upload_record
|
||||
|
||||
upload_record("<local_record_path>", record, strict=True)
|
||||
```
|
||||
@@ -460,7 +459,7 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
intentional baseline replacement.
|
||||
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
|
||||
median comparison, and persistence rules.
|
||||
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
|
||||
- `fastvideo/performance/hf_store.py` — HF sync, record loading,
|
||||
`sanitize()`, and `upload_record()`.
|
||||
- `fastvideo/tests/performance/test_inference_performance.py` — source result
|
||||
JSON schema.
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"description": "Wan2.1 T2V 1.3B inference performance",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
|
||||
@@ -114,6 +114,17 @@ steps:
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Extraction Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
@@ -371,6 +382,21 @@ steps:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "scripts/lora_extraction/**"
|
||||
- "fastvideo/tests/lora_extraction/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/training/training_utils.py"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Extraction Tests"
|
||||
env:
|
||||
- TEST_TYPE=lora_extraction
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
# pull_request_target instead of pull_request: the workflow definition and
|
||||
# the hook config are always taken from the BASE branch, so fork /
|
||||
# first-time-contributor PRs run immediately without a maintainer clicking
|
||||
# "Approve and run". The PR head is checked out as data only.
|
||||
pull_request_target:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
inputs:
|
||||
@@ -15,12 +19,25 @@ permissions:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
|
||||
if: github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
# For PR events, lint the PR head — but keep the hook definitions from
|
||||
# the base branch so an untrusted PR cannot alter what gets executed.
|
||||
- name: Save trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
- uses: actions/checkout@v4
|
||||
if: github.event_name == 'pull_request_target'
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha }}
|
||||
persist-credentials: false
|
||||
- name: Restore trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
@@ -125,7 +125,7 @@ jobs:
|
||||
set -euo pipefail
|
||||
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
|
||||
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
|
||||
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
|
||||
exit 1
|
||||
@@ -136,6 +136,7 @@ jobs:
|
||||
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
|
||||
[ssim]=ssim [training]=training
|
||||
[lora-inference]=inference_lora [lora-training]=training_lora
|
||||
[lora-extraction]=lora_extraction
|
||||
[distillation]=distillation_dmd [self-forcing]=self_forcing
|
||||
[vsa]=training_vsa [vmoba]=inference_vmoba
|
||||
[performance]=performance [api]=api_server
|
||||
|
||||
@@ -13,12 +13,33 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
|
||||
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
|
||||
# push trigger is a sufficient change detector on its own -- no separate
|
||||
# detect-changes/paths-filter job is needed now that there is a single
|
||||
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
|
||||
# rocm Dockerfile stay manual-dispatch only.
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'docker/Dockerfile'
|
||||
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
# One static group, no cancellation: every run of this workflow writes the same
|
||||
# mutable registry tags (latest, py3.12-latest, ...), so runs must serialize —
|
||||
# concurrent push/dispatch runs would race on those tags, and cancelling a run
|
||||
# mid-publish can strand the cu126/cu130 tag families at different commits. An
|
||||
# in-flight superseded build wastes its runner time, but its tags are then
|
||||
# overwritten by the newer queued run. GitHub keeps a single pending run per
|
||||
# group: the newest queued run replaces any older queued one.
|
||||
concurrency:
|
||||
group: infra-build-image
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
# CUDA matrix: Python 3.12 x {12.6.3, 13.0.0} x {amd64, arm64}. Each architecture
|
||||
# builds natively and pushes only by digest; publish-cuda-manifests is the sole
|
||||
@@ -28,7 +49,11 @@ jobs:
|
||||
# aliases; 13.0.0/cu130 is published under explicit versioned tags. Flash-attn
|
||||
# 2.8.3 comes from the architecture-specific prebuilt releases.
|
||||
build-cuda-images:
|
||||
if: ${{ github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
|
||||
# on a push that changed docker/Dockerfile (inputs are null on push). The
|
||||
# repository guard keeps fork syncs from auto-building; manual dispatch
|
||||
# still works in forks.
|
||||
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -75,10 +100,10 @@ jobs:
|
||||
secrets: inherit
|
||||
|
||||
publish-cuda-manifests:
|
||||
# !cancelled(): a failed sibling build leg must not skip the manifests for a
|
||||
# CUDA lane whose own digests all exist; the digest-count check below fails
|
||||
# the incomplete lane loudly instead.
|
||||
if: ${{ !cancelled() && github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
# !cancelled(): publish lanes whose digests exist even if a sibling build
|
||||
# leg failed (the digest-count check fails incomplete lanes); it also
|
||||
# bypasses skipped-needs propagation, hence the explicit skipped check.
|
||||
if: ${{ !cancelled() && needs.build-cuda-images.result != 'skipped' }}
|
||||
needs: build-cuda-images
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
|
||||
@@ -84,9 +84,12 @@ RUN source /opt/venv/bin/activate \
|
||||
FFMPEG_NATIVE_CXX=/usr/bin/g++ \
|
||||
bash /opt/FastVideo/apps/dreamverse/scripts/install_native_ffmpeg.sh
|
||||
|
||||
# FASTVIDEO_FA4: FA4 (flash_attn.cute) is opt-in; this image installs it via
|
||||
# the dreamverse extra and is validated with it, so enable it here.
|
||||
ENV FASTVIDEO_DREAMVERSE_HOME=/var/lib/dreamverse \
|
||||
STREAM_MODE=av_fmp4 \
|
||||
FASTVIDEO_ENABLE_PROMPT_SAFETY=0 \
|
||||
FASTVIDEO_FA4=1 \
|
||||
HF_HOME=/root/.cache/huggingface
|
||||
|
||||
RUN mkdir -p /var/lib/dreamverse
|
||||
|
||||
@@ -12,13 +12,17 @@ Defaults:
|
||||
|
||||
- `HF_REPO_ID=FastVideo/performance-tracking`
|
||||
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
|
||||
- `PERF_MAX_REGRESSION=0.05`
|
||||
|
||||
Records can include source metadata:
|
||||
Records can include source metadata and rolling-baseline policy context:
|
||||
|
||||
- `run_source`: `pr`, `local`, `scheduled_main`, or `unknown`
|
||||
- `baseline_eligible`: only successful scheduled-main records should be true
|
||||
- Buildkite metadata such as branch, PR number, build URL, build ID, and job ID
|
||||
- `regression_thresholds`: per-metric rolling-baseline percent and absolute
|
||||
floors used for recomputed status context
|
||||
|
||||
Dashboard/API metric payloads expose `threshold_exceeded` for raw threshold
|
||||
crossings; `regressed` remains the gated CI-failure signal.
|
||||
|
||||
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
|
||||
configured dataset repo requires authenticated access:
|
||||
@@ -89,7 +93,8 @@ Trend charts show metric-specific axes and exact point details on hover/focus:
|
||||
- PR number, branch, and Buildkite URL when present
|
||||
|
||||
The latest status table uses the stored JSON `success` value. Recomputed
|
||||
baseline context is shown separately and does not override stored status.
|
||||
baseline context applies each metric's percent and absolute regression floors
|
||||
and does not override stored status.
|
||||
|
||||
## API
|
||||
|
||||
|
||||
@@ -440,6 +440,8 @@ export default function App() {
|
||||
<th>Throughput</th>
|
||||
<th>Memory</th>
|
||||
<th>Worst</th>
|
||||
<th>Exceeded</th>
|
||||
<th>Failing</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
@@ -465,6 +467,12 @@ export default function App() {
|
||||
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
|
||||
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
|
||||
<td>
|
||||
{row.threshold_exceeded_metrics.length
|
||||
? row.threshold_exceeded_metrics.join(", ")
|
||||
: "none"}
|
||||
</td>
|
||||
<td>{row.failing_metrics.length ? row.failing_metrics.join(", ") : "none"}</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
|
||||
@@ -2,6 +2,12 @@ export type MetricValue = {
|
||||
current: number | null;
|
||||
baseline: number | null;
|
||||
regression_pct: number | null;
|
||||
absolute_delta: number | null;
|
||||
threshold_percent: number;
|
||||
threshold_absolute: number;
|
||||
gated: boolean;
|
||||
threshold_exceeded: boolean;
|
||||
regressed: boolean;
|
||||
label: string;
|
||||
lower_is_better: boolean;
|
||||
precision: number;
|
||||
@@ -15,7 +21,8 @@ export type SummaryRow = {
|
||||
success: boolean;
|
||||
baseline_n: number;
|
||||
worst_regression_pct: number | null;
|
||||
regression_threshold_pct: number;
|
||||
threshold_exceeded_metrics: string[];
|
||||
failing_metrics: string[];
|
||||
computed_regression_status: "pass" | "fail";
|
||||
status: "pass" | "fail";
|
||||
run_source: RunSource;
|
||||
|
||||
+4
-2
@@ -170,12 +170,14 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
|
||||
# Then verify both survive so a broken overlay fails the build instead of shipping
|
||||
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay and FA4 falls back to FA2.
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
|
||||
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
|
||||
# arm64 and the image runs FA3/FA2 as usual.
|
||||
RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; FA4 falls back to FA2)"; \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
|
||||
else \
|
||||
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
|
||||
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
|
||||
|
||||
@@ -103,6 +103,7 @@ can merge a PR.
|
||||
|---|---|---|
|
||||
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
|
||||
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
|
||||
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
|
||||
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
|
||||
@@ -144,6 +145,7 @@ Valid direct test names:
|
||||
| `/test training` | `training` |
|
||||
| `/test lora-inference` | `inference_lora` |
|
||||
| `/test lora-training` | `training_lora` |
|
||||
| `/test lora-extraction` | `lora_extraction` |
|
||||
| `/test distillation` | `distillation_dmd` |
|
||||
| `/test self-forcing` | `self_forcing` |
|
||||
| `/test vsa` | `training_vsa` |
|
||||
|
||||
@@ -72,7 +72,10 @@ fastvideo/tests/performance/
|
||||
│ writes Markdown summary + (optionally) uploads new records
|
||||
├── dashboard.py
|
||||
│ └── builds time-series Plotly HTML from HF history
|
||||
└── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
|
||||
fastvideo/performance/
|
||||
├── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
└── metric_policy.py # shared rolling-baseline threshold policy
|
||||
```
|
||||
|
||||
The HF dataset (`FastVideo/performance-tracking` by default) holds one
|
||||
@@ -92,25 +95,28 @@ and recipe changes instead of treating all records for a model as equivalent.
|
||||
|
||||
## Metrics
|
||||
|
||||
Each benchmark records six metrics:
|
||||
Each benchmark records six metrics. The rolling-baseline comparator also has a
|
||||
per-metric policy with direction, percent threshold, absolute threshold, and a
|
||||
`gated` flag.
|
||||
|
||||
| Metric | Raw key | Normalized key | Direction |
|
||||
|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better |
|
||||
| Metric | Raw key | Normalized key | Direction | Default rolling policy |
|
||||
|---|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better | 8% and 0.5 s |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better | 8% and 0.05 FPS |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better | 5% and 256 MB |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better | 5% and 0.25 s |
|
||||
|
||||
`test_inference_performance.py` temporarily sets `FASTVIDEO_STAGE_LOGGING=1`
|
||||
while it runs so pipeline stage execution times are available in
|
||||
`generate_video(...).logging_info`. Stage logs use pipeline-unique keys such as
|
||||
`prompt_encoding_stage` so duplicate stage classes do not collide. For
|
||||
`PipelineStage` entries, the extractor maps the `stage_class` field:
|
||||
`TextEncodingStage` maps to `text_encoder_time_s`, `DenoisingStage` and
|
||||
`DmdDenoisingStage` map to `dit_time_s`, and `DecodingStage` maps to
|
||||
`vae_decode_time_s`, with a fallback for older logs that used the class name as
|
||||
`PipelineStage` entries, shared component stage bases emit a stable
|
||||
`component_metric`: text encoding stages map to `text_encoder_time_s`,
|
||||
denoising stages and subclasses map to `dit_time_s`, and decoding stages map to
|
||||
`vae_decode_time_s`. The extractor falls back to known `stage_class` names for
|
||||
older logs that do not include `component_metric` or that used the class name as
|
||||
the stage key. Generator-side timings such as `PostDecodeFrameProcessStage`,
|
||||
`VideoSaveStage`, and `AudioMuxStage` are intentionally ignored. If a pipeline
|
||||
does not report one of the mapped stages, that component metric is stored as
|
||||
@@ -156,9 +162,22 @@ headroom and almost never need touching.
|
||||
|
||||
`compare_baseline.py` loads the last 5 successful, baseline-eligible records
|
||||
for the same `(model_id, gpu_type)` from the HF dataset, computes the median
|
||||
for each available metric, and fails if the current run regresses by more than
|
||||
`PERF_MAX_REGRESSION` (default 5%). For latency, memory, and component times,
|
||||
higher values are regressions. For throughput, lower values are regressions.
|
||||
for each available metric, and evaluates the current run with the metric's
|
||||
rolling regression policy. For latency, memory, and component times, higher
|
||||
values are regressions. For throughput, lower values are regressions.
|
||||
|
||||
A metric exceeds its rolling threshold when both of these are true:
|
||||
|
||||
```text
|
||||
percent_delta > threshold_percent
|
||||
absolute_delta > threshold_absolute
|
||||
```
|
||||
|
||||
Gated metrics fail CI when that threshold crossing happens. Set `gated: false`
|
||||
for metrics that should remain visible in reports and the dashboard without
|
||||
failing CI. Dashboard/API payloads expose `threshold_exceeded` separately from
|
||||
`regressed`, where `regressed` means a gated CI failure. Missing or `null`
|
||||
metrics are skipped.
|
||||
|
||||
This is the **drift detector** — it catches sub-threshold regressions that
|
||||
slowly add up. Only scheduled-main successful records are baseline eligible.
|
||||
@@ -172,6 +191,45 @@ agent skill to advance the rolling median.
|
||||
|
||||
## Schemas
|
||||
|
||||
### Benchmark config (`.buildkite/performance-benchmarks/tests/*.json`)
|
||||
|
||||
Benchmark configs without `config_schema_version` are treated as legacy v1
|
||||
configs and remain loadable. New or migrated configs should use
|
||||
`config_schema_version: 2` and include explicit comparable identity fields:
|
||||
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1
|
||||
}
|
||||
```
|
||||
|
||||
`benchmark_id` is still required in this phase because raw artifact names,
|
||||
generated-video directories, normalized record paths, and the current rolling
|
||||
baseline comparator still depend on it. The v2 identity fields are config
|
||||
metadata that make the measured workload explicit:
|
||||
|
||||
| Field | Purpose |
|
||||
|---|---|
|
||||
| `workload_id` | Stable benchmark family, such as `wan-t2v-1.3b`. |
|
||||
| `variant_id` | Intentional recipe family, such as `canonical`. |
|
||||
| `benchmark_version` | Version of the measurement protocol and comparison policy. |
|
||||
|
||||
If a config declares `config_schema_version: 2`, loading fails clearly when any
|
||||
required v2 identity field is missing. If v2 identity or metadata fields are
|
||||
added without `config_schema_version: 2`, loading also fails so partial
|
||||
migrations do not silently run as v1 configs. Optional v2 metadata fields
|
||||
reserved for follow-up work, such as `recipe`, `metric_threshold_policy`, and
|
||||
`quality_metadata`, must be JSON objects when present.
|
||||
|
||||
Recipe fingerprinting, hardware/software profile IDs, exact-identity
|
||||
comparison, metric-specific threshold policy behavior, promoted baselines, and
|
||||
dashboard regrouping are separate follow-up changes. Until those land, rolling
|
||||
baseline comparison remains keyed by `(model_id, gpu_type)`.
|
||||
|
||||
### Raw record (`results/perf_*.json`)
|
||||
|
||||
Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
@@ -179,6 +237,10 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
|
||||
"device": "NVIDIA L40S",
|
||||
"num_gpus": 2,
|
||||
@@ -196,6 +258,13 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
"max_dit_time_s": 10.0,
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.10,
|
||||
"threshold_absolute": 1.0,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"commit": "<full sha>",
|
||||
"pr_number": "1234",
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
@@ -222,6 +291,13 @@ result, used as the rolling-baseline source of truth.
|
||||
"text_encoder_time_s": 2.141,
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208,
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.08,
|
||||
"threshold_absolute": 0.5,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
@@ -238,18 +314,17 @@ successful main/full-suite uploads and remain eligible for rolling baselines.
|
||||
|
||||
| Variable | Default | Used by | Purpose |
|
||||
|---|---|---|---|
|
||||
| `PERF_MAX_REGRESSION` | `0.05` | `compare_baseline.py` | Per-metric regression fraction that fails the build. |
|
||||
| `PERFORMANCE_TRACKING_ROOT` | `/tmp/perf-tracking` | `compare_baseline.py`, `dashboard.py` | Local directory the HF dataset is synced to. |
|
||||
| `PERF_REPORTS_DIR` | `/root/data/perf_reports` | `compare_baseline.py`, `dashboard.py` | Where the Markdown summary and Plotly HTML get written for Buildkite to pick up. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `fastvideo/performance/hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `fastvideo/performance/hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `PERF_UPLOAD_POLICY` | `never` | `compare_baseline.py` | Upload policy: `never`, `pass`, or `always`. |
|
||||
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
|
||||
| `TEST_SCOPE` | unset | `compare_baseline.py` | CI context used to infer scheduled-main runs together with `BUILDKITE_BRANCH=main`. |
|
||||
| `BUILDKITE_BRANCH`, `BUILDKITE_COMMIT`, `BUILDKITE_PULL_REQUEST` | unset | `compare_baseline.py`, `test_inference_performance.py` | CI metadata stamped into records. |
|
||||
| `DASHBOARD_DAYS` | `30` | `dashboard.py` | Lookback window for the Plotly trend pages. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `fastvideo/performance/hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `FASTVIDEO_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
|
||||
|
||||
## CI integration
|
||||
@@ -261,9 +336,11 @@ Buildkite artifact upload is in
|
||||
`.buildkite/scripts/pr_test.sh:upload_performance_artifacts`.
|
||||
|
||||
Each performance build runs pytest first. If that fixed-threshold phase fails,
|
||||
`compare_baseline.py` is skipped, so Markdown summaries and normalized JSON
|
||||
artifacts are not emitted. The dashboard still runs best-effort for
|
||||
observability. When pytest passes, the rolling-baseline phase emits:
|
||||
PR/direct runs skip `compare_baseline.py` because they only upload passing
|
||||
records. Scheduled-main runs still execute `compare_baseline.py` with
|
||||
`PERF_PYTEST_RC` set so the failed canonical attempt is visible in normalized
|
||||
JSON and dashboard history. The dashboard runs best-effort for observability.
|
||||
When the rolling-baseline phase runs, it emits:
|
||||
|
||||
* **Markdown summary** — appended to `$GITHUB_STEP_SUMMARY` when that variable
|
||||
is set, and written as `perf_<sha>_<ts>.md` for Buildkite upload. Contains a
|
||||
@@ -279,11 +356,16 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
## Adding a new benchmark
|
||||
|
||||
1. Drop a new JSON config into
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. Required keys:
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. New configs should
|
||||
use v2 identity fields:
|
||||
|
||||
```json
|
||||
{
|
||||
"benchmark_id": "<unique-id>",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "<stable-workload-id>",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"model": { "model_path": "...", "model_short_name": "..." },
|
||||
"init_kwargs": { "num_gpus": 1, ... },
|
||||
"generation_kwargs": { "num_frames": 45, ... },
|
||||
@@ -299,9 +381,17 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"default": { "max_generation_time_s": 120.0, "max_peak_memory_mb": 30000.0 }
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": { "threshold_percent": 0.10, "threshold_absolute": 1.0, "gated": true }
|
||||
}
|
||||
}
|
||||
```
|
||||
}
|
||||
|
||||
```
|
||||
|
||||
Legacy v1 configs without `config_schema_version` still load, but should not
|
||||
gain v2 identity or metadata fields until they are migrated to
|
||||
`config_schema_version: 2`.
|
||||
|
||||
2. The pytest test auto-discovers all configs — no test code needed. CI
|
||||
picks it up on the next `/test performance` run.
|
||||
@@ -320,6 +410,13 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
a useful fixed gate. The rolling baseline will still track component times
|
||||
when static component thresholds are omitted.
|
||||
|
||||
6. Omit `regression_thresholds` to use the default rolling-baseline policy, or
|
||||
include only benchmark-specific deviations. Tune these independently from
|
||||
the fixed thresholds when a metric is noisy or should be informational. The
|
||||
fixed `thresholds` block is an absolute pytest ceiling. The
|
||||
`regression_thresholds` block controls rolling-baseline comparisons against
|
||||
recent scheduled-main records.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
**"No baseline for ... Initializing"** — first run for this `(model_id,
|
||||
@@ -336,5 +433,5 @@ pipelines that did not report a mapped component stage.
|
||||
|
||||
**Component timing is `null`** — the generated result did not include a mapped
|
||||
stage in `logging_info.stages`. Check that the pipeline emits stage logging
|
||||
and that the stage name is listed in `STAGE_METRIC_MAP` in
|
||||
`test_inference_performance.py`.
|
||||
and that the stage emits `component_metric` or is covered by the legacy
|
||||
`STAGE_METRIC_MAP` fallback in `test_inference_performance.py`.
|
||||
|
||||
@@ -74,6 +74,23 @@ uv pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
### Flash Attention 4 (opt-in)
|
||||
|
||||
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
|
||||
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
|
||||
runtime on some GPU/shape combinations. To use FA4, install the pinned
|
||||
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
|
||||
|
||||
```bash
|
||||
export FASTVIDEO_FA4=1
|
||||
```
|
||||
|
||||
On GPUs below sm90 a capability gate routes to FlashAttention-2 the calls FA4
|
||||
cannot serve there: grad-enabled (training) attention (FA4's backward requires
|
||||
sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
|
||||
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
|
||||
FastVideo fails loudly instead of silently falling back.
|
||||
|
||||
### FP4 Flash Attention 4 (Blackwell only)
|
||||
|
||||
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
# Causal Consistency Distillation: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# ODE-data-free distillation. A frozen teacher takes a single CFG Euler step;
|
||||
# the student matches an EMA copy of itself at the next timestep, all under
|
||||
# clean-history teacher forcing.
|
||||
#
|
||||
# All three roles initialize from the SAME checkpoint (the teacher-forcing
|
||||
# AR-diffusion model). Point init_from at that checkpoint for a real run.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
ema:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
|
||||
discrete_cd_N: 48
|
||||
guidance_scale: 3.0
|
||||
ema_decay: 0.99
|
||||
ema_start_step: 200
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_cd
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_cd_shift5
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,73 @@
|
||||
# DFSFT (Diffusion-Forcing SFT), frame-wise: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model with a block size of 1 frame
|
||||
# - Training: each frame gets its own independent noise level (frame-wise
|
||||
# diffusion forcing), versus the chunk-wise variant that shares one noise
|
||||
# level across num_frames_per_block frames.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
num_frames_per_block: 1
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 1
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_dfsft_framewise
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_dfsft_framewise_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -0,0 +1,72 @@
|
||||
# TFSFT (Teacher-Forcing SFT): Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model
|
||||
# - Training: inhomogeneous timesteps per chunk, but the causal transformer
|
||||
# denoises the current block while attending to *clean* history (clean_x),
|
||||
# not its own noisy rollout.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_tfsft
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_tfsft_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -12,6 +12,9 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
if(NOT GPU_BACKEND)
|
||||
set(GPU_BACKEND "CUDA")
|
||||
endif()
|
||||
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
@@ -50,7 +53,16 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
|
||||
if(_FASTVIDEO_USER_CUDA_ARCH)
|
||||
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
|
||||
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
|
||||
# Only numeric spellings translate; keywords like "native"/"all" would
|
||||
# otherwise be mangled into nonsense ("nativ.e").
|
||||
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
|
||||
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
|
||||
message(FATAL_ERROR
|
||||
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
|
||||
"supported. Use a numeric arch (e.g. 90a, 121), set "
|
||||
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
|
||||
"auto-detect from the visible GPU.")
|
||||
endif()
|
||||
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
|
||||
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
|
||||
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
|
||||
@@ -173,6 +185,14 @@ else()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
|
||||
# there, and tk's base_types.cuh brace-initializes signed-char vector members
|
||||
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
|
||||
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
|
||||
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
endif()
|
||||
|
||||
if(ENABLE_TK_KERNELS)
|
||||
message(STATUS "ThunderKittens kernels: ENABLED")
|
||||
else()
|
||||
@@ -263,11 +283,6 @@ set(CUDA_FLAGS
|
||||
"--expt-relaxed-constexpr"
|
||||
"-Xcompiler=-fno-strict-aliasing"
|
||||
"-Xcompiler=-fPIC"
|
||||
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
|
||||
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
|
||||
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
|
||||
# is a no-op on x86_64, where char is already signed.
|
||||
"-Xcompiler=-fsigned-char"
|
||||
"-DTORCH_COMPILE"
|
||||
"-Xnvlink=--verbose"
|
||||
"-Xptxas=--verbose"
|
||||
@@ -426,3 +441,14 @@ if(ENABLE_ATTN_QAT_INFER)
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
# One-look answer to "what is this build producing?" — kept last so it is the
|
||||
# final thing configure prints. The per-kernel matrix lives in README.md.
|
||||
message(STATUS "============== fastvideo-kernel build summary ==============")
|
||||
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
|
||||
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
|
||||
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
|
||||
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
|
||||
message(STATUS "============================================================")
|
||||
|
||||
@@ -2,6 +2,42 @@
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Kernel inventory
|
||||
|
||||
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
|
||||
|
||||
| Extension | Kernels | Sources | GPU arch | Build gate |
|
||||
|---|---|---|---|---|
|
||||
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
|
||||
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
|
||||
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
|
||||
|
||||
Runtime-JIT kernels (no build step, ship in every wheel/image):
|
||||
|
||||
| Kernels | Where | Used when |
|
||||
|---|---|---|
|
||||
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
|
||||
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
|
||||
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|
||||
|---|---|---|---|---|---|
|
||||
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
|
||||
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
|
||||
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
|
||||
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
|
||||
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
|
||||
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
|
||||
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
|
||||
|
||||
Notes:
|
||||
|
||||
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
|
||||
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
|
||||
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
|
||||
|
||||
## Installation
|
||||
|
||||
### Standard Installation (Local Development)
|
||||
|
||||
@@ -46,6 +46,23 @@ fi
|
||||
if git rev-parse --git-dir >/dev/null 2>&1; then
|
||||
git submodule update --init --recursive include/cutlass include/tk
|
||||
fi
|
||||
# Fail fast with a clear message if the headers are still missing (e.g. a
|
||||
# Docker context that excluded .git AND the submodule contents) instead of
|
||||
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
|
||||
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
|
||||
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
|
||||
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
|
||||
if [ ! -d include/cutlass/include ]; then
|
||||
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
|
||||
echo " CUTLASS sources must already be present (run" >&2
|
||||
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
|
||||
echo " in the source checkout, or include them in the build context)." >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -d include/tk/include ]; then
|
||||
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
|
||||
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
|
||||
fi
|
||||
|
||||
# Install build dependencies
|
||||
uv pip install scikit-build-core cmake ninja
|
||||
|
||||
@@ -1,15 +1,39 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
from fastvideo import envs
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1, mirroring the
|
||||
# kernel package's FASTVIDEO_VSA_CUTEDSL: its CuTeDSL kernels JIT-compile per
|
||||
# shape family and can fail at runtime on some arch/shape combinations, so it
|
||||
# is never auto-selected just because it is installed. Below sm90 a capability
|
||||
# gate in flash_attn_cute routes to FA2 the calls FA4 cannot serve there:
|
||||
# grad-enabled (its backward asserts sm90+) and GQA (pack_gqa fails CuTeDSL
|
||||
# JIT, observed on sm_89).
|
||||
if envs.FASTVIDEO_FA4:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
fa_version = "4"
|
||||
except ImportError:
|
||||
else:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
@@ -21,6 +45,12 @@ except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
try:
|
||||
if importlib.util.find_spec("flash_attn.cute") is not None:
|
||||
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
|
||||
"set FASTVIDEO_FA4=1 to use it for inference.")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
|
||||
# already a registered torch.library custom op, so dynamo treats it as a
|
||||
@@ -87,9 +117,10 @@ if fa_version in ("2", "3"):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
# FA4 path: `flash_attn_func` (from `flash_attn_cute`) goes through a
|
||||
# registered torch.library custom op (with an FA4 backward on sm90+;
|
||||
# grad-enabled and GQA calls below sm90 route to FA2), so a passthrough
|
||||
# is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
@@ -99,17 +130,6 @@ else:
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_WARNED_NON_FA_DTYPE = False
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
@@ -271,12 +291,8 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
# SP). Cast through bf16 and restore, matching TORCH_SDPA's tolerance.
|
||||
orig_dtype = query.dtype
|
||||
if orig_dtype not in (torch.float16, torch.bfloat16):
|
||||
global _WARNED_NON_FA_DTYPE
|
||||
if not _WARNED_NON_FA_DTYPE:
|
||||
_WARNED_NON_FA_DTYPE = True
|
||||
logger.warning(
|
||||
"FLASH_ATTN received %s inputs; casting to bfloat16 for the "
|
||||
"kernel and restoring on output.", orig_dtype)
|
||||
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
|
||||
f"bfloat16 for the kernel and restoring on output.")
|
||||
query = query.to(torch.bfloat16)
|
||||
key = key.to(torch.bfloat16)
|
||||
value = value.to(torch.bfloat16)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -20,6 +20,7 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: str | None = None
|
||||
FASTVIDEO_FA4: bool = False
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: str | None = None
|
||||
@@ -207,9 +208,18 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
|
||||
# - "SAGE_ATTN": use Sage Attention
|
||||
# - "SAGE_ATTN_THREE": use Sage Attention 3
|
||||
# FLASH_ATTN uses FlashAttention-3/2; to run FlashAttention-4 set
|
||||
# FASTVIDEO_FA4=1 as well (see below).
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
# If set (=1), the FLASH_ATTN backend uses FlashAttention-4
|
||||
# (flash_attn.cute). FA4 is opt-in and never auto-selected just because it
|
||||
# is installed. Below sm90, grad-enabled and GQA calls are routed to FA2
|
||||
# (FA4's backward asserts sm90+ and its pack_gqa fails to JIT there).
|
||||
"FASTVIDEO_FA4":
|
||||
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
|
||||
|
||||
@@ -323,9 +323,9 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q.forward_native(query)
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k.forward_native(key)
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
@@ -437,6 +437,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
# Causal-specific
|
||||
self.block_mask = None
|
||||
self.teacher_forcing_block_mask = None
|
||||
self.num_frame_per_block = config.arch_config.num_frames_per_block
|
||||
assert self.num_frame_per_block <= 3
|
||||
self.independent_first_frame = False
|
||||
@@ -500,6 +501,70 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
return block_mask
|
||||
|
||||
@staticmethod
|
||||
def _prepare_teacher_forcing_mask(
|
||||
device: torch.device | str, num_frames: int = 21,
|
||||
frame_seqlen: int = 1560, num_frame_per_block=1, local_attn_size=-1
|
||||
) -> BlockMask:
|
||||
"""Attention mask for the teacher-forcing ``[clean | noisy]`` sequence.
|
||||
|
||||
A noisy token attends to its own block plus the clean context of all
|
||||
strictly previous blocks; clean tokens are block-wise causal.
|
||||
"""
|
||||
if local_attn_size != -1:
|
||||
raise NotImplementedError(
|
||||
f"Teacher forcing ignores local_attn_size={local_attn_size}: "
|
||||
"unlike the block-wise causal mask, this mask always attends "
|
||||
"to the full clean context. Windowed teacher forcing is not "
|
||||
"implemented; use local_attn_size=-1 for teacher-forcing "
|
||||
"training.")
|
||||
total_length = num_frames * frame_seqlen * 2
|
||||
padded_length = math.ceil(total_length / 128) * 128 - total_length
|
||||
|
||||
clean_ends = num_frames * frame_seqlen
|
||||
context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_context_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_context_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_noise_starts = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
noise_noise_ends = torch.zeros(total_length + padded_length, device=device, dtype=torch.long)
|
||||
|
||||
attention_block_size = frame_seqlen * num_frame_per_block
|
||||
frame_indices = torch.arange(
|
||||
start=0, end=num_frames * frame_seqlen,
|
||||
step=attention_block_size, device=device, dtype=torch.long
|
||||
)
|
||||
for start in frame_indices:
|
||||
context_ends[start:start + attention_block_size] = start + attention_block_size
|
||||
|
||||
noisy_image_start_list = torch.arange(
|
||||
num_frames * frame_seqlen, total_length,
|
||||
step=attention_block_size, device=device, dtype=torch.long
|
||||
)
|
||||
noisy_image_end_list = noisy_image_start_list + attention_block_size
|
||||
for block_index, (start, end) in enumerate(zip(noisy_image_start_list, noisy_image_end_list)):
|
||||
noise_noise_starts[start:end] = start
|
||||
noise_noise_ends[start:end] = end
|
||||
noise_context_ends[start:end] = block_index * attention_block_size
|
||||
|
||||
def attention_mask(b, h, q_idx, kv_idx):
|
||||
clean_mask = (q_idx < clean_ends) & (kv_idx < context_ends[q_idx])
|
||||
c1 = (kv_idx < noise_noise_ends[q_idx]) & (kv_idx >= noise_noise_starts[q_idx])
|
||||
c2 = (kv_idx < noise_context_ends[q_idx]) & (kv_idx >= noise_context_starts[q_idx])
|
||||
noise_mask = (q_idx >= clean_ends) & (c1 | c2)
|
||||
eye_mask = q_idx == kv_idx
|
||||
return eye_mask | clean_mask | noise_mask
|
||||
|
||||
block_mask = create_block_mask(
|
||||
attention_mask, B=None, H=None,
|
||||
Q_LEN=total_length + padded_length, KV_LEN=total_length + padded_length,
|
||||
_compile=False, device=device)
|
||||
|
||||
if not dist.is_initialized() or dist.get_rank() == 0:
|
||||
print(f" cache a teacher-forcing mask with block size of {num_frame_per_block} frames")
|
||||
print(block_mask)
|
||||
|
||||
return block_mask
|
||||
|
||||
def _forward_inference(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -628,9 +693,12 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
start_frame: int = 0,
|
||||
clean_x: torch.Tensor | None = None,
|
||||
aug_t: torch.Tensor | None = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
teacher_forcing = clean_x is not None
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
if isinstance(encoder_hidden_states_image,
|
||||
@@ -663,15 +731,26 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
freqs_cis = (freqs_cos,
|
||||
freqs_sin) if freqs_cos is not None else None
|
||||
|
||||
# Construct blockwise causal attn mask
|
||||
if self.block_mask is None:
|
||||
self.block_mask = self._prepare_blockwise_causal_attn_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size
|
||||
)
|
||||
if teacher_forcing:
|
||||
if self.teacher_forcing_block_mask is None:
|
||||
self.teacher_forcing_block_mask = self._prepare_teacher_forcing_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size,
|
||||
)
|
||||
block_mask = self.teacher_forcing_block_mask
|
||||
else:
|
||||
if self.block_mask is None:
|
||||
self.block_mask = self._prepare_blockwise_causal_attn_mask(
|
||||
device=hidden_states.device,
|
||||
num_frames=num_frames,
|
||||
frame_seqlen=post_patch_height * post_patch_width,
|
||||
num_frame_per_block=self.num_frame_per_block,
|
||||
local_attn_size=self.local_attn_size
|
||||
)
|
||||
block_mask = self.block_mask
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
@@ -679,6 +758,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
encoder_hidden_states_text = encoder_hidden_states
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
@@ -694,18 +774,35 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
|
||||
if teacher_forcing:
|
||||
# Tile RoPE/modulation so clean frame i and noisy frame i share a position.
|
||||
clean_tokens = self.patch_embedding(clean_x).flatten(2).transpose(1, 2)
|
||||
hidden_states = torch.cat([clean_tokens, hidden_states], dim=1)
|
||||
if aug_t is None:
|
||||
aug_t = torch.zeros_like(timestep)
|
||||
_, timestep_proj_clean, _, _ = self.condition_embedder(
|
||||
aug_t.flatten(), encoder_hidden_states_text, None)
|
||||
timestep_proj_clean = timestep_proj_clean.unflatten(
|
||||
1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
|
||||
timestep_proj = torch.cat([timestep_proj_clean, timestep_proj], dim=1)
|
||||
freqs_cis = (torch.cat([freqs_cos, freqs_cos], dim=0),
|
||||
torch.cat([freqs_sin, freqs_sin], dim=0))
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis,
|
||||
block_mask=self.block_mask)
|
||||
block_mask=block_mask)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis,
|
||||
block_mask=self.block_mask)
|
||||
block_mask=block_mask)
|
||||
|
||||
if teacher_forcing:
|
||||
hidden_states = hidden_states[:, hidden_states.shape[1] // 2:]
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
|
||||
|
||||
@@ -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:
|
||||
@@ -312,8 +310,12 @@ def load_records_for_model(
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_NUMERIC_COLS = (
|
||||
"latency", "throughput", "memory",
|
||||
"text_encoder_time_s", "dit_time_s", "vae_decode_time_s",
|
||||
"latency",
|
||||
"throughput",
|
||||
"memory",
|
||||
"text_encoder_time_s",
|
||||
"dit_time_s",
|
||||
"vae_decode_time_s",
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Metric policy for rolling performance baseline comparisons."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricPolicy:
|
||||
key: str
|
||||
label: str
|
||||
precision: int
|
||||
lower_is_better: bool
|
||||
threshold_percent: float
|
||||
threshold_absolute: float
|
||||
gated: bool = True
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricDelta:
|
||||
absolute: float
|
||||
percent: float
|
||||
threshold_exceeded: bool
|
||||
regressed: bool
|
||||
|
||||
|
||||
DEFAULT_METRIC_POLICIES: tuple[MetricPolicy, ...] = (
|
||||
MetricPolicy("latency", "Latency", 3, True, 0.08, 0.5),
|
||||
MetricPolicy("throughput", "Throughput", 3, False, 0.08, 0.05),
|
||||
MetricPolicy("memory", "Memory", 1, True, 0.05, 256.0),
|
||||
MetricPolicy("text_encoder_time_s", "Text Enc", 3, True, 0.05, 0.25),
|
||||
MetricPolicy("dit_time_s", "DiT", 3, True, 0.05, 0.25),
|
||||
MetricPolicy("vae_decode_time_s", "VAE Decode", 3, True, 0.05, 0.25),
|
||||
)
|
||||
|
||||
|
||||
def _optional_float(value: Any) -> float | None:
|
||||
if value is None or isinstance(value, bool):
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _optional_bool(value: Any) -> bool | None:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in {"1", "true", "yes", "on"}:
|
||||
return True
|
||||
if normalized in {"0", "false", "no", "off"}:
|
||||
return False
|
||||
return None
|
||||
|
||||
|
||||
def resolve_metric_policies(threshold_overrides: Mapping[str, Any] | None, ) -> tuple[MetricPolicy, ...]:
|
||||
"""Return default metric policies with optional per-metric overrides."""
|
||||
|
||||
if not isinstance(threshold_overrides, Mapping):
|
||||
threshold_overrides = {}
|
||||
policies: list[MetricPolicy] = []
|
||||
for base_policy in DEFAULT_METRIC_POLICIES:
|
||||
raw_override = threshold_overrides.get(base_policy.key, {})
|
||||
if not isinstance(raw_override, Mapping):
|
||||
raw_override = {}
|
||||
|
||||
threshold_percent = _optional_float(raw_override.get("threshold_percent"))
|
||||
threshold_absolute = _optional_float(raw_override.get("threshold_absolute"))
|
||||
gated = _optional_bool(raw_override.get("gated"))
|
||||
|
||||
policies.append(
|
||||
MetricPolicy(
|
||||
key=base_policy.key,
|
||||
label=base_policy.label,
|
||||
precision=base_policy.precision,
|
||||
lower_is_better=base_policy.lower_is_better,
|
||||
threshold_percent=(base_policy.threshold_percent if threshold_percent is None else threshold_percent),
|
||||
threshold_absolute=(base_policy.threshold_absolute
|
||||
if threshold_absolute is None else threshold_absolute),
|
||||
gated=base_policy.gated if gated is None else gated,
|
||||
))
|
||||
return tuple(policies)
|
||||
|
||||
|
||||
def serialize_metric_thresholds(policies: tuple[MetricPolicy, ...], ) -> dict[str, dict[str, float | bool]]:
|
||||
return {
|
||||
policy.key: {
|
||||
"threshold_percent": policy.threshold_percent,
|
||||
"threshold_absolute": policy.threshold_absolute,
|
||||
"gated": policy.gated,
|
||||
}
|
||||
for policy in policies
|
||||
}
|
||||
|
||||
|
||||
def regression_delta(
|
||||
policy: MetricPolicy,
|
||||
current: float,
|
||||
baseline: float,
|
||||
) -> MetricDelta | None:
|
||||
if baseline <= 0:
|
||||
return None
|
||||
absolute_delta = current - baseline if policy.lower_is_better else baseline - current
|
||||
percent_delta = absolute_delta / baseline
|
||||
threshold_exceeded = (percent_delta > policy.threshold_percent and absolute_delta > policy.threshold_absolute)
|
||||
return MetricDelta(
|
||||
absolute=absolute_delta,
|
||||
percent=percent_delta,
|
||||
threshold_exceeded=threshold_exceeded,
|
||||
regressed=policy.gated and threshold_exceeded,
|
||||
)
|
||||
@@ -13,7 +13,7 @@ from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
|
||||
from fastvideo.tests.performance import hf_store
|
||||
from fastvideo.performance import hf_store
|
||||
|
||||
from .service import build_latest_summary, build_trends, filter_records
|
||||
|
||||
@@ -150,7 +150,6 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type)
|
||||
rows = build_latest_summary(
|
||||
filtered,
|
||||
max_regression=float(os.environ.get("PERF_MAX_REGRESSION", "0.05")),
|
||||
run_source=run_source,
|
||||
)
|
||||
return {
|
||||
|
||||
@@ -1,26 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Metric definitions shared by the performance dashboard backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
from fastvideo.performance.metric_policy import DEFAULT_METRIC_POLICIES
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricDefinition:
|
||||
key: str
|
||||
label: str
|
||||
precision: int
|
||||
lower_is_better: bool
|
||||
|
||||
|
||||
METRICS: tuple[MetricDefinition, ...] = (
|
||||
MetricDefinition("latency", "Latency", 3, True),
|
||||
MetricDefinition("throughput", "Throughput", 3, False),
|
||||
MetricDefinition("memory", "Memory", 1, True),
|
||||
MetricDefinition("text_encoder_time_s", "Text Encoder", 3, True),
|
||||
MetricDefinition("dit_time_s", "DiT", 3, True),
|
||||
MetricDefinition("vae_decode_time_s", "VAE Decode", 3, True),
|
||||
)
|
||||
METRICS = DEFAULT_METRIC_POLICIES
|
||||
|
||||
METRIC_BY_KEY = {metric.key: metric for metric in METRICS}
|
||||
|
||||
@@ -13,9 +13,8 @@ from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.tests.performance.hf_store import is_baseline_eligible_record, safe_float
|
||||
|
||||
from .metrics import METRICS
|
||||
from fastvideo.performance.hf_store import is_baseline_eligible_record, safe_float
|
||||
from fastvideo.performance.metric_policy import regression_delta, resolve_metric_policies
|
||||
|
||||
Record = dict[str, Any]
|
||||
|
||||
@@ -95,19 +94,9 @@ def baseline_value(records: list[Record], metric_key: str) -> float | None:
|
||||
return float(statistics.median(values))
|
||||
|
||||
|
||||
def regression_percent(metric_key: str, current: float | None, baseline: float | None) -> float | None:
|
||||
if current is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
metric = next(metric for metric in METRICS if metric.key == metric_key)
|
||||
if metric.lower_is_better:
|
||||
return (current - baseline) / baseline * 100.0
|
||||
return (baseline - current) / baseline * 100.0
|
||||
|
||||
|
||||
def build_latest_summary(records: list[Record],
|
||||
*,
|
||||
baseline_window: int = 5,
|
||||
max_regression: float = 0.05,
|
||||
run_source: str | None = None) -> list[Record]:
|
||||
rows: list[Record] = []
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
@@ -123,52 +112,58 @@ def build_latest_summary(records: list[Record],
|
||||
if record is not latest and record.get("success", True) and is_baseline_eligible_record(record)
|
||||
]
|
||||
baseline_records = baseline_pool[-baseline_window:]
|
||||
metric_policies = resolve_metric_policies(latest.get("regression_thresholds"))
|
||||
|
||||
metrics: dict[str, Record] = {}
|
||||
regressions: list[float] = []
|
||||
for metric in METRICS:
|
||||
current = safe_float(latest.get(metric.key))
|
||||
baseline = baseline_value(baseline_records, metric.key)
|
||||
regression = regression_percent(metric.key, current, baseline)
|
||||
metrics[metric.key] = {
|
||||
failing_metrics: list[str] = []
|
||||
threshold_exceeded_metrics: list[str] = []
|
||||
for policy in metric_policies:
|
||||
current = safe_float(latest.get(policy.key))
|
||||
baseline = baseline_value(baseline_records, policy.key)
|
||||
delta = None
|
||||
if current is not None and baseline is not None:
|
||||
delta = regression_delta(policy, current, baseline)
|
||||
regression = None if delta is None else delta.percent * 100.0
|
||||
metrics[policy.key] = {
|
||||
"current": current,
|
||||
"baseline": baseline,
|
||||
"regression_pct": regression,
|
||||
"label": metric.label,
|
||||
"lower_is_better": metric.lower_is_better,
|
||||
"precision": metric.precision,
|
||||
"absolute_delta": None if delta is None else delta.absolute,
|
||||
"threshold_percent": policy.threshold_percent * 100.0,
|
||||
"threshold_absolute": policy.threshold_absolute,
|
||||
"gated": policy.gated,
|
||||
"threshold_exceeded": False if delta is None else delta.threshold_exceeded,
|
||||
"regressed": False if delta is None else delta.regressed,
|
||||
"label": policy.label,
|
||||
"lower_is_better": policy.lower_is_better,
|
||||
"precision": policy.precision,
|
||||
}
|
||||
if regression is not None:
|
||||
regressions.append(regression)
|
||||
if delta is not None and delta.threshold_exceeded:
|
||||
threshold_exceeded_metrics.append(policy.key)
|
||||
if delta is not None and delta.regressed:
|
||||
failing_metrics.append(policy.key)
|
||||
|
||||
worst_regression = max(regressions) if regressions else None
|
||||
success = bool(latest.get("success", True))
|
||||
status = "pass" if success else "fail"
|
||||
|
||||
rows.append({
|
||||
"model_id":
|
||||
model_id,
|
||||
"gpu_type":
|
||||
gpu_type,
|
||||
"timestamp":
|
||||
latest.get("timestamp"),
|
||||
"commit_sha":
|
||||
latest.get("commit_sha"),
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"timestamp": latest.get("timestamp"),
|
||||
"commit_sha": latest.get("commit_sha"),
|
||||
**record_metadata(latest),
|
||||
"success":
|
||||
success,
|
||||
"baseline_n":
|
||||
len(baseline_records),
|
||||
"worst_regression_pct":
|
||||
worst_regression,
|
||||
"regression_threshold_pct":
|
||||
max_regression * 100.0,
|
||||
"computed_regression_status":
|
||||
"fail" if worst_regression is not None and worst_regression > max_regression * 100.0 else "pass",
|
||||
"status":
|
||||
status,
|
||||
"metrics":
|
||||
metrics,
|
||||
"success": success,
|
||||
"baseline_n": len(baseline_records),
|
||||
"worst_regression_pct": worst_regression,
|
||||
"threshold_exceeded_metrics": threshold_exceeded_metrics,
|
||||
"failing_metrics": failing_metrics,
|
||||
"computed_regression_status": "fail" if failing_metrics else "pass",
|
||||
"status": status,
|
||||
"metrics": metrics,
|
||||
})
|
||||
|
||||
return sorted(rows, key=lambda row: (row["status"] != "fail", row["model_id"], row["gpu_type"]))
|
||||
@@ -179,14 +174,15 @@ def build_trends(records: list[Record]) -> list[Record]:
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
points = []
|
||||
for record in group:
|
||||
metric_policies = resolve_metric_policies(record.get("regression_thresholds"))
|
||||
point = {
|
||||
"timestamp": record.get("timestamp"),
|
||||
"commit_sha": record.get("commit_sha"),
|
||||
**record_metadata(record),
|
||||
"success": bool(record.get("success", True)),
|
||||
"metrics": {
|
||||
metric.key: safe_float(record.get(metric.key))
|
||||
for metric in METRICS
|
||||
policy.key: safe_float(record.get(policy.key))
|
||||
for policy in metric_policies
|
||||
},
|
||||
}
|
||||
points.append(point)
|
||||
|
||||
@@ -34,6 +34,7 @@ class PipelineStage(ABC):
|
||||
composed with other stages to create a complete pipeline. Each stage is responsible
|
||||
for a specific part of the process, such as prompt encoding, latent preparation, etc.
|
||||
"""
|
||||
performance_component_metric: str | None = None
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""
|
||||
@@ -155,6 +156,9 @@ class PipelineStage(ABC):
|
||||
logger.info("[%s] Execution completed in %s ms", stage_name, execution_time * 1000)
|
||||
batch.logging_info.add_stage_execution_time(stage_key, execution_time)
|
||||
batch.logging_info.add_stage_metric(stage_key, "stage_class", stage_class_name)
|
||||
component_metric = self.performance_component_metric
|
||||
if component_metric is not None:
|
||||
batch.logging_info.add_stage_metric(stage_key, "component_metric", component_metric)
|
||||
except Exception as e:
|
||||
torch.cuda.synchronize()
|
||||
execution_time = time.perf_counter() - start_time
|
||||
|
||||
@@ -28,6 +28,7 @@ class DecodingStage(PipelineStage):
|
||||
This stage handles the decoding of latent representations into the final
|
||||
output format (e.g., pixel values).
|
||||
"""
|
||||
performance_component_metric = "vae_decode_time_s"
|
||||
|
||||
def __init__(self, vae, pipeline=None) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
@@ -51,6 +51,7 @@ class DenoisingStage(PipelineStage):
|
||||
This stage handles the iterative denoising process that transforms
|
||||
the initial noise into the final output.
|
||||
"""
|
||||
performance_component_metric = "dit_time_s"
|
||||
|
||||
def __init__(self, transformer, scheduler, pipeline=None, transformer_2=None, vae=None) -> None:
|
||||
super().__init__()
|
||||
@@ -1187,6 +1188,7 @@ class Cosmos25V2WDenoisingStage(Cosmos25DenoisingStage):
|
||||
|
||||
class Cosmos25AutoDenoisingStage(PipelineStage):
|
||||
"""Route Cosmos 2.5 denoising to T2W vs V2W/I2W."""
|
||||
performance_component_metric = "dit_time_s"
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
|
||||
@@ -24,6 +24,7 @@ class TextEncodingStage(PipelineStage):
|
||||
This stage handles the encoding of text prompts into the embedding space
|
||||
expected by the diffusion model.
|
||||
"""
|
||||
performance_component_metric = "text_encoder_time_s"
|
||||
|
||||
def __init__(self, text_encoders, tokenizers) -> None:
|
||||
"""
|
||||
@@ -350,6 +351,7 @@ class Cosmos25TextEncodingStage(PipelineStage):
|
||||
Cosmos 2.5 uses Reason1 (Qwen2.5-VL) and relies on the encoder's
|
||||
`compute_text_embeddings_online()`.
|
||||
"""
|
||||
performance_component_metric = "text_encoder_time_s"
|
||||
|
||||
def __init__(self, text_encoder) -> None:
|
||||
super().__init__()
|
||||
|
||||
@@ -3,15 +3,16 @@
|
||||
|
||||
Landed in PR #1225 slice 5 (Attn-QAT 5/12). The resolver centralises the
|
||||
varlen-flash-attn import-fallback logic that several backends
|
||||
(``bsa_attn.py``, ``video_sparse_attn.py``) used to duplicate. The fallback
|
||||
(``bsa_attn.py``, ``video_sparse_attn.py``) used to duplicate. The resolution
|
||||
order is:
|
||||
|
||||
1. ``fastvideo.attention.utils.flash_attn_cute``
|
||||
1. ``fastvideo.attention.utils.flash_attn_cute`` -- only when
|
||||
``FASTVIDEO_FA4=1`` (explicit opt-in), and then it must import or the
|
||||
resolver raises RuntimeError instead of falling through
|
||||
2. ``flash_attn_interface``
|
||||
3. ``flash_attn``
|
||||
|
||||
These tests verify that the resolver picks the highest-priority impl
|
||||
available and falls through cleanly on ``ImportError``. CPU-only, no
|
||||
These tests verify the opt-in gate and the FA3/FA2 fallthrough. CPU-only, no
|
||||
flash-attn install required.
|
||||
"""
|
||||
|
||||
@@ -35,8 +36,36 @@ def _reload_resolver_module():
|
||||
return importlib.import_module("fastvideo.attention.utils.flash_attn_no_pad")
|
||||
|
||||
|
||||
def test_resolver_falls_back_when_cute_unavailable(monkeypatch) -> None:
|
||||
"""When ``flash_attn_cute`` is unimportable, resolver tries the next impl."""
|
||||
def test_resolver_skips_cute_without_opt_in(monkeypatch) -> None:
|
||||
"""Without ``FASTVIDEO_FA4=1`` the resolver must not even attempt the cute
|
||||
import."""
|
||||
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
|
||||
attempted: list[str] = []
|
||||
real_import = builtins.__import__
|
||||
|
||||
def spying_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
attempted.append(name)
|
||||
return real_import(name, globals, locals, fromlist, level)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", spying_import)
|
||||
|
||||
mod = _reload_resolver_module()
|
||||
resolved = mod._resolve_flash_attn_varlen_func()
|
||||
assert resolved is not None
|
||||
assert resolved.__name__ == "flash_attn_varlen_func"
|
||||
assert "fastvideo.attention.utils.flash_attn_cute" not in attempted
|
||||
|
||||
|
||||
def test_resolver_raises_when_opted_in_but_cute_unavailable(monkeypatch) -> None:
|
||||
"""With ``FASTVIDEO_FA4=1`` an unimportable cute build fails loudly instead
|
||||
of silently falling through to FA3/FA2.
|
||||
|
||||
The resolver runs at module import time, so the reload itself must raise.
|
||||
It raises RuntimeError (not ImportError) so importers that treat
|
||||
ImportError as "flash-attn not installed" (``bsa_attn.py``) cannot swallow
|
||||
the opted-in failure.
|
||||
"""
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
real_import = builtins.__import__
|
||||
|
||||
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
@@ -46,21 +75,17 @@ def test_resolver_falls_back_when_cute_unavailable(monkeypatch) -> None:
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", patched_import)
|
||||
|
||||
mod = _reload_resolver_module()
|
||||
resolved = mod._resolve_flash_attn_varlen_func()
|
||||
assert resolved is not None
|
||||
assert resolved.__name__ == "flash_attn_varlen_func"
|
||||
with pytest.raises(RuntimeError, match="cute disabled for test"):
|
||||
_reload_resolver_module()
|
||||
|
||||
|
||||
def test_resolver_returns_flash_attn_when_cute_and_interface_unavailable(monkeypatch) -> None:
|
||||
def test_resolver_returns_flash_attn_when_interface_unavailable(monkeypatch) -> None:
|
||||
"""The terminal fallback is the plain ``flash_attn`` import."""
|
||||
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
|
||||
real_import = builtins.__import__
|
||||
|
||||
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
if name in {
|
||||
"fastvideo.attention.utils.flash_attn_cute",
|
||||
"flash_attn_interface",
|
||||
}:
|
||||
if name == "flash_attn_interface":
|
||||
raise ImportError(f"{name} disabled for test")
|
||||
return real_import(name, globals, locals, fromlist, level)
|
||||
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Guard: every test directory must be collected by some CI lane or be on
|
||||
the explicit allowlist below.
|
||||
|
||||
Three separate incidents on 2026-07-05 found test files that no CI lane
|
||||
ever collects (fastvideo/tests/stages/, tests/local_tests/ additions in
|
||||
PR #1509, and this sweep found seven dark directories in total): the tests
|
||||
pass review, merge, and then silently never run. This test makes going
|
||||
dark an explicit, reviewed decision instead of an accident: adding a new
|
||||
test directory fails CI until it is either wired into a lane or
|
||||
allowlisted here with a reason.
|
||||
|
||||
Pure text analysis — no fastvideo imports, no GPU, no torch.
|
||||
"""
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
TESTS_ROOT = REPO_ROOT / "fastvideo" / "tests"
|
||||
|
||||
# Files whose text constitutes "a CI lane references this directory".
|
||||
CI_SOURCES = [
|
||||
TESTS_ROOT / "modal" / "pr_test.py",
|
||||
TESTS_ROOT / "modal" / "ssim_test.py",
|
||||
*sorted((REPO_ROOT / ".buildkite").rglob("*.yml")),
|
||||
*sorted((REPO_ROOT / ".buildkite").rglob("*.sh")),
|
||||
]
|
||||
|
||||
# Directories that intentionally have no CI lane today. Every entry needs a
|
||||
# reason; remove the entry when the directory gets wired into a lane.
|
||||
# State as found on 2026-07-05 — these SHOULD shrink over time, not grow.
|
||||
ALLOWLIST = {
|
||||
"attention": "no lane yet — GPU attention-backend tests, run manually",
|
||||
"audio": "no lane yet — audio encoder tests, run manually",
|
||||
"distributed": "no lane yet — multi-GPU torchrun tests, run manually",
|
||||
"hooks": "no lane yet — run manually",
|
||||
"layers": "no lane yet — torchrun FSDP dispatch tests, run manually",
|
||||
"nightly": "by design: nightly cadence, not per-PR",
|
||||
"ops": "no lane yet — GPU op tests, run manually",
|
||||
"modal": "CI infrastructure itself, not a test suite",
|
||||
}
|
||||
|
||||
|
||||
def _dirs_with_tests() -> list[str]:
|
||||
dirs = []
|
||||
for child in sorted(TESTS_ROOT.iterdir()):
|
||||
if child.is_dir() and any(child.rglob("test_*.py")):
|
||||
dirs.append(child.name)
|
||||
return dirs
|
||||
|
||||
|
||||
def _ci_text() -> str:
|
||||
return "\n".join(
|
||||
src.read_text(errors="replace") for src in CI_SOURCES if src.exists())
|
||||
|
||||
|
||||
def test_every_test_directory_is_collected_or_allowlisted():
|
||||
ci_text = _ci_text()
|
||||
dark = [
|
||||
name for name in _dirs_with_tests()
|
||||
if f"tests/{name}" not in ci_text and name not in ALLOWLIST
|
||||
]
|
||||
assert not dark, (
|
||||
f"Test directories not referenced by any CI lane and not "
|
||||
f"allowlisted: {dark}. Wire them into a lane in "
|
||||
f"fastvideo/tests/modal/pr_test.py (or a Buildkite step), or add an "
|
||||
f"allowlist entry with a reason in {__file__}.")
|
||||
|
||||
|
||||
def test_local_tests_stays_out_of_ci():
|
||||
# tests/local_tests/ (repo root) is developer-local by design (author
|
||||
# decision, 2026-07-05): parity scaffolds and machine-specific checks
|
||||
# that must never gate CI. Fail if any CI source starts collecting it.
|
||||
assert "tests/local_tests" not in _ci_text(), (
|
||||
"tests/local_tests/ is local-only by design; remove the CI "
|
||||
"reference or move the tests into a fastvideo/tests/ lane.")
|
||||
|
||||
|
||||
def test_allowlist_entries_are_still_real_directories():
|
||||
# A stale allowlist hides regressions; entries must track reality.
|
||||
missing = [
|
||||
name for name in ALLOWLIST
|
||||
if name != "modal" and not (TESTS_ROOT / name).is_dir()
|
||||
]
|
||||
assert not missing, (
|
||||
f"Allowlisted directories no longer exist — remove them: {missing}")
|
||||
@@ -0,0 +1,203 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import init_device_mesh
|
||||
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard
|
||||
from torch.distributed.tensor import DTensor
|
||||
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
|
||||
WORLD_SIZE = 2
|
||||
HIDDEN_SIZE = 8
|
||||
SEED = 1379
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
|
||||
|
||||
def _run_torchrun(script_path: Path, mode: str, output_path: Path) -> None:
|
||||
# --standalone binds the rendezvous port atomically, avoiding the
|
||||
# free-port-probe race a hand-picked --master_port would have.
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--standalone",
|
||||
"--nproc_per_node",
|
||||
str(WORLD_SIZE),
|
||||
str(script_path),
|
||||
"--rmsnorm-fsdp-worker",
|
||||
"--mode",
|
||||
mode,
|
||||
"--output",
|
||||
str(output_path),
|
||||
]
|
||||
env = os.environ.copy()
|
||||
env.setdefault("TORCHDYNAMO_DISABLE", "1")
|
||||
try:
|
||||
process = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=env,
|
||||
timeout=120,
|
||||
)
|
||||
except subprocess.TimeoutExpired as error:
|
||||
raise RuntimeError(
|
||||
f"{mode} worker timed out after 120 seconds\n"
|
||||
f"STDOUT:\n{error.stdout}\n"
|
||||
f"STDERR:\n{error.stderr}"
|
||||
) from error
|
||||
if process.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"{mode} worker failed with code {process.returncode}\n"
|
||||
f"STDOUT:\n{process.stdout}\n"
|
||||
f"STDERR:\n{process.stderr}"
|
||||
)
|
||||
|
||||
|
||||
def _summarize_tensor(tensor: torch.Tensor | Any) -> dict[str, Any]:
|
||||
return {
|
||||
"type": type(tensor).__name__,
|
||||
"is_dtensor": isinstance(tensor, DTensor),
|
||||
"shape": list(tensor.shape) if hasattr(tensor, "shape") else None,
|
||||
"device": str(tensor.device) if hasattr(tensor, "device") else None,
|
||||
"dtype": str(tensor.dtype) if hasattr(tensor, "dtype") else None,
|
||||
}
|
||||
|
||||
|
||||
def _run_worker(mode: str, output_path: Path) -> None:
|
||||
if mode not in {
|
||||
"module_no_offload",
|
||||
"direct_no_offload",
|
||||
"module_cpu_offload",
|
||||
"direct_cpu_offload",
|
||||
}:
|
||||
raise ValueError(f"Unsupported mode: {mode}")
|
||||
|
||||
dist.init_process_group("nccl")
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.cuda.set_device(device)
|
||||
torch.manual_seed(SEED + rank)
|
||||
|
||||
try:
|
||||
mesh = init_device_mesh("cuda", (world_size,))
|
||||
norm = RMSNorm(HIDDEN_SIZE, eps=1e-6, has_weight=True).to(device)
|
||||
with torch.no_grad():
|
||||
norm.weight.fill_(1.0)
|
||||
|
||||
fsdp_kwargs: dict[str, Any] = {"mesh": mesh}
|
||||
if mode.endswith("cpu_offload"):
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(pin_memory=False)
|
||||
# fully_shard is applied to the bare RMSNorm to make the hook bypass
|
||||
# observable. Production sharding (fsdp_load.shard_model) only wraps
|
||||
# whole transformer blocks, whose pre-forward all-gather localizes norm
|
||||
# weights before the qk-norm call sites run, so this pins the dispatch
|
||||
# invariant rather than reproducing a production topology.
|
||||
fully_shard(norm, **fsdp_kwargs)
|
||||
|
||||
x = torch.randn(2, 3, HIDDEN_SIZE, device=device, dtype=torch.bfloat16)
|
||||
call_kind = "direct" if mode.startswith("direct") else "module"
|
||||
|
||||
try:
|
||||
if call_kind == "direct":
|
||||
output = norm.forward_native(x)
|
||||
else:
|
||||
output = norm(x)
|
||||
torch.cuda.synchronize(device)
|
||||
result = {
|
||||
"rank": rank,
|
||||
"ok": True,
|
||||
"mode": mode,
|
||||
"weight": _summarize_tensor(norm.weight),
|
||||
"output": _summarize_tensor(output),
|
||||
}
|
||||
except Exception as exc:
|
||||
result = {
|
||||
"rank": rank,
|
||||
"ok": False,
|
||||
"mode": mode,
|
||||
"error_type": type(exc).__name__,
|
||||
"error": str(exc),
|
||||
"weight": _summarize_tensor(norm.weight),
|
||||
}
|
||||
|
||||
gathered = [None for _ in range(world_size)] if rank == 0 else None
|
||||
dist.gather_object(result, object_gather_list=gathered, dst=0)
|
||||
if rank == 0:
|
||||
output_path.write_text(json.dumps(gathered, indent=2), encoding="utf-8")
|
||||
dist.barrier()
|
||||
finally:
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mode", "expect_ok"),
|
||||
[
|
||||
("module_no_offload", True),
|
||||
("direct_no_offload", False),
|
||||
("module_cpu_offload", True),
|
||||
("direct_cpu_offload", False),
|
||||
],
|
||||
)
|
||||
def test_rmsnorm_forward_native_bypasses_fsdp_hooks(mode: str, expect_ok: bool, tmp_path: Path) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("This test requires CUDA.")
|
||||
if torch.cuda.device_count() < WORLD_SIZE:
|
||||
pytest.skip(f"This test requires at least {WORLD_SIZE} CUDA devices.")
|
||||
|
||||
output_path = tmp_path / f"{mode}.json"
|
||||
_run_torchrun(Path(__file__).resolve(), mode, output_path)
|
||||
results = json.loads(output_path.read_text(encoding="utf-8"))
|
||||
print(f"\n{mode} results:\n{json.dumps(results, indent=2)}")
|
||||
|
||||
if expect_ok:
|
||||
failures = [result for result in results if not result["ok"]]
|
||||
assert not failures, json.dumps(results, indent=2)
|
||||
return
|
||||
|
||||
successes = [result for result in results if result["ok"]]
|
||||
assert not successes, json.dumps(results, indent=2)
|
||||
error_text = "\n".join(result.get("error", "") for result in results)
|
||||
# Pin the specific bypassed-hook failure: "got mixed torch.Tensor and
|
||||
# DTensor" ("Tensor" alone is a substring of "DTensor", so it adds nothing).
|
||||
assert "mixed" in error_text and "DTensor" in error_text, json.dumps(results, indent=2)
|
||||
|
||||
|
||||
def test_no_direct_forward_native_calls_in_models() -> None:
|
||||
"""Direct .forward_native(...) calls bypass nn.Module.__call__ and FSDP
|
||||
hooks (issue #1379); model code must use module dispatch instead."""
|
||||
models_dir = REPO_ROOT / "fastvideo" / "models"
|
||||
offenders = [
|
||||
str(path.relative_to(REPO_ROOT))
|
||||
for path in sorted(models_dir.rglob("*.py"))
|
||||
if ".forward_native(" in path.read_text(encoding="utf-8")
|
||||
]
|
||||
assert not offenders, f"Replace .forward_native(...) with module dispatch in: {offenders}"
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--rmsnorm-fsdp-worker", action="store_true")
|
||||
parser.add_argument("--mode", type=str, default=None)
|
||||
parser.add_argument("--output", type=str, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = _parse_args()
|
||||
if not args.rmsnorm_fsdp_worker:
|
||||
raise SystemExit("This module is intended to be run by pytest.")
|
||||
if args.mode is None or args.output is None:
|
||||
raise SystemExit("--mode and --output are required in worker mode.")
|
||||
_run_worker(mode=args.mode, output_path=Path(args.output))
|
||||
@@ -32,7 +32,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -40,6 +41,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App("fastvideo-gpu-job")
|
||||
|
||||
REPO_DIR = "/FastVideo"
|
||||
@@ -72,12 +76,7 @@ local_secrets = modal.Secret.from_dict({
|
||||
# Mutable tags inherit the registry image's baked backend, including custom
|
||||
# FASTVIDEO_MODAL_IMAGE overrides. Explicit CUDA tags also work with older
|
||||
# images that predate the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in IMAGE_TAG.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in IMAGE_TAG.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(IMAGE_TAG)
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(IMAGE_REF, add_python="3.12")
|
||||
@@ -98,6 +97,9 @@ image = (
|
||||
"TOKENIZERS_PARALLELISM": "false",
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); keep CI parity with the seeded
|
||||
# references. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
})
|
||||
)
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ through the ``fastvideo`` package, so they keep working on thin CI hosts that
|
||||
have ``modal`` but not torch.
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import urllib.request
|
||||
|
||||
_REGISTRY = "ghcr.io"
|
||||
@@ -55,3 +56,23 @@ def resolve_image_ref(image_ref: str) -> str:
|
||||
print(f"WARNING: could not resolve {image_ref} to a digest ({error}); "
|
||||
"Modal may reuse a stale cached image for this tag.")
|
||||
return image_ref
|
||||
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
"""UV_TORCH_BACKEND for a launcher image tag.
|
||||
|
||||
A caller-set UV_TORCH_BACKEND always wins. Otherwise sniff the CUDA
|
||||
version from an explicit image tag (cuda13 -> cu130, cuda12.6 -> cu126)
|
||||
so uv resolves torch against the image's toolkit. Mutable tags (e.g.
|
||||
py3.12-latest) return None and inherit the registry image's baked
|
||||
backend, which keeps a latest-tag CUDA transition safe.
|
||||
"""
|
||||
override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if override:
|
||||
return override
|
||||
tag = image_tag.lower()
|
||||
if "cuda13" in tag:
|
||||
return "cu130"
|
||||
if "cuda12.6" in tag:
|
||||
return "cu126"
|
||||
return None
|
||||
|
||||
@@ -5,7 +5,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -13,6 +14,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App()
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
@@ -24,12 +28,7 @@ print(f"Using image: {image_ref}")
|
||||
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
|
||||
# transition safe. Explicit CUDA tags also work with older images that predate
|
||||
# the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
|
||||
|
||||
image = (modal.Image.from_registry(
|
||||
image_ref, add_python="3.12"
|
||||
@@ -61,6 +60,10 @@ image = (modal.Image.from_registry(
|
||||
**({
|
||||
"UV_TORCH_BACKEND": uv_torch_backend_override
|
||||
} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); CI lanes keep it enabled to match the
|
||||
# SSIM/perf baselines. Caller override wins.
|
||||
"FASTVIDEO_FA4":
|
||||
os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
"HF_REPO_ID":
|
||||
"FastVideo/performance-tracking",
|
||||
}))
|
||||
@@ -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/contract/ ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ ./fastvideo/tests/entrypoints/ ./fastvideo/tests/train/ ./fastvideo/tests/stages/ --ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py --ignore=./fastvideo/tests/train/models --ignore=./fastvideo/tests/train/methods -vs"
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -13,7 +13,8 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -21,6 +22,9 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App()
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
@@ -32,12 +36,7 @@ print(f"Using image: {image_ref}")
|
||||
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
|
||||
# transition safe. Explicit CUDA tags also work with older images that predate
|
||||
# the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(image_ref, add_python="3.12")
|
||||
@@ -64,6 +63,9 @@ image = (
|
||||
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"IMAGE_VERSION": image_version,
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); the SSIM references were seeded
|
||||
# with FA4 inference, so keep it enabled in CI. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
@@ -8,8 +8,8 @@ This script:
|
||||
baseline-eligible successful records (filtered by gpu_type),
|
||||
4) writes normalized records back to the HF dataset repo according to
|
||||
PERF_UPLOAD_POLICY,
|
||||
5) exits non-zero if any metric regresses by more than PERF_MAX_REGRESSION
|
||||
(default 5%).
|
||||
5) exits non-zero if any gated metric exceeds both its percent and absolute
|
||||
regression floors.
|
||||
"""
|
||||
|
||||
import glob
|
||||
@@ -21,21 +21,36 @@ from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
try:
|
||||
from .hf_store import (
|
||||
from fastvideo.performance.hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
from fastvideo.performance.metric_policy import (
|
||||
MetricPolicy,
|
||||
regression_delta,
|
||||
resolve_metric_policies,
|
||||
serialize_metric_thresholds,
|
||||
)
|
||||
except ImportError:
|
||||
from hf_store import (
|
||||
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
|
||||
if repo_root not in sys.path:
|
||||
sys.path.insert(0, repo_root)
|
||||
from fastvideo.performance.hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
from fastvideo.performance.metric_policy import (
|
||||
MetricPolicy,
|
||||
regression_delta,
|
||||
resolve_metric_policies,
|
||||
serialize_metric_thresholds,
|
||||
)
|
||||
|
||||
RESULTS_DIR = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)),
|
||||
@@ -46,25 +61,9 @@ TRACKING_ROOT = os.environ.get(
|
||||
"/tmp/perf-tracking",
|
||||
)
|
||||
PERF_REPORTS_DIR = os.environ.get("PERF_REPORTS_DIR", "/root/data/perf_reports")
|
||||
MAX_REGRESSION = float(os.environ.get("PERF_MAX_REGRESSION", "0.05"))
|
||||
UPLOAD_POLICY = os.environ.get("PERF_UPLOAD_POLICY", "never").strip().lower()
|
||||
VALID_UPLOAD_POLICIES = {"never", "pass", "always"}
|
||||
VALID_RUN_SOURCES = {"pr", "local", "scheduled_main", "unknown"}
|
||||
METRICS = (
|
||||
("latency", "Latency", 3),
|
||||
("throughput", "Throughput", 3),
|
||||
("memory", "Memory", 1),
|
||||
("text_encoder_time_s", "Text Enc", 3),
|
||||
("dit_time_s", "DiT", 3),
|
||||
("vae_decode_time_s", "VAE Decode", 3),
|
||||
)
|
||||
LOWER_IS_BETTER_METRICS = {
|
||||
"latency",
|
||||
"memory",
|
||||
"text_encoder_time_s",
|
||||
"dit_time_s",
|
||||
"vae_decode_time_s",
|
||||
}
|
||||
|
||||
|
||||
def _should_persist_tracking() -> bool:
|
||||
@@ -168,6 +167,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
text_encoder_time = safe_float(result.get("text_encoder_time_s"))
|
||||
dit_time = safe_float(result.get("dit_time_s"))
|
||||
vae_decode_time = safe_float(result.get("vae_decode_time_s"))
|
||||
metric_policies = resolve_metric_policies(result.get("regression_thresholds"))
|
||||
|
||||
return {
|
||||
"model_id": model_id,
|
||||
@@ -180,6 +180,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
"text_encoder_time_s": text_encoder_time,
|
||||
"dit_time_s": dit_time,
|
||||
"vae_decode_time_s": vae_decode_time,
|
||||
"regression_thresholds": serialize_metric_thresholds(metric_policies),
|
||||
"success": True,
|
||||
**_record_metadata(_detect_run_source(), result),
|
||||
}
|
||||
@@ -229,55 +230,41 @@ def _baseline_metric(records: list[dict[str, Any]], key: str) -> float | None:
|
||||
return statistics.median(values)
|
||||
|
||||
|
||||
def _metric_policy_summary(policy: MetricPolicy) -> str:
|
||||
gated = "gated" if policy.gated else "info"
|
||||
return (
|
||||
f"{gated}, >{policy.threshold_percent * 100:.1f}% "
|
||||
f"and >{policy.threshold_absolute:.{policy.precision}f}"
|
||||
)
|
||||
|
||||
|
||||
def _check_regressions(
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
) -> list[str]:
|
||||
failures: list[str] = []
|
||||
|
||||
for metric, _label, _precision in METRICS:
|
||||
if metric not in LOWER_IS_BETTER_METRICS:
|
||||
for policy in metric_policies:
|
||||
baseline = _baseline_metric(baseline_records, policy.key)
|
||||
curr = safe_float(current.get(policy.key))
|
||||
if baseline is None or curr is None:
|
||||
continue
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
curr = safe_float(current.get(metric))
|
||||
if baseline is None or curr is None or baseline <= 0:
|
||||
delta = regression_delta(policy, curr, baseline)
|
||||
if delta is None or not delta.regressed:
|
||||
continue
|
||||
regression = (curr - baseline) / baseline
|
||||
if regression > max_regression:
|
||||
failures.append(f"{current['model_id']} {metric} regressed by "
|
||||
f"{regression * 100:.1f}% "
|
||||
f"(current={curr:.3f}, baseline_median={baseline:.3f})")
|
||||
|
||||
baseline_tp = _baseline_metric(baseline_records, "throughput")
|
||||
curr_tp = safe_float(current.get("throughput"))
|
||||
if baseline_tp is not None and curr_tp is not None and baseline_tp > 0:
|
||||
regression = (baseline_tp - curr_tp) / baseline_tp
|
||||
if regression > max_regression:
|
||||
failures.append(f"{current['model_id']} throughput regressed by "
|
||||
f"{regression * 100:.1f}% "
|
||||
f"(current={curr_tp:.3f}, baseline_median={baseline_tp:.3f})")
|
||||
failures.append(
|
||||
f"{current['model_id']} {policy.key} regressed by "
|
||||
f"{delta.percent * 100:.1f}% and "
|
||||
f"{delta.absolute:.{policy.precision}f} "
|
||||
f"(current={curr:.{policy.precision}f}, "
|
||||
f"baseline_median={baseline:.{policy.precision}f}, "
|
||||
f"threshold={_metric_policy_summary(policy)})"
|
||||
)
|
||||
|
||||
return failures
|
||||
|
||||
|
||||
def _metric_delta_percent(
|
||||
metric: str,
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
) -> float | None:
|
||||
curr = safe_float(current.get(metric))
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
if curr is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
|
||||
if metric in LOWER_IS_BETTER_METRICS:
|
||||
return (curr - baseline) / baseline * 100.0
|
||||
if metric == "throughput":
|
||||
return (baseline - curr) / baseline * 100.0
|
||||
return None
|
||||
|
||||
|
||||
def _compact_value(value: float | None, precision: int = 3) -> str:
|
||||
if value is None:
|
||||
return "n/a"
|
||||
@@ -287,23 +274,42 @@ def _compact_value(value: float | None, precision: int = 3) -> str:
|
||||
def _build_summary_row(
|
||||
record: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
has_failed: bool,
|
||||
) -> dict[str, Any]:
|
||||
"""Format a single benchmark result as a row for the Markdown table."""
|
||||
|
||||
metric_values: dict[str, dict[str, float | None]] = {}
|
||||
metric_values: dict[str, dict[str, Any]] = {}
|
||||
regressions: list[float] = []
|
||||
for metric, _label, _precision in METRICS:
|
||||
curr = safe_float(record.get(metric))
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
regression = _metric_delta_percent(metric, record, baseline_records)
|
||||
metric_values[metric] = {
|
||||
failing_metrics: list[str] = []
|
||||
threshold_exceeded_metrics: list[str] = []
|
||||
for policy in metric_policies:
|
||||
curr = safe_float(record.get(policy.key))
|
||||
baseline = _baseline_metric(baseline_records, policy.key)
|
||||
delta = (
|
||||
regression_delta(policy, curr, baseline)
|
||||
if curr is not None and baseline is not None
|
||||
else None
|
||||
)
|
||||
regression = None if delta is None else delta.percent * 100.0
|
||||
absolute_delta = None if delta is None else delta.absolute
|
||||
metric_values[policy.key] = {
|
||||
"curr": curr,
|
||||
"base": baseline,
|
||||
"regression_pct": regression,
|
||||
"absolute_delta": absolute_delta,
|
||||
"threshold_percent": policy.threshold_percent * 100.0,
|
||||
"threshold_absolute": policy.threshold_absolute,
|
||||
"gated": policy.gated,
|
||||
"threshold_exceeded": False if delta is None else delta.threshold_exceeded,
|
||||
"regressed": False if delta is None else delta.regressed,
|
||||
}
|
||||
if regression is not None:
|
||||
regressions.append(regression)
|
||||
if delta is not None and delta.threshold_exceeded:
|
||||
threshold_exceeded_metrics.append(policy.key)
|
||||
if delta is not None and delta.regressed:
|
||||
failing_metrics.append(policy.key)
|
||||
|
||||
worst_regression_pct = max(regressions) if regressions else None
|
||||
|
||||
@@ -313,40 +319,50 @@ def _build_summary_row(
|
||||
"baseline_n": len(baseline_records),
|
||||
"metrics": metric_values,
|
||||
"worst_regression_pct": worst_regression_pct,
|
||||
"threshold_exceeded_metrics": threshold_exceeded_metrics,
|
||||
"failing_metrics": failing_metrics,
|
||||
"failed": has_failed,
|
||||
}
|
||||
|
||||
|
||||
def _build_markdown_summary(
|
||||
summary_rows: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
metric_policies: tuple[MetricPolicy, ...],
|
||||
) -> str:
|
||||
lines = [
|
||||
"## Performance Baseline Comparison",
|
||||
"",
|
||||
f"Threshold: regressions greater than {max_regression * 100:.1f}% fail",
|
||||
"Threshold: gated metrics fail only when both percent and absolute "
|
||||
"regression floors are exceeded.",
|
||||
"",
|
||||
("| Model | GPU | Baseline N | Latency (curr/base) | "
|
||||
"Throughput (curr/base) | Memory (curr/base) | "
|
||||
"Text Enc (curr/base) | DiT (curr/base) | "
|
||||
"VAE Decode (curr/base) | Worst Regression | Status |"),
|
||||
"|---|---|---:|---|---|---|---|---|---|---:|---|",
|
||||
"VAE Decode (curr/base) | Worst Regression | Exceeded Metrics | "
|
||||
"Failing Metrics | Status |"),
|
||||
"|---|---|---:|---|---|---|---|---|---|---:|---|---|---|",
|
||||
]
|
||||
|
||||
for row in summary_rows:
|
||||
metric_cells = []
|
||||
for metric, _label, precision in METRICS:
|
||||
values = row["metrics"][metric]
|
||||
metric_cells.append(f"{_compact_value(values['curr'], precision)} / "
|
||||
f"{_compact_value(values['base'], precision)}")
|
||||
for policy in metric_policies:
|
||||
values = row["metrics"][policy.key]
|
||||
metric_cells.append(f"{_compact_value(values['curr'], policy.precision)} / "
|
||||
f"{_compact_value(values['base'], policy.precision)}")
|
||||
|
||||
worst_reg = ("n/a" if row["worst_regression_pct"] is None else f"{row['worst_regression_pct']:.1f}%")
|
||||
exceeded_metrics = (
|
||||
", ".join(row["threshold_exceeded_metrics"])
|
||||
if row["threshold_exceeded_metrics"]
|
||||
else "none"
|
||||
)
|
||||
failing_metrics = ", ".join(row["failing_metrics"]) if row["failing_metrics"] else "none"
|
||||
status = "FAIL" if row["failed"] else "PASS"
|
||||
|
||||
lines.append(f"| {row['model_id']} | {row['gpu_type']} | "
|
||||
f"{row['baseline_n']} | "
|
||||
f"{' | '.join(metric_cells)} | "
|
||||
f"{worst_reg} | {status} |")
|
||||
f"{worst_reg} | {exceeded_metrics} | {failing_metrics} | {status} |")
|
||||
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
@@ -400,6 +416,7 @@ def main() -> int:
|
||||
|
||||
for raw in current_results:
|
||||
record = _normalize_record(raw)
|
||||
metric_policies = resolve_metric_policies(record.get("regression_thresholds"))
|
||||
|
||||
baseline_records = load_records_for_model(
|
||||
TRACKING_ROOT,
|
||||
@@ -416,7 +433,7 @@ def main() -> int:
|
||||
failures: list[str] = []
|
||||
record["success"] = True
|
||||
else:
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
failures = _check_regressions(record, baseline_records, metric_policies)
|
||||
if static_threshold_failed:
|
||||
failures.append(f"{record['model_id']} fixed-threshold phase failed "
|
||||
f"(PERF_PYTEST_RC={os.environ.get('PERF_PYTEST_RC')})")
|
||||
@@ -434,10 +451,10 @@ def main() -> int:
|
||||
print("Tracking upload skipped for "
|
||||
f"{record['model_id']} ({record['run_source']}, success={record['success']})")
|
||||
|
||||
summary_rows.append(_build_summary_row(record, baseline_records, bool(failures)))
|
||||
summary_rows.append(_build_summary_row(record, baseline_records, metric_policies, bool(failures)))
|
||||
|
||||
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
|
||||
markdown = _build_markdown_summary(summary_rows, MAX_REGRESSION)
|
||||
markdown = _build_markdown_summary(summary_rows, resolve_metric_policies(None))
|
||||
_emit_markdown_summary(markdown, commit_sha)
|
||||
|
||||
if all_failures:
|
||||
|
||||
@@ -1,12 +1,19 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import sys
|
||||
from html import escape
|
||||
from datetime import datetime
|
||||
|
||||
import plotly.express as px
|
||||
import pandas as pd
|
||||
|
||||
from hf_store import sync_from_hf, load_as_dataframe
|
||||
try:
|
||||
from fastvideo.performance.hf_store import load_as_dataframe, sync_from_hf
|
||||
except ImportError:
|
||||
repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
|
||||
if repo_root not in sys.path:
|
||||
sys.path.insert(0, repo_root)
|
||||
from fastvideo.performance.hf_store import load_as_dataframe, sync_from_hf
|
||||
|
||||
TRACKING_ROOT = os.environ.get(
|
||||
"PERFORMANCE_TRACKING_ROOT",
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.tests.performance.test_inference_performance import (
|
||||
_benchmark_display_id,
|
||||
_config_identity_metadata,
|
||||
_is_v2_config,
|
||||
_validate_benchmark_config,
|
||||
)
|
||||
|
||||
|
||||
def _v2_config():
|
||||
return {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
}
|
||||
|
||||
|
||||
def test_v1_benchmark_config_without_schema_version_validates():
|
||||
cfg = {
|
||||
"benchmark_id": "legacy-benchmark",
|
||||
}
|
||||
|
||||
_validate_benchmark_config(cfg, "legacy.json")
|
||||
|
||||
assert _is_v2_config(cfg) is False
|
||||
assert _config_identity_metadata(cfg) == {}
|
||||
assert _benchmark_display_id(cfg) == "legacy-benchmark"
|
||||
|
||||
|
||||
def test_v2_benchmark_config_identity_validates_and_is_preserved():
|
||||
cfg = _v2_config()
|
||||
cfg["quality_metadata"] = {"some": "data"}
|
||||
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
assert _is_v2_config(cfg) is True
|
||||
assert _config_identity_metadata(cfg) == {
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"quality_metadata": {"some": "data"},
|
||||
}
|
||||
|
||||
|
||||
def test_v2_benchmark_config_missing_identity_fields_fails_clearly():
|
||||
cfg = _v2_config()
|
||||
del cfg["variant_id"]
|
||||
del cfg["benchmark_version"]
|
||||
|
||||
expected = "wan.json: missing required v2 identity fields: variant_id, benchmark_version"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "value"),
|
||||
[
|
||||
("workload_id", {}),
|
||||
("workload_id", ""),
|
||||
("workload_id", " "),
|
||||
("variant_id", []),
|
||||
("variant_id", ""),
|
||||
("variant_id", " "),
|
||||
],
|
||||
)
|
||||
def test_v2_benchmark_config_rejects_invalid_string_identity_values(field, value):
|
||||
cfg = _v2_config()
|
||||
cfg[field] = value
|
||||
|
||||
expected = f"wan.json: v2 identity field {field!r} must be a non-empty string"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [None, "1", 1.5, True])
|
||||
def test_v2_benchmark_config_rejects_invalid_benchmark_version_values(value):
|
||||
cfg = _v2_config()
|
||||
cfg["benchmark_version"] = value
|
||||
|
||||
expected = "wan.json: v2 identity field 'benchmark_version' must be an integer"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
def test_partial_v2_identity_requires_schema_version():
|
||||
cfg = {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
}
|
||||
|
||||
expected = "wan.json: v2 benchmark identity fields require config_schema_version=2"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
def test_optional_v2_metadata_fields_must_be_objects():
|
||||
cfg = {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"quality_metadata": ["not", "an", "object"],
|
||||
}
|
||||
|
||||
expected = "wan.json: optional v2 metadata field 'quality_metadata' must be an object"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.tests.performance import compare_baseline
|
||||
from fastvideo.performance.metric_policy import resolve_metric_policies
|
||||
|
||||
|
||||
def _raw_result():
|
||||
@@ -72,9 +73,157 @@ def test_normalized_record_includes_source_metadata(monkeypatch):
|
||||
assert record["job_id"] == "job-1"
|
||||
|
||||
|
||||
def test_normalized_record_includes_effective_regression_thresholds(monkeypatch):
|
||||
monkeypatch.setenv("PERF_RUN_SOURCE", "pr")
|
||||
raw = _raw_result()
|
||||
raw["regression_thresholds"] = {
|
||||
"latency": {
|
||||
"threshold_percent": 0.09,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
},
|
||||
"throughput": {
|
||||
"gated": False,
|
||||
},
|
||||
}
|
||||
|
||||
record = compare_baseline.normalize_performance_result(raw)
|
||||
|
||||
assert record["regression_thresholds"]["latency"] == {
|
||||
"threshold_percent": 0.09,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
}
|
||||
assert record["regression_thresholds"]["throughput"]["gated"] is False
|
||||
|
||||
|
||||
def test_invalid_regression_threshold_container_uses_defaults():
|
||||
policies = resolve_metric_policies(["not", "a", "mapping"])
|
||||
|
||||
latency = next(policy for policy in policies if policy.key == "latency")
|
||||
assert latency.threshold_percent == 0.08
|
||||
assert latency.threshold_absolute == 0.5
|
||||
assert latency.gated is True
|
||||
|
||||
|
||||
def test_boolean_regression_threshold_values_are_ignored():
|
||||
policies = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": True,
|
||||
"threshold_absolute": False,
|
||||
"gated": "false",
|
||||
}
|
||||
})
|
||||
|
||||
latency = next(policy for policy in policies if policy.key == "latency")
|
||||
assert latency.threshold_percent == 0.08
|
||||
assert latency.threshold_absolute == 0.5
|
||||
assert latency.gated is False
|
||||
|
||||
|
||||
def test_baseline_eligibility_only_for_successful_scheduled_main():
|
||||
assert compare_baseline._is_baseline_eligible("scheduled_main", True) is True
|
||||
assert compare_baseline._is_baseline_eligible("scheduled_main", False) is False
|
||||
assert compare_baseline._is_baseline_eligible("pr", True) is False
|
||||
assert compare_baseline._is_baseline_eligible("local", True) is False
|
||||
|
||||
|
||||
def test_latency_regression_requires_percent_and_absolute_floors():
|
||||
baseline = [{"latency": 10.0}]
|
||||
current = {"model_id": "wan", "latency": 10.6}
|
||||
|
||||
percent_only = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.75,
|
||||
}
|
||||
})
|
||||
absolute_only = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.10,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
both = resolve_metric_policies({
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
|
||||
assert compare_baseline._check_regressions(current, baseline, percent_only) == []
|
||||
assert compare_baseline._check_regressions(current, baseline, absolute_only) == []
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, both)
|
||||
assert len(failures) == 1
|
||||
assert "latency regressed by 6.0% and 0.600" in failures[0]
|
||||
|
||||
|
||||
def test_throughput_regression_uses_higher_is_better_direction():
|
||||
baseline = [{"throughput": 10.0}]
|
||||
current = {"model_id": "wan", "throughput": 9.0}
|
||||
policies = resolve_metric_policies({
|
||||
"throughput": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "throughput regressed by 10.0% and 1.000" in failures[0]
|
||||
|
||||
|
||||
def test_memory_regression_uses_metric_specific_absolute_floor():
|
||||
baseline = [{"memory": 10000.0}]
|
||||
current = {"model_id": "wan", "memory": 10600.0}
|
||||
policies = resolve_metric_policies({
|
||||
"memory": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 256.0,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "memory regressed by 6.0% and 600.0" in failures[0]
|
||||
|
||||
|
||||
def test_component_metric_can_gate_independently():
|
||||
baseline = [{"dit_time_s": 8.0}]
|
||||
current = {"model_id": "wan", "dit_time_s": 8.6}
|
||||
policies = resolve_metric_policies({
|
||||
"dit_time_s": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.25,
|
||||
}
|
||||
})
|
||||
|
||||
failures = compare_baseline._check_regressions(current, baseline, policies)
|
||||
|
||||
assert len(failures) == 1
|
||||
assert "dit_time_s regressed by 7.5% and 0.600" in failures[0]
|
||||
|
||||
|
||||
def test_informational_metric_remains_visible_without_failing():
|
||||
baseline = [{"throughput": 10.0}]
|
||||
current = {"model_id": "wan", "gpu_type": "NVIDIA L40S", "throughput": 8.0}
|
||||
policies = resolve_metric_policies({
|
||||
"throughput": {
|
||||
"threshold_percent": 0.01,
|
||||
"threshold_absolute": 0.01,
|
||||
"gated": False,
|
||||
}
|
||||
})
|
||||
|
||||
row = compare_baseline._build_summary_row(current, baseline, policies, False)
|
||||
|
||||
assert compare_baseline._check_regressions(current, baseline, policies) == []
|
||||
assert row["metrics"]["throughput"]["regression_pct"] == 20.0
|
||||
assert row["metrics"]["throughput"]["gated"] is False
|
||||
assert row["metrics"]["throughput"]["threshold_exceeded"] is True
|
||||
assert row["metrics"]["throughput"]["regressed"] is False
|
||||
assert row["threshold_exceeded_metrics"] == ["throughput"]
|
||||
assert row["failing_metrics"] == []
|
||||
|
||||
@@ -68,6 +68,8 @@ def test_summary_endpoint_returns_latest_group_status():
|
||||
assert body["count"] == 1
|
||||
assert body["status_counts"] == {"pass": 1, "fail": 0}
|
||||
assert body["rows"][0]["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert body["rows"][0]["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert body["rows"][0]["threshold_exceeded_metrics"] == ["latency", "throughput"]
|
||||
assert body["rows"][0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.performance import hf_store
|
||||
from fastvideo.performance_dashboard.service import build_latest_summary, build_trends, filter_records
|
||||
from fastvideo.tests.performance import hf_store
|
||||
|
||||
|
||||
def _record(ts, commit, latency, throughput, success=True, **metadata):
|
||||
@@ -28,16 +28,23 @@ def test_build_latest_summary_uses_previous_successful_records_for_baseline():
|
||||
_record("2026-01-03T00:00:00+00:00", "c" * 40, 11.0, 9.0),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05)
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert len(rows) == 1
|
||||
row = rows[0]
|
||||
assert row["baseline_n"] == 1
|
||||
assert row["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert row["metrics"]["latency"]["regression_pct"] == 10.0
|
||||
assert row["metrics"]["latency"]["absolute_delta"] == 1.0
|
||||
assert row["metrics"]["latency"]["threshold_percent"] == 8.0
|
||||
assert row["metrics"]["latency"]["threshold_absolute"] == 0.5
|
||||
assert row["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert row["metrics"]["latency"]["regressed"] is True
|
||||
assert row["metrics"]["throughput"]["regression_pct"] == 10.0
|
||||
assert row["status"] == "pass"
|
||||
assert row["computed_regression_status"] == "fail"
|
||||
assert row["threshold_exceeded_metrics"] == ["latency", "throughput"]
|
||||
assert row["failing_metrics"] == ["latency", "throughput"]
|
||||
|
||||
|
||||
def test_build_latest_summary_status_uses_latest_record_success_field():
|
||||
@@ -73,7 +80,7 @@ def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05, run_source="pr")
|
||||
rows = build_latest_summary(records, run_source="pr")
|
||||
|
||||
assert len(rows) == 1
|
||||
assert rows[0]["run_source"] == "pr"
|
||||
@@ -83,6 +90,60 @@ def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
|
||||
assert rows[0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
def test_build_latest_summary_requires_absolute_floor_for_computed_regression():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record(
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"b" * 40,
|
||||
10.6,
|
||||
10.0,
|
||||
regression_thresholds={
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.75,
|
||||
"gated": True,
|
||||
}
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert round(rows[0]["metrics"]["latency"]["regression_pct"], 1) == 6.0
|
||||
assert round(rows[0]["metrics"]["latency"]["absolute_delta"], 3) == 0.6
|
||||
assert rows[0]["metrics"]["latency"]["threshold_exceeded"] is False
|
||||
assert rows[0]["metrics"]["latency"]["regressed"] is False
|
||||
assert rows[0]["computed_regression_status"] == "pass"
|
||||
|
||||
|
||||
def test_build_latest_summary_separates_informational_threshold_crossing():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record(
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"b" * 40,
|
||||
10.6,
|
||||
10.0,
|
||||
regression_thresholds={
|
||||
"latency": {
|
||||
"threshold_percent": 0.05,
|
||||
"threshold_absolute": 0.5,
|
||||
"gated": False,
|
||||
}
|
||||
},
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert rows[0]["metrics"]["latency"]["threshold_exceeded"] is True
|
||||
assert rows[0]["metrics"]["latency"]["regressed"] is False
|
||||
assert rows[0]["threshold_exceeded_metrics"] == ["latency"]
|
||||
assert rows[0]["failing_metrics"] == []
|
||||
assert rows[0]["computed_regression_status"] == "pass"
|
||||
|
||||
|
||||
def test_filter_records_and_trends_preserve_metric_points():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
|
||||
@@ -28,6 +28,17 @@ STAGE_METRIC_MAP: dict[str, str] = {
|
||||
"DmdDenoisingStage": "dit_time_s",
|
||||
"DecodingStage": "vae_decode_time_s",
|
||||
}
|
||||
V2_CONFIG_SCHEMA_VERSION = 2
|
||||
V2_REQUIRED_IDENTITY_FIELDS = (
|
||||
"workload_id",
|
||||
"variant_id",
|
||||
"benchmark_version",
|
||||
)
|
||||
V2_OPTIONAL_METADATA_FIELDS = (
|
||||
"recipe",
|
||||
"metric_threshold_policy",
|
||||
"quality_metadata",
|
||||
)
|
||||
|
||||
# -- Config discovery -------------------------------------------------------
|
||||
|
||||
@@ -42,6 +53,71 @@ _BENCHMARKS_DIR = os.path.join(
|
||||
)
|
||||
|
||||
|
||||
def _has_v2_fields(cfg):
|
||||
v2_fields = V2_REQUIRED_IDENTITY_FIELDS + V2_OPTIONAL_METADATA_FIELDS
|
||||
return any(field in cfg for field in v2_fields)
|
||||
|
||||
|
||||
def _is_v2_config(cfg):
|
||||
return cfg.get("config_schema_version") == V2_CONFIG_SCHEMA_VERSION
|
||||
|
||||
|
||||
def _validate_non_empty_string(value, field, path):
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError(f"{path}: v2 identity field {field!r} must be a non-empty string")
|
||||
|
||||
|
||||
def _validate_integer(value, field, path):
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"{path}: v2 identity field {field!r} must be an integer")
|
||||
|
||||
|
||||
def _validate_benchmark_config(cfg, path="<memory>"):
|
||||
missing_common = [field for field in ("benchmark_id",) if field not in cfg]
|
||||
if missing_common:
|
||||
raise ValueError(f"{path}: missing required benchmark config fields: {', '.join(missing_common)}")
|
||||
|
||||
schema_version = cfg.get("config_schema_version")
|
||||
if schema_version is None:
|
||||
if _has_v2_fields(cfg):
|
||||
raise ValueError(f"{path}: v2 benchmark identity fields require config_schema_version=2")
|
||||
return
|
||||
|
||||
if schema_version != V2_CONFIG_SCHEMA_VERSION:
|
||||
raise ValueError(f"{path}: unsupported benchmark config_schema_version={schema_version!r}")
|
||||
|
||||
missing_v2 = [field for field in V2_REQUIRED_IDENTITY_FIELDS if field not in cfg]
|
||||
if missing_v2:
|
||||
raise ValueError(f"{path}: missing required v2 identity fields: {', '.join(missing_v2)}")
|
||||
|
||||
_validate_non_empty_string(cfg["workload_id"], "workload_id", path)
|
||||
_validate_non_empty_string(cfg["variant_id"], "variant_id", path)
|
||||
_validate_integer(cfg["benchmark_version"], "benchmark_version", path)
|
||||
|
||||
for field in V2_OPTIONAL_METADATA_FIELDS:
|
||||
if field in cfg and not isinstance(cfg[field], Mapping):
|
||||
raise ValueError(f"{path}: optional v2 metadata field {field!r} must be an object")
|
||||
|
||||
|
||||
def _config_identity_metadata(cfg):
|
||||
if not _is_v2_config(cfg):
|
||||
return {}
|
||||
metadata = {
|
||||
"config_schema_version": cfg["config_schema_version"],
|
||||
"workload_id": cfg["workload_id"],
|
||||
"variant_id": cfg["variant_id"],
|
||||
"benchmark_version": cfg["benchmark_version"],
|
||||
}
|
||||
for field in V2_OPTIONAL_METADATA_FIELDS:
|
||||
if field in cfg:
|
||||
metadata[field] = cfg[field]
|
||||
return metadata
|
||||
|
||||
|
||||
def _benchmark_display_id(cfg):
|
||||
return cfg["benchmark_id"]
|
||||
|
||||
|
||||
def _discover_benchmarks():
|
||||
"""Glob benchmark JSON configs and return list of (id, config) tuples."""
|
||||
pattern = os.path.join(_BENCHMARKS_DIR, "*.json")
|
||||
@@ -49,6 +125,7 @@ def _discover_benchmarks():
|
||||
for path in sorted(glob.glob(pattern)):
|
||||
with open(path) as f:
|
||||
cfg = json.load(f)
|
||||
_validate_benchmark_config(cfg, path)
|
||||
configs.append(cfg)
|
||||
return configs
|
||||
|
||||
@@ -102,7 +179,11 @@ def _extract_component_times(result: dict) -> dict[str, float | None]:
|
||||
logger.debug("Skipping malformed stage '%s' data: %r", stage_name, stage_data)
|
||||
continue
|
||||
stage_class = stage_data.get("stage_class", stage_name)
|
||||
metric_key = STAGE_METRIC_MAP.get(stage_class)
|
||||
component_metric = stage_data.get("component_metric")
|
||||
if isinstance(component_metric, str) and component_metric in component_times:
|
||||
metric_key = component_metric
|
||||
else:
|
||||
metric_key = STAGE_METRIC_MAP.get(stage_class)
|
||||
if metric_key is None:
|
||||
logger.debug("Unmapped stage '%s' class '%s' (%.3fs)",
|
||||
stage_name,
|
||||
@@ -219,6 +300,7 @@ def _run_benchmark(cfg):
|
||||
|
||||
results = {
|
||||
"benchmark_id": cfg["benchmark_id"],
|
||||
**_config_identity_metadata(cfg),
|
||||
"model_short_name": model_info.get("model_short_name", ""),
|
||||
"device": device_name,
|
||||
"num_gpus": init_kwargs.get("num_gpus", 1),
|
||||
@@ -231,6 +313,7 @@ def _run_benchmark(cfg):
|
||||
"max_peak_memory_mb": round(max_peak_memory, 1),
|
||||
"individual_peak_memories_mb": [round(m, 1) for m in peak_memories],
|
||||
"thresholds": thresholds,
|
||||
"regression_thresholds": cfg.get("regression_thresholds", {}),
|
||||
"commit": os.environ.get("BUILDKITE_COMMIT", ""),
|
||||
"pr_number": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
@@ -275,7 +358,7 @@ def _run_benchmark(cfg):
|
||||
@pytest.mark.parametrize(
|
||||
"cfg",
|
||||
_BENCHMARK_CONFIGS,
|
||||
ids=[c["benchmark_id"] for c in _BENCHMARK_CONFIGS],
|
||||
ids=[_benchmark_display_id(c) for c in _BENCHMARK_CONFIGS],
|
||||
)
|
||||
def test_inference_performance(cfg):
|
||||
"""Measure generation latency, peak GPU memory, and component-level timings
|
||||
|
||||
@@ -1,13 +1,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import PipelineLoggingInfo
|
||||
from fastvideo.pipelines.stages.denoising import Cosmos25AutoDenoisingStage, DenoisingStage
|
||||
from fastvideo.pipelines.stages.text_encoding import Cosmos25TextEncodingStage
|
||||
from fastvideo.tests.performance.test_inference_performance import _extract_component_times
|
||||
|
||||
|
||||
class SubclassStyleDenoisingStage(DenoisingStage):
|
||||
pass
|
||||
|
||||
|
||||
def test_extract_component_times_handles_pipeline_logging_info_object():
|
||||
logging_info = PipelineLoggingInfo()
|
||||
logging_info.add_stage_execution_time("prompt_encoding_stage", 1.25)
|
||||
logging_info.add_stage_metric("prompt_encoding_stage", "stage_class", "TextEncodingStage")
|
||||
logging_info.add_stage_metric("prompt_encoding_stage", "component_metric", "text_encoder_time_s")
|
||||
|
||||
assert _extract_component_times({"logging_info": logging_info}) == {
|
||||
"text_encoder_time_s": 1.25,
|
||||
@@ -45,6 +52,35 @@ def test_extract_component_times_uses_stage_class_for_pipeline_stage_keys():
|
||||
}
|
||||
|
||||
|
||||
def test_denoising_stage_subclasses_inherit_component_metric():
|
||||
assert SubclassStyleDenoisingStage.performance_component_metric == "dit_time_s"
|
||||
|
||||
|
||||
def test_cosmos25_direct_pipeline_stages_define_component_metrics():
|
||||
assert Cosmos25TextEncodingStage.performance_component_metric == "text_encoder_time_s"
|
||||
assert Cosmos25AutoDenoisingStage.performance_component_metric == "dit_time_s"
|
||||
|
||||
|
||||
def test_extract_component_times_uses_component_metric_for_stage_subclasses():
|
||||
result = {
|
||||
"logging_info": {
|
||||
"stages": {
|
||||
"denoising_stage": {
|
||||
"execution_time": 4.2,
|
||||
"stage_class": "CosmosDenoisingStage",
|
||||
"component_metric": "dit_time_s",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
assert _extract_component_times(result) == {
|
||||
"text_encoder_time_s": None,
|
||||
"dit_time_s": 4.2,
|
||||
"vae_decode_time_s": None,
|
||||
}
|
||||
|
||||
|
||||
def test_extract_component_times_keeps_legacy_class_name_keys():
|
||||
# Backward-compatibility check for logs produced before pipeline-unique
|
||||
# stage keys carried a separate stage_class field.
|
||||
|
||||
@@ -21,6 +21,26 @@ class FakeTokenizer:
|
||||
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
|
||||
})
|
||||
|
||||
|
||||
class FakeChatTokenizer:
|
||||
def __init__(self):
|
||||
self.last_messages = None
|
||||
self.last_kwargs = None
|
||||
|
||||
def apply_chat_template(self, messages, **kwargs):
|
||||
self.last_messages = messages
|
||||
self.last_kwargs = kwargs
|
||||
assert isinstance(messages[0], list)
|
||||
assert messages[0][0]["role"] == "system"
|
||||
assert messages[0][1]["role"] == "user"
|
||||
B = len(messages)
|
||||
seq_len = int(kwargs.get("max_length", 4))
|
||||
return TensorDict({
|
||||
"input_ids": torch.arange(B * seq_len).view(B, seq_len),
|
||||
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
|
||||
})
|
||||
|
||||
|
||||
class FakeTextEncoder(torch.nn.Module):
|
||||
def __init__(self, hidden_size=8):
|
||||
super().__init__()
|
||||
@@ -38,6 +58,14 @@ class FakeTextEncoder(torch.nn.Module):
|
||||
def id_preprocess(x: str) -> str:
|
||||
return x
|
||||
|
||||
|
||||
def chat_list_preprocess(x: str):
|
||||
return [
|
||||
{"role": "system", "content": "Describe the video."},
|
||||
{"role": "user", "content": x if x else " "},
|
||||
]
|
||||
|
||||
|
||||
def take_mean_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
# [B, T, H] -> [B, H]
|
||||
return outputs.last_hidden_state.mean(dim=1)
|
||||
@@ -156,3 +184,32 @@ def test_encode_text_does_not_force_hidden_states_for_ltx2_prefix():
|
||||
stage.encode_text("a", fastvideo_args, encoder_index=[0])
|
||||
|
||||
assert stage.text_encoders[0].last_output_hidden_states is False
|
||||
|
||||
|
||||
def test_chat_list_preprocess_output_is_not_stripped():
|
||||
fastvideo_args, hidden = make_args(num_encoders=1, text_len=5, hidden_size=8)
|
||||
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[0]
|
||||
encoder_config.is_chat_model = True
|
||||
encoder_config.treat_empty_as_dot = True
|
||||
fastvideo_args.pipeline_config.preprocess_text_funcs = (chat_list_preprocess, )
|
||||
|
||||
tokenizer = FakeChatTokenizer()
|
||||
stage = TextEncodingStage(
|
||||
text_encoders=[FakeTextEncoder(hidden_size=hidden)],
|
||||
tokenizers=[tokenizer],
|
||||
)
|
||||
|
||||
embeds, masks = stage.encode_text(
|
||||
"a robotic arm welding a metal structure",
|
||||
fastvideo_args,
|
||||
encoder_index=[0],
|
||||
return_attention_mask=True,
|
||||
)
|
||||
|
||||
assert embeds[0].shape == (1, hidden)
|
||||
assert masks[0].shape == (1, 5)
|
||||
assert tokenizer.last_messages == [[
|
||||
{"role": "system", "content": "Describe the video."},
|
||||
{"role": "user", "content": "a robotic arm welding a metal structure"},
|
||||
]]
|
||||
assert tokenizer.last_kwargs["return_tensors"] == "pt"
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Minimum config to run a single training step of
|
||||
# CausalConsistencyDistillationMethod on WanCausalModel for the
|
||||
# per-method smoke test. Uses the real Wan 2.1 1.3B checkpoint for both
|
||||
# the trainable student and the frozen teacher (AR Euler-step target),
|
||||
# with tiny synthetic latents.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
ema:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
|
||||
discrete_cd_N: 12
|
||||
guidance_scale: 3.0
|
||||
ema_decay: 0.95
|
||||
|
||||
training:
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
|
||||
data:
|
||||
seed: 42
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
num_latent_t: 6
|
||||
num_height: 64
|
||||
num_width: 64
|
||||
num_frames: 21
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,46 @@
|
||||
# Minimum config to run a single training step of
|
||||
# DiffusionForcingSFTMethod on a frame-wise WanCausalModel
|
||||
# (num_frames_per_block=1, chunk_size=1) for the per-method smoke
|
||||
# test. Uses the real Wan 2.1 1.3B checkpoint with tiny synthetic
|
||||
# latents.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
num_frames_per_block: 1
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 1
|
||||
|
||||
training:
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
|
||||
data:
|
||||
seed: 42
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
num_latent_t: 6
|
||||
num_height: 64
|
||||
num_width: 64
|
||||
num_frames: 21
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,45 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Minimum config to run a single training step of
|
||||
# TeacherForcingSFTMethod on WanCausalModel for the per-method smoke
|
||||
# test. Identical to wan_causal_t2v_dfsft_min.yaml except the method,
|
||||
# which feeds clean history to the causal transformer (teacher forcing).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
dit_precision: bf16
|
||||
|
||||
distributed:
|
||||
num_gpus: 1
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
|
||||
data:
|
||||
seed: 42
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
num_latent_t: 6
|
||||
num_height: 64
|
||||
num_width: 64
|
||||
num_frames: 21
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
pipeline: {}
|
||||
@@ -0,0 +1,147 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-method GPU smoke test: ``WanCausalModel`` + ``CausalConsistencyDistillationMethod``.
|
||||
|
||||
Mirrors ``test_wan_causal_dfsft.py``. Causal consistency distillation
|
||||
bootstraps a consistency MSE between the student's ``x0`` at ``t`` and an EMA
|
||||
copy of the student at ``t_next``, where ``t_next`` is produced online by a
|
||||
single CFG Euler step of a frozen teacher (all under clean-history teacher
|
||||
forcing). This test exercises the full step: finite loss, nonzero student
|
||||
gradients, frozen teacher, and a post-step EMA update.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29519")
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.consistency_model.causal_cd import (
|
||||
CausalConsistencyDistillationMethod, )
|
||||
from fastvideo.train.models.wan import WanCausalModel
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
|
||||
|
||||
_FIXTURE = str(
|
||||
Path(__file__).resolve().parent.parent / "fixtures"
|
||||
/ "wan_causal_t2v_causal_cd_min.yaml")
|
||||
|
||||
|
||||
def _build_synthetic_batch(
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
batch_size = 1
|
||||
return {
|
||||
"text_embedding":
|
||||
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
|
||||
"text_attention_mask":
|
||||
torch.ones(batch_size, 16, device=device),
|
||||
"vae_latent":
|
||||
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_causal_cd_single_train_step(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("requires CUDA")
|
||||
|
||||
cfg = load_run_config(_FIXTURE)
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.train.utils.dataloader."
|
||||
"build_parquet_t2v_train_dataloader",
|
||||
lambda *args, **kwargs: None,
|
||||
)
|
||||
|
||||
student = WanCausalModel(
|
||||
init_from=cfg.models["student"]["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=True,
|
||||
)
|
||||
student.transformer = student.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
teacher = WanCausalModel(
|
||||
init_from=cfg.models["teacher"]["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=False,
|
||||
)
|
||||
teacher.transformer = teacher.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
ema = WanCausalModel(
|
||||
init_from=cfg.models["ema"]["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=False,
|
||||
)
|
||||
ema.transformer = ema.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
method = CausalConsistencyDistillationMethod(
|
||||
cfg=cfg,
|
||||
role_models={"student": student, "teacher": teacher, "ema": ema},
|
||||
)
|
||||
method.on_train_start()
|
||||
|
||||
batch = _build_synthetic_batch(device, dtype)
|
||||
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
|
||||
|
||||
loss = loss_map["total_loss"]
|
||||
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
|
||||
assert torch.isfinite(loss).item(), (
|
||||
f"total_loss is not finite: {loss.item()}")
|
||||
|
||||
method.backward(loss_map, outputs, grad_accum_rounds=1)
|
||||
|
||||
blocks = student.transformer.blocks
|
||||
assert blocks is not None and len(blocks) > 0
|
||||
layer0 = blocks[0]
|
||||
trainable = [p for p in layer0.parameters() if p.requires_grad]
|
||||
assert len(trainable) > 0, "student layer 0 has no trainable parameters"
|
||||
for i, p in enumerate(trainable):
|
||||
assert p.grad is not None, f"student layer 0 param[{i}] has None grad"
|
||||
assert torch.isfinite(p.grad).all().item(), (
|
||||
f"student layer 0 param[{i}] grad contains NaN/Inf")
|
||||
assert any(p.grad.detach().float().norm().item() > 0.0 for p in trainable), (
|
||||
"all student layer-0 grads are exactly zero; consistency loss "
|
||||
"did not reach the first transformer block")
|
||||
|
||||
# Teacher must stay frozen.
|
||||
assert all(not p.requires_grad for p in teacher.transformer.parameters()), (
|
||||
"teacher must be frozen for Causal-CD")
|
||||
|
||||
# The EMA model and student start from the same checkpoint, so the first
|
||||
# parameter must match before any update. FSDP fully_shard params are
|
||||
# DTensors; compare the local shards (torch.equal is unsupported on
|
||||
# DTensor).
|
||||
def _local(p: torch.Tensor) -> torch.Tensor:
|
||||
return p.to_local() if hasattr(p, "to_local") else p
|
||||
|
||||
ema_param = next(ema.transformer.parameters())
|
||||
student_param = next(student.transformer.parameters())
|
||||
assert torch.equal(_local(ema_param), _local(student_param)), (
|
||||
"EMA model should start identical to the student (same checkpoint)")
|
||||
|
||||
# The EMA update must move EMA toward the student. Apply a visibly large
|
||||
# perturbation so the bf16 lerp is well above rounding noise (the real
|
||||
# optimizer step at lr=2e-6 would be sub-ULP in bf16).
|
||||
with torch.no_grad():
|
||||
student_param.add_(1.0)
|
||||
before = _local(ema_param).detach().float().clone()
|
||||
method._update_ema()
|
||||
after = _local(ema_param).detach().float()
|
||||
assert not torch.equal(before, after), (
|
||||
"EMA weights did not move after _update_ema")
|
||||
# EMA = decay*ema + (1-decay)*student moves ~ (1-decay) of the gap.
|
||||
expected = before + (1.0 - method._ema_decay) * (
|
||||
_local(student_param).detach().float() - before)
|
||||
assert torch.allclose(after, expected, atol=1e-2), (
|
||||
"EMA update did not follow the expected lerp")
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-method GPU smoke test: frame-wise ``WanCausalModel`` + ``DiffusionForcingSFTMethod``.
|
||||
|
||||
Mirrors ``test_wan_causal_dfsft.py`` but with a block size of 1 frame
|
||||
(``num_frames_per_block=1`` on the model, ``chunk_size=1`` on the method),
|
||||
so each frame gets its own independent noise level. The test asserts the
|
||||
override took effect and runs one train step: forward, finite loss, and
|
||||
nonzero gradients reaching the first transformer block.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29520")
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.fine_tuning.dfsft import (
|
||||
DiffusionForcingSFTMethod, )
|
||||
from fastvideo.train.models.wan import WanCausalModel
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
|
||||
|
||||
_FIXTURE = str(
|
||||
Path(__file__).resolve().parent.parent / "fixtures"
|
||||
/ "wan_causal_t2v_dfsft_framewise_min.yaml")
|
||||
|
||||
|
||||
def _build_synthetic_batch(
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
batch_size = 1
|
||||
return {
|
||||
"text_embedding":
|
||||
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
|
||||
"text_attention_mask":
|
||||
torch.ones(batch_size, 16, device=device),
|
||||
"vae_latent":
|
||||
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_causal_dfsft_framewise_single_train_step(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("requires CUDA")
|
||||
|
||||
cfg = load_run_config(_FIXTURE)
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.train.utils.dataloader."
|
||||
"build_parquet_t2v_train_dataloader",
|
||||
lambda *args, **kwargs: None,
|
||||
)
|
||||
|
||||
student_cfg = cfg.models["student"]
|
||||
model = WanCausalModel(
|
||||
init_from=student_cfg["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=True,
|
||||
num_frames_per_block=student_cfg.get("num_frames_per_block"),
|
||||
)
|
||||
assert model.transformer.num_frame_per_block == 1, (
|
||||
"frame-wise override did not reach the transformer")
|
||||
model.transformer = model.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
method = DiffusionForcingSFTMethod(
|
||||
cfg=cfg,
|
||||
role_models={"student": model},
|
||||
)
|
||||
method.on_train_start()
|
||||
|
||||
batch = _build_synthetic_batch(device, dtype)
|
||||
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
|
||||
|
||||
loss = loss_map["total_loss"]
|
||||
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
|
||||
assert torch.isfinite(loss).item(), (
|
||||
f"total_loss is not finite: {loss.item()}")
|
||||
|
||||
method.backward(loss_map, outputs, grad_accum_rounds=1)
|
||||
|
||||
blocks = getattr(model.transformer, "blocks", None)
|
||||
assert blocks is not None and len(blocks) > 0
|
||||
layer0 = blocks[0]
|
||||
|
||||
trainable = [p for p in layer0.parameters() if p.requires_grad]
|
||||
assert len(trainable) > 0, "layer 0 has no trainable parameters"
|
||||
|
||||
for i, p in enumerate(trainable):
|
||||
assert p.grad is not None, f"layer 0 param[{i}] has None grad"
|
||||
assert torch.isfinite(p.grad).all().item(), (
|
||||
f"layer 0 param[{i}] grad contains NaN/Inf")
|
||||
|
||||
any_nonzero = any(
|
||||
p.grad.detach().float().norm().item() > 0.0 for p in trainable)
|
||||
assert any_nonzero, (
|
||||
"all layer-0 grads are exactly zero; backward did not "
|
||||
"reach the first transformer block")
|
||||
@@ -0,0 +1,113 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-method GPU smoke test: ``WanCausalModel`` + ``TeacherForcingSFTMethod``.
|
||||
|
||||
Mirrors ``test_wan_causal_dfsft.py``. Teacher forcing concatenates a clean
|
||||
context copy of every frame inside the causal transformer (``clean_x``) and
|
||||
denoises the current block while attending to *clean* history. This test
|
||||
exercises the ``clean_x`` path end-to-end: forward, finite loss, and nonzero
|
||||
gradients reaching the first transformer block.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29518")
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.fine_tuning.tfsft import (
|
||||
TeacherForcingSFTMethod, )
|
||||
from fastvideo.train.models.wan import WanCausalModel
|
||||
from fastvideo.train.utils.config import load_run_config
|
||||
|
||||
|
||||
_FIXTURE = str(
|
||||
Path(__file__).resolve().parent.parent / "fixtures"
|
||||
/ "wan_causal_t2v_tfsft_min.yaml")
|
||||
|
||||
|
||||
def _build_synthetic_batch(
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
batch_size = 1
|
||||
return {
|
||||
"text_embedding":
|
||||
torch.randn(batch_size, 16, 4096, device=device, dtype=dtype),
|
||||
"text_attention_mask":
|
||||
torch.ones(batch_size, 16, device=device),
|
||||
"vae_latent":
|
||||
torch.randn(batch_size, 16, 6, 8, 8, device=device, dtype=dtype),
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_wan_causal_tfsft_single_train_step(
|
||||
monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("requires CUDA")
|
||||
|
||||
cfg = load_run_config(_FIXTURE)
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
|
||||
monkeypatch.setattr(
|
||||
"fastvideo.train.utils.dataloader."
|
||||
"build_parquet_t2v_train_dataloader",
|
||||
lambda *args, **kwargs: None,
|
||||
)
|
||||
|
||||
model = WanCausalModel(
|
||||
init_from=cfg.models["student"]["init_from"],
|
||||
training_config=cfg.training,
|
||||
trainable=True,
|
||||
)
|
||||
model.transformer = model.transformer.to(device=device, dtype=dtype)
|
||||
|
||||
method = TeacherForcingSFTMethod(
|
||||
cfg=cfg,
|
||||
role_models={"student": model},
|
||||
)
|
||||
method.on_train_start()
|
||||
|
||||
batch = _build_synthetic_batch(device, dtype)
|
||||
loss_map, outputs, _metrics = method.single_train_step(batch, iteration=0)
|
||||
|
||||
loss = loss_map["total_loss"]
|
||||
assert torch.is_tensor(loss), "total_loss must be a torch.Tensor"
|
||||
assert torch.isfinite(loss).item(), (
|
||||
f"total_loss is not finite: {loss.item()}")
|
||||
|
||||
method.backward(loss_map, outputs, grad_accum_rounds=1)
|
||||
|
||||
blocks = getattr(model.transformer, "blocks", None)
|
||||
assert blocks is not None and len(blocks) > 0, (
|
||||
"CausalWanTransformer is expected to expose ``.blocks``")
|
||||
layer0 = blocks[0]
|
||||
|
||||
trainable = [p for p in layer0.parameters() if p.requires_grad]
|
||||
assert len(trainable) > 0, "layer 0 has no trainable parameters"
|
||||
|
||||
for i, p in enumerate(trainable):
|
||||
assert p.grad is not None, f"layer 0 param[{i}] has None grad"
|
||||
assert torch.isfinite(p.grad).all().item(), (
|
||||
f"layer 0 param[{i}] grad contains NaN/Inf")
|
||||
|
||||
any_nonzero = any(
|
||||
p.grad.detach().float().norm().item() > 0.0 for p in trainable)
|
||||
assert any_nonzero, (
|
||||
"all layer-0 grads are exactly zero; backward did not "
|
||||
"reach the first transformer block")
|
||||
|
||||
# Teacher forcing must build its own (concatenated) attention mask and
|
||||
# must not have constructed the diffusion-forcing mask.
|
||||
assert model.transformer.teacher_forcing_block_mask is not None, (
|
||||
"teacher-forcing mask was not constructed")
|
||||
assert model.transformer.block_mask is None, (
|
||||
"diffusion-forcing mask should not be built on the TF path")
|
||||
@@ -456,11 +456,16 @@ class ValidationCallback(Callback):
|
||||
None,
|
||||
)
|
||||
|
||||
loaded_modules: dict[str, Any] = {"transformer": transformer}
|
||||
# Distillation methods build the flow-match scheduler their few-step DMD
|
||||
# sampler needs; inject it so the pipeline doesn't fall back to UniPC.
|
||||
method_scheduler = getattr(self.method, "_sf_scheduler", None)
|
||||
if method_scheduler is not None:
|
||||
loaded_modules["scheduler"] = method_scheduler
|
||||
|
||||
kwargs: dict[str, Any] = {
|
||||
"inference_mode": True,
|
||||
"loaded_modules": {
|
||||
"transformer": transformer,
|
||||
},
|
||||
"loaded_modules": loaded_modules,
|
||||
"tp_size": tc.distributed.tp_size,
|
||||
"sp_size": tc.distributed.sp_size,
|
||||
"num_gpus": tc.distributed.num_gpus,
|
||||
@@ -477,12 +482,6 @@ class ValidationCallback(Callback):
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
scheduler = self._pipeline.get_module("scheduler")
|
||||
if (scheduler is not None and type(scheduler).__name__ == "SelfForcingFlowMatchScheduler"):
|
||||
scheduler.sigma_min = 0.0
|
||||
scheduler.extra_one_step = True
|
||||
scheduler.set_timesteps(num_inference_steps=1000, training=True)
|
||||
|
||||
self._pipeline_key = key
|
||||
return self._pipeline
|
||||
|
||||
|
||||
@@ -9,6 +9,8 @@ __all__ = [
|
||||
"KDMethod",
|
||||
"SelfForcingMethod",
|
||||
"DiffusionForcingSFTMethod",
|
||||
"TeacherForcingSFTMethod",
|
||||
"CausalConsistencyDistillationMethod",
|
||||
]
|
||||
|
||||
|
||||
@@ -28,4 +30,11 @@ def __getattr__(name: str) -> object:
|
||||
if name == "DiffusionForcingSFTMethod":
|
||||
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
|
||||
return DiffusionForcingSFTMethod
|
||||
if name == "TeacherForcingSFTMethod":
|
||||
from fastvideo.train.methods.fine_tuning.tfsft import TeacherForcingSFTMethod
|
||||
return TeacherForcingSFTMethod
|
||||
if name == "CausalConsistencyDistillationMethod":
|
||||
from fastvideo.train.methods.consistency_model.causal_cd import (
|
||||
CausalConsistencyDistillationMethod, )
|
||||
return CausalConsistencyDistillationMethod
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -1,3 +1,22 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
__all__: list[str] = []
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.train.methods.consistency_model.causal_cd import (
|
||||
CausalConsistencyDistillationMethod, )
|
||||
|
||||
__all__ = [
|
||||
"CausalConsistencyDistillationMethod",
|
||||
]
|
||||
|
||||
|
||||
def __getattr__(name: str) -> object:
|
||||
if name == "CausalConsistencyDistillationMethod":
|
||||
from fastvideo.train.methods.consistency_model.causal_cd import (
|
||||
CausalConsistencyDistillationMethod, )
|
||||
|
||||
return CausalConsistencyDistillationMethod
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Causal consistency distillation method (algorithm layer)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (SelfForcingFlowMatchScheduler)
|
||||
from fastvideo.train.methods.base import LogScalar, TrainingMethod
|
||||
from fastvideo.train.models.base import ModelBase
|
||||
from fastvideo.train.utils.checkpoint import _FullModelState
|
||||
from fastvideo.train.utils.optimizer import build_optimizer_and_scheduler
|
||||
|
||||
|
||||
class CausalConsistencyDistillationMethod(TrainingMethod):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
cfg: Any,
|
||||
role_models: dict[str, ModelBase],
|
||||
) -> None:
|
||||
super().__init__(cfg=cfg, role_models=role_models)
|
||||
|
||||
for role in ("student", "teacher", "ema"):
|
||||
if role not in role_models:
|
||||
raise ValueError(f"Causal-CD requires role {role!r} "
|
||||
"(student trainable; teacher + ema frozen, "
|
||||
"both initialized from the student's "
|
||||
"checkpoint)")
|
||||
if not self.student._trainable:
|
||||
raise ValueError("Causal-CD requires student to be trainable")
|
||||
self.teacher = role_models["teacher"]
|
||||
self.ema_model = role_models["ema"]
|
||||
|
||||
self._attn_kind = self._infer_attn_kind()
|
||||
self._guidance_scale = float(self.method_config.get("guidance_scale", 3.0))
|
||||
self._discrete_cd_n = int(self.method_config.get("discrete_cd_N", 48))
|
||||
if self._discrete_cd_n < 2:
|
||||
raise ValueError("method.discrete_cd_N must be >= 2")
|
||||
self._ema_decay = float(self.method_config.get("ema_decay", 0.99))
|
||||
self._ema_start_step = int(self.method_config.get("ema_start_step", 200))
|
||||
shift = getattr(self.training_config.pipeline_config, "flow_shift", None)
|
||||
self._flow_shift = float(shift) if shift else 5.0
|
||||
|
||||
self.student.init_preprocessors(self.training_config)
|
||||
self._sf_scheduler = SelfForcingFlowMatchScheduler(
|
||||
num_inference_steps=self._discrete_cd_n,
|
||||
num_train_timesteps=int(self.student.num_train_timesteps),
|
||||
shift=self._flow_shift,
|
||||
sigma_min=0.0,
|
||||
sigma_max=1.0,
|
||||
extra_one_step=True,
|
||||
training=False,
|
||||
)
|
||||
self._init_optimizers_and_schedulers()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@property
|
||||
def _optimizer_dict(self) -> dict[str, Any]:
|
||||
return {"student": self._student_optimizer}
|
||||
|
||||
@property
|
||||
def _lr_scheduler_dict(self) -> dict[str, Any]:
|
||||
return {"student": self._student_lr_scheduler}
|
||||
|
||||
def get_optimizers(self, iteration: int) -> list[torch.optim.Optimizer]:
|
||||
del iteration
|
||||
return [self._student_optimizer]
|
||||
|
||||
def get_lr_schedulers(self, iteration: int) -> list[Any]:
|
||||
del iteration
|
||||
return [self._student_lr_scheduler]
|
||||
|
||||
def checkpoint_state(self) -> dict[str, Any]:
|
||||
# The EMA role is frozen (so the base class skips it) but mutated by
|
||||
# _update_ema every step; without persisting it a resume reloads the
|
||||
# EMA from init_from and the consistency target snaps back to the
|
||||
# base checkpoint. Mirrors DiffusionNFT's frozen "old" role.
|
||||
states = super().checkpoint_state()
|
||||
states["roles.ema.transformer"] = _FullModelState(self.ema_model.transformer)
|
||||
return states
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def single_train_step(
|
||||
self,
|
||||
batch: dict[str, Any],
|
||||
iteration: int,
|
||||
) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, LogScalar]]:
|
||||
del iteration
|
||||
training_batch = self.student.prepare_batch(
|
||||
batch,
|
||||
generator=self.cuda_generator,
|
||||
latents_source="data",
|
||||
)
|
||||
clean_latents = training_batch.latents
|
||||
if not torch.is_tensor(clean_latents) or clean_latents.ndim != 5:
|
||||
raise ValueError("Causal-CD expects [B, T, C, H, W] latents")
|
||||
|
||||
batch_size, num_latents = int(clean_latents.shape[0]), int(clean_latents.shape[1])
|
||||
device = clean_latents.device
|
||||
|
||||
sigmas = self._sf_scheduler.sigmas.to(device)
|
||||
timesteps = self._sf_scheduler.timesteps.to(device)
|
||||
idx = torch.randint(0, self._discrete_cd_n - 1, (1, ), generator=self.cuda_generator, device=device).squeeze(0)
|
||||
t, t_next = timesteps[idx], timesteps[idx + 1]
|
||||
sigma_t, sigma_t_next = sigmas[idx], sigmas[idx + 1]
|
||||
t_pf = t * torch.ones(batch_size, num_latents, device=device)
|
||||
t_next_pf = t_next * torch.ones(batch_size, num_latents, device=device)
|
||||
|
||||
noise = torch.randn(
|
||||
clean_latents.shape,
|
||||
generator=self.cuda_generator,
|
||||
device=device,
|
||||
dtype=clean_latents.dtype,
|
||||
)
|
||||
latent_t = (1.0 - sigma_t) * clean_latents + sigma_t * noise
|
||||
|
||||
# Set before any forward: predict_noise feeds batch.timesteps into
|
||||
# set_forward_context (VSA sparsity gating), so the teacher CFG
|
||||
# passes below must not see the stale timesteps from prepare_batch.
|
||||
training_batch.timesteps = t_pf
|
||||
|
||||
with torch.no_grad():
|
||||
v_cond = self._predict_flow(self.teacher,
|
||||
latent_t,
|
||||
t_pf,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
clean_x=clean_latents)
|
||||
v_uncond = self._predict_flow(self.teacher,
|
||||
latent_t,
|
||||
t_pf,
|
||||
training_batch,
|
||||
conditional=False,
|
||||
clean_x=clean_latents)
|
||||
v_pred = v_uncond + self._guidance_scale * (v_cond - v_uncond)
|
||||
dt = ((t - t_next) / float(self.student.num_train_timesteps))
|
||||
latent_t_next = latent_t - dt * v_pred
|
||||
|
||||
flow_student = self._predict_flow(self.student,
|
||||
latent_t,
|
||||
t_pf,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
clean_x=clean_latents)
|
||||
x0_t = latent_t - sigma_t * flow_student
|
||||
|
||||
with torch.no_grad():
|
||||
flow_ema = self._predict_flow(self.ema_model,
|
||||
latent_t_next,
|
||||
t_next_pf,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
clean_x=clean_latents)
|
||||
x0_t_next = latent_t_next - sigma_t_next * flow_ema
|
||||
|
||||
loss = F.mse_loss(x0_t.float(), x0_t_next.float())
|
||||
|
||||
loss_map = {"total_loss": loss, "causal_cd_loss": loss}
|
||||
attn_metadata = (training_batch.attn_metadata_vsa if self._attn_kind == "vsa" else training_batch.attn_metadata)
|
||||
outputs: dict[str, Any] = {"_fv_backward": (t_pf, attn_metadata)}
|
||||
metrics: dict[str, LogScalar] = {}
|
||||
return loss_map, outputs, metrics
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def backward(
|
||||
self,
|
||||
loss_map: dict[str, torch.Tensor],
|
||||
outputs: dict[str, Any],
|
||||
*,
|
||||
grad_accum_rounds: int = 1,
|
||||
) -> None:
|
||||
grad_accum_rounds = max(1, int(grad_accum_rounds))
|
||||
ctx = outputs.get("_fv_backward")
|
||||
if ctx is None:
|
||||
super().backward(loss_map, outputs, grad_accum_rounds=grad_accum_rounds)
|
||||
return
|
||||
self.student.backward(loss_map["total_loss"], ctx, grad_accum_rounds=grad_accum_rounds)
|
||||
|
||||
def optimizers_schedulers_step(self, iteration: int) -> None:
|
||||
super().optimizers_schedulers_step(iteration)
|
||||
if iteration >= self._ema_start_step:
|
||||
self._update_ema()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Internal helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def _predict_flow(
|
||||
self,
|
||||
model: ModelBase,
|
||||
latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
batch: Any,
|
||||
*,
|
||||
conditional: bool,
|
||||
clean_x: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return model.predict_noise(latents,
|
||||
timestep,
|
||||
batch,
|
||||
conditional=conditional,
|
||||
cfg_uncond=None,
|
||||
attn_kind=self._attn_kind,
|
||||
clean_x=clean_x)
|
||||
|
||||
@torch.no_grad()
|
||||
def _update_ema(self) -> None:
|
||||
decay = self._ema_decay
|
||||
for ema_p, p in zip(self.ema_model.transformer.parameters(), self.student.transformer.parameters(),
|
||||
strict=True):
|
||||
ema_p.mul_(decay).add_(p.detach().to(ema_p.dtype), alpha=1.0 - decay)
|
||||
|
||||
def _init_optimizers_and_schedulers(self) -> None:
|
||||
tc = self.training_config
|
||||
student_lr = float(tc.optimizer.learning_rate)
|
||||
if student_lr <= 0.0:
|
||||
raise ValueError("training.learning_rate must be > 0 for causal-cd")
|
||||
student_params = [p for p in self.student.transformer.parameters() if p.requires_grad]
|
||||
(
|
||||
self._student_optimizer,
|
||||
self._student_lr_scheduler,
|
||||
) = build_optimizer_and_scheduler(
|
||||
params=student_params,
|
||||
optimizer_config=tc.optimizer,
|
||||
loop_config=tc.loop,
|
||||
learning_rate=student_lr,
|
||||
betas=tc.optimizer.betas,
|
||||
scheduler_name=str(tc.optimizer.lr_scheduler),
|
||||
)
|
||||
@@ -7,10 +7,12 @@ from typing import TYPE_CHECKING
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
|
||||
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
|
||||
from fastvideo.train.methods.fine_tuning.tfsft import TeacherForcingSFTMethod
|
||||
|
||||
__all__ = [
|
||||
"DiffusionForcingSFTMethod",
|
||||
"FineTuneMethod",
|
||||
"TeacherForcingSFTMethod",
|
||||
]
|
||||
|
||||
|
||||
@@ -25,4 +27,9 @@ def __getattr__(name: str) -> object:
|
||||
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
|
||||
|
||||
return FineTuneMethod
|
||||
if name == "TeacherForcingSFTMethod":
|
||||
from fastvideo.train.methods.fine_tuning.tfsft import (
|
||||
TeacherForcingSFTMethod, )
|
||||
|
||||
return TeacherForcingSFTMethod
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -135,12 +135,11 @@ class DiffusionForcingSFTMethod(TrainingMethod):
|
||||
t_inhom.flatten(),
|
||||
)
|
||||
|
||||
pred = self.student.predict_noise(
|
||||
pred = self._predict_noise(
|
||||
noisy_latents,
|
||||
t_inhom,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
attn_kind=self._attn_kind,
|
||||
clean_latents,
|
||||
)
|
||||
|
||||
if bool(self.training_config.model.precondition_outputs):
|
||||
@@ -178,6 +177,23 @@ class DiffusionForcingSFTMethod(TrainingMethod):
|
||||
metrics: dict[str, LogScalar] = {}
|
||||
return loss_map, outputs, metrics
|
||||
|
||||
def _predict_noise(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
training_batch: Any,
|
||||
clean_latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
# Unused here; the teacher-forcing subclass overrides this to pass clean_x.
|
||||
del clean_latents
|
||||
return self.student.predict_noise(
|
||||
noisy_latents,
|
||||
timestep,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
attn_kind=self._attn_kind,
|
||||
)
|
||||
|
||||
# TrainingMethod override: backward
|
||||
def backward(
|
||||
self,
|
||||
|
||||
@@ -0,0 +1,30 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Teacher-forcing SFT method (TFSFT; algorithm layer)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.train.methods.fine_tuning.dfsft import (
|
||||
DiffusionForcingSFTMethod, )
|
||||
|
||||
|
||||
class TeacherForcingSFTMethod(DiffusionForcingSFTMethod):
|
||||
|
||||
def _predict_noise(
|
||||
self,
|
||||
noisy_latents: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
training_batch: Any,
|
||||
clean_latents: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return self.student.predict_noise(
|
||||
noisy_latents,
|
||||
timestep,
|
||||
training_batch,
|
||||
conditional=True,
|
||||
attn_kind=self._attn_kind,
|
||||
clean_x=clean_latents,
|
||||
)
|
||||
@@ -10,12 +10,6 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed.checkpoint.state_dict import (
|
||||
StateDictOptions,
|
||||
get_model_state_dict,
|
||||
set_model_state_dict,
|
||||
)
|
||||
from torch.distributed.checkpoint.stateful import Stateful
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
@@ -37,6 +31,7 @@ from fastvideo.train.methods.rl.common import (
|
||||
validation_caption,
|
||||
validation_shard_indices,
|
||||
)
|
||||
from fastvideo.train.utils.checkpoint import _FullModelState
|
||||
from fastvideo.train.utils.config import (
|
||||
get_optional_float,
|
||||
get_optional_int,
|
||||
@@ -78,31 +73,6 @@ class _DiffusionNFTEMAState:
|
||||
self._method._ema_update_count = int(update_count)
|
||||
|
||||
|
||||
class _FullModelState(Stateful):
|
||||
"""DCP wrapper that saves frozen model parameters too.
|
||||
|
||||
The shared ``ModelWrapper`` intentionally filters to ``requires_grad``
|
||||
parameters. DiffusionNFT's old policy is frozen but must be restored on
|
||||
resume, so it needs full model state.
|
||||
"""
|
||||
|
||||
def __init__(self, model: torch.nn.Module) -> None:
|
||||
self.model = model
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return get_model_state_dict(self.model) # type: ignore[no-any-return]
|
||||
|
||||
def load_state_dict(
|
||||
self,
|
||||
state_dict: dict[str, Any],
|
||||
) -> None:
|
||||
set_model_state_dict(
|
||||
self.model,
|
||||
model_state_dict=state_dict,
|
||||
options=StateDictOptions(strict=False),
|
||||
)
|
||||
|
||||
|
||||
class DiffusionNFTMethod(TrainingMethod):
|
||||
"""DiffusionNFT-style RL for diffusion models.
|
||||
|
||||
|
||||
@@ -320,6 +320,8 @@ class WanModel(ModelBase):
|
||||
conditional: bool,
|
||||
cfg_uncond: dict[str, Any] | None = None,
|
||||
attn_kind: Literal["dense", "vsa"] = "dense",
|
||||
clean_x: torch.Tensor | None = None,
|
||||
aug_t: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
device_type = self.device.type
|
||||
dtype = self._get_training_dtype()
|
||||
@@ -347,7 +349,11 @@ class WanModel(ModelBase):
|
||||
current_timestep=batch.timesteps,
|
||||
attn_metadata=attn_metadata,
|
||||
):
|
||||
input_kwargs = (self._build_distill_input_kwargs(noisy_latents, timestep, text_dict))
|
||||
input_kwargs = (self._build_distill_input_kwargs(noisy_latents,
|
||||
timestep,
|
||||
text_dict,
|
||||
clean_x=clean_x,
|
||||
aug_t=aug_t))
|
||||
transformer = self._get_transformer(timestep)
|
||||
pred_noise = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
|
||||
return pred_noise
|
||||
@@ -530,17 +536,24 @@ class WanModel(ModelBase):
|
||||
noise_input: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
text_dict: dict[str, torch.Tensor] | None,
|
||||
clean_x: torch.Tensor | None = None,
|
||||
aug_t: torch.Tensor | None = None,
|
||||
) -> dict[str, Any]:
|
||||
if text_dict is None:
|
||||
raise ValueError("text_dict cannot be None for "
|
||||
"Wan distillation")
|
||||
return {
|
||||
kwargs: dict[str, Any] = {
|
||||
"hidden_states": noise_input.permute(0, 2, 1, 3, 4),
|
||||
"encoder_hidden_states": text_dict["encoder_hidden_states"],
|
||||
"encoder_attention_mask": text_dict["encoder_attention_mask"],
|
||||
"timestep": timestep,
|
||||
"return_dict": False,
|
||||
}
|
||||
if clean_x is not None:
|
||||
# Teacher forcing: clean context latents (+ optional aug timestep).
|
||||
kwargs["clean_x"] = clean_x.permute(0, 2, 1, 3, 4)
|
||||
kwargs["aug_t"] = aug_t
|
||||
return kwargs
|
||||
|
||||
def _get_transformer(self, timestep: torch.Tensor) -> torch.nn.Module:
|
||||
return self.transformer
|
||||
|
||||
@@ -49,6 +49,7 @@ class WanCausalModel(WanModel, CausalModelBase):
|
||||
transformer_override_safetensor: str
|
||||
| None = None,
|
||||
lora: LoraConfig | dict[str, Any] | None = None,
|
||||
num_frames_per_block: int | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
init_from=init_from,
|
||||
@@ -62,6 +63,16 @@ class WanCausalModel(WanModel, CausalModelBase):
|
||||
)
|
||||
self._streaming_caches: (dict[tuple[int, str], _StreamingCaches]) = {}
|
||||
|
||||
if num_frames_per_block is not None:
|
||||
num_frames_per_block = int(num_frames_per_block)
|
||||
if not 1 <= num_frames_per_block <= 3:
|
||||
# Same bound as CausalWanTransformer3DModel's config path
|
||||
# (assert num_frame_per_block <= 3); this override must not
|
||||
# bypass it.
|
||||
raise ValueError("num_frames_per_block must be between 1 and 3, "
|
||||
f"got {num_frames_per_block}")
|
||||
self.transformer.num_frame_per_block = num_frames_per_block
|
||||
|
||||
# --- CausalModelBase override: clear_caches ---
|
||||
def clear_caches(
|
||||
self,
|
||||
|
||||
@@ -15,6 +15,13 @@ import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dcp
|
||||
from torch.distributed.checkpoint.state_dict import (
|
||||
StateDictOptions,
|
||||
get_model_state_dict,
|
||||
set_model_state_dict,
|
||||
)
|
||||
from torch.distributed.checkpoint.stateful import Stateful
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -131,6 +138,32 @@ class _RoleModuleContainer(torch.nn.Module):
|
||||
self.add_module(name, module)
|
||||
|
||||
|
||||
class _FullModelState(Stateful):
|
||||
"""DCP wrapper that saves frozen model parameters too.
|
||||
|
||||
The shared ``ModelWrapper`` intentionally filters to ``requires_grad``
|
||||
parameters. Frozen-but-mutated roles (e.g. DiffusionNFT's old policy,
|
||||
causal-CD's EMA target) must still be restored on resume, so they need
|
||||
full model state.
|
||||
"""
|
||||
|
||||
def __init__(self, model: torch.nn.Module) -> None:
|
||||
self.model = model
|
||||
|
||||
def state_dict(self) -> dict[str, Any]:
|
||||
return get_model_state_dict(self.model) # type: ignore[no-any-return]
|
||||
|
||||
def load_state_dict(
|
||||
self,
|
||||
state_dict: dict[str, Any],
|
||||
) -> None:
|
||||
set_model_state_dict(
|
||||
self.model,
|
||||
model_state_dict=state_dict,
|
||||
options=StateDictOptions(strict=False),
|
||||
)
|
||||
|
||||
|
||||
class _CallbackStateWrapper:
|
||||
"""Wraps a CallbackDict for DCP save/load."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user