Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3ab66290dc | ||
|
|
7fa0fed781 | ||
|
|
9096310b5c | ||
|
|
fce6ed516d | ||
|
|
82ed9fe58d | ||
|
|
3d8cc4f0a0 | ||
|
|
0557f7a7d9 | ||
|
|
dc66cd97ef | ||
|
|
6da206e196 |
@@ -76,7 +76,7 @@ EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
|
||||
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
|
||||
EFFECTIVE_PR=$PR_NUMBER
|
||||
fi
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} BUILDKITE_BUILD_URL=${BUILDKITE_BUILD_URL:-} BUILDKITE_BUILD_ID=${BUILDKITE_BUILD_ID:-} BUILDKITE_JOB_ID=${BUILDKITE_JOB_ID:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
|
||||
@@ -21,22 +21,31 @@ It serves three audiences:
|
||||
# fastvideo/tests/performance/results/
|
||||
pytest fastvideo/tests/performance/ -vs
|
||||
|
||||
# Optional: compare against the rolling HF baseline (read-only outside CI).
|
||||
# Optional: compare against the rolling HF baseline.
|
||||
# PERF_REPORTS_DIR defaults to /root/data/perf_reports for Modal/CI, so
|
||||
# override it when running outside the container.
|
||||
PERF_REPORTS_DIR=/tmp/fastvideo_perf_reports \
|
||||
python fastvideo/tests/performance/compare_baseline.py
|
||||
|
||||
# Optional: explicitly upload a passing local/manual run.
|
||||
HF_TOKEN=hf_... \
|
||||
PERF_RUN_SOURCE=local \
|
||||
PERF_UPLOAD_POLICY=pass \
|
||||
PERF_REPORTS_DIR=/tmp/fastvideo_perf_reports \
|
||||
python fastvideo/tests/performance/compare_baseline.py
|
||||
|
||||
# Optional: build the Plotly dashboard locally.
|
||||
PERF_REPORTS_DIR=/tmp/fastvideo_perf_reports \
|
||||
python fastvideo/tests/performance/dashboard.py
|
||||
```
|
||||
|
||||
The pytest run never uploads anything. `compare_baseline.py` only writes to
|
||||
the HF dataset when `TEST_SCOPE=full` *and* `BUILDKITE_BRANCH=main`, so local
|
||||
runs are always read-only. The report directory default is container-oriented;
|
||||
set `PERF_REPORTS_DIR` to a writable local path when generating dashboards or
|
||||
when you want local Markdown/normalized-result artifacts from the comparator.
|
||||
The pytest run never uploads anything. `compare_baseline.py` uploads only when
|
||||
`PERF_UPLOAD_POLICY` is set. Local uploads are explicit opt-in and require HF
|
||||
credentials. PR/direct performance runs upload passing records for dashboard
|
||||
visibility, while scheduled-main runs upload both pass and fail records. The
|
||||
report directory default is container-oriented; set `PERF_REPORTS_DIR` to a
|
||||
writable local path when generating dashboards or when you want local
|
||||
Markdown/normalized-result artifacts from the comparator.
|
||||
`compare_baseline.py` reads every `perf_*.json` currently present in
|
||||
`fastvideo/tests/performance/results/`; remove stale result files if you only
|
||||
want to compare the latest local run.
|
||||
@@ -68,7 +77,9 @@ fastvideo/tests/performance/
|
||||
|
||||
The HF dataset (`FastVideo/performance-tracking` by default) holds one
|
||||
normalized JSON per `(model_id, gpu_type, run)` tuple. The rolling baseline is
|
||||
the median of the last 5 successful records for that model+GPU.
|
||||
the median of the last 5 successful, baseline-eligible records for that
|
||||
model+GPU. PR and local records are visible in the dashboard but are not
|
||||
baseline eligible.
|
||||
|
||||
## Planned Coverage
|
||||
|
||||
@@ -94,11 +105,16 @@ Each benchmark records six metrics:
|
||||
|
||||
`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`. It maps `TextEncodingStage` to
|
||||
`text_encoder_time_s`, `DenoisingStage` and `DmdDenoisingStage` to
|
||||
`dit_time_s`, and `DecodingStage` to `vae_decode_time_s`. If a pipeline does
|
||||
not report one of those stages, that component metric is stored as `null` and
|
||||
is skipped by the static threshold and rolling baseline checks.
|
||||
`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
|
||||
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
|
||||
`null` and is skipped by the static threshold and rolling baseline checks.
|
||||
|
||||
## The two gates
|
||||
|
||||
@@ -138,16 +154,16 @@ headroom and almost never need touching.
|
||||
|
||||
### Rolling baseline (per `(model_id, gpu_type)`)
|
||||
|
||||
`compare_baseline.py` loads the last 5 successful 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
|
||||
`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.
|
||||
|
||||
This is the **drift detector** — it catches sub-threshold regressions that
|
||||
slowly add up. It only persists new records when running the full suite on
|
||||
`main`. Local and pull-request runs can compare against the HF baseline, but
|
||||
they do not update it.
|
||||
slowly add up. Only scheduled-main successful records are baseline eligible.
|
||||
Local and pull-request runs can upload dashboard-visible records, but they do
|
||||
not update future gating baselines.
|
||||
|
||||
When the baseline shifts for a legitimate reason (torch upgrade, kernel
|
||||
change, etc.) and CI starts failing, use the
|
||||
@@ -215,6 +231,8 @@ result, used as the rolling-baseline source of truth.
|
||||
Older records in the HF dataset may not have component timing fields. The
|
||||
comparator ignores missing or `null` metrics when computing a median, and the
|
||||
dashboard lists skipped plots for metric series that have no non-null values.
|
||||
Records missing both `run_source` and `baseline_eligible` are treated as legacy
|
||||
successful main/full-suite uploads and remain eligible for rolling baselines.
|
||||
|
||||
## Environment variable reference
|
||||
|
||||
@@ -224,8 +242,11 @@ dashboard lists skipped plots for metric series that have no non-null values.
|
||||
| `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` | unset | `hf_store.py` | Required for upload (main-branch full-suite only); reads work without it. |
|
||||
| `TEST_SCOPE` | unset | `compare_baseline.py` | Set to `full` together with `BUILDKITE_BRANCH=main` to enable HF persistence. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `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. |
|
||||
|
||||
@@ -11,6 +11,9 @@ This page describes the various options for speeding up generation times in Fast
|
||||
- [Sliding Tile Attention (Archived)](#sliding-tile-attention-archived)
|
||||
- [Sage Attention](#sage-attention)
|
||||
- [Sage Attention 3](#sage-attention-3)
|
||||
|
||||
- [FP8 Weight Quantization](#fp8-weight-quantization)
|
||||
|
||||
- [Adaptive Guidance (CFG gating)](#adaptive-guidance-cfg-gating)
|
||||
|
||||
- [torch.compile](#torch-compile)
|
||||
@@ -218,6 +221,50 @@ These backends are model-specific and require the corresponding kernels and
|
||||
dependencies. Use the support matrix and model examples to confirm compatibility
|
||||
before enabling them.
|
||||
|
||||
## FP8 Weight Quantization
|
||||
|
||||
**`transformer_quant="FP8"`**
|
||||
|
||||
Quantizes DiT linear layers (attention projections and FFN) to FP8 e4m3.
|
||||
|
||||
On GPUs older than sm89, the FP8 matmul falls back to a bf16 dequant path
|
||||
automatically.
|
||||
|
||||
### Requirements
|
||||
|
||||
- **GPU**: sm89+ (H100, L40S, RTX 4090, or newer) for hardware FP8 compute
|
||||
- No additional packages required beyond the base FastVideo install
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
|
||||
gen = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# Pass an instance — the bare string is not resolved on the from_pretrained path.
|
||||
transformer_quant=get_quantization_config("FP8")(), # per-tensor (default)
|
||||
# transformer_quant=get_quantization_config("FP8")(granularity="channel"), # slower, higher accuracy
|
||||
)
|
||||
gen.generate(request={"prompt": "A raccoon in sunflowers", "output": {"save_video": True}})
|
||||
```
|
||||
|
||||
Or run the example script:
|
||||
|
||||
```bash
|
||||
python examples/inference/optimizations/fp8_wan2_1_1_3b.py
|
||||
python examples/inference/optimizations/fp8_wan2_1_1_3b.py --granularity channel
|
||||
python examples/inference/optimizations/fp8_wan2_1_1_3b.py --bf16 # baseline
|
||||
```
|
||||
|
||||
### Granularity
|
||||
|
||||
| Mode | Weight scales | Activation scales | Speed | Accuracy |
|
||||
|------|--------------|-------------------|-------|----------|
|
||||
| `tensor` (default) | per-tensor | per-tensor | faster | lower |
|
||||
| `channel` | per-output-channel | per-token (rowwise) | slower | higher |
|
||||
|
||||
<a id="torch-compile"></a>
|
||||
|
||||
## torch.compile
|
||||
|
||||
@@ -0,0 +1,286 @@
|
||||
"""Fast NVFP4 linear inference for Wan2.1-T2V-1.3B with TAEHV decoding.
|
||||
|
||||
This is the FP4-linear fast path from ``fp4_linear_wan2_1_1_3b.py`` with the
|
||||
heavy Wan VAE swapped out for TAEHV -- a tiny autoencoder that decodes Wan2.1
|
||||
latents directly (no denormalization) and is dramatically faster / lighter.
|
||||
|
||||
How it works: the generator runs with ``output_type="latent"`` so the pipeline
|
||||
returns raw denoised latents instead of pixels (the Wan VAE is offloaded and
|
||||
never used). We then decode those latents with TAEHV in this script and save
|
||||
the frames ourselves. This mirrors the FastVideo-Quantization
|
||||
``quantization_example_taehv.py`` proof-of-concept, but kept clean: TAEHV is a
|
||||
pip package (no ``sys.path`` hacks), the latent->uint8 conversion is vectorized,
|
||||
and there is no dead profiler / sanitization code.
|
||||
|
||||
Requirements:
|
||||
- Blackwell GPU (B200/B300, sm100a/sm103a) for the FP4 linear path
|
||||
- flashinfer (``pip install flashinfer-python``)
|
||||
- TAEHV weights ``taew2_1.pth`` (https://github.com/madebyollin/taehv)
|
||||
|
||||
Usage:
|
||||
python fp4_linear_taehv_wan2_1_1_3b.py # FP4 + TAEHV + compile
|
||||
python fp4_linear_taehv_wan2_1_1_3b.py --no-taehv # FP4 + full Wan VAE
|
||||
python fp4_linear_taehv_wan2_1_1_3b.py --no-compile # eager
|
||||
python fp4_linear_taehv_wan2_1_1_3b.py --baseline # dense bf16 reference
|
||||
python fp4_linear_taehv_wan2_1_1_3b.py --distilled_model '' # base Wan2.1 weights
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
|
||||
import imageio
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.layers.quantization.nvfp4_qat_config import NVFP4QATConfig
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
# Distilled, quantization-aware (QAD) transformer for Wan2.1-1.3B (3 steps,
|
||||
# guidance 1.0). Loaded on top of the base Wan2.1 pipeline; pass
|
||||
# ``--distilled_model ''`` to run the base weights instead.
|
||||
DEFAULT_DISTILLED_MODEL = "FastVideo/FastWan-QAD-1.3B"
|
||||
DISTILLED_WEIGHTS_FILE = (
|
||||
"generator_inference_transformer/diffusion_pytorch_model.safetensors"
|
||||
)
|
||||
|
||||
# TAEHV checkpoint for Wan2.1. Clone https://github.com/madebyollin/taehv to get
|
||||
# ``taew2_1.pth`` (Wan 2.1 / Wan 2.2-14B / Qwen-Image all use this VAE).
|
||||
DEFAULT_TAEHV_CHECKPOINT = "/root/taehv/taew2_1.pth"
|
||||
|
||||
PROMPT = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
|
||||
class TaehvDecoder:
|
||||
"""Thin wrapper around the TAEHV tiny autoencoder for Wan2.1 latents.
|
||||
|
||||
TAEHV consumes the *normalized* latents the diffusion model produces (the
|
||||
same representation FastVideo carries internally), so no denormalization is
|
||||
needed -- unlike the full Wan VAE path.
|
||||
"""
|
||||
|
||||
def __init__(self, checkpoint_path: str, device: str = "cuda",
|
||||
dtype: torch.dtype = torch.float16) -> None:
|
||||
from taehv import TAEHV # pip-installed; no sys.path manipulation
|
||||
self.device = device
|
||||
self.dtype = dtype
|
||||
print(f"Loading TAEHV from {checkpoint_path} ...")
|
||||
self.model = TAEHV(checkpoint_path=checkpoint_path).to(device, dtype).eval()
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latents: torch.Tensor):
|
||||
"""Decode FastVideo latents into uint8 RGB frames.
|
||||
|
||||
Args:
|
||||
latents: ``[B, C, T, H, W]`` (NCTHW) normalized latent tensor.
|
||||
|
||||
Returns:
|
||||
A ``(T, H, W, 3)`` uint8 numpy array ready for ``imageio.mimsave``.
|
||||
"""
|
||||
# NCTHW -> NTCHW (TAEHV's expected layout), on the TAEHV device/dtype.
|
||||
latents = latents.permute(0, 2, 1, 3, 4).to(self.device, self.dtype)
|
||||
decoded = self.model.decode_video(
|
||||
latents, parallel=True, show_progress_bar=False)
|
||||
# decoded: [B, T, 3, H, W] in [0, 1]. Take batch 0, vectorize to uint8.
|
||||
frames = (decoded[0].clamp(0, 1) * 255).to(torch.uint8)
|
||||
return frames.permute(0, 2, 3, 1).cpu().numpy()
|
||||
|
||||
|
||||
def resolve_distilled_weights(hf_id: str) -> str:
|
||||
"""Return a local path to the distilled transformer safetensors."""
|
||||
if os.path.exists(hf_id):
|
||||
return hf_id
|
||||
from huggingface_hub import hf_hub_download
|
||||
return hf_hub_download(repo_id=hf_id, filename=DISTILLED_WEIGHTS_FILE)
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def silence_request_log():
|
||||
"""Quiet ``VideoGenerator.generate``'s per-request config printout.
|
||||
|
||||
Each ``generate(...)`` call logs a multi-line debug block (height/width/
|
||||
prompt/steps/...) at INFO via ``logger.info`` in
|
||||
``fastvideo.entrypoints.video_generator``. There is no built-in switch,
|
||||
so this context manager raises that logger's level to WARNING while the
|
||||
warmup calls run, then restores it for the timed run.
|
||||
"""
|
||||
vg_logger = logging.getLogger("fastvideo.entrypoints.video_generator")
|
||||
prev_level = vg_logger.level
|
||||
vg_logger.setLevel(logging.WARNING)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
vg_logger.setLevel(prev_level)
|
||||
|
||||
|
||||
def resolve_taehv_checkpoint(path: str) -> str:
|
||||
"""Validate the TAEHV checkpoint path, with a helpful error if missing."""
|
||||
if os.path.exists(path):
|
||||
return path
|
||||
raise FileNotFoundError(
|
||||
f"TAEHV checkpoint not found at {path!r}. Clone the weights with:\n"
|
||||
" git clone https://github.com/madebyollin/taehv\n"
|
||||
"and pass --taehv_checkpoint <repo>/taew2_1.pth")
|
||||
|
||||
|
||||
def build_generator(args: argparse.Namespace) -> VideoGenerator:
|
||||
model_id = args.model
|
||||
|
||||
# Half precision everywhere; DiT linears are additionally NVFP4-quantized
|
||||
# via dit_config.quant_config below.
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_id)
|
||||
pipeline_config.dit_precision = "bf16"
|
||||
pipeline_config.vae_precision = "bf16"
|
||||
pipeline_config.text_encoder_precisions = ("bf16",)
|
||||
|
||||
if not args.baseline:
|
||||
pipeline_config.dit_config.quant_config = NVFP4QATConfig()
|
||||
|
||||
compile_enabled = not args.no_compile
|
||||
|
||||
extra_kwargs = {}
|
||||
if args.distilled_model:
|
||||
weights_path = resolve_distilled_weights(args.distilled_model)
|
||||
print(f"Using distilled weights: {args.distilled_model} -> {weights_path}")
|
||||
extra_kwargs["init_weights_from_safetensors"] = weights_path
|
||||
|
||||
if args.taehv:
|
||||
# Skip the in-pipeline VAE decode entirely: the pipeline returns raw
|
||||
# latents, the Wan VAE is offloaded to CPU (and not compiled) since we
|
||||
# decode with TAEHV in this script instead.
|
||||
extra_kwargs["output_type"] = "latent"
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_id,
|
||||
pipeline_config=pipeline_config,
|
||||
num_gpus=args.num_gpus,
|
||||
# Keep everything resident on the GPU -- no offloading, except the
|
||||
# unused Wan VAE when TAEHV handles decoding.
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
dit_layerwise_offload=False,
|
||||
vae_cpu_offload=args.taehv,
|
||||
text_encoder_cpu_offload=False,
|
||||
pin_cpu_memory=False,
|
||||
enable_torch_compile=compile_enabled,
|
||||
enable_torch_compile_text_encoder=compile_enabled,
|
||||
enable_torch_compile_vae=compile_enabled and not args.taehv,
|
||||
**extra_kwargs,
|
||||
)
|
||||
return generator
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="FP4 linear Wan2.1-1.3B with TAEHV decoding benchmark")
|
||||
parser.add_argument("--model", default="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
help="Model path or HuggingFace ID")
|
||||
parser.add_argument("--baseline", action="store_true",
|
||||
help="Run dense bf16 instead of FP4 linear")
|
||||
parser.add_argument("--no-compile", action="store_true",
|
||||
help="Disable torch.compile (eager)")
|
||||
parser.add_argument("--taehv", action=argparse.BooleanOptionalAction,
|
||||
default=True,
|
||||
help="Decode with TAEHV instead of the full Wan VAE "
|
||||
"(use --no-taehv for the Wan VAE path)")
|
||||
parser.add_argument("--taehv_checkpoint", default=DEFAULT_TAEHV_CHECKPOINT,
|
||||
help="Path to the TAEHV taew2_1.pth checkpoint")
|
||||
parser.add_argument("--distilled_model", default=DEFAULT_DISTILLED_MODEL,
|
||||
help="HuggingFace ID (or local path) of a distilled "
|
||||
"transformer checkpoint to load on top of --model. "
|
||||
"Pass '' to use the base --model weights instead.")
|
||||
parser.add_argument("--num_gpus", type=int, default=1)
|
||||
parser.add_argument("--infer_steps", type=int, default=3)
|
||||
parser.add_argument("--guidance_scale", type=float, default=1.0)
|
||||
args = parser.parse_args()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
raise SystemExit("CUDA is required for FP4 inference.")
|
||||
|
||||
cap = torch.cuda.get_device_capability()
|
||||
print(f"GPU: {torch.cuda.get_device_name()} (capability {cap[0]}.{cap[1]})")
|
||||
if not args.baseline and cap[0] < 10:
|
||||
print("Warning: NVFP4 requires Blackwell (capability 10.0+); "
|
||||
"FP4 kernels may be unavailable on this GPU.")
|
||||
|
||||
mode = "bf16" if args.baseline else "fp4_linear"
|
||||
mode += "_taehv" if args.taehv else "_wanvae"
|
||||
if not args.no_compile:
|
||||
mode += "_compile"
|
||||
print(f"Mode: {mode.upper()}")
|
||||
|
||||
# Load TAEHV before the (slow) generator build so a bad checkpoint path
|
||||
# fails fast.
|
||||
taehv = TaehvDecoder(resolve_taehv_checkpoint(args.taehv_checkpoint)) \
|
||||
if args.taehv else None
|
||||
|
||||
generator = build_generator(args)
|
||||
|
||||
os.makedirs(OUTPUT_PATH, exist_ok=True)
|
||||
|
||||
# Warmup: with compile enabled the first call(s) pay the DiT compilation
|
||||
# cost. When using TAEHV we also decode the warmup latents so the timed
|
||||
# decode below is warm -- TAEHV's decoder is all conv/upsample, so the
|
||||
# first call otherwise pays cuDNN algo selection + allocator growth
|
||||
# (~0.2s), which is exactly the cold-start overhead we want to exclude.
|
||||
n_warmup = 2 if not args.no_compile else 1
|
||||
with silence_request_log():
|
||||
for _ in range(n_warmup):
|
||||
warm = generator.generate(request={
|
||||
"prompt": PROMPT,
|
||||
"sampling": {"num_inference_steps": 2, "guidance_scale": args.guidance_scale},
|
||||
"output": {"save_video": False, "return_frames": args.taehv},
|
||||
})
|
||||
if args.taehv:
|
||||
taehv.decode(warm.samples)
|
||||
|
||||
output_path = os.path.join(OUTPUT_PATH, f"raccoon_{mode}.mp4")
|
||||
torch.cuda.synchronize()
|
||||
start = time.perf_counter()
|
||||
result = generator.generate(request={
|
||||
"prompt": PROMPT,
|
||||
"sampling": {
|
||||
"num_inference_steps": args.infer_steps,
|
||||
"guidance_scale": args.guidance_scale,
|
||||
},
|
||||
# When using TAEHV we need the latents back and save manually; the Wan
|
||||
# VAE path lets the pipeline decode and save the mp4 itself.
|
||||
"output": {
|
||||
"save_video": not args.taehv,
|
||||
"return_frames": args.taehv,
|
||||
"output_path": output_path,
|
||||
},
|
||||
})
|
||||
torch.cuda.synchronize()
|
||||
denoise_elapsed = time.perf_counter() - start
|
||||
|
||||
if args.taehv:
|
||||
torch.cuda.synchronize()
|
||||
decode_start = time.perf_counter()
|
||||
frames = taehv.decode(result.samples)
|
||||
torch.cuda.synchronize()
|
||||
decode_elapsed = time.perf_counter() - decode_start
|
||||
|
||||
imageio.mimsave(output_path, frames, fps=16, format="mp4")
|
||||
total = denoise_elapsed + decode_elapsed
|
||||
print(f"[{mode.upper()}] denoise {denoise_elapsed:.2f}s + TAEHV decode "
|
||||
f"{decode_elapsed:.2f}s = {total:.2f}s "
|
||||
f"({frames.shape[0]} frames @ {tuple(frames.shape[1:3])})")
|
||||
print(f"Saved video to {output_path}")
|
||||
else:
|
||||
print(f"[{mode.upper()}] {args.infer_steps} steps in {denoise_elapsed:.2f}s "
|
||||
f"({args.infer_steps / denoise_elapsed:.2f} it/s)")
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,140 @@
|
||||
"""FP8 weight quantization inference example.
|
||||
|
||||
Runs Wan2.1-T2V-1.3B with FP8 e4m3 quantized DiT linear layers (attention
|
||||
projections and FFN). Weights are quantized in-place after loading; activations
|
||||
are quantized dynamically at runtime. Reduces GPU memory relative to BF16 and
|
||||
can improve throughput on sm89+ GPUs.
|
||||
|
||||
Requirements:
|
||||
- GPU: sm89+ (H100, L40S, RTX 4090, Ada Lovelace, or newer)
|
||||
Falls back to a bf16 dequant path on older GPUs.
|
||||
- TAEHV (optional): Follow install instructions at https://github.com/madebyollin/taehv
|
||||
|
||||
Usage:
|
||||
python fp8_wan2_1_1_3b.py # FP8 per-tensor (default)
|
||||
python fp8_wan2_1_1_3b.py --bf16 # BF16 baseline
|
||||
python fp8_wan2_1_1_3b.py --granularity channel # per-channel (higher accuracy but slower)
|
||||
python fp8_wan2_1_1_3b.py --taehv-checkpoint /path/to/taew2_1.pth
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
import torch
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
|
||||
def load_taehv(checkpoint_path, device="cuda", dtype=torch.float16):
|
||||
repo_dir = os.path.dirname(checkpoint_path)
|
||||
if repo_dir not in sys.path:
|
||||
sys.path.insert(0, repo_dir)
|
||||
from taehv import TAEHV
|
||||
print(f"Loading TAEHV from {checkpoint_path}...")
|
||||
model = TAEHV(checkpoint_path=checkpoint_path).to(device, dtype)
|
||||
print("TAEHV loaded.")
|
||||
return model
|
||||
|
||||
|
||||
@torch.no_grad() # type: ignore[misc]
|
||||
def decode_with_taehv(taehv_model, latents):
|
||||
latents = latents.permute(0, 2, 1, 3, 4)
|
||||
latents = latents.to(device=next(taehv_model.parameters()).device,
|
||||
dtype=next(taehv_model.parameters()).dtype)
|
||||
decoded = taehv_model.decode_video(latents, parallel=False, show_progress_bar=False)
|
||||
frames = []
|
||||
for frame in decoded[0]:
|
||||
frame_np = (frame.clamp(0, 1) * 255).byte().cpu().permute(1, 2, 0).numpy()
|
||||
frames.append(frame_np)
|
||||
return frames
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FP8 video generation benchmark")
|
||||
parser.add_argument("--bf16", action="store_true",
|
||||
help="BF16 baseline (no FP8 quantization)")
|
||||
parser.add_argument("--granularity", choices=["tensor", "channel"], default="tensor",
|
||||
help="FP8 weight scale granularity: tensor (faster) or channel (more accurate)")
|
||||
parser.add_argument("--taehv-checkpoint", default=None, metavar="PATH",
|
||||
help="Path to taew2_1.pth; enables TAEHV tiny autoencoder decoding")
|
||||
parser.add_argument("--model", default="FastVideo/FastWan-QAD-FP8-1.3B",
|
||||
help="Model path or HuggingFace ID")
|
||||
parser.add_argument("--no-compile", action="store_true", help="Disable torch.compile for the DiT")
|
||||
parser.add_argument("--num_gpus", type=int, default=1)
|
||||
parser.add_argument("--infer_steps", type=int, default=3)
|
||||
args = parser.parse_args()
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "SAGE_ATTN")
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
|
||||
mode = "bf16" if args.bf16 else f"fp8_{args.granularity}"
|
||||
if not args.no_compile:
|
||||
mode += "_compile"
|
||||
use_taehv = args.taehv_checkpoint is not None
|
||||
print(f"Mode: {mode.upper()}" + (" decoder=TAEHV" if use_taehv else " decoder=VAE"))
|
||||
|
||||
taehv_model = load_taehv(args.taehv_checkpoint) if use_taehv else None
|
||||
|
||||
# transformer_quant needs a QuantizationConfig *instance* — the bare string
|
||||
# is not resolved on the from_pretrained kwarg path.
|
||||
extra = {} if args.bf16 else {
|
||||
"transformer_quant": get_quantization_config("FP8")(granularity=args.granularity)
|
||||
}
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
args.model,
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
dit_layerwise_offload=False,
|
||||
vae_cpu_offload=use_taehv,
|
||||
text_encoder_cpu_offload=False,
|
||||
pin_cpu_memory=False,
|
||||
enable_torch_compile=not args.no_compile,
|
||||
enable_torch_compile_vae=not args.no_compile and not use_taehv,
|
||||
output_type="latent" if use_taehv else "pil",
|
||||
**extra,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
|
||||
n_warmup = 1 if not args.no_compile else 0
|
||||
for _ in range(n_warmup):
|
||||
generator.generate(request={"prompt": prompt, "sampling": {"num_inference_steps": 3, "guidance_scale": 1.0},
|
||||
"output": {"save_video": False}})
|
||||
|
||||
os.makedirs(OUTPUT_PATH, exist_ok=True)
|
||||
start = time.time()
|
||||
if use_taehv:
|
||||
result = generator.generate(request={
|
||||
"prompt": prompt,
|
||||
"sampling": {"num_inference_steps": args.infer_steps, "guidance_scale": 1.0},
|
||||
"output": {"save_video": False},
|
||||
})
|
||||
import imageio
|
||||
frames = decode_with_taehv(taehv_model, result.samples)
|
||||
video_path = os.path.join(OUTPUT_PATH, f"raccoon_{mode}.mp4")
|
||||
imageio.mimsave(video_path, frames, fps=16, format="mp4")
|
||||
print(f"Saved TAEHV-decoded video to: {video_path}")
|
||||
else:
|
||||
generator.generate(request={
|
||||
"prompt": prompt,
|
||||
"sampling": {"num_inference_steps": args.infer_steps, "guidance_scale": 1.0},
|
||||
"output": {"save_video": True, "output_path": os.path.join(OUTPUT_PATH, f"raccoon_{mode}.mp4")},
|
||||
})
|
||||
elapsed = time.time() - start
|
||||
print(f"[{mode.upper()}] {args.infer_steps} steps in {elapsed:.2f}s "
|
||||
f"({args.infer_steps / elapsed:.2f} it/s)")
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -103,6 +103,11 @@ if [ "${GPU_BACKEND}" = "CUDA" ]; then
|
||||
if [ -z "${TORCH_CUDA_ARCH_LIST:-}" ]; then
|
||||
if [ "${cc_major}" = "9" ] && [ "${cc_minor}" = "0" ]; then
|
||||
export TORCH_CUDA_ARCH_LIST="9.0a"
|
||||
elif [ "${cc_major}" = "12" ] && [ "${cc_minor}" = "0" ]; then
|
||||
# Blackwell sm_120 needs the arch-conditional 'a' suffix so CMake's
|
||||
# AUTO gate (matches 12.0a/120a/sm_120a) builds the attn_qat_infer
|
||||
# (modified SageAttention3 FP4) kernels instead of silently skipping.
|
||||
export TORCH_CUDA_ARCH_LIST="12.0a"
|
||||
else
|
||||
export TORCH_CUDA_ARCH_LIST="${cc_major}.${cc_minor}"
|
||||
fi
|
||||
|
||||
@@ -55,6 +55,20 @@ class SageAttention3Impl(AttentionImpl):
|
||||
self.softmax_scale = softmax_scale
|
||||
self.dropout = extra_impl_args.get("dropout_p", 0.0)
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Transpose stacked QKV from [3B, L, H, D] to [3B, H, L, D].
|
||||
|
||||
Single bulk permute+contiguous on the entire stacked tensor rather than
|
||||
three separate transposed views for Q, K, V. The .contiguous() is
|
||||
required: sageattn_blackwell's fake kernel returns empty_like(q), so the
|
||||
op's output strides must match contiguous q under torch.compile.
|
||||
"""
|
||||
return qkv.permute(0, 2, 1, 3).contiguous()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
@@ -62,9 +76,15 @@ class SageAttention3Impl(AttentionImpl):
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
"""Call sageattn3_blackwell directly. Input is already [B, H, L, D]
|
||||
and contiguous from preprocess_qkv."""
|
||||
output = sageattn3_blackwell(query, key, value, is_causal=self.causal)
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Transpose output from [B, H, L, D] back to [B, L, H, D]."""
|
||||
return output.permute(0, 2, 1, 3).contiguous()
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
@@ -13,6 +15,26 @@ from fastvideo.utils import get_compute_dtype
|
||||
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
|
||||
|
||||
|
||||
def _attention_compile_disabled() -> bool:
|
||||
"""Whether to keep attention ``forward`` out of the torch.compile graph.
|
||||
|
||||
Defaults to ``True`` (the historical behavior: attention runs eager via
|
||||
``torch.compiler.disable``). Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=0``
|
||||
to let attention be traced/compiled into the surrounding graph.
|
||||
"""
|
||||
val = os.environ.get("FASTVIDEO_DISABLE_ATTENTION_COMPILE")
|
||||
if val is None:
|
||||
return True
|
||||
return val.strip().lower() not in ("0", "false", "no", "off", "")
|
||||
|
||||
|
||||
def _maybe_compiler_disable(fn):
|
||||
"""Apply ``torch.compiler.disable`` unless disabled via env var."""
|
||||
if _attention_compile_disabled():
|
||||
return torch.compiler.disable(fn)
|
||||
return fn
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
"""Distributed attention layer.
|
||||
"""
|
||||
@@ -56,7 +78,7 @@ class DistributedAttention(nn.Module):
|
||||
self.backend = backend_name_to_enum(attn_backend.get_name())
|
||||
self.dtype = dtype
|
||||
|
||||
@torch.compiler.disable
|
||||
@_maybe_compiler_disable
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
@@ -146,7 +168,7 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
"""Distributed attention layer with VSA support.
|
||||
"""
|
||||
|
||||
@torch.compiler.disable
|
||||
@_maybe_compiler_disable
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
|
||||
@@ -273,11 +273,17 @@ class FastVideoArgs:
|
||||
dit_config = getattr(self.pipeline_config, "dit_config", None)
|
||||
if dit_config is None:
|
||||
return
|
||||
# Resolve a registry name (e.g. "nvfp4_qat_train" from the CLI) to a
|
||||
# QuantizationConfig instance; a bare string has no get_quant_method.
|
||||
tq = self.transformer_quant
|
||||
if isinstance(tq, str):
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
tq = get_quantization_config(tq)()
|
||||
# Don't overwrite if the caller already set it explicitly on
|
||||
# dit_config (e.g. via ``pipeline_config.dit_config.quant_config = NVFP4Config()``);
|
||||
# the explicit setter wins.
|
||||
if getattr(dit_config, "quant_config", None) is None:
|
||||
dit_config.quant_config = self.transformer_quant
|
||||
dit_config.quant_config = tq
|
||||
|
||||
def _resolve_refine_args(self) -> None:
|
||||
"""Map generic refine_* args to LTX-2-specific refine fields."""
|
||||
@@ -1018,6 +1024,12 @@ class TrainingArgs(FastVideoArgs):
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
parser.add_argument("--data-path", type=str, required=True, help="Path to parquet files")
|
||||
parser.add_argument("--transformer-quant",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Quantization config name for the DiT (e.g. nvfp4_qat_train for "
|
||||
"QAT-finetune FP4 linear with a straight-through estimator). "
|
||||
"Resolved to a QuantizationConfig and pinned on dit_config.quant_config.")
|
||||
parser.add_argument("--dataloader-num-workers",
|
||||
type=int,
|
||||
required=True,
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FP8 quantization-aware training for linear layers.
|
||||
|
||||
Mirror of ``fp4linear.py`` but for FP8 (e4m3). The forward pass quantizes both
|
||||
activations and weights to FP8 and runs ``torch._scaled_mm``; the backward pass
|
||||
is a bf16 straight-through estimator so the high-precision master weights stay
|
||||
trainable. Falls back to a bf16 fake-quant forward on GPUs older than sm89.
|
||||
"""
|
||||
import torch
|
||||
|
||||
FP8_DTYPE = torch.float8_e4m3fn
|
||||
FP8_MAX = float(torch.finfo(FP8_DTYPE).max) # 448.0
|
||||
FP8_MIN_SCALE = 1.0 / (FP8_MAX * 512.0)
|
||||
|
||||
|
||||
def _supports_fp8_compute() -> bool:
|
||||
"""Whether the active device supports FP8 ``_scaled_mm`` (sm89+)."""
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
cap = torch.cuda.get_device_capability()
|
||||
return cap[0] > 8 or (cap[0] == 8 and cap[1] >= 9)
|
||||
|
||||
|
||||
def _quantize_tensorwise(x_2d: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Returns ``(x_fp8 [M, K], x_scale [1] float32)``."""
|
||||
x_absmax = x_2d.abs().amax().float()
|
||||
x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
return x_fp8, x_scale.view(1)
|
||||
|
||||
|
||||
def _quantize_rowwise(x_2d: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``."""
|
||||
x_absmax = x_2d.abs().amax(dim=-1, keepdim=True).float()
|
||||
x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
return x_fp8, x_scale
|
||||
|
||||
|
||||
def _fake_quant(x_2d: torch.Tensor, granularity: str) -> torch.Tensor:
|
||||
"""bf16 fake-quant (quantize then dequantize) for pre-sm89 fallback."""
|
||||
if granularity == "channel":
|
||||
x_fp8, x_scale = _quantize_rowwise(x_2d)
|
||||
else:
|
||||
x_fp8, x_scale = _quantize_tensorwise(x_2d)
|
||||
return x_fp8.to(x_2d.dtype) * x_scale.to(x_2d.dtype)
|
||||
|
||||
|
||||
class _LinearFWD8BWD16Fn(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, x, weight, bias, granularity="tensor"):
|
||||
# assert/normalize activation dtype
|
||||
if x.dtype not in (torch.float16, torch.bfloat16):
|
||||
x = x.to(dtype=torch.bfloat16)
|
||||
|
||||
# cast params (can be fp32) to activation dtype for quantization
|
||||
weight_cast = weight.to(dtype=x.dtype)
|
||||
bias_cast = bias.to(dtype=x.dtype) if bias is not None else None
|
||||
|
||||
orig_shape = x.shape
|
||||
k = weight_cast.shape[1]
|
||||
n = weight_cast.shape[0]
|
||||
x2d = x.reshape(-1, k).contiguous()
|
||||
|
||||
if not _supports_fp8_compute():
|
||||
# bf16 fake-quant fallback: simulate the FP8 rounding error but
|
||||
# compute the matmul in bf16.
|
||||
x_fq = _fake_quant(x2d, granularity)
|
||||
w_fq = _fake_quant(weight_cast, granularity)
|
||||
out2d = x_fq.matmul(w_fq.t())
|
||||
if bias_cast is not None:
|
||||
out2d = out2d + bias_cast
|
||||
ctx.save_for_backward(x2d, weight, bias)
|
||||
ctx.n = n
|
||||
ctx.orig_shape = orig_shape
|
||||
return out2d.reshape(*orig_shape[:-1], n)
|
||||
|
||||
if granularity == "channel":
|
||||
x_fp8, x_scale = _quantize_rowwise(x2d)
|
||||
w_fp8, w_scale = _quantize_rowwise(weight_cast)
|
||||
scale_b = w_scale.view(1, -1)
|
||||
else:
|
||||
x_fp8, x_scale = _quantize_tensorwise(x2d)
|
||||
w_fp8, w_scale = _quantize_tensorwise(weight_cast)
|
||||
scale_b = w_scale
|
||||
|
||||
out2d = torch._scaled_mm(
|
||||
x_fp8,
|
||||
w_fp8.t(),
|
||||
scale_a=x_scale,
|
||||
scale_b=scale_b,
|
||||
out_dtype=x.dtype,
|
||||
)
|
||||
if isinstance(out2d, tuple):
|
||||
out2d = out2d[0]
|
||||
|
||||
if bias_cast is not None:
|
||||
out2d = out2d + bias_cast
|
||||
|
||||
# save tensors for backward (keep original dtypes)
|
||||
ctx.save_for_backward(x2d, weight, bias)
|
||||
ctx.n = n
|
||||
ctx.orig_shape = orig_shape
|
||||
return out2d.reshape(*orig_shape[:-1], n)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_out):
|
||||
x2d, weight, bias = ctx.saved_tensors
|
||||
M = x2d.shape[0]
|
||||
n = ctx.n
|
||||
|
||||
grad_out_2d = grad_out.reshape(M, n).contiguous()
|
||||
|
||||
# bf16 straight-through estimator: gradients flow through the
|
||||
# full-precision master weights, not the FP8 quantized values.
|
||||
weight_cast = weight.to(dtype=grad_out.dtype)
|
||||
x_cast = x2d.to(dtype=grad_out.dtype)
|
||||
|
||||
grad_x = grad_out_2d.matmul(weight_cast).reshape(*ctx.orig_shape)
|
||||
grad_w = grad_out_2d.t().matmul(x_cast)
|
||||
grad_b = grad_out_2d.sum(dim=0) if bias is not None else None
|
||||
|
||||
# None for the extra forward arg (granularity)
|
||||
return grad_x, grad_w, grad_b, None
|
||||
|
||||
|
||||
def fp8_linear_forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# pass config **positionally**; autograd.Function.apply ignores kwargs
|
||||
return _LinearFWD8BWD16Fn.apply(x, self.weight, self.bias, "tensor"), None
|
||||
@@ -2,7 +2,7 @@ from typing import Literal, get_args
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig
|
||||
|
||||
QuantizationMethods = Literal[None, "AbsMaxFP8", "NVFP4", "nvfp4_qat"]
|
||||
QuantizationMethods = Literal[None, "AbsMaxFP8", "FP8", "NVFP4", "nvfp4_qat", "nvfp4_qat_train", "fp8_qat_train"]
|
||||
|
||||
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
|
||||
|
||||
@@ -51,13 +51,19 @@ def get_quantization_config(quantization: str) -> type[QuantizationConfig]:
|
||||
|
||||
# lazy import to avoid triggering `torch.compile` too early
|
||||
from .absmax_fp8 import AbsMaxFP8Config
|
||||
from .fp8_config import FP8Config
|
||||
from .nvfp4_config import NVFP4Config
|
||||
from .nvfp4_qat_config import NVFP4QATConfig
|
||||
from .nvfp4_qat_train_config import NVFP4QATTrainConfig
|
||||
from .fp8_qat_train_config import FP8QATTrainConfig
|
||||
|
||||
method_to_config: dict[str, type[QuantizationConfig]] = {
|
||||
"AbsMaxFP8": AbsMaxFP8Config,
|
||||
"FP8": FP8Config,
|
||||
"NVFP4": NVFP4Config,
|
||||
"nvfp4_qat": NVFP4QATConfig,
|
||||
"nvfp4_qat_train": NVFP4QATTrainConfig,
|
||||
"fp8_qat_train": FP8QATTrainConfig,
|
||||
}
|
||||
# Update the `method_to_config` with customized quantization methods.
|
||||
method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG)
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Generic FP8 quantization backed by ``torch._scaled_mm``.
|
||||
|
||||
Matches linear layers by suffix (``to_q/k/v/to_out``, ``ffn.fc_in/fc_out``).
|
||||
Supports per-tensor (default, fast) and per-channel (higher accuracy) granularity.
|
||||
Falls back to bf16 dequant on GPUs older than sm89.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
FP8_DTYPE = torch.float8_e4m3fn
|
||||
FP8_MAX = float(torch.finfo(FP8_DTYPE).max) # 448.0
|
||||
FP8_MIN_SCALE = 1.0 / (FP8_MAX * 512.0)
|
||||
|
||||
_FP8_SUFFIXES = (
|
||||
"ffn.fc_in",
|
||||
"ffn.fc_out",
|
||||
"to_q",
|
||||
"to_k",
|
||||
"to_v",
|
||||
"to_out",
|
||||
)
|
||||
|
||||
|
||||
def _supports_fp8_compute() -> bool:
|
||||
"""Whether the active device supports FP8 ``_scaled_mm`` (sm89+)."""
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
cap = torch.cuda.get_device_capability()
|
||||
return cap[0] > 8 or (cap[0] == 8 and cap[1] >= 9)
|
||||
|
||||
|
||||
def _quantize_tensorwise(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Returns ``(x_fp8 [M, K], x_scale [1] float32)``."""
|
||||
x_absmax = x_2d.abs().amax().float()
|
||||
x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
return x_fp8, x_scale.view(1)
|
||||
|
||||
|
||||
def _quantize_rowwise(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``."""
|
||||
x_absmax = x_2d.abs().amax(dim=-1, keepdim=True).float()
|
||||
x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
return x_fp8, x_scale
|
||||
|
||||
|
||||
class FP8QuantizeMethod(QuantizeMethodBase):
|
||||
"""FP8 linear method.
|
||||
|
||||
``granularity='tensor'`` (default): per-tensor weight + per-tensor
|
||||
dynamic activation scales — the fast tensorwise ``_scaled_mm`` path.
|
||||
``granularity='channel'``: per-output-channel weight + per-token
|
||||
activation scales (rowwise) — higher accuracy but slower ``_scaled_mm``.
|
||||
"""
|
||||
|
||||
def __init__(self, granularity: str = "tensor"):
|
||||
super().__init__()
|
||||
self.granularity = granularity
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
weight = Parameter(
|
||||
torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, None]:
|
||||
"""Pre-quantize an activation for reuse across q/k/v projections."""
|
||||
assert x.dtype in (torch.bfloat16, torch.float16), (f"only allow bf16/fp16 inputs to fp8 linear, got {x.dtype}")
|
||||
x_2d = x.view(-1, x.shape[-1])
|
||||
if self.granularity == "channel":
|
||||
x_fp8, x_scale = _quantize_rowwise(x_2d)
|
||||
else:
|
||||
x_fp8, x_scale = _quantize_tensorwise(x_2d)
|
||||
return x_fp8, x_scale, None
|
||||
|
||||
def wants_prequantized_input(self) -> bool:
|
||||
return True
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
pre_quantized: tuple[torch.Tensor, torch.Tensor, Any] | None = None,
|
||||
) -> torch.Tensor:
|
||||
out_dim = layer._fp8_weight.shape[0]
|
||||
original_shape = x.shape
|
||||
|
||||
if not _supports_fp8_compute():
|
||||
return self._apply_dequant(layer, x, bias)
|
||||
|
||||
if pre_quantized is not None:
|
||||
x_fp8, x_scale, _ = pre_quantized
|
||||
if x_fp8.dim() > 2:
|
||||
x_fp8 = x_fp8.reshape(-1, x_fp8.shape[-1])
|
||||
if x_scale.dim() > 2:
|
||||
x_scale = x_scale.reshape(-1, x_scale.shape[-1])
|
||||
elif self.granularity == "channel":
|
||||
x_fp8, x_scale = _quantize_rowwise(x.reshape(-1, x.shape[-1]))
|
||||
else:
|
||||
x_fp8, x_scale = _quantize_tensorwise(x.reshape(-1, x.shape[-1]))
|
||||
|
||||
w_fp8 = layer._fp8_weight
|
||||
w_scale = layer._fp8_weight_scale
|
||||
scale_b = w_scale.view(1, -1) if self.granularity == "channel" else w_scale
|
||||
|
||||
out = torch._scaled_mm(
|
||||
x_fp8,
|
||||
w_fp8.t(),
|
||||
scale_a=x_scale,
|
||||
scale_b=scale_b,
|
||||
out_dtype=torch.bfloat16,
|
||||
)
|
||||
if isinstance(out, tuple):
|
||||
out = out[0]
|
||||
if bias is not None:
|
||||
out = out + bias
|
||||
return out.view(*original_shape[:-1], out_dim)
|
||||
|
||||
def _apply_dequant(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""bf16 fallback for pre-sm89 GPUs."""
|
||||
out_dim = layer._fp8_weight.shape[0]
|
||||
original_shape = x.shape
|
||||
w_fp8 = layer._fp8_weight
|
||||
w_scale = layer._fp8_weight_scale.to(x.dtype)
|
||||
weight = w_fp8.to(x.dtype) * w_scale.unsqueeze(1)
|
||||
out = F.linear(x, weight, bias)
|
||||
return out.view(*original_shape[:-1], out_dim)
|
||||
|
||||
|
||||
class FP8Config(QuantizationConfig):
|
||||
"""FP8 (e4m3) quantization via suffix matching on standard linear layer names."""
|
||||
|
||||
def __init__(self, granularity: str = "tensor"):
|
||||
super().__init__()
|
||||
if granularity not in ("tensor", "channel"):
|
||||
raise ValueError(f"granularity must be 'tensor' or 'channel', got {granularity!r}")
|
||||
self.granularity = granularity
|
||||
|
||||
def get_name(self) -> str:
|
||||
return "FP8"
|
||||
|
||||
def get_supported_act_dtypes(self) -> list[torch.dtype]:
|
||||
return [torch.bfloat16, torch.float16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return 89
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames() -> list[str]:
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> FP8Config:
|
||||
return cls(granularity=config.get("granularity", "tensor"))
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
from fastvideo.layers.linear import LinearBase
|
||||
|
||||
if isinstance(layer, LinearBase) and any(s in prefix for s in _FP8_SUFFIXES):
|
||||
return FP8QuantizeMethod(granularity=self.granularity)
|
||||
return None
|
||||
|
||||
|
||||
def convert_model_to_fp8(model: torch.nn.Module) -> None:
|
||||
"""Quantize all FP8-tagged linear layers in-place after weights are loaded."""
|
||||
import gc
|
||||
from torch.distributed.tensor import DTensor # type: ignore
|
||||
|
||||
with torch.no_grad():
|
||||
for mod in model.modules():
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
if not isinstance(qm, FP8QuantizeMethod):
|
||||
continue
|
||||
weight = getattr(mod, "weight", None)
|
||||
if weight is None:
|
||||
continue
|
||||
weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type]
|
||||
if getattr(qm, "granularity", "tensor") == "channel":
|
||||
w_absmax = weight_local.detach().abs().amax(dim=1).nan_to_num().float()
|
||||
w_scale = (w_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
w_fp8 = (weight_local / w_scale.to(weight_local.dtype).unsqueeze(1)).clamp(-FP8_MAX,
|
||||
FP8_MAX).to(FP8_DTYPE)
|
||||
else:
|
||||
w_absmax = weight_local.detach().abs().amax().nan_to_num().to(torch.float32)
|
||||
w_scale = (w_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE).view(1)
|
||||
w_fp8 = (weight_local / w_scale.to(weight_local.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
mod.register_buffer("_fp8_weight", w_fp8.contiguous(), persistent=False)
|
||||
mod.register_buffer("_fp8_weight_scale", w_scale.to(torch.float32), persistent=False)
|
||||
removed_weight = mod._parameters.pop("weight", None)
|
||||
if removed_weight is not None:
|
||||
removed_weight.grad = None
|
||||
del removed_weight, weight, weight_local, w_absmax, w_scale, w_fp8
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"FP8Config",
|
||||
"FP8QuantizeMethod",
|
||||
"convert_model_to_fp8",
|
||||
]
|
||||
@@ -0,0 +1,79 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FP8 (e4m3) quantization-aware *training* linear method (straight-through estimator).
|
||||
|
||||
Mirror of ``nvfp4_qat_train_config.py`` but for FP8. The weight stays a trainable
|
||||
bf16/fp32 master that is fake-quantized to FP8 on every forward, with a
|
||||
full-precision backward (STE), so the model learns to absorb FP8 linear error.
|
||||
|
||||
The STE lives in ``fastvideo.layers.fp8linear._LinearFWD8BWD16Fn`` (FP8 forward
|
||||
via ``torch._scaled_mm`` on sm89+, with a bf16 fake-quant fallback on older GPUs;
|
||||
full-precision backward). This method bridges it into the standard
|
||||
``quant_config`` path, so it activates via ``transformer_quant="fp8_qat_train"``
|
||||
on the same Wan-2.1 layers as the FP4 path (to_q/k/v/out + ffn). No conversion is
|
||||
needed: the weight is kept in full precision and quantized on the fly each step.
|
||||
|
||||
Unlike the FP4 path this needs no flashinfer and runs on any sm89+ GPU (and even
|
||||
older ones via the bf16 fallback), not just Blackwell.
|
||||
"""
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig, QuantizeMethodBase
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FP8QATTrainQuantizeMethod(QuantizeMethodBase):
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, output_partition_sizes: list[int],
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype, **extra_weight_attrs):
|
||||
# Trainable master weight, fake-quantized to FP8 on each forward.
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=True)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
# FP8 forward + full-precision backward (STE).
|
||||
from fastvideo.layers.fp8linear import _LinearFWD8BWD16Fn
|
||||
return _LinearFWD8BWD16Fn.apply(x, layer.weight, bias, "tensor")
|
||||
|
||||
|
||||
class FP8QATTrainConfig(QuantizationConfig):
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_name(self):
|
||||
return "fp8_qat_train"
|
||||
|
||||
def get_supported_act_dtypes(self):
|
||||
return [torch.bfloat16, torch.float16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls):
|
||||
return 89
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames():
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "FP8QATTrainConfig":
|
||||
return cls()
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
from fastvideo.layers.linear import LinearBase
|
||||
fp8_layers = ["ffn.fc_in", "ffn.fc_out", "to_q", "to_k", "to_v", "to_out"]
|
||||
if isinstance(layer, LinearBase) and any(layer_name in prefix for layer_name in fp8_layers):
|
||||
return FP8QATTrainQuantizeMethod()
|
||||
return None
|
||||
@@ -1,38 +1,90 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""NVFP4 quantization-aware (QAD) linear method, inference path.
|
||||
|
||||
Quantizes every targeted linear's weight to NVFP4 once at load time and
|
||||
runs each forward as a registered flashinfer-backed FP4 matmul. The
|
||||
original fp16/bf16 weight is *popped* immediately after quantization so
|
||||
the half-precision copy does not keep occupying GPU memory — that's
|
||||
what lets a Wan-2.1 pipeline stay fully resident on a single GPU
|
||||
without any CPU offloading.
|
||||
|
||||
The quantize / matmul custom ops are owned by
|
||||
:mod:`fastvideo.layers.quantization.nvfp4_config` and registered under
|
||||
the ``fastvideo_fp4::`` namespace. We reuse them here for two reasons:
|
||||
|
||||
1. Re-registering the same op name in a second module would raise.
|
||||
2. The registered ops have ``register_fake`` shape/dtype kernels, which
|
||||
is what makes the inference pipeline's per-block ``torch.compile``
|
||||
trace through without graph breaks. Calling raw flashinfer functions
|
||||
(the old behavior of this file, plus a ``@torch.compile`` on
|
||||
``apply``) graph-breaks at every quantize and every matmul.
|
||||
|
||||
For QAT *training*, see ``nvfp4_qat_train_config`` which keeps the
|
||||
weight trainable and fake-quantizes on the fly via a straight-through
|
||||
estimator.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig, QuantizeMethodBase
|
||||
from fastvideo.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from fastvideo.layers.quantization.nvfp4_config import (
|
||||
_mm_fp4,
|
||||
_nvfp4_quantize,
|
||||
_require_flashinfer,
|
||||
)
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
try:
|
||||
import flashinfer
|
||||
except ImportError:
|
||||
flashinfer = None
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Wan-style attention + FFN projection layers. Matched as substrings of the
|
||||
# layer prefix (e.g. "blocks.0.attn1.to_q" contains "to_q").
|
||||
DEFAULT_FP4_LAYERS = (
|
||||
"ffn.fc_in",
|
||||
"ffn.fc_out",
|
||||
"to_q",
|
||||
"to_k",
|
||||
"to_v",
|
||||
"to_out",
|
||||
)
|
||||
|
||||
def _require_flashinfer() -> Any:
|
||||
if flashinfer is None:
|
||||
raise ImportError("flashinfer is required for NVFP4 QAT quantization. "
|
||||
"Please install flashinfer to use the nvfp4_qat quantization backend.")
|
||||
return flashinfer
|
||||
|
||||
def _layout_128x4() -> Any:
|
||||
SfLayout, _, _ = _require_flashinfer()
|
||||
return SfLayout.layout_128x4
|
||||
|
||||
|
||||
class NVFP4QATQuantizeMethod(QuantizeMethodBase):
|
||||
"""Inference-only NVFP4 linear method with weight popping.
|
||||
|
||||
The dense ``weight`` parameter is materialized at load time only so
|
||||
that :func:`convert_model_to_fp4` can read it once; the loader then
|
||||
removes it via ``mod._parameters.pop('weight')``. From that point
|
||||
forward, ``apply`` reads only ``_fp4_weight`` / ``_fp4_weight_scale``
|
||||
/ ``_weight_global_sf``.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.weight_fp4 = None
|
||||
self.weight_scale = None
|
||||
# Static input global scale factor. Matches the FastVideo-Quantization
|
||||
# production path; recomputing it per-call via a ``.max()`` reduction
|
||||
# (the previous behavior) adds a sync point, costs a kernel launch,
|
||||
# and produces a data-dependent value that prevents CUDA-graph
|
||||
# capture under ``torch.compile(mode='reduce-overhead')``.
|
||||
self.x_global_sf = torch.tensor(1.0, device="cuda", dtype=torch.float32)
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, output_partition_sizes: list[int],
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype, **extra_weight_attrs):
|
||||
"""Create weights for a linear layer. Note the corrected signature to match LinearMethodBase."""
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype, **extra_weight_attrs) -> None:
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
@@ -43,28 +95,27 @@ class NVFP4QATQuantizeMethod(QuantizeMethodBase):
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
@torch.compile
|
||||
def apply(self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
"""Apply NVFP4 QAT quantized computation."""
|
||||
flashinfer_mod = _require_flashinfer()
|
||||
out_dim = layer.weight.shape[0]
|
||||
# ``_fp4_weight`` carries the (out, in/2) packed fp4 weight, so its
|
||||
# row count is the output dim even after the dense weight is popped.
|
||||
out_dim = layer._fp4_weight.shape[0]
|
||||
original_shape = x.shape
|
||||
assert x.dtype == torch.bfloat16 or x.dtype == torch.float16, f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}"
|
||||
|
||||
assert x.dtype in (torch.bfloat16, torch.float16), (f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}")
|
||||
x = x.view(-1, x.shape[-1])
|
||||
|
||||
x_global_sf = (448 * 6) / x.float().abs().nan_to_num().max()
|
||||
x_fp4, x_scale = flashinfer_mod.nvfp4_quantize(
|
||||
x_global_sf = self.x_global_sf
|
||||
x_fp4, x_scale = _nvfp4_quantize(
|
||||
x,
|
||||
x_global_sf,
|
||||
sfLayout=flashinfer_mod.SfLayout.layout_128x4,
|
||||
sfLayout=_layout_128x4(),
|
||||
do_shuffle=False,
|
||||
)
|
||||
|
||||
weight_fp4 = layer._fp4_weight
|
||||
weight_scale = layer._fp4_weight_scale
|
||||
weight_global_sf = layer._weight_global_sf
|
||||
|
||||
out = flashinfer_mod.mm_fp4(
|
||||
out = _mm_fp4(
|
||||
x_fp4,
|
||||
weight_fp4.T,
|
||||
x_scale,
|
||||
@@ -76,67 +127,106 @@ class NVFP4QATQuantizeMethod(QuantizeMethodBase):
|
||||
)
|
||||
|
||||
if bias is not None:
|
||||
if bias.device != out.device or bias.dtype != out.dtype:
|
||||
bias = bias.to(device=out.device, dtype=out.dtype)
|
||||
out = out + bias
|
||||
|
||||
if len(original_shape) == 3:
|
||||
out = out.view(original_shape[0], original_shape[1], out_dim)
|
||||
|
||||
out = out.view(*original_shape[:-1], out_dim)
|
||||
return out
|
||||
|
||||
|
||||
class NVFP4QATConfig(QuantizationConfig):
|
||||
"""NVFP4 (Wan-style) linear quantization, inference.
|
||||
|
||||
def __init__(self) -> None:
|
||||
Args:
|
||||
target_layers: Substrings matched against each linear layer's
|
||||
prefix. A layer is quantized if any substring is contained in
|
||||
its prefix. Defaults to the standard Wan attention + FFN
|
||||
projections (:data:`DEFAULT_FP4_LAYERS`).
|
||||
"""
|
||||
|
||||
def __init__(self, target_layers: tuple[str, ...] | None = None) -> None:
|
||||
super().__init__()
|
||||
self.target_layers = (tuple(target_layers) if target_layers else DEFAULT_FP4_LAYERS)
|
||||
|
||||
def get_name(self):
|
||||
def get_name(self) -> str:
|
||||
return "nvfp4_qat"
|
||||
|
||||
def get_supported_act_dtypes(self):
|
||||
def get_supported_act_dtypes(self) -> list[torch.dtype]:
|
||||
return [torch.bfloat16, torch.float16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls):
|
||||
def get_min_capability(cls) -> int:
|
||||
return 100
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames():
|
||||
def get_config_filenames() -> list[str]:
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "NVFP4QATConfig":
|
||||
return cls()
|
||||
def from_config(cls, config: dict[str, Any]) -> NVFP4QATConfig:
|
||||
target_layers = config.get("target_layers")
|
||||
if target_layers is not None:
|
||||
target_layers = tuple(target_layers)
|
||||
return cls(target_layers=target_layers)
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
from fastvideo.layers.linear import LinearBase
|
||||
fp4_layers = ["ffn.fc_in", "ffn.fc_out", "to_q", "to_k", "to_v", "to_out"]
|
||||
if isinstance(layer, LinearBase) and any(layer_name in prefix for layer_name in fp4_layers):
|
||||
if isinstance(layer, LinearBase) and any(name in prefix for name in self.target_layers):
|
||||
return NVFP4QATQuantizeMethod()
|
||||
return None
|
||||
|
||||
|
||||
@torch.compile
|
||||
def convert_model_to_fp4(model: torch.nn.Module):
|
||||
flashinfer_mod = _require_flashinfer()
|
||||
def convert_model_to_fp4(model: torch.nn.Module) -> None:
|
||||
"""Prequantize every FP4-tagged linear and drop its dense weight.
|
||||
|
||||
Walks the module tree, and for each layer whose ``quant_method`` is
|
||||
an :class:`NVFP4QATQuantizeMethod`, computes the NVFP4 packed weight
|
||||
/ scale / global-scale buffers, then pops the original fp16/bf16
|
||||
``weight`` parameter so it no longer occupies GPU memory.
|
||||
"""
|
||||
SfLayout, _, _ = _require_flashinfer()
|
||||
from torch.distributed.tensor import DTensor # type: ignore
|
||||
for mod in model.modules():
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
if isinstance(qm, NVFP4QATQuantizeMethod):
|
||||
|
||||
with torch.no_grad():
|
||||
for mod in model.modules():
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
if not isinstance(qm, NVFP4QATQuantizeMethod):
|
||||
continue
|
||||
|
||||
weight = getattr(mod, "weight", None)
|
||||
if weight is None:
|
||||
continue
|
||||
|
||||
weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type]
|
||||
weight_global_sf = (448 * 6) / weight_local.float().abs().nan_to_num().max()
|
||||
fp4_w, fp4_s = flashinfer_mod.nvfp4_quantize(
|
||||
|
||||
# Only the reduced scalar needs fp32; avoid a full fp32 copy.
|
||||
weight_absmax = (weight_local.detach().abs().nan_to_num().amax().to(dtype=torch.float32))
|
||||
weight_global_sf = (448 * 6) / weight_absmax
|
||||
fp4_w, fp4_s = _nvfp4_quantize(
|
||||
weight_local,
|
||||
weight_global_sf,
|
||||
sfLayout=flashinfer_mod.SfLayout.layout_128x4,
|
||||
sfLayout=SfLayout.layout_128x4,
|
||||
do_shuffle=False,
|
||||
)
|
||||
mod.register_buffer("_fp4_weight", fp4_w, persistent=False)
|
||||
mod.register_buffer("_fp4_weight_scale", fp4_s, persistent=False)
|
||||
mod.register_buffer("_weight_global_sf",
|
||||
torch.tensor(weight_global_sf, dtype=torch.bfloat16),
|
||||
persistent=False)
|
||||
mod.register_buffer(
|
||||
"_weight_global_sf",
|
||||
weight_global_sf.to(dtype=torch.bfloat16),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
# Drop the dense weight as soon as the fp4 buffers are installed
|
||||
# so it cannot keep occupying GPU memory.
|
||||
removed_weight = mod._parameters.pop("weight", None)
|
||||
if removed_weight is not None:
|
||||
removed_weight.grad = None
|
||||
del removed_weight, weight, weight_local, weight_absmax
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"NVFP4QATConfig",
|
||||
"NVFP4QATQuantizeMethod",
|
||||
"convert_model_to_fp4",
|
||||
"DEFAULT_FP4_LAYERS",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""NVFP4 quantization-aware *training* linear method (straight-through estimator).
|
||||
|
||||
The inference ``nvfp4_qat`` config quantizes each weight to FP4 once at load time
|
||||
(``convert_model_to_fp4``) and has no gradient path — it is inference only. For
|
||||
QAT *finetuning* the weight must stay a trainable bf16/fp32 master that is
|
||||
fake-quantized to FP4 on every forward, with a full-precision backward (a
|
||||
straight-through estimator), so the model learns to absorb FP4 linear error.
|
||||
|
||||
That STE already exists in ``fastvideo.layers.fp4linear._LinearFWD4BWD16Fn``
|
||||
(FP4 forward, full-precision backward) but is otherwise unwired. This method
|
||||
bridges it into the standard ``quant_config`` path, so it activates via
|
||||
``transformer_quant="nvfp4_qat_train"`` on the same Wan-2.1 layers as nvfp4_qat
|
||||
(to_q/k/v/out + ffn). No ``convert_model_to_fp4`` is needed: the weight is kept
|
||||
in full precision and quantized on the fly each step.
|
||||
"""
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig, QuantizeMethodBase
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class NVFP4QATTrainQuantizeMethod(QuantizeMethodBase):
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, output_partition_sizes: list[int],
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype, **extra_weight_attrs):
|
||||
# Trainable master weight, fake-quantized to FP4 on each forward.
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=True)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
# FP4 forward + full-precision backward (STE).
|
||||
from fastvideo.layers.fp4linear import _LinearFWD4BWD16Fn
|
||||
return _LinearFWD4BWD16Fn.apply(x, layer.weight, bias, "cutlass", 16, True)
|
||||
|
||||
|
||||
class NVFP4QATTrainConfig(QuantizationConfig):
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_name(self):
|
||||
return "nvfp4_qat_train"
|
||||
|
||||
def get_supported_act_dtypes(self):
|
||||
return [torch.bfloat16, torch.float16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls):
|
||||
return 100
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames():
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "NVFP4QATTrainConfig":
|
||||
return cls()
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
from fastvideo.layers.linear import LinearBase
|
||||
fp4_layers = ["ffn.fc_in", "ffn.fc_out", "to_q", "to_k", "to_v", "to_out"]
|
||||
if isinstance(layer, LinearBase) and any(layer_name in prefix for layer_name in fp4_layers):
|
||||
return NVFP4QATTrainQuantizeMethod()
|
||||
return None
|
||||
@@ -28,31 +28,29 @@ from fastvideo.utils import set_mixed_precision_policy, is_pin_memory_available
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _maybe_convert_model_to_nvfp4(model: nn.Module) -> None:
|
||||
"""Quantize NVFP4-tagged linear layers in-place after weights are loaded.
|
||||
def _maybe_quantize_model(model: nn.Module) -> None:
|
||||
"""Quantize NVFP4- or FP8-tagged linear layers in-place after weights are loaded.
|
||||
|
||||
Walks the module tree once, looking for layers whose ``quant_method``
|
||||
is an :class:`NVFP4QuantizeMethod` (attached at construction time by
|
||||
:meth:`NVFP4Config.get_quant_method`). When at least one such layer
|
||||
exists, calls :func:`convert_model_to_nvfp4` to register the
|
||||
``_nvfp4_weight*`` / ``_nvfp4_alpha`` / ``_weight_global_sf`` buffers
|
||||
on each targeted layer.
|
||||
is an :class:`NVFP4QuantizeMethod` or :class:`FP8QuantizeMethod` (attached
|
||||
at construction time by the respective ``get_quant_method``). When at least
|
||||
one such layer exists, calls the matching conversion function to register
|
||||
quantized weight buffers on each targeted layer.
|
||||
|
||||
The walk returns on the first NVFP4 layer found so non-NVFP4 callers
|
||||
pay only an ``isinstance`` check per module. flashinfer is imported
|
||||
lazily inside :func:`convert_model_to_nvfp4` so this helper is a
|
||||
no-op on hosts without the NVFP4 backend.
|
||||
The walk returns on the first quantized layer found so unquantized callers
|
||||
pay only an ``isinstance`` check per module. Both imports are deferred so
|
||||
this is a no-op on hosts without the relevant backends.
|
||||
"""
|
||||
# Defer the import: nvfp4_config imports heavy diffusers /
|
||||
# torch.distributed symbols at module-load time, and unconditional
|
||||
# import would penalize every loader call regardless of whether
|
||||
# NVFP4 is wired.
|
||||
# Defer imports: these modules pull in heavy symbols at module-load time.
|
||||
from fastvideo.layers.quantization.nvfp4_config import (
|
||||
NVFP4QuantizeMethod, convert_model_to_nvfp4,
|
||||
)
|
||||
from fastvideo.layers.quantization.nvfp4_qat_config import (
|
||||
NVFP4QATQuantizeMethod, convert_model_to_fp4,
|
||||
)
|
||||
from fastvideo.layers.quantization.fp8_config import (
|
||||
FP8QuantizeMethod, convert_model_to_fp8,
|
||||
)
|
||||
|
||||
for mod in model.modules():
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
@@ -63,6 +61,9 @@ def _maybe_convert_model_to_nvfp4(model: nn.Module) -> None:
|
||||
if isinstance(qm, NVFP4QATQuantizeMethod):
|
||||
logger.info("Converting loaded model weights for NVFP4-QAT linear layers")
|
||||
convert_model_to_fp4(model)
|
||||
if isinstance(qm, FP8QuantizeMethod):
|
||||
logger.info("Converting loaded model weights for FP8 linear layers")
|
||||
convert_model_to_fp8(model)
|
||||
return
|
||||
|
||||
|
||||
@@ -196,14 +197,13 @@ def maybe_load_fsdp_model(
|
||||
if isinstance(p, torch.nn.Parameter):
|
||||
p.requires_grad = False
|
||||
|
||||
# NVFP4 weight prequantization. We detect by the registered
|
||||
# ``quant_method`` on linear layers rather than by a separate flag —
|
||||
# construction-time ``NVFP4Config.get_quant_method`` already attached
|
||||
# ``NVFP4QuantizeMethod`` to every targeted layer, so the loader's
|
||||
# responsibility is just to materialize the per-layer nvfp4 weight /
|
||||
# scale buffers from the freshly-loaded bf16 weights. No-op when
|
||||
# ``flashinfer`` is not installed (lazy import inside the helper).
|
||||
_maybe_convert_model_to_nvfp4(model)
|
||||
# Post-load weight quantization. We detect the active scheme by the
|
||||
# ``quant_method`` attached to each linear layer at construction time
|
||||
# (via ``QuantizationConfig.get_quant_method``). The loader's
|
||||
# responsibility is just to materialize the quantized weight buffers
|
||||
# from the freshly-loaded bf16 weights. No-op when no quantized layers
|
||||
# are present (lazy imports inside the helper).
|
||||
_maybe_quantize_model(model)
|
||||
|
||||
compile_in_loader = enable_torch_compile and training_mode
|
||||
if compile_in_loader:
|
||||
|
||||
@@ -111,10 +111,17 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
days: int = Query(DEFAULT_DAYS, ge=1, le=3650),
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
run_source: str | None = None,
|
||||
success: bool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
loaded = data_store.load_records(days=days)
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type, success=success)
|
||||
filtered = filter_records(
|
||||
loaded,
|
||||
model_id=model_id,
|
||||
gpu_type=gpu_type,
|
||||
run_source=run_source,
|
||||
success=success,
|
||||
)
|
||||
return {
|
||||
"records": filtered,
|
||||
"count": len(filtered),
|
||||
@@ -122,6 +129,7 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
"days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"run_source": run_source,
|
||||
"success": success,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
@@ -132,6 +140,7 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
days: int = Query(DEFAULT_DAYS, ge=1, le=3650),
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
run_source: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
# Latest status should be stable when users change the trend window.
|
||||
# Use all cached records for latest/baseline computation; the ``days``
|
||||
@@ -139,7 +148,11 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
# endpoints without affecting the summary semantics.
|
||||
loaded = data_store.load_records(days=None)
|
||||
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")))
|
||||
rows = build_latest_summary(
|
||||
filtered,
|
||||
max_regression=float(os.environ.get("PERF_MAX_REGRESSION", "0.05")),
|
||||
run_source=run_source,
|
||||
)
|
||||
return {
|
||||
"rows": rows,
|
||||
"count": len(rows),
|
||||
@@ -152,6 +165,7 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
"trend_window_days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"run_source": run_source,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
}
|
||||
@@ -161,9 +175,10 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
days: int = Query(DEFAULT_DAYS, ge=1, le=3650),
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
run_source: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
loaded = data_store.load_records(days=days)
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type)
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type, run_source=run_source)
|
||||
groups = build_trends(filtered)
|
||||
return {
|
||||
"groups": groups,
|
||||
@@ -172,6 +187,7 @@ def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
"days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"run_source": run_source,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
}
|
||||
|
||||
@@ -13,7 +13,7 @@ from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.tests.performance.hf_store import safe_float
|
||||
from fastvideo.tests.performance.hf_store import is_baseline_eligible_record, safe_float
|
||||
|
||||
from .metrics import METRICS
|
||||
|
||||
@@ -45,6 +45,7 @@ def filter_records(
|
||||
*,
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
run_source: str | None = None,
|
||||
success: bool | None = None,
|
||||
) -> list[Record]:
|
||||
filtered = records
|
||||
@@ -52,11 +53,31 @@ def filter_records(
|
||||
filtered = [record for record in filtered if record.get("model_id") == model_id]
|
||||
if gpu_type:
|
||||
filtered = [record for record in filtered if record.get("gpu_type") == gpu_type]
|
||||
if run_source:
|
||||
filtered = [record for record in filtered if record_run_source(record) == run_source]
|
||||
if success is not None:
|
||||
filtered = [record for record in filtered if bool(record.get("success", True)) == success]
|
||||
return sorted(filtered, key=record_sort_key)
|
||||
|
||||
|
||||
def record_run_source(record: Record) -> str:
|
||||
value = str(record.get("run_source") or "unknown")
|
||||
return value if value in {"pr", "local", "scheduled_main", "unknown"} else "unknown"
|
||||
|
||||
|
||||
def record_metadata(record: Record) -> Record:
|
||||
return {
|
||||
"run_source": record_run_source(record),
|
||||
"baseline_eligible": is_baseline_eligible_record(record),
|
||||
"branch": record.get("branch") or "",
|
||||
"pr_number": record.get("pr_number") or "",
|
||||
"test_scope": record.get("test_scope") or "",
|
||||
"build_url": record.get("build_url") or "",
|
||||
"build_id": record.get("build_id") or "",
|
||||
"job_id": record.get("job_id") or "",
|
||||
}
|
||||
|
||||
|
||||
def group_by_model_gpu(records: list[Record]) -> dict[tuple[str, str], list[Record]]:
|
||||
groups: dict[tuple[str, str], list[Record]] = defaultdict(list)
|
||||
for record in records:
|
||||
@@ -86,12 +107,22 @@ def regression_percent(metric_key: str, current: float | None, baseline: float |
|
||||
def build_latest_summary(records: list[Record],
|
||||
*,
|
||||
baseline_window: int = 5,
|
||||
max_regression: float = 0.05) -> list[Record]:
|
||||
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():
|
||||
latest = group[-1]
|
||||
earlier_successes = [record for record in group[:-1] if record.get("success", True)]
|
||||
baseline_records = earlier_successes[-baseline_window:]
|
||||
latest_candidates = group
|
||||
if run_source:
|
||||
latest_candidates = [record for record in group if record_run_source(record) == run_source]
|
||||
if not latest_candidates:
|
||||
continue
|
||||
|
||||
latest = latest_candidates[-1]
|
||||
baseline_pool = [
|
||||
record for record in group
|
||||
if record is not latest and record.get("success", True) and is_baseline_eligible_record(record)
|
||||
]
|
||||
baseline_records = baseline_pool[-baseline_window:]
|
||||
|
||||
metrics: dict[str, Record] = {}
|
||||
regressions: list[float] = []
|
||||
@@ -123,6 +154,7 @@ def build_latest_summary(records: list[Record],
|
||||
latest.get("timestamp"),
|
||||
"commit_sha":
|
||||
latest.get("commit_sha"),
|
||||
**record_metadata(latest),
|
||||
"success":
|
||||
success,
|
||||
"baseline_n":
|
||||
@@ -150,6 +182,7 @@ def build_trends(records: list[Record]) -> list[Record]:
|
||||
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))
|
||||
|
||||
@@ -26,6 +26,12 @@ image = (modal.Image.from_registry(
|
||||
os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"BUILDKITE_BRANCH":
|
||||
os.environ.get("BUILDKITE_BRANCH", ""),
|
||||
"BUILDKITE_BUILD_URL":
|
||||
os.environ.get("BUILDKITE_BUILD_URL", ""),
|
||||
"BUILDKITE_BUILD_ID":
|
||||
os.environ.get("BUILDKITE_BUILD_ID", ""),
|
||||
"BUILDKITE_JOB_ID":
|
||||
os.environ.get("BUILDKITE_JOB_ID", ""),
|
||||
"TEST_SCOPE":
|
||||
os.environ.get("TEST_SCOPE", ""),
|
||||
"IMAGE_VERSION":
|
||||
@@ -337,18 +343,30 @@ def run_lora_extraction_tests():
|
||||
],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_performance_tests():
|
||||
# compare_baseline.py runs only after pytest passes, so normalized_perf_*.json
|
||||
# artifacts are emitted for rolling-baseline failures, not fixed-threshold
|
||||
# pytest failures. dashboard.py still runs on red CI for observability.
|
||||
# PR/direct records are uploaded only on pass; scheduled main uploads pass
|
||||
# and fail so the dashboard records every canonical baseline attempt.
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && "
|
||||
"export PERFORMANCE_TRACKING_ROOT='/tmp/perf-tracking' && "
|
||||
"hf auth login --token $HF_API_KEY && "
|
||||
"if [ \"${BUILDKITE_BRANCH:-}\" = 'main' ] && [ \"${TEST_SCOPE:-}\" = 'full' ]; then "
|
||||
"export PERF_RUN_SOURCE='scheduled_main'; "
|
||||
"export PERF_UPLOAD_POLICY='always'; "
|
||||
"elif [ -n \"${BUILDKITE_PULL_REQUEST:-}\" ] && [ \"${BUILDKITE_PULL_REQUEST:-false}\" != 'false' ]; then "
|
||||
"export PERF_RUN_SOURCE='pr'; "
|
||||
"export PERF_UPLOAD_POLICY='pass'; "
|
||||
"elif [ \"${TEST_SCOPE:-}\" = 'direct' ]; then "
|
||||
"export PERF_RUN_SOURCE='unknown'; "
|
||||
"export PERF_UPLOAD_POLICY='pass'; "
|
||||
"else "
|
||||
"export PERF_RUN_SOURCE='unknown'; "
|
||||
"export PERF_UPLOAD_POLICY='never'; "
|
||||
"fi; "
|
||||
"pytest ./fastvideo/tests/performance -vs; "
|
||||
"PYTEST_RC=$?; "
|
||||
"PERF_RC=0; "
|
||||
"if [ $PYTEST_RC -eq 0 ]; then "
|
||||
"python ./fastvideo/tests/performance/compare_baseline.py; "
|
||||
"if [ $PYTEST_RC -eq 0 ] || [ \"$PERF_UPLOAD_POLICY\" = 'always' ]; then "
|
||||
"PERF_PYTEST_RC=$PYTEST_RC python ./fastvideo/tests/performance/compare_baseline.py; "
|
||||
"PERF_RC=$?; "
|
||||
"fi; "
|
||||
"python ./fastvideo/tests/performance/dashboard.py || true; "
|
||||
|
||||
@@ -4,10 +4,10 @@
|
||||
This script:
|
||||
1) reads current benchmark results from fastvideo/tests/performance/results,
|
||||
2) syncs the canonical baseline from the configured HF dataset repo,
|
||||
3) compares each current record against the median of up to 5 prior records
|
||||
(filtered by gpu_type, successful only),
|
||||
4) on persist runs (full-suite on main branch), writes the normalized record
|
||||
back to the HF dataset repo,
|
||||
3) compares each current record against the median of up to 5 prior
|
||||
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%).
|
||||
"""
|
||||
@@ -20,13 +20,22 @@ import sys
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
try:
|
||||
from .hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
except ImportError:
|
||||
from hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
|
||||
RESULTS_DIR = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)),
|
||||
@@ -38,6 +47,9 @@ TRACKING_ROOT = os.environ.get(
|
||||
)
|
||||
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),
|
||||
@@ -56,9 +68,74 @@ LOWER_IS_BETTER_METRICS = {
|
||||
|
||||
|
||||
def _should_persist_tracking() -> bool:
|
||||
test_scope = os.environ.get("TEST_SCOPE", "")
|
||||
branch = os.environ.get("BUILDKITE_BRANCH", "")
|
||||
return test_scope == "full" and branch == "main"
|
||||
return _normalized_upload_policy() != "never"
|
||||
|
||||
|
||||
def _normalized_upload_policy() -> str:
|
||||
if UPLOAD_POLICY in VALID_UPLOAD_POLICIES:
|
||||
return UPLOAD_POLICY
|
||||
print(f"Invalid PERF_UPLOAD_POLICY={UPLOAD_POLICY!r}; using 'never'")
|
||||
return "never"
|
||||
|
||||
|
||||
def _truthy_pr_number(value: str | None) -> bool:
|
||||
return bool(value and value not in {"false", "0", "None", "none"})
|
||||
|
||||
|
||||
def _detect_run_source() -> str:
|
||||
explicit = os.environ.get("PERF_RUN_SOURCE", "").strip().lower()
|
||||
if explicit in VALID_RUN_SOURCES:
|
||||
return explicit
|
||||
if explicit:
|
||||
print(f"Invalid PERF_RUN_SOURCE={explicit!r}; inferring run source")
|
||||
|
||||
if _truthy_pr_number(os.environ.get("BUILDKITE_PULL_REQUEST")):
|
||||
return "pr"
|
||||
if os.environ.get("BUILDKITE_BRANCH") == "main" and os.environ.get("TEST_SCOPE") == "full":
|
||||
return "scheduled_main"
|
||||
if not os.environ.get("BUILDKITE_COMMIT"):
|
||||
return "local"
|
||||
return "unknown"
|
||||
|
||||
|
||||
def _is_baseline_eligible(run_source: str, success: bool) -> bool:
|
||||
return run_source == "scheduled_main" and success
|
||||
|
||||
|
||||
def _upload_allowed(record: dict[str, Any]) -> bool:
|
||||
policy = _normalized_upload_policy()
|
||||
if policy == "always":
|
||||
return True
|
||||
if policy == "pass":
|
||||
return bool(record.get("success", True))
|
||||
return False
|
||||
|
||||
|
||||
def _result_failed_static_thresholds() -> bool:
|
||||
value = os.environ.get("PERF_PYTEST_RC", "")
|
||||
if not value:
|
||||
return False
|
||||
try:
|
||||
return int(value) != 0
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def _record_metadata(run_source: str, result: dict[str, Any]) -> dict[str, Any]:
|
||||
pr_number = result.get("pr_number") or os.environ.get("BUILDKITE_PULL_REQUEST", "")
|
||||
if not _truthy_pr_number(str(pr_number)):
|
||||
pr_number = ""
|
||||
return {
|
||||
"run_source": run_source,
|
||||
"baseline_eligible": False,
|
||||
"branch": os.environ.get("BUILDKITE_BRANCH", ""),
|
||||
"pr_number": pr_number,
|
||||
"test_scope": os.environ.get("TEST_SCOPE", ""),
|
||||
"build_url": os.environ.get("BUILDKITE_BUILD_URL", ""),
|
||||
"build_id": os.environ.get("BUILDKITE_BUILD_ID", ""),
|
||||
"job_id": os.environ.get("BUILDKITE_JOB_ID", ""),
|
||||
}
|
||||
|
||||
|
||||
def _load_current_results() -> list[dict[str, Any]]:
|
||||
pattern = os.path.join(RESULTS_DIR, "perf_*.json")
|
||||
@@ -104,6 +181,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
"dit_time_s": dit_time,
|
||||
"vae_decode_time_s": vae_decode_time,
|
||||
"success": True,
|
||||
**_record_metadata(_detect_run_source(), result),
|
||||
}
|
||||
|
||||
|
||||
@@ -297,8 +375,11 @@ def _emit_markdown_summary(markdown: str, commit_sha: str) -> None:
|
||||
|
||||
def main() -> int:
|
||||
persist_tracking = _should_persist_tracking()
|
||||
upload_policy = _normalized_upload_policy()
|
||||
static_threshold_failed = _result_failed_static_thresholds()
|
||||
|
||||
# Strict on persist: a silent sync failure would pollute the baseline.
|
||||
# Strict on upload-enabled runs: silent sync failure would make comparison
|
||||
# and upload state ambiguous.
|
||||
sync_from_hf(TRACKING_ROOT, strict=persist_tracking)
|
||||
|
||||
current_results = _load_current_results()
|
||||
@@ -310,10 +391,12 @@ def main() -> int:
|
||||
summary_rows: list[dict[str, Any]] = []
|
||||
|
||||
if persist_tracking:
|
||||
print("Tracking persistence enabled: full-suite run on main branch")
|
||||
print(f"Tracking persistence enabled: PERF_UPLOAD_POLICY={upload_policy}")
|
||||
else:
|
||||
print("Tracking persistence disabled: "
|
||||
"only full-suite runs on main branch are persisted")
|
||||
print("Tracking persistence disabled: PERF_UPLOAD_POLICY=never")
|
||||
|
||||
if static_threshold_failed:
|
||||
print(f"Static-threshold phase failed: PERF_PYTEST_RC={os.environ.get('PERF_PYTEST_RC')}")
|
||||
|
||||
for raw in current_results:
|
||||
record = _normalize_record(raw)
|
||||
@@ -324,6 +407,7 @@ def main() -> int:
|
||||
record["gpu_type"],
|
||||
last_n=5,
|
||||
successful_only=True,
|
||||
baseline_eligible_only=True,
|
||||
)
|
||||
|
||||
if not baseline_records:
|
||||
@@ -333,15 +417,22 @@ def main() -> int:
|
||||
record["success"] = True
|
||||
else:
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
record["success"] = not failures
|
||||
all_failures.extend(failures)
|
||||
if static_threshold_failed:
|
||||
failures.append(f"{record['model_id']} fixed-threshold phase failed "
|
||||
f"(PERF_PYTEST_RC={os.environ.get('PERF_PYTEST_RC')})")
|
||||
|
||||
record["success"] = not failures
|
||||
record["baseline_eligible"] = _is_baseline_eligible(record["run_source"], record["success"])
|
||||
all_failures.extend(failures)
|
||||
|
||||
_write_normalized_artifact(record)
|
||||
|
||||
# Strict upload: a silent failure would freeze the rolling baseline.
|
||||
if persist_tracking:
|
||||
if _upload_allowed(record):
|
||||
current_path = _write_tracking_record(record)
|
||||
upload_record(current_path, record, strict=True)
|
||||
else:
|
||||
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)))
|
||||
|
||||
|
||||
@@ -48,6 +48,20 @@ def safe_float(value: Any) -> float | None:
|
||||
return None
|
||||
|
||||
|
||||
def is_baseline_eligible_record(record: dict[str, Any]) -> bool:
|
||||
"""Return whether *record* may contribute to rolling baselines.
|
||||
|
||||
Legacy records predate ``baseline_eligible`` and ``run_source``. They were
|
||||
uploaded only by the old successful main/full-suite path, so keep them
|
||||
eligible until the HF history naturally rolls forward.
|
||||
"""
|
||||
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
|
||||
|
||||
|
||||
def resolve_hf_token() -> str | None:
|
||||
"""Return the first configured Hugging Face token env var."""
|
||||
for env_var in HF_TOKEN_ENV_VARS:
|
||||
@@ -197,6 +211,7 @@ def load_records(
|
||||
*,
|
||||
days: int | None = None,
|
||||
successful_only: bool = False,
|
||||
baseline_eligible_only: bool = False,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return raw JSON dicts from *local_dir*.
|
||||
|
||||
@@ -206,6 +221,9 @@ def load_records(
|
||||
many days. Records with a missing/unparsable timestamp are kept.
|
||||
successful_only: When True, only records with ``success=True`` are
|
||||
returned. Useful when building a regression baseline.
|
||||
baseline_eligible_only: When True, only baseline-eligible records are
|
||||
returned. Legacy records missing both ``baseline_eligible`` and
|
||||
``run_source`` are treated as eligible.
|
||||
|
||||
Returns:
|
||||
List of raw dicts sorted by ``timestamp`` ascending (records that could
|
||||
@@ -227,6 +245,9 @@ def load_records(
|
||||
if successful_only and not data.get("success", True):
|
||||
continue
|
||||
|
||||
if baseline_eligible_only and not is_baseline_eligible_record(data):
|
||||
continue
|
||||
|
||||
if cutoff is not None:
|
||||
raw_ts = data.get("timestamp")
|
||||
if raw_ts:
|
||||
@@ -251,6 +272,7 @@ def load_records_for_model(
|
||||
*,
|
||||
last_n: int | None = None,
|
||||
successful_only: bool = True,
|
||||
baseline_eligible_only: bool = False,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Return records for a specific *model_id*, optionally filtered by GPU.
|
||||
|
||||
@@ -261,6 +283,7 @@ def load_records_for_model(
|
||||
last_n: When set, return only the most recent *n* records (after all
|
||||
other filters). Useful for sliding-window baseline calculations.
|
||||
successful_only: Passed through to :func:`load_records`.
|
||||
baseline_eligible_only: Passed through to :func:`load_records`.
|
||||
|
||||
Returns:
|
||||
List of matching dicts sorted by timestamp ascending.
|
||||
@@ -269,7 +292,11 @@ def load_records_for_model(
|
||||
if not os.path.isdir(model_dir):
|
||||
return []
|
||||
|
||||
records = load_records(model_dir, successful_only=successful_only)
|
||||
records = load_records(
|
||||
model_dir,
|
||||
successful_only=successful_only,
|
||||
baseline_eligible_only=baseline_eligible_only,
|
||||
)
|
||||
|
||||
if gpu_type is not None:
|
||||
records = [r for r in records if r.get("gpu_type") == gpu_type]
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.tests.performance import compare_baseline
|
||||
|
||||
|
||||
def _raw_result():
|
||||
return {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"device": "NVIDIA L40S",
|
||||
"avg_generation_time_s": 10.0,
|
||||
"throughput_fps": 4.5,
|
||||
"max_peak_memory_mb": 10000.0,
|
||||
"commit": "a" * 40,
|
||||
"timestamp": "2026-06-16T00:00:00+00:00",
|
||||
"pr_number": "123",
|
||||
}
|
||||
|
||||
|
||||
def test_detect_run_source_prefers_explicit_env(monkeypatch):
|
||||
monkeypatch.setenv("PERF_RUN_SOURCE", "local")
|
||||
monkeypatch.setenv("BUILDKITE_PULL_REQUEST", "123")
|
||||
|
||||
assert compare_baseline._detect_run_source() == "local"
|
||||
|
||||
|
||||
def test_detect_run_source_infers_pr(monkeypatch):
|
||||
monkeypatch.delenv("PERF_RUN_SOURCE", raising=False)
|
||||
monkeypatch.setenv("BUILDKITE_PULL_REQUEST", "123")
|
||||
|
||||
assert compare_baseline._detect_run_source() == "pr"
|
||||
|
||||
|
||||
def test_detect_run_source_infers_scheduled_main(monkeypatch):
|
||||
monkeypatch.delenv("PERF_RUN_SOURCE", raising=False)
|
||||
monkeypatch.setenv("BUILDKITE_PULL_REQUEST", "false")
|
||||
monkeypatch.setenv("BUILDKITE_BRANCH", "main")
|
||||
monkeypatch.setenv("TEST_SCOPE", "full")
|
||||
|
||||
assert compare_baseline._detect_run_source() == "scheduled_main"
|
||||
|
||||
|
||||
def test_upload_policy_pass_requires_success(monkeypatch):
|
||||
monkeypatch.setattr(compare_baseline, "UPLOAD_POLICY", "pass")
|
||||
|
||||
assert compare_baseline._upload_allowed({"success": True}) is True
|
||||
assert compare_baseline._upload_allowed({"success": False}) is False
|
||||
|
||||
|
||||
def test_upload_policy_always_uploads_failures(monkeypatch):
|
||||
monkeypatch.setattr(compare_baseline, "UPLOAD_POLICY", "always")
|
||||
|
||||
assert compare_baseline._upload_allowed({"success": False}) is True
|
||||
|
||||
|
||||
def test_normalized_record_includes_source_metadata(monkeypatch):
|
||||
monkeypatch.setenv("PERF_RUN_SOURCE", "pr")
|
||||
monkeypatch.setenv("BUILDKITE_BRANCH", "feature/perf")
|
||||
monkeypatch.setenv("TEST_SCOPE", "direct")
|
||||
monkeypatch.setenv("BUILDKITE_BUILD_URL", "https://buildkite.example/build")
|
||||
monkeypatch.setenv("BUILDKITE_BUILD_ID", "build-1")
|
||||
monkeypatch.setenv("BUILDKITE_JOB_ID", "job-1")
|
||||
|
||||
record = compare_baseline.normalize_performance_result(_raw_result())
|
||||
|
||||
assert record["run_source"] == "pr"
|
||||
assert record["baseline_eligible"] is False
|
||||
assert record["branch"] == "feature/perf"
|
||||
assert record["pr_number"] == "123"
|
||||
assert record["test_scope"] == "direct"
|
||||
assert record["build_url"] == "https://buildkite.example/build"
|
||||
assert record["build_id"] == "build-1"
|
||||
assert record["job_id"] == "job-1"
|
||||
|
||||
|
||||
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
|
||||
|
||||
@@ -36,8 +36,8 @@ class FakeStore(PerformanceDataStore):
|
||||
return records
|
||||
|
||||
|
||||
def _record(model_id, gpu_type, ts, commit, latency, throughput, success=True):
|
||||
return {
|
||||
def _record(model_id, gpu_type, ts, commit, latency, throughput, success=True, **metadata):
|
||||
record = {
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"timestamp": ts,
|
||||
@@ -50,6 +50,8 @@ def _record(model_id, gpu_type, ts, commit, latency, throughput, success=True):
|
||||
"vae_decode_time_s": 3.0,
|
||||
"success": success,
|
||||
}
|
||||
record.update(metadata)
|
||||
return record
|
||||
|
||||
|
||||
def test_summary_endpoint_returns_latest_group_status():
|
||||
@@ -87,6 +89,46 @@ def test_summary_status_is_independent_of_days_window():
|
||||
assert len(trends["groups"][0]["points"]) == 1
|
||||
|
||||
|
||||
def test_dashboard_endpoints_filter_and_return_run_source_metadata():
|
||||
app = create_app(FakeStore([
|
||||
_record(
|
||||
"wan",
|
||||
"NVIDIA L40S",
|
||||
"2026-01-01T00:00:00+00:00",
|
||||
"a" * 40,
|
||||
10.0,
|
||||
10.0,
|
||||
run_source="pr",
|
||||
pr_number="123",
|
||||
branch="feature/perf",
|
||||
baseline_eligible=False,
|
||||
),
|
||||
_record(
|
||||
"wan",
|
||||
"NVIDIA L40S",
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"b" * 40,
|
||||
11.0,
|
||||
9.0,
|
||||
run_source="scheduled_main",
|
||||
baseline_eligible=True,
|
||||
),
|
||||
]))
|
||||
client = TestClient(app)
|
||||
|
||||
summary = client.get("/api/performance/summary", params={"run_source": "pr"}).json()
|
||||
trends = client.get("/api/performance/trends", params={"run_source": "pr"}).json()
|
||||
|
||||
assert summary["count"] == 1
|
||||
assert summary["rows"][0]["run_source"] == "pr"
|
||||
assert summary["rows"][0]["pr_number"] == "123"
|
||||
assert summary["rows"][0]["baseline_n"] == 1
|
||||
assert summary["rows"][0]["metrics"]["latency"]["baseline"] == 11.0
|
||||
assert summary["filters"]["run_source"] == "pr"
|
||||
assert trends["count"] == 1
|
||||
assert trends["groups"][0]["points"][0]["run_source"] == "pr"
|
||||
|
||||
|
||||
def test_records_and_trends_endpoints_filter_by_model_and_gpu():
|
||||
app = create_app(FakeStore([
|
||||
_record("wan", "NVIDIA L40S", "2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
|
||||
@@ -3,8 +3,8 @@ from fastvideo.performance_dashboard.service import build_latest_summary, build_
|
||||
from fastvideo.tests.performance import hf_store
|
||||
|
||||
|
||||
def _record(ts, commit, latency, throughput, success=True):
|
||||
return {
|
||||
def _record(ts, commit, latency, throughput, success=True, **metadata):
|
||||
record = {
|
||||
"model_id": "wan-t2v-1.3b-2gpu",
|
||||
"gpu_type": "NVIDIA L40S",
|
||||
"timestamp": ts,
|
||||
@@ -17,6 +17,8 @@ def _record(ts, commit, latency, throughput, success=True):
|
||||
"vae_decode_time_s": 3.0,
|
||||
"success": success,
|
||||
}
|
||||
record.update(metadata)
|
||||
return record
|
||||
|
||||
|
||||
def test_build_latest_summary_uses_previous_successful_records_for_baseline():
|
||||
@@ -50,6 +52,37 @@ def test_build_latest_summary_status_uses_latest_record_success_field():
|
||||
assert rows[0]["success"] is False
|
||||
|
||||
|
||||
def test_build_latest_summary_run_source_filter_keeps_canonical_baseline():
|
||||
records = [
|
||||
_record(
|
||||
"2026-01-01T00:00:00+00:00",
|
||||
"a" * 40,
|
||||
10.0,
|
||||
10.0,
|
||||
run_source="scheduled_main",
|
||||
baseline_eligible=True,
|
||||
),
|
||||
_record(
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"b" * 40,
|
||||
11.0,
|
||||
9.0,
|
||||
run_source="pr",
|
||||
baseline_eligible=False,
|
||||
pr_number="123",
|
||||
),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05, run_source="pr")
|
||||
|
||||
assert len(rows) == 1
|
||||
assert rows[0]["run_source"] == "pr"
|
||||
assert rows[0]["pr_number"] == "123"
|
||||
assert rows[0]["baseline_n"] == 1
|
||||
assert rows[0]["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert rows[0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
def test_filter_records_and_trends_preserve_metric_points():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
@@ -64,6 +97,34 @@ def test_filter_records_and_trends_preserve_metric_points():
|
||||
assert trends[0]["points"][1]["metrics"]["latency"] == 12.0
|
||||
|
||||
|
||||
def test_trends_include_source_metadata_with_legacy_defaults():
|
||||
records = [
|
||||
_record(
|
||||
"2026-01-01T00:00:00+00:00",
|
||||
"a" * 40,
|
||||
10.0,
|
||||
10.0,
|
||||
run_source="pr",
|
||||
baseline_eligible=False,
|
||||
pr_number="123",
|
||||
branch="feature/dashboard",
|
||||
build_url="https://buildkite.example/build",
|
||||
),
|
||||
_record("2026-01-02T00:00:00+00:00", "b" * 40, 12.0, 8.0),
|
||||
]
|
||||
|
||||
filtered = filter_records(records, run_source="pr")
|
||||
trends = build_trends(records)
|
||||
|
||||
assert len(filtered) == 1
|
||||
assert trends[0]["points"][0]["run_source"] == "pr"
|
||||
assert trends[0]["points"][0]["pr_number"] == "123"
|
||||
assert trends[0]["points"][0]["branch"] == "feature/dashboard"
|
||||
assert trends[0]["points"][0]["build_url"] == "https://buildkite.example/build"
|
||||
assert trends[0]["points"][1]["run_source"] == "unknown"
|
||||
assert trends[0]["points"][1]["baseline_eligible"] is True
|
||||
|
||||
|
||||
def test_hf_token_resolution_accepts_standard_env_names(monkeypatch):
|
||||
for env_var in hf_store.HF_TOKEN_ENV_VARS:
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
@@ -71,3 +132,28 @@ def test_hf_token_resolution_accepts_standard_env_names(monkeypatch):
|
||||
monkeypatch.setenv("HF_TOKEN", "hf_local")
|
||||
|
||||
assert hf_store.resolve_hf_token() == "hf_local"
|
||||
|
||||
|
||||
def test_load_records_can_filter_baseline_eligible_records(tmp_path):
|
||||
model_dir = tmp_path / "wan"
|
||||
model_dir.mkdir()
|
||||
(model_dir / "pr.json").write_text(
|
||||
'{"timestamp": "2026-01-01T00:00:00+00:00", "success": true, "baseline_eligible": false}',
|
||||
encoding="utf-8",
|
||||
)
|
||||
(model_dir / "main.json").write_text(
|
||||
'{"timestamp": "2026-01-02T00:00:00+00:00", "success": true, "baseline_eligible": true}',
|
||||
encoding="utf-8",
|
||||
)
|
||||
(model_dir / "legacy.json").write_text(
|
||||
'{"timestamp": "2026-01-03T00:00:00+00:00", "success": true}',
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
records = hf_store.load_records(str(tmp_path), successful_only=True, baseline_eligible_only=True)
|
||||
|
||||
assert len(records) == 2
|
||||
assert {record["timestamp"] for record in records} == {
|
||||
"2026-01-02T00:00:00+00:00",
|
||||
"2026-01-03T00:00:00+00:00",
|
||||
}
|
||||
|
||||
@@ -98,10 +98,15 @@ def _extract_component_times(result: dict) -> dict[str, float | None]:
|
||||
return component_times
|
||||
logger.info("Discovered pipeline stages: %s", list(stages.keys()))
|
||||
for stage_name, stage_data in stages.items():
|
||||
metric_key = STAGE_METRIC_MAP.get(stage_name)
|
||||
if not isinstance(stage_data, Mapping):
|
||||
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)
|
||||
if metric_key is None:
|
||||
logger.debug("Unmapped stage '%s' (%.3fs)",
|
||||
logger.debug("Unmapped stage '%s' class '%s' (%.3fs)",
|
||||
stage_name,
|
||||
stage_class,
|
||||
stage_data.get("execution_time", 0))
|
||||
continue
|
||||
elapsed = stage_data.get("execution_time")
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import PipelineLoggingInfo
|
||||
from fastvideo.tests.performance.test_inference_performance import _extract_component_times
|
||||
|
||||
|
||||
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")
|
||||
|
||||
assert _extract_component_times({"logging_info": logging_info}) == {
|
||||
"text_encoder_time_s": 1.25,
|
||||
"dit_time_s": None,
|
||||
"vae_decode_time_s": None,
|
||||
}
|
||||
|
||||
|
||||
def test_extract_component_times_uses_stage_class_for_pipeline_stage_keys():
|
||||
# Regression guard for #1377: pre-fix code looked up the pipeline stage key
|
||||
# and returned all component metrics as None for this shape.
|
||||
result = {
|
||||
"logging_info": {
|
||||
"stages": {
|
||||
"prompt_encoding_stage": {
|
||||
"execution_time": 1.2,
|
||||
"stage_class": "TextEncodingStage",
|
||||
},
|
||||
"denoising_stage": {
|
||||
"execution_time": 3.4,
|
||||
"stage_class": "DenoisingStage",
|
||||
},
|
||||
"decoding_stage": {
|
||||
"execution_time": 0.8,
|
||||
"stage_class": "DecodingStage",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
assert _extract_component_times(result) == {
|
||||
"text_encoder_time_s": 1.2,
|
||||
"dit_time_s": 3.4,
|
||||
"vae_decode_time_s": 0.8,
|
||||
}
|
||||
|
||||
|
||||
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.
|
||||
result = {
|
||||
"logging_info": {
|
||||
"stages": {
|
||||
"TextEncodingStage": {"execution_time": 1.0},
|
||||
"DenoisingStage": {"execution_time": 2.0},
|
||||
"DecodingStage": {"execution_time": 3.0},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
assert _extract_component_times(result) == {
|
||||
"text_encoder_time_s": 1.0,
|
||||
"dit_time_s": 2.0,
|
||||
"vae_decode_time_s": 3.0,
|
||||
}
|
||||
|
||||
|
||||
def test_extract_component_times_accumulates_duplicate_component_classes():
|
||||
result = {
|
||||
"logging_info": {
|
||||
"stages": {
|
||||
"base_denoising_stage": {
|
||||
"execution_time": 2.0,
|
||||
"stage_class": "DenoisingStage",
|
||||
},
|
||||
"refine_denoising_stage": {
|
||||
"execution_time": 3.5,
|
||||
"stage_class": "DenoisingStage",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
assert _extract_component_times(result) == {
|
||||
"text_encoder_time_s": None,
|
||||
"dit_time_s": 5.5,
|
||||
"vae_decode_time_s": None,
|
||||
}
|
||||
|
||||
|
||||
def test_extract_component_times_ignores_unmapped_stages():
|
||||
# Generator-side bookkeeping timings are intentionally excluded from the
|
||||
# component gates.
|
||||
result = {
|
||||
"logging_info": {
|
||||
"stages": {
|
||||
"PostDecodeFrameProcessStage": {"execution_time": 0.2},
|
||||
"VideoSaveStage": {"execution_time": 0.4},
|
||||
"AudioMuxStage": {"execution_time": 0.1},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
assert _extract_component_times(result) == {
|
||||
"text_encoder_time_s": None,
|
||||
"dit_time_s": None,
|
||||
"vae_decode_time_s": None,
|
||||
}
|
||||
|
||||
|
||||
def test_extract_component_times_skips_malformed_stage_data():
|
||||
result = {
|
||||
"logging_info": {
|
||||
"stages": {
|
||||
"prompt_encoding_stage": None,
|
||||
"denoising_stage": "not-a-stage-metric-dict",
|
||||
"decoding_stage": {
|
||||
"execution_time": 0.8,
|
||||
"stage_class": "DecodingStage",
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
assert _extract_component_times(result) == {
|
||||
"text_encoder_time_s": None,
|
||||
"dit_time_s": None,
|
||||
"vae_decode_time_s": 0.8,
|
||||
}
|
||||
@@ -14,6 +14,12 @@ Defaults:
|
||||
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
|
||||
- `PERF_MAX_REGRESSION=0.05`
|
||||
|
||||
Records can include source metadata:
|
||||
|
||||
- `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
|
||||
|
||||
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
|
||||
configured dataset repo requires authenticated access:
|
||||
|
||||
@@ -68,13 +74,31 @@ ngrok http 8000
|
||||
The ngrok URL will serve the dashboard UI and all `/api/performance/*`
|
||||
endpoints from the same local port.
|
||||
|
||||
## Dashboard Behavior
|
||||
|
||||
The dashboard supports model, GPU, source, and day-window filters.
|
||||
|
||||
Trend charts show metric-specific axes and exact point details on hover/focus:
|
||||
|
||||
- metric value and unit
|
||||
- timestamp
|
||||
- commit SHA
|
||||
- run source
|
||||
- stored status
|
||||
- baseline eligibility
|
||||
- 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.
|
||||
|
||||
## API
|
||||
|
||||
- `GET /api/performance/health`
|
||||
- `POST /api/performance/refresh`
|
||||
- `GET /api/performance/summary?days=90`
|
||||
- `GET /api/performance/trends?days=90`
|
||||
- `GET /api/performance/records?days=90`
|
||||
- `GET /api/performance/summary?days=90&run_source=pr`
|
||||
- `GET /api/performance/trends?days=90&run_source=scheduled_main`
|
||||
- `GET /api/performance/records?days=90&run_source=local`
|
||||
|
||||
The current v1 grouping key is `(model_id, gpu_type)`. Baselines are computed
|
||||
from the latest five previous successful records in each group.
|
||||
from the latest five previous successful records in each group for dashboard
|
||||
context. CI gating uses only records marked `baseline_eligible=true`.
|
||||
|
||||
@@ -1,8 +1,63 @@
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
|
||||
import { fetchSummary, fetchTrends, refreshData, SummaryResponse, TrendGroup } from "./api";
|
||||
import { fetchSummary, fetchTrends, refreshData, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
|
||||
|
||||
const METRIC_KEYS = ["latency", "throughput", "memory", "text_encoder_time_s", "dit_time_s", "vae_decode_time_s"];
|
||||
const RUN_SOURCES: Array<{ value: "" | RunSource; label: string }> = [
|
||||
{ value: "", label: "All sources" },
|
||||
{ value: "scheduled_main", label: "Scheduled main" },
|
||||
{ value: "pr", label: "PR" },
|
||||
{ value: "local", label: "Local" },
|
||||
{ value: "unknown", label: "Unknown" }
|
||||
];
|
||||
|
||||
const METRIC_DEFINITIONS: Record<
|
||||
string,
|
||||
{
|
||||
label: string;
|
||||
unit: string;
|
||||
precision: number;
|
||||
tooltipPrecision: number;
|
||||
secondary?: (value: number) => string;
|
||||
}
|
||||
> = {
|
||||
latency: {
|
||||
label: "Latency",
|
||||
unit: "s",
|
||||
precision: 2,
|
||||
tooltipPrecision: 3,
|
||||
secondary: (value) => `${formatNumber(value * 1000, 0)} ms`
|
||||
},
|
||||
throughput: { label: "Throughput", unit: "FPS", precision: 2, tooltipPrecision: 3 },
|
||||
memory: {
|
||||
label: "Memory",
|
||||
unit: "MB",
|
||||
precision: 0,
|
||||
tooltipPrecision: 1,
|
||||
secondary: (value) => `${formatNumber(value / 1024, 2)} GB`
|
||||
},
|
||||
text_encoder_time_s: {
|
||||
label: "Text Encoder",
|
||||
unit: "s",
|
||||
precision: 2,
|
||||
tooltipPrecision: 3,
|
||||
secondary: (value) => `${formatNumber(value * 1000, 0)} ms`
|
||||
},
|
||||
dit_time_s: {
|
||||
label: "DiT",
|
||||
unit: "s",
|
||||
precision: 2,
|
||||
tooltipPrecision: 3,
|
||||
secondary: (value) => `${formatNumber(value * 1000, 0)} ms`
|
||||
},
|
||||
vae_decode_time_s: {
|
||||
label: "VAE Decode",
|
||||
unit: "s",
|
||||
precision: 2,
|
||||
tooltipPrecision: 3,
|
||||
secondary: (value) => `${formatNumber(value * 1000, 0)} ms`
|
||||
}
|
||||
};
|
||||
|
||||
function formatNumber(value: number | null | undefined, precision = 2) {
|
||||
if (value === null || value === undefined || Number.isNaN(value)) {
|
||||
@@ -26,52 +81,195 @@ function formatTime(value: string | null | undefined) {
|
||||
return date.toLocaleString();
|
||||
}
|
||||
|
||||
function formatDate(value: string | null | undefined) {
|
||||
if (!value) {
|
||||
return "unknown";
|
||||
}
|
||||
const date = new Date(value);
|
||||
if (Number.isNaN(date.getTime())) {
|
||||
return value;
|
||||
}
|
||||
return date.toLocaleDateString(undefined, { month: "short", day: "numeric" });
|
||||
}
|
||||
|
||||
function runSourceLabel(value: string | null | undefined) {
|
||||
if (value === "scheduled_main") {
|
||||
return "Scheduled main";
|
||||
}
|
||||
if (value === "pr") {
|
||||
return "PR";
|
||||
}
|
||||
if (value === "local") {
|
||||
return "Local";
|
||||
}
|
||||
return "Unknown";
|
||||
}
|
||||
|
||||
function metricLabel(metricKey: string) {
|
||||
return METRIC_DEFINITIONS[metricKey]?.label ?? metricKey;
|
||||
}
|
||||
|
||||
function formatMetricValue(metricKey: string, value: number | null | undefined, tooltip = false) {
|
||||
const definition = METRIC_DEFINITIONS[metricKey];
|
||||
if (!definition) {
|
||||
return formatNumber(value, tooltip ? 3 : 2);
|
||||
}
|
||||
const formatted = formatNumber(value, tooltip ? definition.tooltipPrecision : definition.precision);
|
||||
return formatted === "n/a" ? formatted : `${formatted} ${definition.unit}`;
|
||||
}
|
||||
|
||||
type ChartPoint = {
|
||||
plotIndex: number;
|
||||
value: number;
|
||||
point: TrendPoint;
|
||||
x: number;
|
||||
y: number;
|
||||
};
|
||||
|
||||
function TrendChart({ group, metricKey }: { group: TrendGroup; metricKey: string }) {
|
||||
const [activePoint, setActivePoint] = useState<ChartPoint | null>(null);
|
||||
const points = group.points
|
||||
.map((point, index) => ({
|
||||
index,
|
||||
value: point.metrics[metricKey],
|
||||
success: point.success
|
||||
.map((point) => ({
|
||||
point,
|
||||
value: point.metrics[metricKey]
|
||||
}))
|
||||
.filter((point) => point.value !== null && point.value !== undefined) as Array<{
|
||||
index: number;
|
||||
point: TrendPoint;
|
||||
value: number;
|
||||
success: boolean;
|
||||
}>;
|
||||
|
||||
if (points.length === 0) {
|
||||
return <div className="empty-chart">No data</div>;
|
||||
}
|
||||
|
||||
const width = 280;
|
||||
const height = 96;
|
||||
const pad = 12;
|
||||
const width = 360;
|
||||
const height = 190;
|
||||
const margin = { top: 16, right: 18, bottom: 34, left: 54 };
|
||||
const plotWidth = width - margin.left - margin.right;
|
||||
const plotHeight = height - margin.top - margin.bottom;
|
||||
const min = Math.min(...points.map((point) => point.value));
|
||||
const max = Math.max(...points.map((point) => point.value));
|
||||
const span = max - min || 1;
|
||||
const maxIndex = Math.max(...points.map((point) => point.index)) || 1;
|
||||
const xy = (point: { index: number; value: number }) => {
|
||||
const x = pad + (point.index / maxIndex) * (width - pad * 2);
|
||||
const y = height - pad - ((point.value - min) / span) * (height - pad * 2);
|
||||
return `${x},${y}`;
|
||||
};
|
||||
const xDenominator = Math.max(points.length - 1, 1);
|
||||
const yTicks = [max, min + span / 2, min];
|
||||
const chartPoints: ChartPoint[] = points.map((point, plotIndex) => {
|
||||
const x = margin.left + (plotIndex / xDenominator) * plotWidth;
|
||||
const y = margin.top + (1 - (point.value - min) / span) * plotHeight;
|
||||
return { ...point, plotIndex, x, y };
|
||||
});
|
||||
const rawXTicks = chartPoints.length === 1
|
||||
? [chartPoints[0]]
|
||||
: [chartPoints[0], chartPoints[Math.floor((chartPoints.length - 1) / 2)], chartPoints[chartPoints.length - 1]];
|
||||
const xTicks = rawXTicks.filter(
|
||||
(point, index, items) => items.findIndex((candidate) => candidate.plotIndex === point.plotIndex) === index
|
||||
);
|
||||
const metric = METRIC_DEFINITIONS[metricKey];
|
||||
const selectedPoint = activePoint ?? chartPoints[chartPoints.length - 1];
|
||||
const activePointStyle = activePoint
|
||||
? {
|
||||
left: `${(activePoint.x / width) * 100}%`,
|
||||
top: `${(activePoint.y / height) * 100}%`
|
||||
}
|
||||
: undefined;
|
||||
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}`;
|
||||
|
||||
return (
|
||||
<svg className="trend-chart" viewBox={`0 0 ${width} ${height}`} role="img">
|
||||
<polyline points={points.map(xy).join(" ")} fill="none" stroke="currentColor" strokeWidth="2.2" />
|
||||
{points.map((point) => {
|
||||
const [cx, cy] = xy(point).split(",");
|
||||
return (
|
||||
<circle
|
||||
key={`${point.index}-${point.value}`}
|
||||
cx={cx}
|
||||
cy={cy}
|
||||
r="3"
|
||||
className={point.success ? "point-pass" : "point-fail"}
|
||||
/>
|
||||
);
|
||||
})}
|
||||
</svg>
|
||||
<div className="chart-shell">
|
||||
<svg className="trend-chart" viewBox={`0 0 ${width} ${height}`} role="img" aria-label={ariaLabel}>
|
||||
<line className="axis-line" x1={margin.left} y1={margin.top} x2={margin.left} y2={height - margin.bottom} />
|
||||
<line
|
||||
className="axis-line"
|
||||
x1={margin.left}
|
||||
y1={height - margin.bottom}
|
||||
x2={width - margin.right}
|
||||
y2={height - margin.bottom}
|
||||
/>
|
||||
{yTicks.map((tick) => {
|
||||
const y = margin.top + (1 - (tick - min) / span) * plotHeight;
|
||||
return (
|
||||
<g key={`y-${tick}`}>
|
||||
<line className="grid-line" x1={margin.left} y1={y} x2={width - margin.right} y2={y} />
|
||||
<text className="axis-label" x={margin.left - 8} y={y + 4} textAnchor="end">
|
||||
{formatMetricValue(metricKey, tick)}
|
||||
</text>
|
||||
</g>
|
||||
);
|
||||
})}
|
||||
{xTicks.map((point) => (
|
||||
<text
|
||||
className="axis-label"
|
||||
key={`x-${point.plotIndex}-${point.point.timestamp ?? ""}`}
|
||||
x={point.x}
|
||||
y={height - 10}
|
||||
textAnchor={point.plotIndex === 0 ? "start" : point.plotIndex === chartPoints.length - 1 ? "end" : "middle"}
|
||||
>
|
||||
{formatDate(point.point.timestamp)}
|
||||
</text>
|
||||
))}
|
||||
<polyline
|
||||
points={chartPoints.map((point) => `${point.x},${point.y}`).join(" ")}
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
strokeWidth="2.2"
|
||||
/>
|
||||
{chartPoints.map((point) => {
|
||||
const pointLabel = `${metricLabel(metricKey)} ${formatMetricValue(metricKey, point.value, true)} at ${formatTime(
|
||||
point.point.timestamp
|
||||
)}, commit ${shortSha(point.point.commit_sha)}, ${runSourceLabel(point.point.run_source)}`;
|
||||
return (
|
||||
<g
|
||||
key={`${point.plotIndex}-${point.value}-${point.point.commit_sha ?? ""}`}
|
||||
onMouseEnter={() => setActivePoint(point)}
|
||||
onMouseLeave={() => setActivePoint(null)}
|
||||
>
|
||||
<title>{pointLabel}</title>
|
||||
<circle
|
||||
className="point-hit-area"
|
||||
cx={point.x}
|
||||
cy={point.y}
|
||||
r="12"
|
||||
tabIndex={0}
|
||||
aria-label={pointLabel}
|
||||
onBlur={() => setActivePoint(null)}
|
||||
onFocus={() => setActivePoint(point)}
|
||||
/>
|
||||
<circle
|
||||
cx={point.x}
|
||||
cy={point.y}
|
||||
r={activePoint?.plotIndex === point.plotIndex ? 5 : 4}
|
||||
className={point.point.success ? "point-pass point-marker" : "point-fail point-marker"}
|
||||
/>
|
||||
</g>
|
||||
);
|
||||
})}
|
||||
</svg>
|
||||
{activePoint ? (
|
||||
<div className="hover-tooltip" style={activePointStyle} role="tooltip">
|
||||
<strong>{formatMetricValue(metricKey, activePoint.value, true)}</strong>
|
||||
{metric?.secondary ? <span>{metric.secondary(activePoint.value)}</span> : null}
|
||||
<span>{shortSha(activePoint.point.commit_sha)}</span>
|
||||
<span>{runSourceLabel(activePoint.point.run_source)}</span>
|
||||
</div>
|
||||
) : null}
|
||||
<div className="point-tooltip" aria-live="polite">
|
||||
<strong>
|
||||
{formatMetricValue(metricKey, selectedPoint.value, true)}
|
||||
{metric?.secondary ? <span> ({metric.secondary(selectedPoint.value)})</span> : null}
|
||||
</strong>
|
||||
<span>{formatTime(selectedPoint.point.timestamp)}</span>
|
||||
<span>Commit {shortSha(selectedPoint.point.commit_sha)}</span>
|
||||
<span>{runSourceLabel(selectedPoint.point.run_source)}</span>
|
||||
<span>{selectedPoint.point.success ? "Stored status: pass" : "Stored status: fail"}</span>
|
||||
<span>{selectedPoint.point.baseline_eligible ? "Baseline eligible" : "Not baseline eligible"}</span>
|
||||
{selectedPoint.point.pr_number ? <span>PR #{selectedPoint.point.pr_number}</span> : null}
|
||||
{selectedPoint.point.branch ? <span>Branch {selectedPoint.point.branch}</span> : null}
|
||||
{selectedPoint.point.build_url ? (
|
||||
<a href={selectedPoint.point.build_url} target="_blank" rel="noreferrer">
|
||||
Buildkite
|
||||
</a>
|
||||
) : null}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
@@ -79,6 +277,7 @@ export default function App() {
|
||||
const [days, setDays] = useState(90);
|
||||
const [modelFilter, setModelFilter] = useState("");
|
||||
const [gpuFilter, setGpuFilter] = useState("");
|
||||
const [sourceFilter, setSourceFilter] = useState<"" | RunSource>("");
|
||||
const [summary, setSummary] = useState<SummaryResponse | null>(null);
|
||||
const [trends, setTrends] = useState<TrendGroup[]>([]);
|
||||
const [loading, setLoading] = useState(true);
|
||||
@@ -90,8 +289,8 @@ export default function App() {
|
||||
setError(null);
|
||||
try {
|
||||
const [summaryData, trendData] = await Promise.all([
|
||||
fetchSummary(days, modelFilter || undefined, gpuFilter || undefined),
|
||||
fetchTrends(days, modelFilter || undefined, gpuFilter || undefined)
|
||||
fetchSummary(days, modelFilter || undefined, gpuFilter || undefined, sourceFilter || undefined),
|
||||
fetchTrends(days, modelFilter || undefined, gpuFilter || undefined, sourceFilter || undefined)
|
||||
]);
|
||||
setSummary(summaryData);
|
||||
setTrends(trendData.groups);
|
||||
@@ -119,7 +318,7 @@ export default function App() {
|
||||
load();
|
||||
const interval = window.setInterval(load, 5 * 60 * 1000);
|
||||
return () => window.clearInterval(interval);
|
||||
}, [days, modelFilter, gpuFilter]);
|
||||
}, [days, modelFilter, gpuFilter, sourceFilter]);
|
||||
|
||||
const models = useMemo(() => {
|
||||
const values = new Set(summary?.rows.map((row) => row.model_id) ?? []);
|
||||
@@ -182,6 +381,16 @@ export default function App() {
|
||||
))}
|
||||
</select>
|
||||
</label>
|
||||
<label>
|
||||
Source
|
||||
<select value={sourceFilter} onChange={(event) => setSourceFilter(event.target.value as "" | RunSource)}>
|
||||
{RUN_SOURCES.map((source) => (
|
||||
<option key={source.value || "all"} value={source.value}>
|
||||
{source.label}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</label>
|
||||
</section>
|
||||
|
||||
{error && <div className="notice error">Failed to load dashboard data: {error}</div>}
|
||||
@@ -224,6 +433,8 @@ export default function App() {
|
||||
<th>Model</th>
|
||||
<th>GPU</th>
|
||||
<th>Commit</th>
|
||||
<th>Source</th>
|
||||
<th>Baseline</th>
|
||||
<th>Baseline N</th>
|
||||
<th>Latency</th>
|
||||
<th>Throughput</th>
|
||||
@@ -245,6 +456,10 @@ export default function App() {
|
||||
<td>{row.model_id}</td>
|
||||
<td>{row.gpu_type}</td>
|
||||
<td>{shortSha(row.commit_sha)}</td>
|
||||
<td>
|
||||
<span className={`source-badge source-${row.run_source}`}>{runSourceLabel(row.run_source)}</span>
|
||||
</td>
|
||||
<td>{row.baseline_eligible ? "eligible" : "excluded"}</td>
|
||||
<td>{row.baseline_n}</td>
|
||||
<td>{formatNumber(row.metrics.latency?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
|
||||
@@ -274,7 +489,7 @@ export default function App() {
|
||||
METRIC_KEYS.map((metricKey) => (
|
||||
<article className="trend-card" key={`${group.model_id}-${group.gpu_type}-${metricKey}`}>
|
||||
<div>
|
||||
<h3>{summary?.rows[0]?.metrics[metricKey]?.label ?? metricKey}</h3>
|
||||
<h3>{metricLabel(metricKey)}</h3>
|
||||
<p>
|
||||
{group.model_id} | {group.gpu_type}
|
||||
</p>
|
||||
|
||||
@@ -18,9 +18,19 @@ export type SummaryRow = {
|
||||
regression_threshold_pct: number;
|
||||
computed_regression_status: "pass" | "fail";
|
||||
status: "pass" | "fail";
|
||||
run_source: RunSource;
|
||||
baseline_eligible: boolean;
|
||||
branch: string;
|
||||
pr_number: string;
|
||||
test_scope: string;
|
||||
build_url: string;
|
||||
build_id: string;
|
||||
job_id: string;
|
||||
metrics: Record<string, MetricValue>;
|
||||
};
|
||||
|
||||
export type RunSource = "pr" | "local" | "scheduled_main" | "unknown";
|
||||
|
||||
export type SummaryResponse = {
|
||||
rows: SummaryRow[];
|
||||
count: number;
|
||||
@@ -29,9 +39,11 @@ export type SummaryResponse = {
|
||||
fail: number;
|
||||
};
|
||||
filters: {
|
||||
days: number;
|
||||
days: number | null;
|
||||
trend_window_days?: number;
|
||||
model_id: string | null;
|
||||
gpu_type: string | null;
|
||||
run_source: string | null;
|
||||
};
|
||||
sync: SyncState;
|
||||
};
|
||||
@@ -40,6 +52,14 @@ export type TrendPoint = {
|
||||
timestamp: string | null;
|
||||
commit_sha: string | null;
|
||||
success: boolean;
|
||||
run_source: RunSource;
|
||||
baseline_eligible: boolean;
|
||||
branch: string;
|
||||
pr_number: string;
|
||||
test_scope: string;
|
||||
build_url: string;
|
||||
build_id: string;
|
||||
job_id: string;
|
||||
metrics: Record<string, number | null>;
|
||||
};
|
||||
|
||||
@@ -85,15 +105,15 @@ async function getJson<T>(path: string): Promise<T> {
|
||||
return response.json() as Promise<T>;
|
||||
}
|
||||
|
||||
export async function fetchSummary(days = 90, modelId?: string, gpuType?: string) {
|
||||
export async function fetchSummary(days = 90, modelId?: string, gpuType?: string, runSource?: string) {
|
||||
return getJson<SummaryResponse>(
|
||||
`/api/performance/summary?${params({ days, model_id: modelId, gpu_type: gpuType })}`
|
||||
`/api/performance/summary?${params({ days, model_id: modelId, gpu_type: gpuType, run_source: runSource })}`
|
||||
);
|
||||
}
|
||||
|
||||
export async function fetchTrends(days = 90, modelId?: string, gpuType?: string) {
|
||||
export async function fetchTrends(days = 90, modelId?: string, gpuType?: string, runSource?: string) {
|
||||
return getJson<TrendsResponse>(
|
||||
`/api/performance/trends?${params({ days, model_id: modelId, gpu_type: gpuType })}`
|
||||
`/api/performance/trends?${params({ days, model_id: modelId, gpu_type: gpuType, run_source: runSource })}`
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -84,7 +84,7 @@ h3 {
|
||||
|
||||
.filters {
|
||||
display: grid;
|
||||
grid-template-columns: 120px minmax(220px, 1fr) minmax(220px, 1fr);
|
||||
grid-template-columns: 120px minmax(200px, 1fr) minmax(200px, 1fr) minmax(180px, 0.8fr);
|
||||
gap: 14px;
|
||||
margin-bottom: 18px;
|
||||
}
|
||||
@@ -186,7 +186,7 @@ h3 {
|
||||
|
||||
table {
|
||||
width: 100%;
|
||||
min-width: 900px;
|
||||
min-width: 1120px;
|
||||
border-collapse: collapse;
|
||||
}
|
||||
|
||||
@@ -235,6 +235,33 @@ td {
|
||||
opacity: 0.78;
|
||||
}
|
||||
|
||||
.source-badge {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
border-radius: 999px;
|
||||
padding: 4px 9px;
|
||||
color: #1f2933;
|
||||
background: #e8edf2;
|
||||
font-size: 0.75rem;
|
||||
font-weight: 800;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.source-scheduled_main {
|
||||
color: #065f46;
|
||||
background: #d1fae5;
|
||||
}
|
||||
|
||||
.source-pr {
|
||||
color: #1d4ed8;
|
||||
background: #dbeafe;
|
||||
}
|
||||
|
||||
.source-local {
|
||||
color: #7c2d12;
|
||||
background: #ffedd5;
|
||||
}
|
||||
|
||||
.empty {
|
||||
padding: 28px 16px;
|
||||
color: #607080;
|
||||
@@ -246,7 +273,7 @@ td {
|
||||
|
||||
.trend-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fill, minmax(280px, 1fr));
|
||||
grid-template-columns: repeat(auto-fill, minmax(340px, 1fr));
|
||||
gap: 12px;
|
||||
padding: 14px;
|
||||
}
|
||||
@@ -260,16 +287,101 @@ td {
|
||||
|
||||
.trend-chart {
|
||||
width: 100%;
|
||||
min-height: 96px;
|
||||
min-height: 190px;
|
||||
color: #0f6b8f;
|
||||
overflow: visible;
|
||||
}
|
||||
|
||||
.chart-shell {
|
||||
position: relative;
|
||||
display: grid;
|
||||
gap: 10px;
|
||||
min-width: 0;
|
||||
}
|
||||
|
||||
.axis-line {
|
||||
stroke: #9aa8b6;
|
||||
stroke-width: 1;
|
||||
}
|
||||
|
||||
.grid-line {
|
||||
stroke: #e4e9ee;
|
||||
stroke-width: 1;
|
||||
}
|
||||
|
||||
.axis-label {
|
||||
fill: #667789;
|
||||
font-size: 10px;
|
||||
}
|
||||
|
||||
.point-pass {
|
||||
fill: #0f6b8f;
|
||||
outline: none;
|
||||
}
|
||||
|
||||
.point-fail {
|
||||
fill: #d94f4f;
|
||||
outline: none;
|
||||
}
|
||||
|
||||
.point-hit-area {
|
||||
fill: transparent;
|
||||
outline: none;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.point-marker {
|
||||
pointer-events: none;
|
||||
}
|
||||
|
||||
.point-hit-area:focus + .point-marker {
|
||||
stroke: #17212b;
|
||||
stroke-width: 2;
|
||||
}
|
||||
|
||||
.hover-tooltip {
|
||||
position: absolute;
|
||||
z-index: 2;
|
||||
display: grid;
|
||||
gap: 2px;
|
||||
min-width: 132px;
|
||||
max-width: 190px;
|
||||
border: 1px solid #22313f;
|
||||
border-radius: 6px;
|
||||
padding: 8px 10px;
|
||||
color: #ffffff;
|
||||
background: #17212b;
|
||||
font-size: 0.78rem;
|
||||
pointer-events: none;
|
||||
transform: translate(10px, -100%);
|
||||
box-shadow: 0 10px 24px rgb(15 23 42 / 22%);
|
||||
}
|
||||
|
||||
.hover-tooltip strong {
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
|
||||
.point-tooltip {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(2, minmax(0, 1fr));
|
||||
gap: 4px 10px;
|
||||
border: 1px solid #d5dde5;
|
||||
border-radius: 6px;
|
||||
padding: 10px;
|
||||
color: #263646;
|
||||
background: #f8fafc;
|
||||
font-size: 0.78rem;
|
||||
}
|
||||
|
||||
.point-tooltip strong {
|
||||
grid-column: 1 / -1;
|
||||
color: #132232;
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
|
||||
.point-tooltip a {
|
||||
color: #0f6b8f;
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.empty-chart {
|
||||
|
||||
Reference in New Issue
Block a user