Compare commits

...
25 Commits
Author SHA1 Message Date
SolitaryThinker 4d8713f6e4 [bugfix]: bound WanCausalModel num_frames_per_block override to <= 3
The constructor override only checked >= 1, bypassing the <= 3 limit the
config path enforces in CausalWanTransformer3DModel. Apply the same upper
bound with a clear error message.
2026-07-05 15:01:25 -07:00
SolitaryThinker fda43fc610 [bugfix]: refuse teacher forcing with a local attention window
_prepare_teacher_forcing_mask silently ignored local_attn_size while the
block-wise causal mask honors it, so a configured attention window was
dropped on the teacher-forcing path. Raise NotImplementedError instead of
training with a mask that contradicts the config.
2026-07-05 15:01:09 -07:00
SolitaryThinker 5a8329ddf6 [bugfix]: set causal-CD timesteps before the teacher CFG forwards
training_batch.timesteps was assigned t_pf only after the two teacher
CFG passes, so their set_forward_context(current_timestep=...) carried
the stale random timesteps from prepare_batch — wrong VSA sparsity
gating when attn_kind == "vsa". Assign before any forward.
2026-07-05 15:00:47 -07:00
SolitaryThinker fcb5b465c5 [bugfix]: checkpoint causal-CD EMA consistency target on save/resume
The base TrainingMethod.checkpoint_state() only persists trainable roles,
so CausalConsistencyDistillationMethod's EMA target (frozen but mutated by
_update_ema every step) was never saved; a resume reloaded it from
init_from, snapping the consistency target back to the base checkpoint.

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