Compare commits

...
9 Commits
Author SHA1 Message Date
SolitaryThinkerandClaude Opus 4.8 3ab66290dc [bugfix] QAD 5090: emit sm_120a in build.sh so attn_qat_infer kernels build
build.sh auto-detected Blackwell (sm_120) and exported TORCH_CUDA_ARCH_LIST=12.0 without the arch-conditional 'a' suffix. CMake's AUTO gate for the attn_qat_infer (modified SageAttention3 FP4) kernels only matches 12.0a/120a/sm_120a, so fp4attn_cuda/fp4quant_cuda were silently skipped and the ATTN_QAT_INFER backend fell back to Flash Attention at runtime. Exporting the env var also bypassed CMake's local-GPU fallback that would otherwise have enabled them.

Mirror the existing 9.0 -> 9.0a Hopper handling for 12.0 -> 12.0a.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-23 07:29:53 +00:00
7fa0fed781 [feat] QAD 5090: env-gate attention torch.compile via FASTVIDEO_DISABLE_ATTENTION_COMPILE
DistributedAttention.forward (and the VSA subclass) are hard-decorated with
@torch.compiler.disable, which keeps attention out of the surrounding
torch.compile graph unconditionally. That blocks the inference compile path
even after the FP4 linear and SageAttention3 graph-break fixes land, since
the attention forward itself can never be traced.

Make the disable conditional on FASTVIDEO_DISABLE_ATTENTION_COMPILE:
- unset / "1" / "true" (default): keep torch.compiler.disable — current behavior
- "0" / "false" / "no" / "off": drop it so attention can fold into the graph

The env var is read at import time (decorators are applied at class
definition), which is the right granularity for the multiproc spawn path:
each worker re-imports and inherits the parent's env.

Co-authored-by: Loay Rashid <42599591+loaydatrain@users.noreply.github.com>
Co-authored-by: Kaiqin Kong <k1kong@ucsd.edu>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2026-06-23 06:44:05 +00:00
loaydatrain 9096310b5c precommit stuff 2026-06-23 06:44:05 +00:00
loaydatrain fce6ed516d Adding TAEHV script, sage_attn3 more torch compile friendly, weight popping+torch compile+single level quant changes to nvfp4_qat_config 2026-06-23 06:44:05 +00:00
Kevin Lin 82ed9fe58d [feat] QAD 5090: FP8 linear layer inference (#1465) 2026-06-22 18:25:06 -07:00
Satyam Srivastava 3d8cc4f0a0 [bugfix] Fix performance component timing extraction (#1473) 2026-06-22 13:05:35 -07:00
Satyam Srivastava 0557f7a7d9 [ci] Add performance dashboard metadata and visualizations (#1470) 2026-06-19 14:27:01 -07:00
dc66cd97ef [feat] QAD 5090: FP8 QAT linear training (14/12) (#1464)
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
Co-authored-by: Loay Rashid <42599591+loaydatrain@users.noreply.github.com>
Co-authored-by: Kaiqin Kong <k1kong@ucsd.edu>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-06-19 01:03:51 -07:00
6da206e196 [feat] QAD 5090: FP4 QAT linear STE for training (13/12) (#1463)
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
Co-authored-by: Loay Rashid <42599591+loaydatrain@users.noreply.github.com>
Co-authored-by: Kaiqin Kong <k1kong@ucsd.edu>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-18 15:55:47 -07:00
30 changed files with 2267 additions and 192 deletions
+1 -1
View File
@@ -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=""
+41 -20
View File
@@ -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. |
+47
View File
@@ -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()
+5
View File
@@ -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
+24 -4
View File
@@ -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()
+24 -2
View File
@@ -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,
+13 -1
View File
@@ -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,
+130
View File
@@ -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
+7 -1
View File
@@ -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)
+241
View File
@@ -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
+140 -50
View File
@@ -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
+23 -23
View File
@@ -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:
+19 -3
View File
@@ -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(),
}
+38 -5
View File
@@ -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))
+23 -5
View File
@@ -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; "
+113 -22
View File
@@ -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)))
+28 -1
View File
@@ -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,
}
+28 -4
View File
@@ -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`.
+250 -35
View File
@@ -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>
+25 -5
View File
@@ -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 })}`
);
}
+116 -4
View File
@@ -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 {