Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3ab66290dc | ||
|
|
7fa0fed781 | ||
|
|
9096310b5c | ||
|
|
fce6ed516d | ||
|
|
82ed9fe58d | ||
|
|
3d8cc4f0a0 | ||
|
|
0557f7a7d9 | ||
|
|
dc66cd97ef | ||
|
|
6da206e196 | ||
|
|
87f98c9b8b | ||
|
|
e60601df7f | ||
|
|
eed9c4bfbf | ||
|
|
1dee77f4a4 | ||
|
|
b80148819c | ||
|
|
c3b971488e | ||
|
|
88e753f281 | ||
|
|
77832059cc |
@@ -0,0 +1,331 @@
|
||||
# Performance Dashboard Memory
|
||||
|
||||
Date: 2026-06-16
|
||||
Branch: `ci/dashboard`
|
||||
|
||||
## Purpose
|
||||
|
||||
This branch adds a local live dashboard for FastVideo performance benchmark
|
||||
history. It is intended for maintainer/operator use: inspect latest benchmark
|
||||
status, compare current values with recent baseline context, and view trends
|
||||
from the Hugging Face performance-tracking dataset.
|
||||
|
||||
The dashboard is a FastAPI + React app. It is separate from the existing
|
||||
Svelte `ui/` app.
|
||||
|
||||
## Main Files Added Or Changed
|
||||
|
||||
Backend:
|
||||
|
||||
- `fastvideo/performance_dashboard/__init__.py`
|
||||
- `fastvideo/performance_dashboard/__main__.py`
|
||||
- `fastvideo/performance_dashboard/api.py`
|
||||
- `fastvideo/performance_dashboard/metrics.py`
|
||||
- `fastvideo/performance_dashboard/service.py`
|
||||
|
||||
Frontend:
|
||||
|
||||
- `performance_dashboard/frontend/package.json`
|
||||
- `performance_dashboard/frontend/package-lock.json`
|
||||
- `performance_dashboard/frontend/tsconfig.json`
|
||||
- `performance_dashboard/frontend/vite.config.ts`
|
||||
- `performance_dashboard/frontend/index.html`
|
||||
- `performance_dashboard/frontend/scripts/build.mjs`
|
||||
- `performance_dashboard/frontend/src/main.tsx`
|
||||
- `performance_dashboard/frontend/src/api.ts`
|
||||
- `performance_dashboard/frontend/src/App.tsx`
|
||||
- `performance_dashboard/frontend/src/styles.css`
|
||||
|
||||
Docs/tests:
|
||||
|
||||
- `performance_dashboard/README.md`
|
||||
- `docs/contributing/performance_benchmarks.md`
|
||||
- `fastvideo/tests/performance/test_dashboard_service.py`
|
||||
- `fastvideo/tests/performance/test_dashboard_api.py`
|
||||
|
||||
Shared HF utility change:
|
||||
|
||||
- `fastvideo/tests/performance/hf_store.py`
|
||||
|
||||
## Data Source
|
||||
|
||||
The source of truth remains the Hugging Face dataset repo used by existing
|
||||
performance CI:
|
||||
|
||||
```text
|
||||
HF_REPO_ID=FastVideo/performance-tracking
|
||||
```
|
||||
|
||||
The dataset stores normalized JSON records emitted by
|
||||
`fastvideo/tests/performance/compare_baseline.py`. The current v1 normalized
|
||||
schema includes:
|
||||
|
||||
- `model_id`
|
||||
- `timestamp`
|
||||
- `commit_sha`
|
||||
- `gpu_type`
|
||||
- `latency`
|
||||
- `throughput`
|
||||
- `memory`
|
||||
- `text_encoder_time_s`
|
||||
- `dit_time_s`
|
||||
- `vae_decode_time_s`
|
||||
- `success`
|
||||
|
||||
Records are grouped by `(model_id, gpu_type)` for v1 dashboard behavior.
|
||||
|
||||
## Local Cache
|
||||
|
||||
The backend syncs the HF dataset to a local cache directory:
|
||||
|
||||
```text
|
||||
PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard
|
||||
```
|
||||
|
||||
If `PERFORMANCE_TRACKING_ROOT` is not set, the dashboard defaults to:
|
||||
|
||||
```text
|
||||
/tmp/fastvideo-perf-dashboard
|
||||
```
|
||||
|
||||
The sync is performed through the existing helper:
|
||||
|
||||
```python
|
||||
fastvideo.tests.performance.hf_store.sync_from_hf(...)
|
||||
```
|
||||
|
||||
The dashboard then loads JSON files from the local cache through:
|
||||
|
||||
```python
|
||||
fastvideo.tests.performance.hf_store.load_records(...)
|
||||
```
|
||||
|
||||
## Authentication
|
||||
|
||||
Originally `hf_store.py` only read `HF_API_KEY`. This caused local dashboard
|
||||
runs to fail when users had standard Hugging Face token variables set.
|
||||
|
||||
`hf_store.py` now resolves tokens from the first available variable in:
|
||||
|
||||
```text
|
||||
HF_API_KEY
|
||||
HUGGINGFACE_HUB_TOKEN
|
||||
HF_TOKEN
|
||||
```
|
||||
|
||||
For local use:
|
||||
|
||||
```bash
|
||||
export HF_TOKEN=hf_...
|
||||
```
|
||||
|
||||
If the HF repo is private or gated, the token must have dataset read access.
|
||||
|
||||
## Backend API
|
||||
|
||||
The FastAPI app is created by:
|
||||
|
||||
```python
|
||||
fastvideo.performance_dashboard.api:create_app
|
||||
```
|
||||
|
||||
The module-level app is:
|
||||
|
||||
```python
|
||||
fastvideo.performance_dashboard.api:app
|
||||
```
|
||||
|
||||
Endpoints:
|
||||
|
||||
- `GET /api/performance/health`
|
||||
- `POST /api/performance/refresh`
|
||||
- `GET /api/performance/records?days=90`
|
||||
- `GET /api/performance/summary?days=90`
|
||||
- `GET /api/performance/trends?days=90`
|
||||
|
||||
`POST /api/performance/refresh` forces a fresh HF sync.
|
||||
|
||||
## Status Semantics
|
||||
|
||||
Important: the dashboard intentionally separates stored CI status from
|
||||
recomputed context.
|
||||
|
||||
Stored status:
|
||||
|
||||
- Comes directly from the latest JSON record's `success` field.
|
||||
- This is what the dashboard displays as `Stored Status`.
|
||||
- This is the primary latest status.
|
||||
|
||||
Recomputed status:
|
||||
|
||||
- Calculated locally from cached records for explanatory context.
|
||||
- Uses the latest record's metric values compared to the median of the latest
|
||||
five previous successful records in the same `(model_id, gpu_type)` group.
|
||||
- Displayed separately as `Recomputed`.
|
||||
- Does not override the stored JSON `success` status.
|
||||
|
||||
This distinction was added after observing that recomputing pass/fail from the
|
||||
local cache can disagree with the status originally uploaded by CI.
|
||||
|
||||
## Time Window Behavior
|
||||
|
||||
The default dashboard time window is 90 days.
|
||||
|
||||
The selected `days` value affects:
|
||||
|
||||
- trend charts
|
||||
- record browsing/filtering
|
||||
|
||||
The selected `days` value does not affect:
|
||||
|
||||
- latest stored status
|
||||
- latest summary baseline context
|
||||
|
||||
Reason: latest status should not change when users widen or narrow the trend
|
||||
window. The API keeps `days` on `/summary` only for shared frontend filter
|
||||
state, but summary loading uses all cached records.
|
||||
|
||||
This fixed a bug where changing from roughly 35 days to 42 days could change
|
||||
the latest status from pass to fail because older records entered the local
|
||||
baseline window.
|
||||
|
||||
## Metric Logic
|
||||
|
||||
Dashboard metric definitions live in:
|
||||
|
||||
```text
|
||||
fastvideo/performance_dashboard/metrics.py
|
||||
```
|
||||
|
||||
Tracked metrics:
|
||||
|
||||
- `latency` lower is better
|
||||
- `throughput` higher is better
|
||||
- `memory` lower is better
|
||||
- `text_encoder_time_s` lower is better
|
||||
- `dit_time_s` lower is better
|
||||
- `vae_decode_time_s` lower is better
|
||||
|
||||
Baseline context uses the median of up to five previous successful records for
|
||||
the same `(model_id, gpu_type)`.
|
||||
|
||||
## Frontend Behavior
|
||||
|
||||
The React app:
|
||||
|
||||
- fetches `/api/performance/summary`
|
||||
- fetches `/api/performance/trends`
|
||||
- displays summary cards
|
||||
- displays latest rows by model/GPU
|
||||
- displays native SVG trend charts
|
||||
- has model/GPU/day filters
|
||||
- includes a refresh button
|
||||
- auto-refreshes every five minutes
|
||||
|
||||
The UI is implemented without a charting library. Trend charts are native SVG
|
||||
in `performance_dashboard/frontend/src/App.tsx`.
|
||||
|
||||
The production frontend build uses `esbuild` through
|
||||
`performance_dashboard/frontend/scripts/build.mjs`. Vite is still used for the
|
||||
dev server and `/api` proxy.
|
||||
|
||||
Why esbuild for production build:
|
||||
|
||||
- Vite/Rollup hit a local macOS native optional dependency code-signing issue
|
||||
in this environment.
|
||||
- Direct esbuild worked reliably and is sufficient for this small dashboard.
|
||||
|
||||
## Static Serving
|
||||
|
||||
After frontend build, the FastAPI server serves:
|
||||
|
||||
- static JS/CSS from `performance_dashboard/frontend/dist/assets`
|
||||
- `performance_dashboard/frontend/dist/index.html` for the dashboard page
|
||||
|
||||
This allows a single local port to serve both the API and UI.
|
||||
|
||||
## Local Run Workflow
|
||||
|
||||
Build frontend:
|
||||
|
||||
```bash
|
||||
cd performance_dashboard/frontend
|
||||
conda run -n fastvideo env PATH=/Applications/Codex.app/Contents/Resources/cua_node/bin:/usr/local/bin:/usr/bin:/bin \
|
||||
/Applications/Codex.app/Contents/Resources/cua_node/bin/npm install
|
||||
conda run -n fastvideo env PATH=/Applications/Codex.app/Contents/Resources/cua_node/bin:/usr/local/bin:/usr/bin:/bin \
|
||||
/Applications/Codex.app/Contents/Resources/cua_node/bin/npm run build
|
||||
```
|
||||
|
||||
Run dashboard:
|
||||
|
||||
```bash
|
||||
export HF_TOKEN=hf_...
|
||||
python -m fastvideo.performance_dashboard --host 0.0.0.0 --port 8000
|
||||
```
|
||||
|
||||
Open locally:
|
||||
|
||||
```text
|
||||
http://127.0.0.1:8000
|
||||
```
|
||||
|
||||
## ngrok Workflow
|
||||
|
||||
`python -m fastvideo.performance_dashboard --host 0.0.0.0 --port 8000`
|
||||
starts the actual local dashboard server.
|
||||
|
||||
`ngrok http 8000` does not start the dashboard. It exposes the already-running
|
||||
local server through a temporary public URL.
|
||||
|
||||
Typical flow:
|
||||
|
||||
```bash
|
||||
python -m fastvideo.performance_dashboard --host 0.0.0.0 --port 8000
|
||||
ngrok http 8000
|
||||
```
|
||||
|
||||
Use the HTTPS URL printed by ngrok to view the dashboard remotely.
|
||||
|
||||
## Verification Commands
|
||||
|
||||
Backend tests:
|
||||
|
||||
```bash
|
||||
conda run -n fastvideo python -m pytest \
|
||||
fastvideo/tests/performance/test_dashboard_service.py \
|
||||
fastvideo/tests/performance/test_dashboard_api.py \
|
||||
-q
|
||||
```
|
||||
|
||||
Expected after latest changes:
|
||||
|
||||
```text
|
||||
8 passed
|
||||
```
|
||||
|
||||
Frontend build:
|
||||
|
||||
```bash
|
||||
cd performance_dashboard/frontend
|
||||
conda run -n fastvideo env PATH=/Applications/Codex.app/Contents/Resources/cua_node/bin:/usr/local/bin:/usr/bin:/bin \
|
||||
/Applications/Codex.app/Contents/Resources/cua_node/bin/npm run build
|
||||
```
|
||||
|
||||
Expected:
|
||||
|
||||
```text
|
||||
tsc && node scripts/build.mjs
|
||||
```
|
||||
|
||||
with exit code 0.
|
||||
|
||||
## Known Notes
|
||||
|
||||
- `performance_dashboard/frontend/node_modules/` and
|
||||
`performance_dashboard/frontend/dist/` are ignored by git.
|
||||
- `npm install` reported two high-severity audit findings in dependency tree.
|
||||
`npm audit fix --force` was not run because it can introduce breaking
|
||||
dependency upgrades.
|
||||
- Existing `fastvideo` package imports may emit platform warnings such as NPU
|
||||
or macOS torch distributed messages. These are not dashboard-specific errors.
|
||||
|
||||
@@ -76,7 +76,7 @@ EFFECTIVE_PR=${BUILDKITE_PULL_REQUEST:-false}
|
||||
if [ "$EFFECTIVE_PR" = "false" ] && [ -n "${PR_NUMBER:-}" ]; then
|
||||
EFFECTIVE_PR=$PR_NUMBER
|
||||
fi
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$EFFECTIVE_PR BUILDKITE_BRANCH=${BUILDKITE_BRANCH:-} TEST_SCOPE=${TEST_SCOPE:-} BUILDKITE_BUILD_URL=${BUILDKITE_BUILD_URL:-} BUILDKITE_BUILD_ID=${BUILDKITE_BUILD_ID:-} BUILDKITE_JOB_ID=${BUILDKITE_JOB_ID:-} IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
|
||||
@@ -21,26 +21,42 @@ 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.
|
||||
|
||||
## Local live dashboard
|
||||
|
||||
For an app-style local dashboard backed by the same HF performance-tracking
|
||||
records, see `performance_dashboard/README.md`. The dashboard provides a
|
||||
FastAPI API plus a React UI and can be exposed with `ngrok` after building the
|
||||
frontend.
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
@@ -61,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
|
||||
|
||||
@@ -87,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
|
||||
|
||||
@@ -131,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
|
||||
@@ -208,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
|
||||
|
||||
@@ -217,8 +242,11 @@ dashboard lists skipped plots for metric series that have no non-null values.
|
||||
| `PERFORMANCE_TRACKING_ROOT` | `/tmp/perf-tracking` | `compare_baseline.py`, `dashboard.py` | Local directory the HF dataset is synced to. |
|
||||
| `PERF_REPORTS_DIR` | `/root/data/perf_reports` | `compare_baseline.py`, `dashboard.py` | Where the Markdown summary and Plotly HTML get written for Buildkite to pick up. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY` | unset | `hf_store.py` | Required for upload (main-branch full-suite only); reads work without it. |
|
||||
| `TEST_SCOPE` | unset | `compare_baseline.py` | Set to `full` together with `BUILDKITE_BRANCH=main` to enable HF persistence. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `PERF_UPLOAD_POLICY` | `never` | `compare_baseline.py` | Upload policy: `never`, `pass`, or `always`. |
|
||||
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
|
||||
| `TEST_SCOPE` | unset | `compare_baseline.py` | CI context used to infer scheduled-main runs together with `BUILDKITE_BRANCH=main`. |
|
||||
| `BUILDKITE_BRANCH`, `BUILDKITE_COMMIT`, `BUILDKITE_PULL_REQUEST` | unset | `compare_baseline.py`, `test_inference_performance.py` | CI metadata stamped into records. |
|
||||
| `DASHBOARD_DAYS` | `30` | `dashboard.py` | Lookback window for the Plotly trend pages. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
|
||||
@@ -11,6 +11,9 @@ This page describes the various options for speeding up generation times in Fast
|
||||
- [Sliding Tile Attention (Archived)](#sliding-tile-attention-archived)
|
||||
- [Sage Attention](#sage-attention)
|
||||
- [Sage Attention 3](#sage-attention-3)
|
||||
|
||||
- [FP8 Weight Quantization](#fp8-weight-quantization)
|
||||
|
||||
- [Adaptive Guidance (CFG gating)](#adaptive-guidance-cfg-gating)
|
||||
|
||||
- [torch.compile](#torch-compile)
|
||||
@@ -24,6 +27,7 @@ This page describes the various options for speeding up generation times in Fast
|
||||
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
|
||||
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
|
||||
- Sage Attention 3: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN_THREE`
|
||||
- Attn-QAT inference (modified SageAttention3 FP4, sm_120/RTX 5090): `FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER`
|
||||
- Video MoBA Attention: `FASTVIDEO_ATTENTION_BACKEND=VMOBA_ATTN`
|
||||
- Sparse Linear Attention: `FASTVIDEO_ATTENTION_BACKEND=SLA_ATTN`
|
||||
- SageSLA Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_SLA_ATTN`
|
||||
@@ -122,6 +126,45 @@ gen.generate_video(prompt="A raccoon in sunflowers", save_video=True)
|
||||
- Per-call cosine similarity vs BF16: ~0.99 (slight quantization error accumulates over denoising steps)
|
||||
- Only supports `headdim >= 128`
|
||||
|
||||
### NVFP4 + Attn-QAT (modified SageAttention3, Blackwell sm_120)
|
||||
|
||||
**`ATTN_QAT_INFER`** with **`transformer_quant=nvfp4_qat`**
|
||||
|
||||
Runs the DiT fully in 4-bit: NVFP4 linear layers (activations quantized on the
|
||||
fly) plus the modified SageAttention3 FP4 attention backend. This is the
|
||||
inference half of the Quantization-Aware Distillation (QAD) recipe and the path
|
||||
used for the RTX 5090 release.
|
||||
|
||||
The `attn_qat_infer` kernel hard-gates on **sm_120 (consumer Blackwell / RTX
|
||||
5090)**; on other GPUs the backend logs a notice and falls back to Flash
|
||||
Attention. See the [Attn-QAT paper](https://arxiv.org/abs/2603.00040).
|
||||
|
||||
Enable both halves — attention via the env var, linear via `transformer_quant`:
|
||||
|
||||
```python
|
||||
import os
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "ATTN_QAT_INFER"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
gen = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
# Wan-2.1 uses the nvfp4_qat config (NVFP4 is LTX2-specific). Pass an
|
||||
# instance — the bare string is not resolved on the from_pretrained path.
|
||||
transformer_quant=get_quantization_config("nvfp4_qat")(),
|
||||
use_fsdp_inference=False, # FSDP shards invalidate the FP4 tensor pointers
|
||||
)
|
||||
gen.generate(request={"prompt": "A raccoon in sunflowers", "output": {"save_video": True}})
|
||||
```
|
||||
|
||||
Or run the example script:
|
||||
|
||||
```bash
|
||||
python examples/inference/optimizations/nvfp4_qat_wan2_1_1_3b.py
|
||||
python examples/inference/optimizations/nvfp4_qat_wan2_1_1_3b.py --bf16 # baseline
|
||||
```
|
||||
|
||||
### Sliding Tile Attention (Archived)
|
||||
|
||||
**`SLIDING_TILE_ATTN`**
|
||||
@@ -178,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()
|
||||
@@ -0,0 +1,95 @@
|
||||
"""NVFP4 + Attn-QAT (modified SageAttention3) inference on Blackwell.
|
||||
|
||||
Runs Wan2.1-T2V-1.3B fully in 4-bit: NVFP4 linear layers (activations
|
||||
quantized on the fly) together with the modified SageAttention3 FP4 attention
|
||||
backend (``ATTN_QAT_INFER``). This is the inference half of the
|
||||
Quantization-Aware Distillation (QAD) recipe.
|
||||
|
||||
Requirements:
|
||||
- RTX 5090 / consumer Blackwell (sm_120a). The attn_qat_infer kernel hard
|
||||
gates on sm_120; on other GPUs it falls back to Flash Attention.
|
||||
- The attn_qat_infer kernel built into fastvideo-kernel (see #1455) and
|
||||
flashinfer for the NVFP4 linear matmuls.
|
||||
|
||||
Usage:
|
||||
python nvfp4_qat_wan2_1_1_3b.py # NVFP4 linear + Attn-QAT attn
|
||||
python nvfp4_qat_wan2_1_1_3b.py --bf16 # BF16 baseline
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import time
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="NVFP4 + Attn-QAT video generation")
|
||||
parser.add_argument("--bf16", action="store_true",
|
||||
help="BF16 baseline (no NVFP4 linear, default attention)")
|
||||
parser.add_argument("--model", default="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
help="Model path or HuggingFace ID")
|
||||
parser.add_argument("--quant-method", default="nvfp4_qat", choices=["nvfp4_qat", "NVFP4"],
|
||||
help="Linear quantization config. Wan-2.1 uses nvfp4_qat (matches its "
|
||||
"to_q/k/v/out + ffn layers); NVFP4 is LTX2-specific and will NOT "
|
||||
"quantize Wan.")
|
||||
parser.add_argument("--compile", action="store_true", help="Enable torch.compile for the DiT")
|
||||
parser.add_argument("--num_gpus", type=int, default=1)
|
||||
parser.add_argument("--infer_steps", type=int, default=50)
|
||||
args = parser.parse_args()
|
||||
|
||||
# The attention backend is selected via env var before the engine starts.
|
||||
# ATTN_QAT_INFER -> AttnQatInferBackend (modified SageAttention3 FP4).
|
||||
if not args.bf16:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "ATTN_QAT_INFER"
|
||||
|
||||
# Import after the env var so the platform picks up the selection.
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
|
||||
mode = "bf16" if args.bf16 else args.quant_method
|
||||
if args.compile:
|
||||
mode += "_compile"
|
||||
print(f"Mode: {mode.upper()}")
|
||||
|
||||
# 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(args.quant_method)()}
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
args.model,
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=args.bf16,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
enable_torch_compile=args.compile,
|
||||
**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 = 2 if args.compile else 1
|
||||
for _ in range(n_warmup):
|
||||
generator.generate(request={"prompt": prompt, "sampling": {"num_inference_steps": 2},
|
||||
"output": {"save_video": False}})
|
||||
|
||||
os.makedirs(OUTPUT_PATH, exist_ok=True)
|
||||
start = time.time()
|
||||
generator.generate(request={
|
||||
"prompt": prompt,
|
||||
"sampling": {"num_inference_steps": args.infer_steps},
|
||||
"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()
|
||||
@@ -0,0 +1,119 @@
|
||||
# MixKit training data (QAD 5090 recipe)
|
||||
|
||||
The QAD 5090 models are distilled from Wan2.1-T2V-1.3B on a MixKit subset at
|
||||
**480×832, 77 frames, 16 fps**. FastVideo training consumes **Parquet** shards of
|
||||
precomputed VAE latents + text embeddings (no text encoder / VAE needed at train
|
||||
time).
|
||||
|
||||
## Option A — download the preprocessed data (recommended)
|
||||
|
||||
The encoded dataset is published on the Hugging Face Hub, ready to train:
|
||||
|
||||
```bash
|
||||
# from the repo root
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
|
||||
```
|
||||
|
||||
This pulls [`weizhou03/HD-Mixkit-Finetune-Wan`](https://huggingface.co/datasets/weizhou03/HD-Mixkit-Finetune-Wan)
|
||||
into `data/HD-Mixkit-Finetune-Wan/`:
|
||||
|
||||
```
|
||||
data/HD-Mixkit-Finetune-Wan/
|
||||
├── combined_parquet_dataset/ # training shards -> point --data_path here
|
||||
│ └── worker_0/data_chunk_*.parquet
|
||||
└── validation_parquet_dataset/ # validation shards
|
||||
└── worker_0/data_chunk_0.parquet
|
||||
```
|
||||
|
||||
Each Parquet row holds the VAE latent bytes + text-embedding bytes (plus
|
||||
shape/dtype metadata), matching FastVideo's standard preprocessing output.
|
||||
|
||||
## Option B — build the Parquet from raw videos
|
||||
|
||||
If you want to reproduce the encoding from your own MixKit videos, arrange them as
|
||||
a `merged` dataset (videos + a captions JSON), then run FastVideo's standard
|
||||
preprocessing to VAE-encode and text-embed them into Parquet:
|
||||
|
||||
```bash
|
||||
GPU_NUM=2
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
|
||||
--model_path "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" \
|
||||
--mode preprocess \
|
||||
--workload_type t2v \
|
||||
--preprocess.video_loader_type torchvision \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path "data/mixkit_raw/" \
|
||||
--preprocess.dataset_output_dir "data/HD-Mixkit-Finetune-Wan/" \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
--preprocess.num_frames 77 \
|
||||
--preprocess.train_fps 16 \
|
||||
--preprocess.samples_per_file 8
|
||||
```
|
||||
|
||||
The raw videos are full-HD MixKit clips (≈1080p/30fps); preprocessing resizes to
|
||||
480×832, resamples to 16 fps, and extracts 77 frames per clip. See
|
||||
[`docs/training/data_preprocess.md`](../../../../../docs/training/data_preprocess.md)
|
||||
for the full parameter reference.
|
||||
|
||||
## Train (QAT finetune)
|
||||
|
||||
With the data in place, run the quantization-aware finetune. The 4-bit attention
|
||||
path is **config-driven** — selected purely by an env var, no monkey-patching:
|
||||
|
||||
```bash
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
|
||||
# or point at your own parquet dir / GPU count:
|
||||
NUM_GPUS=4 bash .../mixkit/finetune_qat.sh data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/
|
||||
```
|
||||
|
||||
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` routes attention through the
|
||||
fake-quantized Triton kernel (straight-through estimator), so the DiT learns to
|
||||
absorb FP4 attention error. This kernel is Triton, so it runs on both `sm_100`
|
||||
(B200/GB200) and `sm_120` (RTX 5090).
|
||||
|
||||
## Train stage 2 (QAT DMD distillation to 3 steps)
|
||||
|
||||
Distill the QAT-finetuned generator down to **3 sampling steps**. Only the
|
||||
generator is quantized (Attn-QAT); the teacher (`real_score`) and critic
|
||||
(`fake_score`) stay full precision. This is enforced in the loader
|
||||
(`component_loader.py`, via the `_loading_teacher_critic_model` flag), so the
|
||||
same global `ATTN_QAT_TRAIN` env reaches **only** the generator — no per-model
|
||||
flags or monkey-patching.
|
||||
|
||||
```bash
|
||||
# generator init = the stage-1 finetune checkpoint
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/distill_dmd_qat.sh \
|
||||
data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/ \
|
||||
checkpoints/wan_t2v_qat_finetune/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors
|
||||
```
|
||||
|
||||
DMD runs a double loop (critic every step, generator every
|
||||
`generator_update_interval`), and validation samples the distilled student at
|
||||
3 steps — the final 4-bit-attention model.
|
||||
|
||||
## Inference (NVFP4 4-bit linear)
|
||||
|
||||
For Wan-2.1, enable the FP4 linear layers with the **`nvfp4_qat`** quantization
|
||||
config (it matches Wan's `to_q/k/v/out` + `ffn` layers; the plain `NVFP4` config
|
||||
is LTX2-specific and will not quantize Wan):
|
||||
|
||||
```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", num_gpus=1,
|
||||
transformer_quant=get_quantization_config("nvfp4_qat")(), # a config instance, not the string
|
||||
use_fsdp_inference=False,
|
||||
)
|
||||
gen.generate(request={"prompt": "...", "output": {"save_video": True}})
|
||||
```
|
||||
|
||||
The loader converts the tagged linear weights to FP4 at load time
|
||||
(`_maybe_convert_model_to_nvfp4`). Combine with
|
||||
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER` on an RTX 5090 (`sm_120`) for the
|
||||
full 4-bit path; on other GPUs the attention falls back to Flash while the FP4
|
||||
linear layers still run. `flashinfer` (and a host C++ compiler for its FP4
|
||||
kernel JIT) are required.
|
||||
@@ -0,0 +1,51 @@
|
||||
#!/bin/bash
|
||||
# QAD recipe stage 2 — quantization-aware DMD distillation of Wan2.1-T2V-1.3B
|
||||
# down to 3 sampling steps, with the GENERATOR in fake-quant Attn-QAT and the
|
||||
# teacher (real_score) + critic (fake_score) at full precision.
|
||||
#
|
||||
# Generator-only QAT is config-driven: FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN
|
||||
# is applied to the generator only, because the loader masks it (and the
|
||||
# nvfp4_qat quant) for the teacher/critic via the `_loading_teacher_critic_model`
|
||||
# flag (see fastvideo/models/loader/component_loader.py). No monkey-patching.
|
||||
#
|
||||
# Init the generator from the stage-1 finetune checkpoint (finetune_qat.sh).
|
||||
# Data: run download_mixkit_data.sh first.
|
||||
#
|
||||
# Verified end-to-end on Blackwell (GB200/sm_100): generator loads with
|
||||
# ATTN_QAT_TRAIN while teacher/critic load full-precision; the DMD double loop
|
||||
# runs (generator updates every generator_update_interval steps, critic every
|
||||
# step), 3-step validation generates videos, checkpoint saved.
|
||||
set -euo pipefail
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN # generator-only (loader-gated)
|
||||
export WANDB_MODE=${WANDB_MODE:-online}
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
BASE="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=${1:-"data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/"}
|
||||
# Generator init weights = the stage-1 QAT-finetune checkpoint.
|
||||
INIT_WEIGHTS=${2:-"checkpoints/wan_t2v_qat_finetune/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors"}
|
||||
VALIDATION_FILE="$(dirname "$0")/../crush_smol/validation.json"
|
||||
NUM_GPUS=${NUM_GPUS:-4}
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
--num_gpus "${NUM_GPUS}" --sp_size 1 --tp_size 1 \
|
||||
--hsdp_replicate_dim "${NUM_GPUS}" --hsdp_shard_dim 1 \
|
||||
--model_path "${BASE}" --pretrained_model_name_or_path "${BASE}" \
|
||||
--real_score_model_path "${BASE}" --fake_score_model_path "${BASE}" \
|
||||
--init_weights_from_safetensors "${INIT_WEIGHTS}" \
|
||||
--data_path "${DATA_DIR}" --dataloader_num_workers 4 \
|
||||
--max_train_steps 2000 --train_batch_size 1 --train_sp_batch_size 1 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--num_latent_t 20 --num_height 480 --num_width 832 --num_frames 77 \
|
||||
--enable_gradient_checkpointing_type full \
|
||||
--log_validation --validation_dataset_file "${VALIDATION_FILE}" \
|
||||
--validation_steps 200 --validation_sampling_steps 3 --validation_guidance_scale 6.0 \
|
||||
--learning_rate 2e-6 --mixed_precision bf16 --weight_decay 0.01 --max_grad_norm 1.0 \
|
||||
--weight_only_checkpointing_steps 500 --training_state_checkpointing_steps 500 \
|
||||
--tracker_project_name wan_t2v_distill_dmd_qat \
|
||||
--output_dir checkpoints/wan_t2v_distill_dmd_qat \
|
||||
--inference_mode False --dit_precision fp32 --ema_start_step 0 --training_cfg_rate 0.0 \
|
||||
--generator_update_interval 5 --real_score_guidance_scale 2.0 \
|
||||
--dmd_denoising_steps '1000,757,522' --min_timestep_ratio 0.02 --max_timestep_ratio 0.98
|
||||
@@ -0,0 +1,22 @@
|
||||
#!/bin/bash
|
||||
# Download the preprocessed MixKit finetune dataset used for the QAD 5090 recipe.
|
||||
#
|
||||
# This is the MixKit subset already VAE-encoded (Wan2.1-T2V-1.3B) and text-embedded
|
||||
# into Parquet shards, so it can be fed straight to training with no further
|
||||
# preprocessing. To build the Parquet from raw videos yourself, see README.md.
|
||||
#
|
||||
# Usage (run from the repo root):
|
||||
# bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh [DATA_ROOT]
|
||||
set -euo pipefail
|
||||
|
||||
DATA_ROOT=${1:-data/HD-Mixkit-Finetune-Wan}
|
||||
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id "weizhou03/HD-Mixkit-Finetune-Wan" \
|
||||
--local_dir "${DATA_ROOT}" \
|
||||
--repo_type "dataset"
|
||||
|
||||
echo "Done."
|
||||
echo " Train data: ${DATA_ROOT}/combined_parquet_dataset"
|
||||
echo " Validation data: ${DATA_ROOT}/validation_parquet_dataset"
|
||||
echo "Point your training script's data path at the combined_parquet_dataset directory."
|
||||
@@ -0,0 +1,44 @@
|
||||
#!/bin/bash
|
||||
# QAD recipe — quantization-aware finetune of Wan2.1-T2V-1.3B with fake-quant
|
||||
# (Attn-QAT) attention.
|
||||
#
|
||||
# The 4-bit attention path is selected purely by env var (config-driven, no
|
||||
# monkey-patching): FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN routes attention
|
||||
# through the fake-quantized Triton kernel (straight-through estimator), so the
|
||||
# DiT learns to absorb FP4 attention error instead of fighting it.
|
||||
#
|
||||
# Data: run download_mixkit_data.sh first (preprocessed Parquet).
|
||||
#
|
||||
# Verified end-to-end on Blackwell (GB200/sm_100): the ATTN_QAT_TRAIN backend is
|
||||
# selected (not a fallback), forward+backward run, loss/grad are healthy, and
|
||||
# validation generates videos. The kernel is Triton so it runs on sm_100 and
|
||||
# sm_120 alike (the FP4 inference kernel, by contrast, is sm_120-only).
|
||||
set -euo pipefail
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN # <-- enables Attn-QAT training
|
||||
export WANDB_MODE=${WANDB_MODE:-online}
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=${1:-"data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/"}
|
||||
VALIDATION_FILE="$(dirname "$0")/../crush_smol/validation.json"
|
||||
NUM_GPUS=${NUM_GPUS:-4}
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
--num_gpus "${NUM_GPUS}" --sp_size "${NUM_GPUS}" --tp_size 1 \
|
||||
--hsdp_replicate_dim 1 --hsdp_shard_dim "${NUM_GPUS}" \
|
||||
--model_path "${MODEL_PATH}" --pretrained_model_name_or_path "${MODEL_PATH}" \
|
||||
--data_path "${DATA_DIR}" --dataloader_num_workers 1 \
|
||||
--max_train_steps 2000 --train_batch_size 1 --train_sp_batch_size 1 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--num_latent_t 20 --num_height 480 --num_width 832 --num_frames 77 \
|
||||
--enable_gradient_checkpointing_type full \
|
||||
--log_validation --validation_dataset_file "${VALIDATION_FILE}" \
|
||||
--validation_steps 200 --validation_sampling_steps 50 --validation_guidance_scale 3.0 \
|
||||
--learning_rate 5e-5 --mixed_precision bf16 --weight_decay 1e-4 --max_grad_norm 1.0 \
|
||||
--weight_only_checkpointing_steps 500 --training_state_checkpointing_steps 500 \
|
||||
--tracker_project_name wan_t2v_qat_finetune --output_dir checkpoints/wan_t2v_qat_finetune \
|
||||
--inference_mode False --training_cfg_rate 0.1 --not_apply_cfg_solver \
|
||||
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
@@ -50,6 +50,21 @@ include_directories(
|
||||
set(FASTVIDEO_KERNEL_BUILD_TK "AUTO" CACHE STRING "Build ThunderKittens kernels: AUTO/ON/OFF")
|
||||
set_property(CACHE FASTVIDEO_KERNEL_BUILD_TK PROPERTY STRINGS AUTO ON OFF)
|
||||
|
||||
set(_FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER_DEFAULT "AUTO")
|
||||
if(DEFINED FASTVIDEO_KERNEL_BUILD_MODIFIED_SAGE3 AND NOT DEFINED CACHE{FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER})
|
||||
set(_FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER_DEFAULT "${FASTVIDEO_KERNEL_BUILD_MODIFIED_SAGE3}")
|
||||
endif()
|
||||
|
||||
set(FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER "${_FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER_DEFAULT}" CACHE STRING
|
||||
"Build attn_qat_infer Blackwell inference kernels: AUTO/ON/OFF")
|
||||
set_property(CACHE FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER PROPERTY STRINGS AUTO ON OFF)
|
||||
|
||||
if(DEFINED FASTVIDEO_KERNEL_BUILD_MODIFIED_SAGE3)
|
||||
message(DEPRECATION
|
||||
"FASTVIDEO_KERNEL_BUILD_MODIFIED_SAGE3 is deprecated. "
|
||||
"Use FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER instead.")
|
||||
endif()
|
||||
|
||||
# Prefer environment variable (used by CI) if CMake var is not explicitly set.
|
||||
if(NOT DEFINED TORCH_CUDA_ARCH_LIST AND DEFINED ENV{TORCH_CUDA_ARCH_LIST})
|
||||
set(TORCH_CUDA_ARCH_LIST "$ENV{TORCH_CUDA_ARCH_LIST}")
|
||||
@@ -57,6 +72,7 @@ endif()
|
||||
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST (cmake/env): ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "FASTVIDEO_KERNEL_BUILD_TK: ${FASTVIDEO_KERNEL_BUILD_TK}")
|
||||
message(STATUS "FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER: ${FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER}")
|
||||
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
if(FASTVIDEO_KERNEL_BUILD_TK STREQUAL "ON")
|
||||
@@ -91,6 +107,54 @@ else()
|
||||
message(STATUS "ThunderKittens kernels: DISABLED (will use Triton fallbacks at runtime)")
|
||||
endif()
|
||||
|
||||
set(ENABLE_ATTN_QAT_INFER OFF)
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
message(STATUS "attn_qat_infer kernels: DISABLED (ROCm build)")
|
||||
else()
|
||||
set(_WANTS_ATTN_QAT_INFER OFF)
|
||||
if(FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER STREQUAL "ON")
|
||||
set(_WANTS_ATTN_QAT_INFER ON)
|
||||
elseif(FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER STREQUAL "AUTO")
|
||||
if(TORCH_CUDA_ARCH_LIST)
|
||||
string(REGEX MATCH
|
||||
"(^|[; ,])((12\\.0a)|(120a)|(sm_120a))([; ,]|$)"
|
||||
_HAS_120A "${TORCH_CUDA_ARCH_LIST}")
|
||||
if(_HAS_120A)
|
||||
set(_WANTS_ATTN_QAT_INFER ON)
|
||||
endif()
|
||||
else()
|
||||
execute_process(
|
||||
COMMAND "${Python_EXECUTABLE}" -c
|
||||
"import torch; print('1' if (torch.cuda.is_available() and torch.version.cuda and torch.cuda.get_device_capability()[0] >= 12) else '0')"
|
||||
OUTPUT_VARIABLE _LOCAL_HAS_BLACKWELL
|
||||
OUTPUT_STRIP_TRAILING_WHITESPACE
|
||||
ERROR_QUIET
|
||||
)
|
||||
if(_LOCAL_HAS_BLACKWELL STREQUAL "1")
|
||||
set(_WANTS_ATTN_QAT_INFER ON)
|
||||
endif()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if(_WANTS_ATTN_QAT_INFER)
|
||||
if(CUDAToolkit_VERSION VERSION_LESS 12.8)
|
||||
message(WARNING
|
||||
"attn_qat_infer kernels require CUDA Toolkit 12.8+. "
|
||||
"Skipping because CUDAToolkit_VERSION=${CUDAToolkit_VERSION}.")
|
||||
else()
|
||||
set(ENABLE_ATTN_QAT_INFER ON)
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if(ENABLE_ATTN_QAT_INFER)
|
||||
message(STATUS "attn_qat_infer kernels: ENABLED")
|
||||
else()
|
||||
message(STATUS
|
||||
"attn_qat_infer kernels: DISABLED "
|
||||
"(requires CUDA 12.8+ and Blackwell sm_120a)")
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# Always try to build the extension if CUDA is available, but conditionally add sources/flags
|
||||
set(BUILD_CXX_KERNELS ON)
|
||||
|
||||
@@ -183,3 +247,74 @@ if(BUILD_CXX_KERNELS)
|
||||
install(TARGETS fastvideo_kernel_ops LIBRARY DESTINATION fastvideo_kernel/_C)
|
||||
endif()
|
||||
|
||||
if(ENABLE_ATTN_QAT_INFER)
|
||||
set(ATTN_QAT_INFER_DIR ${CMAKE_SOURCE_DIR}/attn_qat_infer)
|
||||
set(ATTN_QAT_INFER_INCLUDE_DIRS
|
||||
${ATTN_QAT_INFER_DIR}
|
||||
${CMAKE_SOURCE_DIR}/include/cutlass/include
|
||||
${CMAKE_SOURCE_DIR}/include/cutlass/tools/util/include
|
||||
${TORCH_INCLUDE_DIRS}
|
||||
)
|
||||
set(ATTN_QAT_INFER_CUDA_FLAGS
|
||||
"-O3"
|
||||
"-std=c++17"
|
||||
"-U__CUDA_NO_HALF_OPERATORS__"
|
||||
"-U__CUDA_NO_HALF_CONVERSIONS__"
|
||||
"-U__CUDA_NO_BFLOAT16_OPERATORS__"
|
||||
"-U__CUDA_NO_BFLOAT16_CONVERSIONS__"
|
||||
"-U__CUDA_NO_BFLOAT162_OPERATORS__"
|
||||
"-U__CUDA_NO_BFLOAT162_CONVERSIONS__"
|
||||
"--expt-relaxed-constexpr"
|
||||
"--expt-extended-lambda"
|
||||
"--use_fast_math"
|
||||
"--ptxas-options=--verbose,--warn-on-local-memory-usage"
|
||||
"-lineinfo"
|
||||
"-DCUTLASS_DEBUG_TRACE_LEVEL=0"
|
||||
"-DNDEBUG"
|
||||
"-DQBLKSIZE=128"
|
||||
"-DKBLKSIZE=128"
|
||||
"-DCTA256"
|
||||
"-DDQINRMEM"
|
||||
)
|
||||
|
||||
Python_add_library(fp4attn_cuda MODULE WITH_SOABI
|
||||
attn_qat_infer/blackwell/api.cu
|
||||
)
|
||||
target_include_directories(fp4attn_cuda PRIVATE ${ATTN_QAT_INFER_INCLUDE_DIRS})
|
||||
target_compile_definitions(fp4attn_cuda PRIVATE TORCH_EXTENSION_NAME=fp4attn_cuda)
|
||||
target_compile_options(fp4attn_cuda PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-O3 -std=c++17>
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:${ATTN_QAT_INFER_CUDA_FLAGS}>
|
||||
)
|
||||
set_target_properties(fp4attn_cuda PROPERTIES
|
||||
CUDA_ARCHITECTURES "120a"
|
||||
CXX_STANDARD 17
|
||||
CUDA_STANDARD 17
|
||||
)
|
||||
target_link_libraries(fp4attn_cuda PRIVATE ${TORCH_LIBRARIES} CUDA::cudart CUDA::cuda_driver)
|
||||
|
||||
Python_add_library(fp4quant_cuda MODULE WITH_SOABI
|
||||
attn_qat_infer/quantization/fp4_quantization_4d.cu
|
||||
)
|
||||
target_include_directories(fp4quant_cuda PRIVATE ${ATTN_QAT_INFER_INCLUDE_DIRS})
|
||||
target_compile_definitions(fp4quant_cuda PRIVATE TORCH_EXTENSION_NAME=fp4quant_cuda)
|
||||
target_compile_options(fp4quant_cuda PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-O3 -std=c++17>
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:${ATTN_QAT_INFER_CUDA_FLAGS}>
|
||||
)
|
||||
set_target_properties(fp4quant_cuda PROPERTIES
|
||||
CUDA_ARCHITECTURES "120a"
|
||||
CXX_STANDARD 17
|
||||
CUDA_STANDARD 17
|
||||
)
|
||||
target_link_libraries(fp4quant_cuda PRIVATE ${TORCH_LIBRARIES} CUDA::cudart CUDA::cuda_driver)
|
||||
|
||||
if(TORCH_PYTHON_LIBRARY_PATH)
|
||||
target_link_libraries(fp4attn_cuda PRIVATE "${TORCH_PYTHON_LIBRARY_PATH}")
|
||||
target_link_libraries(fp4quant_cuda PRIVATE "${TORCH_PYTHON_LIBRARY_PATH}")
|
||||
endif()
|
||||
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
|
||||
@@ -2,5 +2,6 @@ include LICENSE
|
||||
include README.md
|
||||
include pyproject.toml
|
||||
recursive-include python/fastvideo_kernel *.py
|
||||
recursive-include attn_qat_infer *.py *.cu *.cuh *.cpp *.h
|
||||
recursive-include csrc *.cu *.cuh *.cpp *.h
|
||||
recursive-include include/tk *.cu *.cuh *.cpp *.h *.src
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
from .api import sageattn_blackwell
|
||||
@@ -0,0 +1,189 @@
|
||||
# Modified from the original SageATtention3 code
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch.nn.functional as F
|
||||
from typing import Tuple
|
||||
from torch.nn.functional import scaled_dot_product_attention as sdpa
|
||||
import fp4attn_cuda
|
||||
import fp4quant_cuda
|
||||
|
||||
# Centralized block size configuration for sageattn_blackwell kernels
|
||||
# These should match the values in fastvideo/attention/backends/sageattn/blackwell/block_config.h
|
||||
BLOCK_M = 128 # Block size for M dimension (query sequence length)
|
||||
BLOCK_N = 128 # Block size for N dimension (key/value sequence length)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def group_mean_kernel(
|
||||
q_ptr,
|
||||
q_out_ptr,
|
||||
qm_out_ptr,
|
||||
B, H, L, D: tl.constexpr,
|
||||
stride_qb, stride_qh, stride_ql, stride_qd,
|
||||
stride_qmb, stride_qmh, stride_qml, stride_qmd,
|
||||
GROUP_SIZE: tl.constexpr
|
||||
):
|
||||
pid_b = tl.program_id(0)
|
||||
pid_h = tl.program_id(1)
|
||||
pid_group = tl.program_id(2)
|
||||
|
||||
group_start = pid_group * GROUP_SIZE
|
||||
offsets = group_start + tl.arange(0, GROUP_SIZE)
|
||||
|
||||
q_offsets = pid_b * stride_qb + pid_h * stride_qh + offsets[:, None] * stride_ql + tl.arange(0, D)[None, :] * stride_qd
|
||||
q_group = tl.load(q_ptr + q_offsets)
|
||||
|
||||
qm_group = tl.sum(q_group, axis=0) / GROUP_SIZE
|
||||
|
||||
q_group = q_group - qm_group
|
||||
tl.store(q_out_ptr + q_offsets, q_group)
|
||||
|
||||
qm_offset = pid_b * stride_qmb + pid_h * stride_qmh + pid_group * stride_qml + tl.arange(0, D) * stride_qmd
|
||||
tl.store(qm_out_ptr + qm_offset, qm_group)
|
||||
|
||||
|
||||
def triton_group_mean(q: torch.Tensor):
|
||||
B, H, L, D = q.shape
|
||||
GROUP_SIZE = BLOCK_M
|
||||
num_groups = L // GROUP_SIZE
|
||||
|
||||
q_out = torch.empty_like(q) # [B, H, L, D]
|
||||
qm = torch.empty(B, H, num_groups, D, device=q.device, dtype=q.dtype)
|
||||
|
||||
grid = (B, H, num_groups)
|
||||
|
||||
group_mean_kernel[grid](
|
||||
q, q_out, qm,
|
||||
B, H, L, D,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
qm.stride(0), qm.stride(1), qm.stride(2), qm.stride(3),
|
||||
GROUP_SIZE=GROUP_SIZE
|
||||
)
|
||||
return q_out, qm
|
||||
|
||||
|
||||
def preprocess_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, per_block_mean: bool = True, enable_smoothing_q: bool = False, enable_smoothing_k: bool = False):
|
||||
|
||||
def pad_to_block_size(x):
|
||||
L = x.size(2)
|
||||
pad_len = (BLOCK_M - L % BLOCK_M) % BLOCK_M
|
||||
if pad_len == 0:
|
||||
return x.contiguous()
|
||||
return F.pad(x, (0, 0, 0, pad_len), value=0).contiguous()
|
||||
|
||||
if enable_smoothing_k:
|
||||
k -= k.mean(dim=-2, keepdim=True)
|
||||
q, k, v = map(lambda x: pad_to_block_size(x), [q, k, v])
|
||||
if per_block_mean and enable_smoothing_q:
|
||||
q, qm = triton_group_mean(q)
|
||||
elif enable_smoothing_q:
|
||||
qm = q.mean(dim=-2, keepdim=True)
|
||||
q = q - qm
|
||||
if enable_smoothing_q:
|
||||
delta_s = torch.matmul(qm, k.transpose(-2, -1)).to(torch.float32).contiguous()
|
||||
else: # used to disable q smoothing
|
||||
B, H, L, D = q.shape
|
||||
delta_s = torch.zeros((B, H, L // BLOCK_M, k.shape[2]), device=q.device, dtype=torch.float32)
|
||||
|
||||
return q, k, v, delta_s
|
||||
|
||||
def scale_and_quant_fp4(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.ndim == 4
|
||||
B, H, N, D = x.shape
|
||||
packed_fp4 = torch.empty((B, H, N, D // 2), device=x.device, dtype=torch.uint8)
|
||||
fp8_scale = torch.empty((B, H, N, D // 16), device=x.device, dtype=torch.float8_e4m3fn)
|
||||
fp4quant_cuda.scaled_fp4_quant(x, packed_fp4, fp8_scale, 1)
|
||||
return packed_fp4, fp8_scale
|
||||
|
||||
def scale_and_quant_fp4_permute(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.ndim == 4
|
||||
B, H, N, D = x.shape
|
||||
packed_fp4 = torch.empty((B, H, N, D // 2), device=x.device, dtype=torch.uint8)
|
||||
fp8_scale = torch.empty((B, H, N, D // 16), device=x.device, dtype=torch.float8_e4m3fn)
|
||||
fp4quant_cuda.scaled_fp4_quant_permute(x, packed_fp4, fp8_scale, 1)
|
||||
return packed_fp4, fp8_scale
|
||||
|
||||
def scale_and_quant_fp4_transpose(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.ndim == 4
|
||||
B, H, N, D = x.shape
|
||||
packed_fp4 = torch.empty((B, H, D, N // 2), device=x.device, dtype=torch.uint8)
|
||||
fp8_scale = torch.empty((B, H, D, N // 16), device=x.device, dtype=torch.float8_e4m3fn)
|
||||
fp4quant_cuda.scaled_fp4_quant_trans(x, packed_fp4, fp8_scale, 1)
|
||||
return packed_fp4, fp8_scale
|
||||
|
||||
def blockscaled_fp4_attn(qlist: Tuple,
|
||||
klist: Tuple,
|
||||
vlist: Tuple,
|
||||
delta_s: torch.Tensor,
|
||||
KL: int,
|
||||
is_causal: bool = False,
|
||||
per_block_mean: bool = True,
|
||||
is_bf16: bool = True,
|
||||
single_level_p_quant: bool = False,
|
||||
sm_scale: float | None = None
|
||||
):
|
||||
softmax_scale = sm_scale if sm_scale is not None else (qlist[0].shape[-1] * 2) ** (-0.5)
|
||||
return fp4attn_cuda.fwd(qlist[0], klist[0], vlist[0], qlist[1], klist[1], vlist[1], delta_s, KL, None, softmax_scale, is_causal, per_block_mean, is_bf16, single_level_p_quant)
|
||||
|
||||
|
||||
def sageattn_blackwell(q, k, v, attn_mask = None, is_causal = False, per_block_mean = True, single_level_p_quant = True, sm_scale: float | None = None, **kwargs):
|
||||
"""
|
||||
SageAttention3 Blackwell kernel for FP4 attention.
|
||||
|
||||
Args:
|
||||
q: Query tensor [B, H, L, D]
|
||||
k: Key tensor [B, H, L, D]
|
||||
v: Value tensor [B, H, L, D]
|
||||
attn_mask: Attention mask (not used)
|
||||
is_causal: Whether to use causal masking
|
||||
per_block_mean: Whether to use per-block mean for Q smoothing
|
||||
single_level_p_quant: If True, use single-level quantization: s_P2, P̂_2 = φ(P̃) directly
|
||||
(standard per-block FP4 quantization like V, no s_P1).
|
||||
If False (default), use two-level quantization:
|
||||
s_P1 = rowmax(P̃)/(448×6), then s_P2, P̂_2 = φ(P̃/s_P1).
|
||||
sm_scale: Softmax scale to pass through to the CUDA kernel. If None,
|
||||
defaults to the kernel's 1/sqrt(D) scale.
|
||||
**kwargs: Additional arguments (ignored)
|
||||
|
||||
Returns:
|
||||
Output tensor [B, H, L, D]
|
||||
"""
|
||||
if q.size(-1) >= 256:
|
||||
print(f"Unsupported Headdim {q.size(-1)}")
|
||||
return sdpa(q, k, v, is_causal = is_causal)
|
||||
QL = q.size(2)
|
||||
KL = k.size(2)
|
||||
is_bf16 = q.dtype == torch.bfloat16
|
||||
q, k, v, delta_s = preprocess_qkv(q, k, v, per_block_mean)
|
||||
qlist_from_cuda = scale_and_quant_fp4(q)
|
||||
klist_from_cuda = scale_and_quant_fp4_permute(k)
|
||||
vlist_from_cuda = scale_and_quant_fp4_transpose(v)
|
||||
o_fp4 = blockscaled_fp4_attn(
|
||||
qlist_from_cuda,
|
||||
klist_from_cuda,
|
||||
vlist_from_cuda,
|
||||
delta_s,
|
||||
KL,
|
||||
is_causal,
|
||||
per_block_mean,
|
||||
is_bf16,
|
||||
single_level_p_quant,
|
||||
sm_scale
|
||||
)[0][:, :, :QL, :].contiguous()
|
||||
return o_fp4
|
||||
@@ -0,0 +1 @@
|
||||
__version__ = "3.0.0.b1"
|
||||
@@ -0,0 +1,347 @@
|
||||
// Modified from the original SageAttention3 code
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
// Include these 2 headers instead of torch/extension.h since we don't need all of the torch headers.
|
||||
#include <torch/python.h>
|
||||
#include <torch/nn/functional.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <cutlass/numeric_types.h>
|
||||
|
||||
#include "params.h"
|
||||
#include "launch.h"
|
||||
#include "static_switch.h"
|
||||
#include "block_config.h"
|
||||
|
||||
#define CHECK_DEVICE(x) TORCH_CHECK(x.is_cuda(), #x " must be on CUDA")
|
||||
#define CHECK_SHAPE(x, ...) TORCH_CHECK(x.sizes() == torch::IntArrayRef({__VA_ARGS__}), #x " must have shape (" #__VA_ARGS__ ")")
|
||||
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
|
||||
|
||||
|
||||
void set_params_fprop(Flash_fwd_params ¶ms,
|
||||
// sizes
|
||||
const size_t b,
|
||||
const size_t seqlen_q,
|
||||
const size_t seqlen_k,
|
||||
const size_t unpadded_seqlen_k,
|
||||
const size_t seqlen_q_rounded,
|
||||
const size_t seqlen_k_rounded,
|
||||
const size_t h,
|
||||
const size_t h_k,
|
||||
const size_t d,
|
||||
const size_t d_rounded,
|
||||
// device pointers
|
||||
const at::Tensor q,
|
||||
const at::Tensor k,
|
||||
const at::Tensor v,
|
||||
const at::Tensor delta_s,
|
||||
at::Tensor out,
|
||||
const at::Tensor sfq,
|
||||
const at::Tensor sfk,
|
||||
const at::Tensor sfv,
|
||||
void *cu_seqlens_q_d,
|
||||
void *cu_seqlens_k_d,
|
||||
void *seqused_k,
|
||||
void *p_d,
|
||||
void *softmax_lse_d,
|
||||
float p_dropout,
|
||||
float softmax_scale,
|
||||
int window_size_left,
|
||||
int window_size_right,
|
||||
bool per_block_mean,
|
||||
bool is_bf16,
|
||||
bool single_level_p_quant=false,
|
||||
bool seqlenq_ngroups_swapped=false) {
|
||||
|
||||
// Reset the parameters
|
||||
params = {};
|
||||
// Set the pointers and strides.
|
||||
params.q_ptr = q.data_ptr();
|
||||
params.k_ptr = k.data_ptr();
|
||||
params.v_ptr = v.data_ptr();
|
||||
params.delta_s_ptr = delta_s.data_ptr();
|
||||
params.sfq_ptr = sfq.data_ptr();
|
||||
params.sfk_ptr = sfk.data_ptr();
|
||||
params.sfv_ptr = sfv.data_ptr();
|
||||
|
||||
// All stride are in elements, not bytes.
|
||||
params.q_row_stride = q.stride(-2) * 2;
|
||||
params.k_row_stride = k.stride(-2) * 2;
|
||||
params.v_row_stride = v.stride(-2) * 2;;
|
||||
params.q_head_stride = q.stride(-3) * 2;
|
||||
params.k_head_stride = k.stride(-3) * 2;
|
||||
params.v_head_stride = v.stride(-3) * 2; // for packed q k v
|
||||
|
||||
params.ds_row_stride = delta_s.stride(-2);
|
||||
params.ds_head_stride = delta_s.stride(-3);
|
||||
|
||||
params.sfq_row_stride = sfq.stride(-2);
|
||||
params.sfk_row_stride = sfk.stride(-2);
|
||||
params.sfv_row_stride = sfv.stride(-2);
|
||||
params.sfq_head_stride = sfq.stride(-3);
|
||||
params.sfk_head_stride = sfk.stride(-3);
|
||||
params.sfv_head_stride = sfv.stride(-3);
|
||||
params.o_ptr = out.data_ptr();
|
||||
params.o_row_stride = out.stride(-2);
|
||||
params.o_head_stride = out.stride(-3);
|
||||
|
||||
if (cu_seqlens_q_d == nullptr) {
|
||||
params.q_batch_stride = q.stride(0) * 2;
|
||||
params.k_batch_stride = k.stride(0) * 2;
|
||||
params.v_batch_stride = v.stride(0) * 2;
|
||||
params.ds_batch_stride = delta_s.stride(0);
|
||||
params.sfq_batch_stride = sfq.stride(0);
|
||||
params.sfk_batch_stride = sfk.stride(0);
|
||||
params.sfv_batch_stride = sfv.stride(0);
|
||||
params.o_batch_stride = out.stride(0);
|
||||
if (seqlenq_ngroups_swapped) {
|
||||
params.q_batch_stride *= seqlen_q;
|
||||
params.o_batch_stride *= seqlen_q;
|
||||
}
|
||||
}
|
||||
|
||||
params.cu_seqlens_q = static_cast<int *>(cu_seqlens_q_d);
|
||||
params.cu_seqlens_k = static_cast<int *>(cu_seqlens_k_d);
|
||||
params.seqused_k = static_cast<int *>(seqused_k);
|
||||
|
||||
// P = softmax(QK^T)
|
||||
params.p_ptr = p_d;
|
||||
|
||||
// Softmax sum
|
||||
params.softmax_lse_ptr = softmax_lse_d;
|
||||
|
||||
// Set the dimensions.
|
||||
params.b = b;
|
||||
params.h = h;
|
||||
params.h_k = h_k;
|
||||
params.h_h_k_ratio = h / h_k;
|
||||
params.seqlen_q = seqlen_q;
|
||||
params.seqlen_k = seqlen_k;
|
||||
params.unpadded_seqlen_k = unpadded_seqlen_k;
|
||||
params.seqlen_q_rounded = seqlen_q_rounded;
|
||||
params.seqlen_k_rounded = seqlen_k_rounded;
|
||||
params.d = d;
|
||||
params.d_rounded = d_rounded;
|
||||
|
||||
params.head_divmod = cutlass::FastDivmod(int(h));
|
||||
|
||||
// Set the different scale values.
|
||||
params.scale_softmax = softmax_scale;
|
||||
params.scale_softmax_log2 = softmax_scale * M_LOG2E;
|
||||
__half scale_softmax_log2_half = __float2half(params.scale_softmax_log2);
|
||||
__half2 scale_softmax_log2_half2 = __half2(scale_softmax_log2_half, scale_softmax_log2_half);
|
||||
params.scale_softmax_log2_half2 = reinterpret_cast<uint32_t&>(scale_softmax_log2_half2);
|
||||
|
||||
// Set this to probability of keeping an element to simplify things.
|
||||
params.p_dropout = 1.f - p_dropout;
|
||||
// Convert p from float to int so we don't have to convert the random uint to float to compare.
|
||||
// [Minor] We want to round down since when we do the comparison we use <= instead of <
|
||||
// params.p_dropout_in_uint = uint32_t(std::floor(params.p_dropout * 4294967295.0));
|
||||
// params.p_dropout_in_uint16_t = uint16_t(std::floor(params.p_dropout * 65535.0));
|
||||
params.p_dropout_in_uint8_t = uint8_t(std::floor(params.p_dropout * 255.0));
|
||||
params.rp_dropout = 1.f / params.p_dropout;
|
||||
params.scale_softmax_rp_dropout = params.rp_dropout * params.scale_softmax;
|
||||
TORCH_CHECK(p_dropout < 1.f);
|
||||
#ifdef FLASHATTENTION_DISABLE_DROPOUT
|
||||
TORCH_CHECK(p_dropout == 0.0f, "This flash attention build does not support dropout.");
|
||||
#endif
|
||||
|
||||
// Causal is the special case where window_size_right == 0 and window_size_left < 0.
|
||||
// Local is the more general case where window_size_right >= 0 or window_size_left >= 0.
|
||||
params.is_causal = window_size_left < 0 && window_size_right == 0;
|
||||
params.per_block_mean = per_block_mean;
|
||||
if (per_block_mean) {
|
||||
params.seqlen_s = seqlen_q;
|
||||
} else {
|
||||
params.seqlen_s = flash::BLOCK_M; // size of BLOCK_M
|
||||
}
|
||||
if (window_size_left < 0 && window_size_right >= 0) { window_size_left = seqlen_k; }
|
||||
if (window_size_left >= 0 && window_size_right < 0) { window_size_right = seqlen_k; }
|
||||
params.window_size_left = window_size_left;
|
||||
params.window_size_right = window_size_right;
|
||||
|
||||
#ifdef FLASHATTENTION_DISABLE_LOCAL
|
||||
TORCH_CHECK(params.is_causal || (window_size_left < 0 && window_size_right < 0),
|
||||
"This flash attention build does not support local attention.");
|
||||
#endif
|
||||
|
||||
params.is_seqlens_k_cumulative = true;
|
||||
params.is_bf16 = is_bf16;
|
||||
params.single_level_p_quant = single_level_p_quant;
|
||||
#ifdef FLASHATTENTION_DISABLE_UNEVEN_K
|
||||
TORCH_CHECK(d == d_rounded, "This flash attention build does not support headdim not being a multiple of 32.");
|
||||
#endif
|
||||
}
|
||||
|
||||
template<bool IsBF16>
|
||||
void run_mha_fwd_dispatch_dtype(Flash_fwd_params ¶ms, cudaStream_t stream) {
|
||||
using OType = std::conditional_t<IsBF16, cutlass::bfloat16_t, cutlass::half_t>;
|
||||
if (params.d == 64) {
|
||||
run_mha_fwd_<cutlass::nv_float4_t<cutlass::float_e2m1_t>, 64, OType>(params, stream);
|
||||
} else if (params.d == 128) {
|
||||
run_mha_fwd_<cutlass::nv_float4_t<cutlass::float_e2m1_t>, 128, OType>(params, stream);
|
||||
}
|
||||
}
|
||||
|
||||
void run_mha_fwd(Flash_fwd_params ¶ms, cudaStream_t stream, bool force_split_kernel = false) {
|
||||
BOOL_SWITCH(params.is_bf16, IsBF16, ([&] {
|
||||
run_mha_fwd_dispatch_dtype<IsBF16>(params, stream);
|
||||
}));
|
||||
}
|
||||
|
||||
std::vector<at::Tensor>
|
||||
mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size // 2)
|
||||
const at::Tensor &k, // batch_size x seqlen_k x num_heads_k x (head_size // 2)
|
||||
const at::Tensor &v, // batch_size x seqlen_k x num_heads_k x (head_size // 2)
|
||||
const at::Tensor &sfq,
|
||||
const at::Tensor &sfk,
|
||||
const at::Tensor &sfv,
|
||||
const at::Tensor &delta_s,
|
||||
int unpadded_k,
|
||||
c10::optional<at::Tensor> &out_, // batch_size x seqlen_q x num_heads x head_size
|
||||
const float softmax_scale,
|
||||
bool is_causal,
|
||||
bool per_block_mean,
|
||||
bool is_bf16,
|
||||
bool single_level_p_quant=false // If true, use only per-row scale s_P2 (no per-block s_P1)
|
||||
) {
|
||||
|
||||
auto dprops = at::cuda::getCurrentDeviceProperties();
|
||||
bool is_blackwell_or_newer = dprops->major >= 12;
|
||||
TORCH_CHECK(is_blackwell_or_newer, "only supports Blackwell GPUs or newer.");
|
||||
|
||||
auto q_dtype = q.dtype();
|
||||
auto sfq_dtype = sfq.dtype();
|
||||
TORCH_CHECK(q_dtype == torch::kUInt8, "q dtype must be uint8");
|
||||
TORCH_CHECK(k.dtype() == q_dtype, "query and key must have the same dtype");
|
||||
TORCH_CHECK(v.dtype() == q_dtype, "query and value must have the same dtype");
|
||||
CHECK_DEVICE(q); CHECK_DEVICE(k); CHECK_DEVICE(v);
|
||||
|
||||
TORCH_CHECK(sfq_dtype == torch::kFloat8_e4m3fn, "q dtype must be uint8");
|
||||
TORCH_CHECK(sfk.dtype() == sfq_dtype, "query and key must have the same dtype");
|
||||
TORCH_CHECK(sfv.dtype() == sfq_dtype, "query and value must have the same dtype");
|
||||
CHECK_DEVICE(sfq); CHECK_DEVICE(sfk); CHECK_DEVICE(sfv);
|
||||
|
||||
TORCH_CHECK(q.stride(-1) == 1, "Input tensor must have contiguous last dimension");
|
||||
TORCH_CHECK(k.stride(-1) == 1, "Input tensor must have contiguous last dimension");
|
||||
TORCH_CHECK(v.stride(-1) == 1, "Input tensor must have contiguous last dimension");
|
||||
TORCH_CHECK(delta_s.stride(-1) == 1, "Input tensor must have contiguous last dimension");
|
||||
|
||||
TORCH_CHECK(q.is_contiguous(), "Input tensor must be contiguous");
|
||||
TORCH_CHECK(k.is_contiguous(), "Input tensor must be contiguous");
|
||||
TORCH_CHECK(v.is_contiguous(), "Input tensor must be contiguous");
|
||||
|
||||
const auto sizes = q.sizes();
|
||||
auto opts = q.options();
|
||||
const int batch_size = sizes[0];
|
||||
int seqlen_q = sizes[2];
|
||||
int num_heads = sizes[1];
|
||||
const int head_size_og = sizes[3];
|
||||
const int unpacked_head_size = head_size_og * 2;
|
||||
const int seqlen_k = k.size(2);
|
||||
const int num_heads_k = k.size(1);
|
||||
|
||||
TORCH_CHECK(batch_size > 0, "batch size must be postive");
|
||||
TORCH_CHECK(unpacked_head_size <= 256, "FlashAttention forward only supports head dimension at most 256");
|
||||
TORCH_CHECK(num_heads % num_heads_k == 0, "Number of heads in key/value must divide number of heads in query");
|
||||
TORCH_CHECK(num_heads == num_heads_k, "We do not support MQA/GQA yet");
|
||||
|
||||
TORCH_CHECK(unpacked_head_size == 64 || unpacked_head_size == 128 || unpacked_head_size == 256, "Only support head size 64, 128, and 256 for now");
|
||||
|
||||
CHECK_SHAPE(q, batch_size, num_heads, seqlen_q, head_size_og);
|
||||
CHECK_SHAPE(k, batch_size, num_heads_k, seqlen_k, head_size_og);
|
||||
CHECK_SHAPE(v, batch_size, num_heads_k, unpacked_head_size, seqlen_k/2);
|
||||
// CHECK_SHAPE(delta_s, batch_size, num_heads, seqlen_q / 128, seqlen_k);
|
||||
// CHECK_SHAPE(sfq, batch_size, seqlen_q, num_heads, unpacked_head_size);
|
||||
// CHECK_SHAPE(sfk, batch_size, seqlen_k, num_heads_k, unpacked_head_size);
|
||||
// CHECK_SHAPE(sfv, batch_size, unpacked_head_size, num_heads_k, seqlen_k);
|
||||
TORCH_CHECK(unpacked_head_size % 8 == 0, "head_size must be a multiple of 8");
|
||||
|
||||
auto dtype = is_bf16 ? at::ScalarType::BFloat16 : at::ScalarType::Half;
|
||||
at::Tensor out = torch::empty({batch_size, num_heads, seqlen_q, unpacked_head_size}, opts.dtype(dtype));
|
||||
|
||||
auto round_multiple = [](int x, int m) { return (x + m - 1) / m * m; };
|
||||
// const int head_size = round_multiple(head_size_og, 8);
|
||||
// const int head_size_rounded = round_multiple(head_size, 32);
|
||||
const int seqlen_q_rounded = round_multiple(seqlen_q, flash::BLOCK_M);
|
||||
const int seqlen_k_rounded = round_multiple(seqlen_k, flash::BLOCK_N);
|
||||
|
||||
// Otherwise the kernel will be launched from cuda:0 device
|
||||
// Cast to char to avoid compiler warning about narrowing
|
||||
at::cuda::CUDAGuard device_guard{(char)q.get_device()};
|
||||
|
||||
|
||||
|
||||
auto softmax_lse = torch::empty({batch_size, num_heads, seqlen_q}, opts.dtype(at::kFloat));
|
||||
at::Tensor p;
|
||||
|
||||
Flash_fwd_params params;
|
||||
set_params_fprop(params,
|
||||
batch_size,
|
||||
seqlen_q, seqlen_k, unpadded_k,
|
||||
seqlen_q_rounded, seqlen_k_rounded,
|
||||
num_heads, num_heads_k,
|
||||
unpacked_head_size, unpacked_head_size,
|
||||
q, k, v, delta_s, out,
|
||||
sfq, sfk, sfv,
|
||||
/*cu_seqlens_q_d=*/nullptr,
|
||||
/*cu_seqlens_k_d=*/nullptr,
|
||||
/*seqused_k=*/nullptr,
|
||||
nullptr,
|
||||
softmax_lse.data_ptr(),
|
||||
/*p_dropout=*/0.f,
|
||||
softmax_scale,
|
||||
/*window_size_left=*/-1,
|
||||
/*window_size_right=*/is_causal ? 0 : -1,
|
||||
per_block_mean,
|
||||
is_bf16,
|
||||
single_level_p_quant
|
||||
);
|
||||
// StaticPersistentTileScheduler does not use tile_count_semaphore; avoid a
|
||||
// stack-local tensor whose data pointer would dangle after mha_fwd returns
|
||||
// while the async kernel may still be running.
|
||||
params.tile_count_semaphore = nullptr;
|
||||
|
||||
if (seqlen_k > 0) {
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
run_mha_fwd(params, stream);
|
||||
} else {
|
||||
// If seqlen_k == 0, then we have an empty tensor. We need to set the output to 0.
|
||||
out.zero_();
|
||||
softmax_lse.fill_(std::numeric_limits<float>::infinity());
|
||||
}
|
||||
|
||||
// at::Tensor out_padded = out;
|
||||
// if (head_size_og % 8 != 0) {
|
||||
// out = out.index({"...", torch::indexing::Slice(torch::indexing::None, head_size_og)});
|
||||
// if (out_.has_value()) { out_.value().copy_(out); }
|
||||
// }
|
||||
|
||||
// return {out, q_padded, k_padded, v_padded, out_padded, softmax_lse, p};
|
||||
// cudaDeviceSynchronize();
|
||||
// auto err = cudaGetLastError();
|
||||
// printf("%s\n", cudaGetErrorString(err));
|
||||
return {out, softmax_lse};
|
||||
}
|
||||
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "FlashAttention";
|
||||
m.def("fwd", &mha_fwd, "Forward pass");
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
// Centralized block size configuration for sageattn_blackwell kernels
|
||||
// Block sizes for M and N dimensions
|
||||
namespace flash {
|
||||
// Block size for M dimension (query sequence length)
|
||||
static constexpr int BLOCK_M = 128;
|
||||
|
||||
// Block size for N dimension (key/value sequence length)
|
||||
static constexpr int BLOCK_N = 128;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* This code is based on code from FlashAttention3, https://github.com/Dao-AILab/flash-attention
|
||||
* Copyright (c) 2024, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
namespace flash {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<bool Varlen=true>
|
||||
struct BlockInfo {
|
||||
|
||||
template<typename Params>
|
||||
__device__ BlockInfo(const Params ¶ms, const int bidb)
|
||||
: sum_s_q(!Varlen || params.cu_seqlens_q == nullptr ? -1 : params.cu_seqlens_q[bidb])
|
||||
, sum_s_k(!Varlen || params.cu_seqlens_k == nullptr || !params.is_seqlens_k_cumulative ? -1 : params.cu_seqlens_k[bidb])
|
||||
, actual_seqlen_q(!Varlen || params.cu_seqlens_q == nullptr ? params.seqlen_q : params.cu_seqlens_q[bidb + 1] - sum_s_q)
|
||||
// If is_seqlens_k_cumulative, then seqlen_k is cu_seqlens_k[bidb + 1] - cu_seqlens_k[bidb].
|
||||
// Otherwise it's cu_seqlens_k[bidb], i.e., we use cu_seqlens_k to store the sequence lengths of K.
|
||||
, seqlen_k_cache(!Varlen || params.cu_seqlens_k == nullptr ? params.seqlen_k : (params.is_seqlens_k_cumulative ? params.cu_seqlens_k[bidb + 1] - sum_s_k : params.cu_seqlens_k[bidb]))
|
||||
, actual_seqlen_k(params.seqused_k ? params.seqused_k[bidb] : seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : params.seqlen_knew))
|
||||
{
|
||||
}
|
||||
|
||||
template <typename index_t>
|
||||
__forceinline__ __device__ index_t q_offset(const index_t batch_stride, const index_t row_stride, const int bidb) const {
|
||||
return sum_s_q == -1 ? bidb * batch_stride : uint32_t(sum_s_q) * row_stride;
|
||||
}
|
||||
|
||||
template <typename index_t>
|
||||
__forceinline__ __device__ index_t k_offset(const index_t batch_stride, const index_t row_stride, const int bidb) const {
|
||||
return sum_s_k == -1 ? bidb * batch_stride : uint32_t(sum_s_k) * row_stride;
|
||||
}
|
||||
|
||||
const int sum_s_q;
|
||||
const int sum_s_k;
|
||||
const int actual_seqlen_q;
|
||||
// We have to have seqlen_k_cache declared before actual_seqlen_k, otherwise actual_seqlen_k is set to 0.
|
||||
const int seqlen_k_cache;
|
||||
const int actual_seqlen_k;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,149 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief Blocked Scale configs specific for SM100 BlockScaled MMA
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
|
||||
#include "cute/int_tuple.hpp"
|
||||
#include "cute/atom/mma_traits_sm100.hpp"
|
||||
|
||||
namespace flash {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
using namespace cute;
|
||||
|
||||
template<int SFVecSize, UMMA::Major major = UMMA::Major::K>
|
||||
struct BlockScaledBasicChunk {
|
||||
|
||||
using Blk_MN = _64;
|
||||
using Blk_SF = _4;
|
||||
|
||||
using SfAtom = Layout< Shape< Shape<_16,_4>, Shape<Int<SFVecSize>, _4>>,
|
||||
Stride<Stride<_16,_4>, Stride< _0, _1>>>;
|
||||
};
|
||||
|
||||
template<int SFVecSize_>
|
||||
struct BlockScaledConfig {
|
||||
// We are creating the SFA and SFB tensors' layouts in the collective since they always have the same layout.
|
||||
// k-major order
|
||||
static constexpr int SFVecSize = SFVecSize_;
|
||||
static constexpr int MMA_NSF = 4; // SFVecSize, MMA_NSF
|
||||
using BlkScaledChunk = BlockScaledBasicChunk<SFVecSize>;
|
||||
using Blk_MN = _64;
|
||||
using Blk_SF = _4;
|
||||
using mnBasicBlockShape = Shape<_16,_4>;
|
||||
using mnBasicBlockStride = Stride<_16,_4>;
|
||||
using kBasicBlockShape = Shape<Int<SFVecSize>, Int<MMA_NSF>>; // SFVecSize, MMA_NSF
|
||||
using kBasicBlockStride = Stride<_0, _1>;
|
||||
using SfAtom = Layout< Shape< mnBasicBlockShape, kBasicBlockShape>,
|
||||
Stride<mnBasicBlockStride, kBasicBlockStride>>;
|
||||
|
||||
using LayoutSF = decltype(blocked_product(SfAtom{},
|
||||
make_layout(
|
||||
make_shape(int32_t(0), int32_t(0), int32_t(0), int32_t(0)),
|
||||
make_stride(int32_t(0), _1{}, int32_t(0), int32_t(0)))));
|
||||
// A single indivisible block will hold 4 scale factors of 64 rows/columns (A/B matrix).
|
||||
// 4 is chosen to make consecutive 32bits of data to have scale factors for only a single row (col). 32bits corresponds to the TMEM word size
|
||||
using Blk_Elems = decltype(Blk_MN{} * Blk_SF{});
|
||||
using sSF_strideMN = decltype(prepend(Blk_Elems{}, mnBasicBlockStride{}));
|
||||
|
||||
|
||||
// The following function is provided for user fill dynamic problem size to the layout_SFA.
|
||||
template < class ProblemShape>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
tile_atom_to_shape_SFQKV(ProblemShape problem_shape) {
|
||||
auto [Seqlen, Dim, HeadNum, Batch] = problem_shape;
|
||||
return tile_to_shape(SfAtom{}, make_shape(Seqlen, Dim, HeadNum, Batch), Step<_2,_1,_3,_4>{});
|
||||
}
|
||||
|
||||
// The following function is provided for user fill dynamic problem size to the layout_SFB.
|
||||
template <class ProblemShape>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
tile_atom_to_shape_SFVt(ProblemShape problem_shape) {
|
||||
auto [Dim, Seqlen, HeadNum, Batch] = problem_shape;
|
||||
return tile_to_shape(SfAtom{}, make_shape(Dim, Seqlen, HeadNum, Batch), Step<_2,_1,_3,_4>{});
|
||||
}
|
||||
|
||||
template<class TiledMma, class TileShape_MNK>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
deduce_smem_layoutSFQ(TiledMma tiled_mma, TileShape_MNK tileshape_mnk) {
|
||||
|
||||
using sSFQ_shapeK = decltype(prepend(make_shape(Blk_SF{}/Int<MMA_NSF>{}, size<2>(TileShape_MNK{}) / Int<SFVecSize>{} / Blk_SF{}), kBasicBlockShape{}));
|
||||
using sSFQ_shapeM = decltype(prepend(size<0>(TileShape_MNK{}) / Blk_MN{}, mnBasicBlockShape{}));
|
||||
using sSFQ_strideM = sSF_strideMN;
|
||||
using sSFQ_strideK = decltype(prepend(make_stride(Int<MMA_NSF>{}, size<0>(TileShape_MNK{}) / Blk_MN{} * Blk_Elems{}), kBasicBlockStride{}));
|
||||
using sSFQ_shape = decltype(make_shape(sSFQ_shapeM{}, sSFQ_shapeK{}));
|
||||
using sSFQ_stride = decltype(make_stride(sSFQ_strideM{}, sSFQ_strideK{}));
|
||||
using SmemLayoutAtomSFQ = decltype(make_layout(sSFQ_shape{}, sSFQ_stride{}));
|
||||
return SmemLayoutAtomSFQ{};
|
||||
}
|
||||
|
||||
template<class TiledMma, class TileShape_MNK>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
deduce_smem_layoutSFKV(TiledMma tiled_mma, TileShape_MNK tileshape_mnk) {
|
||||
|
||||
using sSFK_shapeK = decltype(prepend(make_shape(Blk_SF{}/Int<MMA_NSF>{}, size<2>(TileShape_MNK{}) / Int<SFVecSize>{} / Blk_SF{}), kBasicBlockShape{}));
|
||||
using sSFK_shapeN = decltype(prepend(size<1>(TileShape_MNK{}) / Blk_MN{}, mnBasicBlockShape{}));
|
||||
using sSFK_strideN = sSF_strideMN;
|
||||
using sSFK_strideK = decltype(prepend(make_stride(Int<MMA_NSF>{}, size<1>(TileShape_MNK{}) / Blk_MN{} * Blk_Elems{}), kBasicBlockStride{}));
|
||||
using sSFK_shape = decltype(make_shape(sSFK_shapeN{}, sSFK_shapeK{}));
|
||||
using sSFK_stride = decltype(make_stride(sSFK_strideN{}, sSFK_strideK{}));
|
||||
using SmemLayoutAtomSFK = decltype(make_layout(sSFK_shape{}, sSFK_stride{}));
|
||||
return SmemLayoutAtomSFK{};
|
||||
}
|
||||
|
||||
template<class TiledMma, class TileShape_MNK>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
deduce_smem_layoutSFVt(TiledMma tiled_mma, TileShape_MNK tileshape_mnk) {
|
||||
|
||||
using sSFVt_shapeK = decltype(prepend(make_shape(Blk_SF{}/Int<MMA_NSF>{}, size<2>(TileShape_MNK{}) / Int<SFVecSize>{} / Blk_SF{}), kBasicBlockShape{}));
|
||||
using sSFVt_shapeN = decltype(prepend(size<1>(TileShape_MNK{}) / Blk_MN{}, mnBasicBlockShape{}));
|
||||
using sSFVt_strideN = sSF_strideMN;
|
||||
using sSFVt_strideK = decltype(prepend(make_stride(Int<MMA_NSF>{}, size<1>(TileShape_MNK{}) / Blk_MN{} * Blk_Elems{}), kBasicBlockStride{}));
|
||||
using sSFVt_shape = decltype(make_shape(sSFVt_shapeN{}, sSFVt_shapeK{}));
|
||||
using sSFVt_stride = decltype(make_stride(sSFVt_strideN{}, sSFVt_strideK{}));
|
||||
using SmemLayoutAtomSFVt = decltype(make_layout(sSFVt_shape{}, sSFVt_stride{}));
|
||||
return SmemLayoutAtomSFVt{};
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,327 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#include "cute/arch/mma_sm120.hpp"
|
||||
#include "cute/atom/mma_traits_sm120.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/float8.h"
|
||||
#include "cutlass/float_subbyte.h"
|
||||
|
||||
namespace cute::SM120::BLOCKSCALED {
|
||||
|
||||
using cutlass::float_e2m1_t;
|
||||
using cutlass::float_ue4m3_t;
|
||||
|
||||
// MMA.SF 16x32x64 TN E2M1 x E2M1 with SF E4M3
|
||||
struct SM120_16x32x64_TN_VS_NVFP4 {
|
||||
using DRegisters = float[16];
|
||||
using ARegisters = uint32_t[4];
|
||||
using BRegisters = uint32_t[8];
|
||||
using CRegisters = float[16];
|
||||
|
||||
static constexpr int SFBits = 32;
|
||||
using RegTypeSF = cute::uint_bit_t<SFBits>;
|
||||
|
||||
using SFARegisters = RegTypeSF[1];
|
||||
using SFBRegisters = RegTypeSF[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(float & d0 , float & d1 , float & d2 , float & d3 ,
|
||||
float & d4 , float & d5 , float & d6 , float & d7 ,
|
||||
float & d8 , float & d9 , float & d10, float & d11,
|
||||
float & d12, float & d13, float & d14, float & d15,
|
||||
uint32_t const& a0 , uint32_t const& a1 , uint32_t const& a2 , uint32_t const& a3 ,
|
||||
uint32_t const& b0 , uint32_t const& b1 , uint32_t const& b2 , uint32_t const& b3 ,
|
||||
uint32_t const& b4 , uint32_t const& b5 , uint32_t const& b6 , uint32_t const& b7 ,
|
||||
float const & c0 , float const & c1 , float const & c2 , float const & c3 ,
|
||||
float const & c4 , float const & c5 , float const & c6 , float const & c7 ,
|
||||
float const & c8 , float const & c9 , float const & c10 , float const & c11,
|
||||
float const & c12, float const & c13, float const & c14, float const & c15,
|
||||
RegTypeSF const& sfa0,
|
||||
RegTypeSF const& sfb0)
|
||||
{
|
||||
static constexpr uint16_t tidA = 0;
|
||||
static constexpr uint16_t bidA = 0;
|
||||
static constexpr uint16_t bidB = 0;
|
||||
static constexpr uint16_t tidB0 = 0;
|
||||
static constexpr uint16_t tidB1 = 1;
|
||||
static constexpr uint16_t tidB2 = 2;
|
||||
static constexpr uint16_t tidB3 = 3;
|
||||
|
||||
#if defined(CUTE_ARCH_MXF4NVF4_4X_UE4M3_MMA_ENABLED)
|
||||
asm volatile(
|
||||
"mma.sync.aligned.kind::mxf4nvf4.block_scale.scale_vec::4X.m16n8k64.row.col.f32.e2m1.e2m1.f32.ue4m3 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13},"
|
||||
"{%14},"
|
||||
"{%15, %16},"
|
||||
"{%17},"
|
||||
"{%18, %19};\n"
|
||||
: "=f"(d0), "=f"(d1), "=f"(d8), "=f"(d9)
|
||||
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b0), "r"(b1),
|
||||
"f"(c0), "f"(c1), "f"(c8), "f"(c9),
|
||||
"r"(uint32_t(sfa0)) , "h"(bidA), "h"(tidA),
|
||||
"r"(uint32_t(sfb0)) , "h"(bidB), "h"(tidB0));
|
||||
|
||||
asm volatile(
|
||||
"mma.sync.aligned.kind::mxf4nvf4.block_scale.scale_vec::4X.m16n8k64.row.col.f32.e2m1.e2m1.f32.ue4m3 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13},"
|
||||
"{%14},"
|
||||
"{%15, %16},"
|
||||
"{%17},"
|
||||
"{%18, %19};\n"
|
||||
: "=f"(d2), "=f"(d3), "=f"(d10), "=f"(d11)
|
||||
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b2), "r"(b3),
|
||||
"f"(c2), "f"(c3), "f"(c10), "f"(c11),
|
||||
"r"(uint32_t(sfa0)) , "h"(bidA), "h"(tidA),
|
||||
"r"(uint32_t(sfb0)) , "h"(bidB), "h"(tidB1));
|
||||
|
||||
asm volatile(
|
||||
"mma.sync.aligned.kind::mxf4nvf4.block_scale.scale_vec::4X.m16n8k64.row.col.f32.e2m1.e2m1.f32.ue4m3 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13},"
|
||||
"{%14},"
|
||||
"{%15, %16},"
|
||||
"{%17},"
|
||||
"{%18, %19};\n"
|
||||
: "=f"(d4), "=f"(d5), "=f"(d12), "=f"(d13)
|
||||
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b4), "r"(b5),
|
||||
"f"(c4), "f"(c5), "f"(c12), "f"(c13),
|
||||
"r"(uint32_t(sfa0)) , "h"(bidA), "h"(tidA),
|
||||
"r"(uint32_t(sfb0)) , "h"(bidB), "h"(tidB2));
|
||||
|
||||
asm volatile(
|
||||
"mma.sync.aligned.kind::mxf4nvf4.block_scale.scale_vec::4X.m16n8k64.row.col.f32.e2m1.e2m1.f32.ue4m3 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13},"
|
||||
"{%14},"
|
||||
"{%15, %16},"
|
||||
"{%17},"
|
||||
"{%18, %19};\n"
|
||||
: "=f"(d6), "=f"(d7), "=f"(d14), "=f"(d15)
|
||||
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b6), "r"(b7),
|
||||
"f"(c6), "f"(c7), "f"(c14), "f"(c15),
|
||||
"r"(uint32_t(sfa0)) , "h"(bidA), "h"(tidA),
|
||||
"r"(uint32_t(sfb0)) , "h"(bidB), "h"(tidB3));
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM120::BLOCKSCALED::SM120_16x8x64_TN_VS without CUTE_ARCH_MXF4NVF4_4X_UE4M3_MMA_ENABLED");
|
||||
#endif
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cute::SM120::BLOCKSCALED
|
||||
|
||||
namespace cute {
|
||||
|
||||
// MMA NVFP4 16x32x64 TN
|
||||
template <>
|
||||
struct MMA_Traits<SM120::BLOCKSCALED::SM120_16x32x64_TN_VS_NVFP4>
|
||||
{
|
||||
// The MMA accepts 4-bit inputs regardless of the types for A and B
|
||||
using ValTypeA = uint4_t;
|
||||
using ValTypeB = uint4_t;
|
||||
|
||||
using ValTypeD = float;
|
||||
using ValTypeC = float;
|
||||
|
||||
using ValTypeSF = cutlass::float_ue4m3_t;
|
||||
constexpr static int SFVecSize = 16;
|
||||
|
||||
using Shape_MNK = Shape<_16,_32,_64>;
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// (T32,V32) -> (M16,K64)
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _8,_2, _2>>,
|
||||
Stride<Stride<_128,_1>,Stride<_16,_8,_512>>>;
|
||||
// (T32,V64) -> (N32,K64)
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_8, _2, _4>>,
|
||||
Stride<Stride<_256,_1>,Stride<_32,_1024, _8>>>;
|
||||
// (T32,V64) -> (M16,K64)
|
||||
using SFALayout = Layout<Shape <Shape <_2,_2,_8>,_64>,
|
||||
Stride<Stride<_8,_0,_1>,_16>>;
|
||||
// (T32,V64) -> (N32,K64)
|
||||
using SFBLayout = Layout<Shape <Shape <_4,_8>,_64>,
|
||||
Stride<Stride<_8,_1>, _32>>;
|
||||
// (T32,V16) -> (M16,N32)
|
||||
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < Shape<_2, _4>,_2>>,
|
||||
Stride<Stride<_32,_1>,Stride<Stride<_16, _128>,_8>>>;
|
||||
};
|
||||
|
||||
|
||||
template <class SFATensor, class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_SFA(SFATensor&& sfatensor, TiledMMA<Atom, TiledThr, TiledPerm>& mma)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(sfatensor) >= Int<2>{});
|
||||
|
||||
using AtomShape_MNK = typename Atom::Shape_MNK;
|
||||
using AtomLayoutSFA_TV = typename Atom::Traits::SFALayout;
|
||||
|
||||
auto permutation_mnk = TiledPerm{};
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<0>(permutation_mnk),
|
||||
get<2>(permutation_mnk));
|
||||
auto t_tensor = logical_divide(sfatensor, t_tile); // (PermM,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
auto a_tile = make_tile(make_layout(size<0>(AtomShape_MNK{})),
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto a_tensor = zipped_divide(t_tensor, a_tile); // ((AtomM,AtomK),(RestM,RestK))
|
||||
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = a_tensor.compose(AtomLayoutSFA_TV{},_); // ((ThrV,FrgV),(RestM,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<1>(thr_layout_vmnk)),
|
||||
make_layout(size<3>(thr_layout_vmnk))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrM,ThrK)),(FrgV,(RestM,RestK)))
|
||||
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
template <class SFBTensor, class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_SFB(SFBTensor&& sfbtensor, TiledMMA<Atom, TiledThr, TiledPerm>& mma)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(sfbtensor) >= Int<2>{});
|
||||
|
||||
using AtomShape_MNK = typename Atom::Shape_MNK;
|
||||
using AtomLayoutSFB_TV = typename Atom::Traits::SFBLayout;
|
||||
|
||||
auto permutation_mnk = TiledPerm{};
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<1>(permutation_mnk),
|
||||
get<2>(permutation_mnk));
|
||||
auto t_tensor = logical_divide(sfbtensor, t_tile); // (PermN,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
auto a_tile = make_tile(make_layout(size<1>(AtomShape_MNK{})),
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto a_tensor = zipped_divide(t_tensor, a_tile); // ((AtomN,AtomK),(RestN,RestK))
|
||||
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = a_tensor.compose(AtomLayoutSFB_TV{},_); // ((ThrV,FrgV),(RestN,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<2>(thr_layout_vmnk)),
|
||||
make_layout(size<3>(thr_layout_vmnk))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrN,ThrK)),(FrgV,(RestN,RestK)))
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
template <class SFATensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_SFA(SFATensor&& sfatensor, ThrMma& thread_mma) {
|
||||
auto thr_tensor = make_tensor(static_cast<SFATensor&&>(sfatensor).data(), thrfrg_SFA(sfatensor.layout(),thread_mma));
|
||||
auto thr_vmnk = thread_mma.thr_vmnk_;
|
||||
auto thr_vmk = make_coord(get<0>(thr_vmnk), make_coord(get<1>(thr_vmnk), get<3>(thr_vmnk)));
|
||||
return thr_tensor(thr_vmk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
}
|
||||
|
||||
template <class SFATensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_fragment_SFA(SFATensor&& sfatensor, ThrMma& thread_mma) {
|
||||
using ValTypeSF = typename ThrMma::Atom::Traits::ValTypeSF;
|
||||
return make_fragment_like<ValTypeSF>(partition_SFA(sfatensor, thread_mma));
|
||||
}
|
||||
|
||||
template <class SFBTensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_SFB(SFBTensor&& sfbtensor, ThrMma& thread_mma) {
|
||||
auto thr_tensor = make_tensor(static_cast<SFBTensor&&>(sfbtensor).data(), thrfrg_SFB(sfbtensor.layout(),thread_mma));
|
||||
auto thr_vmnk = thread_mma.thr_vmnk_;
|
||||
auto thr_vnk = make_coord(get<0>(thr_vmnk), make_coord(get<2>(thr_vmnk), get<3>(thr_vmnk)));
|
||||
return thr_tensor(thr_vnk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
}
|
||||
|
||||
template <class SFBTensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_fragment_SFB(SFBTensor&& sfbtensor, ThrMma& thread_mma) {
|
||||
using ValTypeSF = typename ThrMma::Atom::Traits::ValTypeSF;
|
||||
return make_fragment_like<ValTypeSF>(partition_SFB(sfbtensor, thread_mma));
|
||||
}
|
||||
|
||||
template<class TiledMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutSFA_TV(TiledMma& mma)
|
||||
{
|
||||
// (M,K) -> (M,K)
|
||||
auto tile_shape_mnk = tile_shape(mma);
|
||||
auto ref_A = make_layout(make_shape(size<0>(tile_shape_mnk), size<2>(tile_shape_mnk)));
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto atile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk), size<2>(thr_layout_vmnk)),
|
||||
make_stride( Int<1>{} , Int<0>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk);
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_SFA(ref_A, mma).compose(atile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
template<class TiledMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutSFB_TV(TiledMma& mma)
|
||||
{
|
||||
// (N,K) -> (N,K)
|
||||
auto tile_shape_mnk = tile_shape(mma);
|
||||
auto ref_B = make_layout(make_shape(size<1>(tile_shape_mnk), size<2>(tile_shape_mnk)));
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto btile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk), size<2>(thr_layout_vmnk)),
|
||||
make_stride( Int<0>{} , Int<1>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk);
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_SFB(ref_B, mma).compose(btile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
} // namespace cute
|
||||
@@ -0,0 +1,222 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cutlass/cutlass.h>
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "named_barrier.h"
|
||||
#include "utils.h"
|
||||
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <typename Ktraits>
|
||||
struct CollectiveEpilogueFwd{
|
||||
|
||||
using Element = typename Ktraits::ElementOut;
|
||||
static constexpr int kBlockM = Ktraits::kBlockM;
|
||||
static constexpr int kBlockN = Ktraits::kBlockN;
|
||||
static constexpr int kHeadDim = Ktraits::kHeadDim;
|
||||
using TileShape_MNK = Shape<Int<kBlockM>, Int<kBlockN>, Int<kHeadDim>>;
|
||||
static constexpr int kNWarps = Ktraits::kNWarps;
|
||||
static constexpr int kNThreads = kNWarps * cutlass::NumThreadsPerWarp;
|
||||
static constexpr int NumMmaThreads = kNThreads - cutlass::NumThreadsPerWarpGroup;
|
||||
|
||||
using GmemTiledCopyOTMA = cute::SM90_TMA_STORE;
|
||||
|
||||
// These are for storing the output tensor without TMA (e.g., for setting output to zero)
|
||||
static constexpr int kGmemElemsPerLoad = sizeof(cute::uint128_t) / sizeof(Element);
|
||||
static_assert(kHeadDim % kGmemElemsPerLoad == 0, "kHeadDim must be a multiple of kGmemElemsPerLoad");
|
||||
static constexpr int kGmemThreadsPerRow = kHeadDim / kGmemElemsPerLoad;
|
||||
static_assert(NumMmaThreads % kGmemThreadsPerRow == 0, "NumMmaThreads must be a multiple of kGmemThreadsPerRow");
|
||||
using GmemLayoutAtom = Layout<Shape <Int<NumMmaThreads / kGmemThreadsPerRow>, Int<kGmemThreadsPerRow>>,
|
||||
Stride<Int<kGmemThreadsPerRow>, _1>>;
|
||||
using GmemTiledCopyO = decltype(
|
||||
make_tiled_copy(Copy_Atom<DefaultCopy, Element>{},
|
||||
GmemLayoutAtom{},
|
||||
Layout<Shape<_1, Int<kGmemElemsPerLoad>>>{})); // Val layout, 8 or 16 vals per store
|
||||
|
||||
using SmemLayoutO = typename Ktraits::SmemLayoutO;
|
||||
|
||||
using SmemCopyAtomO = Copy_Atom<SM90_U32x2_STSM_N, Element>;
|
||||
using SharedStorage = cute::array_aligned<Element, cute::cosize_v<SmemLayoutO>>;
|
||||
|
||||
using ShapeO = cute::Shape<int32_t, int32_t, int32_t, int32_t>; // (seqlen_q, d, head, batch)
|
||||
using StrideO = cute::Stride<int64_t, _1, int64_t, int64_t>;
|
||||
using StrideLSE = cute::Stride<_1, int64_t, int64_t>; // (seqlen_q, head, batch)
|
||||
|
||||
using TMA_O = decltype(make_tma_copy(
|
||||
GmemTiledCopyOTMA{},
|
||||
make_tensor(make_gmem_ptr(static_cast<Element*>(nullptr)), repeat_like(StrideO{}, int32_t(0)), StrideO{}),
|
||||
SmemLayoutO{},
|
||||
select<0, 2>(TileShape_MNK{}),
|
||||
_1{})); // no mcast for O
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
Element* ptr_O;
|
||||
ShapeO const shape_O;
|
||||
StrideO const stride_O;
|
||||
float* ptr_LSE;
|
||||
StrideLSE const stride_LSE;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
Element* ptr_O;
|
||||
ShapeO const shape_O;
|
||||
StrideO const stride_O;
|
||||
float* ptr_LSE;
|
||||
StrideLSE const stride_LSE;
|
||||
TMA_O tma_store_O;
|
||||
};
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
Tensor mO = make_tensor(make_gmem_ptr(args.ptr_O), args.shape_O, args.stride_O);
|
||||
TMA_O tma_store_O = make_tma_copy(
|
||||
GmemTiledCopyOTMA{},
|
||||
mO,
|
||||
SmemLayoutO{},
|
||||
select<0, 2>(TileShape_MNK{}),
|
||||
_1{}); // no mcast for O
|
||||
return {args.ptr_O, args.shape_O, args.stride_O, args.ptr_LSE, args.stride_LSE, tma_store_O};
|
||||
}
|
||||
|
||||
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
|
||||
CUTLASS_DEVICE
|
||||
static void prefetch_tma_descriptors(Params const& epilogue_params) {
|
||||
cute::prefetch_tma_descriptor(epilogue_params.tma_store_O.get_tma_descriptor());
|
||||
}
|
||||
|
||||
template <typename SharedStorage, typename FrgTensorO, typename TiledMma>
|
||||
CUTLASS_DEVICE void
|
||||
mma_store(
|
||||
SharedStorage& shared_storage,
|
||||
TiledMma tiled_mma,
|
||||
FrgTensorO const& tOrO,
|
||||
int thread_idx
|
||||
){
|
||||
Tensor sO = cute::as_position_independent_swizzle_tensor(make_tensor(make_smem_ptr(shared_storage.smem_o.begin()), SmemLayoutO{}));
|
||||
auto smem_tiled_copy_O = make_tiled_copy_C(SmemCopyAtomO{}, tiled_mma);
|
||||
auto smem_thr_copy_O = smem_tiled_copy_O.get_thread_slice(thread_idx);
|
||||
constexpr int numel = decltype(size(tOrO))::value;
|
||||
cutlass::NumericArrayConverter<Element, float, numel> convert_op;
|
||||
// HACK: this requires tensor to be "contiguous"
|
||||
auto frag = convert_op(*reinterpret_cast<const cutlass::Array<float, numel> *>(tOrO.data()));
|
||||
auto tOrO_out = make_tensor(make_rmem_ptr<Element>(&frag), tOrO.layout());
|
||||
Tensor taccOrO = smem_thr_copy_O.retile_S(tOrO_out); // ((Atom,AtomNum), MMA_M, MMA_N)
|
||||
Tensor taccOsO = smem_thr_copy_O.partition_D(sO); // ((Atom,AtomNum),PIPE_M,PIPE_N)
|
||||
cute::copy(smem_tiled_copy_O, taccOrO, taccOsO);
|
||||
cutlass::arch::fence_view_async_shared(); // ensure smem writes are visible to TMA
|
||||
}
|
||||
|
||||
template<typename SharedStorage, typename Params, typename WorkTileInfo, typename SchedulerParams>
|
||||
CUTLASS_DEVICE void
|
||||
tma_store(
|
||||
SharedStorage& shared_storage,
|
||||
Params const& epilogue_params,
|
||||
WorkTileInfo work_tile_info,
|
||||
SchedulerParams const& scheduler_params,
|
||||
int thread_idx
|
||||
) {
|
||||
auto [m_block, bidh, bidb] = work_tile_info.get_block_coord(scheduler_params);
|
||||
Tensor sO = cute::as_position_independent_swizzle_tensor(make_tensor(make_smem_ptr(shared_storage.smem_o.begin()), SmemLayoutO{}));
|
||||
Tensor mO = epilogue_params.tma_store_O.get_tma_tensor(epilogue_params.shape_O);
|
||||
Tensor gO = local_tile(mO(_, _, bidh, bidb), select<0, 2>(TileShape_MNK{}), make_coord(m_block, _0{})); // (M, K)
|
||||
auto block_tma_O = epilogue_params.tma_store_O.get_slice(_0{});
|
||||
Tensor tOgO = block_tma_O.partition_D(gO); // (TMA, TMA_M, TMA_K)
|
||||
Tensor tOsO = block_tma_O.partition_S(sO); // (TMA, TMA_M, TMA_K)
|
||||
|
||||
// auto shape_LSE = select<0, 2, 3>(epilogue_params.shape_O);
|
||||
// Tensor mLSE = make_tensor(make_gmem_ptr(epilogue_params.ptr_LSE), shape_LSE, epilogue_params.stride_LSE);
|
||||
// Tensor gLSE = local_tile(mLSE(_, bidh, bidb), Shape<Int<kBlockM>>{}, make_coord(m_block));
|
||||
|
||||
// Tensor caccO = cute::make_identity_tensor(select<0, 2>(TileShape_MNK{}));
|
||||
// auto thread_mma = tiled_mma.get_thread_slice(thread_idx);
|
||||
// Tensor taccOcO = thread_mma.partition_C(caccO); // (MMA,MMA_M,MMA_K)
|
||||
// static_assert(decltype(size<0, 0>(taccOcO))::value == 2);
|
||||
// static_assert(decltype(size<0, 1>(taccOcO))::value == 2);
|
||||
// // // // taccOcO has shape ((2, 2, V), MMA_M, MMA_K), we only take only the row indices.
|
||||
// Tensor taccOcO_row = taccOcO(make_coord(_0{}, _), _, _0{});
|
||||
// CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
|
||||
// if (get<1>(taccOcO_row(_0{})) == 0) {
|
||||
// #pragma unroll
|
||||
// for (int mi = 0; mi < size(lse); ++mi) {
|
||||
// const int row = get<0>(taccOcO_row(mi));
|
||||
// if (row < get<0>(shape_LSE) - m_block * kBlockM) { gLSE(row) = lse(mi); }
|
||||
// }
|
||||
// }
|
||||
|
||||
// if (cutlass::canonical_warp_idx_sync() == kNWarps - 1) {
|
||||
// cutlass::arch::NamedBarrier::sync(NumMmaThreads + cutlass::NumThreadsPerWarp,
|
||||
// static_cast<uint32_t>(FP4NamedBarriers::EpilogueBarrier));
|
||||
// int const lane_predicate = cute::elect_one_sync();
|
||||
// if (lane_predicate) {
|
||||
// cute::copy(epilogue_params.tma_store_O, tOsO, tOgO);
|
||||
// tma_store_arrive();
|
||||
// }
|
||||
// }
|
||||
cute::copy(epilogue_params.tma_store_O, tOsO, tOgO);
|
||||
tma_store_arrive();
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
store_tail() {
|
||||
tma_store_wait<0>();
|
||||
}
|
||||
|
||||
// Write 0 to output and -inf to LSE
|
||||
CUTLASS_DEVICE void
|
||||
store_zero(
|
||||
Params const& epilogue_params,
|
||||
int thread_idx,
|
||||
cute::tuple<int32_t, int32_t, int32_t> const& block_coord
|
||||
) {
|
||||
auto [m_block, bidh, bidb] = block_coord;
|
||||
Tensor mO = make_tensor(make_gmem_ptr(epilogue_params.ptr_O), epilogue_params.shape_O, epilogue_params.stride_O);
|
||||
Tensor gO = local_tile(mO(_, _, bidh, bidb), select<0, 2>(TileShape_MNK{}), make_coord(m_block, _0{})); // (M, K)
|
||||
auto shape_LSE = select<0, 2, 3>(epilogue_params.shape_O);
|
||||
Tensor mLSE = make_tensor(make_gmem_ptr(epilogue_params.ptr_LSE), shape_LSE, epilogue_params.stride_LSE);
|
||||
Tensor gLSE = local_tile(mLSE(_, bidh, bidb), Shape<Int<kBlockM>>{}, make_coord(m_block));
|
||||
|
||||
GmemTiledCopyO gmem_tiled_copy_O;
|
||||
auto gmem_thr_copy_O = gmem_tiled_copy_O.get_thread_slice(thread_idx);
|
||||
Tensor tOgO = gmem_thr_copy_O.partition_D(gO);
|
||||
Tensor tOrO = make_fragment_like(tOgO);
|
||||
clear(tOrO);
|
||||
// Construct identity layout for sO
|
||||
Tensor cO = cute::make_identity_tensor(select<0, 2>(TileShape_MNK{})); // (BLK_M,BLK_K) -> (blk_m,blk_k)
|
||||
// Repeat the partitioning with identity layouts
|
||||
Tensor tOcO = gmem_thr_copy_O.partition_D(cO);
|
||||
Tensor tOpO = make_tensor<bool>(make_shape(size<2>(tOgO)));
|
||||
#pragma unroll
|
||||
for (int k = 0; k < size(tOpO); ++k) { tOpO(k) = get<1>(tOcO(_0{}, _0{}, k)) < get<1>(epilogue_params.shape_O); }
|
||||
// Clear_OOB_K must be false since we don't want to write zeros to gmem
|
||||
flash::copy</*Is_even_MN=*/false, /*Is_even_K=*/false, /*Clear_OOB_MN=*/false, /*Clear_OOB_K=*/false>(
|
||||
gmem_tiled_copy_O, tOrO, tOgO, tOcO, tOpO, get<0>(epilogue_params.shape_O) - m_block * kBlockM
|
||||
);
|
||||
static_assert(kBlockM <= NumMmaThreads);
|
||||
if (thread_idx < get<0>(shape_LSE) - m_block * kBlockM) { gLSE(thread_idx) = INFINITY; }
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,202 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cute/algorithm/copy.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/layout/layout.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
|
||||
#include "blockscaled_layout.h"
|
||||
#include "cute_extension.h"
|
||||
#include "named_barrier.h"
|
||||
using namespace cute;
|
||||
|
||||
template <
|
||||
int kStages,
|
||||
int EpiStages,
|
||||
typename Element,
|
||||
typename ElementSF,
|
||||
typename OutputType,
|
||||
typename SmemLayoutQ,
|
||||
typename SmemLayoutK,
|
||||
typename SmemLayoutV,
|
||||
typename SmemLayoutDS,
|
||||
typename SmemLayoutO,
|
||||
typename SmemLayoutSFQ,
|
||||
typename SmemLayoutSFK,
|
||||
typename SmemLayoutSFV
|
||||
>
|
||||
struct SharedStorageQKVOwithSF : cute::aligned_struct<128, _0>{
|
||||
|
||||
alignas(1024) cute::ArrayEngine<Element, cute::cosize_v<SmemLayoutQ>> smem_q;
|
||||
alignas(1024) cute::ArrayEngine<Element, cute::cosize_v<SmemLayoutK>> smem_k;
|
||||
cute::ArrayEngine<ElementSF, cute::cosize_v<SmemLayoutSFQ>> smem_SFQ;
|
||||
cute::ArrayEngine<ElementSF, cute::cosize_v<SmemLayoutSFK>> smem_SFK;
|
||||
cute::ArrayEngine<ElementSF, cute::cosize_v<SmemLayoutSFV>> smem_SFV;
|
||||
alignas(1024) cute::ArrayEngine<float, cute::cosize_v<SmemLayoutDS>> smem_ds;
|
||||
alignas(1024) cute::ArrayEngine<Element, cute::cosize_v<SmemLayoutV>> smem_v;
|
||||
alignas(1024) cute::ArrayEngine<OutputType, cute::cosize_v<SmemLayoutO>> smem_o;
|
||||
|
||||
struct {
|
||||
alignas(16) typename cutlass::PipelineTmaAsync<1>::SharedStorage pipeline_q;
|
||||
alignas(16) typename cutlass::PipelineTmaAsync<kStages>::SharedStorage pipeline_k;
|
||||
alignas(16) typename cutlass::PipelineTmaAsync<kStages>::SharedStorage pipeline_v;
|
||||
alignas(16) typename flash::OrderedSequenceBarrierVarGroupSize<EpiStages, 2>::SharedStorage barrier_o;
|
||||
int tile_count_semaphore;
|
||||
};
|
||||
};
|
||||
|
||||
template <
|
||||
int kHeadDim_,
|
||||
int kBlockM_,
|
||||
int kBlockN_,
|
||||
int kStages_,
|
||||
int kClusterM_,
|
||||
bool BlockMean_,
|
||||
typename ElementPairType_ = cutlass::nv_float4_t<cutlass::float_e2m1_t>,
|
||||
typename ElementOut_ = cutlass::bfloat16_t
|
||||
>
|
||||
struct Flash_fwd_kernel_traits {
|
||||
static constexpr int kBlockM = kBlockM_;
|
||||
static constexpr int kBlockN = kBlockN_;
|
||||
static constexpr int kHeadDim = kHeadDim_;
|
||||
static constexpr bool BlockMean = BlockMean_;
|
||||
static constexpr bool SmoothQ = true;
|
||||
static_assert(kHeadDim % 32 == 0);
|
||||
static_assert(kBlockM == 64 || kBlockM == 128);
|
||||
static constexpr int kNWarps = kBlockM == 128 ? 12 : 8;
|
||||
static constexpr int kNThreads = kNWarps * cutlass::NumThreadsPerWarp;
|
||||
static constexpr int kClusterM = kClusterM_;
|
||||
static constexpr int kStages = kStages_;
|
||||
static constexpr int EpiStages = 1;
|
||||
static constexpr int NumSFQK = kHeadDim / 16;
|
||||
static constexpr int NumSFPV = kBlockN / 16;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using Element = cutlass::float_e2m1_t;
|
||||
using ElementAccum = float;
|
||||
using ElementOut = ElementOut_;
|
||||
using index_t = int64_t;
|
||||
static constexpr auto SFVectorSize = 16;
|
||||
using TileShape_MNK = Shape<Int<kBlockM>, Int<kBlockN>, Int<kHeadDim>>;
|
||||
using ClusterShape_MNK = Shape<_1, _1, _1>;
|
||||
using PermTileM = decltype(cute::min(size<0>(TileShape_MNK{}), _128{}));
|
||||
using PermTileN = _32;
|
||||
using PermTileK = Int<kHeadDim>;
|
||||
|
||||
using ElementQMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element<Element>());
|
||||
using ElementKMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element<Element>());
|
||||
|
||||
using AtomLayoutMNK = std::conditional_t<kBlockM == 128,
|
||||
Layout<Shape<_8, _1, _1>>,
|
||||
Layout<Shape<_4, _1, _1>>
|
||||
>;
|
||||
using TiledMmaQK = decltype(cute::make_tiled_mma(
|
||||
cute::SM120::BLOCKSCALED::SM120_16x32x64_TN_VS_NVFP4{},
|
||||
AtomLayoutMNK{},
|
||||
Tile<PermTileM, PermTileN, PermTileK>{}
|
||||
));
|
||||
|
||||
using TiledMmaPV = decltype(cute::make_tiled_mma(
|
||||
cute::SM120::BLOCKSCALED::SM120_16x32x64_TN_VS_NVFP4{},
|
||||
AtomLayoutMNK{},
|
||||
Tile<PermTileM, _32, PermTileK>{}
|
||||
));
|
||||
|
||||
static constexpr int MMA_NSF = size<2>(typename TiledMmaQK::AtomShape_MNK{}) / SFVectorSize;
|
||||
|
||||
using GmemTiledCopy = SM90_TMA_LOAD;
|
||||
using GmemTiledCopySF = SM90_TMA_LOAD;
|
||||
|
||||
using SmemLayoutAtomQ = decltype(cutlass::gemm::collective::detail::sm120_rr_smem_selector<Element, decltype(size<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutAtomK = decltype(cutlass::gemm::collective::detail::sm120_rr_smem_selector<Element, decltype(size<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutAtomV = decltype(cutlass::gemm::collective::detail::sm120_rr_smem_selector<Element, decltype(size<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutAtomVt = decltype(cutlass::gemm::collective::detail::sm120_rr_smem_selector<Element, decltype(size<1>(TileShape_MNK{}))>());
|
||||
using SmemLayoutQ = decltype(tile_to_shape(SmemLayoutAtomQ{}, select<0, 2>(TileShape_MNK{})));
|
||||
using SmemLayoutK =
|
||||
decltype(tile_to_shape(SmemLayoutAtomK{},
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{}), Int<kStages>{})));
|
||||
using SmemLayoutV =
|
||||
decltype(tile_to_shape(SmemLayoutAtomV{},
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{}), Int<kStages>{})));
|
||||
using SmemLayoutVt =
|
||||
decltype(tile_to_shape(SmemLayoutAtomVt{},
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{}), Int<kStages>{})));
|
||||
using SmemLayoutAtomDS = Layout<Shape<Int<kBlockM>, Int<kBlockN>>, Stride<_0, _1>>;
|
||||
using SmemLayoutDS =
|
||||
decltype(tile_to_shape(SmemLayoutAtomDS{},
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<1>(TileShape_MNK{}), Int<kStages>{})));
|
||||
|
||||
using SmemCopyAtomQ = Copy_Atom<SM75_U32x4_LDSM_N, Element>;
|
||||
using SmemCopyAtomKV = Copy_Atom<SM75_U32x4_LDSM_N, Element>;
|
||||
using SmemCopyAtomSF = Copy_Atom<UniversalCopy<ElementSF>, ElementSF>;
|
||||
using SmemCopyAtomDS = Copy_Atom<UniversalCopy<float>, float>;
|
||||
|
||||
using BlkScaledConfig = flash::BlockScaledConfig<SFVectorSize>;
|
||||
using LayoutSF = typename BlkScaledConfig::LayoutSF;
|
||||
using SfAtom = typename BlkScaledConfig::SfAtom;
|
||||
using SmemLayoutAtomSFQ = decltype(BlkScaledConfig::deduce_smem_layoutSFQ(TiledMmaQK{}, TileShape_MNK{}));
|
||||
using SmemLayoutAtomSFK = decltype(BlkScaledConfig::deduce_smem_layoutSFKV(TiledMmaQK{}, TileShape_MNK{}));
|
||||
using SmemLayoutAtomSFV = decltype(BlkScaledConfig::deduce_smem_layoutSFKV(TiledMmaPV{}, TileShape_MNK{}));
|
||||
using SmemLayoutAtomSFVt = decltype(BlkScaledConfig::deduce_smem_layoutSFVt(TiledMmaPV{}, Shape<Int<kBlockM>, Int<kHeadDim>, Int<kBlockN>>{}));
|
||||
using LayoutSFP = decltype(
|
||||
make_layout(
|
||||
make_shape(make_shape(_16{}, _4{}), _1{}, Int<kBlockN / 64>{}),
|
||||
make_stride(make_stride(_0{}, _1{}), _0{}, _4{})
|
||||
)
|
||||
);
|
||||
using LayoutP = decltype(
|
||||
make_layout(
|
||||
make_shape(make_shape(_8{}, _2{}, _2{}), _1{}, Int<kBlockN / 64>{}),
|
||||
make_stride(make_stride(_1{}, _8{}, _16{}), _0{}, _32{})
|
||||
)
|
||||
);
|
||||
using SmemLayoutSFQ = decltype(make_layout(
|
||||
shape(SmemLayoutAtomSFQ{}),
|
||||
stride(SmemLayoutAtomSFQ{})
|
||||
));
|
||||
using SmemLayoutSFK = decltype(make_layout(
|
||||
append(shape(SmemLayoutAtomSFK{}), Int<kStages>{}),
|
||||
append(stride(SmemLayoutAtomSFK{}), size(filter_zeros(SmemLayoutAtomSFK{})))
|
||||
));
|
||||
using SmemLayoutSFV = decltype(make_layout(
|
||||
append(shape(SmemLayoutAtomSFV{}), Int<kStages>{}),
|
||||
append(stride(SmemLayoutAtomSFV{}), size(filter_zeros(SmemLayoutAtomSFV{})))
|
||||
));
|
||||
using SmemLayoutSFVt = decltype(make_layout(
|
||||
append(shape(SmemLayoutAtomSFVt{}), Int<kStages>{}),
|
||||
append(stride(SmemLayoutAtomSFVt{}), size(filter_zeros(SmemLayoutAtomSFVt{})))
|
||||
));
|
||||
|
||||
using SmemLayoutAtomO = decltype(cutlass::gemm::collective::detail::ss_smem_selector<GMMA::Major::K, ElementOut,
|
||||
decltype(cute::get<0>(TileShape_MNK{})), decltype(cute::get<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutO = decltype(tile_to_shape(SmemLayoutAtomO{}, select<0, 2>(TileShape_MNK{}), Step<_1, _2>{}));
|
||||
using SharedStorage = SharedStorageQKVOwithSF<kStages, EpiStages, Element, ElementSF, ElementOut,
|
||||
SmemLayoutQ, SmemLayoutK, SmemLayoutV, SmemLayoutDS,
|
||||
SmemLayoutO, SmemLayoutSFQ, SmemLayoutSFK, SmemLayoutSFVt>;
|
||||
using MainloopPipeline = typename cutlass::PipelineTmaAsync<kStages>;
|
||||
using PipelineState = typename cutlass::PipelineState<kStages>;
|
||||
using MainloopPipelineQ = cutlass::PipelineTmaAsync<1>;
|
||||
using PipelineParamsQ = typename MainloopPipelineQ::Params;
|
||||
using PipelineStateQ = typename cutlass::PipelineState<1>;
|
||||
using EpilogueBarrier = typename flash::OrderedSequenceBarrierVarGroupSize<EpiStages, 2>;
|
||||
};
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
// Modified from the original SageAttention3 code
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/arch/reg_reconfig.h>
|
||||
#include <cutlass/array.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
#include <cutlass/numeric_conversion.h>
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
|
||||
#include "params.h"
|
||||
#include "utils.h"
|
||||
#include "tile_scheduler.h"
|
||||
#include "mainloop_tma_ws.h"
|
||||
#include "epilogue_tma_ws.h"
|
||||
#include "named_barrier.h"
|
||||
#include "softmax_fused.h"
|
||||
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <typename Ktraits, bool Is_causal, typename TileScheduler>
|
||||
__global__ void __launch_bounds__(Ktraits::kNWarps * cutlass::NumThreadsPerWarp, 1)
|
||||
compute_attn_ws(CUTE_GRID_CONSTANT Flash_fwd_params const params,
|
||||
CUTE_GRID_CONSTANT typename CollectiveMainloopFwd<Ktraits, Is_causal>::Params const mainloop_params,
|
||||
CUTE_GRID_CONSTANT typename CollectiveEpilogueFwd<Ktraits>::Params const epilogue_params,
|
||||
CUTE_GRID_CONSTANT typename TileScheduler::Params const scheduler_params
|
||||
) {
|
||||
|
||||
using Element = typename Ktraits::Element;
|
||||
using ElementAccum = typename Ktraits::ElementAccum;
|
||||
using SoftType = ElementAccum;
|
||||
using TileShape_MNK = typename Ktraits::TileShape_MNK;
|
||||
using ClusterShape = typename Ktraits::ClusterShape_MNK;
|
||||
|
||||
static constexpr int NumMmaThreads = size(typename Ktraits::TiledMmaQK{});
|
||||
static constexpr int NumCopyThreads = cutlass::NumThreadsPerWarpGroup;
|
||||
static constexpr int kBlockM = Ktraits::kBlockM;
|
||||
|
||||
using CollectiveMainloop = CollectiveMainloopFwd<Ktraits, Is_causal>;
|
||||
using CollectiveEpilogue = CollectiveEpilogueFwd<Ktraits>;
|
||||
|
||||
using MainloopPipeline = typename Ktraits::MainloopPipeline;
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
using PipelineState = typename MainloopPipeline::PipelineState;
|
||||
using MainloopPipelineQ = typename Ktraits::MainloopPipelineQ;
|
||||
using PipelineParamsQ = typename Ktraits::PipelineParamsQ;
|
||||
using PipelineStateQ = typename Ktraits::PipelineStateQ;
|
||||
using EpilogueBarrier = typename Ktraits::EpilogueBarrier;
|
||||
|
||||
|
||||
enum class WarpGroupRole {
|
||||
Producer = 0,
|
||||
Consumer0 = 1,
|
||||
Consumer1 = 2
|
||||
};
|
||||
enum class ProducerWarpRole {
|
||||
Mainloop = 0,
|
||||
Epilogue = 1,
|
||||
Warp2 = 2,
|
||||
Warp3 = 3
|
||||
};
|
||||
|
||||
extern __shared__ char shared_memory[];
|
||||
auto &shared_storage = *reinterpret_cast<typename Ktraits::SharedStorage*>(shared_memory);
|
||||
|
||||
int const lane_predicate = cute::elect_one_sync();
|
||||
int const warp_idx = cutlass::canonical_warp_idx_sync();
|
||||
int warp_group_idx = cutlass::canonical_warp_group_idx();
|
||||
int const warp_group_thread_idx = threadIdx.x % cutlass::NumThreadsPerWarpGroup;
|
||||
int warp_idx_in_warp_group = warp_idx % cutlass::NumWarpsPerWarpGroup;
|
||||
auto warp_group_role = WarpGroupRole(warp_group_idx);
|
||||
auto producer_warp_role = ProducerWarpRole(warp_idx_in_warp_group);
|
||||
|
||||
// Issue Tma Descriptor Prefetch from a single thread
|
||||
if (warp_idx == 0 && lane_predicate) {
|
||||
CollectiveMainloop::prefetch_tma_descriptors(mainloop_params);
|
||||
CollectiveEpilogue::prefetch_tma_descriptors(epilogue_params);
|
||||
}
|
||||
|
||||
// Obtain warp index
|
||||
|
||||
PipelineParams pipeline_params_v;
|
||||
pipeline_params_v.transaction_bytes = CollectiveMainloop::TmaTransactionBytesV;
|
||||
pipeline_params_v.role = warp_group_role == WarpGroupRole::Producer
|
||||
? MainloopPipeline::ThreadCategory::Producer
|
||||
: MainloopPipeline::ThreadCategory::Consumer;
|
||||
pipeline_params_v.is_leader = warp_group_thread_idx == 0;
|
||||
pipeline_params_v.num_consumers = NumMmaThreads;
|
||||
|
||||
PipelineParams pipeline_params_k;
|
||||
pipeline_params_k.transaction_bytes = CollectiveMainloop::TmaTransactionBytesK;
|
||||
pipeline_params_k.role = warp_group_role == WarpGroupRole::Producer
|
||||
? MainloopPipeline::ThreadCategory::Producer
|
||||
: MainloopPipeline::ThreadCategory::Consumer;
|
||||
pipeline_params_k.is_leader = warp_group_thread_idx == 0;
|
||||
pipeline_params_k.num_consumers = NumMmaThreads;
|
||||
|
||||
PipelineParamsQ pipeline_params_q;
|
||||
pipeline_params_q.transaction_bytes = CollectiveMainloop::TmaTransactionBytesQ;
|
||||
pipeline_params_q.role = warp_group_role == WarpGroupRole::Producer
|
||||
? MainloopPipelineQ::ThreadCategory::Producer
|
||||
: MainloopPipelineQ::ThreadCategory::Consumer;
|
||||
pipeline_params_q.is_leader = warp_group_thread_idx == 0;
|
||||
pipeline_params_q.num_consumers = NumMmaThreads;
|
||||
|
||||
// We're counting on pipeline_k to call cutlass::arch::fence_barrier_init();
|
||||
MainloopPipelineQ pipeline_q(shared_storage.pipeline_q, pipeline_params_q, ClusterShape{});
|
||||
MainloopPipeline pipeline_k(shared_storage.pipeline_k, pipeline_params_k, ClusterShape{});
|
||||
MainloopPipeline pipeline_v(shared_storage.pipeline_v, pipeline_params_v, ClusterShape{});
|
||||
|
||||
uint32_t epilogue_barrier_group_size_list[2] = {cutlass::NumThreadsPerWarp, NumMmaThreads};
|
||||
typename EpilogueBarrier::Params params_epilogue_barrier;
|
||||
params_epilogue_barrier.group_id = (warp_group_role == WarpGroupRole::Producer);
|
||||
params_epilogue_barrier.group_size_list = epilogue_barrier_group_size_list;
|
||||
EpilogueBarrier barrier_o(shared_storage.barrier_o, params_epilogue_barrier);
|
||||
|
||||
CollectiveMainloop collective_mainloop;
|
||||
CollectiveEpilogue collective_epilogue;
|
||||
__syncthreads();
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
cutlass::arch::warpgroup_reg_dealloc<24>();
|
||||
TileScheduler scheduler;
|
||||
|
||||
if (producer_warp_role == ProducerWarpRole::Mainloop) { // Load Q, K, V
|
||||
PipelineStateQ smem_pipe_write_q = cutlass::make_producer_start_state<MainloopPipelineQ>();
|
||||
PipelineState smem_pipe_write_k = cutlass::make_producer_start_state<MainloopPipeline>();
|
||||
PipelineState smem_pipe_write_v = cutlass::make_producer_start_state<MainloopPipeline>();
|
||||
|
||||
int work_idx = 0;
|
||||
for (auto work_tile_info = scheduler.get_initial_work(); work_tile_info.is_valid(scheduler_params); work_tile_info = scheduler.get_next_work(scheduler_params, work_tile_info)) {
|
||||
int tile_count_semaphore = 0;
|
||||
collective_mainloop.load(mainloop_params, scheduler_params,
|
||||
pipeline_q, pipeline_k, pipeline_v,
|
||||
smem_pipe_write_q, smem_pipe_write_k, smem_pipe_write_v,
|
||||
shared_storage, work_tile_info, work_idx, tile_count_semaphore);
|
||||
}
|
||||
collective_mainloop.load_tail(pipeline_q, pipeline_k, pipeline_v,
|
||||
smem_pipe_write_q, smem_pipe_write_k, smem_pipe_write_v);
|
||||
} else if (producer_warp_role == ProducerWarpRole::Epilogue) {
|
||||
for (auto work_tile_info = scheduler.get_initial_work(); work_tile_info.is_valid(scheduler_params); work_tile_info = scheduler.get_next_work(scheduler_params, work_tile_info)) {
|
||||
barrier_o.wait();
|
||||
collective_epilogue.tma_store(shared_storage, epilogue_params, work_tile_info, scheduler_params, threadIdx.x);
|
||||
collective_epilogue.store_tail();
|
||||
barrier_o.arrive();
|
||||
}
|
||||
|
||||
}
|
||||
} else if (warp_group_role == WarpGroupRole::Consumer0 || warp_group_role == WarpGroupRole::Consumer1) {
|
||||
cutlass::arch::warpgroup_reg_alloc<232>();
|
||||
typename Ktraits::TiledMmaPV tiled_mma_pv;
|
||||
TileScheduler scheduler{};
|
||||
PipelineState smem_pipe_read_k, smem_pipe_read_v;
|
||||
PipelineStateQ smem_pipe_read_q;
|
||||
|
||||
int work_idx = 0;
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for (auto work_tile_info = scheduler.get_initial_work(); work_tile_info.is_valid(scheduler_params); work_tile_info = scheduler.get_next_work(scheduler_params, work_tile_info)) {
|
||||
// Attention output (GEMM-II) accumulator.
|
||||
Tensor tOrO = partition_fragment_C(tiled_mma_pv, select<0, 2>(TileShape_MNK{}));
|
||||
// flash::Softmax<2 * (2 * kBlockM / NumMmaThreads)> softmax;
|
||||
// Pass single_level_p_quant flag to control P quantization mode
|
||||
flash::SoftmaxFused<2 * (2 * kBlockM / NumMmaThreads)> softmax_fused(params.single_level_p_quant);
|
||||
auto block_coord = work_tile_info.get_block_coord(scheduler_params);
|
||||
auto [m_block, bidh, bidb] = block_coord;
|
||||
|
||||
int n_block_max = collective_mainloop.get_n_block_max(mainloop_params, m_block);
|
||||
if (Is_causal && n_block_max <= 0) { // We exit early and write 0 to gO and -inf to gLSE.
|
||||
collective_epilogue.store_zero(epilogue_params, threadIdx.x - NumCopyThreads, block_coord);
|
||||
continue;
|
||||
}
|
||||
|
||||
collective_mainloop.mma(mainloop_params, pipeline_q, pipeline_k, pipeline_v, smem_pipe_read_q, smem_pipe_read_k, smem_pipe_read_v,
|
||||
tOrO, softmax_fused, n_block_max, threadIdx.x - NumCopyThreads, work_idx, m_block, shared_storage);
|
||||
barrier_o.wait();
|
||||
collective_epilogue.mma_store(shared_storage, tiled_mma_pv, tOrO, threadIdx.x - NumCopyThreads);
|
||||
barrier_o.arrive();
|
||||
++work_idx;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,114 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include "cutlass/cluster_launch.hpp"
|
||||
|
||||
#include "static_switch.h"
|
||||
#include "params.h"
|
||||
#include "tile_scheduler.h"
|
||||
#include "kernel_ws.h"
|
||||
#include "kernel_traits.h"
|
||||
#include "block_config.h"
|
||||
|
||||
|
||||
template<typename Kernel_traits, bool Is_causal>
|
||||
void run_flash_fwd(Flash_fwd_params ¶ms, cudaStream_t stream) {
|
||||
using Element = typename Kernel_traits::Element;
|
||||
using ElementSF = typename Kernel_traits::ElementSF;
|
||||
using ElementOut = typename Kernel_traits::ElementOut;
|
||||
using TileShape_MNK = typename Kernel_traits::TileShape_MNK;
|
||||
using ClusterShape = typename Kernel_traits::ClusterShape_MNK;
|
||||
using CollectiveMainloop = flash::CollectiveMainloopFwd<Kernel_traits, Is_causal>;
|
||||
using CollectiveEpilogue = flash::CollectiveEpilogueFwd<Kernel_traits>;
|
||||
// using Scheduler = flash::SingleTileScheduler;
|
||||
using Scheduler = flash::StaticPersistentTileScheduler;
|
||||
typename CollectiveMainloop::Params mainloop_params =
|
||||
CollectiveMainloop::to_underlying_arguments({
|
||||
static_cast<Element const*>(params.q_ptr),
|
||||
{params.seqlen_q, params.d, params.h, params.b}, // shape_Q
|
||||
{params.q_row_stride, _1{}, params.q_head_stride, params.q_batch_stride}, // stride_Q
|
||||
static_cast<Element const*>(params.k_ptr),
|
||||
{params.seqlen_k, params.d, params.h_k, params.b}, // shape_K
|
||||
{params.k_row_stride, _1{}, params.k_head_stride, params.k_batch_stride}, // stride_K
|
||||
{params.unpadded_seqlen_k, params.d, params.h_k, params.b}, // shape_K
|
||||
static_cast<Element const*>(params.v_ptr),
|
||||
{params.d, params.seqlen_k, params.h_k, params.b}, // shape_Vt
|
||||
{params.v_row_stride, _1{}, params.v_head_stride, params.v_batch_stride}, // stride_Vt
|
||||
static_cast<ElementSF const*>(params.sfq_ptr),
|
||||
{params.seqlen_q, params.d, params.h, params.b}, // shape_SFQ
|
||||
static_cast<ElementSF const*>(params.sfk_ptr),
|
||||
{params.seqlen_k, params.d, params.h_k, params.b}, // shape_SFK
|
||||
static_cast<ElementSF const*>(params.sfv_ptr),
|
||||
{params.d, params.seqlen_k, params.h_k, params.b}, // shape_SFVt
|
||||
static_cast<float const*>(params.delta_s_ptr),
|
||||
{params.seqlen_s, params.seqlen_k, params.h_k, params.b},
|
||||
{params.ds_row_stride, _1{}, params.ds_head_stride, params.ds_batch_stride},
|
||||
params.scale_softmax_log2
|
||||
});
|
||||
typename CollectiveEpilogue::Params epilogue_params =
|
||||
CollectiveEpilogue::to_underlying_arguments({
|
||||
static_cast<ElementOut*>(params.o_ptr),
|
||||
{params.seqlen_q, params.d, params.h, params.b}, // shape_O
|
||||
{params.o_row_stride, _1{}, params.o_head_stride, params.o_batch_stride}, // stride_O
|
||||
static_cast<float*>(params.softmax_lse_ptr),
|
||||
{_1{}, params.seqlen_q, params.h * params.seqlen_q}, // stride_LSE
|
||||
});
|
||||
|
||||
int num_blocks_m = cutlass::ceil_div(params.seqlen_q, Kernel_traits::kBlockM);
|
||||
num_blocks_m = cutlass::ceil_div(num_blocks_m, size<0>(ClusterShape{})) * size<0>(ClusterShape{});
|
||||
typename Scheduler::Arguments scheduler_args = {num_blocks_m, params.h, params.b};
|
||||
typename Scheduler::Params scheduler_params = Scheduler::to_underlying_arguments(scheduler_args);
|
||||
// Get the ptr to kernel function.
|
||||
void *kernel;
|
||||
kernel = (void *)flash::compute_attn_ws<Kernel_traits, Is_causal, Scheduler>;
|
||||
int smem_size = sizeof(typename Kernel_traits::SharedStorage);
|
||||
if (smem_size >= 48 * 1024) {
|
||||
C10_CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
|
||||
}
|
||||
static constexpr int ctaSize = Kernel_traits::kNWarps * 32;
|
||||
params.m_block_divmod = cutlass::FastDivmod(num_blocks_m);
|
||||
params.total_blocks = num_blocks_m * params.h * params.b;
|
||||
dim3 grid_dims = Scheduler::get_grid_dim(scheduler_args, 170);
|
||||
dim3 block_dims(ctaSize);
|
||||
dim3 cluster_dims(size<0>(ClusterShape{}), size<1>(ClusterShape{}), size<2>(ClusterShape{}));
|
||||
cutlass::ClusterLaunchParams launch_params{grid_dims, block_dims, cluster_dims, smem_size, stream};
|
||||
cutlass::launch_kernel_on_cluster(launch_params, kernel, params, mainloop_params, epilogue_params, scheduler_params);
|
||||
|
||||
C10_CUDA_KERNEL_LAUNCH_CHECK();
|
||||
}
|
||||
|
||||
|
||||
template<typename T, int Headdim, typename O = cutlass::bfloat16_t>
|
||||
void run_mha_fwd_(Flash_fwd_params ¶ms, cudaStream_t stream) {
|
||||
BOOL_SWITCH(params.is_causal, Is_causal, [&] {
|
||||
BOOL_SWITCH(params.per_block_mean, per_block, [&] {
|
||||
if constexpr (Headdim == 64 || Headdim == 128) {
|
||||
run_flash_fwd<
|
||||
Flash_fwd_kernel_traits<Headdim, flash::BLOCK_M, flash::BLOCK_N, 3, 1, per_block, T, O>,
|
||||
Is_causal
|
||||
>(params, stream);
|
||||
} else {
|
||||
static_assert(Headdim == 64 || Headdim == 128, "Unsupported Headdim");
|
||||
}
|
||||
});
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,920 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/array.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
#include <cutlass/numeric_conversion.h>
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
|
||||
#include "utils.h"
|
||||
#include "named_barrier.h"
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <typename Ktraits, bool Is_causal>
|
||||
struct CollectiveMainloopFwd {
|
||||
|
||||
using Element = typename Ktraits::Element;
|
||||
using ElementSF = typename Ktraits::ElementSF;
|
||||
// using TMAElement = Element;
|
||||
// using TMAElementSF = typename Ktraits::ElementSF;
|
||||
using TileShape_MNK = typename Ktraits::TileShape_MNK;
|
||||
using ClusterShape = typename Ktraits::ClusterShape_MNK;
|
||||
|
||||
static constexpr int kStages = Ktraits::kStages;
|
||||
static constexpr int kHeadDim = Ktraits::kHeadDim;
|
||||
static constexpr int BlockMean = Ktraits::BlockMean;
|
||||
using GmemTiledCopy = typename Ktraits::GmemTiledCopy;
|
||||
using SmemLayoutQ = typename Ktraits::SmemLayoutQ;
|
||||
using SmemLayoutK = typename Ktraits::SmemLayoutK;
|
||||
using SmemLayoutV = typename Ktraits::SmemLayoutV;
|
||||
using SmemLayoutVt = typename Ktraits::SmemLayoutVt;
|
||||
using SmemLayoutDS = typename Ktraits::SmemLayoutDS;
|
||||
using SmemLayoutAtomDS = typename Ktraits::SmemLayoutAtomDS;
|
||||
using LayoutDS = decltype(
|
||||
blocked_product(
|
||||
SmemLayoutAtomDS{},
|
||||
make_layout(
|
||||
make_shape(int32_t(0), int32_t(0), int32_t(0), int32_t(0)),
|
||||
make_stride(int32_t(0), _1{}, int32_t(0), int32_t(0)))
|
||||
)
|
||||
);
|
||||
using ShapeQKV = cute::Shape<int32_t, int32_t, int32_t, int32_t>; // (seqlen, d, head, batch)
|
||||
using StrideQKV = cute::Stride<int64_t, _1, int64_t, int64_t>;
|
||||
using ShapeSF = cute::Shape<int32_t, int32_t, int32_t, int32_t>; // (seqlen, d // 16, head, batch)
|
||||
using LayoutSF = typename Ktraits::LayoutSF;
|
||||
using LayoutP = typename Ktraits::LayoutP;
|
||||
using LayoutSFP = typename Ktraits::LayoutSFP;
|
||||
using SfAtom = typename Ktraits::SfAtom;
|
||||
using TMA_Q = decltype(make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
make_tensor(make_gmem_ptr(static_cast<Element const*>(nullptr)), repeat_like(StrideQKV{}, int32_t(0)), StrideQKV{}),
|
||||
SmemLayoutQ{},
|
||||
select<0, 2>(TileShape_MNK{}),
|
||||
_1{}));
|
||||
|
||||
using TMA_KV = decltype(make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
make_tensor(make_gmem_ptr(static_cast<Element const*>(nullptr)), repeat_like(StrideQKV{}, int32_t(0)), StrideQKV{}),
|
||||
take<0, 2>(SmemLayoutK{}),
|
||||
select<1, 2>(TileShape_MNK{}),
|
||||
_1{}));
|
||||
|
||||
using TMA_Vt = decltype(make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
make_tensor(make_gmem_ptr(static_cast<Element const*>(nullptr)), repeat_like(StrideQKV{}, int32_t(0)), StrideQKV{}),
|
||||
take<0, 2>(SmemLayoutVt{}),
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{}));
|
||||
|
||||
using TMA_DS = decltype(make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
make_tensor(make_gmem_ptr(static_cast<float const*>(nullptr)), LayoutDS{}),
|
||||
take<0, 2>(SmemLayoutDS{}),
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{}));
|
||||
|
||||
using BlkScaledConfig = typename Ktraits::BlkScaledConfig;
|
||||
using GmemTiledCopySF = typename Ktraits::GmemTiledCopySF;
|
||||
using SmemLayoutSFQ = typename Ktraits::SmemLayoutSFQ;
|
||||
using SmemLayoutSFK = typename Ktraits::SmemLayoutSFK;
|
||||
using SmemLayoutSFV = typename Ktraits::SmemLayoutSFV;
|
||||
using SmemLayoutSFVt = typename Ktraits::SmemLayoutSFVt;
|
||||
|
||||
using TMA_SFQ = decltype(make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
make_tensor(static_cast<ElementSF const*>(nullptr), LayoutSF{}),
|
||||
SmemLayoutSFQ{},
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<2>(TileShape_MNK{})),
|
||||
_1{})); // No programmatic multicast
|
||||
|
||||
|
||||
using TMA_SFKV = decltype(make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
make_tensor(static_cast<ElementSF const*>(nullptr), LayoutSF{}),
|
||||
SmemLayoutSFK{}(_,_,cute::Int<0>{}),
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{})),
|
||||
_1{}));
|
||||
|
||||
using TMA_SFVt = decltype(make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
make_tensor(static_cast<ElementSF const*>(nullptr), LayoutSF{}),
|
||||
SmemLayoutSFVt{}(_,_,cute::Int<0>{}),
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{}));
|
||||
|
||||
using SmemCopyAtomQ = typename Ktraits::SmemCopyAtomQ;
|
||||
using SmemCopyAtomKV = typename Ktraits::SmemCopyAtomKV;
|
||||
using SmemCopyAtomSF = typename Ktraits::SmemCopyAtomSF;
|
||||
using TiledMmaQK = typename Ktraits::TiledMmaQK;
|
||||
using TiledMmaPV = typename Ktraits::TiledMmaPV;
|
||||
static constexpr int NumMmaThreads = size(TiledMmaQK{});
|
||||
using MainloopPipeline = typename Ktraits::MainloopPipeline;
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
using PipelineState = typename MainloopPipeline::PipelineState;
|
||||
using MainloopPipelineQ = typename Ktraits::MainloopPipelineQ;
|
||||
using PipelineParamsQ = typename Ktraits::PipelineParamsQ;
|
||||
using PipelineStateQ = typename Ktraits::PipelineStateQ;
|
||||
using EpilogueBarrier = typename Ktraits::EpilogueBarrier;
|
||||
|
||||
// Set the bytes transferred in this TMA transaction (may involve multiple issues)
|
||||
static constexpr uint32_t TmaTransactionBytesQ = static_cast<uint32_t>(
|
||||
cutlass::bits_to_bytes(cosize((SmemLayoutSFQ{})) * cute::sizeof_bits_v<ElementSF>) +
|
||||
cutlass::bits_to_bytes(size((SmemLayoutQ{})) * sizeof_bits<Element>::value));
|
||||
|
||||
static constexpr uint32_t TmaTransactionBytesK = static_cast<uint32_t>(
|
||||
cutlass::bits_to_bytes(cosize(take<0,2>(SmemLayoutSFK{})) * cute::sizeof_bits_v<ElementSF>) +
|
||||
cutlass::bits_to_bytes(cosize(take<0,2>(SmemLayoutDS{})) * cute::sizeof_bits_v<float>) +
|
||||
cutlass::bits_to_bytes(size(take<0,2>(SmemLayoutK{})) * sizeof_bits<Element>::value));
|
||||
|
||||
static constexpr uint32_t TmaTransactionBytesV = static_cast<uint32_t>(
|
||||
cutlass::bits_to_bytes(cosize(take<0,2>(SmemLayoutSFVt{})) * cute::sizeof_bits_v<ElementSF>) +
|
||||
cutlass::bits_to_bytes(size(take<0,2>(SmemLayoutVt{})) * sizeof_bits<Element>::value));
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
Element const* ptr_Q;
|
||||
ShapeQKV const shape_Q;
|
||||
StrideQKV const stride_Q;
|
||||
Element const* ptr_K;
|
||||
ShapeQKV const shape_K;
|
||||
StrideQKV const stride_K;
|
||||
ShapeQKV const unpadded_shape_K;
|
||||
Element const* ptr_Vt;
|
||||
ShapeQKV const shape_Vt;
|
||||
StrideQKV const stride_Vt;
|
||||
ElementSF const* ptr_SFQ{nullptr};
|
||||
ShapeSF const shape_SFQ{};
|
||||
ElementSF const* ptr_SFK{nullptr};
|
||||
ShapeSF const shape_SFK{};
|
||||
ElementSF const* ptr_SFVt{nullptr};
|
||||
ShapeSF const shape_SFVt{};
|
||||
float const* ptr_ds;
|
||||
ShapeQKV const shape_ds;
|
||||
StrideQKV const stride_ds;
|
||||
float const softmax_scale_log2;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
ShapeQKV const shape_Q;
|
||||
LayoutSF const layout_SFQ;
|
||||
ShapeQKV const shape_K;
|
||||
ShapeQKV const unpadded_shape_K;
|
||||
LayoutSF const layout_SFK;
|
||||
ShapeQKV const shape_Vt;
|
||||
LayoutSF const layout_SFVt;
|
||||
LayoutDS const layout_DS;
|
||||
TMA_Q tma_load_Q;
|
||||
TMA_SFQ tma_load_SFQ;
|
||||
TMA_KV tma_load_K;
|
||||
TMA_SFKV tma_load_SFK;
|
||||
TMA_Vt tma_load_Vt;
|
||||
TMA_SFVt tma_load_SFVt;
|
||||
TMA_DS tma_load_DS;
|
||||
float const softmax_scale_log2;
|
||||
};
|
||||
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
Tensor mQ = make_tensor(make_gmem_ptr(args.ptr_Q), args.shape_Q, args.stride_Q);
|
||||
TMA_Q tma_load_Q = make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
mQ,
|
||||
SmemLayoutQ{},
|
||||
select<0, 2>(TileShape_MNK{}),
|
||||
_1{}); // no mcast for Q
|
||||
Tensor mK = make_tensor(make_gmem_ptr(args.ptr_K), args.shape_K, args.stride_K);
|
||||
TMA_KV tma_load_K = make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
mK,
|
||||
SmemLayoutK{}(_, _, _0{}),
|
||||
select<1, 2>(TileShape_MNK{}),
|
||||
_1{}); // mcast along M mode for this N load, if any
|
||||
Tensor mVt = make_tensor(make_gmem_ptr(args.ptr_Vt), args.shape_Vt, args.stride_Vt);
|
||||
TMA_Vt tma_load_Vt = make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
mVt,
|
||||
SmemLayoutVt{}(_, _, _0{}),
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{}); // mcast along M mode for this N load, if any
|
||||
auto [Seqlen_Q, Seqlen_K, HeadNum, Batch] = args.shape_ds;
|
||||
LayoutDS layout_ds = tile_to_shape(SmemLayoutAtomDS{}, make_shape(Seqlen_Q, Seqlen_K, HeadNum, Batch), Step<_2,_1,_3,_4>{});
|
||||
Tensor mDS = make_tensor(make_gmem_ptr(args.ptr_ds), layout_ds);
|
||||
TMA_DS tma_load_ds = make_tma_copy (
|
||||
GmemTiledCopy{},
|
||||
mDS,
|
||||
SmemLayoutDS{}(_, _, _0{}),
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{});
|
||||
LayoutSF layout_sfq = BlkScaledConfig::tile_atom_to_shape_SFQKV(args.shape_SFQ);
|
||||
Tensor mSFQ = make_tensor(make_gmem_ptr(args.ptr_SFQ), layout_sfq);
|
||||
TMA_SFQ tma_load_sfq = make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
mSFQ,
|
||||
SmemLayoutSFQ{},
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<2>(TileShape_MNK{})),
|
||||
_1{});
|
||||
LayoutSF layout_sfk = BlkScaledConfig::tile_atom_to_shape_SFQKV(args.shape_SFK);
|
||||
Tensor mSFK = make_tensor(make_gmem_ptr(args.ptr_SFK), layout_sfk);
|
||||
TMA_SFKV tma_load_sfk = make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
mSFK,
|
||||
SmemLayoutSFK{}(_, _, _0{}),
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{})),
|
||||
_1{});
|
||||
LayoutSF layout_sfvt = BlkScaledConfig::tile_atom_to_shape_SFVt(args.shape_SFVt);
|
||||
Tensor mSFVt = make_tensor(make_gmem_ptr(args.ptr_SFVt), layout_sfvt);
|
||||
TMA_SFVt tma_load_sfvt = make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
mSFVt,
|
||||
SmemLayoutSFVt{}(_, _, _0{}),
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{});
|
||||
return {args.shape_Q, layout_sfq,
|
||||
args.shape_K, args.unpadded_shape_K, layout_sfk,
|
||||
args.shape_Vt, layout_sfvt,
|
||||
layout_ds,
|
||||
tma_load_Q, tma_load_sfq,
|
||||
tma_load_K, tma_load_sfk,
|
||||
tma_load_Vt, tma_load_sfvt,
|
||||
tma_load_ds,
|
||||
args.softmax_scale_log2};
|
||||
}
|
||||
|
||||
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
|
||||
CUTLASS_DEVICE
|
||||
static void prefetch_tma_descriptors(Params const& mainloop_params) {
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_Q.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_K.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_Vt.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_SFQ.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_SFK.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_SFVt.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_DS.get_tma_descriptor());
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
int get_n_block_max(Params const& mainloop_params, int m_block) {
|
||||
static constexpr int kBlockM = get<0>(TileShape_MNK{});
|
||||
static constexpr int kBlockN = get<1>(TileShape_MNK{});
|
||||
int const seqlen_q = get<0>(mainloop_params.shape_Q);
|
||||
int const seqlen_k = get<0>(mainloop_params.shape_K);
|
||||
int n_block_max = cute::ceil_div(seqlen_k, kBlockN);
|
||||
if constexpr (Is_causal) {
|
||||
n_block_max = std::min(n_block_max,
|
||||
cute::ceil_div((m_block + 1) * kBlockM + seqlen_k - seqlen_q, kBlockN));
|
||||
}
|
||||
return n_block_max;
|
||||
}
|
||||
|
||||
template <class SFATensor, class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_SFA(SFATensor&& sfatensor, TiledMMA<Atom, TiledThr, TiledPerm>& mma)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(sfatensor) >= Int<2>{});
|
||||
|
||||
using AtomShape_MNK = typename Atom::Shape_MNK;
|
||||
using AtomLayoutSFA_TV = typename Atom::Traits::SFALayout;
|
||||
|
||||
auto permutation_mnk = TiledPerm{};
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<0>(permutation_mnk),
|
||||
get<2>(permutation_mnk));
|
||||
auto t_tensor = logical_divide(sfatensor, t_tile); // (PermM,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
auto a_tile = make_tile(make_layout(size<0>(AtomShape_MNK{})),
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto a_tensor = zipped_divide(t_tensor, a_tile); // ((AtomM,AtomK),(RestM,RestK))
|
||||
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = a_tensor.compose(AtomLayoutSFA_TV{},_); // ((ThrV,FrgV),(RestM,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<1>(thr_layout_vmnk)),
|
||||
make_layout(size<3>(thr_layout_vmnk))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrM,ThrK)),(FrgV,(RestM,RestK)))
|
||||
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
template <class SFBTensor, class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_SFB(SFBTensor&& sfbtensor, TiledMMA<Atom, TiledThr, TiledPerm>& mma)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(sfbtensor) >= Int<2>{});
|
||||
|
||||
using AtomShape_MNK = typename Atom::Shape_MNK;
|
||||
using AtomLayoutSFB_TV = typename Atom::Traits::SFBLayout;
|
||||
|
||||
auto permutation_mnk = TiledPerm{};
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<1>(permutation_mnk),
|
||||
get<2>(permutation_mnk));
|
||||
auto t_tensor = logical_divide(sfbtensor, t_tile); // (PermN,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
auto a_tile = make_tile(make_layout(size<1>(AtomShape_MNK{})),
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto a_tensor = zipped_divide(t_tensor, a_tile); // ((AtomN,AtomK),(RestN,RestK))
|
||||
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = a_tensor.compose(AtomLayoutSFB_TV{},_); // ((ThrV,FrgV),(RestN,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<2>(thr_layout_vmnk)),
|
||||
make_layout(size<3>(thr_layout_vmnk))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrN,ThrK)),(FrgV,(RestN,RestK)))
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
template <class SFATensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_fragment_SFA(SFATensor&& sfatensor, ThrMma& thread_mma)
|
||||
{
|
||||
using ValTypeSF = typename ThrMma::Atom::Traits::ValTypeSF;
|
||||
auto thr_tensor = make_tensor(static_cast<SFATensor&&>(sfatensor).data(), thrfrg_SFA(sfatensor.layout(),thread_mma));
|
||||
auto thr_vmnk = thread_mma.thr_vmnk_;
|
||||
auto thr_vmk = make_coord(get<0>(thr_vmnk), make_coord(get<1>(thr_vmnk), get<3>(thr_vmnk)));
|
||||
auto partition_SFA = thr_tensor(thr_vmk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
return make_fragment_like<ValTypeSF>(partition_SFA);
|
||||
}
|
||||
|
||||
template <class SFBTensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_fragment_SFB(SFBTensor&& sfbtensor, ThrMma& thread_mma)
|
||||
{
|
||||
using ValTypeSF = typename ThrMma::Atom::Traits::ValTypeSF;
|
||||
auto thr_tensor = make_tensor(static_cast<SFBTensor&&>(sfbtensor).data(), thrfrg_SFB(sfbtensor.layout(),thread_mma));
|
||||
auto thr_vmnk = thread_mma.thr_vmnk_;
|
||||
auto thr_vnk = make_coord(get<0>(thr_vmnk), make_coord(get<2>(thr_vmnk), get<3>(thr_vmnk)));
|
||||
auto partition_SFB = thr_tensor(thr_vnk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
return make_fragment_like<ValTypeSF>(partition_SFB);
|
||||
}
|
||||
|
||||
template<class TiledMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutSFA_TV(TiledMma& mma)
|
||||
{
|
||||
// (M,K) -> (M,K)
|
||||
auto tile_shape_mnk = tile_shape(mma);
|
||||
auto ref_A = make_layout(make_shape(size<0>(tile_shape_mnk), size<2>(tile_shape_mnk)));
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto atile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk), size<2>(thr_layout_vmnk)),
|
||||
make_stride( Int<1>{} , Int<0>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk);
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_SFA(ref_A, mma).compose(atile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
template<class TiledMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutSFB_TV(TiledMma& mma)
|
||||
{
|
||||
// (N,K) -> (N,K)
|
||||
auto tile_shape_mnk = tile_shape(mma);
|
||||
auto ref_B = make_layout(make_shape(size<1>(tile_shape_mnk), size<2>(tile_shape_mnk)));
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto btile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk), size<2>(thr_layout_vmnk)),
|
||||
make_stride( Int<0>{} , Int<1>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk);
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_SFB(ref_B, mma).compose(btile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
template <typename SchedulerParams, typename SharedStorage, typename WorkTileInfo>
|
||||
CUTLASS_DEVICE void
|
||||
load(Params const& mainloop_params,
|
||||
SchedulerParams const& scheduler_params,
|
||||
MainloopPipelineQ pipeline_q,
|
||||
MainloopPipeline pipeline_k,
|
||||
MainloopPipeline pipeline_v,
|
||||
PipelineStateQ& smem_pipe_write_q,
|
||||
PipelineState& smem_pipe_write_k,
|
||||
PipelineState& smem_pipe_write_v,
|
||||
SharedStorage &shared_storage,
|
||||
WorkTileInfo work_tile_info,
|
||||
int& work_idx,
|
||||
int& tile_count_semaphore
|
||||
) {
|
||||
|
||||
static constexpr int kBlockM = get<0>(TileShape_MNK{});
|
||||
static constexpr int kBlockN = get<1>(TileShape_MNK{});
|
||||
|
||||
auto [m_block, bidh, bidb] = work_tile_info.get_block_coord(scheduler_params);
|
||||
|
||||
int n_block_max = get_n_block_max(mainloop_params, m_block);
|
||||
|
||||
Tensor sQ = make_tensor(make_smem_ptr(shared_storage.smem_q.begin()), SmemLayoutQ{});
|
||||
Tensor sK = make_tensor(make_smem_ptr(shared_storage.smem_k.begin()), SmemLayoutK{});
|
||||
Tensor sVt = make_tensor(make_smem_ptr(shared_storage.smem_v.begin()), SmemLayoutVt{});
|
||||
Tensor sSFQ = make_tensor(make_smem_ptr(shared_storage.smem_SFQ.begin()), SmemLayoutSFQ{});
|
||||
Tensor sSFK = make_tensor(make_smem_ptr(shared_storage.smem_SFK.begin()), SmemLayoutSFK{});
|
||||
Tensor sSFVt = make_tensor(make_smem_ptr(shared_storage.smem_SFV.begin()), SmemLayoutSFVt{});
|
||||
Tensor sDS = make_tensor(make_smem_ptr(shared_storage.smem_ds.begin()), SmemLayoutDS{});
|
||||
|
||||
Tensor mQ = mainloop_params.tma_load_Q.get_tma_tensor(mainloop_params.shape_Q);
|
||||
Tensor mK = mainloop_params.tma_load_K.get_tma_tensor(mainloop_params.shape_K);
|
||||
Tensor mVt = mainloop_params.tma_load_Vt.get_tma_tensor(mainloop_params.shape_Vt);
|
||||
Tensor mDS = mainloop_params.tma_load_DS.get_tma_tensor(shape(mainloop_params.layout_DS));
|
||||
Tensor mSFQ = mainloop_params.tma_load_SFQ.get_tma_tensor(shape(mainloop_params.layout_SFQ));
|
||||
Tensor mSFK = mainloop_params.tma_load_SFK.get_tma_tensor(shape(mainloop_params.layout_SFK));
|
||||
Tensor mSFVt = mainloop_params.tma_load_SFVt.get_tma_tensor(shape(mainloop_params.layout_SFVt));
|
||||
uint32_t block_rank_in_cluster = cute::block_rank_in_cluster();
|
||||
constexpr uint32_t cluster_shape_x = get<0>(ClusterShape());
|
||||
uint2 cluster_local_block_id = {block_rank_in_cluster % cluster_shape_x, block_rank_in_cluster / cluster_shape_x};
|
||||
Tensor gQ = local_tile(mQ(_, _, bidh, bidb), select<0, 2>(TileShape_MNK{}), make_coord(m_block, _0{})); // (M, K)
|
||||
Tensor gK = local_tile(mK(_, _, bidh, bidb), select<1, 2>(TileShape_MNK{}), make_coord(_, _0{})); // (N, K, _)
|
||||
Tensor gVt = local_tile(mVt(_, _, bidh, bidb), make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})), make_coord(_0{}, _)); // (N, K, _)
|
||||
Tensor gDS = [&] {
|
||||
if constexpr (BlockMean) {
|
||||
return local_tile(mDS(_, _, bidh, bidb), select<0, 1>(TileShape_MNK{}), make_coord(m_block, _));
|
||||
} else {
|
||||
return local_tile(mDS(_, _, bidh, bidb), select<0, 1>(TileShape_MNK{}), make_coord(_0{}, _));
|
||||
}
|
||||
}();
|
||||
Tensor gSFQ = local_tile(mSFQ(_, _, bidh, bidb), select<0, 2>(TileShape_MNK{}), make_coord(m_block, _0{}));
|
||||
Tensor gSFK = local_tile(mSFK(_, _, bidh, bidb), select<1, 2>(TileShape_MNK{}), make_coord(_, _0{}));
|
||||
Tensor gSFVt = local_tile(mSFVt(_, _, bidh, bidb), make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})), make_coord(_0{}, _));
|
||||
auto block_tma_q = mainloop_params.tma_load_Q.get_slice(_0{});
|
||||
Tensor tQgQ = block_tma_q.partition_S(gQ);
|
||||
Tensor tQsQ = block_tma_q.partition_D(sQ);
|
||||
auto block_tma_sfq = mainloop_params.tma_load_SFQ.get_slice(_0{});
|
||||
Tensor tQgSFQ = block_tma_sfq.partition_S(gSFQ);
|
||||
Tensor tQsSFQ = block_tma_sfq.partition_D(sSFQ);
|
||||
auto block_tma_k = mainloop_params.tma_load_K.get_slice(cluster_local_block_id.x);
|
||||
Tensor tKgK = group_modes<0, 3>(block_tma_k.partition_S(gK));
|
||||
Tensor tKsK = group_modes<0, 3>(block_tma_k.partition_D(sK));
|
||||
auto block_tma_sfk = mainloop_params.tma_load_SFK.get_slice(cluster_local_block_id.x);
|
||||
Tensor tKgSFK = group_modes<0, 3>(block_tma_sfk.partition_S(gSFK));
|
||||
Tensor tKsSFK = group_modes<0, 3>(block_tma_sfk.partition_D(sSFK));
|
||||
auto block_tma_vt = mainloop_params.tma_load_Vt.get_slice(cluster_local_block_id.x);
|
||||
Tensor tVgVt = group_modes<0, 3>(block_tma_vt.partition_S(gVt));
|
||||
Tensor tVsVt = group_modes<0, 3>(block_tma_vt.partition_D(sVt));
|
||||
auto block_tma_sfvt = mainloop_params.tma_load_SFVt.get_slice(cluster_local_block_id.x);
|
||||
Tensor tVgSFVt = group_modes<0, 3>(block_tma_sfvt.partition_S(gSFVt));
|
||||
Tensor tVsSFVt = group_modes<0, 3>(block_tma_sfvt.partition_D(sSFVt));
|
||||
auto block_tma_ds = mainloop_params.tma_load_DS.get_slice(cluster_local_block_id.x);
|
||||
Tensor tDSgDS = group_modes<0, 3>(block_tma_ds.partition_S(gDS));
|
||||
Tensor tDSsDS = group_modes<0, 3>(block_tma_ds.partition_D(sDS));
|
||||
uint16_t mcast_mask_kv = 0;
|
||||
|
||||
int n_block = n_block_max - 1;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
if (lane_predicate) {
|
||||
pipeline_q.producer_acquire(smem_pipe_write_q);
|
||||
copy(mainloop_params.tma_load_Q.with(*pipeline_q.producer_get_barrier(smem_pipe_write_q), 0), tQgQ, tQsQ);
|
||||
copy(mainloop_params.tma_load_SFQ.with(*pipeline_q.producer_get_barrier(smem_pipe_write_q), 0), tQgSFQ, tQsSFQ);
|
||||
++smem_pipe_write_q;
|
||||
pipeline_k.producer_acquire(smem_pipe_write_k);
|
||||
copy(mainloop_params.tma_load_K.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tKgK(_, n_block), tKsK(_, smem_pipe_write_k.index()));
|
||||
copy(mainloop_params.tma_load_SFK.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tKgSFK(_, n_block), tKsSFK(_, smem_pipe_write_k.index()));
|
||||
copy(mainloop_params.tma_load_DS.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tDSgDS(_, n_block), tDSsDS(_, smem_pipe_write_k.index()));
|
||||
++smem_pipe_write_k;
|
||||
pipeline_v.producer_acquire(smem_pipe_write_v);
|
||||
copy(mainloop_params.tma_load_Vt.with(*pipeline_v.producer_get_barrier(smem_pipe_write_v), mcast_mask_kv),
|
||||
tVgVt(_, n_block), tVsVt(_, smem_pipe_write_v.index()));
|
||||
copy(mainloop_params.tma_load_SFVt.with(*pipeline_v.producer_get_barrier(smem_pipe_write_v), mcast_mask_kv),
|
||||
tVgSFVt(_, n_block), tVsSFVt(_, smem_pipe_write_v.index()));
|
||||
++smem_pipe_write_v;
|
||||
}
|
||||
|
||||
n_block--;
|
||||
if (lane_predicate) {
|
||||
// CUTLASS_PRAGMA_NO_UNROLL
|
||||
#pragma unroll 2
|
||||
for (; n_block >= 0; --n_block) {
|
||||
pipeline_k.producer_acquire(smem_pipe_write_k);
|
||||
copy(mainloop_params.tma_load_K.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tKgK(_, n_block), tKsK(_, smem_pipe_write_k.index()));
|
||||
copy(mainloop_params.tma_load_SFK.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tKgSFK(_, n_block), tKsSFK(_, smem_pipe_write_k.index()));
|
||||
copy(mainloop_params.tma_load_DS.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tDSgDS(_, n_block), tDSsDS(_, smem_pipe_write_k.index()));
|
||||
++smem_pipe_write_k;
|
||||
pipeline_v.producer_acquire(smem_pipe_write_v);
|
||||
copy(mainloop_params.tma_load_Vt.with(*pipeline_v.producer_get_barrier(smem_pipe_write_v), mcast_mask_kv),
|
||||
tVgVt(_, n_block), tVsVt(_, smem_pipe_write_v.index()));
|
||||
copy(mainloop_params.tma_load_SFVt.with(*pipeline_v.producer_get_barrier(smem_pipe_write_v), mcast_mask_kv),
|
||||
tVgSFVt(_, n_block), tVsSFVt(_, smem_pipe_write_v.index()));
|
||||
++smem_pipe_write_v;
|
||||
}
|
||||
}
|
||||
++work_idx;
|
||||
}
|
||||
|
||||
/// Perform a Producer Epilogue to prevent early exit of blocks in a Cluster
|
||||
CUTLASS_DEVICE void
|
||||
load_tail(MainloopPipelineQ pipeline_q,
|
||||
MainloopPipeline pipeline_k,
|
||||
MainloopPipeline pipeline_v,
|
||||
PipelineStateQ& smem_pipe_write_q,
|
||||
PipelineState& smem_pipe_write_k,
|
||||
PipelineState& smem_pipe_write_v) {
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
// Issue the epilogue waits
|
||||
if (lane_predicate) {
|
||||
pipeline_q.producer_tail(smem_pipe_write_q);
|
||||
pipeline_k.producer_tail(smem_pipe_write_k);
|
||||
pipeline_v.producer_tail(smem_pipe_write_v);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename SharedStorage, typename FrgTensorO, typename SoftmaxFused>
|
||||
CUTLASS_DEVICE void
|
||||
mma(Params const& mainloop_params,
|
||||
MainloopPipelineQ pipeline_q,
|
||||
MainloopPipeline pipeline_k,
|
||||
MainloopPipeline pipeline_v,
|
||||
PipelineStateQ& smem_pipe_read_q,
|
||||
PipelineState& smem_pipe_read_k,
|
||||
PipelineState& smem_pipe_read_v,
|
||||
FrgTensorO& tOrO_store,
|
||||
SoftmaxFused& softmax_fused,
|
||||
int n_block_count,
|
||||
int thread_idx,
|
||||
int work_idx,
|
||||
int m_block,
|
||||
SharedStorage& shared_storage
|
||||
) {
|
||||
|
||||
static_assert(is_rmem<FrgTensorO>::value, "O tensor must be rmem resident.");
|
||||
|
||||
static constexpr int kBlockM = get<0>(TileShape_MNK{});
|
||||
static constexpr int kBlockN = get<1>(TileShape_MNK{});
|
||||
static constexpr int kBlockK = get<2>(TileShape_MNK{});
|
||||
Tensor sQ = make_tensor(make_smem_ptr(shared_storage.smem_q.begin()), SmemLayoutQ{});
|
||||
Tensor sK = make_tensor(make_smem_ptr(shared_storage.smem_k.begin()), SmemLayoutK{});
|
||||
Tensor sVt = make_tensor(make_smem_ptr(shared_storage.smem_v.begin()), SmemLayoutVt{});
|
||||
Tensor sDS = make_tensor(make_smem_ptr(shared_storage.smem_ds.begin()), SmemLayoutDS{});
|
||||
Tensor sSFQ = make_tensor(make_smem_ptr(shared_storage.smem_SFQ.begin()), SmemLayoutSFQ{});
|
||||
Tensor sSFK = make_tensor(make_smem_ptr(shared_storage.smem_SFK.begin()), SmemLayoutSFK{});
|
||||
Tensor sSFVt = make_tensor(make_smem_ptr(shared_storage.smem_SFV.begin()), SmemLayoutSFVt{});
|
||||
|
||||
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ)));
|
||||
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK)));
|
||||
TiledMmaQK tiled_mma_qk;
|
||||
TiledMmaPV tiled_mma_pv;
|
||||
auto thread_mma_qk = tiled_mma_qk.get_thread_slice(thread_idx);
|
||||
auto thread_mma_pv = tiled_mma_pv.get_thread_slice(thread_idx);
|
||||
|
||||
Tensor tSrQ = thread_mma_qk.partition_fragment_A(sQ);
|
||||
Tensor tSrK = thread_mma_qk.partition_fragment_B(sK(_,_,Int<0>{}));
|
||||
Tensor tOrVt = thread_mma_pv.partition_fragment_B(sVt(_,_,Int<0>{}));
|
||||
Tensor tOrP = make_tensor_like<Element>(LayoutP{});
|
||||
Tensor tSrSFQ = partition_fragment_SFA(sSFQ, thread_mma_qk);
|
||||
Tensor tSrSFK = partition_fragment_SFB(sSFK(_,_,Int<0>{}), thread_mma_qk);
|
||||
Tensor tOrSFVt = partition_fragment_SFB(sSFVt(_,_,Int<0>{}), thread_mma_pv);
|
||||
Tensor tOrSFP = make_tensor<ElementSF>(LayoutSFP{});
|
||||
Tensor tOrSFP_flt = filter_zeros(tOrSFP);
|
||||
Tensor tSrDS = make_tensor<float>(make_shape(_8{}, _4{}), make_stride(_1{}, _8{}));
|
||||
// copy qk and sf from smem to rmem
|
||||
auto smem_tiled_copy_Q = make_tiled_copy_A(SmemCopyAtomQ{}, tiled_mma_qk);
|
||||
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(thread_idx);
|
||||
Tensor tSsQ = smem_thr_copy_Q.partition_S(as_position_independent_swizzle_tensor(sQ));
|
||||
Tensor tSrQ_copy_view = smem_thr_copy_Q.retile_D(tSrQ);
|
||||
|
||||
auto smem_tiled_copy_K = make_tiled_copy_B(SmemCopyAtomKV{}, tiled_mma_qk);
|
||||
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(thread_idx);
|
||||
Tensor tSsK = smem_thr_copy_K.partition_S(as_position_independent_swizzle_tensor(sK));
|
||||
Tensor tSrK_copy_view = smem_thr_copy_K.retile_D(tSrK);
|
||||
|
||||
auto smem_tiled_copy_V = make_tiled_copy_B(SmemCopyAtomKV{}, tiled_mma_pv);
|
||||
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(thread_idx);
|
||||
Tensor tOsVt = smem_thr_copy_V.partition_S(as_position_independent_swizzle_tensor(sVt));
|
||||
Tensor tOrVt_copy_view = smem_thr_copy_V.retile_D(tOrVt);
|
||||
|
||||
auto tile_shape_mnk = tile_shape(tiled_mma_qk);
|
||||
auto smem_tiled_copy_SFQ = make_tiled_copy_impl(SmemCopyAtomSF{},
|
||||
get_layoutSFA_TV(tiled_mma_qk),
|
||||
make_shape(size<0>(tile_shape_mnk), size<2>(tile_shape_mnk))
|
||||
);
|
||||
auto smem_thr_copy_SFQ = smem_tiled_copy_SFQ.get_thread_slice(thread_idx);
|
||||
Tensor tSsSFQ = smem_thr_copy_SFQ.partition_S(as_position_independent_swizzle_tensor(sSFQ));
|
||||
Tensor tSrSFQ_copy_view = smem_thr_copy_SFQ.retile_D(tSrSFQ);
|
||||
|
||||
auto smem_tiled_copy_SFK = make_tiled_copy_impl(SmemCopyAtomSF{},
|
||||
get_layoutSFB_TV(tiled_mma_qk),
|
||||
make_shape(size<1>(tile_shape_mnk), size<2>(tile_shape_mnk))
|
||||
);
|
||||
auto smem_thr_copy_SFK = smem_tiled_copy_SFK.get_thread_slice(thread_idx);
|
||||
Tensor tSsSFK = smem_thr_copy_SFK.partition_S(as_position_independent_swizzle_tensor(sSFK));
|
||||
Tensor tSrSFK_copy_view = smem_thr_copy_SFK.retile_D(tSrSFK);
|
||||
|
||||
auto smem_tiled_copy_SFV = make_tiled_copy_impl(SmemCopyAtomSF{},
|
||||
get_layoutSFB_TV(tiled_mma_pv),
|
||||
make_shape(size<1>(tile_shape_mnk), size<2>(tile_shape_mnk))
|
||||
);
|
||||
auto smem_thr_copy_SFV = smem_tiled_copy_SFV.get_thread_slice(thread_idx);
|
||||
Tensor tOsSFVt = smem_thr_copy_SFV.partition_S(as_position_independent_swizzle_tensor(sSFVt));
|
||||
Tensor tOrSFVt_copy_view = smem_thr_copy_SFV.retile_D(tOrSFVt);
|
||||
|
||||
auto consumer_wait = [](auto& pipeline, auto& smem_pipe_read) {
|
||||
auto barrier_token = pipeline.consumer_try_wait(smem_pipe_read);
|
||||
pipeline.consumer_wait(smem_pipe_read, barrier_token);
|
||||
};
|
||||
|
||||
int const seqlen_q = get<0>(mainloop_params.shape_Q);
|
||||
int const seqlen_k = get<0>(mainloop_params.shape_K);
|
||||
int const unpadded_seqlen_k = get<0>(mainloop_params.unpadded_shape_K);
|
||||
int n_block = n_block_count - 1;
|
||||
|
||||
auto copy_k_block = [&](auto block_id) {
|
||||
auto tSsK_stage = tSsK(_, _, _, smem_pipe_read_k.index());
|
||||
auto tSsSFK_stage = tSsSFK(_, _, _, smem_pipe_read_k.index());
|
||||
copy(smem_tiled_copy_K, tSsK_stage(_, _, block_id), tSrK_copy_view(_, _, block_id));
|
||||
copy(smem_tiled_copy_SFK, tSsSFK_stage(_, _, block_id), tSrSFK_copy_view(_, _, block_id));
|
||||
};
|
||||
|
||||
auto copy_v_block = [&](auto block_id) {
|
||||
auto tOsVt_stage = tOsVt(_, _, _, smem_pipe_read_v.index());
|
||||
auto tOsSFVt_stage = tOsSFVt(_, _, _, smem_pipe_read_v.index());
|
||||
copy(smem_tiled_copy_V, tOsVt_stage(_, _, block_id), tOrVt_copy_view(_, _, block_id));
|
||||
copy(smem_tiled_copy_SFV, tOsSFVt_stage(_, _, block_id), tOrSFVt_copy_view(_, _, block_id));
|
||||
};
|
||||
// auto gemm_qk = [&](auto block_id) {
|
||||
// cute::gemm(tiled_mma_qk, make_zip_tensor(tSrQ(_, _, block_id), tSrSFQ(_, _, block_id)), make_zip_tensor(tSrK(_, _, block_id), tSrSFK(_, _, block_id)), tSrS);
|
||||
// };
|
||||
// auto gemm_pv = [&](auto block_id) {
|
||||
// cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, block_id), tOrSFP(_, _, block_id)), make_zip_tensor(tOrVt(_, _, block_id), tOrSFVt(_, _, block_id)), tOrO);
|
||||
// };
|
||||
auto add_delta_s = [&](auto& acc) {
|
||||
// The MMA atom composites 4 sub-MMA m16n8k64 covering N=0-7, 8-15, 16-23, 24-31.
|
||||
// Each float4 register group spans two sub-MMAs, so N positions are scattered
|
||||
// (e.g., {2t, 2t+1, 8+2t, 9+2t}), not consecutive.
|
||||
float const* ds_ptr = reinterpret_cast<float const*>(
|
||||
&sDS(_0{}, _0{}, smem_pipe_read_k.index()));
|
||||
auto acc_float4 = recast<float4>(acc);
|
||||
int tid = threadIdx.x % 4;
|
||||
for (int i = 0; i < 4; i++) {
|
||||
int base_n = i * 32 + tid * 2;
|
||||
float4 delta_s_0 = make_float4(
|
||||
ds_ptr[base_n], ds_ptr[base_n + 1],
|
||||
ds_ptr[base_n + 8], ds_ptr[base_n + 9]);
|
||||
float4 delta_s_1 = make_float4(
|
||||
ds_ptr[base_n + 16], ds_ptr[base_n + 17],
|
||||
ds_ptr[base_n + 24], ds_ptr[base_n + 25]);
|
||||
acc_float4(make_coord(make_coord(_0{}, _0{}), _0{}), _0{}, i) = delta_s_0;
|
||||
acc_float4(make_coord(make_coord(_0{}, _0{}), _1{}), _0{}, i) = delta_s_0;
|
||||
acc_float4(make_coord(make_coord(_0{}, _1{}), _0{}), _0{}, i) = delta_s_1;
|
||||
acc_float4(make_coord(make_coord(_0{}, _1{}), _1{}), _0{}, i) = delta_s_1;
|
||||
}
|
||||
};
|
||||
consumer_wait(pipeline_q, smem_pipe_read_q);
|
||||
copy(smem_tiled_copy_Q, tSsQ, tSrQ_copy_view);
|
||||
copy(smem_tiled_copy_SFQ, tSsSFQ, tSrSFQ_copy_view);
|
||||
pipeline_q.consumer_release(smem_pipe_read_q);
|
||||
++smem_pipe_read_q;
|
||||
|
||||
Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout()));
|
||||
Tensor AbsMaxP = make_tensor_like<float>(
|
||||
make_layout(shape(group<1, 4>(flatten(tSrS_converion_view.layout()(make_coord(_0{}, _), _, _)))))
|
||||
);
|
||||
consumer_wait(pipeline_k, smem_pipe_read_k);
|
||||
copy_k_block(_0{});
|
||||
add_delta_s(tSrS);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tSrQ); ++k_block) {
|
||||
cute::gemm(tiled_mma_qk, make_zip_tensor(tSrQ(_, _, k_block), tSrSFQ(_, _, k_block)),
|
||||
make_zip_tensor(tSrK(_, _, k_block), tSrSFK(_, _, k_block)), tSrS);
|
||||
if (k_block < size<2>(tSrQ) - 1) {
|
||||
copy_k_block(k_block + 1);
|
||||
} else {
|
||||
pipeline_k.consumer_release(smem_pipe_read_k);
|
||||
++smem_pipe_read_k;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
auto col_limit_causal = [&](int row, int n_block) {
|
||||
return row + 1 + seqlen_k - n_block * kBlockN - seqlen_q + m_block * kBlockM;
|
||||
};
|
||||
{
|
||||
Tensor cS = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tScS = thread_mma_qk.partition_C(cS);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(tSrS); ++i) {
|
||||
if constexpr (!Is_causal) { // Just masking based on col
|
||||
if (int(get<1>(tScS(i))) >= int(unpadded_seqlen_k - n_block * kBlockN)) { tSrS(i) = -INFINITY; }
|
||||
} else {
|
||||
if (int(get<1>(tScS(i))) >= std::min(seqlen_k - n_block * kBlockN,
|
||||
col_limit_causal(int(get<0>(tScS(i))), n_block))) {
|
||||
tSrS(i) = -INFINITY;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
auto quantize = [&](auto mma_k, auto acc_conversion_view) {
|
||||
Tensor AbsMaxP_stagek = AbsMaxP(_, make_coord(_, _, mma_k));
|
||||
Tensor acc_conversion_stagek = acc_conversion_view(_, _, mma_k);
|
||||
Tensor SFP = make_tensor_like<cutlass::float_ue4m3_t>(AbsMaxP_stagek.layout());
|
||||
Tensor SFP_uint32_view = recast<uint32_t>(SFP);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(AbsMaxP_stagek); i += 4) {
|
||||
uint32_t& tmp = SFP_uint32_view(i / 4);
|
||||
flash::packed_float_to_ue4m3(
|
||||
AbsMaxP_stagek(i),
|
||||
AbsMaxP_stagek(i + 1),
|
||||
AbsMaxP_stagek(i + 2),
|
||||
AbsMaxP_stagek(i + 3),
|
||||
tmp
|
||||
);
|
||||
}
|
||||
int const quad_id = threadIdx.x & 3;
|
||||
uint32_t MASK = (0xFF00FF) << ((quad_id & 1) * 8);
|
||||
Tensor tOrSFP_uint32_view = recast<uint32_t>(tOrSFP(_, _, mma_k));
|
||||
Tensor tOrP_uint32_view = recast<uint32_t>(tOrP(_, _, mma_k));
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_m = 0; mma_m < size<1>(tOrP); ++mma_m) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
flash::packed_float_to_e2m1(
|
||||
acc_conversion_stagek(make_coord(_0{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_1{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_2{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_3{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_4{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_5{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_6{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_7{}, i), mma_m),
|
||||
tOrP_uint32_view(i, mma_m)
|
||||
);
|
||||
}
|
||||
uint32_t local_sfp = SFP_uint32_view(_0{}, _0{}, mma_m);
|
||||
uint32_t peer_sfp = __shfl_xor_sync(int32_t(-1), local_sfp, 2);
|
||||
if ((quad_id & 1) == 0) {
|
||||
uint32_t sfp = (local_sfp & MASK) | ((peer_sfp & MASK) << 8);
|
||||
tOrSFP_uint32_view(_0{}, mma_m) = sfp;
|
||||
} else {
|
||||
uint32_t sfp = (peer_sfp & MASK) | ((local_sfp & MASK) >> 8);
|
||||
tOrSFP_uint32_view(_0{}, mma_m) = sfp;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
softmax_fused.template online_softmax_with_quant</*Is_first=*/true>(tSrS, AbsMaxP, mainloop_params.softmax_scale_log2);
|
||||
|
||||
consumer_wait(pipeline_v, smem_pipe_read_v);
|
||||
copy_v_block(_0{});
|
||||
quantize(_0{}, tSrS_converion_view);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v_block = 0; v_block < size<2>(tOrP); ++v_block) {
|
||||
cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)),
|
||||
make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO_store);
|
||||
if (v_block < size<2>(tOrP) - 1) {
|
||||
copy_v_block(v_block + 1);
|
||||
quantize(v_block + 1, tSrS_converion_view);
|
||||
} else {
|
||||
pipeline_v.consumer_release(smem_pipe_read_v);
|
||||
++smem_pipe_read_v;
|
||||
}
|
||||
}
|
||||
|
||||
n_block--;
|
||||
constexpr int n_masking_steps = !Is_causal ? 1 : cute::ceil_div(kBlockM, kBlockN) + 1;
|
||||
// // Only go through these if Is_causal, since n_masking_steps = 1 when !Is_causal
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int masking_step = 0; masking_step < n_masking_steps - 1 && n_block >= 0; ++masking_step, --n_block) {
|
||||
Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout()));
|
||||
consumer_wait(pipeline_k, smem_pipe_read_k);
|
||||
copy_k_block(_0{});
|
||||
add_delta_s(tSrS);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tSrQ); ++k_block) {
|
||||
cute::gemm(tiled_mma_qk, make_zip_tensor(tSrQ(_, _, k_block), tSrSFQ(_, _, k_block)),
|
||||
make_zip_tensor(tSrK(_, _, k_block), tSrSFK(_, _, k_block)), tSrS);
|
||||
if (k_block < size<2>(tSrQ) - 1) {
|
||||
copy_k_block(k_block + 1);
|
||||
}
|
||||
}
|
||||
pipeline_k.consumer_release(smem_pipe_read_k); // release K
|
||||
++smem_pipe_read_k;
|
||||
Tensor cS = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tScS = thread_mma_qk.partition_C(cS);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < size(tSrS); ++i) {
|
||||
if (int(get<1>(tScS(i))) >= col_limit_causal(int(get<0>(tScS(i))), n_block)) {
|
||||
tSrS(i) = -INFINITY;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
softmax_fused.template online_softmax_with_quant</*Is_first=*/false>(tSrS, AbsMaxP, mainloop_params.softmax_scale_log2);
|
||||
Tensor tOrO = make_fragment_like(tOrO_store);
|
||||
consumer_wait(pipeline_v, smem_pipe_read_v);
|
||||
copy_v_block(_0{});
|
||||
quantize(_0{}, tSrS_converion_view);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v_block = 0; v_block < size<2>(tOrP); ++v_block) {
|
||||
cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)),
|
||||
make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO);
|
||||
if (v_block < size<2>(tOrP) - 1) {
|
||||
copy_v_block(v_block + 1);
|
||||
quantize(v_block + 1, tSrS_converion_view);
|
||||
}
|
||||
}
|
||||
pipeline_v.consumer_release(smem_pipe_read_v);
|
||||
++smem_pipe_read_v;
|
||||
if (masking_step > 0) { softmax_fused.rescale_o(tOrO_store, tOrO); }
|
||||
}
|
||||
|
||||
#pragma unroll 1
|
||||
for (; n_block >= 0; --n_block) {
|
||||
Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout()));
|
||||
consumer_wait(pipeline_k, smem_pipe_read_k);
|
||||
copy_k_block(_0{});
|
||||
add_delta_s(tSrS);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tSrQ); ++k_block) {
|
||||
cute::gemm(tiled_mma_qk, make_zip_tensor(tSrQ(_, _, k_block), tSrSFQ(_, _, k_block)),
|
||||
make_zip_tensor(tSrK(_, _, k_block), tSrSFK(_, _, k_block)), tSrS);
|
||||
if (k_block < size<2>(tSrQ) - 1) {
|
||||
copy_k_block(k_block + 1);
|
||||
} else {
|
||||
pipeline_k.consumer_release(smem_pipe_read_k);
|
||||
++smem_pipe_read_k;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
softmax_fused.template online_softmax_with_quant</*Is_first=*/false>(tSrS, AbsMaxP, mainloop_params.softmax_scale_log2);
|
||||
Tensor tOrO = make_fragment_like(tOrO_store);
|
||||
consumer_wait(pipeline_v, smem_pipe_read_v);
|
||||
copy_v_block(_0{});
|
||||
quantize(_0{}, tSrS_converion_view);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v_block = 0; v_block < size<2>(tOrP); ++v_block) {
|
||||
cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)),
|
||||
make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO);
|
||||
if (v_block < size<2>(tOrP) - 1) {
|
||||
copy_v_block(v_block + 1);
|
||||
quantize(v_block + 1, tSrS_converion_view);
|
||||
} else {
|
||||
pipeline_v.consumer_release(smem_pipe_read_v);
|
||||
++smem_pipe_read_v;
|
||||
}
|
||||
}
|
||||
softmax_fused.rescale_o(tOrO_store, tOrO);
|
||||
}
|
||||
softmax_fused.finalize(tOrO_store);
|
||||
return;
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // namespace flash
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/arch/barrier.h"
|
||||
#include "cutlass/pipeline/sm90_pipeline.hpp"
|
||||
|
||||
namespace flash {
|
||||
|
||||
enum class FP4NamedBarriers {
|
||||
QueryEmpty = 1,
|
||||
WarpSpecializedConsumer = 2,
|
||||
WarpSpecializedPingPongConsumer1 = 3,
|
||||
WarpSpecializedPingPongConsumer2 = 4,
|
||||
ProducerEnd = 5,
|
||||
ConsumerEnd = 6,
|
||||
EpilogueBarrier = 7
|
||||
};
|
||||
|
||||
template<int SequenceDepth, int SequenceLength>
|
||||
struct OrderedSequenceBarrierVarGroupSizeSharedStorage {
|
||||
using Barrier = cutlass::arch::ClusterBarrier;
|
||||
Barrier barrier_[SequenceDepth][SequenceLength];
|
||||
};
|
||||
|
||||
template<int SequenceDepth_, int SequenceLength_>
|
||||
class OrderedSequenceBarrierVarGroupSize {
|
||||
public:
|
||||
static constexpr int SequenceDepth = SequenceDepth_;
|
||||
static constexpr int SequenceLength = SequenceLength_;
|
||||
using Barrier = cutlass::arch::ClusterBarrier;
|
||||
using SharedStorage = flash::OrderedSequenceBarrierVarGroupSizeSharedStorage<SequenceDepth, SequenceLength>;
|
||||
|
||||
|
||||
struct Params {
|
||||
uint32_t group_id;
|
||||
uint32_t* group_size_list;
|
||||
};
|
||||
|
||||
private :
|
||||
// In future this Params object can be replaced easily with a CG object
|
||||
Params params_;
|
||||
Barrier *barrier_ptr_;
|
||||
cutlass::PipelineState<SequenceDepth> stage_;
|
||||
|
||||
static constexpr int Depth = SequenceDepth;
|
||||
static constexpr int Length = SequenceLength;
|
||||
|
||||
public:
|
||||
OrderedSequenceBarrierVarGroupSize() = delete;
|
||||
OrderedSequenceBarrierVarGroupSize(const OrderedSequenceBarrierVarGroupSize&) = delete;
|
||||
OrderedSequenceBarrierVarGroupSize(OrderedSequenceBarrierVarGroupSize&&) = delete;
|
||||
OrderedSequenceBarrierVarGroupSize& operator=(const OrderedSequenceBarrierVarGroupSize&) = delete;
|
||||
OrderedSequenceBarrierVarGroupSize& operator=(OrderedSequenceBarrierVarGroupSize&&) = delete;
|
||||
~OrderedSequenceBarrierVarGroupSize() = default;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
OrderedSequenceBarrierVarGroupSize(SharedStorage& storage, Params const& params) :
|
||||
params_(params),
|
||||
barrier_ptr_(&storage.barrier_[0][0]),
|
||||
// Group 0 - starts with an opposite phase
|
||||
stage_({0, params.group_id == 0, 0}) {
|
||||
int warp_idx = cutlass::canonical_warp_idx_sync();
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
// Barrier FULL, EMPTY init
|
||||
// Init is done only by the one elected thread of the block
|
||||
if (warp_idx == 0 && lane_predicate) {
|
||||
for (int d = 0; d < Depth; ++d) {
|
||||
for (int l = 0; l < Length; ++l) {
|
||||
barrier_ptr_[d * Length + l].init(*(params.group_size_list + l));
|
||||
}
|
||||
}
|
||||
}
|
||||
cutlass::arch::fence_barrier_init();
|
||||
}
|
||||
|
||||
// Wait on a stage to be unlocked
|
||||
CUTLASS_DEVICE
|
||||
void wait() {
|
||||
get_barrier_for_current_stage(params_.group_id).wait(stage_.phase());
|
||||
}
|
||||
|
||||
// Signal completion of Stage and move to the next stage
|
||||
// (group_id) signals to (group_id+1)
|
||||
CUTLASS_DEVICE
|
||||
void arrive() {
|
||||
int signalling_id = (params_.group_id + 1) % Length;
|
||||
get_barrier_for_current_stage(signalling_id).arrive();
|
||||
++stage_;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void advance() {
|
||||
++stage_;
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
Barrier& get_barrier_for_current_stage(int group_id) {
|
||||
return barrier_ptr_[stage_.index() * Length + group_id];
|
||||
}
|
||||
};
|
||||
|
||||
} // flash
|
||||
@@ -0,0 +1,180 @@
|
||||
// Modified from the original SageAttention3 code
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include <vector>
|
||||
|
||||
#ifdef OLD_GENERATOR_PATH
|
||||
#include <ATen/CUDAGeneratorImpl.h>
|
||||
#else
|
||||
#include <ATen/cuda/CUDAGeneratorImpl.h>
|
||||
#endif
|
||||
|
||||
#include <ATen/cuda/CUDAGraphsUtils.cuh> // For at::cuda::philox::unpack
|
||||
|
||||
#include "cutlass/fast_math.h" // For cutlass::FastDivmod
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct Qkv_params {
|
||||
using index_t = int64_t;
|
||||
// The QKV matrices.
|
||||
void *__restrict__ q_ptr;
|
||||
void *__restrict__ k_ptr;
|
||||
void *__restrict__ v_ptr;
|
||||
void *__restrict__ delta_s_ptr;
|
||||
// The QKV scale factor matrices.
|
||||
void *__restrict__ sfq_ptr;
|
||||
void *__restrict__ sfk_ptr;
|
||||
void *__restrict__ sfv_ptr;
|
||||
// The stride between rows of the Q, K and V matrices.
|
||||
index_t q_batch_stride;
|
||||
index_t k_batch_stride;
|
||||
index_t v_batch_stride;
|
||||
index_t q_row_stride;
|
||||
index_t k_row_stride;
|
||||
index_t v_row_stride;
|
||||
index_t q_head_stride;
|
||||
index_t k_head_stride;
|
||||
index_t v_head_stride;
|
||||
index_t ds_batch_stride;
|
||||
index_t ds_row_stride;
|
||||
index_t ds_head_stride;
|
||||
// The stride of the Q, K and V scale factor matrices.
|
||||
index_t sfq_batch_stride;
|
||||
index_t sfk_batch_stride;
|
||||
index_t sfv_batch_stride;
|
||||
index_t sfq_row_stride;
|
||||
index_t sfk_row_stride;
|
||||
index_t sfv_row_stride;
|
||||
index_t sfq_head_stride;
|
||||
index_t sfk_head_stride;
|
||||
index_t sfv_head_stride;
|
||||
|
||||
// The number of heads.
|
||||
int h, h_k;
|
||||
// In the case of multi-query and grouped-query attention (MQA/GQA), nheads_k could be
|
||||
// different from nheads (query).
|
||||
int h_h_k_ratio; // precompute h / h_k,
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct Flash_fwd_params : public Qkv_params {
|
||||
|
||||
// The O matrix (output).
|
||||
void * __restrict__ o_ptr;
|
||||
void * __restrict__ oaccum_ptr;
|
||||
void * __restrict__ s_ptr;
|
||||
|
||||
// The stride between rows of O.
|
||||
index_t o_batch_stride;
|
||||
index_t o_row_stride;
|
||||
index_t o_head_stride;
|
||||
|
||||
// The pointer to the P matrix.
|
||||
void * __restrict__ p_ptr;
|
||||
|
||||
// The pointer to the softmax sum.
|
||||
void * __restrict__ softmax_lse_ptr;
|
||||
void * __restrict__ softmax_lseaccum_ptr;
|
||||
|
||||
// The dimensions.
|
||||
int b, seqlen_q, seqlen_k, seqlen_knew, d, seqlen_q_rounded, seqlen_k_rounded, d_rounded, rotary_dim, unpadded_seqlen_k;
|
||||
cutlass::FastDivmod head_divmod, m_block_divmod;
|
||||
int total_blocks;
|
||||
int seqlen_s;
|
||||
|
||||
// The scaling factors for the kernel.
|
||||
float scale_softmax;
|
||||
float scale_softmax_log2;
|
||||
uint32_t scale_softmax_log2_half2;
|
||||
|
||||
// array of length b+1 holding starting offset of each sequence.
|
||||
int * __restrict__ cu_seqlens_q;
|
||||
int * __restrict__ cu_seqlens_k;
|
||||
|
||||
// If provided, the actual length of each k sequence.
|
||||
int * __restrict__ seqused_k;
|
||||
|
||||
int *__restrict__ blockmask;
|
||||
|
||||
// The K_new and V_new matrices.
|
||||
void * __restrict__ knew_ptr;
|
||||
void * __restrict__ vnew_ptr;
|
||||
|
||||
// The stride between rows of the Q, K and V matrices.
|
||||
index_t knew_batch_stride;
|
||||
index_t vnew_batch_stride;
|
||||
index_t knew_row_stride;
|
||||
index_t vnew_row_stride;
|
||||
index_t knew_head_stride;
|
||||
index_t vnew_head_stride;
|
||||
|
||||
// The cos and sin matrices for rotary embedding.
|
||||
void * __restrict__ rotary_cos_ptr;
|
||||
void * __restrict__ rotary_sin_ptr;
|
||||
|
||||
// The indices to index into the KV cache.
|
||||
int * __restrict__ cache_batch_idx;
|
||||
|
||||
// Paged KV cache
|
||||
int * __restrict__ block_table;
|
||||
index_t block_table_batch_stride;
|
||||
int page_block_size;
|
||||
|
||||
// The dropout probability (probability of keeping an activation).
|
||||
float p_dropout;
|
||||
// uint32_t p_dropout_in_uint;
|
||||
// uint16_t p_dropout_in_uint16_t;
|
||||
uint8_t p_dropout_in_uint8_t;
|
||||
|
||||
// Scale factor of 1 / (1 - p_dropout).
|
||||
float rp_dropout;
|
||||
float scale_softmax_rp_dropout;
|
||||
|
||||
// Local window size
|
||||
int window_size_left, window_size_right;
|
||||
|
||||
// Random state.
|
||||
at::PhiloxCudaState philox_args;
|
||||
|
||||
// Pointer to the RNG seed (idx 0) and offset (idx 1).
|
||||
uint64_t * rng_state;
|
||||
|
||||
bool is_bf16;
|
||||
bool is_e4m3;
|
||||
bool is_causal;
|
||||
bool per_block_mean;
|
||||
bool single_level_p_quant; // If true, use single-level 1x16 block scale quantization for P (like V), instead of two-level quantization
|
||||
// If is_seqlens_k_cumulative, then seqlen_k is cu_seqlens_k[bidb + 1] - cu_seqlens_k[bidb].
|
||||
// Otherwise it's cu_seqlens_k[bidb], i.e., we use cu_seqlens_k to store the sequence lengths of K.
|
||||
bool is_seqlens_k_cumulative;
|
||||
|
||||
bool is_rotary_interleaved;
|
||||
|
||||
int num_splits; // For split-KV version
|
||||
|
||||
void * __restrict__ alibi_slopes_ptr;
|
||||
index_t alibi_slopes_batch_stride;
|
||||
|
||||
int * __restrict__ tile_count_semaphore;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,190 @@
|
||||
// Modified from the original SageAttention3 code
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#pragma once
|
||||
|
||||
#include <cmath>
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "utils.h"
|
||||
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <int Rows>
|
||||
struct SoftmaxFused{
|
||||
|
||||
using TensorT = decltype(make_fragment_like<float>(Shape<Int<Rows>>{}));
|
||||
TensorT row_sum, row_max, scores_scale;
|
||||
static constexpr float fp8_scalexfp4_scale = 1.f / (448 * 6);
|
||||
static constexpr float fp8_scalexfp4_scale_log2 = -11.392317422778762f; //log2f(fp8_scalexfp4_scale)
|
||||
static constexpr float fp4_scale_log2 = -2.584962500721156f; // log2f(fp4_scale)
|
||||
static constexpr int RowReductionThr = 4;
|
||||
|
||||
// If true, use single-level quantization: s_P2, P̂_2 = φ(P̃) directly (standard per-block FP4 quantization like V)
|
||||
// If false (default), use two-level quantization: s_P1 = rowmax(P̃)/(448×6), then s_P2, P̂_2 = φ(P̃/s_P1)
|
||||
bool single_level_p_quant;
|
||||
|
||||
CUTLASS_DEVICE SoftmaxFused(bool single_level = false) : single_level_p_quant(single_level) {};
|
||||
|
||||
template<bool FirstTile, bool InfCheck = false, typename TensorAcc, typename TensorMax>
|
||||
CUTLASS_DEVICE auto online_softmax_with_quant(
|
||||
TensorAcc& acc,
|
||||
TensorMax& AbsMaxP,
|
||||
const float softmax_scale_log2
|
||||
) {
|
||||
Tensor acc_reduction_view = make_tensor(acc.data(), flash::convert_to_reduction_layout(acc.layout()));
|
||||
Tensor acc_conversion_view = make_tensor(acc.data(), flash::convert_to_conversion_layout(acc.layout()));
|
||||
Tensor acc_conversion_flatten = group_modes<1, 5>(group_modes<0, 2>(flatten(acc_conversion_view)));
|
||||
|
||||
if constexpr (FirstTile) {
|
||||
fill(row_max, -INFINITY);
|
||||
clear(row_sum);
|
||||
fill(scores_scale, 1.f);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size<0>(acc_reduction_view); mi++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1, 1>(acc_reduction_view); ni++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ei = 0; ei < size<1, 0>(acc_reduction_view); ei++) {
|
||||
AbsMaxP(mi, ni) = fmaxf(AbsMaxP(mi, ni), acc_reduction_view(mi, make_coord(ei, ni)));
|
||||
}
|
||||
float max_recv = __shfl_xor_sync(int32_t(-1), AbsMaxP(mi, ni), 1); // exchange max with neighbour thread of 8 elements
|
||||
AbsMaxP(mi, ni) = fmaxf(AbsMaxP(mi, ni), max_recv);
|
||||
row_max(mi) = fmaxf(row_max(mi), AbsMaxP(mi, ni));
|
||||
}
|
||||
|
||||
float max_recv = __shfl_xor_sync(int32_t(-1), row_max(mi), 2); // exchange max in a quad in a row
|
||||
row_max(mi) = fmaxf(row_max(mi), max_recv);
|
||||
|
||||
// Two-level P quantization (default): s_P1 = rowmax(P̃)/(448×6), then s_P2,P̂_2 = φ(P̃/s_P1)
|
||||
// - Pre-scales P to [0, 448×6] range before φ, output scaled by s_P1
|
||||
// Single-level P quantization: s_P2, P̂_2 = φ(P̃) directly (like V quantization)
|
||||
// - No s_P1, just standard per-block FP4 quantization φ
|
||||
const float s_P1_offset = single_level_p_quant ? 0.f : fp8_scalexfp4_scale_log2;
|
||||
const float max_scaled = InfCheck
|
||||
? (row_max(mi) == -INFINITY ? 0.f : (row_max(mi) * softmax_scale_log2 + s_P1_offset))
|
||||
: (row_max(mi) * softmax_scale_log2 + s_P1_offset);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(acc_reduction_view); ni++) {
|
||||
acc_reduction_view(mi, ni) = flash::ptx_exp2(acc_reduction_view(mi, ni) * softmax_scale_log2 - max_scaled);
|
||||
}
|
||||
// s_P2 = max(P_block)/6 — per-block scale factor from φ function (same formula for both modes)
|
||||
// The difference is in max_scaled: two-level includes 448×6 pre-scaling, single-level doesn't
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int sfi = 0; sfi < size<1>(AbsMaxP); sfi++) {
|
||||
AbsMaxP(mi, sfi) = flash::ptx_exp2(AbsMaxP(mi, sfi) * softmax_scale_log2 - max_scaled + fp4_scale_log2);
|
||||
}
|
||||
}
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size<0>(acc_reduction_view); mi++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(acc_reduction_view); ni++) {
|
||||
row_sum(mi) += acc_reduction_view(mi, ni);
|
||||
}
|
||||
}
|
||||
}
|
||||
else {
|
||||
Tensor scores_max_prev = make_fragment_like(row_max);
|
||||
cute::copy(row_max, scores_max_prev);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size<0>(acc_reduction_view); mi++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1, 1>(acc_reduction_view); ni++) {
|
||||
float local_max = -INFINITY;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ei = 0; ei < size<1, 0>(acc_reduction_view); ei++) {
|
||||
local_max = fmaxf(local_max, acc_reduction_view(mi, make_coord(ei, ni)));
|
||||
}
|
||||
float max_recv = __shfl_xor_sync(int32_t(-1), local_max, 1); // exchange max with neighbour thread of 8 elements
|
||||
AbsMaxP(mi, ni) = fmaxf(local_max, max_recv);
|
||||
row_max(mi) = fmaxf(row_max(mi), AbsMaxP(mi, ni));
|
||||
}
|
||||
|
||||
float max_recv = __shfl_xor_sync(int32_t(-1), row_max(mi), 2); // exchange max in a quad in a row
|
||||
row_max(mi) = fmaxf(row_max(mi), max_recv);
|
||||
|
||||
float scores_max_cur = !InfCheck
|
||||
? row_max(mi)
|
||||
: (row_max(mi) == -INFINITY ? 0.0f : row_max(mi));
|
||||
scores_scale(mi) = flash::ptx_exp2((scores_max_prev(mi) - scores_max_cur) * softmax_scale_log2);
|
||||
|
||||
// Two-level P quantization (default): s_P1 = rowmax(P̃)/(448×6), then s_P2,P̂_2 = φ(P̃/s_P1)
|
||||
// Single-level P quantization: s_P2, P̂_2 = φ(P̃) directly (like V quantization)
|
||||
const float s_P1_offset = single_level_p_quant ? 0.f : fp8_scalexfp4_scale_log2;
|
||||
const float max_scaled = InfCheck
|
||||
? (row_max(mi) == -INFINITY ? 0.f : (row_max(mi) * softmax_scale_log2 + s_P1_offset))
|
||||
: (row_max(mi) * softmax_scale_log2 + s_P1_offset);
|
||||
row_sum(mi) = row_sum(mi) * scores_scale(mi);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(acc_reduction_view); ni++) {
|
||||
acc_reduction_view(mi, ni) = flash::ptx_exp2(acc_reduction_view(mi, ni) * softmax_scale_log2 - max_scaled);
|
||||
row_sum(mi) += acc_reduction_view(mi, ni);
|
||||
}
|
||||
// s_P2 = max(P_block)/6 — per-block scale factor from φ function
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int sfi = 0; sfi < size<1>(AbsMaxP); sfi++) {
|
||||
AbsMaxP(mi, sfi) = flash::ptx_exp2(AbsMaxP(mi, sfi) * softmax_scale_log2 - max_scaled + fp4_scale_log2);
|
||||
}
|
||||
// scores_scale(mi) = max_scaled;
|
||||
}
|
||||
}
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(AbsMaxP); ++i) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < size<0>(acc_conversion_flatten); ++j)
|
||||
acc_conversion_flatten(j, i) /= AbsMaxP(i);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TensorAcc>
|
||||
CUTLASS_DEVICE void finalize(TensorAcc& o_store) {
|
||||
Tensor o_store_reduction_view = make_tensor(o_store.data(), flash::convert_to_reduction_layout(o_store.layout()));
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size(row_max); ++mi) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 1; i < RowReductionThr; i <<= 1) {
|
||||
float sum_recv = __shfl_xor_sync(int32_t(-1), row_sum(mi), i);
|
||||
row_sum(mi) += sum_recv;
|
||||
}
|
||||
float sum = row_sum(mi);
|
||||
float inv_sum = (sum == 0.f || sum != sum) ? 0.f : 1 / sum;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(o_store_reduction_view); ++ni) {
|
||||
o_store_reduction_view(mi, ni) *= inv_sum;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TensorAcc>
|
||||
CUTLASS_DEVICE void rescale_o(TensorAcc& o_store, TensorAcc const& o_tmp) {
|
||||
Tensor o_store_reduction_view = make_tensor(o_store.data(), flash::convert_to_reduction_layout(o_store.layout()));
|
||||
Tensor o_tmp_reduction_view = make_tensor(o_tmp.data(), flash::convert_to_reduction_layout(o_tmp.layout()));
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size(row_max); ++mi) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(o_store_reduction_view); ++ni) {
|
||||
o_store_reduction_view(mi, ni) = o_store_reduction_view(mi, ni) * scores_scale(mi) + o_tmp_reduction_view(mi, ni);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
|
||||
};
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,83 @@
|
||||
// Inspired by
|
||||
// https://github.com/NVIDIA/DALI/blob/main/include/dali/core/static_switch.h
|
||||
// and https://github.com/pytorch/pytorch/blob/master/aten/src/ATen/Dispatch.h
|
||||
|
||||
#pragma once
|
||||
|
||||
/// @param COND - a boolean expression to switch by
|
||||
/// @param CONST_NAME - a name given for the constexpr bool variable.
|
||||
/// @param ... - code to execute for true and false
|
||||
///
|
||||
/// Usage:
|
||||
/// ```
|
||||
/// BOOL_SWITCH(flag, BoolConst, [&] {
|
||||
/// some_function<BoolConst>(...);
|
||||
/// });
|
||||
/// ```
|
||||
//
|
||||
|
||||
#define BOOL_SWITCH(COND, CONST_NAME, ...) \
|
||||
[&] { \
|
||||
if (COND) { \
|
||||
constexpr static bool CONST_NAME = true; \
|
||||
return __VA_ARGS__(); \
|
||||
} else { \
|
||||
constexpr static bool CONST_NAME = false; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
}()
|
||||
|
||||
#define PREC_SWITCH(PRECTYPE, ...) \
|
||||
[&] { \
|
||||
if (PRECTYPE == 1) { \
|
||||
using kPrecType = cutlass::half_t; \
|
||||
constexpr static bool kSoftFp16 = false; \
|
||||
constexpr static bool kHybrid = false; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (PRECTYPE == 2) { \
|
||||
using kPrecType = cutlass::float_e4m3_t; \
|
||||
constexpr static bool kSoftFp16 = false; \
|
||||
constexpr static bool kHybrid = false; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (PRECTYPE == 3) { \
|
||||
using kPrecType = cutlass::float_e4m3_t; \
|
||||
constexpr static bool kSoftFp16 = false; \
|
||||
constexpr static bool kHybrid = true; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (PRECTYPE == 4) { \
|
||||
using kPrecType = cutlass::float_e4m3_t; \
|
||||
constexpr static bool kSoftFp16 = true; \
|
||||
constexpr static bool kHybrid = false; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
}()
|
||||
|
||||
#define HEADDIM_SWITCH(HEADDIM, ...) \
|
||||
[&] { \
|
||||
if (HEADDIM == 64) { \
|
||||
constexpr static int kHeadSize = 64; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (HEADDIM == 128) { \
|
||||
constexpr static int kHeadSize = 128; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (HEADDIM == 256) { \
|
||||
constexpr static int kHeadSize = 256; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
}()
|
||||
|
||||
#define SEQLEN_SWITCH(USE_VAR_SEQ_LEN, SEQ_LEN_OUT_OF_BOUND_CHECK, ...) \
|
||||
[&] { \
|
||||
if (!USE_VAR_SEQ_LEN) { \
|
||||
if (SEQ_LEN_OUT_OF_BOUND_CHECK) { \
|
||||
using kSeqLenTraitsType = FixedSeqLenTraits<true>; \
|
||||
return __VA_ARGS__(); \
|
||||
} else { \
|
||||
using kSeqLenTraitsType = FixedSeqLenTraits<false>; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
} else { \
|
||||
using kSeqLenTraitsType = VarSeqLenTraits; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
}()
|
||||
@@ -0,0 +1,304 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* This code is based on code from FlashAttention3, https://github.com/Dao-AILab/flash-attention
|
||||
* Copyright (c) 2024, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/fast_math.h"
|
||||
|
||||
namespace flash {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
class StaticPersistentTileSchedulerOld {
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
private:
|
||||
int current_work_linear_idx_;
|
||||
cutlass::FastDivmod const &m_block_divmod, &head_divmod;
|
||||
int const total_blocks;
|
||||
|
||||
public:
|
||||
struct WorkTileInfo {
|
||||
int M_idx = 0;
|
||||
int H_idx = 0;
|
||||
int B_idx = 0;
|
||||
bool is_valid_tile = false;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
is_valid() const {
|
||||
return is_valid_tile;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static WorkTileInfo
|
||||
invalid_work_tile() {
|
||||
return {-1, -1, -1, false};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE explicit StaticPersistentTileSchedulerOld(cutlass::FastDivmod const &m_block_divmod_,
|
||||
cutlass::FastDivmod const &head_divmod_,
|
||||
int const total_blocks_) :
|
||||
m_block_divmod(m_block_divmod_), head_divmod(head_divmod_), total_blocks(total_blocks_) {
|
||||
|
||||
// MSVC requires protecting use of CUDA-specific nonstandard syntax,
|
||||
// like blockIdx and gridDim, with __CUDA_ARCH__.
|
||||
#if defined(__CUDA_ARCH__)
|
||||
// current_work_linear_idx_ = blockIdx.x + blockIdx.y * gridDim.x + blockIdx.z * gridDim.x * gridDim.y;
|
||||
current_work_linear_idx_ = blockIdx.x;
|
||||
#else
|
||||
CUTLASS_ASSERT(false && "This line should never be reached");
|
||||
#endif
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_current_work() const {
|
||||
return get_current_work_for_linear_idx(current_work_linear_idx_);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_current_work_for_linear_idx(int linear_idx) const {
|
||||
if (linear_idx >= total_blocks) {
|
||||
return WorkTileInfo::invalid_work_tile();
|
||||
}
|
||||
|
||||
// Map worker's linear index into the CTA tiled problem shape to the corresponding MHB indices
|
||||
int M_idx, H_idx, B_idx;
|
||||
int quotient = m_block_divmod.divmod(M_idx, linear_idx);
|
||||
B_idx = head_divmod.divmod(H_idx, quotient);
|
||||
return {M_idx, H_idx, B_idx, true};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
// advance_to_next_work(int advance_count = 1) {
|
||||
advance_to_next_work() {
|
||||
// current_work_linear_idx_ += int(gridDim.x * gridDim.y * gridDim.z);
|
||||
current_work_linear_idx_ += int(gridDim.x);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
fetch_next_work() {
|
||||
WorkTileInfo new_work_tile_info;
|
||||
advance_to_next_work();
|
||||
new_work_tile_info = get_current_work();
|
||||
return new_work_tile_info;
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
class SingleTileScheduler {
|
||||
|
||||
public:
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
int const num_blocks_m, num_head, num_batch;
|
||||
int const* tile_count_semaphore = nullptr;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {};
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
return {};
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_grid_dim(Arguments const& args, int num_sm) {
|
||||
return {uint32_t(args.num_blocks_m), uint32_t(args.num_head), uint32_t(args.num_batch)};
|
||||
}
|
||||
|
||||
struct WorkTileInfo {
|
||||
int M_idx = 0;
|
||||
int H_idx = 0;
|
||||
int B_idx = 0;
|
||||
bool is_valid_tile = false;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool
|
||||
is_valid(Params const& params) const {
|
||||
return is_valid_tile;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
cute::tuple<int32_t, int32_t, int32_t>
|
||||
get_block_coord(Params const& params) const {
|
||||
return {M_idx, H_idx, B_idx};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_next_work(Params const& params) const {
|
||||
return {-1, -1, -1, false};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_initial_work() const {
|
||||
return {int(blockIdx.x), int(blockIdx.y), int(blockIdx.z), true};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_next_work(Params const& params, WorkTileInfo const& current_work) const {
|
||||
return {-1, -1, -1, false};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
class StaticPersistentTileScheduler {
|
||||
|
||||
public:
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
int const num_blocks_m, num_head, num_batch;
|
||||
int const* tile_count_semaphore = nullptr;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
int total_blocks;
|
||||
cutlass::FastDivmod m_block_divmod, head_divmod;
|
||||
};
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
return {args.num_blocks_m * args.num_head * args.num_batch,
|
||||
cutlass::FastDivmod(args.num_blocks_m), cutlass::FastDivmod(args.num_head)};
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_grid_dim(Arguments const& args, int num_sm) {
|
||||
return {uint32_t(num_sm)};
|
||||
}
|
||||
|
||||
struct WorkTileInfo {
|
||||
int tile_idx;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool
|
||||
is_valid(Params const& params) const {
|
||||
return tile_idx < params.total_blocks;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
cute::tuple<int32_t, int32_t, int32_t>
|
||||
get_block_coord(Params const& params) const {
|
||||
int m_block, bidh, bidb;
|
||||
bidb = params.head_divmod.divmod(bidh, params.m_block_divmod.divmod(m_block, tile_idx));
|
||||
return {m_block, bidh, bidb};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_initial_work() const {
|
||||
return {int(blockIdx.x)};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_next_work(Params const& params, WorkTileInfo const& current_work) const {
|
||||
return {current_work.tile_idx + int(gridDim.x)};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
class DynamicPersistentTileScheduler {
|
||||
|
||||
public:
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
int const num_blocks_m, num_head, num_batch;
|
||||
int const* tile_count_semaphore;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
int const total_blocks;
|
||||
cutlass::FastDivmod const m_block_divmod, head_divmod;
|
||||
int const* tile_count_semaphore;
|
||||
};
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
return {args.num_blocks_m * args.num_head * args.num_batch,
|
||||
cutlass::FastDivmod(args.num_blocks_m), cutlass::FastDivmod(args.num_head),
|
||||
args.tile_count_semaphore};
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_grid_dim(Arguments const& args, int num_sm) {
|
||||
return {uint32_t(num_sm)};
|
||||
}
|
||||
|
||||
using WorkTileInfo = StaticPersistentTileScheduler::WorkTileInfo;
|
||||
// struct WorkTileInfo {
|
||||
// int tile_idx;
|
||||
|
||||
// CUTLASS_DEVICE
|
||||
// bool
|
||||
// is_valid(Params const& params) const {
|
||||
// return tile_idx < params.total_blocks;
|
||||
// }
|
||||
|
||||
// CUTLASS_DEVICE
|
||||
// cute::tuple<int32_t, int32_t, int32_t>
|
||||
// get_block_coord(Params const& params) const {
|
||||
// int m_block, bidh, bidb;
|
||||
// bidb = params.head_divmod.divmod(bidh, params.m_block_divmod.divmod(m_block, tile_idx));
|
||||
// return {m_block, bidh, bidb};
|
||||
// }
|
||||
|
||||
// };
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_initial_work() const {
|
||||
return {int(blockIdx.x)};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_next_work(Params const& params, WorkTileInfo const& current_work) const {
|
||||
return {current_work.tile_idx + int(gridDim.x)};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // flash
|
||||
@@ -0,0 +1,408 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <assert.h>
|
||||
#include <stdint.h>
|
||||
#include <stdlib.h>
|
||||
|
||||
#include <cuda_fp16.h>
|
||||
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
|
||||
#include <cuda_bf16.h>
|
||||
#endif
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
|
||||
#include <cutlass/array.h>
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/numeric_conversion.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<typename T>
|
||||
struct MaxOp {
|
||||
__device__ __forceinline__ T operator()(T const & x, T const & y) { return x > y ? x : y; }
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MaxOp<float> {
|
||||
// This is slightly faster
|
||||
__device__ __forceinline__ float operator()(float const &x, float const &y) { return max(x, y); }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<typename T>
|
||||
struct SumOp {
|
||||
__device__ __forceinline__ T operator()(T const & x, T const & y) { return x + y; }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<int THREADS>
|
||||
struct Allreduce {
|
||||
static_assert(THREADS == 32 || THREADS == 16 || THREADS == 8 || THREADS == 4);
|
||||
template<typename T, typename Operator>
|
||||
static __device__ __forceinline__ T run(T x, Operator &op) {
|
||||
constexpr int OFFSET = THREADS / 2;
|
||||
x = op(x, __shfl_xor_sync(uint32_t(-1), x, OFFSET));
|
||||
return Allreduce<OFFSET>::run(x, op);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<>
|
||||
struct Allreduce<2> {
|
||||
template<typename T, typename Operator>
|
||||
static __device__ __forceinline__ T run(T x, Operator &op) {
|
||||
x = op(x, __shfl_xor_sync(uint32_t(-1), x, 1));
|
||||
return x;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<bool zero_init=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1, typename Operator>
|
||||
__device__ __forceinline__ void thread_reduce_(Tensor<Engine0, Layout0> const &tensor, Tensor<Engine1, Layout1> &summary, Operator &op) {
|
||||
static_assert(Layout0::rank == 2, "Only support 2D Tensor");
|
||||
static_assert(Layout1::rank == 1, "Only support 1D Tensor");
|
||||
CUTE_STATIC_ASSERT_V(size<0>(summary) == size<0>(tensor));
|
||||
#pragma unroll
|
||||
for (int mi = 0; mi < size<0>(tensor); mi++) {
|
||||
summary(mi) = zero_init ? tensor(mi, 0) : op(summary(mi), tensor(mi, 0));
|
||||
#pragma unroll
|
||||
for (int ni = 1; ni < size<1>(tensor); ni++) {
|
||||
summary(mi) = op(summary(mi), tensor(mi, ni));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<typename Engine0, typename Layout0, typename Engine1, typename Layout1, typename Operator>
|
||||
__device__ __forceinline__ void quad_allreduce_(Tensor<Engine0, Layout0> &dst, Tensor<Engine1, Layout1> &src, Operator &op) {
|
||||
CUTE_STATIC_ASSERT_V(size(dst) == size(src));
|
||||
#pragma unroll
|
||||
for (int i = 0; i < size(dst); i++){
|
||||
dst(i) = Allreduce<4>::run(src(i), op);
|
||||
}
|
||||
}
|
||||
|
||||
template<bool zero_init=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1, typename Operator>
|
||||
__device__ __forceinline__ void reduce_(Tensor<Engine0, Layout0> const& tensor, Tensor<Engine1, Layout1> &summary, Operator &op) {
|
||||
thread_reduce_<zero_init>(tensor, summary, op);
|
||||
quad_allreduce_(summary, summary, op);
|
||||
}
|
||||
|
||||
template<bool zero_init=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1>
|
||||
__device__ __forceinline__ void reduce_max(Tensor<Engine0, Layout0> const& tensor, Tensor<Engine1, Layout1> &max){
|
||||
MaxOp<float> max_op;
|
||||
reduce_<zero_init>(tensor, max, max_op);
|
||||
}
|
||||
|
||||
template<bool zero_init=true, bool warp_reduce=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1>
|
||||
__device__ __forceinline__ void reduce_sum(Tensor<Engine0, Layout0> const& tensor, Tensor<Engine1, Layout1> &sum){
|
||||
SumOp<float> sum_op;
|
||||
thread_reduce_<zero_init>(tensor, sum, sum_op);
|
||||
if constexpr (warp_reduce) { quad_allreduce_(sum, sum, sum_op); }
|
||||
}
|
||||
|
||||
__forceinline__ __device__ __half2 half_exp(__half2 x) {
|
||||
uint32_t tmp_out, tmp_in;
|
||||
tmp_in = reinterpret_cast<uint32_t&>(x);
|
||||
asm ("ex2.approx.f16x2 %0, %1;\n"
|
||||
: "=r"(tmp_out)
|
||||
: "r"(tmp_in));
|
||||
__half2 out = reinterpret_cast<__half2&>(tmp_out);
|
||||
return out;
|
||||
}
|
||||
|
||||
// Apply the exp to all the elements.
|
||||
template <bool zero_init=false, typename Engine0, typename Layout0, typename Engine1, typename Layout1>
|
||||
__forceinline__ __device__ void max_scale_exp2_sum(Tensor<Engine0, Layout0> &tensor, Tensor<Engine1, Layout1> &max, Tensor<Engine1, Layout1> &sum, const float scale) {
|
||||
static_assert(Layout0::rank == 2, "Only support 2D Tensor"); static_assert(Layout1::rank == 1, "Only support 1D Tensor"); CUTE_STATIC_ASSERT_V(size<0>(max) == size<0>(tensor));
|
||||
#pragma unroll
|
||||
for (int mi = 0; mi < size<0>(tensor); ++mi) {
|
||||
MaxOp<float> max_op;
|
||||
max(mi) = zero_init ? tensor(mi, 0) : max_op(max(mi), tensor(mi, 0));
|
||||
#pragma unroll
|
||||
for (int ni = 1; ni < size<1>(tensor); ni++) {
|
||||
max(mi) = max_op(max(mi), tensor(mi, ni));
|
||||
}
|
||||
max(mi) = Allreduce<4>::run(max(mi), max_op);
|
||||
// If max is -inf, then all elements must have been -inf (possibly due to masking).
|
||||
// We don't want (-inf - (-inf)) since that would give NaN.
|
||||
const float max_scaled = max(mi) == -INFINITY ? 0.f : max(mi) * scale;
|
||||
sum(mi) = 0;
|
||||
#pragma unroll
|
||||
for (int ni = 0; ni < size<1>(tensor); ++ni) {
|
||||
// Instead of computing exp(x - max), we compute exp2(x * log_2(e) -
|
||||
// max * log_2(e)) This allows the compiler to use the ffma
|
||||
// instruction instead of fadd and fmul separately.
|
||||
tensor(mi, ni) = exp2f(tensor(mi, ni) * scale - max_scaled);
|
||||
sum(mi) += tensor(mi, ni);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Apply the exp to all the elements.
|
||||
template <bool Scale_max=true, bool Check_inf=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1>
|
||||
__forceinline__ __device__ void scale_apply_exp2(Tensor<Engine0, Layout0> &tensor, Tensor<Engine1, Layout1> const &max, const float scale) {
|
||||
static_assert(Layout0::rank == 2, "Only support 2D Tensor");
|
||||
static_assert(Layout1::rank == 1, "Only support 1D Tensor");
|
||||
CUTE_STATIC_ASSERT_V(size<0>(max) == size<0>(tensor));
|
||||
#pragma unroll
|
||||
for (int mi = 0; mi < size<0>(tensor); ++mi) {
|
||||
// If max is -inf, then all elements must have been -inf (possibly due to masking).
|
||||
// We don't want (-inf - (-inf)) since that would give NaN.
|
||||
// If we don't have float around M_LOG2E the multiplication is done in fp64.
|
||||
const float max_scaled = Check_inf
|
||||
? (max(mi) == -INFINITY ? 0.f : (max(mi) * (Scale_max ? scale : float(M_LOG2E))))
|
||||
: (max(mi) * (Scale_max ? scale : float(M_LOG2E)));
|
||||
#pragma unroll
|
||||
for (int ni = 0; ni < size<1>(tensor); ++ni) {
|
||||
// Instead of computing exp(x - max), we compute exp2(x * log_2(e) -
|
||||
// max * log_2(e)) This allows the compiler to use the ffma
|
||||
// instruction instead of fadd and fmul separately.
|
||||
tensor(mi, ni) = exp2f(tensor(mi, ni) * scale - max_scaled);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
__forceinline__ __device__ float ptx_exp2(float x) {
|
||||
float y;
|
||||
asm volatile("ex2.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
|
||||
return y;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
packed_float_to_ue4m3(
|
||||
float const &f0, float const &f1, float const &f2, float const &f3,
|
||||
uint32_t &out
|
||||
) {
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
".reg .b16 lo;\n" \
|
||||
".reg .b16 hi;\n" \
|
||||
"cvt.rn.satfinite.e4m3x2.f32 lo, %2, %1;\n" \
|
||||
"cvt.rn.satfinite.e4m3x2.f32 hi, %4, %3;\n" \
|
||||
"mov.b32 %0, {lo, hi};\n" \
|
||||
"}" \
|
||||
: "=r"(out) : "f"(f0), "f"(f1), "f"(f2), "f"(f3));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
packed_float_to_e2m1(
|
||||
float const &f0, float const &f1, float const &f2, float const& f3,
|
||||
float const &f4, float const &f5, float const &f6, float const& f7,
|
||||
uint32_t &out
|
||||
) {
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
".reg .b8 byte0;\n" \
|
||||
".reg .b8 byte1;\n" \
|
||||
".reg .b8 byte2;\n" \
|
||||
".reg .b8 byte3;\n" \
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte0, %2, %1;\n" \
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte1, %4, %3;\n" \
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte2, %6, %5;\n" \
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte3, %8, %7;\n" \
|
||||
"mov.b32 %0, {byte0, byte1, byte2, byte3};\n" \
|
||||
"}" \
|
||||
: "=r"(out) : "f"(f0), "f"(f1), "f"(f2), "f"(f3),
|
||||
"f"(f4), "f"(f5), "f"(f6), "f"(f7));
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
add(float2 & c,
|
||||
float2 const& a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("add.f32x2 %0, %1, %2;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t &>(c))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
add_inplace(float2 &a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("add.f32x2 %0, %0, %1;\n"
|
||||
: "+l"(reinterpret_cast<uint64_t &>(a)) // a: input/output
|
||||
: "l"(reinterpret_cast<uint64_t const&>(b)) // b: input
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
sub(float2 & c,
|
||||
float2 const& a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("sub.f32x2 %0, %1, %2;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t &>(c))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
sub_inplace(float2 &a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("sub.f32x2 %0, %0, %1;\n"
|
||||
: "+l"(reinterpret_cast<uint64_t &>(a)) // a: input/output
|
||||
: "l"(reinterpret_cast<uint64_t const&>(b)) // b: input
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
mul(float2 & c,
|
||||
float2 const& a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("mul.f32x2 %0, %1, %2;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t &>(c))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
fma(float2 & d,
|
||||
float2 const& a,
|
||||
float2 const& b,
|
||||
float2 const& c)
|
||||
{
|
||||
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t &>(d))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(c)));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
fma_inplace(float2 &a,
|
||||
float2 const& b,
|
||||
float2 const& c)
|
||||
{
|
||||
asm volatile("fma.rn.f32x2 %0, %0, %1, %2;\n"
|
||||
: "+l"(reinterpret_cast<uint64_t &>(a))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(b)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(c)));
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class Layout
|
||||
>
|
||||
CUTLASS_DEVICE constexpr
|
||||
auto convert_to_reduction_layout(Layout mma_layout) {
|
||||
static_assert(rank(mma_layout) == 3, "Mma Layout should be (MmaAtom, MmaM, MmaN)");
|
||||
static_assert(rank(get<0>(shape(mma_layout))) == 2, "MmaAtom should be (AtomN, AtomM)");
|
||||
|
||||
return make_layout(
|
||||
make_layout(get<0,1>(mma_layout), get<1>(mma_layout)),
|
||||
make_layout(get<0,0>(mma_layout), get<2>(mma_layout))
|
||||
);
|
||||
}
|
||||
|
||||
template <
|
||||
class Tensor
|
||||
>
|
||||
CUTLASS_DEVICE constexpr
|
||||
auto convert_to_reduction_tensor(Tensor mma_tensor) {
|
||||
return make_tensor(mma_tensor.data(), convert_to_reduction_layout(mma_tensor.layout()));
|
||||
}
|
||||
|
||||
|
||||
template <
|
||||
class Layout
|
||||
>
|
||||
CUTLASS_DEVICE constexpr
|
||||
auto convert_to_conversion_layout(Layout mma_layout) {
|
||||
static_assert(rank(mma_layout) == 3, "Mma Layout should be (MmaAtom, MmaM, MmaN)");
|
||||
static_assert(rank(get<0>(shape(mma_layout))) == 2, "MmaAtom should be (AtomN, AtomM)");
|
||||
|
||||
constexpr int MmaAtomN = size<0, 0>(mma_layout);
|
||||
constexpr int MmaAtomM = size<0, 1>(mma_layout);
|
||||
constexpr int MmaM = size<1>(mma_layout);
|
||||
constexpr int MmaN = size<2>(mma_layout);
|
||||
|
||||
static_assert(MmaAtomN == 8, "MmaAtomN should be 8.");
|
||||
static_assert(MmaAtomM == 2, "MmaAtomM should be 2.");
|
||||
static_assert(MmaN % 2 == 0, "MmaN should be multiple of 2.");
|
||||
|
||||
auto mma_n_division = zipped_divide(
|
||||
layout<2>(mma_layout), make_tile(_2{})
|
||||
);
|
||||
return make_layout(
|
||||
make_layout(layout<0,0>(mma_layout), make_layout(layout<0,1>(mma_layout), layout<0>(mma_n_division))),
|
||||
layout<1>(mma_layout), layout<1>(mma_n_division)
|
||||
);
|
||||
}
|
||||
|
||||
template <
|
||||
class Tensor
|
||||
>
|
||||
CUTLASS_DEVICE constexpr
|
||||
auto convert_to_conversion_tensor(Tensor mma_tensor) {
|
||||
return make_tensor(mma_tensor.data(), convert_to_conversion_layout(mma_tensor.layout()));
|
||||
}
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <bool Is_even_MN=true, bool Is_even_K=true, bool Clear_OOB_MN=false, bool Clear_OOB_K=true,
|
||||
typename TiledCopy, typename Engine0, typename Layout0, typename Engine1, typename Layout1,
|
||||
typename Engine2, typename Layout2, typename Engine3, typename Layout3>
|
||||
CUTLASS_DEVICE void copy(TiledCopy tiled_copy, Tensor<Engine0, Layout0> const &S,
|
||||
Tensor<Engine1, Layout1> &D, Tensor<Engine2, Layout2> const &identity_MN,
|
||||
Tensor<Engine3, Layout3> const &predicate_K, const int max_MN=0) {
|
||||
CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{});
|
||||
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA
|
||||
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M
|
||||
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K
|
||||
// There's no case where !Clear_OOB_K && Clear_OOB_MN
|
||||
static_assert(!(Clear_OOB_MN && !Clear_OOB_K));
|
||||
#pragma unroll
|
||||
for (int m = 0; m < size<1>(S); ++m) {
|
||||
if (Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN) {
|
||||
#pragma unroll
|
||||
for (int k = 0; k < size<2>(S); ++k) {
|
||||
if (Is_even_K || predicate_K(k)) {
|
||||
cute::copy(tiled_copy, S(_, m, k), D(_, m, k));
|
||||
} else if (Clear_OOB_K) {
|
||||
cute::clear(D(_, m, k));
|
||||
}
|
||||
}
|
||||
} else if (Clear_OOB_MN) {
|
||||
cute::clear(D(_, m, _));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1 @@
|
||||
__version__ = "3.0.0.b1"
|
||||
@@ -0,0 +1,90 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import torch
|
||||
import fp4quant
|
||||
from triton.tools.mxfp import MXFP4Tensor
|
||||
|
||||
from bench_utils import bench_kineto
|
||||
b = 1
|
||||
h = 32
|
||||
n = 16384
|
||||
d = 128
|
||||
|
||||
def test():
|
||||
q = torch.randn((b, h, n, d), device="cuda", dtype=torch.float16)
|
||||
o = torch.empty((b, h, n, d // 2), device="cuda", dtype=torch.uint8)
|
||||
o_s = torch.empty((b, h, n, d // 16), device="cuda", dtype=torch.float8_e4m3fn)
|
||||
fp4quant.scaled_fp4_quant_permute(q, o, o_s, 1)
|
||||
|
||||
test()
|
||||
|
||||
t = bench_kineto(test, "scaled_fp4_quant_kernel", suppress_kineto_output=True)
|
||||
|
||||
IO = b * h * n * d * 2 + b * h * n * d * 0.5 + b * h * n * d // 16 * 1
|
||||
throughput = IO / t * 1e-9
|
||||
|
||||
print(f"Throughput: {throughput:.2f} GB/s")
|
||||
|
||||
def scale_and_fp4_tensor(x: torch.Tensor, packed_dim: int = 3, all_ones: bool = False, permuted: bool = False):
|
||||
assert x.is_contiguous() and x.ndim == 4 and x.shape[-1] % 16 == 0
|
||||
B, H, M, N = x.shape
|
||||
x = x.view(B, H, M, N // 16, 16)
|
||||
scales = (x.abs().amax(dim=-1, keepdim=True) / 6).to(torch.float32)
|
||||
if all_ones:
|
||||
scales = torch.ones_like(scales)
|
||||
x_scaled = x / scales
|
||||
packed_fp4 = MXFP4Tensor(x_scaled.flatten(start_dim=-2)).to_packed_tensor(dim=packed_dim)
|
||||
dequant_x = (MXFP4Tensor(x_scaled).to(torch.float32) * scales.to(torch.float8_e4m3fn).to(torch.float32)).flatten(start_dim=-2)
|
||||
fp8_scale = scales.flatten(start_dim=-2).to(torch.float8_e4m3fn)
|
||||
permuted_fp8_scale = None
|
||||
if permuted:
|
||||
scales = scales.view(B, H // 64, 4, 16, M, N // 16).permute(0, 1, 3, 2, 4, 5).reshape(B, H, M, N // 16)
|
||||
permuted_fp8_scale = scales.view(B, H // 64, 64, M, N // 64, 4).permute(0, 1, 4, 3, 2, 5).reshape(B, H, M, N // 16).to(torch.float8_e4m3fn)
|
||||
return fp8_scale, packed_fp4, dequant_x, permuted_fp8_scale
|
||||
|
||||
b = 2
|
||||
h = 4
|
||||
n = 251
|
||||
n_padded = (n + 127) // 128 * 128
|
||||
d = 128
|
||||
|
||||
q = torch.randn(b, h, n, d, dtype=torch.float16, device='cuda')
|
||||
o = torch.empty((b, h, n, d // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s = torch.empty((b, h, n, d // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
|
||||
fp4quant.scaled_fp4_quant(q, o, o_s, 1)
|
||||
|
||||
k_permute = [0, 1, 8, 9, 16, 17, 24, 25, 2, 3, 10, 11, 18, 19, 26, 27, 4, 5, 12, 13, 20, 21, 28, 29, 6, 7, 14, 15, 22, 23, 30, 31]
|
||||
o_permuted = torch.empty((b, h, n_padded, d // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s_permuted = torch.empty((b, h, n_padded, d // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
fp4quant.scaled_fp4_quant_permute(q, o_permuted, o_s_permuted, 1)
|
||||
|
||||
# padding
|
||||
if n % 128 != 0:
|
||||
o_permuted_gt = torch.cat([o, torch.zeros((b, h, n_padded - n, d // 2), dtype=torch.uint8, device='cuda')], dim=2)
|
||||
o_s_permuted_gt = torch.cat([o_s, torch.zeros((b, h, n_padded - n, d // 16), dtype=torch.float8_e4m3fn, device='cuda')], dim=2)
|
||||
else:
|
||||
o_permuted_gt = o
|
||||
o_s_permuted_gt = o_s
|
||||
|
||||
# use scale_and_fp4_tensor + torch permutation to get the ground truth
|
||||
o_permuted_gt = o_permuted_gt.reshape(b, h, n_padded // 32, 32, d // 2)[:, :, :, k_permute, :].reshape(b, h, n_padded, d // 2)
|
||||
o_s_permuted_gt = o_s_permuted_gt.reshape(b, h, n_padded // 32, 32, d // 16)[:, :, :, k_permute, :].reshape(b, h, n_padded, d // 16)
|
||||
|
||||
assert((o_permuted - o_permuted_gt).abs().max() == 0)
|
||||
assert((o_s_permuted.float() - o_s_permuted_gt.float()).abs().max() == 0)
|
||||
|
||||
print("All tests passed!")
|
||||
@@ -0,0 +1,86 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import torch
|
||||
import fp4quant
|
||||
from triton.tools.mxfp import MXFP4Tensor
|
||||
|
||||
from bench_utils import bench_kineto
|
||||
b = 1
|
||||
h = 32
|
||||
n = 16384
|
||||
d = 128
|
||||
|
||||
def test():
|
||||
q = torch.randn((b, h, n, d), device="cuda", dtype=torch.float16)
|
||||
o = torch.empty((b, h, n, d // 2), device="cuda", dtype=torch.uint8)
|
||||
o_s = torch.empty((b, h, n, d // 16), device="cuda", dtype=torch.float8_e4m3fn)
|
||||
fp4quant.scaled_fp4_quant(q, o, o_s, 1)
|
||||
|
||||
test()
|
||||
|
||||
t = bench_kineto(test, "scaled_fp4_quant_kernel", suppress_kineto_output=True)
|
||||
|
||||
IO = b * h * n * d * 2 + b * h * n * d * 0.5 + b * h * n * d // 16 * 1
|
||||
throughput = IO / t * 1e-9
|
||||
|
||||
print(f"Throughput: {throughput:.2f} GB/s")
|
||||
|
||||
def scale_and_fp4_tensor(x: torch.Tensor, packed_dim: int = 3, all_ones: bool = False, permuted: bool = False):
|
||||
assert x.is_contiguous() and x.ndim == 4 and x.shape[-1] % 16 == 0
|
||||
B, H, M, N = x.shape
|
||||
x = x.view(B, H, M, N // 16, 16)
|
||||
scales = (x.abs().amax(dim=-1, keepdim=True) / 6).to(torch.float32)
|
||||
if all_ones:
|
||||
scales = torch.ones_like(scales)
|
||||
x_scaled = x / scales
|
||||
packed_fp4 = MXFP4Tensor(x_scaled.flatten(start_dim=-2)).to_packed_tensor(dim=packed_dim)
|
||||
dequant_x = (MXFP4Tensor(x_scaled).to(torch.float32) * scales.to(torch.float8_e4m3fn).to(torch.float32)).flatten(start_dim=-2)
|
||||
fp8_scale = scales.flatten(start_dim=-2).to(torch.float8_e4m3fn)
|
||||
permuted_fp8_scale = None
|
||||
if permuted:
|
||||
scales = scales.view(B, H // 64, 4, 16, M, N // 16).permute(0, 1, 3, 2, 4, 5).reshape(B, H, M, N // 16)
|
||||
permuted_fp8_scale = scales.view(B, H // 64, 64, M, N // 64, 4).permute(0, 1, 4, 3, 2, 5).reshape(B, H, M, N // 16).to(torch.float8_e4m3fn)
|
||||
return fp8_scale, packed_fp4, dequant_x, permuted_fp8_scale
|
||||
|
||||
b = 2
|
||||
h = 4
|
||||
n = 251
|
||||
d = 128
|
||||
|
||||
q = torch.randn(b, h, n, d, dtype=torch.float16, device='cuda')
|
||||
o = torch.empty((b, h, n, d // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s = torch.empty((b, h, n, d // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
|
||||
fp4quant.scaled_fp4_quant(q, o, o_s, 1)
|
||||
|
||||
fp8_scale, packed_fp4, dequant_x, permuted_fp8_scale = scale_and_fp4_tensor(q, packed_dim=3)
|
||||
|
||||
assert((fp8_scale.float() - o_s.float()).abs().max() == 0)
|
||||
|
||||
o_binary = [
|
||||
(int(bin_str[:4], 2), int(bin_str[4:], 2))
|
||||
for bin_str in [format(x.item(), '08b') for x in o.view(-1)]
|
||||
]
|
||||
o_binary_gt = [
|
||||
(int(bin_str[:4], 2), int(bin_str[4:], 2))
|
||||
for bin_str in [format(x.item(), '08b') for x in packed_fp4.view(-1)]
|
||||
]
|
||||
for i in range(len(o_binary)):
|
||||
# check contiguous 4 bits. Difference should be at most one
|
||||
assert(abs(o_binary[i][0] - o_binary_gt[i][0]) <= 1)
|
||||
assert(abs(o_binary[i][1] - o_binary_gt[i][1]) <= 1)
|
||||
|
||||
print("All tests passed!")
|
||||
@@ -0,0 +1,86 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import torch
|
||||
import fp4quant
|
||||
from triton.tools.mxfp import MXFP4Tensor
|
||||
|
||||
from bench_utils import bench_kineto
|
||||
b = 1
|
||||
h = 32
|
||||
n = 16384
|
||||
d = 128
|
||||
|
||||
def test():
|
||||
q = torch.randn((b, h, n, d), device="cuda", dtype=torch.float16)
|
||||
o = torch.empty((b, h, d, n // 2), device="cuda", dtype=torch.uint8)
|
||||
o_s = torch.empty((b, h, d, n // 16), device="cuda", dtype=torch.float8_e4m3fn)
|
||||
fp4quant.scaled_fp4_quant_trans(q, o, o_s, 1)
|
||||
|
||||
test()
|
||||
|
||||
t = bench_kineto(test, "scaled_fp4_quant_trans_kernel", suppress_kineto_output=True)
|
||||
|
||||
IO = b * h * n * d * 2 + b * h * n * d * 0.5 + b * h * n * d // 16 * 1
|
||||
throughput = IO / t * 1e-9
|
||||
|
||||
print(f"Throughput: {throughput:.2f} GB/s")
|
||||
|
||||
def scale_and_fp4_tensor(x: torch.Tensor, packed_dim: int = 3, all_ones: bool = False, permuted: bool = False):
|
||||
assert x.is_contiguous() and x.ndim == 4 and x.shape[-1] % 16 == 0
|
||||
B, H, M, N = x.shape
|
||||
x = x.view(B, H, M, N // 16, 16)
|
||||
scales = (x.abs().amax(dim=-1, keepdim=True) / 6).to(torch.float32)
|
||||
if all_ones:
|
||||
scales = torch.ones_like(scales)
|
||||
x_scaled = x / scales
|
||||
packed_fp4 = MXFP4Tensor(x_scaled.flatten(start_dim=-2)).to_packed_tensor(dim=packed_dim)
|
||||
dequant_x = (MXFP4Tensor(x_scaled).to(torch.float32) * scales.to(torch.float8_e4m3fn).to(torch.float32)).flatten(start_dim=-2)
|
||||
fp8_scale = scales.flatten(start_dim=-2).to(torch.float8_e4m3fn)
|
||||
permuted_fp8_scale = None
|
||||
if permuted:
|
||||
scales = scales.view(B, H // 64, 4, 16, M, N // 16).permute(0, 1, 3, 2, 4, 5).reshape(B, H, M, N // 16)
|
||||
permuted_fp8_scale = scales.view(B, H // 64, 64, M, N // 64, 4).permute(0, 1, 4, 3, 2, 5).reshape(B, H, M, N // 16).to(torch.float8_e4m3fn)
|
||||
return fp8_scale, packed_fp4, dequant_x, permuted_fp8_scale
|
||||
|
||||
b = 2
|
||||
h = 4
|
||||
n = 491
|
||||
n_padded = (n + 127) // 128 * 128
|
||||
d = 128
|
||||
|
||||
q = torch.randn(b, h, n, d, dtype=torch.float16, device='cuda')
|
||||
o = torch.empty((b, h, d, n_padded // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s = torch.empty((b, h, d, n_padded // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
|
||||
fp4quant.scaled_fp4_quant_trans(q, o, o_s, 1)
|
||||
|
||||
if n % 128 != 0:
|
||||
q_padded = torch.cat([q, torch.zeros((b, h, n_padded - n, d), dtype=torch.float16, device='cuda')], dim=2)
|
||||
else:
|
||||
q_padded = q
|
||||
|
||||
# use torch transpose + scaled_fp4_quant to get the ground truth
|
||||
q_padded = q_padded.transpose(2, 3).reshape(b, h, n_padded, d).contiguous()
|
||||
o_gt = torch.empty((b, h, n_padded, d // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s_gt = torch.empty((b, h, n_padded, d // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
fp4quant.scaled_fp4_quant(q_padded, o_gt, o_s_gt, 1)
|
||||
o_gt = o_gt.reshape(b, h, d, n_padded // 2).contiguous()
|
||||
o_s_gt = o_s_gt.reshape(b, h, d, n_padded // 16).contiguous()
|
||||
|
||||
assert((o_s_gt.float() - o_s.float()).abs().max() == 0)
|
||||
assert((o_gt - o).abs().max() == 0)
|
||||
|
||||
print("All tests passed!")
|
||||
@@ -0,0 +1,169 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
def bench(fn, num_warmups: int = 5, num_tests: int = 10,
|
||||
high_precision: bool = False):
|
||||
# Flush L2 cache with 256 MB data
|
||||
torch.cuda.synchronize()
|
||||
cache = torch.empty(int(256e6 // 4), dtype=torch.int, device='cuda')
|
||||
cache.zero_()
|
||||
|
||||
# Warmup
|
||||
for _ in range(num_warmups):
|
||||
fn()
|
||||
|
||||
# Add a large kernel to eliminate the CPU launch overhead
|
||||
if high_precision:
|
||||
x = torch.randn((8192, 8192), dtype=torch.float, device='cuda')
|
||||
y = torch.randn((8192, 8192), dtype=torch.float, device='cuda')
|
||||
x @ y
|
||||
|
||||
# Testing
|
||||
start_event = torch.cuda.Event(enable_timing=True)
|
||||
end_event = torch.cuda.Event(enable_timing=True)
|
||||
start_event.record()
|
||||
for i in range(num_tests):
|
||||
fn()
|
||||
end_event.record()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
return start_event.elapsed_time(end_event) / num_tests
|
||||
|
||||
|
||||
class empty_suppress:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
pass
|
||||
|
||||
|
||||
class suppress_stdout_stderr:
|
||||
def __enter__(self):
|
||||
self.outnull_file = open(os.devnull, 'w')
|
||||
self.errnull_file = open(os.devnull, 'w')
|
||||
|
||||
self.old_stdout_fileno_undup = sys.stdout.fileno()
|
||||
self.old_stderr_fileno_undup = sys.stderr.fileno()
|
||||
|
||||
self.old_stdout_fileno = os.dup(sys.stdout.fileno())
|
||||
self.old_stderr_fileno = os.dup(sys.stderr.fileno())
|
||||
|
||||
self.old_stdout = sys.stdout
|
||||
self.old_stderr = sys.stderr
|
||||
|
||||
os.dup2(self.outnull_file.fileno(), self.old_stdout_fileno_undup)
|
||||
os.dup2(self.errnull_file.fileno(), self.old_stderr_fileno_undup)
|
||||
|
||||
sys.stdout = self.outnull_file
|
||||
sys.stderr = self.errnull_file
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
sys.stdout = self.old_stdout
|
||||
sys.stderr = self.old_stderr
|
||||
|
||||
os.dup2(self.old_stdout_fileno, self.old_stdout_fileno_undup)
|
||||
os.dup2(self.old_stderr_fileno, self.old_stderr_fileno_undup)
|
||||
|
||||
os.close(self.old_stdout_fileno)
|
||||
os.close(self.old_stderr_fileno)
|
||||
|
||||
self.outnull_file.close()
|
||||
self.errnull_file.close()
|
||||
|
||||
|
||||
def bench_kineto(fn, kernel_names, num_tests: int = 30, suppress_kineto_output: bool = False,
|
||||
trace_path: str = None, barrier_comm_profiling: bool = False, flush_l2: bool = False):
|
||||
# Conflict with Nsight Systems
|
||||
using_nsys = os.environ.get('DG_NSYS_PROFILING', False)
|
||||
|
||||
# For some auto-tuning kernels with prints
|
||||
fn()
|
||||
|
||||
# Profile
|
||||
suppress = suppress_stdout_stderr if suppress_kineto_output and not using_nsys else empty_suppress
|
||||
with suppress():
|
||||
schedule = torch.profiler.schedule(wait=0, warmup=1, active=1, repeat=1) if not using_nsys else None
|
||||
profiler = torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CUDA], schedule=schedule) if not using_nsys else empty_suppress()
|
||||
with profiler:
|
||||
for i in range(2):
|
||||
# NOTES: use a large kernel and a barrier to eliminate the unbalanced CPU launch overhead
|
||||
if barrier_comm_profiling:
|
||||
lhs = torch.randn((8192, 8192), dtype=torch.float, device='cuda')
|
||||
rhs = torch.randn((8192, 8192), dtype=torch.float, device='cuda')
|
||||
lhs @ rhs
|
||||
dist.all_reduce(torch.ones(1, dtype=torch.float, device='cuda'))
|
||||
for _ in range(num_tests):
|
||||
if flush_l2:
|
||||
torch.empty(int(256e6 // 4), dtype=torch.int, device='cuda').zero_()
|
||||
fn()
|
||||
|
||||
if not using_nsys:
|
||||
profiler.step()
|
||||
|
||||
# Return 1 if using Nsight Systems
|
||||
if using_nsys:
|
||||
return 1
|
||||
|
||||
# Parse the profiling table
|
||||
assert isinstance(kernel_names, str) or isinstance(kernel_names, tuple)
|
||||
is_tupled = isinstance(kernel_names, tuple)
|
||||
prof_lines = profiler.key_averages().table(sort_by='cuda_time_total', max_name_column_width=100).split('\n')
|
||||
kernel_names = (kernel_names, ) if isinstance(kernel_names, str) else kernel_names
|
||||
assert all([isinstance(name, str) for name in kernel_names])
|
||||
for name in kernel_names:
|
||||
assert sum([name in line for line in prof_lines]) == 1, f'Errors of the kernel {name} in the profiling table'
|
||||
|
||||
# Save chrome traces
|
||||
if trace_path is not None:
|
||||
profiler.export_chrome_trace(trace_path)
|
||||
|
||||
# Return average kernel times
|
||||
units = {'ms': 1e3, 'us': 1e6}
|
||||
kernel_times = []
|
||||
for name in kernel_names:
|
||||
for line in prof_lines:
|
||||
if name in line:
|
||||
time_str = line.split()[-2]
|
||||
for unit, scale in units.items():
|
||||
if unit in time_str:
|
||||
kernel_times.append(float(time_str.replace(unit, '')) / scale)
|
||||
break
|
||||
break
|
||||
return tuple(kernel_times) if is_tupled else kernel_times[0]
|
||||
|
||||
|
||||
def calc_diff(x, y):
|
||||
x, y = x.double(), y.double()
|
||||
denominator = (x * x + y * y).sum()
|
||||
sim = 2 * (x * y).sum() / denominator
|
||||
return 1 - sim
|
||||
|
||||
|
||||
def count_bytes(tensors):
|
||||
total = 0
|
||||
for t in tensors:
|
||||
if isinstance(t, tuple):
|
||||
total += count_bytes(t)
|
||||
else:
|
||||
total += t.numel() * t.element_size()
|
||||
return total
|
||||
@@ -0,0 +1,52 @@
|
||||
#pragma once
|
||||
|
||||
#include <stdio.h>
|
||||
|
||||
#if defined(__HIPCC__)
|
||||
#define HOST_DEVICE_INLINE __host__ __device__
|
||||
#define DEVICE_INLINE __device__
|
||||
#define HOST_INLINE __host__
|
||||
#elif defined(__CUDACC__) || defined(_NVHPC_CUDA)
|
||||
#define HOST_DEVICE_INLINE __host__ __device__ __forceinline__
|
||||
#define DEVICE_INLINE __device__ __forceinline__
|
||||
#define HOST_INLINE __host__ __forceinline__
|
||||
#else
|
||||
#define HOST_DEVICE_INLINE inline
|
||||
#define DEVICE_INLINE inline
|
||||
#define HOST_INLINE inline
|
||||
#endif
|
||||
|
||||
#define CUDA_CHECK(cmd) \
|
||||
do { \
|
||||
cudaError_t e = cmd; \
|
||||
if (e != cudaSuccess) { \
|
||||
printf("Failed: Cuda error %s:%d '%s'\n", __FILE__, __LINE__, \
|
||||
cudaGetErrorString(e)); \
|
||||
exit(EXIT_FAILURE); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
int64_t get_device_attribute(int64_t attribute, int64_t device_id) {
|
||||
static int value = [=]() {
|
||||
int device = static_cast<int>(device_id);
|
||||
if (device < 0) {
|
||||
CUDA_CHECK(cudaGetDevice(&device));
|
||||
}
|
||||
int value;
|
||||
CUDA_CHECK(cudaDeviceGetAttribute(
|
||||
&value, static_cast<cudaDeviceAttr>(attribute), device));
|
||||
return static_cast<int>(value);
|
||||
}();
|
||||
|
||||
return value;
|
||||
}
|
||||
|
||||
namespace cuda_utils {
|
||||
|
||||
template <typename T>
|
||||
HOST_DEVICE_INLINE constexpr std::enable_if_t<std::is_integral_v<T>, T>
|
||||
ceil_div(T a, T b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
}; // namespace cuda_utils
|
||||
@@ -0,0 +1,644 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#include <torch/all.h>
|
||||
#include <torch/python.h>
|
||||
#include <torch/nn/functional.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <cuda_runtime_api.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <cuda_fp8.h>
|
||||
|
||||
#include "cuda_utils.h"
|
||||
#include "../blackwell/block_config.h"
|
||||
|
||||
#define DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(pytorch_dtype, c_type, ...) \
|
||||
if (pytorch_dtype == at::ScalarType::Half) { \
|
||||
using c_type = half; \
|
||||
__VA_ARGS__ \
|
||||
} else if (pytorch_dtype == at::ScalarType::BFloat16) { \
|
||||
using c_type = nv_bfloat16; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
std::ostringstream oss; \
|
||||
oss << __PRETTY_FUNCTION__ << " failed to dispatch data type " << pytorch_dtype; \
|
||||
TORCH_CHECK(false, oss.str()); \
|
||||
}
|
||||
|
||||
#define DISPATCH_HEAD_DIM(head_dim, HEAD_DIM, ...) \
|
||||
if (head_dim == 64) { \
|
||||
constexpr int HEAD_DIM = 64; \
|
||||
__VA_ARGS__ \
|
||||
} else if (head_dim == 128) { \
|
||||
constexpr int HEAD_DIM = 128; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
std::ostringstream err_msg; \
|
||||
err_msg << "Unsupported head dim: " << int(head_dim); \
|
||||
throw std::invalid_argument(err_msg.str()); \
|
||||
}
|
||||
|
||||
#define CHECK_CUDA(x) \
|
||||
TORCH_CHECK(x.is_cuda(), "Tensor " #x " must be on CUDA")
|
||||
#define CHECK_DTYPE(x, true_dtype) \
|
||||
TORCH_CHECK(x.dtype() == true_dtype, \
|
||||
"Tensor " #x " must have dtype (" #true_dtype ")")
|
||||
#define CHECK_DIMS(x, true_dim) \
|
||||
TORCH_CHECK(x.dim() == true_dim, \
|
||||
"Tensor " #x " must have dimension number (" #true_dim ")")
|
||||
#define CHECK_SHAPE(x, ...) \
|
||||
TORCH_CHECK(x.sizes() == torch::IntArrayRef({__VA_ARGS__}), \
|
||||
"Tensor " #x " must have shape (" #__VA_ARGS__ ")")
|
||||
#define CHECK_CONTIGUOUS(x) \
|
||||
TORCH_CHECK(x.is_contiguous(), "Tensor " #x " must be contiguous")
|
||||
#define CHECK_LASTDIM_CONTIGUOUS(x) \
|
||||
TORCH_CHECK(x.stride(-1) == 1, \
|
||||
"Tensor " #x " must be contiguous at the last dimension")
|
||||
|
||||
constexpr int CVT_FP4_ELTS_PER_THREAD = 16;
|
||||
|
||||
// Convert 4 float2 values into 8 e2m1 values (represented as one uint32_t).
|
||||
inline __device__ uint32_t fp32_vec_to_e2m1(float2 *array) {
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000)
|
||||
uint32_t val;
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .b8 byte0;\n"
|
||||
".reg .b8 byte1;\n"
|
||||
".reg .b8 byte2;\n"
|
||||
".reg .b8 byte3;\n"
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte0, %2, %1;\n"
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte1, %4, %3;\n"
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte2, %6, %5;\n"
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte3, %8, %7;\n"
|
||||
"mov.b32 %0, {byte0, byte1, byte2, byte3};\n"
|
||||
"}"
|
||||
: "=r"(val)
|
||||
: "f"(array[0].x), "f"(array[0].y), "f"(array[1].x), "f"(array[1].y),
|
||||
"f"(array[2].x), "f"(array[2].y), "f"(array[3].x), "f"(array[3].y));
|
||||
return val;
|
||||
#else
|
||||
return 0;
|
||||
#endif
|
||||
}
|
||||
|
||||
// Get type2 from type or vice versa (applied to half and bfloat16)
|
||||
template <typename T>
|
||||
struct TypeConverter {
|
||||
using Type = half2;
|
||||
}; // keep for generality
|
||||
|
||||
template <>
|
||||
struct TypeConverter<half2> {
|
||||
using Type = half;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct TypeConverter<half> {
|
||||
using Type = half2;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct TypeConverter<__nv_bfloat162> {
|
||||
using Type = __nv_bfloat16;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct TypeConverter<__nv_bfloat16> {
|
||||
using Type = __nv_bfloat162;
|
||||
};
|
||||
|
||||
// Define a 32 bytes packed data type.
|
||||
template <class Type>
|
||||
struct PackedVec {
|
||||
typename TypeConverter<Type>::Type elts[8];
|
||||
};
|
||||
|
||||
template <uint32_t head_dim, uint32_t BLOCK_SIZE, bool permute, typename T>
|
||||
__global__ void scaled_fp4_quant_kernel(
|
||||
const T* input, uint8_t* output, uint8_t* output_sf,
|
||||
int batch_size, int num_heads, int num_tokens,
|
||||
int stride_bz_input, int stride_h_input, int stride_seq_input,
|
||||
int stride_bz_output, int stride_h_output, int stride_seq_output,
|
||||
int stride_bz_output_sf, int stride_h_output_sf, int stride_seq_output_sf) {
|
||||
static_assert(std::is_same<T, half>::value || std::is_same<T, nv_bfloat16>::value, "Only half and bfloat16 input are supported");
|
||||
using PackedVec = PackedVec<T>;
|
||||
|
||||
const int batch_id = blockIdx.y;
|
||||
const int head_id = blockIdx.z;
|
||||
const int token_block_id = blockIdx.x;
|
||||
|
||||
static_assert(CVT_FP4_ELTS_PER_THREAD == 8 || CVT_FP4_ELTS_PER_THREAD == 16,
|
||||
"CVT_FP4_ELTS_PER_THREAD must be 8 or 16");
|
||||
static_assert(sizeof(PackedVec) == sizeof(T) * CVT_FP4_ELTS_PER_THREAD,
|
||||
"Vec size is not matched.");
|
||||
|
||||
constexpr uint32_t NUM_THREADS_PER_TOKEN = head_dim / CVT_FP4_ELTS_PER_THREAD;
|
||||
|
||||
// load input
|
||||
const int token_id = token_block_id * BLOCK_SIZE + threadIdx.x / NUM_THREADS_PER_TOKEN;
|
||||
|
||||
int load_token_id;
|
||||
if constexpr (!permute) {
|
||||
load_token_id = token_id;
|
||||
} else {
|
||||
int local_token_id = threadIdx.x / NUM_THREADS_PER_TOKEN;
|
||||
int local_token_id_residue = local_token_id % 32;
|
||||
// [0, 1, 8, 9, 16, 17, 24, 25, 2, 3, 10, 11, 18, 19, 26, 27, 4, 5, 12, 13, 20, 21, 28, 29, 6, 7, 14, 15, 22, 23, 30, 31]
|
||||
load_token_id = token_block_id * BLOCK_SIZE + (local_token_id / 32) * 32 +
|
||||
(local_token_id_residue / 8) * 2 +
|
||||
((local_token_id_residue % 8) / 2) * 8 +
|
||||
(local_token_id_residue % 8) % 2;
|
||||
}
|
||||
|
||||
PackedVec in_vec;
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
reinterpret_cast<uint32_t&>(in_vec.elts[i]) = 0;
|
||||
}
|
||||
|
||||
if (load_token_id < num_tokens) {
|
||||
in_vec = reinterpret_cast<PackedVec const*>(input +
|
||||
batch_id * stride_bz_input + // batch dim
|
||||
head_id * stride_h_input + // head dim
|
||||
load_token_id * stride_seq_input + // seq dim
|
||||
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD)[0]; // feature dim
|
||||
}
|
||||
|
||||
// calculate max of every consecutive 16 elements
|
||||
auto localMax = __habs2(in_vec.elts[0]);
|
||||
#pragma unroll
|
||||
for (int i = 1; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) { // local max
|
||||
localMax = __hmax2(localMax, __habs2(in_vec.elts[i]));
|
||||
}
|
||||
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 8) { // shuffle across two threads
|
||||
localMax = __hmax2(__shfl_xor_sync(0xffffffff, localMax, 1, 32), localMax);
|
||||
}
|
||||
|
||||
float vecMax = float(__hmax(localMax.x, localMax.y));
|
||||
|
||||
// scaling factor
|
||||
float SFValue = vecMax / 6.0f;
|
||||
uint8_t SFValueFP8;
|
||||
reinterpret_cast<__nv_fp8_e4m3&>(SFValueFP8) = __nv_fp8_e4m3(SFValue);
|
||||
SFValue = float(reinterpret_cast<__nv_fp8_e4m3&>(SFValueFP8));
|
||||
|
||||
float SFValueInv = (SFValue == 0.0f) ? 0.0f : 1.0f / SFValue;
|
||||
|
||||
// convert input to float2 and apply scale
|
||||
float2 fp2Vals[CVT_FP4_ELTS_PER_THREAD / 2];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
if constexpr (std::is_same<T, half>::value) {
|
||||
fp2Vals[i] = __half22float2(in_vec.elts[i]);
|
||||
} else {
|
||||
fp2Vals[i] = __bfloat1622float2(in_vec.elts[i]);
|
||||
}
|
||||
fp2Vals[i].x = fp2Vals[i].x * SFValueInv;
|
||||
fp2Vals[i].y = fp2Vals[i].y * SFValueInv;
|
||||
}
|
||||
|
||||
// convert to e2m1
|
||||
uint32_t e2m1Vals[CVT_FP4_ELTS_PER_THREAD / 8];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 8; i++) {
|
||||
e2m1Vals[i] = fp32_vec_to_e2m1(fp2Vals + i * 4);
|
||||
}
|
||||
|
||||
// Skip out-of-range tokens: never write past the sequence when num_tokens
|
||||
// is not a multiple of BLOCK_SIZE (out-of-bounds write fix).
|
||||
if (token_id >= num_tokens) return;
|
||||
|
||||
// save
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 8) {
|
||||
reinterpret_cast<uint32_t*>(output +
|
||||
batch_id * stride_bz_output +
|
||||
head_id * stride_h_output +
|
||||
token_id * stride_seq_output +
|
||||
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD / 2)[0] = e2m1Vals[0];
|
||||
} else {
|
||||
reinterpret_cast<uint64_t*>(output +
|
||||
batch_id * stride_bz_output +
|
||||
head_id * stride_h_output +
|
||||
token_id * stride_seq_output +
|
||||
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD / 2)[0] = reinterpret_cast<uint64_t*>(e2m1Vals)[0];
|
||||
}
|
||||
|
||||
uint8_t* output_sf_save_base = output_sf + batch_id * stride_bz_output_sf + head_id * stride_h_output_sf + (token_id / 64) * 64 * stride_seq_output_sf;
|
||||
uint32_t token_id_local = token_id % 64;
|
||||
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 16) {
|
||||
uint32_t col_id_local = threadIdx.x % NUM_THREADS_PER_TOKEN;
|
||||
uint32_t offset_local = (col_id_local / 4) * 256 + (col_id_local % 4) +
|
||||
(token_id_local / 16) * 4 + (token_id_local % 16) * 16;
|
||||
reinterpret_cast<uint8_t*>(output_sf_save_base + offset_local)[0] = SFValueFP8;
|
||||
} else {
|
||||
if (threadIdx.x % 2 == 0) {
|
||||
uint32_t col_id_local = (threadIdx.x % NUM_THREADS_PER_TOKEN) / 2;
|
||||
uint32_t offset_local = (col_id_local / 4) * 256 + (col_id_local % 4) +
|
||||
(token_id_local / 16) * 4 + (token_id_local % 16) * 16;
|
||||
reinterpret_cast<uint8_t*>(output_sf_save_base + offset_local)[0] = SFValueFP8;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <uint32_t head_dim, uint32_t BLOCK_SIZE, typename T>
|
||||
__global__ void scaled_fp4_quant_trans_kernel(
|
||||
const T* input, uint8_t* output, uint8_t* output_sf,
|
||||
int batch_size, int num_heads, int num_tokens,
|
||||
int stride_bz_input, int stride_h_input, int stride_seq_input,
|
||||
int stride_bz_output, int stride_h_output, int stride_d_output,
|
||||
int stride_bz_output_sf, int stride_h_output_sf, int stride_d_output_sf) {
|
||||
static_assert(std::is_same<T, half>::value || std::is_same<T, nv_bfloat16>::value, "Only half and bfloat16 input are supported");
|
||||
using PackedVec = PackedVec<T>;
|
||||
|
||||
const int batch_id = blockIdx.y;
|
||||
const int head_id = blockIdx.z;
|
||||
const int token_block_id = blockIdx.x;
|
||||
|
||||
static_assert(CVT_FP4_ELTS_PER_THREAD == 8 || CVT_FP4_ELTS_PER_THREAD == 16,
|
||||
"CVT_FP4_ELTS_PER_THREAD must be 8 or 16");
|
||||
static_assert(sizeof(PackedVec) == sizeof(T) * CVT_FP4_ELTS_PER_THREAD,
|
||||
"Vec size is not matched.");
|
||||
|
||||
constexpr uint32_t NUM_THREADS_PER_TOKEN = head_dim / CVT_FP4_ELTS_PER_THREAD;
|
||||
constexpr uint32_t NUM_THREADS_PER_SEQ = BLOCK_SIZE / CVT_FP4_ELTS_PER_THREAD;
|
||||
|
||||
// load input
|
||||
const int token_id = token_block_id * BLOCK_SIZE + threadIdx.x / NUM_THREADS_PER_TOKEN;
|
||||
// Permute V rows within each 32-element block so the PV MMA K-indexed
|
||||
// access reads the correct CLayout N-indexed values (Edenzzzz causal fix).
|
||||
const int k_intra = token_id & 31;
|
||||
const int load_token_id = (token_id & ~31)
|
||||
| ((k_intra & 6) << 2) | ((k_intra & 24) >> 2) | (k_intra & 1);
|
||||
|
||||
PackedVec in_vec;
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
reinterpret_cast<uint32_t&>(in_vec.elts[i]) = 0;
|
||||
}
|
||||
|
||||
if (load_token_id < num_tokens) {
|
||||
in_vec = reinterpret_cast<PackedVec const*>(input +
|
||||
batch_id * stride_bz_input + // batch dim
|
||||
head_id * stride_h_input + // head dim
|
||||
load_token_id * stride_seq_input + // seq dim (permuted)
|
||||
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD)[0]; // feature dim
|
||||
}
|
||||
|
||||
// transpose
|
||||
__shared__ T shared_input[BLOCK_SIZE * head_dim];
|
||||
reinterpret_cast<PackedVec*>(shared_input)[threadIdx.x] = in_vec;
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
in_vec.elts[i].x = shared_input[(threadIdx.x / NUM_THREADS_PER_SEQ) + ((threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD + 2 * i) * head_dim];
|
||||
in_vec.elts[i].y = shared_input[(threadIdx.x / NUM_THREADS_PER_SEQ) + ((threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD + 2 * i + 1) * head_dim];
|
||||
}
|
||||
|
||||
// calculate max of every consecutive 16 elements
|
||||
auto localMax = __habs2(in_vec.elts[0]);
|
||||
#pragma unroll
|
||||
for (int i = 1; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) { // local max
|
||||
localMax = __hmax2(localMax, __habs2(in_vec.elts[i]));
|
||||
}
|
||||
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 8) { // shuffle across two threads
|
||||
localMax = __hmax2(__shfl_xor_sync(0xffffffff, localMax, 1, 32), localMax);
|
||||
}
|
||||
|
||||
float vecMax = float(__hmax(localMax.x, localMax.y));
|
||||
|
||||
// scaling factor
|
||||
float SFValue = vecMax / 6.0f;
|
||||
uint8_t SFValueFP8;
|
||||
reinterpret_cast<__nv_fp8_e4m3&>(SFValueFP8) = __nv_fp8_e4m3(SFValue);
|
||||
SFValue = float(reinterpret_cast<__nv_fp8_e4m3&>(SFValueFP8));
|
||||
|
||||
float SFValueInv = (SFValue == 0.0f) ? 0.0f : 1.0f / SFValue;
|
||||
|
||||
// convert input to float2 and apply scale
|
||||
float2 fp2Vals[CVT_FP4_ELTS_PER_THREAD / 2];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
if constexpr (std::is_same<T, half>::value) {
|
||||
fp2Vals[i] = __half22float2(in_vec.elts[i]);
|
||||
} else {
|
||||
fp2Vals[i] = __bfloat1622float2(in_vec.elts[i]);
|
||||
}
|
||||
fp2Vals[i].x = fp2Vals[i].x * SFValueInv;
|
||||
fp2Vals[i].y = fp2Vals[i].y * SFValueInv;
|
||||
}
|
||||
|
||||
// convert to e2m1
|
||||
uint32_t e2m1Vals[CVT_FP4_ELTS_PER_THREAD / 8];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 8; i++) {
|
||||
e2m1Vals[i] = fp32_vec_to_e2m1(fp2Vals + i * 4);
|
||||
}
|
||||
|
||||
// Skip out-of-range tokens: never write past the sequence when num_tokens
|
||||
// is not a multiple of BLOCK_SIZE (out-of-bounds write fix).
|
||||
const int write_token_id = token_block_id * BLOCK_SIZE +
|
||||
(threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD;
|
||||
if (write_token_id >= num_tokens) return;
|
||||
|
||||
// save
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 8) {
|
||||
reinterpret_cast<uint32_t*>(output +
|
||||
batch_id * stride_bz_output +
|
||||
head_id * stride_h_output +
|
||||
(threadIdx.x / NUM_THREADS_PER_SEQ) * stride_d_output +
|
||||
(token_block_id * BLOCK_SIZE + (threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD) / 2)[0] = e2m1Vals[0];
|
||||
} else {
|
||||
reinterpret_cast<uint64_t*>(output +
|
||||
batch_id * stride_bz_output +
|
||||
head_id * stride_h_output +
|
||||
(threadIdx.x / NUM_THREADS_PER_SEQ) * stride_d_output +
|
||||
(token_block_id * BLOCK_SIZE + (threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD) / 2)[0] = reinterpret_cast<uint64_t*>(e2m1Vals)[0];
|
||||
}
|
||||
|
||||
uint8_t *output_sf_save_base = output_sf +
|
||||
batch_id * stride_bz_output_sf +
|
||||
head_id * stride_h_output_sf +
|
||||
(threadIdx.x / NUM_THREADS_PER_SEQ / 64) * 64 * stride_d_output_sf;
|
||||
uint32_t row_id_local = (threadIdx.x / NUM_THREADS_PER_SEQ) % 64;
|
||||
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 16) {
|
||||
uint32_t col_id_local = token_block_id * BLOCK_SIZE / CVT_FP4_ELTS_PER_THREAD + threadIdx.x % NUM_THREADS_PER_SEQ;
|
||||
uint32_t offset_local = (col_id_local / 4) * 256 + (col_id_local % 4) +
|
||||
(row_id_local / 16) * 4 + (row_id_local % 16) * 16;
|
||||
reinterpret_cast<uint8_t*>(output_sf_save_base + offset_local)[0] = SFValueFP8;
|
||||
} else {
|
||||
if (threadIdx.x % 2 == 0) {
|
||||
uint32_t col_id_local = token_block_id * BLOCK_SIZE / CVT_FP4_ELTS_PER_THREAD + (threadIdx.x % NUM_THREADS_PER_SEQ) / 2;
|
||||
uint32_t offset_local = (col_id_local / 4) * 256 + (col_id_local % 4) +
|
||||
(row_id_local / 16) * 4 + (row_id_local % 16) * 16;
|
||||
reinterpret_cast<uint8_t*>(output_sf_save_base + offset_local)[0] = SFValueFP8;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void scaled_fp4_quant(torch::Tensor const& input,
|
||||
torch::Tensor const& output,
|
||||
torch::Tensor const& output_sf,
|
||||
int tensor_layout) {
|
||||
constexpr int BLOCK_SIZE = flash::BLOCK_M;
|
||||
|
||||
CHECK_CUDA(input);
|
||||
CHECK_CUDA(output);
|
||||
CHECK_CUDA(output_sf);
|
||||
|
||||
CHECK_LASTDIM_CONTIGUOUS(input);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output_sf);
|
||||
|
||||
CHECK_DTYPE(output, at::ScalarType::Byte);
|
||||
CHECK_DTYPE(output_sf, at::ScalarType::Float8_e4m3fn);
|
||||
|
||||
CHECK_DIMS(input, 4);
|
||||
CHECK_DIMS(output, 4);
|
||||
CHECK_DIMS(output_sf, 4);
|
||||
|
||||
const int batch_size = input.size(0);
|
||||
const int head_dim = input.size(3);
|
||||
|
||||
const int stride_bz_input = input.stride(0);
|
||||
const int stride_bz_output = output.stride(0);
|
||||
const int stride_bz_output_sf = output_sf.stride(0);
|
||||
|
||||
int num_tokens, num_heads;
|
||||
int stride_seq_input, stride_seq_output, stride_seq_output_sf;
|
||||
int stride_h_input, stride_h_output, stride_h_output_sf;
|
||||
if (tensor_layout == 0) {
|
||||
num_tokens = input.size(1);
|
||||
num_heads = input.size(2);
|
||||
stride_seq_input = input.stride(1);
|
||||
stride_seq_output = output.stride(1);
|
||||
stride_seq_output_sf = output_sf.stride(1);
|
||||
stride_h_input = input.stride(2);
|
||||
stride_h_output = output.stride(2);
|
||||
stride_h_output_sf = output_sf.stride(2);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, num_tokens, num_heads, head_dim / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, num_tokens, num_heads, head_dim / 16);
|
||||
} else {
|
||||
num_tokens = input.size(2);
|
||||
num_heads = input.size(1);
|
||||
stride_seq_input = input.stride(2);
|
||||
stride_seq_output = output.stride(2);
|
||||
stride_seq_output_sf = output_sf.stride(2);
|
||||
stride_h_input = input.stride(1);
|
||||
stride_h_output = output.stride(1);
|
||||
stride_h_output_sf = output_sf.stride(1);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, num_heads, num_tokens, head_dim / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, num_heads, num_tokens, head_dim / 16);
|
||||
}
|
||||
|
||||
auto input_dtype = input.scalar_type();
|
||||
auto stream = at::cuda::getCurrentCUDAStream(input.get_device());
|
||||
|
||||
DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(input_dtype, c_type, {
|
||||
DISPATCH_HEAD_DIM(head_dim, HEAD_DIM, {
|
||||
dim3 block(BLOCK_SIZE * HEAD_DIM / CVT_FP4_ELTS_PER_THREAD, 1, 1);
|
||||
dim3 grid((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE, batch_size, num_heads);
|
||||
|
||||
scaled_fp4_quant_kernel<HEAD_DIM, BLOCK_SIZE, false, c_type>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
reinterpret_cast<c_type*>(input.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output_sf.data_ptr()),
|
||||
batch_size, num_heads, num_tokens,
|
||||
stride_bz_input, stride_h_input, stride_seq_input,
|
||||
stride_bz_output, stride_h_output, stride_seq_output,
|
||||
stride_bz_output_sf, stride_h_output_sf, stride_seq_output_sf);
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
void scaled_fp4_quant_permute(torch::Tensor const& input,
|
||||
torch::Tensor const& output,
|
||||
torch::Tensor const& output_sf,
|
||||
int tensor_layout) {
|
||||
constexpr int BLOCK_SIZE = flash::BLOCK_M;
|
||||
|
||||
CHECK_CUDA(input);
|
||||
CHECK_CUDA(output);
|
||||
CHECK_CUDA(output_sf);
|
||||
|
||||
CHECK_LASTDIM_CONTIGUOUS(input);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output_sf);
|
||||
|
||||
CHECK_DTYPE(output, at::ScalarType::Byte);
|
||||
CHECK_DTYPE(output_sf, at::ScalarType::Float8_e4m3fn);
|
||||
|
||||
CHECK_DIMS(input, 4);
|
||||
CHECK_DIMS(output, 4);
|
||||
CHECK_DIMS(output_sf, 4);
|
||||
|
||||
const int batch_size = input.size(0);
|
||||
const int head_dim = input.size(3);
|
||||
|
||||
const int stride_bz_input = input.stride(0);
|
||||
const int stride_bz_output = output.stride(0);
|
||||
const int stride_bz_output_sf = output_sf.stride(0);
|
||||
|
||||
int num_tokens, num_heads;
|
||||
int stride_seq_input, stride_seq_output, stride_seq_output_sf;
|
||||
int stride_h_input, stride_h_output, stride_h_output_sf;
|
||||
if (tensor_layout == 0) {
|
||||
num_tokens = input.size(1);
|
||||
num_heads = input.size(2);
|
||||
stride_seq_input = input.stride(1);
|
||||
stride_seq_output = output.stride(1);
|
||||
stride_seq_output_sf = output_sf.stride(1);
|
||||
stride_h_input = input.stride(2);
|
||||
stride_h_output = output.stride(2);
|
||||
stride_h_output_sf = output_sf.stride(2);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE, num_heads, head_dim / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE, num_heads, head_dim / 16);
|
||||
} else {
|
||||
num_tokens = input.size(2);
|
||||
num_heads = input.size(1);
|
||||
stride_seq_input = input.stride(2);
|
||||
stride_seq_output = output.stride(2);
|
||||
stride_seq_output_sf = output_sf.stride(2);
|
||||
stride_h_input = input.stride(1);
|
||||
stride_h_output = output.stride(1);
|
||||
stride_h_output_sf = output_sf.stride(1);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, num_heads, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE, head_dim / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, num_heads, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE, head_dim / 16);
|
||||
}
|
||||
|
||||
auto input_dtype = input.scalar_type();
|
||||
auto stream = at::cuda::getCurrentCUDAStream(input.get_device());
|
||||
|
||||
DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(input_dtype, c_type, {
|
||||
DISPATCH_HEAD_DIM(head_dim, HEAD_DIM, {
|
||||
constexpr int BLOCK_SIZE = flash::BLOCK_M;
|
||||
dim3 block(BLOCK_SIZE * HEAD_DIM / CVT_FP4_ELTS_PER_THREAD, 1, 1);
|
||||
dim3 grid((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE, batch_size, num_heads);
|
||||
|
||||
scaled_fp4_quant_kernel<HEAD_DIM, BLOCK_SIZE, true, c_type>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
reinterpret_cast<c_type*>(input.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output_sf.data_ptr()),
|
||||
batch_size, num_heads, num_tokens,
|
||||
stride_bz_input, stride_h_input, stride_seq_input,
|
||||
stride_bz_output, stride_h_output, stride_seq_output,
|
||||
stride_bz_output_sf, stride_h_output_sf, stride_seq_output_sf);
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
void scaled_fp4_quant_trans(torch::Tensor const& input,
|
||||
torch::Tensor const& output,
|
||||
torch::Tensor const& output_sf,
|
||||
int tensor_layout) {
|
||||
constexpr int BLOCK_SIZE = flash::BLOCK_M;
|
||||
|
||||
CHECK_CUDA(input);
|
||||
CHECK_CUDA(output);
|
||||
CHECK_CUDA(output_sf);
|
||||
|
||||
CHECK_LASTDIM_CONTIGUOUS(input);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output_sf);
|
||||
|
||||
CHECK_DTYPE(output, at::ScalarType::Byte);
|
||||
CHECK_DTYPE(output_sf, at::ScalarType::Float8_e4m3fn);
|
||||
|
||||
CHECK_DIMS(input, 4);
|
||||
CHECK_DIMS(output, 4);
|
||||
CHECK_DIMS(output_sf, 4);
|
||||
|
||||
const int batch_size = input.size(0);
|
||||
const int head_dim = input.size(3);
|
||||
|
||||
const int stride_bz_input = input.stride(0);
|
||||
const int stride_bz_output = output.stride(0);
|
||||
const int stride_bz_output_sf = output_sf.stride(0);
|
||||
|
||||
int num_tokens, num_heads;
|
||||
int stride_seq_input;
|
||||
int stride_d_output, stride_d_output_sf;
|
||||
int stride_h_input, stride_h_output, stride_h_output_sf;
|
||||
if (tensor_layout == 0) {
|
||||
num_tokens = input.size(1);
|
||||
num_heads = input.size(2);
|
||||
stride_seq_input = input.stride(1);
|
||||
stride_d_output = output.stride(1);
|
||||
stride_d_output_sf = output_sf.stride(1);
|
||||
stride_h_input = input.stride(2);
|
||||
stride_h_output = output.stride(2);
|
||||
stride_h_output_sf = output_sf.stride(2);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, head_dim, num_heads, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, head_dim, num_heads, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE / 16);
|
||||
} else {
|
||||
num_tokens = input.size(2);
|
||||
num_heads = input.size(1);
|
||||
stride_seq_input = input.stride(2);
|
||||
stride_d_output = output.stride(2);
|
||||
stride_d_output_sf = output_sf.stride(2);
|
||||
stride_h_input = input.stride(1);
|
||||
stride_h_output = output.stride(1);
|
||||
stride_h_output_sf = output_sf.stride(1);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, num_heads, head_dim, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, num_heads, head_dim, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE / 16);
|
||||
}
|
||||
|
||||
auto input_dtype = input.scalar_type();
|
||||
auto stream = at::cuda::getCurrentCUDAStream(input.get_device());
|
||||
|
||||
DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(input_dtype, c_type, {
|
||||
DISPATCH_HEAD_DIM(head_dim, HEAD_DIM, {
|
||||
dim3 block(BLOCK_SIZE * HEAD_DIM / CVT_FP4_ELTS_PER_THREAD, 1, 1);
|
||||
dim3 grid((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE, batch_size, num_heads);
|
||||
|
||||
scaled_fp4_quant_trans_kernel<HEAD_DIM, BLOCK_SIZE, c_type>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
reinterpret_cast<c_type*>(input.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output_sf.data_ptr()),
|
||||
batch_size, num_heads, num_tokens,
|
||||
stride_bz_input, stride_h_input, stride_seq_input,
|
||||
stride_bz_output, stride_h_output, stride_d_output,
|
||||
stride_bz_output_sf, stride_h_output_sf, stride_d_output_sf);
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("scaled_fp4_quant", &scaled_fp4_quant);
|
||||
m.def("scaled_fp4_quant_permute", &scaled_fp4_quant_permute);
|
||||
m.def("scaled_fp4_quant_trans", &scaled_fp4_quant_trans);
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -32,4 +32,4 @@ dependencies = [
|
||||
[tool.scikit-build]
|
||||
cmake.build-type = "Release"
|
||||
minimum-version = "build-system.requires"
|
||||
wheel.packages = ["python/fastvideo_kernel"]
|
||||
wheel.packages = ["python/fastvideo_kernel", "attn_qat_infer"]
|
||||
|
||||
@@ -25,12 +25,17 @@ from fastvideo_kernel.turbodiffusion_ops import (
|
||||
int8_quant,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.block_sparse_attn_varlen import (
|
||||
block_sparse_attn_varlen,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"sliding_tile_attention",
|
||||
"video_sparse_attn",
|
||||
"video_sparse_attn_bshd",
|
||||
"block_sparse_attn",
|
||||
"block_sparse_attn_from_indices",
|
||||
"block_sparse_attn_varlen",
|
||||
"moba_attn_varlen",
|
||||
"process_moba_input",
|
||||
"process_moba_output",
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
"""Variable-length block-sparse attention via sequence packing.
|
||||
|
||||
Packs multiple variable-length sequences into a single [1, H, T_total, D]
|
||||
tensor and delegates to the existing block_sparse_attn_from_indices kernel
|
||||
in a single launch. No kernel modifications required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Sequence
|
||||
|
||||
import torch
|
||||
|
||||
from .block_sparse_attn import block_sparse_attn_from_indices
|
||||
|
||||
BLOCK_SIZE = 64
|
||||
|
||||
|
||||
def _scatter_to_padded(
|
||||
src: torch.Tensor,
|
||||
block_sizes: torch.Tensor,
|
||||
block_size: int,
|
||||
dst: torch.Tensor,
|
||||
dst_offset: int,
|
||||
src_start: int,
|
||||
src_end: int,
|
||||
) -> None:
|
||||
"""Copy tokens from a flat source into block-aligned positions in dst.
|
||||
|
||||
Each block occupies exactly `block_size` slots in dst. The first
|
||||
`block_sizes[b]` slots of block *b* receive real tokens; the remainder
|
||||
stays zero (padding the kernel expects).
|
||||
|
||||
src: [total_tokens, H, D]
|
||||
dst: [1, H, total_padded, D]
|
||||
block_sizes: [num_blocks] int32, actual token count per block.
|
||||
"""
|
||||
src_pos = src_start
|
||||
dst_pos = dst_offset
|
||||
sizes = block_sizes.cpu().tolist()
|
||||
for actual in sizes:
|
||||
actual = min(actual, src_end - src_pos)
|
||||
if actual > 0:
|
||||
dst[:, :, dst_pos:dst_pos + actual, :] = (
|
||||
src[src_pos:src_pos + actual].transpose(0, 1).unsqueeze(0)
|
||||
)
|
||||
src_pos += actual
|
||||
dst_pos += block_size
|
||||
|
||||
|
||||
def _gather_from_padded(
|
||||
src: torch.Tensor,
|
||||
block_sizes: torch.Tensor,
|
||||
block_size: int,
|
||||
dst: torch.Tensor,
|
||||
src_offset: int,
|
||||
dst_start: int,
|
||||
dst_end: int,
|
||||
) -> None:
|
||||
"""Inverse of _scatter_to_padded: extract real tokens from padded blocks.
|
||||
|
||||
src: [1, H, total_padded, D]
|
||||
dst: [total_tokens, H, D]
|
||||
"""
|
||||
src_pos = src_offset
|
||||
dst_pos = dst_start
|
||||
sizes = block_sizes.cpu().tolist()
|
||||
for actual in sizes:
|
||||
actual = min(actual, dst_end - dst_pos)
|
||||
if actual > 0:
|
||||
dst[dst_pos:dst_pos + actual] = (
|
||||
src[0, :, src_pos:src_pos + actual, :].transpose(0, 1)
|
||||
)
|
||||
dst_pos += actual
|
||||
src_pos += block_size
|
||||
|
||||
|
||||
def block_sparse_attn_varlen(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
cu_seqlens_q: torch.Tensor,
|
||||
cu_seqlens_kv: torch.Tensor,
|
||||
q2k_idx_list: Sequence[torch.Tensor],
|
||||
q2k_num_list: Sequence[torch.Tensor],
|
||||
variable_block_sizes_list: Sequence[torch.Tensor],
|
||||
q_variable_block_sizes_list: Sequence[torch.Tensor] | None = None,
|
||||
block_size: int = BLOCK_SIZE,
|
||||
) -> torch.Tensor:
|
||||
"""Block-sparse attention over packed variable-length sequences.
|
||||
|
||||
Args:
|
||||
q: [total_q_tokens, H, D] packed query tensor.
|
||||
k: [total_kv_tokens, H, D] packed key tensor.
|
||||
v: [total_kv_tokens, H, D] packed value tensor.
|
||||
cu_seqlens_q: [N+1] int32, cumulative Q token offsets.
|
||||
cu_seqlens_kv: [N+1] int32, cumulative KV token offsets.
|
||||
q2k_idx_list: Per-sequence q2k_idx tensors, each [1, H, Nq_i, Mk].
|
||||
q2k_num_list: Per-sequence q2k_num tensors, each [1, H, Nq_i].
|
||||
variable_block_sizes_list: Per-sequence KV block sizes, each [Nkv_i].
|
||||
q_variable_block_sizes_list: Per-sequence Q block sizes, each [Nq_i].
|
||||
If None, each Q block is assumed to be exactly `block_size` tokens.
|
||||
block_size: Attention block size (default 64).
|
||||
|
||||
Returns:
|
||||
out: [total_q_tokens, H, D] packed output tensor.
|
||||
"""
|
||||
device = q.device
|
||||
dtype = q.dtype
|
||||
num_heads = q.shape[1]
|
||||
head_dim = q.shape[2]
|
||||
num_seqs = cu_seqlens_q.shape[0] - 1
|
||||
|
||||
cu_q = cu_seqlens_q.cpu().tolist()
|
||||
cu_kv = cu_seqlens_kv.cpu().tolist()
|
||||
|
||||
padded_q_lens = []
|
||||
padded_kv_lens = []
|
||||
q_block_offsets = [0]
|
||||
kv_block_offsets = [0]
|
||||
q_vbs_resolved = []
|
||||
|
||||
for i in range(num_seqs):
|
||||
n_q_blocks = q2k_num_list[i].shape[-1]
|
||||
n_kv_blocks = variable_block_sizes_list[i].numel()
|
||||
padded_q_lens.append(n_q_blocks * block_size)
|
||||
padded_kv_lens.append(n_kv_blocks * block_size)
|
||||
q_block_offsets.append(q_block_offsets[-1] + n_q_blocks)
|
||||
kv_block_offsets.append(kv_block_offsets[-1] + n_kv_blocks)
|
||||
|
||||
if q_variable_block_sizes_list is not None:
|
||||
q_vbs_resolved.append(q_variable_block_sizes_list[i])
|
||||
else:
|
||||
q_vbs_resolved.append(
|
||||
torch.full((n_q_blocks,), block_size, dtype=torch.int32)
|
||||
)
|
||||
|
||||
total_padded_q = sum(padded_q_lens)
|
||||
total_padded_kv = sum(padded_kv_lens)
|
||||
|
||||
q_packed = torch.zeros(1, num_heads, total_padded_q, head_dim, device=device, dtype=dtype)
|
||||
k_packed = torch.zeros(1, num_heads, total_padded_kv, head_dim, device=device, dtype=dtype)
|
||||
v_packed = torch.zeros(1, num_heads, total_padded_kv, head_dim, device=device, dtype=dtype)
|
||||
|
||||
q_offset = 0
|
||||
kv_offset = 0
|
||||
for i in range(num_seqs):
|
||||
_scatter_to_padded(
|
||||
q, q_vbs_resolved[i], block_size,
|
||||
q_packed, q_offset, cu_q[i], cu_q[i + 1],
|
||||
)
|
||||
_scatter_to_padded(
|
||||
k, variable_block_sizes_list[i], block_size,
|
||||
k_packed, kv_offset, cu_kv[i], cu_kv[i + 1],
|
||||
)
|
||||
_scatter_to_padded(
|
||||
v, variable_block_sizes_list[i], block_size,
|
||||
v_packed, kv_offset, cu_kv[i], cu_kv[i + 1],
|
||||
)
|
||||
q_offset += padded_q_lens[i]
|
||||
kv_offset += padded_kv_lens[i]
|
||||
|
||||
total_q_blocks = q_block_offsets[-1]
|
||||
max_kv_per_q = max(t.shape[-1] for t in q2k_idx_list)
|
||||
|
||||
global_q2k_idx = torch.zeros(
|
||||
1, num_heads, total_q_blocks, max_kv_per_q,
|
||||
dtype=torch.int32, device=device,
|
||||
)
|
||||
global_q2k_num = torch.zeros(
|
||||
1, num_heads, total_q_blocks,
|
||||
dtype=torch.int32, device=device,
|
||||
)
|
||||
global_vbs_parts = []
|
||||
|
||||
for i in range(num_seqs):
|
||||
qb_start = q_block_offsets[i]
|
||||
qb_end = q_block_offsets[i + 1]
|
||||
n_q_blocks = qb_end - qb_start
|
||||
kv_offset_blocks = kv_block_offsets[i]
|
||||
|
||||
idx = q2k_idx_list[i]
|
||||
num = q2k_num_list[i]
|
||||
vbs = variable_block_sizes_list[i]
|
||||
|
||||
mk = idx.shape[-1]
|
||||
global_q2k_idx[:, :, qb_start:qb_end, :mk] = idx[:, :, :n_q_blocks, :] + kv_offset_blocks
|
||||
global_q2k_num[:, :, qb_start:qb_end] = num[:, :, :n_q_blocks]
|
||||
global_vbs_parts.append(vbs)
|
||||
|
||||
global_vbs = torch.cat(global_vbs_parts, dim=0).to(torch.int32).contiguous()
|
||||
|
||||
out_packed, _ = block_sparse_attn_from_indices(
|
||||
q_packed, k_packed, v_packed,
|
||||
global_q2k_idx, global_q2k_num, global_vbs,
|
||||
)
|
||||
|
||||
out = torch.zeros(cu_q[-1], num_heads, head_dim, device=device, dtype=dtype)
|
||||
q_offset = 0
|
||||
for i in range(num_seqs):
|
||||
_gather_from_padded(
|
||||
out_packed, q_vbs_resolved[i], block_size,
|
||||
out, q_offset, cu_q[i], cu_q[i + 1],
|
||||
)
|
||||
q_offset += padded_q_lens[i]
|
||||
|
||||
return out
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,55 @@
|
||||
"""Compatibility shim for the legacy non-QAT Triton attention import path.
|
||||
|
||||
Historically callers imported
|
||||
``fastvideo_kernel.triton_kernels.fused_attention`` directly. The shared
|
||||
implementation now lives in ``attn_qat_train.py`` and is parameterized by the
|
||||
``IS_QAT`` flag. This module preserves the original public API for tests and
|
||||
downstream users while always dispatching to the non-QAT configuration.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from .attn_qat_train import attention as _attention
|
||||
|
||||
|
||||
def attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
causal: bool,
|
||||
sm_scale: float,
|
||||
warp_specialize: bool = True,
|
||||
) -> torch.Tensor:
|
||||
"""Run the shared Triton attention kernel in non-QAT mode."""
|
||||
use_qat_qkv_backward = True
|
||||
smooth_k = False
|
||||
is_qat = False
|
||||
two_level_quant_p = False
|
||||
fake_quant_p = False
|
||||
use_high_prec_o = False
|
||||
smooth_q = False
|
||||
use_global_sf_p = False
|
||||
use_global_sf_qkv = False
|
||||
|
||||
return _attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
causal,
|
||||
sm_scale,
|
||||
use_qat_qkv_backward,
|
||||
smooth_k,
|
||||
warp_specialize,
|
||||
is_qat,
|
||||
two_level_quant_p,
|
||||
fake_quant_p,
|
||||
use_high_prec_o,
|
||||
smooth_q,
|
||||
use_global_sf_p,
|
||||
use_global_sf_qkv,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["attention"]
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/triton-lang/triton/blob/main/python/triton_kernels/triton_kernels/numerics_details/mxfp_details/_upcast_from_mxfp.py
|
||||
# and https://github.com/triton-lang/triton/blob/main/python/triton_kernels/triton_kernels/numerics_details/mxfp_details/_downcast_to_mxfp.py
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from triton.language.target_info import cuda_capability_geq
|
||||
|
||||
MXFP_BLOCK_SIZE = tl.constexpr(16)
|
||||
|
||||
@triton.jit
|
||||
def _compute_quant_and_scale(
|
||||
src_tensor,
|
||||
valid_src_mask,
|
||||
mx_tensor_dtype: tl.constexpr = tl.uint8,
|
||||
use_global_sf=True,
|
||||
two_level_quant_P=False,
|
||||
):
|
||||
BLOCK_SIZE_OUT_DIM: tl.constexpr = src_tensor.shape[0]
|
||||
BLOCK_SIZE_QUANT_DIM: tl.constexpr = src_tensor.shape[1]
|
||||
BLOCK_SIZE_QUANT_MX_SCALE: tl.constexpr = src_tensor.shape[1] // MXFP_BLOCK_SIZE
|
||||
is_fp4: tl.constexpr = mx_tensor_dtype == tl.uint8
|
||||
|
||||
tl.static_assert(
|
||||
is_fp4
|
||||
or mx_tensor_dtype == tl.float8e4nv
|
||||
or mx_tensor_dtype == tl.float8e5,
|
||||
"mx_tensor_dtype must be uint8, float8e4nv, or float8e5",
|
||||
)
|
||||
|
||||
# Explicit cast to fp32 since most ops are not supported on bfloat16. We avoid needless conversions to and from bf16
|
||||
f32_tensor = src_tensor.to(tl.float32)
|
||||
abs_tensor = tl.abs(f32_tensor)
|
||||
abs_tensor = tl.where(valid_src_mask, abs_tensor, -1.0) # Don't consider padding tensors in scale computation
|
||||
|
||||
if two_level_quant_P:
|
||||
# row max from SageAttn3 paper
|
||||
global_max_val = tl.max(f32_tensor, axis=1, keep_dims=True) # (BLOCK_SIZE_OUT_DIM, 1)
|
||||
global_max_val = tl.maximum(global_max_val, 1e-8)
|
||||
s_enc = ((6 * 448) / global_max_val).reshape([BLOCK_SIZE_OUT_DIM, 1, 1])
|
||||
s_dec = (1 / s_enc)
|
||||
|
||||
abs_tensor = tl.reshape(abs_tensor, [BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, MXFP_BLOCK_SIZE])
|
||||
|
||||
if use_global_sf and not two_level_quant_P:
|
||||
global_max_val = tl.max(abs_tensor)
|
||||
# Avoid division by zero: if all values are padding (max is 0), use a default scale
|
||||
global_max_val = tl.maximum(global_max_val, 1e-8)
|
||||
s_enc = (6 * 448) / global_max_val
|
||||
s_dec = (1 / s_enc)
|
||||
elif not two_level_quant_P and not use_global_sf:
|
||||
s_dec = 1.0
|
||||
s_enc = 1.0
|
||||
|
||||
max_val = tl.max(abs_tensor, axis=2, keep_dims=True) # (BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1) # per block maxima
|
||||
s_dec_b = max_val / 6 # (BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1)
|
||||
s_dec_b_e4m3 = (s_dec_b * s_enc).to(tl.float8e4nv) # (BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1)
|
||||
s_enc_b = 1 / (s_dec_b_e4m3.to(tl.float32) * s_dec) # (BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1)
|
||||
|
||||
f32_tensor = tl.reshape(f32_tensor, [BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, MXFP_BLOCK_SIZE])
|
||||
quant_tensor = f32_tensor * s_enc_b
|
||||
|
||||
# Reshape the tensors after scaling
|
||||
quant_tensor = quant_tensor.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_DIM])
|
||||
# Set the invalid portions of the tensor to 0. This will ensure that any padding tensors are 0 in the mx format.
|
||||
quant_tensor = tl.where(valid_src_mask, quant_tensor, 0.0)
|
||||
dequant_scale = s_dec_b_e4m3.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE])
|
||||
|
||||
if is_fp4 and cuda_capability_geq(10, 0):
|
||||
# Convert scaled values to two f32 lanes and use PTX cvt to e2m1x2 with two f32 operands.
|
||||
pairs = tl.reshape(quant_tensor, [BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_DIM // 2, 2])
|
||||
lo_f, hi_f = tl.split(pairs)
|
||||
lo_f32 = lo_f.to(tl.float32)
|
||||
hi_f32 = hi_f.to(tl.float32)
|
||||
|
||||
# Inline PTX: cvt.rn.satfinite.e2m1x2.f32 takes two f32 sources and produces one .b8 packed e2m1x2.
|
||||
out_tensor = tl.inline_asm_elementwise(
|
||||
"""
|
||||
{
|
||||
.reg .b8 r;
|
||||
cvt.rn.satfinite.e2m1x2.f32 r, $1, $2;
|
||||
mov.b32 $0, {r, r, r, r};
|
||||
}
|
||||
""",
|
||||
constraints="=r,f,f",
|
||||
args=[hi_f32, lo_f32],
|
||||
dtype=tl.uint8,
|
||||
is_pure=True,
|
||||
pack=1,
|
||||
)
|
||||
elif is_fp4:
|
||||
quant_tensor = quant_tensor.to(tl.uint32, bitcast=True)
|
||||
signs = quant_tensor & 0x80000000
|
||||
exponents = (quant_tensor >> 23) & 0xFF
|
||||
mantissas_orig = (quant_tensor & 0x7FFFFF)
|
||||
|
||||
# For RTNE: 0.25 < x < 0.75 maps to 0.5 (denormal); exactly 0.25 maps to 0.0
|
||||
E8_BIAS = 127
|
||||
E2_BIAS = 1
|
||||
# Move implicit bit 1 at the beginning to mantissa for denormals
|
||||
is_subnormal = exponents < E8_BIAS
|
||||
adjusted_exponents = tl.core.sub(E8_BIAS, exponents + 1, sanitize_overflow=False)
|
||||
mantissas_pre = (0x400000 | (mantissas_orig >> 1))
|
||||
mantissas = tl.where(is_subnormal, mantissas_pre >> adjusted_exponents, mantissas_orig)
|
||||
|
||||
# For normal numbers, we change the bias from 127 to 1, and for subnormals, we keep exponent as 0.
|
||||
exponents = tl.maximum(exponents, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
|
||||
|
||||
# Combine sign, exponent, and mantissa, while saturating
|
||||
# Round to nearest, ties to even (RTNE): use guard/sticky and LSB to decide increment
|
||||
m2bits = mantissas >> 21
|
||||
lsb_keep = (m2bits >> 1) & 0x1
|
||||
guard = m2bits & 0x1
|
||||
IS_SRC_FP32: tl.constexpr = src_tensor.dtype == tl.float32
|
||||
if IS_SRC_FP32:
|
||||
bit0_dropped = (mantissas_orig & 0x1) != 0
|
||||
mask = (1 << tl.minimum(adjusted_exponents, 31)) - 1
|
||||
dropped_post = (mantissas_pre & mask) != 0
|
||||
sticky = is_subnormal & (bit0_dropped | dropped_post)
|
||||
sticky |= ((mantissas & 0x1FFFFF) != 0).to(tl.uint32)
|
||||
else:
|
||||
sticky = ((mantissas & 0x1FFFFF) != 0).to(tl.uint32)
|
||||
round_inc = guard & (sticky | lsb_keep)
|
||||
e2m1_tmp = tl.minimum((((exponents << 2) | m2bits) + round_inc) >> 1, 0x7)
|
||||
e2m1_value = ((signs >> 28) | e2m1_tmp).to(tl.uint8)
|
||||
|
||||
e2m1_value = tl.reshape(e2m1_value, [BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_DIM // 2, 2])
|
||||
evens, odds = tl.split(e2m1_value)
|
||||
out_tensor = evens | (odds << 4)
|
||||
else:
|
||||
out_tensor = quant_tensor.to(mx_tensor_dtype)
|
||||
|
||||
return out_tensor, dequant_scale, s_dec
|
||||
|
||||
@triton.jit
|
||||
def _compute_dequant(
|
||||
mx_tensor,
|
||||
scale,
|
||||
s_dec,
|
||||
BLOCK_SIZE_OUT_DIM: tl.constexpr,
|
||||
BLOCK_SIZE_QUANT_DIM: tl.constexpr,
|
||||
dst_dtype: tl.constexpr,
|
||||
):
|
||||
tl.static_assert(BLOCK_SIZE_QUANT_DIM % MXFP_BLOCK_SIZE == 0, f"Block size along quantization block must be a multiple of {MXFP_BLOCK_SIZE=}")
|
||||
# uint8 signifies two fp4 e2m1 values packed into a single byte
|
||||
mx_tensor_dtype: tl.constexpr = mx_tensor.dtype
|
||||
tl.static_assert(dst_dtype == tl.float16 or dst_dtype == tl.bfloat16 or dst_dtype == tl.float32)
|
||||
tl.static_assert(
|
||||
mx_tensor_dtype == tl.uint8
|
||||
or ((mx_tensor_dtype == tl.float8e4nv or mx_tensor_dtype == tl.float8e5) or mx_tensor_dtype == dst_dtype),
|
||||
"mx_tensor_ptr must be uint8 or float8 or dst_dtype")
|
||||
tl.static_assert(scale.dtype == tl.float8e4nv, "scale must be float8e4nv")
|
||||
|
||||
# Determine if we are dealing with fp8 types.
|
||||
is_fp4: tl.constexpr = mx_tensor_dtype == tl.uint8
|
||||
BLOCK_SIZE_QUANT_MX_SCALE: tl.constexpr = BLOCK_SIZE_QUANT_DIM // MXFP_BLOCK_SIZE
|
||||
|
||||
# Upcast the scale to the destination type.
|
||||
if dst_dtype == tl.bfloat16:
|
||||
dst_scale = scale.to(tl.bfloat16)
|
||||
else:
|
||||
dst_scale = scale.to(tl.float32)
|
||||
if dst_dtype == tl.float16:
|
||||
dst_scale = dst_scale.to(tl.float16)
|
||||
|
||||
# Now upcast the tensor.
|
||||
intermediate_dtype: tl.constexpr = tl.bfloat16 if dst_dtype == tl.float32 else dst_dtype
|
||||
if cuda_capability_geq(10, 0):
|
||||
assert is_fp4
|
||||
packed_u32 = tl.inline_asm_elementwise(
|
||||
asm="""
|
||||
{
|
||||
.reg .b8 in_8;
|
||||
.reg .f16x2 out;
|
||||
cvt.u8.u32 in_8, $1;
|
||||
cvt.rn.f16x2.e2m1x2 out, in_8;
|
||||
mov.b32 $0, out;
|
||||
}
|
||||
""",
|
||||
constraints="=r,r",
|
||||
args=[mx_tensor], # tl.uint8 passed in as a 32-bit reg with value in low 8 bits
|
||||
dtype=tl.uint32,
|
||||
is_pure=True,
|
||||
pack=1,
|
||||
)
|
||||
lo_u16 = (packed_u32 & 0xFFFF).to(tl.uint16)
|
||||
hi_u16 = (packed_u32 >> 16).to(tl.uint16)
|
||||
lo_f16 = lo_u16.to(tl.float16, bitcast=True)
|
||||
hi_f16 = hi_u16.to(tl.float16, bitcast=True)
|
||||
|
||||
if intermediate_dtype == tl.float16:
|
||||
x0, x1 = lo_f16, hi_f16
|
||||
else:
|
||||
x0 = lo_f16.to(intermediate_dtype)
|
||||
x1 = hi_f16.to(intermediate_dtype)
|
||||
|
||||
dst_tensor = tl.interleave(x0, x1)
|
||||
|
||||
else:
|
||||
assert is_fp4
|
||||
dst_bias: tl.constexpr = 127 if intermediate_dtype == tl.bfloat16 else 15 # exponent bias
|
||||
dst_0p5: tl.constexpr = 16128 if intermediate_dtype == tl.bfloat16 else 0x3800
|
||||
dst_m_bits: tl.constexpr = 7 if intermediate_dtype == tl.bfloat16 else 10 # mantissa bits
|
||||
# e2m1
|
||||
em0 = mx_tensor & 0x07
|
||||
em1 = mx_tensor & 0x70
|
||||
x0 = (em0.to(tl.uint16) << (dst_m_bits - 1)) | ((mx_tensor & 0x08).to(tl.uint16) << 12)
|
||||
x1 = (em1.to(tl.uint16) << (dst_m_bits - 5)) | ((mx_tensor & 0x80).to(tl.uint16) << 8)
|
||||
# Three cases:
|
||||
# 1) x is normal and non-zero: Correct bias
|
||||
x0 = tl.where((em0 & 0x06) != 0, x0 + ((dst_bias - 1) << dst_m_bits), x0)
|
||||
x1 = tl.where((em1 & 0x60) != 0, x1 + ((dst_bias - 1) << dst_m_bits), x1)
|
||||
# 2) x is subnormal (x == 0bs001 where s is the sign): Map to +-0.5 in the dst type
|
||||
x0 = tl.where(em0 == 0x01, dst_0p5 | (x0 & 0x8000), x0)
|
||||
x1 = tl.where(em1 == 0x10, dst_0p5 | (x1 & 0x8000), x1)
|
||||
# 3) x is zero, do nothing
|
||||
dst_tensor = tl.interleave(x0, x1).to(intermediate_dtype, bitcast=True)
|
||||
|
||||
dst_tensor = dst_tensor.to(dst_dtype)
|
||||
|
||||
# Reshape for proper broadcasting: the scale was stored with a 16‐sized “inner” grouping.
|
||||
dst_tensor = dst_tensor.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, MXFP_BLOCK_SIZE])
|
||||
dst_scale = dst_scale.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1])
|
||||
scale = scale.reshape(dst_scale.shape)
|
||||
|
||||
out_tensor = dst_tensor * dst_scale * s_dec # NVFP4 has the additional global scale factor
|
||||
if dst_dtype == tl.float32:
|
||||
max_fin = 3.4028234663852886e+38
|
||||
elif dst_dtype == tl.bfloat16:
|
||||
max_fin = 3.3895313892515355e+38
|
||||
else:
|
||||
tl.static_assert(dst_dtype == tl.float16)
|
||||
max_fin = 65504
|
||||
out_tensor = tl.clamp(out_tensor, min=-max_fin, max=max_fin)
|
||||
out_tensor = out_tensor.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_DIM])
|
||||
out_tensor = out_tensor.to(dst_dtype)
|
||||
return out_tensor
|
||||
@@ -0,0 +1,80 @@
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from .nvfp4_utils import _compute_quant_and_scale, _compute_dequant
|
||||
|
||||
@triton.jit
|
||||
def fake_quantize(src_tensor, valid_src_mask, BLOCK_SIZE_OUT_DIM: tl.constexpr,
|
||||
BLOCK_SIZE_QUANT_DIM: tl.constexpr,
|
||||
dst_dtype: tl.constexpr,
|
||||
mx_tensor_dtype: tl.constexpr = tl.uint8,
|
||||
use_global_sf: tl.constexpr = True,
|
||||
two_level_quant_P: tl.constexpr = False):
|
||||
high_prec_src_tensor = src_tensor
|
||||
src_tensor, src_scale, src_s_dec = _compute_quant_and_scale(src_tensor=src_tensor,
|
||||
valid_src_mask=valid_src_mask,
|
||||
mx_tensor_dtype=mx_tensor_dtype,
|
||||
use_global_sf=use_global_sf,
|
||||
two_level_quant_P=two_level_quant_P)
|
||||
src_tensor = _compute_dequant(mx_tensor=src_tensor,
|
||||
scale=src_scale,
|
||||
s_dec=src_s_dec,
|
||||
BLOCK_SIZE_OUT_DIM=BLOCK_SIZE_OUT_DIM,
|
||||
BLOCK_SIZE_QUANT_DIM=BLOCK_SIZE_QUANT_DIM,
|
||||
dst_dtype=dst_dtype)
|
||||
return src_tensor, high_prec_src_tensor.to(src_tensor.dtype)
|
||||
|
||||
@triton.jit
|
||||
def fake_quantize_q(Q, fake_Q, stride_z_q, stride_h_q,
|
||||
stride_tok_q, stride_d_q,
|
||||
fake_stride_z_q, fake_stride_h_q,
|
||||
fake_stride_tok_q, fake_stride_d_q,
|
||||
H, N_CTX_Q,
|
||||
BLOCK_M: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
use_global_sf: tl.constexpr = True):
|
||||
bhid = tl.program_id(1)
|
||||
adj_q = (stride_h_q * (bhid % H) + stride_z_q * (bhid // H))
|
||||
fake_adj_q = (fake_stride_h_q * (bhid % H) + fake_stride_z_q * (bhid // H))
|
||||
Q += adj_q
|
||||
fake_Q += fake_adj_q
|
||||
|
||||
pid = tl.program_id(0)
|
||||
start_m = pid * BLOCK_M
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
q_valid = offs_m < N_CTX_Q
|
||||
q = tl.load(Q + offs_m[:, None] * stride_tok_q + offs_k[None, :] * stride_d_q, mask=q_valid[:, None], other=0.0)
|
||||
q, _ = fake_quantize(src_tensor=q, valid_src_mask=q_valid[:, None], BLOCK_SIZE_OUT_DIM=BLOCK_M, BLOCK_SIZE_QUANT_DIM=HEAD_DIM, dst_dtype=q.dtype, use_global_sf=use_global_sf)
|
||||
tl.store(fake_Q + offs_m[:, None] * fake_stride_tok_q + offs_k[None, :] * fake_stride_d_q, q, mask=q_valid[:, None])
|
||||
|
||||
@triton.jit
|
||||
def fake_quantize_kv(K, V, fake_K, fake_V, stride_z_kv, stride_h_kv,
|
||||
stride_tok_kv, stride_d_kv,
|
||||
fake_stride_z_kv, fake_stride_h_kv,
|
||||
fake_stride_tok_kv, fake_stride_d_kv,
|
||||
H, N_CTX_KV,
|
||||
BLOCK_N: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
use_global_sf: tl.constexpr = True):
|
||||
bhid = tl.program_id(1)
|
||||
adj_kv = (stride_h_kv * (bhid % H) + stride_z_kv * (bhid // H))
|
||||
fake_adj_kv = (fake_stride_h_kv * (bhid % H) + fake_stride_z_kv * (bhid // H))
|
||||
K += adj_kv
|
||||
V += adj_kv
|
||||
fake_K += fake_adj_kv
|
||||
fake_V += fake_adj_kv
|
||||
|
||||
pid = tl.program_id(0)
|
||||
start_n = pid * BLOCK_N
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
kv_valid = offs_n < N_CTX_KV
|
||||
k_block = tl.load(K + offs_n[:, None] * stride_tok_kv + offs_k[None, :] * stride_d_kv, mask=kv_valid[:, None], other=0.0)
|
||||
v_block = tl.load(V + offs_n[:, None] * stride_tok_kv + offs_k[None, :] * stride_d_kv, mask=kv_valid[:, None], other=0.0)
|
||||
k, _ = fake_quantize(src_tensor=k_block, valid_src_mask=kv_valid[:, None], BLOCK_SIZE_OUT_DIM=BLOCK_N, BLOCK_SIZE_QUANT_DIM=HEAD_DIM, dst_dtype=k_block.dtype, use_global_sf=use_global_sf)
|
||||
v, _ = fake_quantize(src_tensor=v_block, valid_src_mask=kv_valid[:, None], BLOCK_SIZE_OUT_DIM=BLOCK_N, BLOCK_SIZE_QUANT_DIM=HEAD_DIM, dst_dtype=v_block.dtype, use_global_sf=use_global_sf)
|
||||
tl.store(fake_K + offs_n[:, None] * fake_stride_tok_kv + offs_k[None, :] * fake_stride_d_kv, k, mask=kv_valid[:, None])
|
||||
tl.store(fake_V + offs_n[:, None] * fake_stride_tok_kv + offs_k[None, :] * fake_stride_d_kv, v, mask=kv_valid[:, None])
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,434 @@
|
||||
"""Correctness tests for variable-length block-sparse attention.
|
||||
|
||||
Reference: per-sequence calls to block_sparse_attn_from_indices.
|
||||
Test: single-launch via block_sparse_attn_varlen.
|
||||
Tests cover both forward and backward (gradient) correctness.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
|
||||
from .test_vsa import (
|
||||
BLOCK_M,
|
||||
generate_variable_block_sizes,
|
||||
get_non_pad_index,
|
||||
vsa_pad,
|
||||
generate_tensor,
|
||||
)
|
||||
from .utils import generate_block_sparse_mask_for_function
|
||||
from fastvideo_kernel.block_sparse_attn import (
|
||||
block_sparse_attn_from_indices,
|
||||
_map_to_index,
|
||||
)
|
||||
from fastvideo_kernel.block_sparse_attn_varlen import block_sparse_attn_varlen
|
||||
|
||||
|
||||
def _reference_per_sequence(
|
||||
q_list, k_list, v_list,
|
||||
block_masks, vbs_list,
|
||||
non_pad_q_list, non_pad_kv_list,
|
||||
q_nblocks_list, kv_nblocks_list,
|
||||
):
|
||||
"""Run per-sequence block_sparse_attn and concat outputs."""
|
||||
outs = []
|
||||
for i in range(len(q_list)):
|
||||
q_pad = vsa_pad(q_list[i], non_pad_q_list[i], q_nblocks_list[i], BLOCK_M)
|
||||
k_pad = vsa_pad(k_list[i], non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
|
||||
v_pad = vsa_pad(v_list[i], non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
|
||||
|
||||
q2k_idx, q2k_num = _map_to_index(block_masks[i].unsqueeze(0))
|
||||
o_pad, _ = block_sparse_attn_from_indices(
|
||||
q_pad, k_pad, v_pad, q2k_idx, q2k_num, vbs_list[i],
|
||||
)
|
||||
o = o_pad[:, :, non_pad_q_list[i], :]
|
||||
outs.append(o.squeeze(0).transpose(0, 1))
|
||||
return torch.cat(outs, dim=0)
|
||||
|
||||
|
||||
def _run_varlen_test(
|
||||
seq_configs: list,
|
||||
h: int = 8,
|
||||
d: int = 64,
|
||||
topk: int = 2,
|
||||
atol: float = 0.05,
|
||||
rtol: float = 0.02,
|
||||
):
|
||||
"""Core test: compare varlen vs per-sequence reference.
|
||||
|
||||
seq_configs: list of (num_q_blocks, num_kv_blocks) per sequence.
|
||||
"""
|
||||
device = "cuda"
|
||||
num_seqs = len(seq_configs)
|
||||
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
block_masks = []
|
||||
vbs_list = []
|
||||
q_vbs_list = []
|
||||
non_pad_q_list = []
|
||||
non_pad_kv_list = []
|
||||
q_nblocks_list = []
|
||||
kv_nblocks_list = []
|
||||
q2k_idx_list = []
|
||||
q2k_num_list = []
|
||||
q_vbs_for_varlen = []
|
||||
|
||||
cu_q = [0]
|
||||
cu_kv = [0]
|
||||
|
||||
for nq, nkv in seq_configs:
|
||||
vbs_kv = generate_variable_block_sizes(nkv, device=device)
|
||||
vbs_q = generate_variable_block_sizes(nq, device=device)
|
||||
sq = int(vbs_q.sum().item())
|
||||
skv = int(vbs_kv.sum().item())
|
||||
|
||||
q = generate_tensor((1, h, sq, d), torch.bfloat16, device)
|
||||
k = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
v = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
|
||||
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
|
||||
npq = get_non_pad_index(vbs_q, nq, BLOCK_M)
|
||||
npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M)
|
||||
|
||||
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
|
||||
|
||||
q_list.append(q)
|
||||
k_list.append(k)
|
||||
v_list.append(v)
|
||||
block_masks.append(mask)
|
||||
vbs_list.append(vbs_kv)
|
||||
q_vbs_list.append(vbs_q)
|
||||
non_pad_q_list.append(npq)
|
||||
non_pad_kv_list.append(npkv)
|
||||
q_nblocks_list.append(nq)
|
||||
kv_nblocks_list.append(nkv)
|
||||
q2k_idx_list.append(q2k_idx)
|
||||
q2k_num_list.append(q2k_num)
|
||||
q_vbs_for_varlen.append(vbs_q)
|
||||
|
||||
cu_q.append(cu_q[-1] + sq)
|
||||
cu_kv.append(cu_kv[-1] + skv)
|
||||
|
||||
ref_out = _reference_per_sequence(
|
||||
q_list, k_list, v_list,
|
||||
block_masks, vbs_list,
|
||||
non_pad_q_list, non_pad_kv_list,
|
||||
q_nblocks_list, kv_nblocks_list,
|
||||
)
|
||||
|
||||
q_packed = torch.cat(
|
||||
[qi.squeeze(0).transpose(0, 1) for qi in q_list], dim=0,
|
||||
)
|
||||
k_packed = torch.cat(
|
||||
[ki.squeeze(0).transpose(0, 1) for ki in k_list], dim=0,
|
||||
)
|
||||
v_packed = torch.cat(
|
||||
[vi.squeeze(0).transpose(0, 1) for vi in v_list], dim=0,
|
||||
)
|
||||
|
||||
cu_seqlens_q = torch.tensor(cu_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_kv = torch.tensor(cu_kv, dtype=torch.int32, device=device)
|
||||
|
||||
varlen_out = block_sparse_attn_varlen(
|
||||
q_packed, k_packed, v_packed,
|
||||
cu_seqlens_q, cu_seqlens_kv,
|
||||
q2k_idx_list, q2k_num_list,
|
||||
vbs_list,
|
||||
q_variable_block_sizes_list=q_vbs_for_varlen,
|
||||
)
|
||||
|
||||
max_abs = (ref_out - varlen_out).abs().max().item()
|
||||
mean_abs = ref_out.abs().mean().item()
|
||||
max_rel = max_abs / (mean_abs + 1e-8)
|
||||
|
||||
print(f" seqs={[c for c in seq_configs]}, max_abs={max_abs:.4e}, max_rel={max_rel:.4e}")
|
||||
assert max_rel < rtol, f"max relative error {max_rel:.4e} exceeds threshold {rtol}"
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
class TestVSAVarlen:
|
||||
|
||||
def test_equal_length(self):
|
||||
"""Two sequences with same number of blocks."""
|
||||
_run_varlen_test([(4, 4), (4, 4)], h=8, d=64)
|
||||
|
||||
def test_different_lengths(self):
|
||||
"""Three sequences with different block counts."""
|
||||
_run_varlen_test([(2, 3), (5, 4), (3, 6)], h=8, d=64)
|
||||
|
||||
def test_single_sequence(self):
|
||||
"""Degenerate case: single sequence should match non-varlen path."""
|
||||
_run_varlen_test([(8, 8)], h=8, d=64)
|
||||
|
||||
def test_many_heads(self):
|
||||
"""More heads to stress the packing logic."""
|
||||
_run_varlen_test([(3, 4), (5, 3)], h=16, d=128)
|
||||
|
||||
def test_many_sequences(self):
|
||||
"""Stress test: 8 sequences with varying block counts."""
|
||||
configs = [(i + 2, i + 3) for i in range(8)]
|
||||
_run_varlen_test(configs, h=8, d=64)
|
||||
|
||||
def test_topk_equals_num_blocks(self):
|
||||
"""Edge: topk covers all KV blocks (dense attention)."""
|
||||
_run_varlen_test([(3, 3), (4, 4)], h=8, d=64, topk=8)
|
||||
|
||||
def test_single_block_per_sequence(self):
|
||||
"""Minimal: each sequence has exactly 1 Q block and 1 KV block."""
|
||||
_run_varlen_test([(1, 1), (1, 1), (1, 1)], h=8, d=64, topk=1)
|
||||
|
||||
def test_asymmetric_q_kv(self):
|
||||
"""Q and KV have very different block counts."""
|
||||
_run_varlen_test([(1, 8), (8, 1)], h=8, d=64, topk=1)
|
||||
|
||||
def test_without_q_vbs(self):
|
||||
"""Test the default path where q_variable_block_sizes_list is None.
|
||||
|
||||
Uses full block_size=64 for Q blocks so the None path is valid.
|
||||
"""
|
||||
device = "cuda"
|
||||
h, d, topk = 4, 64, 2
|
||||
nq, nkv = 3, 4
|
||||
|
||||
vbs_kv = generate_variable_block_sizes(nkv, device=device)
|
||||
sq = nq * BLOCK_M
|
||||
skv = int(vbs_kv.sum().item())
|
||||
|
||||
q = generate_tensor((1, h, sq, d), torch.bfloat16, device)
|
||||
k = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
v = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
|
||||
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
|
||||
npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M)
|
||||
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
|
||||
|
||||
k_pad = vsa_pad(k, npkv, nkv, BLOCK_M)
|
||||
v_pad = vsa_pad(v, npkv, nkv, BLOCK_M)
|
||||
ref_out, _ = block_sparse_attn_from_indices(q, k_pad, v_pad, q2k_idx, q2k_num, vbs_kv)
|
||||
ref_flat = ref_out.squeeze(0).transpose(0, 1)
|
||||
|
||||
q_flat = q.squeeze(0).transpose(0, 1)
|
||||
k_flat = k.squeeze(0).transpose(0, 1)
|
||||
v_flat = v.squeeze(0).transpose(0, 1)
|
||||
cu_q = torch.tensor([0, sq], dtype=torch.int32, device=device)
|
||||
cu_kv = torch.tensor([0, skv], dtype=torch.int32, device=device)
|
||||
|
||||
varlen_out = block_sparse_attn_varlen(
|
||||
q_flat, k_flat, v_flat,
|
||||
cu_q, cu_kv,
|
||||
[q2k_idx], [q2k_num], [vbs_kv],
|
||||
)
|
||||
|
||||
max_abs = (ref_flat - varlen_out).abs().max().item()
|
||||
mean_abs = ref_flat.abs().mean().item()
|
||||
max_rel = max_abs / (mean_abs + 1e-8)
|
||||
print(f" without_q_vbs: max_abs={max_abs:.4e}, max_rel={max_rel:.4e}")
|
||||
assert max_rel < 0.02, f"max relative error {max_rel:.4e} exceeds threshold"
|
||||
|
||||
|
||||
def _run_varlen_backward_test(
|
||||
seq_configs: list,
|
||||
h: int = 8,
|
||||
d: int = 64,
|
||||
topk: int = 2,
|
||||
grad_rtol: float = 0.05,
|
||||
):
|
||||
"""Backward correctness: compare dQ/dK/dV from varlen vs per-sequence reference.
|
||||
|
||||
Both paths use the same underlying block_sparse_attn_from_indices kernel
|
||||
(which has registered autograd). The varlen wrapper's scatter/gather must
|
||||
correctly propagate gradients through PyTorch's in-place slice assignment.
|
||||
"""
|
||||
device = "cuda"
|
||||
num_seqs = len(seq_configs)
|
||||
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
block_masks = []
|
||||
vbs_list = []
|
||||
q_vbs_list = []
|
||||
non_pad_q_list = []
|
||||
non_pad_kv_list = []
|
||||
q_nblocks_list = []
|
||||
kv_nblocks_list = []
|
||||
q2k_idx_list = []
|
||||
q2k_num_list = []
|
||||
|
||||
cu_q = [0]
|
||||
cu_kv = [0]
|
||||
|
||||
for nq, nkv in seq_configs:
|
||||
vbs_kv = generate_variable_block_sizes(nkv, device=device)
|
||||
vbs_q = generate_variable_block_sizes(nq, device=device)
|
||||
sq = int(vbs_q.sum().item())
|
||||
skv = int(vbs_kv.sum().item())
|
||||
|
||||
q = generate_tensor((1, h, sq, d), torch.bfloat16, device)
|
||||
k = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
v = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
|
||||
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
|
||||
npq = get_non_pad_index(vbs_q, nq, BLOCK_M)
|
||||
npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M)
|
||||
|
||||
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
|
||||
|
||||
q_list.append(q)
|
||||
k_list.append(k)
|
||||
v_list.append(v)
|
||||
block_masks.append(mask)
|
||||
vbs_list.append(vbs_kv)
|
||||
q_vbs_list.append(vbs_q)
|
||||
non_pad_q_list.append(npq)
|
||||
non_pad_kv_list.append(npkv)
|
||||
q_nblocks_list.append(nq)
|
||||
kv_nblocks_list.append(nkv)
|
||||
q2k_idx_list.append(q2k_idx)
|
||||
q2k_num_list.append(q2k_num)
|
||||
|
||||
cu_q.append(cu_q[-1] + sq)
|
||||
cu_kv.append(cu_kv[-1] + skv)
|
||||
|
||||
# --- Reference: per-sequence backward ---
|
||||
ref_q_grads = []
|
||||
ref_k_grads = []
|
||||
ref_v_grads = []
|
||||
ref_outs = []
|
||||
for i in range(num_seqs):
|
||||
qi = q_list[i].detach().requires_grad_(True)
|
||||
ki = k_list[i].detach().requires_grad_(True)
|
||||
vi = v_list[i].detach().requires_grad_(True)
|
||||
|
||||
q_pad = vsa_pad(qi, non_pad_q_list[i], q_nblocks_list[i], BLOCK_M)
|
||||
k_pad = vsa_pad(ki, non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
|
||||
v_pad = vsa_pad(vi, non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
|
||||
|
||||
q2k_idx, q2k_num = _map_to_index(block_masks[i].unsqueeze(0))
|
||||
o_pad, _ = block_sparse_attn_from_indices(
|
||||
q_pad, k_pad, v_pad, q2k_idx, q2k_num, vbs_list[i],
|
||||
)
|
||||
o = o_pad[:, :, non_pad_q_list[i], :]
|
||||
o_flat = o.squeeze(0).transpose(0, 1)
|
||||
ref_outs.append(o_flat)
|
||||
|
||||
dO = torch.ones_like(o_flat)
|
||||
o_flat.backward(dO)
|
||||
|
||||
ref_q_grads.append(qi.grad.squeeze(0).transpose(0, 1))
|
||||
ref_k_grads.append(ki.grad.squeeze(0).transpose(0, 1))
|
||||
ref_v_grads.append(vi.grad.squeeze(0).transpose(0, 1))
|
||||
|
||||
ref_dq = torch.cat(ref_q_grads, dim=0)
|
||||
ref_dk = torch.cat(ref_k_grads, dim=0)
|
||||
ref_dv = torch.cat(ref_v_grads, dim=0)
|
||||
|
||||
# --- Varlen backward ---
|
||||
q_packed = torch.cat(
|
||||
[qi.squeeze(0).transpose(0, 1) for qi in q_list], dim=0,
|
||||
).detach().requires_grad_(True)
|
||||
k_packed = torch.cat(
|
||||
[ki.squeeze(0).transpose(0, 1) for ki in k_list], dim=0,
|
||||
).detach().requires_grad_(True)
|
||||
v_packed = torch.cat(
|
||||
[vi.squeeze(0).transpose(0, 1) for vi in v_list], dim=0,
|
||||
).detach().requires_grad_(True)
|
||||
|
||||
cu_seqlens_q = torch.tensor(cu_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_kv = torch.tensor(cu_kv, dtype=torch.int32, device=device)
|
||||
|
||||
varlen_out = block_sparse_attn_varlen(
|
||||
q_packed, k_packed, v_packed,
|
||||
cu_seqlens_q, cu_seqlens_kv,
|
||||
q2k_idx_list, q2k_num_list,
|
||||
vbs_list,
|
||||
q_variable_block_sizes_list=q_vbs_list,
|
||||
)
|
||||
|
||||
dO = torch.ones_like(varlen_out)
|
||||
varlen_out.backward(dO)
|
||||
|
||||
varlen_dq = q_packed.grad
|
||||
varlen_dk = k_packed.grad
|
||||
varlen_dv = v_packed.grad
|
||||
|
||||
for name, ref, actual in [
|
||||
("dQ", ref_dq, varlen_dq),
|
||||
("dK", ref_dk, varlen_dk),
|
||||
("dV", ref_dv, varlen_dv),
|
||||
]:
|
||||
assert actual is not None, f"{name}: gradient is None (autograd chain broken)"
|
||||
max_abs = (ref - actual).abs().max().item()
|
||||
mean_abs = ref.abs().mean().item()
|
||||
max_rel = max_abs / (mean_abs + 1e-8)
|
||||
print(f" {name}: max_abs={max_abs:.4e}, max_rel={max_rel:.4e}")
|
||||
assert max_rel < grad_rtol, (
|
||||
f"{name}: max relative error {max_rel:.4e} exceeds threshold {grad_rtol}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
class TestVSAVarlenBackward:
|
||||
|
||||
def test_backward_equal_length(self):
|
||||
"""Backward: two sequences with same number of blocks."""
|
||||
_run_varlen_backward_test([(4, 4), (4, 4)], h=8, d=64)
|
||||
|
||||
def test_backward_different_lengths(self):
|
||||
"""Backward: three sequences with different block counts."""
|
||||
_run_varlen_backward_test([(2, 3), (5, 4), (3, 6)], h=8, d=64)
|
||||
|
||||
def test_backward_single_sequence(self):
|
||||
"""Backward: single sequence should match non-varlen gradient path."""
|
||||
_run_varlen_backward_test([(8, 8)], h=8, d=64)
|
||||
|
||||
def test_backward_many_heads(self):
|
||||
"""Backward: more heads to stress gradient routing."""
|
||||
_run_varlen_backward_test([(3, 4), (5, 3)], h=16, d=128)
|
||||
|
||||
def test_backward_asymmetric_q_kv(self):
|
||||
"""Backward: Q and KV have very different block counts."""
|
||||
_run_varlen_backward_test([(1, 8), (8, 1)], h=8, d=64, topk=1)
|
||||
|
||||
def test_backward_grad_nonzero(self):
|
||||
"""Smoke test: gradients are non-zero (autograd chain is connected)."""
|
||||
device = "cuda"
|
||||
h, d, topk = 4, 64, 2
|
||||
nq, nkv = 3, 4
|
||||
|
||||
vbs_kv = generate_variable_block_sizes(nkv, device=device)
|
||||
vbs_q = generate_variable_block_sizes(nq, device=device)
|
||||
sq = int(vbs_q.sum().item())
|
||||
skv = int(vbs_kv.sum().item())
|
||||
|
||||
q = torch.randn(sq, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
|
||||
k = torch.randn(skv, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
|
||||
v = torch.randn(skv, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
|
||||
|
||||
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
|
||||
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
|
||||
|
||||
cu_q = torch.tensor([0, sq], dtype=torch.int32, device=device)
|
||||
cu_kv = torch.tensor([0, skv], dtype=torch.int32, device=device)
|
||||
|
||||
out = block_sparse_attn_varlen(
|
||||
q, k, v,
|
||||
cu_q, cu_kv,
|
||||
[q2k_idx], [q2k_num], [vbs_kv],
|
||||
q_variable_block_sizes_list=[vbs_q],
|
||||
)
|
||||
|
||||
loss = out.sum()
|
||||
loss.backward()
|
||||
|
||||
assert q.grad is not None, "q.grad is None"
|
||||
assert k.grad is not None, "k.grad is None"
|
||||
assert v.grad is not None, "v.grad is None"
|
||||
assert q.grad.abs().sum().item() > 0, "q.grad is all zeros"
|
||||
assert k.grad.abs().sum().item() > 0, "k.grad is all zeros"
|
||||
assert v.grad.abs().sum().item() > 0, "v.grad is all zeros"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v", "-s"])
|
||||
@@ -1,624 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Typed omni request plane (design.md §6.1).
|
||||
|
||||
This module introduces the request-plane vocabulary the next-generation
|
||||
runtime is built on:
|
||||
|
||||
- :class:`OmniRequest` — a typed multimodal request whose ``task`` is
|
||||
*declared, never inferred*, and whose inputs are typed
|
||||
:class:`ModalPart` s instead of per-model fields on a god-object.
|
||||
- :class:`OmniOutput` — a typed multimodal output whose modalities are
|
||||
named :class:`Artifact` slots carrying provenance, replacing the
|
||||
``extra["audio"]`` escape hatch (design.md P3).
|
||||
- :data:`OmniEvent` — one streaming-event union (progress / chunk /
|
||||
final), the single channel that ``LoopStage.step`` emits through.
|
||||
|
||||
It is deliberately additive and engine-agnostic: nothing here imports
|
||||
``torch`` or the pipeline machinery, and adapters bridge to/from today's
|
||||
:class:`~fastvideo.api.schema.GenerationRequest` /
|
||||
:class:`~fastvideo.api.results.GenerationResult` so the new types are
|
||||
usable through the existing ``VideoGenerator`` before the engine itself
|
||||
is rebuilt (plan.md M1). Later milestones evolve ``GenerationRequest``
|
||||
into ``OmniRequest`` in place and drop the adapters.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, ClassVar
|
||||
from uuid import uuid4
|
||||
|
||||
from fastvideo.api.results import (
|
||||
GenerationResult,
|
||||
VideoEvent,
|
||||
VideoFinalEvent,
|
||||
VideoPartialEvent,
|
||||
VideoProgressEvent,
|
||||
)
|
||||
from fastvideo.api.schema import (
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
RequestRuntimeConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
class Modality(str, Enum):
|
||||
"""A media modality. ``str`` mixin keeps it JSON/serialization friendly."""
|
||||
|
||||
TEXT = "text"
|
||||
IMAGE = "image"
|
||||
VIDEO = "video"
|
||||
AUDIO = "audio"
|
||||
ACTION = "action"
|
||||
LATENT = "latent"
|
||||
|
||||
|
||||
class TaskType(str, Enum):
|
||||
"""The declared task (design.md §6.1 / kills P7).
|
||||
|
||||
The pipeline graph branches on ``request.task``. Heuristics may only
|
||||
*suggest* a default at the API boundary (see :func:`infer_task`); they
|
||||
never decide control flow inside the runtime.
|
||||
"""
|
||||
|
||||
T2V = "t2v" # text -> video
|
||||
I2V = "i2v" # image (+text) -> video
|
||||
TI2V = "ti2v" # text + init image -> video
|
||||
V2V = "v2v" # video -> video (edit / restyle)
|
||||
V2W = "v2w" # video -> world (continue a world-model rollout)
|
||||
T2I = "t2i" # text -> image
|
||||
I2I = "i2i" # image -> image (edit)
|
||||
T2A = "t2a" # text -> audio
|
||||
T2VS = "t2vs" # text -> video + sound (joint A/V)
|
||||
A2W = "a2w" # action -> world (interactive world model)
|
||||
REASON = "reason" # text -> text (AR reasoner / prompt upsampling)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Inputs: typed modality parts (replaces InputConfig's per-model fields, P3).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModalPart:
|
||||
"""Base for a typed input part.
|
||||
|
||||
``modality`` is intrinsic to the concrete subclass (a ``ClassVar``, not
|
||||
an instance field). ``role`` disambiguates several parts of one
|
||||
modality — e.g. ``"prompt"`` vs ``"negative"`` text, ``"init"`` vs
|
||||
``"conditioning"`` image — and is keyword-only so subclass payloads stay
|
||||
positional.
|
||||
"""
|
||||
|
||||
modality: ClassVar[Modality]
|
||||
role: str | None = field(default=None, kw_only=True)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextPart(ModalPart):
|
||||
modality: ClassVar[Modality] = Modality.TEXT
|
||||
text: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class ImagePart(ModalPart):
|
||||
modality: ClassVar[Modality] = Modality.IMAGE
|
||||
image: Any | None = None # PIL.Image / ndarray / tensor
|
||||
path: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoPart(ModalPart):
|
||||
modality: ClassVar[Modality] = Modality.VIDEO
|
||||
video: Any | None = None
|
||||
path: str | None = None
|
||||
fps: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class AudioPart(ModalPart):
|
||||
modality: ClassVar[Modality] = Modality.AUDIO
|
||||
audio: Any | None = None
|
||||
path: str | None = None
|
||||
sample_rate: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ActionPart(ModalPart):
|
||||
modality: ClassVar[Modality] = Modality.ACTION
|
||||
action: Any | None = None # mouse / keyboard / camera tensors
|
||||
kind: str | None = None # "mouse" | "keyboard" | "camera" | ...
|
||||
|
||||
|
||||
@dataclass
|
||||
class LatentPart(ModalPart):
|
||||
modality: ClassVar[Modality] = Modality.LATENT
|
||||
latents: Any | None = None
|
||||
of_modality: Modality = Modality.VIDEO # modality these latents decode to
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Per-call parameters: AR sampling vs diffusion, separated (design.md §6.1).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingParams:
|
||||
"""AR decode knobs.
|
||||
|
||||
Distinct from the legacy diffusion ``fastvideo.api.SamplingParam``: these
|
||||
drive ``ARDecodeLoop`` (Cosmos3 reasoner, omni thinkers/talkers, codec
|
||||
decode), not denoising.
|
||||
"""
|
||||
|
||||
max_tokens: int = 512
|
||||
temperature: float = 1.0
|
||||
top_p: float = 1.0
|
||||
top_k: int | None = None
|
||||
stop: list[str] = field(default_factory=list)
|
||||
seed: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class DiffusionParams:
|
||||
"""Denoise knobs. ``guidance_per_modality`` carries per-modality CFG
|
||||
scales for joint A/V denoise (LTX-2, Cosmos3 t2vs); ``guidance_scale`` is
|
||||
the scalar default."""
|
||||
|
||||
steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_per_modality: dict[Modality, float] = field(default_factory=dict)
|
||||
sigmas: list[float] | None = None
|
||||
flow_shift: float | None = None
|
||||
height: int | None = None
|
||||
width: int | None = None
|
||||
num_frames: int | None = None
|
||||
fps: int | None = None
|
||||
seed: int | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Outputs spec: requested modalities + streaming + capture flags.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamSpec:
|
||||
"""Per-modality streaming policy. ``chunk_ms`` applies to audio chunks,
|
||||
``per_chunk`` to chunked-causal video; ``enabled`` alone covers token
|
||||
text."""
|
||||
|
||||
enabled: bool = True
|
||||
chunk_ms: int | None = None
|
||||
per_chunk: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class OutputSpec:
|
||||
"""Requested output modalities, streaming, and capture flags."""
|
||||
|
||||
modalities: list[Modality] = field(default_factory=lambda: [Modality.VIDEO])
|
||||
stream: dict[Modality, StreamSpec] = field(default_factory=dict)
|
||||
return_latents: bool = False
|
||||
return_trajectory: bool = False
|
||||
|
||||
@property
|
||||
def streaming(self) -> bool:
|
||||
"""True if any modality is requested as a stream."""
|
||||
return any(spec.enabled for spec in self.stream.values())
|
||||
|
||||
|
||||
@dataclass
|
||||
class NodeOverrides:
|
||||
"""Per-graph-node parameter overrides (design.md §6.1).
|
||||
|
||||
For a multi-loop graph, ``node_params["refine"].steps`` overrides the
|
||||
refine loop's step count without leaking ``refine_*`` onto the universal
|
||||
schema. Validation against each node's declared schema arrives with
|
||||
``PipelineSpec`` (a later milestone); for now this is a typed bag with
|
||||
``get`` / item / attribute access.
|
||||
"""
|
||||
|
||||
params: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def get(self, key: str, default: Any = None) -> Any:
|
||||
return self.params.get(key, default)
|
||||
|
||||
def __getitem__(self, key: str) -> Any:
|
||||
return self.params[key]
|
||||
|
||||
def __getattr__(self, key: str) -> Any:
|
||||
# Only invoked when normal attribute lookup fails, so ``params``
|
||||
# itself resolves through ``__dict__`` and never recurses.
|
||||
try:
|
||||
return self.__dict__["params"][key]
|
||||
except KeyError as exc:
|
||||
raise AttributeError(key) from exc
|
||||
|
||||
|
||||
@dataclass
|
||||
class OmniRequest:
|
||||
"""A typed multimodal request (design.md §6.1)."""
|
||||
|
||||
task: TaskType
|
||||
inputs: list[ModalPart] = field(default_factory=list)
|
||||
sampling: SamplingParams = field(default_factory=SamplingParams)
|
||||
diffusion: DiffusionParams = field(default_factory=DiffusionParams)
|
||||
outputs: OutputSpec = field(default_factory=OutputSpec)
|
||||
node_params: dict[str, NodeOverrides] = field(default_factory=dict)
|
||||
priority: int = 0
|
||||
request_id: str = field(default_factory=lambda: uuid4().hex)
|
||||
|
||||
# -- accessors ----------------------------------------------------------
|
||||
|
||||
def parts(self, modality: Modality) -> list[ModalPart]:
|
||||
return [p for p in self.inputs if p.modality is modality]
|
||||
|
||||
@property
|
||||
def prompt(self) -> str | None:
|
||||
for part in self.inputs:
|
||||
if isinstance(part, TextPart) and part.role in (None, "prompt"):
|
||||
return part.text
|
||||
return None
|
||||
|
||||
@property
|
||||
def negative_prompt(self) -> str | None:
|
||||
for part in self.inputs:
|
||||
if isinstance(part, TextPart) and part.role in ("negative", "negative_prompt"):
|
||||
return part.text
|
||||
return None
|
||||
|
||||
# -- constructors / adapters -------------------------------------------
|
||||
|
||||
@classmethod
|
||||
def from_prompt(
|
||||
cls,
|
||||
prompt: str | None,
|
||||
task: TaskType,
|
||||
*,
|
||||
negative_prompt: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> OmniRequest:
|
||||
"""Build a request from a bare prompt — what ``generate_video`` calls
|
||||
internally so the offline shim constructs an ``OmniRequest`` (G5)."""
|
||||
inputs: list[ModalPart] = []
|
||||
if prompt is not None:
|
||||
inputs.append(TextPart(prompt))
|
||||
if negative_prompt is not None:
|
||||
inputs.append(TextPart(negative_prompt, role="negative"))
|
||||
return cls(task=task, inputs=inputs, **kwargs)
|
||||
|
||||
@classmethod
|
||||
def from_generation_request(
|
||||
cls,
|
||||
request: GenerationRequest,
|
||||
task: TaskType | None = None,
|
||||
) -> OmniRequest:
|
||||
"""Lift a legacy typed :class:`GenerationRequest` into an
|
||||
``OmniRequest``; ``task`` defaults to the boundary heuristic."""
|
||||
resolved = task if task is not None else infer_task(request)
|
||||
inputs: list[ModalPart] = []
|
||||
prompt = request.prompt
|
||||
if isinstance(prompt, list):
|
||||
prompt = prompt[0] if prompt else None
|
||||
if prompt is not None:
|
||||
inputs.append(TextPart(prompt))
|
||||
if request.negative_prompt is not None:
|
||||
inputs.append(TextPart(request.negative_prompt, role="negative"))
|
||||
|
||||
inp = request.inputs
|
||||
image_path = inp.image_path if isinstance(inp.image_path, str) else None
|
||||
if image_path is not None or inp.pil_image is not None:
|
||||
inputs.append(ImagePart(image=inp.pil_image, path=image_path))
|
||||
video_path = inp.video_path if isinstance(inp.video_path, str) else None
|
||||
if video_path is not None:
|
||||
inputs.append(VideoPart(path=video_path))
|
||||
if inp.mouse_cond is not None or inp.keyboard_cond is not None:
|
||||
inputs.append(ActionPart(action=inp.mouse_cond, kind="mouse"))
|
||||
|
||||
sampling = request.sampling
|
||||
diffusion = DiffusionParams(
|
||||
steps=sampling.num_inference_steps,
|
||||
guidance_scale=sampling.guidance_scale,
|
||||
sigmas=sampling.sigmas,
|
||||
height=sampling.height,
|
||||
width=sampling.width,
|
||||
num_frames=sampling.num_frames,
|
||||
fps=sampling.fps,
|
||||
seed=sampling.seed,
|
||||
)
|
||||
outputs = OutputSpec(
|
||||
return_trajectory=request.runtime.return_trajectory_latents,
|
||||
return_latents=request.runtime.return_trajectory_decoded,
|
||||
)
|
||||
node_params = {
|
||||
node: NodeOverrides(params=dict(overrides))
|
||||
for node, overrides in request.stage_overrides.items() if isinstance(overrides, dict)
|
||||
}
|
||||
return cls(
|
||||
task=resolved,
|
||||
inputs=inputs,
|
||||
diffusion=diffusion,
|
||||
outputs=outputs,
|
||||
node_params=node_params,
|
||||
)
|
||||
|
||||
def to_generation_request(self) -> GenerationRequest:
|
||||
"""Lower to a legacy :class:`GenerationRequest` so today's
|
||||
``VideoGenerator`` can execute an ``OmniRequest`` unchanged (M1)."""
|
||||
diff = self.diffusion
|
||||
sampling = SamplingConfig()
|
||||
sampling.num_inference_steps = diff.steps
|
||||
sampling.guidance_scale = diff.guidance_scale
|
||||
sampling.sigmas = diff.sigmas
|
||||
if diff.seed is not None:
|
||||
sampling.seed = diff.seed
|
||||
if diff.num_frames is not None:
|
||||
sampling.num_frames = diff.num_frames
|
||||
if diff.height is not None:
|
||||
sampling.height = diff.height
|
||||
if diff.width is not None:
|
||||
sampling.width = diff.width
|
||||
if diff.fps is not None:
|
||||
sampling.fps = diff.fps
|
||||
|
||||
image_path: str | None = None
|
||||
pil_image: Any | None = None
|
||||
video_path: str | None = None
|
||||
for part in self.inputs:
|
||||
if isinstance(part, ImagePart):
|
||||
image_path = image_path or part.path
|
||||
pil_image = pil_image if pil_image is not None else part.image
|
||||
elif isinstance(part, VideoPart):
|
||||
video_path = video_path or part.path
|
||||
inputs = InputConfig(image_path=image_path, pil_image=pil_image, video_path=video_path)
|
||||
|
||||
runtime = RequestRuntimeConfig(
|
||||
return_trajectory_latents=self.outputs.return_trajectory,
|
||||
return_trajectory_decoded=self.outputs.return_latents,
|
||||
)
|
||||
output = OutputConfig(return_frames=Modality.VIDEO in self.outputs.modalities)
|
||||
stage_overrides = {node: dict(ov.params) for node, ov in self.node_params.items()}
|
||||
return GenerationRequest(
|
||||
prompt=self.prompt,
|
||||
negative_prompt=self.negative_prompt,
|
||||
inputs=inputs,
|
||||
sampling=sampling,
|
||||
runtime=runtime,
|
||||
output=output,
|
||||
stage_overrides=stage_overrides,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Outputs: named artifacts with provenance (kills extra["audio"], P3).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class Artifact:
|
||||
"""Base output artifact. ``source_node`` records which graph node
|
||||
produced it (provenance, design.md §6.1)."""
|
||||
|
||||
modality: ClassVar[Modality]
|
||||
source_node: str | None = field(default=None, kw_only=True)
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoArtifact(Artifact):
|
||||
modality: ClassVar[Modality] = Modality.VIDEO
|
||||
frames: Any | None = None # numpy (N, H, W, 3) uint8
|
||||
tensor: Any | None = None # raw sample tensor
|
||||
path: str | None = None
|
||||
fps: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class AudioArtifact(Artifact):
|
||||
modality: ClassVar[Modality] = Modality.AUDIO
|
||||
audio: Any | None = None
|
||||
sample_rate: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextArtifact(Artifact):
|
||||
modality: ClassVar[Modality] = Modality.TEXT
|
||||
text: str = ""
|
||||
token_ids: list[int] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class TensorArtifact(Artifact):
|
||||
"""Action tensors and other raw tensor outputs."""
|
||||
|
||||
modality: ClassVar[Modality] = Modality.ACTION
|
||||
tensor: Any | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class LatentArtifact(Artifact):
|
||||
modality: ClassVar[Modality] = Modality.LATENT
|
||||
latents: Any | None = None
|
||||
timesteps: Any | None = None
|
||||
of_modality: Modality = Modality.VIDEO
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestMetrics:
|
||||
generation_time: float | None = None
|
||||
peak_memory_mb: float | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class OmniOutput:
|
||||
"""Typed multimodal output (design.md §6.1)."""
|
||||
|
||||
request_id: str
|
||||
artifacts: dict[str, Artifact] = field(default_factory=dict)
|
||||
metrics: RequestMetrics = field(default_factory=RequestMetrics)
|
||||
|
||||
def get(self, name: str) -> Artifact | None:
|
||||
return self.artifacts.get(name)
|
||||
|
||||
@property
|
||||
def video(self) -> VideoArtifact | None:
|
||||
art = self.artifacts.get("video")
|
||||
return art if isinstance(art, VideoArtifact) else None
|
||||
|
||||
@property
|
||||
def audio(self) -> AudioArtifact | None:
|
||||
art = self.artifacts.get("audio")
|
||||
return art if isinstance(art, AudioArtifact) else None
|
||||
|
||||
@classmethod
|
||||
def from_generation_result(
|
||||
cls,
|
||||
result: GenerationResult,
|
||||
request_id: str = "",
|
||||
) -> OmniOutput:
|
||||
"""Map a legacy :class:`GenerationResult` into named artifacts.
|
||||
|
||||
Audio becomes a first-class :class:`AudioArtifact` carrying its
|
||||
sample rate, instead of riding in ``extra["audio"]`` (P3).
|
||||
"""
|
||||
artifacts: dict[str, Artifact] = {}
|
||||
if result.frames is not None or result.samples is not None or result.video_path is not None:
|
||||
artifacts["video"] = VideoArtifact(
|
||||
frames=result.frames,
|
||||
tensor=result.samples,
|
||||
path=result.video_path,
|
||||
source_node="decode",
|
||||
)
|
||||
if result.audio is not None:
|
||||
artifacts["audio"] = AudioArtifact(
|
||||
audio=result.audio,
|
||||
sample_rate=result.audio_sample_rate,
|
||||
source_node="audio_decode",
|
||||
)
|
||||
if result.trajectory is not None or result.trajectory_decoded is not None:
|
||||
artifacts["latents"] = LatentArtifact(
|
||||
latents=result.trajectory,
|
||||
timesteps=result.trajectory_timesteps,
|
||||
source_node="denoise",
|
||||
)
|
||||
return cls(
|
||||
request_id=request_id,
|
||||
artifacts=artifacts,
|
||||
metrics=RequestMetrics(
|
||||
generation_time=result.generation_time,
|
||||
peak_memory_mb=result.peak_memory_mb,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Streaming: one event union (evolves api/results.py's Video*Event).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class OmniProgressEvent:
|
||||
"""Per-step progress telemetry."""
|
||||
|
||||
step: int
|
||||
total_steps: int
|
||||
node: str = "denoise"
|
||||
|
||||
|
||||
@dataclass
|
||||
class OmniChunkEvent:
|
||||
"""A streamed artifact chunk — the universal ``StepResult.emit`` channel
|
||||
(design.md §6.2.2): text tokens, audio chunks, or decoded frame chunks.
|
||||
|
||||
``pts`` optionally carries a presentation timestamp for raw-frame
|
||||
streaming over WebRTC (design.md §9.1)."""
|
||||
|
||||
modality: Modality
|
||||
index: int
|
||||
payload: Any = None
|
||||
pts: float | None = None
|
||||
node: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class OmniFinalEvent:
|
||||
"""Terminal event carrying the full :class:`OmniOutput`."""
|
||||
|
||||
output: OmniOutput
|
||||
|
||||
|
||||
OmniEvent = OmniProgressEvent | OmniChunkEvent | OmniFinalEvent
|
||||
"""Union of every event the engine streams; consumers match by ``isinstance``."""
|
||||
|
||||
|
||||
def infer_task(request: GenerationRequest) -> TaskType:
|
||||
"""Best-effort boundary heuristic for a legacy request (design.md §6.1).
|
||||
|
||||
A *suggestion* only — the runtime always branches on the explicit
|
||||
``OmniRequest.task``, never on this. Single-frame requests are images;
|
||||
a video input implies edit; an image input implies image-to-video.
|
||||
"""
|
||||
inp = request.inputs
|
||||
has_image = bool(inp.image_path or inp.pil_image)
|
||||
has_video = bool(inp.video_path or inp.stage1_video)
|
||||
if request.sampling.num_frames == 1:
|
||||
return TaskType.I2I if has_image else TaskType.T2I
|
||||
if has_video:
|
||||
return TaskType.V2V
|
||||
if has_image:
|
||||
return TaskType.I2V
|
||||
return TaskType.T2V
|
||||
|
||||
|
||||
def omni_event_from_video_event(event: VideoEvent, request_id: str = "") -> OmniEvent:
|
||||
"""Adapt a legacy :data:`~fastvideo.api.results.VideoEvent` to an
|
||||
:data:`OmniEvent` so the streaming surface can migrate incrementally."""
|
||||
if isinstance(event, VideoProgressEvent):
|
||||
return OmniProgressEvent(step=event.step, total_steps=event.total_steps, node=event.stage)
|
||||
if isinstance(event, VideoPartialEvent):
|
||||
return OmniChunkEvent(modality=Modality.VIDEO, index=event.index, payload=event.frames, node="decode")
|
||||
if isinstance(event, VideoFinalEvent):
|
||||
if event.result is not None:
|
||||
output = OmniOutput.from_generation_result(event.result, request_id)
|
||||
else:
|
||||
output = OmniOutput(request_id=request_id)
|
||||
if event.frames is not None:
|
||||
output.artifacts["video"] = VideoArtifact(frames=event.frames, source_node="decode")
|
||||
return OmniFinalEvent(output=output)
|
||||
raise TypeError(f"unknown VideoEvent type: {type(event).__name__}")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ActionPart",
|
||||
"Artifact",
|
||||
"AudioArtifact",
|
||||
"AudioPart",
|
||||
"DiffusionParams",
|
||||
"ImagePart",
|
||||
"LatentArtifact",
|
||||
"LatentPart",
|
||||
"Modality",
|
||||
"ModalPart",
|
||||
"NodeOverrides",
|
||||
"OmniChunkEvent",
|
||||
"OmniEvent",
|
||||
"OmniFinalEvent",
|
||||
"OmniOutput",
|
||||
"OmniProgressEvent",
|
||||
"OmniRequest",
|
||||
"OutputSpec",
|
||||
"RequestMetrics",
|
||||
"SamplingParams",
|
||||
"StreamSpec",
|
||||
"TaskType",
|
||||
"TensorArtifact",
|
||||
"TextArtifact",
|
||||
"TextPart",
|
||||
"VideoArtifact",
|
||||
"VideoPart",
|
||||
"infer_task",
|
||||
"omni_event_from_video_event",
|
||||
]
|
||||
@@ -49,6 +49,10 @@ def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
|
||||
return _attn_qat_train_attention
|
||||
|
||||
|
||||
def is_attn_qat_train_available() -> bool:
|
||||
return _get_attn_qat_train_attention() is not None
|
||||
|
||||
|
||||
def attn_qat_train(q_BLHD: torch.Tensor,
|
||||
k_BLHD: torch.Tensor,
|
||||
v_BLHD: torch.Tensor,
|
||||
|
||||
@@ -55,6 +55,20 @@ class SageAttention3Impl(AttentionImpl):
|
||||
self.softmax_scale = softmax_scale
|
||||
self.dropout = extra_impl_args.get("dropout_p", 0.0)
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Transpose stacked QKV from [3B, L, H, D] to [3B, H, L, D].
|
||||
|
||||
Single bulk permute+contiguous on the entire stacked tensor rather than
|
||||
three separate transposed views for Q, K, V. The .contiguous() is
|
||||
required: sageattn_blackwell's fake kernel returns empty_like(q), so the
|
||||
op's output strides must match contiguous q under torch.compile.
|
||||
"""
|
||||
return qkv.permute(0, 2, 1, 3).contiguous()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
@@ -62,9 +76,15 @@ class SageAttention3Impl(AttentionImpl):
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
"""Call sageattn3_blackwell directly. Input is already [B, H, L, D]
|
||||
and contiguous from preprocess_qkv."""
|
||||
output = sageattn3_blackwell(query, key, value, is_causal=self.causal)
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Transpose output from [B, H, L, D] back to [B, L, H, D]."""
|
||||
return output.permute(0, 2, 1, 3).contiguous()
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
@@ -13,6 +15,26 @@ from fastvideo.utils import get_compute_dtype
|
||||
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
|
||||
|
||||
|
||||
def _attention_compile_disabled() -> bool:
|
||||
"""Whether to keep attention ``forward`` out of the torch.compile graph.
|
||||
|
||||
Defaults to ``True`` (the historical behavior: attention runs eager via
|
||||
``torch.compiler.disable``). Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=0``
|
||||
to let attention be traced/compiled into the surrounding graph.
|
||||
"""
|
||||
val = os.environ.get("FASTVIDEO_DISABLE_ATTENTION_COMPILE")
|
||||
if val is None:
|
||||
return True
|
||||
return val.strip().lower() not in ("0", "false", "no", "off", "")
|
||||
|
||||
|
||||
def _maybe_compiler_disable(fn):
|
||||
"""Apply ``torch.compiler.disable`` unless disabled via env var."""
|
||||
if _attention_compile_disabled():
|
||||
return torch.compiler.disable(fn)
|
||||
return fn
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
"""Distributed attention layer.
|
||||
"""
|
||||
@@ -56,7 +78,7 @@ class DistributedAttention(nn.Module):
|
||||
self.backend = backend_name_to_enum(attn_backend.get_name())
|
||||
self.dtype = dtype
|
||||
|
||||
@torch.compiler.disable
|
||||
@_maybe_compiler_disable
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
@@ -146,7 +168,7 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
"""Distributed attention layer with VSA support.
|
||||
"""
|
||||
|
||||
@torch.compiler.disable
|
||||
@_maybe_compiler_disable
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
|
||||
@@ -24,7 +24,9 @@ class DiTArchConfig(ArchConfig):
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN, AttentionBackendEnum.SAGE_ATTN_THREE,
|
||||
AttentionBackendEnum.SLA_ATTN, AttentionBackendEnum.SAGE_SLA_ATTN)
|
||||
AttentionBackendEnum.ATTN_QAT_INFER,
|
||||
AttentionBackendEnum.ATTN_QAT_TRAIN, AttentionBackendEnum.SLA_ATTN,
|
||||
AttentionBackendEnum.SAGE_SLA_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -273,11 +273,17 @@ class FastVideoArgs:
|
||||
dit_config = getattr(self.pipeline_config, "dit_config", None)
|
||||
if dit_config is None:
|
||||
return
|
||||
# Resolve a registry name (e.g. "nvfp4_qat_train" from the CLI) to a
|
||||
# QuantizationConfig instance; a bare string has no get_quant_method.
|
||||
tq = self.transformer_quant
|
||||
if isinstance(tq, str):
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
tq = get_quantization_config(tq)()
|
||||
# Don't overwrite if the caller already set it explicitly on
|
||||
# dit_config (e.g. via ``pipeline_config.dit_config.quant_config = NVFP4Config()``);
|
||||
# the explicit setter wins.
|
||||
if getattr(dit_config, "quant_config", None) is None:
|
||||
dit_config.quant_config = self.transformer_quant
|
||||
dit_config.quant_config = tq
|
||||
|
||||
def _resolve_refine_args(self) -> None:
|
||||
"""Map generic refine_* args to LTX-2-specific refine fields."""
|
||||
@@ -1018,6 +1024,12 @@ class TrainingArgs(FastVideoArgs):
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
parser.add_argument("--data-path", type=str, required=True, help="Path to parquet files")
|
||||
parser.add_argument("--transformer-quant",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Quantization config name for the DiT (e.g. nvfp4_qat_train for "
|
||||
"QAT-finetune FP4 linear with a straight-through estimator). "
|
||||
"Resolved to a QuantizationConfig and pinned on dit_config.quant_config.")
|
||||
parser.add_argument("--dataloader-num-workers",
|
||||
type=int,
|
||||
required=True,
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FP8 quantization-aware training for linear layers.
|
||||
|
||||
Mirror of ``fp4linear.py`` but for FP8 (e4m3). The forward pass quantizes both
|
||||
activations and weights to FP8 and runs ``torch._scaled_mm``; the backward pass
|
||||
is a bf16 straight-through estimator so the high-precision master weights stay
|
||||
trainable. Falls back to a bf16 fake-quant forward on GPUs older than sm89.
|
||||
"""
|
||||
import torch
|
||||
|
||||
FP8_DTYPE = torch.float8_e4m3fn
|
||||
FP8_MAX = float(torch.finfo(FP8_DTYPE).max) # 448.0
|
||||
FP8_MIN_SCALE = 1.0 / (FP8_MAX * 512.0)
|
||||
|
||||
|
||||
def _supports_fp8_compute() -> bool:
|
||||
"""Whether the active device supports FP8 ``_scaled_mm`` (sm89+)."""
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
cap = torch.cuda.get_device_capability()
|
||||
return cap[0] > 8 or (cap[0] == 8 and cap[1] >= 9)
|
||||
|
||||
|
||||
def _quantize_tensorwise(x_2d: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Returns ``(x_fp8 [M, K], x_scale [1] float32)``."""
|
||||
x_absmax = x_2d.abs().amax().float()
|
||||
x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
return x_fp8, x_scale.view(1)
|
||||
|
||||
|
||||
def _quantize_rowwise(x_2d: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``."""
|
||||
x_absmax = x_2d.abs().amax(dim=-1, keepdim=True).float()
|
||||
x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
return x_fp8, x_scale
|
||||
|
||||
|
||||
def _fake_quant(x_2d: torch.Tensor, granularity: str) -> torch.Tensor:
|
||||
"""bf16 fake-quant (quantize then dequantize) for pre-sm89 fallback."""
|
||||
if granularity == "channel":
|
||||
x_fp8, x_scale = _quantize_rowwise(x_2d)
|
||||
else:
|
||||
x_fp8, x_scale = _quantize_tensorwise(x_2d)
|
||||
return x_fp8.to(x_2d.dtype) * x_scale.to(x_2d.dtype)
|
||||
|
||||
|
||||
class _LinearFWD8BWD16Fn(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, x, weight, bias, granularity="tensor"):
|
||||
# assert/normalize activation dtype
|
||||
if x.dtype not in (torch.float16, torch.bfloat16):
|
||||
x = x.to(dtype=torch.bfloat16)
|
||||
|
||||
# cast params (can be fp32) to activation dtype for quantization
|
||||
weight_cast = weight.to(dtype=x.dtype)
|
||||
bias_cast = bias.to(dtype=x.dtype) if bias is not None else None
|
||||
|
||||
orig_shape = x.shape
|
||||
k = weight_cast.shape[1]
|
||||
n = weight_cast.shape[0]
|
||||
x2d = x.reshape(-1, k).contiguous()
|
||||
|
||||
if not _supports_fp8_compute():
|
||||
# bf16 fake-quant fallback: simulate the FP8 rounding error but
|
||||
# compute the matmul in bf16.
|
||||
x_fq = _fake_quant(x2d, granularity)
|
||||
w_fq = _fake_quant(weight_cast, granularity)
|
||||
out2d = x_fq.matmul(w_fq.t())
|
||||
if bias_cast is not None:
|
||||
out2d = out2d + bias_cast
|
||||
ctx.save_for_backward(x2d, weight, bias)
|
||||
ctx.n = n
|
||||
ctx.orig_shape = orig_shape
|
||||
return out2d.reshape(*orig_shape[:-1], n)
|
||||
|
||||
if granularity == "channel":
|
||||
x_fp8, x_scale = _quantize_rowwise(x2d)
|
||||
w_fp8, w_scale = _quantize_rowwise(weight_cast)
|
||||
scale_b = w_scale.view(1, -1)
|
||||
else:
|
||||
x_fp8, x_scale = _quantize_tensorwise(x2d)
|
||||
w_fp8, w_scale = _quantize_tensorwise(weight_cast)
|
||||
scale_b = w_scale
|
||||
|
||||
out2d = torch._scaled_mm(
|
||||
x_fp8,
|
||||
w_fp8.t(),
|
||||
scale_a=x_scale,
|
||||
scale_b=scale_b,
|
||||
out_dtype=x.dtype,
|
||||
)
|
||||
if isinstance(out2d, tuple):
|
||||
out2d = out2d[0]
|
||||
|
||||
if bias_cast is not None:
|
||||
out2d = out2d + bias_cast
|
||||
|
||||
# save tensors for backward (keep original dtypes)
|
||||
ctx.save_for_backward(x2d, weight, bias)
|
||||
ctx.n = n
|
||||
ctx.orig_shape = orig_shape
|
||||
return out2d.reshape(*orig_shape[:-1], n)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_out):
|
||||
x2d, weight, bias = ctx.saved_tensors
|
||||
M = x2d.shape[0]
|
||||
n = ctx.n
|
||||
|
||||
grad_out_2d = grad_out.reshape(M, n).contiguous()
|
||||
|
||||
# bf16 straight-through estimator: gradients flow through the
|
||||
# full-precision master weights, not the FP8 quantized values.
|
||||
weight_cast = weight.to(dtype=grad_out.dtype)
|
||||
x_cast = x2d.to(dtype=grad_out.dtype)
|
||||
|
||||
grad_x = grad_out_2d.matmul(weight_cast).reshape(*ctx.orig_shape)
|
||||
grad_w = grad_out_2d.t().matmul(x_cast)
|
||||
grad_b = grad_out_2d.sum(dim=0) if bias is not None else None
|
||||
|
||||
# None for the extra forward arg (granularity)
|
||||
return grad_x, grad_w, grad_b, None
|
||||
|
||||
|
||||
def fp8_linear_forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
# pass config **positionally**; autograd.Function.apply ignores kwargs
|
||||
return _LinearFWD8BWD16Fn.apply(x, self.weight, self.bias, "tensor"), None
|
||||
@@ -2,7 +2,7 @@ from typing import Literal, get_args
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig
|
||||
|
||||
QuantizationMethods = Literal[None, "AbsMaxFP8", "NVFP4", "nvfp4_qat"]
|
||||
QuantizationMethods = Literal[None, "AbsMaxFP8", "FP8", "NVFP4", "nvfp4_qat", "nvfp4_qat_train", "fp8_qat_train"]
|
||||
|
||||
QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods))
|
||||
|
||||
@@ -51,13 +51,19 @@ def get_quantization_config(quantization: str) -> type[QuantizationConfig]:
|
||||
|
||||
# lazy import to avoid triggering `torch.compile` too early
|
||||
from .absmax_fp8 import AbsMaxFP8Config
|
||||
from .fp8_config import FP8Config
|
||||
from .nvfp4_config import NVFP4Config
|
||||
from .nvfp4_qat_config import NVFP4QATConfig
|
||||
from .nvfp4_qat_train_config import NVFP4QATTrainConfig
|
||||
from .fp8_qat_train_config import FP8QATTrainConfig
|
||||
|
||||
method_to_config: dict[str, type[QuantizationConfig]] = {
|
||||
"AbsMaxFP8": AbsMaxFP8Config,
|
||||
"FP8": FP8Config,
|
||||
"NVFP4": NVFP4Config,
|
||||
"nvfp4_qat": NVFP4QATConfig,
|
||||
"nvfp4_qat_train": NVFP4QATTrainConfig,
|
||||
"fp8_qat_train": FP8QATTrainConfig,
|
||||
}
|
||||
# Update the `method_to_config` with customized quantization methods.
|
||||
method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG)
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Generic FP8 quantization backed by ``torch._scaled_mm``.
|
||||
|
||||
Matches linear layers by suffix (``to_q/k/v/to_out``, ``ffn.fc_in/fc_out``).
|
||||
Supports per-tensor (default, fast) and per-channel (higher accuracy) granularity.
|
||||
Falls back to bf16 dequant on GPUs older than sm89.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
FP8_DTYPE = torch.float8_e4m3fn
|
||||
FP8_MAX = float(torch.finfo(FP8_DTYPE).max) # 448.0
|
||||
FP8_MIN_SCALE = 1.0 / (FP8_MAX * 512.0)
|
||||
|
||||
_FP8_SUFFIXES = (
|
||||
"ffn.fc_in",
|
||||
"ffn.fc_out",
|
||||
"to_q",
|
||||
"to_k",
|
||||
"to_v",
|
||||
"to_out",
|
||||
)
|
||||
|
||||
|
||||
def _supports_fp8_compute() -> bool:
|
||||
"""Whether the active device supports FP8 ``_scaled_mm`` (sm89+)."""
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
cap = torch.cuda.get_device_capability()
|
||||
return cap[0] > 8 or (cap[0] == 8 and cap[1] >= 9)
|
||||
|
||||
|
||||
def _quantize_tensorwise(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Returns ``(x_fp8 [M, K], x_scale [1] float32)``."""
|
||||
x_absmax = x_2d.abs().amax().float()
|
||||
x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
return x_fp8, x_scale.view(1)
|
||||
|
||||
|
||||
def _quantize_rowwise(x_2d: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Returns ``(x_fp8 [M, K], x_scale [M, 1] float32)``."""
|
||||
x_absmax = x_2d.abs().amax(dim=-1, keepdim=True).float()
|
||||
x_scale = (x_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
x_fp8 = (x_2d / x_scale.to(x_2d.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
return x_fp8, x_scale
|
||||
|
||||
|
||||
class FP8QuantizeMethod(QuantizeMethodBase):
|
||||
"""FP8 linear method.
|
||||
|
||||
``granularity='tensor'`` (default): per-tensor weight + per-tensor
|
||||
dynamic activation scales — the fast tensorwise ``_scaled_mm`` path.
|
||||
``granularity='channel'``: per-output-channel weight + per-token
|
||||
activation scales (rowwise) — higher accuracy but slower ``_scaled_mm``.
|
||||
"""
|
||||
|
||||
def __init__(self, granularity: str = "tensor"):
|
||||
super().__init__()
|
||||
self.granularity = granularity
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
):
|
||||
weight = Parameter(
|
||||
torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def quantize_input(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, None]:
|
||||
"""Pre-quantize an activation for reuse across q/k/v projections."""
|
||||
assert x.dtype in (torch.bfloat16, torch.float16), (f"only allow bf16/fp16 inputs to fp8 linear, got {x.dtype}")
|
||||
x_2d = x.view(-1, x.shape[-1])
|
||||
if self.granularity == "channel":
|
||||
x_fp8, x_scale = _quantize_rowwise(x_2d)
|
||||
else:
|
||||
x_fp8, x_scale = _quantize_tensorwise(x_2d)
|
||||
return x_fp8, x_scale, None
|
||||
|
||||
def wants_prequantized_input(self) -> bool:
|
||||
return True
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
pre_quantized: tuple[torch.Tensor, torch.Tensor, Any] | None = None,
|
||||
) -> torch.Tensor:
|
||||
out_dim = layer._fp8_weight.shape[0]
|
||||
original_shape = x.shape
|
||||
|
||||
if not _supports_fp8_compute():
|
||||
return self._apply_dequant(layer, x, bias)
|
||||
|
||||
if pre_quantized is not None:
|
||||
x_fp8, x_scale, _ = pre_quantized
|
||||
if x_fp8.dim() > 2:
|
||||
x_fp8 = x_fp8.reshape(-1, x_fp8.shape[-1])
|
||||
if x_scale.dim() > 2:
|
||||
x_scale = x_scale.reshape(-1, x_scale.shape[-1])
|
||||
elif self.granularity == "channel":
|
||||
x_fp8, x_scale = _quantize_rowwise(x.reshape(-1, x.shape[-1]))
|
||||
else:
|
||||
x_fp8, x_scale = _quantize_tensorwise(x.reshape(-1, x.shape[-1]))
|
||||
|
||||
w_fp8 = layer._fp8_weight
|
||||
w_scale = layer._fp8_weight_scale
|
||||
scale_b = w_scale.view(1, -1) if self.granularity == "channel" else w_scale
|
||||
|
||||
out = torch._scaled_mm(
|
||||
x_fp8,
|
||||
w_fp8.t(),
|
||||
scale_a=x_scale,
|
||||
scale_b=scale_b,
|
||||
out_dtype=torch.bfloat16,
|
||||
)
|
||||
if isinstance(out, tuple):
|
||||
out = out[0]
|
||||
if bias is not None:
|
||||
out = out + bias
|
||||
return out.view(*original_shape[:-1], out_dim)
|
||||
|
||||
def _apply_dequant(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""bf16 fallback for pre-sm89 GPUs."""
|
||||
out_dim = layer._fp8_weight.shape[0]
|
||||
original_shape = x.shape
|
||||
w_fp8 = layer._fp8_weight
|
||||
w_scale = layer._fp8_weight_scale.to(x.dtype)
|
||||
weight = w_fp8.to(x.dtype) * w_scale.unsqueeze(1)
|
||||
out = F.linear(x, weight, bias)
|
||||
return out.view(*original_shape[:-1], out_dim)
|
||||
|
||||
|
||||
class FP8Config(QuantizationConfig):
|
||||
"""FP8 (e4m3) quantization via suffix matching on standard linear layer names."""
|
||||
|
||||
def __init__(self, granularity: str = "tensor"):
|
||||
super().__init__()
|
||||
if granularity not in ("tensor", "channel"):
|
||||
raise ValueError(f"granularity must be 'tensor' or 'channel', got {granularity!r}")
|
||||
self.granularity = granularity
|
||||
|
||||
def get_name(self) -> str:
|
||||
return "FP8"
|
||||
|
||||
def get_supported_act_dtypes(self) -> list[torch.dtype]:
|
||||
return [torch.bfloat16, torch.float16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return 89
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames() -> list[str]:
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> FP8Config:
|
||||
return cls(granularity=config.get("granularity", "tensor"))
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
from fastvideo.layers.linear import LinearBase
|
||||
|
||||
if isinstance(layer, LinearBase) and any(s in prefix for s in _FP8_SUFFIXES):
|
||||
return FP8QuantizeMethod(granularity=self.granularity)
|
||||
return None
|
||||
|
||||
|
||||
def convert_model_to_fp8(model: torch.nn.Module) -> None:
|
||||
"""Quantize all FP8-tagged linear layers in-place after weights are loaded."""
|
||||
import gc
|
||||
from torch.distributed.tensor import DTensor # type: ignore
|
||||
|
||||
with torch.no_grad():
|
||||
for mod in model.modules():
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
if not isinstance(qm, FP8QuantizeMethod):
|
||||
continue
|
||||
weight = getattr(mod, "weight", None)
|
||||
if weight is None:
|
||||
continue
|
||||
weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type]
|
||||
if getattr(qm, "granularity", "tensor") == "channel":
|
||||
w_absmax = weight_local.detach().abs().amax(dim=1).nan_to_num().float()
|
||||
w_scale = (w_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE)
|
||||
w_fp8 = (weight_local / w_scale.to(weight_local.dtype).unsqueeze(1)).clamp(-FP8_MAX,
|
||||
FP8_MAX).to(FP8_DTYPE)
|
||||
else:
|
||||
w_absmax = weight_local.detach().abs().amax().nan_to_num().to(torch.float32)
|
||||
w_scale = (w_absmax / FP8_MAX).clamp(min=FP8_MIN_SCALE).view(1)
|
||||
w_fp8 = (weight_local / w_scale.to(weight_local.dtype)).clamp(-FP8_MAX, FP8_MAX).to(FP8_DTYPE)
|
||||
mod.register_buffer("_fp8_weight", w_fp8.contiguous(), persistent=False)
|
||||
mod.register_buffer("_fp8_weight_scale", w_scale.to(torch.float32), persistent=False)
|
||||
removed_weight = mod._parameters.pop("weight", None)
|
||||
if removed_weight is not None:
|
||||
removed_weight.grad = None
|
||||
del removed_weight, weight, weight_local, w_absmax, w_scale, w_fp8
|
||||
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"FP8Config",
|
||||
"FP8QuantizeMethod",
|
||||
"convert_model_to_fp8",
|
||||
]
|
||||
@@ -0,0 +1,79 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FP8 (e4m3) quantization-aware *training* linear method (straight-through estimator).
|
||||
|
||||
Mirror of ``nvfp4_qat_train_config.py`` but for FP8. The weight stays a trainable
|
||||
bf16/fp32 master that is fake-quantized to FP8 on every forward, with a
|
||||
full-precision backward (STE), so the model learns to absorb FP8 linear error.
|
||||
|
||||
The STE lives in ``fastvideo.layers.fp8linear._LinearFWD8BWD16Fn`` (FP8 forward
|
||||
via ``torch._scaled_mm`` on sm89+, with a bf16 fake-quant fallback on older GPUs;
|
||||
full-precision backward). This method bridges it into the standard
|
||||
``quant_config`` path, so it activates via ``transformer_quant="fp8_qat_train"``
|
||||
on the same Wan-2.1 layers as the FP4 path (to_q/k/v/out + ffn). No conversion is
|
||||
needed: the weight is kept in full precision and quantized on the fly each step.
|
||||
|
||||
Unlike the FP4 path this needs no flashinfer and runs on any sm89+ GPU (and even
|
||||
older ones via the bf16 fallback), not just Blackwell.
|
||||
"""
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig, QuantizeMethodBase
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class FP8QATTrainQuantizeMethod(QuantizeMethodBase):
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, output_partition_sizes: list[int],
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype, **extra_weight_attrs):
|
||||
# Trainable master weight, fake-quantized to FP8 on each forward.
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=True)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
# FP8 forward + full-precision backward (STE).
|
||||
from fastvideo.layers.fp8linear import _LinearFWD8BWD16Fn
|
||||
return _LinearFWD8BWD16Fn.apply(x, layer.weight, bias, "tensor")
|
||||
|
||||
|
||||
class FP8QATTrainConfig(QuantizationConfig):
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_name(self):
|
||||
return "fp8_qat_train"
|
||||
|
||||
def get_supported_act_dtypes(self):
|
||||
return [torch.bfloat16, torch.float16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls):
|
||||
return 89
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames():
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "FP8QATTrainConfig":
|
||||
return cls()
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
from fastvideo.layers.linear import LinearBase
|
||||
fp8_layers = ["ffn.fc_in", "ffn.fc_out", "to_q", "to_k", "to_v", "to_out"]
|
||||
if isinstance(layer, LinearBase) and any(layer_name in prefix for layer_name in fp8_layers):
|
||||
return FP8QATTrainQuantizeMethod()
|
||||
return None
|
||||
@@ -1,38 +1,90 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""NVFP4 quantization-aware (QAD) linear method, inference path.
|
||||
|
||||
Quantizes every targeted linear's weight to NVFP4 once at load time and
|
||||
runs each forward as a registered flashinfer-backed FP4 matmul. The
|
||||
original fp16/bf16 weight is *popped* immediately after quantization so
|
||||
the half-precision copy does not keep occupying GPU memory — that's
|
||||
what lets a Wan-2.1 pipeline stay fully resident on a single GPU
|
||||
without any CPU offloading.
|
||||
|
||||
The quantize / matmul custom ops are owned by
|
||||
:mod:`fastvideo.layers.quantization.nvfp4_config` and registered under
|
||||
the ``fastvideo_fp4::`` namespace. We reuse them here for two reasons:
|
||||
|
||||
1. Re-registering the same op name in a second module would raise.
|
||||
2. The registered ops have ``register_fake`` shape/dtype kernels, which
|
||||
is what makes the inference pipeline's per-block ``torch.compile``
|
||||
trace through without graph breaks. Calling raw flashinfer functions
|
||||
(the old behavior of this file, plus a ``@torch.compile`` on
|
||||
``apply``) graph-breaks at every quantize and every matmul.
|
||||
|
||||
For QAT *training*, see ``nvfp4_qat_train_config`` which keeps the
|
||||
weight trainable and fake-quantizes on the fly via a straight-through
|
||||
estimator.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig, QuantizeMethodBase
|
||||
from fastvideo.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from fastvideo.layers.quantization.nvfp4_config import (
|
||||
_mm_fp4,
|
||||
_nvfp4_quantize,
|
||||
_require_flashinfer,
|
||||
)
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
try:
|
||||
import flashinfer
|
||||
except ImportError:
|
||||
flashinfer = None
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Wan-style attention + FFN projection layers. Matched as substrings of the
|
||||
# layer prefix (e.g. "blocks.0.attn1.to_q" contains "to_q").
|
||||
DEFAULT_FP4_LAYERS = (
|
||||
"ffn.fc_in",
|
||||
"ffn.fc_out",
|
||||
"to_q",
|
||||
"to_k",
|
||||
"to_v",
|
||||
"to_out",
|
||||
)
|
||||
|
||||
def _require_flashinfer() -> Any:
|
||||
if flashinfer is None:
|
||||
raise ImportError("flashinfer is required for NVFP4 QAT quantization. "
|
||||
"Please install flashinfer to use the nvfp4_qat quantization backend.")
|
||||
return flashinfer
|
||||
|
||||
def _layout_128x4() -> Any:
|
||||
SfLayout, _, _ = _require_flashinfer()
|
||||
return SfLayout.layout_128x4
|
||||
|
||||
|
||||
class NVFP4QATQuantizeMethod(QuantizeMethodBase):
|
||||
"""Inference-only NVFP4 linear method with weight popping.
|
||||
|
||||
The dense ``weight`` parameter is materialized at load time only so
|
||||
that :func:`convert_model_to_fp4` can read it once; the loader then
|
||||
removes it via ``mod._parameters.pop('weight')``. From that point
|
||||
forward, ``apply`` reads only ``_fp4_weight`` / ``_fp4_weight_scale``
|
||||
/ ``_weight_global_sf``.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.weight_fp4 = None
|
||||
self.weight_scale = None
|
||||
# Static input global scale factor. Matches the FastVideo-Quantization
|
||||
# production path; recomputing it per-call via a ``.max()`` reduction
|
||||
# (the previous behavior) adds a sync point, costs a kernel launch,
|
||||
# and produces a data-dependent value that prevents CUDA-graph
|
||||
# capture under ``torch.compile(mode='reduce-overhead')``.
|
||||
self.x_global_sf = torch.tensor(1.0, device="cuda", dtype=torch.float32)
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, output_partition_sizes: list[int],
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype, **extra_weight_attrs):
|
||||
"""Create weights for a linear layer. Note the corrected signature to match LinearMethodBase."""
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype, **extra_weight_attrs) -> None:
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
@@ -43,28 +95,27 @@ class NVFP4QATQuantizeMethod(QuantizeMethodBase):
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
@torch.compile
|
||||
def apply(self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
"""Apply NVFP4 QAT quantized computation."""
|
||||
flashinfer_mod = _require_flashinfer()
|
||||
out_dim = layer.weight.shape[0]
|
||||
# ``_fp4_weight`` carries the (out, in/2) packed fp4 weight, so its
|
||||
# row count is the output dim even after the dense weight is popped.
|
||||
out_dim = layer._fp4_weight.shape[0]
|
||||
original_shape = x.shape
|
||||
assert x.dtype == torch.bfloat16 or x.dtype == torch.float16, f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}"
|
||||
|
||||
assert x.dtype in (torch.bfloat16, torch.float16), (f"only allow bf16/fp16 inputs to fp4 linear, got {x.dtype}")
|
||||
x = x.view(-1, x.shape[-1])
|
||||
|
||||
x_global_sf = (448 * 6) / x.float().abs().nan_to_num().max()
|
||||
x_fp4, x_scale = flashinfer_mod.nvfp4_quantize(
|
||||
x_global_sf = self.x_global_sf
|
||||
x_fp4, x_scale = _nvfp4_quantize(
|
||||
x,
|
||||
x_global_sf,
|
||||
sfLayout=flashinfer_mod.SfLayout.layout_128x4,
|
||||
sfLayout=_layout_128x4(),
|
||||
do_shuffle=False,
|
||||
)
|
||||
|
||||
weight_fp4 = layer._fp4_weight
|
||||
weight_scale = layer._fp4_weight_scale
|
||||
weight_global_sf = layer._weight_global_sf
|
||||
|
||||
out = flashinfer_mod.mm_fp4(
|
||||
out = _mm_fp4(
|
||||
x_fp4,
|
||||
weight_fp4.T,
|
||||
x_scale,
|
||||
@@ -76,67 +127,106 @@ class NVFP4QATQuantizeMethod(QuantizeMethodBase):
|
||||
)
|
||||
|
||||
if bias is not None:
|
||||
if bias.device != out.device or bias.dtype != out.dtype:
|
||||
bias = bias.to(device=out.device, dtype=out.dtype)
|
||||
out = out + bias
|
||||
|
||||
if len(original_shape) == 3:
|
||||
out = out.view(original_shape[0], original_shape[1], out_dim)
|
||||
|
||||
out = out.view(*original_shape[:-1], out_dim)
|
||||
return out
|
||||
|
||||
|
||||
class NVFP4QATConfig(QuantizationConfig):
|
||||
"""NVFP4 (Wan-style) linear quantization, inference.
|
||||
|
||||
def __init__(self) -> None:
|
||||
Args:
|
||||
target_layers: Substrings matched against each linear layer's
|
||||
prefix. A layer is quantized if any substring is contained in
|
||||
its prefix. Defaults to the standard Wan attention + FFN
|
||||
projections (:data:`DEFAULT_FP4_LAYERS`).
|
||||
"""
|
||||
|
||||
def __init__(self, target_layers: tuple[str, ...] | None = None) -> None:
|
||||
super().__init__()
|
||||
self.target_layers = (tuple(target_layers) if target_layers else DEFAULT_FP4_LAYERS)
|
||||
|
||||
def get_name(self):
|
||||
def get_name(self) -> str:
|
||||
return "nvfp4_qat"
|
||||
|
||||
def get_supported_act_dtypes(self):
|
||||
def get_supported_act_dtypes(self) -> list[torch.dtype]:
|
||||
return [torch.bfloat16, torch.float16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls):
|
||||
def get_min_capability(cls) -> int:
|
||||
return 100
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames():
|
||||
def get_config_filenames() -> list[str]:
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "NVFP4QATConfig":
|
||||
return cls()
|
||||
def from_config(cls, config: dict[str, Any]) -> NVFP4QATConfig:
|
||||
target_layers = config.get("target_layers")
|
||||
if target_layers is not None:
|
||||
target_layers = tuple(target_layers)
|
||||
return cls(target_layers=target_layers)
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
from fastvideo.layers.linear import LinearBase
|
||||
fp4_layers = ["ffn.fc_in", "ffn.fc_out", "to_q", "to_k", "to_v", "to_out"]
|
||||
if isinstance(layer, LinearBase) and any(layer_name in prefix for layer_name in fp4_layers):
|
||||
if isinstance(layer, LinearBase) and any(name in prefix for name in self.target_layers):
|
||||
return NVFP4QATQuantizeMethod()
|
||||
return None
|
||||
|
||||
|
||||
@torch.compile
|
||||
def convert_model_to_fp4(model: torch.nn.Module):
|
||||
flashinfer_mod = _require_flashinfer()
|
||||
def convert_model_to_fp4(model: torch.nn.Module) -> None:
|
||||
"""Prequantize every FP4-tagged linear and drop its dense weight.
|
||||
|
||||
Walks the module tree, and for each layer whose ``quant_method`` is
|
||||
an :class:`NVFP4QATQuantizeMethod`, computes the NVFP4 packed weight
|
||||
/ scale / global-scale buffers, then pops the original fp16/bf16
|
||||
``weight`` parameter so it no longer occupies GPU memory.
|
||||
"""
|
||||
SfLayout, _, _ = _require_flashinfer()
|
||||
from torch.distributed.tensor import DTensor # type: ignore
|
||||
for mod in model.modules():
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
if isinstance(qm, NVFP4QATQuantizeMethod):
|
||||
|
||||
with torch.no_grad():
|
||||
for mod in model.modules():
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
if not isinstance(qm, NVFP4QATQuantizeMethod):
|
||||
continue
|
||||
|
||||
weight = getattr(mod, "weight", None)
|
||||
if weight is None:
|
||||
continue
|
||||
|
||||
weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type]
|
||||
weight_global_sf = (448 * 6) / weight_local.float().abs().nan_to_num().max()
|
||||
fp4_w, fp4_s = flashinfer_mod.nvfp4_quantize(
|
||||
|
||||
# Only the reduced scalar needs fp32; avoid a full fp32 copy.
|
||||
weight_absmax = (weight_local.detach().abs().nan_to_num().amax().to(dtype=torch.float32))
|
||||
weight_global_sf = (448 * 6) / weight_absmax
|
||||
fp4_w, fp4_s = _nvfp4_quantize(
|
||||
weight_local,
|
||||
weight_global_sf,
|
||||
sfLayout=flashinfer_mod.SfLayout.layout_128x4,
|
||||
sfLayout=SfLayout.layout_128x4,
|
||||
do_shuffle=False,
|
||||
)
|
||||
mod.register_buffer("_fp4_weight", fp4_w, persistent=False)
|
||||
mod.register_buffer("_fp4_weight_scale", fp4_s, persistent=False)
|
||||
mod.register_buffer("_weight_global_sf",
|
||||
torch.tensor(weight_global_sf, dtype=torch.bfloat16),
|
||||
persistent=False)
|
||||
mod.register_buffer(
|
||||
"_weight_global_sf",
|
||||
weight_global_sf.to(dtype=torch.bfloat16),
|
||||
persistent=False,
|
||||
)
|
||||
|
||||
# Drop the dense weight as soon as the fp4 buffers are installed
|
||||
# so it cannot keep occupying GPU memory.
|
||||
removed_weight = mod._parameters.pop("weight", None)
|
||||
if removed_weight is not None:
|
||||
removed_weight.grad = None
|
||||
del removed_weight, weight, weight_local, weight_absmax
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"NVFP4QATConfig",
|
||||
"NVFP4QATQuantizeMethod",
|
||||
"convert_model_to_fp4",
|
||||
"DEFAULT_FP4_LAYERS",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""NVFP4 quantization-aware *training* linear method (straight-through estimator).
|
||||
|
||||
The inference ``nvfp4_qat`` config quantizes each weight to FP4 once at load time
|
||||
(``convert_model_to_fp4``) and has no gradient path — it is inference only. For
|
||||
QAT *finetuning* the weight must stay a trainable bf16/fp32 master that is
|
||||
fake-quantized to FP4 on every forward, with a full-precision backward (a
|
||||
straight-through estimator), so the model learns to absorb FP4 linear error.
|
||||
|
||||
That STE already exists in ``fastvideo.layers.fp4linear._LinearFWD4BWD16Fn``
|
||||
(FP4 forward, full-precision backward) but is otherwise unwired. This method
|
||||
bridges it into the standard ``quant_config`` path, so it activates via
|
||||
``transformer_quant="nvfp4_qat_train"`` on the same Wan-2.1 layers as nvfp4_qat
|
||||
(to_q/k/v/out + ffn). No ``convert_model_to_fp4`` is needed: the weight is kept
|
||||
in full precision and quantized on the fly each step.
|
||||
"""
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig, QuantizeMethodBase
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class NVFP4QATTrainQuantizeMethod(QuantizeMethodBase):
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module, input_size_per_partition: int, output_partition_sizes: list[int],
|
||||
input_size: int, output_size: int, params_dtype: torch.dtype, **extra_weight_attrs):
|
||||
# Trainable master weight, fake-quantized to FP4 on each forward.
|
||||
weight = Parameter(torch.empty(
|
||||
sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype,
|
||||
),
|
||||
requires_grad=True)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, extra_weight_attrs)
|
||||
|
||||
def apply(self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None) -> torch.Tensor:
|
||||
# FP4 forward + full-precision backward (STE).
|
||||
from fastvideo.layers.fp4linear import _LinearFWD4BWD16Fn
|
||||
return _LinearFWD4BWD16Fn.apply(x, layer.weight, bias, "cutlass", 16, True)
|
||||
|
||||
|
||||
class NVFP4QATTrainConfig(QuantizationConfig):
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def get_name(self):
|
||||
return "nvfp4_qat_train"
|
||||
|
||||
def get_supported_act_dtypes(self):
|
||||
return [torch.bfloat16, torch.float16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls):
|
||||
return 100
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames():
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "NVFP4QATTrainConfig":
|
||||
return cls()
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
from fastvideo.layers.linear import LinearBase
|
||||
fp4_layers = ["ffn.fc_in", "ffn.fc_out", "to_q", "to_k", "to_v", "to_out"]
|
||||
if isinstance(layer, LinearBase) and any(layer_name in prefix for layer_name in fp4_layers):
|
||||
return NVFP4QATTrainQuantizeMethod()
|
||||
return None
|
||||
@@ -942,6 +942,20 @@ class TransformerLoader(ComponentLoader):
|
||||
dit_config = deepcopy(fastvideo_args.pipeline_config.dit_config)
|
||||
dit_config.update_model_arch(config)
|
||||
|
||||
# Generator-only QAT for DMD distillation: the teacher (real_score) and
|
||||
# critic (fake_score) transformers load with this flag set and must stay
|
||||
# full precision. Drop the nvfp4_qat quant from their copied config, and
|
||||
# mask the global ATTN_QAT_TRAIN env so their attention falls back to dense
|
||||
# (the backend is read globally at build time). The generator loads without
|
||||
# the flag and keeps both.
|
||||
_qat_generator_only = hasattr(fastvideo_args, "_loading_teacher_critic_model")
|
||||
_qat_prev_attn_env = None
|
||||
if _qat_generator_only:
|
||||
dit_config.quant_config = None
|
||||
from fastvideo.attention.selector import _cached_get_attn_backend
|
||||
_qat_prev_attn_env = os.environ.pop("FASTVIDEO_ATTENTION_BACKEND", None)
|
||||
_cached_get_attn_backend.cache_clear()
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||
|
||||
# Find all safetensors files
|
||||
@@ -1028,6 +1042,12 @@ class TransformerLoader(ComponentLoader):
|
||||
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs,
|
||||
)
|
||||
|
||||
if _qat_generator_only:
|
||||
from fastvideo.attention.selector import _cached_get_attn_backend
|
||||
if _qat_prev_attn_env is not None:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = _qat_prev_attn_env
|
||||
_cached_get_attn_backend.cache_clear()
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
|
||||
@@ -28,35 +28,43 @@ 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():
|
||||
if isinstance(getattr(mod, "quant_method", None),
|
||||
NVFP4QuantizeMethod):
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
if isinstance(qm, NVFP4QuantizeMethod):
|
||||
logger.info("Converting loaded model weights for NVFP4 linear layers")
|
||||
convert_model_to_nvfp4(model)
|
||||
return
|
||||
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
|
||||
|
||||
|
||||
# TODO(PY): move this to utils elsewhere
|
||||
@@ -189,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:
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Local dashboard helpers for FastVideo performance tracking."""
|
||||
@@ -0,0 +1,26 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run the local performance dashboard server."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
|
||||
import uvicorn
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Run the FastVideo performance dashboard")
|
||||
parser.add_argument("--host", default="127.0.0.1")
|
||||
parser.add_argument("--port", type=int, default=8000)
|
||||
parser.add_argument("--reload", action="store_true")
|
||||
args = parser.parse_args()
|
||||
uvicorn.run(
|
||||
"fastvideo.performance_dashboard.api:app",
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
reload=args.reload,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,207 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastAPI app for the local performance dashboard."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import threading
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import FastAPI, Query
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
|
||||
from fastvideo.tests.performance import hf_store
|
||||
|
||||
from .service import build_latest_summary, build_trends, filter_records
|
||||
|
||||
DEFAULT_TRACKING_ROOT = "/tmp/fastvideo-perf-dashboard"
|
||||
DEFAULT_DAYS = 90
|
||||
FRONTEND_DIST = os.path.abspath(
|
||||
os.path.join(os.path.dirname(__file__), "..", "..", "performance_dashboard", "frontend", "dist"))
|
||||
|
||||
|
||||
class PerformanceDataStore:
|
||||
|
||||
def __init__(self, tracking_root: str | None = None) -> None:
|
||||
self.tracking_root = tracking_root or os.environ.get("PERFORMANCE_TRACKING_ROOT", DEFAULT_TRACKING_ROOT)
|
||||
self._lock = threading.RLock()
|
||||
self.last_sync_at: str | None = None
|
||||
self.last_sync_error: str | None = None
|
||||
|
||||
@property
|
||||
def repo_id(self) -> str:
|
||||
return hf_store.HF_REPO_ID
|
||||
|
||||
def sync(self) -> dict[str, Any]:
|
||||
with self._lock:
|
||||
try:
|
||||
local_dir = hf_store.sync_from_hf(self.tracking_root, reuse_existing=False)
|
||||
self.last_sync_at = datetime.now(timezone.utc).isoformat()
|
||||
self.last_sync_error = None
|
||||
return {
|
||||
"ok": True,
|
||||
"repo_id": self.repo_id,
|
||||
"tracking_root": local_dir,
|
||||
"last_sync_at": self.last_sync_at,
|
||||
"last_sync_error": None,
|
||||
}
|
||||
except Exception as exc:
|
||||
self.last_sync_error = str(exc)
|
||||
return {
|
||||
"ok": False,
|
||||
"repo_id": self.repo_id,
|
||||
"tracking_root": self.tracking_root,
|
||||
"last_sync_at": self.last_sync_at,
|
||||
"last_sync_error": self.last_sync_error,
|
||||
}
|
||||
|
||||
def ensure_synced(self) -> None:
|
||||
if self.last_sync_at is not None:
|
||||
return
|
||||
with self._lock:
|
||||
if self.last_sync_at is None:
|
||||
hf_store.sync_from_hf(self.tracking_root, reuse_existing=True)
|
||||
self.last_sync_at = datetime.now(timezone.utc).isoformat()
|
||||
|
||||
def load_records(self, *, days: int | None = None, successful_only: bool = False) -> list[dict[str, Any]]:
|
||||
self.ensure_synced()
|
||||
return hf_store.load_records(self.tracking_root, days=days, successful_only=successful_only)
|
||||
|
||||
def health(self) -> dict[str, Any]:
|
||||
return {
|
||||
"ok": self.last_sync_error is None,
|
||||
"repo_id": self.repo_id,
|
||||
"tracking_root": self.tracking_root,
|
||||
"last_sync_at": self.last_sync_at,
|
||||
"last_sync_error": self.last_sync_error,
|
||||
}
|
||||
|
||||
|
||||
def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
data_store = store or PerformanceDataStore()
|
||||
app = FastAPI(title="FastVideo Performance Dashboard")
|
||||
app.state.performance_store = data_store
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=[
|
||||
"http://localhost:5173",
|
||||
"http://127.0.0.1:5173",
|
||||
"http://localhost:3000",
|
||||
"http://127.0.0.1:3000",
|
||||
],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
@app.get("/api/performance/health")
|
||||
def health() -> dict[str, Any]:
|
||||
return data_store.health()
|
||||
|
||||
@app.post("/api/performance/refresh")
|
||||
def refresh() -> dict[str, Any]:
|
||||
return data_store.sync()
|
||||
|
||||
@app.get("/api/performance/records")
|
||||
def records(
|
||||
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,
|
||||
run_source=run_source,
|
||||
success=success,
|
||||
)
|
||||
return {
|
||||
"records": filtered,
|
||||
"count": len(filtered),
|
||||
"filters": {
|
||||
"days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"run_source": run_source,
|
||||
"success": success,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
}
|
||||
|
||||
@app.get("/api/performance/summary")
|
||||
def summary(
|
||||
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``
|
||||
# query is kept only so the frontend can share filter state across
|
||||
# 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")),
|
||||
run_source=run_source,
|
||||
)
|
||||
return {
|
||||
"rows": rows,
|
||||
"count": len(rows),
|
||||
"status_counts": {
|
||||
"pass": sum(1 for row in rows if row["status"] == "pass"),
|
||||
"fail": sum(1 for row in rows if row["status"] == "fail"),
|
||||
},
|
||||
"filters": {
|
||||
"days": None,
|
||||
"trend_window_days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"run_source": run_source,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
}
|
||||
|
||||
@app.get("/api/performance/trends")
|
||||
def trends(
|
||||
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, run_source=run_source)
|
||||
groups = build_trends(filtered)
|
||||
return {
|
||||
"groups": groups,
|
||||
"count": len(groups),
|
||||
"filters": {
|
||||
"days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"run_source": run_source,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
}
|
||||
|
||||
assets_dir = os.path.join(FRONTEND_DIST, "assets")
|
||||
index_file = os.path.join(FRONTEND_DIST, "index.html")
|
||||
if os.path.isdir(assets_dir) and os.path.isfile(index_file):
|
||||
app.mount("/assets", StaticFiles(directory=assets_dir), name="performance-dashboard-assets")
|
||||
|
||||
@app.get("/{full_path:path}", include_in_schema=False)
|
||||
def frontend(full_path: str) -> FileResponse:
|
||||
return FileResponse(index_file)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
@@ -0,0 +1,26 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Metric definitions shared by the performance dashboard backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricDefinition:
|
||||
key: str
|
||||
label: str
|
||||
precision: int
|
||||
lower_is_better: bool
|
||||
|
||||
|
||||
METRICS: tuple[MetricDefinition, ...] = (
|
||||
MetricDefinition("latency", "Latency", 3, True),
|
||||
MetricDefinition("throughput", "Throughput", 3, False),
|
||||
MetricDefinition("memory", "Memory", 1, True),
|
||||
MetricDefinition("text_encoder_time_s", "Text Encoder", 3, True),
|
||||
MetricDefinition("dit_time_s", "DiT", 3, True),
|
||||
MetricDefinition("vae_decode_time_s", "VAE Decode", 3, True),
|
||||
)
|
||||
|
||||
METRIC_BY_KEY = {metric.key: metric for metric in METRICS}
|
||||
@@ -0,0 +1,198 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pure data transforms for the local performance dashboard.
|
||||
|
||||
The functions in this module operate on normalized records from
|
||||
``fastvideo/tests/performance/compare_baseline.py``. They intentionally avoid
|
||||
network and FastAPI concerns so they can be tested with in-memory fixtures.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import statistics
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.tests.performance.hf_store import is_baseline_eligible_record, safe_float
|
||||
|
||||
from .metrics import METRICS
|
||||
|
||||
Record = dict[str, Any]
|
||||
|
||||
|
||||
def parse_timestamp(value: Any) -> datetime | None:
|
||||
if not value:
|
||||
return None
|
||||
if isinstance(value, datetime):
|
||||
ts = value
|
||||
else:
|
||||
try:
|
||||
ts = datetime.fromisoformat(str(value))
|
||||
except ValueError:
|
||||
return None
|
||||
if ts.tzinfo is None:
|
||||
return ts.replace(tzinfo=timezone.utc)
|
||||
return ts.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def record_sort_key(record: Record) -> tuple[datetime, str]:
|
||||
ts = parse_timestamp(record.get("timestamp"))
|
||||
return (ts or datetime.min.replace(tzinfo=timezone.utc), str(record.get("commit_sha") or ""))
|
||||
|
||||
|
||||
def filter_records(
|
||||
records: list[Record],
|
||||
*,
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
run_source: str | None = None,
|
||||
success: bool | None = None,
|
||||
) -> list[Record]:
|
||||
filtered = records
|
||||
if model_id:
|
||||
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:
|
||||
model_id = str(record.get("model_id") or "unknown")
|
||||
gpu_type = str(record.get("gpu_type") or "unknown")
|
||||
groups[(model_id, gpu_type)].append(record)
|
||||
return {key: sorted(value, key=record_sort_key) for key, value in groups.items()}
|
||||
|
||||
|
||||
def baseline_value(records: list[Record], metric_key: str) -> float | None:
|
||||
values = [safe_float(record.get(metric_key)) for record in records]
|
||||
values = [value for value in values if value is not None]
|
||||
if not values:
|
||||
return None
|
||||
return float(statistics.median(values))
|
||||
|
||||
|
||||
def regression_percent(metric_key: str, current: float | None, baseline: float | None) -> float | None:
|
||||
if current is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
metric = next(metric for metric in METRICS if metric.key == metric_key)
|
||||
if metric.lower_is_better:
|
||||
return (current - baseline) / baseline * 100.0
|
||||
return (baseline - current) / baseline * 100.0
|
||||
|
||||
|
||||
def build_latest_summary(records: list[Record],
|
||||
*,
|
||||
baseline_window: int = 5,
|
||||
max_regression: float = 0.05,
|
||||
run_source: str | None = None) -> list[Record]:
|
||||
rows: list[Record] = []
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
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] = []
|
||||
for metric in METRICS:
|
||||
current = safe_float(latest.get(metric.key))
|
||||
baseline = baseline_value(baseline_records, metric.key)
|
||||
regression = regression_percent(metric.key, current, baseline)
|
||||
metrics[metric.key] = {
|
||||
"current": current,
|
||||
"baseline": baseline,
|
||||
"regression_pct": regression,
|
||||
"label": metric.label,
|
||||
"lower_is_better": metric.lower_is_better,
|
||||
"precision": metric.precision,
|
||||
}
|
||||
if regression is not None:
|
||||
regressions.append(regression)
|
||||
|
||||
worst_regression = max(regressions) if regressions else None
|
||||
success = bool(latest.get("success", True))
|
||||
status = "pass" if success else "fail"
|
||||
|
||||
rows.append({
|
||||
"model_id":
|
||||
model_id,
|
||||
"gpu_type":
|
||||
gpu_type,
|
||||
"timestamp":
|
||||
latest.get("timestamp"),
|
||||
"commit_sha":
|
||||
latest.get("commit_sha"),
|
||||
**record_metadata(latest),
|
||||
"success":
|
||||
success,
|
||||
"baseline_n":
|
||||
len(baseline_records),
|
||||
"worst_regression_pct":
|
||||
worst_regression,
|
||||
"regression_threshold_pct":
|
||||
max_regression * 100.0,
|
||||
"computed_regression_status":
|
||||
"fail" if worst_regression is not None and worst_regression > max_regression * 100.0 else "pass",
|
||||
"status":
|
||||
status,
|
||||
"metrics":
|
||||
metrics,
|
||||
})
|
||||
|
||||
return sorted(rows, key=lambda row: (row["status"] != "fail", row["model_id"], row["gpu_type"]))
|
||||
|
||||
|
||||
def build_trends(records: list[Record]) -> list[Record]:
|
||||
trends: list[Record] = []
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
points = []
|
||||
for record in group:
|
||||
point = {
|
||||
"timestamp": record.get("timestamp"),
|
||||
"commit_sha": record.get("commit_sha"),
|
||||
**record_metadata(record),
|
||||
"success": bool(record.get("success", True)),
|
||||
"metrics": {
|
||||
metric.key: safe_float(record.get(metric.key))
|
||||
for metric in METRICS
|
||||
},
|
||||
}
|
||||
points.append(point)
|
||||
trends.append({
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"points": points,
|
||||
})
|
||||
return sorted(trends, key=lambda trend: (trend["model_id"], trend["gpu_type"]))
|
||||
@@ -140,6 +140,23 @@ class CudaPlatformBase(Platform):
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info("Sage Attention 3 backend is not installed. Fall back to Flash Attention.")
|
||||
elif selected_backend == AttentionBackendEnum.ATTN_QAT_INFER:
|
||||
from fastvideo.attention.backends.attn_qat_infer import ( # noqa: F401
|
||||
AttnQatInferBackend, is_attn_qat_infer_available)
|
||||
if is_attn_qat_infer_available():
|
||||
logger.info("Using Attn-QAT inference (modified SageAttention3 FP4) backend.")
|
||||
return "fastvideo.attention.backends.attn_qat_infer.AttnQatInferBackend"
|
||||
logger.info("Attn-QAT inference kernel is not built. Fall back to Flash Attention.")
|
||||
elif selected_backend == AttentionBackendEnum.ATTN_QAT_TRAIN:
|
||||
from fastvideo.attention.backends.attn_qat_train import ( # noqa: F401
|
||||
AttnQatTrainBackend, is_attn_qat_train_available)
|
||||
if is_attn_qat_train_available():
|
||||
logger.info("Using Attn-QAT training (fake-quantized attention) backend.")
|
||||
return "fastvideo.attention.backends.attn_qat_train.AttnQatTrainBackend"
|
||||
raise ImportError(
|
||||
"ATTN_QAT_TRAIN selected but fastvideo_kernel.triton_kernels.attn_qat_train is not built. "
|
||||
"Silent fallback would produce a non-QAT training run; refusing to proceed. "
|
||||
"Install the training kernel or pick a different FASTVIDEO_ATTENTION_BACKEND.")
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from fastvideo_kernel import video_sparse_attn # noqa: F401
|
||||
|
||||
@@ -15,6 +15,8 @@ class AttentionBackendEnum(enum.Enum):
|
||||
TORCH_SDPA = enum.auto()
|
||||
SAGE_ATTN = enum.auto()
|
||||
SAGE_ATTN_THREE = enum.auto()
|
||||
ATTN_QAT_INFER = enum.auto()
|
||||
ATTN_QAT_TRAIN = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
BSA_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
|
||||
@@ -1,168 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the typed omni request plane (fastvideo/api/omni.py, design.md §6.1)."""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.api.omni import (
|
||||
AudioArtifact,
|
||||
DiffusionParams,
|
||||
ImagePart,
|
||||
Modality,
|
||||
NodeOverrides,
|
||||
OmniChunkEvent,
|
||||
OmniFinalEvent,
|
||||
OmniOutput,
|
||||
OmniProgressEvent,
|
||||
OmniRequest,
|
||||
OutputSpec,
|
||||
StreamSpec,
|
||||
TaskType,
|
||||
TextPart,
|
||||
VideoArtifact,
|
||||
infer_task,
|
||||
omni_event_from_video_event,
|
||||
)
|
||||
from fastvideo.api.results import (
|
||||
GenerationResult,
|
||||
VideoFinalEvent,
|
||||
VideoPartialEvent,
|
||||
VideoProgressEvent,
|
||||
)
|
||||
from fastvideo.api.schema import GenerationRequest, InputConfig, SamplingConfig
|
||||
|
||||
|
||||
def test_enums_are_strings():
|
||||
assert TaskType.T2V == "t2v"
|
||||
assert Modality.AUDIO == "audio"
|
||||
# str mixin keeps them usable anywhere a string is expected
|
||||
assert f"{TaskType.REASON}" == "TaskType.REASON" or TaskType.REASON.value == "reason"
|
||||
assert TaskType("i2v") is TaskType.I2V
|
||||
|
||||
|
||||
def test_request_id_autogenerated_and_unique():
|
||||
a = OmniRequest(task=TaskType.T2V)
|
||||
b = OmniRequest(task=TaskType.T2V)
|
||||
assert a.request_id and b.request_id
|
||||
assert a.request_id != b.request_id
|
||||
|
||||
|
||||
def test_prompt_and_negative_accessors():
|
||||
req = OmniRequest.from_prompt("a cat", TaskType.T2V, negative_prompt="blurry")
|
||||
assert req.prompt == "a cat"
|
||||
assert req.negative_prompt == "blurry"
|
||||
assert req.parts(Modality.TEXT)
|
||||
assert req.parts(Modality.IMAGE) == []
|
||||
|
||||
|
||||
def test_positional_text_part_is_text_not_role():
|
||||
# role is keyword-only, so the positional arg is the payload
|
||||
part = TextPart("hello")
|
||||
assert part.text == "hello"
|
||||
assert part.role is None
|
||||
|
||||
|
||||
def test_round_trip_through_generation_request():
|
||||
req = OmniRequest(
|
||||
task=TaskType.T2V,
|
||||
inputs=[TextPart("a fox"), TextPart("ugly", role="negative")],
|
||||
diffusion=DiffusionParams(steps=30, guidance_scale=6.0, height=480, width=832, num_frames=49),
|
||||
)
|
||||
legacy = req.to_generation_request()
|
||||
assert isinstance(legacy, GenerationRequest)
|
||||
assert legacy.prompt == "a fox"
|
||||
assert legacy.negative_prompt == "ugly"
|
||||
assert legacy.sampling.num_inference_steps == 30
|
||||
assert legacy.sampling.guidance_scale == 6.0
|
||||
assert legacy.sampling.height == 480
|
||||
assert legacy.sampling.num_frames == 49
|
||||
|
||||
back = OmniRequest.from_generation_request(legacy, task=TaskType.T2V)
|
||||
assert back.prompt == "a fox"
|
||||
assert back.negative_prompt == "ugly"
|
||||
assert back.diffusion.steps == 30
|
||||
assert back.diffusion.guidance_scale == 6.0
|
||||
assert back.diffusion.height == 480
|
||||
assert back.diffusion.num_frames == 49
|
||||
|
||||
|
||||
def test_image_input_lowers_and_lifts():
|
||||
req = OmniRequest(task=TaskType.I2V, inputs=[TextPart("dance"), ImagePart(path="/tmp/x.png")])
|
||||
legacy = req.to_generation_request()
|
||||
assert legacy.inputs.image_path == "/tmp/x.png"
|
||||
back = OmniRequest.from_generation_request(legacy)
|
||||
assert back.task == TaskType.I2V # inferred from the image input
|
||||
assert back.parts(Modality.IMAGE)[0].path == "/tmp/x.png" # type: ignore[attr-defined]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("inputs", "sampling", "expected"),
|
||||
[
|
||||
(InputConfig(), SamplingConfig(num_frames=49), TaskType.T2V),
|
||||
(InputConfig(image_path="/a.png"), SamplingConfig(num_frames=49), TaskType.I2V),
|
||||
(InputConfig(video_path="/a.mp4"), SamplingConfig(num_frames=49), TaskType.V2V),
|
||||
(InputConfig(), SamplingConfig(num_frames=1), TaskType.T2I),
|
||||
(InputConfig(image_path="/a.png"), SamplingConfig(num_frames=1), TaskType.I2I),
|
||||
],
|
||||
)
|
||||
def test_infer_task_heuristic(inputs, sampling, expected):
|
||||
req = GenerationRequest(prompt="x", inputs=inputs, sampling=sampling)
|
||||
assert infer_task(req) == expected
|
||||
|
||||
|
||||
def test_named_artifacts_replace_extra_audio():
|
||||
result = GenerationResult(
|
||||
prompt="song",
|
||||
frames="FRAMES",
|
||||
audio="WAV",
|
||||
audio_sample_rate=44100,
|
||||
generation_time=1.5,
|
||||
peak_memory_mb=1234.0,
|
||||
)
|
||||
out = OmniOutput.from_generation_result(result, request_id="r1")
|
||||
assert out.request_id == "r1"
|
||||
assert isinstance(out.video, VideoArtifact)
|
||||
assert out.video.frames == "FRAMES"
|
||||
# audio is a first-class artifact carrying its sample rate, not extra["audio"]
|
||||
assert isinstance(out.audio, AudioArtifact)
|
||||
assert out.audio.sample_rate == 44100
|
||||
assert out.audio.source_node == "audio_decode"
|
||||
assert out.metrics.generation_time == 1.5
|
||||
assert out.metrics.peak_memory_mb == 1234.0
|
||||
|
||||
|
||||
def test_omni_event_from_video_event():
|
||||
prog = omni_event_from_video_event(VideoProgressEvent(step=3, total_steps=50, stage="refine"))
|
||||
assert isinstance(prog, OmniProgressEvent)
|
||||
assert prog.step == 3 and prog.node == "refine"
|
||||
|
||||
chunk = omni_event_from_video_event(VideoPartialEvent(frames="CHUNK", index=2))
|
||||
assert isinstance(chunk, OmniChunkEvent)
|
||||
assert chunk.modality == Modality.VIDEO and chunk.index == 2 and chunk.payload == "CHUNK"
|
||||
|
||||
final = omni_event_from_video_event(
|
||||
VideoFinalEvent(result=GenerationResult(frames="F", audio="A", audio_sample_rate=24000)),
|
||||
request_id="rid",
|
||||
)
|
||||
assert isinstance(final, OmniFinalEvent)
|
||||
assert final.output.request_id == "rid"
|
||||
assert final.output.audio.sample_rate == 24000
|
||||
|
||||
|
||||
def test_node_overrides_access():
|
||||
req = OmniRequest(task=TaskType.T2V, node_params={"refine": NodeOverrides(params={"steps": 8})})
|
||||
assert req.node_params["refine"].steps == 8 # attribute access
|
||||
assert req.node_params["refine"]["steps"] == 8 # item access
|
||||
assert req.node_params["refine"].get("missing", 0) == 0
|
||||
with pytest.raises(AttributeError):
|
||||
_ = req.node_params["refine"].nonexistent
|
||||
# node_params survive the round-trip into stage_overrides
|
||||
legacy = req.to_generation_request()
|
||||
assert legacy.stage_overrides == {"refine": {"steps": 8}}
|
||||
|
||||
|
||||
def test_output_spec_streaming_flag():
|
||||
spec = OutputSpec(modalities=[Modality.VIDEO, Modality.AUDIO])
|
||||
assert spec.streaming is False
|
||||
spec.stream[Modality.AUDIO] = StreamSpec(enabled=True, chunk_ms=200)
|
||||
assert spec.streaming is True
|
||||
@@ -26,6 +26,12 @@ image = (modal.Image.from_registry(
|
||||
os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"BUILDKITE_BRANCH":
|
||||
os.environ.get("BUILDKITE_BRANCH", ""),
|
||||
"BUILDKITE_BUILD_URL":
|
||||
os.environ.get("BUILDKITE_BUILD_URL", ""),
|
||||
"BUILDKITE_BUILD_ID":
|
||||
os.environ.get("BUILDKITE_BUILD_ID", ""),
|
||||
"BUILDKITE_JOB_ID":
|
||||
os.environ.get("BUILDKITE_JOB_ID", ""),
|
||||
"TEST_SCOPE":
|
||||
os.environ.get("TEST_SCOPE", ""),
|
||||
"IMAGE_VERSION":
|
||||
@@ -337,18 +343,30 @@ def run_lora_extraction_tests():
|
||||
],
|
||||
volumes={"/root/data": model_vol})
|
||||
def run_performance_tests():
|
||||
# compare_baseline.py runs only after pytest passes, so normalized_perf_*.json
|
||||
# artifacts are emitted for rolling-baseline failures, not fixed-threshold
|
||||
# pytest failures. dashboard.py still runs on red CI for observability.
|
||||
# PR/direct records are uploaded only on pass; scheduled main uploads pass
|
||||
# and fail so the dashboard records every canonical baseline attempt.
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && "
|
||||
"export PERFORMANCE_TRACKING_ROOT='/tmp/perf-tracking' && "
|
||||
"hf auth login --token $HF_API_KEY && "
|
||||
"if [ \"${BUILDKITE_BRANCH:-}\" = 'main' ] && [ \"${TEST_SCOPE:-}\" = 'full' ]; then "
|
||||
"export PERF_RUN_SOURCE='scheduled_main'; "
|
||||
"export PERF_UPLOAD_POLICY='always'; "
|
||||
"elif [ -n \"${BUILDKITE_PULL_REQUEST:-}\" ] && [ \"${BUILDKITE_PULL_REQUEST:-false}\" != 'false' ]; then "
|
||||
"export PERF_RUN_SOURCE='pr'; "
|
||||
"export PERF_UPLOAD_POLICY='pass'; "
|
||||
"elif [ \"${TEST_SCOPE:-}\" = 'direct' ]; then "
|
||||
"export PERF_RUN_SOURCE='unknown'; "
|
||||
"export PERF_UPLOAD_POLICY='pass'; "
|
||||
"else "
|
||||
"export PERF_RUN_SOURCE='unknown'; "
|
||||
"export PERF_UPLOAD_POLICY='never'; "
|
||||
"fi; "
|
||||
"pytest ./fastvideo/tests/performance -vs; "
|
||||
"PYTEST_RC=$?; "
|
||||
"PERF_RC=0; "
|
||||
"if [ $PYTEST_RC -eq 0 ]; then "
|
||||
"python ./fastvideo/tests/performance/compare_baseline.py; "
|
||||
"if [ $PYTEST_RC -eq 0 ] || [ \"$PERF_UPLOAD_POLICY\" = 'always' ]; then "
|
||||
"PERF_PYTEST_RC=$PYTEST_RC python ./fastvideo/tests/performance/compare_baseline.py; "
|
||||
"PERF_RC=$?; "
|
||||
"fi; "
|
||||
"python ./fastvideo/tests/performance/dashboard.py || true; "
|
||||
|
||||
@@ -4,10 +4,10 @@
|
||||
This script:
|
||||
1) reads current benchmark results from fastvideo/tests/performance/results,
|
||||
2) syncs the canonical baseline from the configured HF dataset repo,
|
||||
3) compares each current record against the median of up to 5 prior records
|
||||
(filtered by gpu_type, successful only),
|
||||
4) on persist runs (full-suite on main branch), writes the normalized record
|
||||
back to the HF dataset repo,
|
||||
3) compares each current record against the median of up to 5 prior
|
||||
baseline-eligible successful records (filtered by gpu_type),
|
||||
4) writes normalized records back to the HF dataset repo according to
|
||||
PERF_UPLOAD_POLICY,
|
||||
5) exits non-zero if any metric regresses by more than PERF_MAX_REGRESSION
|
||||
(default 5%).
|
||||
"""
|
||||
@@ -20,13 +20,22 @@ import sys
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
try:
|
||||
from .hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
except ImportError:
|
||||
from hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
|
||||
RESULTS_DIR = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)),
|
||||
@@ -38,6 +47,9 @@ TRACKING_ROOT = os.environ.get(
|
||||
)
|
||||
PERF_REPORTS_DIR = os.environ.get("PERF_REPORTS_DIR", "/root/data/perf_reports")
|
||||
MAX_REGRESSION = float(os.environ.get("PERF_MAX_REGRESSION", "0.05"))
|
||||
UPLOAD_POLICY = os.environ.get("PERF_UPLOAD_POLICY", "never").strip().lower()
|
||||
VALID_UPLOAD_POLICIES = {"never", "pass", "always"}
|
||||
VALID_RUN_SOURCES = {"pr", "local", "scheduled_main", "unknown"}
|
||||
METRICS = (
|
||||
("latency", "Latency", 3),
|
||||
("throughput", "Throughput", 3),
|
||||
@@ -56,9 +68,74 @@ LOWER_IS_BETTER_METRICS = {
|
||||
|
||||
|
||||
def _should_persist_tracking() -> bool:
|
||||
test_scope = os.environ.get("TEST_SCOPE", "")
|
||||
branch = os.environ.get("BUILDKITE_BRANCH", "")
|
||||
return test_scope == "full" and branch == "main"
|
||||
return _normalized_upload_policy() != "never"
|
||||
|
||||
|
||||
def _normalized_upload_policy() -> str:
|
||||
if UPLOAD_POLICY in VALID_UPLOAD_POLICIES:
|
||||
return UPLOAD_POLICY
|
||||
print(f"Invalid PERF_UPLOAD_POLICY={UPLOAD_POLICY!r}; using 'never'")
|
||||
return "never"
|
||||
|
||||
|
||||
def _truthy_pr_number(value: str | None) -> bool:
|
||||
return bool(value and value not in {"false", "0", "None", "none"})
|
||||
|
||||
|
||||
def _detect_run_source() -> str:
|
||||
explicit = os.environ.get("PERF_RUN_SOURCE", "").strip().lower()
|
||||
if explicit in VALID_RUN_SOURCES:
|
||||
return explicit
|
||||
if explicit:
|
||||
print(f"Invalid PERF_RUN_SOURCE={explicit!r}; inferring run source")
|
||||
|
||||
if _truthy_pr_number(os.environ.get("BUILDKITE_PULL_REQUEST")):
|
||||
return "pr"
|
||||
if os.environ.get("BUILDKITE_BRANCH") == "main" and os.environ.get("TEST_SCOPE") == "full":
|
||||
return "scheduled_main"
|
||||
if not os.environ.get("BUILDKITE_COMMIT"):
|
||||
return "local"
|
||||
return "unknown"
|
||||
|
||||
|
||||
def _is_baseline_eligible(run_source: str, success: bool) -> bool:
|
||||
return run_source == "scheduled_main" and success
|
||||
|
||||
|
||||
def _upload_allowed(record: dict[str, Any]) -> bool:
|
||||
policy = _normalized_upload_policy()
|
||||
if policy == "always":
|
||||
return True
|
||||
if policy == "pass":
|
||||
return bool(record.get("success", True))
|
||||
return False
|
||||
|
||||
|
||||
def _result_failed_static_thresholds() -> bool:
|
||||
value = os.environ.get("PERF_PYTEST_RC", "")
|
||||
if not value:
|
||||
return False
|
||||
try:
|
||||
return int(value) != 0
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def _record_metadata(run_source: str, result: dict[str, Any]) -> dict[str, Any]:
|
||||
pr_number = result.get("pr_number") or os.environ.get("BUILDKITE_PULL_REQUEST", "")
|
||||
if not _truthy_pr_number(str(pr_number)):
|
||||
pr_number = ""
|
||||
return {
|
||||
"run_source": run_source,
|
||||
"baseline_eligible": False,
|
||||
"branch": os.environ.get("BUILDKITE_BRANCH", ""),
|
||||
"pr_number": pr_number,
|
||||
"test_scope": os.environ.get("TEST_SCOPE", ""),
|
||||
"build_url": os.environ.get("BUILDKITE_BUILD_URL", ""),
|
||||
"build_id": os.environ.get("BUILDKITE_BUILD_ID", ""),
|
||||
"job_id": os.environ.get("BUILDKITE_JOB_ID", ""),
|
||||
}
|
||||
|
||||
|
||||
def _load_current_results() -> list[dict[str, Any]]:
|
||||
pattern = os.path.join(RESULTS_DIR, "perf_*.json")
|
||||
@@ -104,6 +181,7 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
"dit_time_s": dit_time,
|
||||
"vae_decode_time_s": vae_decode_time,
|
||||
"success": True,
|
||||
**_record_metadata(_detect_run_source(), result),
|
||||
}
|
||||
|
||||
|
||||
@@ -297,8 +375,11 @@ def _emit_markdown_summary(markdown: str, commit_sha: str) -> None:
|
||||
|
||||
def main() -> int:
|
||||
persist_tracking = _should_persist_tracking()
|
||||
upload_policy = _normalized_upload_policy()
|
||||
static_threshold_failed = _result_failed_static_thresholds()
|
||||
|
||||
# Strict on persist: a silent sync failure would pollute the baseline.
|
||||
# Strict on upload-enabled runs: silent sync failure would make comparison
|
||||
# and upload state ambiguous.
|
||||
sync_from_hf(TRACKING_ROOT, strict=persist_tracking)
|
||||
|
||||
current_results = _load_current_results()
|
||||
@@ -310,10 +391,12 @@ def main() -> int:
|
||||
summary_rows: list[dict[str, Any]] = []
|
||||
|
||||
if persist_tracking:
|
||||
print("Tracking persistence enabled: full-suite run on main branch")
|
||||
print(f"Tracking persistence enabled: PERF_UPLOAD_POLICY={upload_policy}")
|
||||
else:
|
||||
print("Tracking persistence disabled: "
|
||||
"only full-suite runs on main branch are persisted")
|
||||
print("Tracking persistence disabled: PERF_UPLOAD_POLICY=never")
|
||||
|
||||
if static_threshold_failed:
|
||||
print(f"Static-threshold phase failed: PERF_PYTEST_RC={os.environ.get('PERF_PYTEST_RC')}")
|
||||
|
||||
for raw in current_results:
|
||||
record = _normalize_record(raw)
|
||||
@@ -324,6 +407,7 @@ def main() -> int:
|
||||
record["gpu_type"],
|
||||
last_n=5,
|
||||
successful_only=True,
|
||||
baseline_eligible_only=True,
|
||||
)
|
||||
|
||||
if not baseline_records:
|
||||
@@ -333,15 +417,22 @@ def main() -> int:
|
||||
record["success"] = True
|
||||
else:
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
record["success"] = not failures
|
||||
all_failures.extend(failures)
|
||||
if static_threshold_failed:
|
||||
failures.append(f"{record['model_id']} fixed-threshold phase failed "
|
||||
f"(PERF_PYTEST_RC={os.environ.get('PERF_PYTEST_RC')})")
|
||||
|
||||
record["success"] = not failures
|
||||
record["baseline_eligible"] = _is_baseline_eligible(record["run_source"], record["success"])
|
||||
all_failures.extend(failures)
|
||||
|
||||
_write_normalized_artifact(record)
|
||||
|
||||
# Strict upload: a silent failure would freeze the rolling baseline.
|
||||
if persist_tracking:
|
||||
if _upload_allowed(record):
|
||||
current_path = _write_tracking_record(record)
|
||||
upload_record(current_path, record, strict=True)
|
||||
else:
|
||||
print("Tracking upload skipped for "
|
||||
f"{record['model_id']} ({record['run_source']}, success={record['success']})")
|
||||
|
||||
summary_rows.append(_build_summary_row(record, baseline_records, bool(failures)))
|
||||
|
||||
|
||||
@@ -24,7 +24,7 @@ from huggingface_hub import HfApi, snapshot_download
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
HF_REPO_ID: str = os.environ.get("HF_REPO_ID", "FastVideo/performance-tracking")
|
||||
HF_TOKEN: str | None = os.environ.get("HF_API_KEY")
|
||||
HF_TOKEN_ENV_VARS = ("HF_API_KEY", "HUGGINGFACE_HUB_TOKEN", "HF_TOKEN")
|
||||
SYNC_MARKER = ".hf_sync_complete"
|
||||
SYNC_REUSE_TTL_SECONDS = int(os.environ.get("PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS", "3600"))
|
||||
|
||||
@@ -48,6 +48,29 @@ 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:
|
||||
token = os.environ.get(env_var)
|
||||
if token:
|
||||
return token
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HF I/O
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -120,7 +143,7 @@ def sync_from_hf(
|
||||
repo_id=HF_REPO_ID,
|
||||
repo_type="dataset",
|
||||
local_dir=local_dir,
|
||||
token=HF_TOKEN,
|
||||
token=resolve_hf_token(),
|
||||
allow_patterns="*.json",
|
||||
)
|
||||
os.makedirs(local_dir, exist_ok=True)
|
||||
@@ -150,8 +173,9 @@ def upload_record(
|
||||
that must not silently lose records — otherwise the rolling baseline can
|
||||
stop advancing without any signal in the build log.
|
||||
"""
|
||||
if not HF_TOKEN:
|
||||
msg = "hf_store: HF_API_KEY not set"
|
||||
token = resolve_hf_token()
|
||||
if not token:
|
||||
msg = f"hf_store: none of {', '.join(HF_TOKEN_ENV_VARS)} set"
|
||||
if strict:
|
||||
raise RuntimeError(f"{msg}; cannot upload.")
|
||||
print(f"{msg}, skipping upload.")
|
||||
@@ -161,7 +185,7 @@ def upload_record(
|
||||
path_in_repo = f"{sanitize(model_id)}/{os.path.basename(local_path)}"
|
||||
commit_sha = (record.get("commit_sha") or "unknown")[:7]
|
||||
|
||||
api = HfApi(token=HF_TOKEN)
|
||||
api = HfApi(token=token)
|
||||
try:
|
||||
api.upload_file(
|
||||
path_or_fileobj=local_path,
|
||||
@@ -187,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*.
|
||||
|
||||
@@ -196,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
|
||||
@@ -217,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:
|
||||
@@ -241,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.
|
||||
|
||||
@@ -251,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.
|
||||
@@ -259,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
|
||||
|
||||
@@ -0,0 +1,155 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastapi.testclient import TestClient
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from fastvideo.performance_dashboard.api import PerformanceDataStore, create_app
|
||||
|
||||
|
||||
class FakeStore(PerformanceDataStore):
|
||||
def __init__(self, records):
|
||||
super().__init__(tracking_root="/tmp/fake-fastvideo-perf-dashboard")
|
||||
self._records = records
|
||||
self.last_sync_at = "2026-01-03T00:00:00+00:00"
|
||||
|
||||
@property
|
||||
def repo_id(self):
|
||||
return "FastVideo/performance-tracking"
|
||||
|
||||
def sync(self):
|
||||
return {
|
||||
"ok": True,
|
||||
"repo_id": self.repo_id,
|
||||
"tracking_root": self.tracking_root,
|
||||
"last_sync_at": self.last_sync_at,
|
||||
"last_sync_error": None,
|
||||
}
|
||||
|
||||
def load_records(self, *, days=None, successful_only=False):
|
||||
records = list(self._records)
|
||||
if days is not None:
|
||||
latest_ts = max(datetime.fromisoformat(record["timestamp"]) for record in records) if records else None
|
||||
if latest_ts is not None:
|
||||
cutoff = latest_ts - timedelta(days=days)
|
||||
records = [record for record in records if datetime.fromisoformat(record["timestamp"]) >= cutoff]
|
||||
if successful_only:
|
||||
return [record for record in records if record.get("success", True)]
|
||||
return records
|
||||
|
||||
|
||||
def _record(model_id, gpu_type, ts, commit, latency, throughput, success=True, **metadata):
|
||||
record = {
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"timestamp": ts,
|
||||
"commit_sha": commit,
|
||||
"latency": latency,
|
||||
"throughput": throughput,
|
||||
"memory": 10000.0,
|
||||
"text_encoder_time_s": 2.0,
|
||||
"dit_time_s": 8.0,
|
||||
"vae_decode_time_s": 3.0,
|
||||
"success": success,
|
||||
}
|
||||
record.update(metadata)
|
||||
return record
|
||||
|
||||
|
||||
def test_summary_endpoint_returns_latest_group_status():
|
||||
app = create_app(FakeStore([
|
||||
_record("wan", "NVIDIA L40S", "2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("wan", "NVIDIA L40S", "2026-01-02T00:00:00+00:00", "b" * 40, 11.0, 9.0),
|
||||
]))
|
||||
client = TestClient(app)
|
||||
|
||||
response = client.get("/api/performance/summary")
|
||||
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["count"] == 1
|
||||
assert body["status_counts"] == {"pass": 1, "fail": 0}
|
||||
assert body["rows"][0]["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert body["rows"][0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
def test_summary_status_is_independent_of_days_window():
|
||||
app = create_app(FakeStore([
|
||||
_record("wan", "NVIDIA L40S", "2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("wan", "NVIDIA L40S", "2026-02-15T00:00:00+00:00", "b" * 40, 11.0, 9.0),
|
||||
]))
|
||||
client = TestClient(app)
|
||||
|
||||
narrow = client.get("/api/performance/summary", params={"days": 1}).json()
|
||||
wide = client.get("/api/performance/summary", params={"days": 365}).json()
|
||||
trends = client.get("/api/performance/trends", params={"days": 1}).json()
|
||||
|
||||
assert narrow["rows"][0]["status"] == wide["rows"][0]["status"]
|
||||
assert narrow["rows"][0]["baseline_n"] == wide["rows"][0]["baseline_n"] == 1
|
||||
assert narrow["filters"]["days"] is None
|
||||
assert narrow["filters"]["trend_window_days"] == 1
|
||||
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),
|
||||
_record("ltx", "NVIDIA A100", "2026-01-01T00:00:00+00:00", "b" * 40, 20.0, 5.0),
|
||||
]))
|
||||
client = TestClient(app)
|
||||
|
||||
records = client.get("/api/performance/records", params={"model_id": "wan"}).json()
|
||||
trends = client.get("/api/performance/trends", params={"gpu_type": "NVIDIA L40S"}).json()
|
||||
|
||||
assert records["count"] == 1
|
||||
assert records["records"][0]["model_id"] == "wan"
|
||||
assert trends["count"] == 1
|
||||
assert trends["groups"][0]["gpu_type"] == "NVIDIA L40S"
|
||||
|
||||
|
||||
def test_refresh_endpoint_reports_sync_metadata():
|
||||
app = create_app(FakeStore([]))
|
||||
client = TestClient(app)
|
||||
|
||||
response = client.post("/api/performance/refresh")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["ok"] is True
|
||||
@@ -0,0 +1,159 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.performance_dashboard.service import build_latest_summary, build_trends, filter_records
|
||||
from fastvideo.tests.performance import hf_store
|
||||
|
||||
|
||||
def _record(ts, commit, latency, throughput, success=True, **metadata):
|
||||
record = {
|
||||
"model_id": "wan-t2v-1.3b-2gpu",
|
||||
"gpu_type": "NVIDIA L40S",
|
||||
"timestamp": ts,
|
||||
"commit_sha": commit,
|
||||
"latency": latency,
|
||||
"throughput": throughput,
|
||||
"memory": 10000.0,
|
||||
"text_encoder_time_s": None,
|
||||
"dit_time_s": 8.0,
|
||||
"vae_decode_time_s": 3.0,
|
||||
"success": success,
|
||||
}
|
||||
record.update(metadata)
|
||||
return record
|
||||
|
||||
|
||||
def test_build_latest_summary_uses_previous_successful_records_for_baseline():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("2026-01-02T00:00:00+00:00", "b" * 40, 12.0, 8.0, success=False),
|
||||
_record("2026-01-03T00:00:00+00:00", "c" * 40, 11.0, 9.0),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05)
|
||||
|
||||
assert len(rows) == 1
|
||||
row = rows[0]
|
||||
assert row["baseline_n"] == 1
|
||||
assert row["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert row["metrics"]["latency"]["regression_pct"] == 10.0
|
||||
assert row["metrics"]["throughput"]["regression_pct"] == 10.0
|
||||
assert row["status"] == "pass"
|
||||
assert row["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
def test_build_latest_summary_status_uses_latest_record_success_field():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("2026-01-02T00:00:00+00:00", "b" * 40, 10.0, 10.0, success=False),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert rows[0]["status"] == "fail"
|
||||
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),
|
||||
_record("2026-01-02T00:00:00+00:00", "b" * 40, 12.0, 8.0, success=False),
|
||||
]
|
||||
|
||||
failed = filter_records(records, success=False)
|
||||
trends = build_trends(records)
|
||||
|
||||
assert [record["commit_sha"] for record in failed] == ["b" * 40]
|
||||
assert len(trends) == 1
|
||||
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)
|
||||
|
||||
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,
|
||||
}
|
||||
@@ -0,0 +1,104 @@
|
||||
# FastVideo Performance Dashboard
|
||||
|
||||
Local FastAPI + React dashboard for records stored in the Hugging Face
|
||||
performance tracking dataset.
|
||||
|
||||
## Data Source
|
||||
|
||||
The dashboard reads the same normalized JSON records used by
|
||||
`fastvideo/tests/performance/compare_baseline.py`.
|
||||
|
||||
Defaults:
|
||||
|
||||
- `HF_REPO_ID=FastVideo/performance-tracking`
|
||||
- `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:
|
||||
|
||||
```bash
|
||||
export HF_TOKEN=hf_...
|
||||
```
|
||||
|
||||
If Hugging Face returns `401 Unauthorized`, confirm that `HF_REPO_ID` points to
|
||||
the dataset repo you expect and that your token has access to it.
|
||||
|
||||
## Development
|
||||
|
||||
Run the API:
|
||||
|
||||
```bash
|
||||
python -m fastvideo.performance_dashboard --host 127.0.0.1 --port 8000 --reload
|
||||
```
|
||||
|
||||
Run the React dev server:
|
||||
|
||||
```bash
|
||||
cd performance_dashboard/frontend
|
||||
npm install
|
||||
npm run dev
|
||||
```
|
||||
|
||||
Open `http://127.0.0.1:5173`. Vite proxies `/api/*` to the FastAPI server on
|
||||
port 8000.
|
||||
|
||||
## Single-Port Mode For ngrok
|
||||
|
||||
Build the frontend:
|
||||
|
||||
```bash
|
||||
cd performance_dashboard/frontend
|
||||
npm install
|
||||
npm run build
|
||||
```
|
||||
|
||||
Serve API and built frontend from one FastAPI process:
|
||||
|
||||
```bash
|
||||
python -m fastvideo.performance_dashboard --host 0.0.0.0 --port 8000
|
||||
```
|
||||
|
||||
Expose it:
|
||||
|
||||
```bash
|
||||
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&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 for dashboard
|
||||
context. CI gating uses only records marked `baseline_eligible=true`.
|
||||
@@ -0,0 +1,3 @@
|
||||
node_modules/
|
||||
dist/
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
<!doctype html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>FastVideo Performance Dashboard</title>
|
||||
</head>
|
||||
<body>
|
||||
<div id="root"></div>
|
||||
<script type="module" src="/src/main.tsx"></script>
|
||||
</body>
|
||||
</html>
|
||||
|
||||
+1824
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,23 @@
|
||||
{
|
||||
"name": "fastvideo-performance-dashboard",
|
||||
"version": "0.1.0",
|
||||
"private": true,
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"dev": "vite dev --host 0.0.0.0",
|
||||
"build": "tsc && node scripts/build.mjs",
|
||||
"preview": "vite preview --host 0.0.0.0"
|
||||
},
|
||||
"dependencies": {
|
||||
"react": "^19.2.3",
|
||||
"react-dom": "^19.2.3"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@types/react": "^19.2.7",
|
||||
"@types/react-dom": "^19.2.3",
|
||||
"@vitejs/plugin-react": "^5.1.1",
|
||||
"esbuild": "^0.25.12",
|
||||
"typescript": "^5.9.3",
|
||||
"vite": "^6.4.2"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
import * as esbuild from "esbuild";
|
||||
import { mkdir, writeFile } from "node:fs/promises";
|
||||
import { resolve } from "node:path";
|
||||
|
||||
const root = resolve(import.meta.dirname, "..");
|
||||
const dist = resolve(root, "dist");
|
||||
const assets = resolve(dist, "assets");
|
||||
|
||||
await mkdir(assets, { recursive: true });
|
||||
|
||||
await esbuild.build({
|
||||
entryPoints: [resolve(root, "src/main.tsx")],
|
||||
bundle: true,
|
||||
format: "esm",
|
||||
minify: true,
|
||||
sourcemap: true,
|
||||
target: ["es2020"],
|
||||
outfile: resolve(assets, "dashboard.js"),
|
||||
loader: {
|
||||
".svg": "file"
|
||||
}
|
||||
});
|
||||
|
||||
await writeFile(
|
||||
resolve(dist, "index.html"),
|
||||
`<!doctype html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>FastVideo Performance Dashboard</title>
|
||||
<link rel="stylesheet" href="/assets/dashboard.css" />
|
||||
</head>
|
||||
<body>
|
||||
<div id="root"></div>
|
||||
<script type="module" src="/assets/dashboard.js"></script>
|
||||
</body>
|
||||
</html>
|
||||
`
|
||||
);
|
||||
@@ -0,0 +1,506 @@
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
|
||||
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)) {
|
||||
return "n/a";
|
||||
}
|
||||
return value.toFixed(precision);
|
||||
}
|
||||
|
||||
function shortSha(value: string | null | undefined) {
|
||||
return value ? value.slice(0, 7) : "unknown";
|
||||
}
|
||||
|
||||
function formatTime(value: string | null | undefined) {
|
||||
if (!value) {
|
||||
return "never";
|
||||
}
|
||||
const date = new Date(value);
|
||||
if (Number.isNaN(date.getTime())) {
|
||||
return value;
|
||||
}
|
||||
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) => ({
|
||||
point,
|
||||
value: point.metrics[metricKey]
|
||||
}))
|
||||
.filter((point) => point.value !== null && point.value !== undefined) as Array<{
|
||||
point: TrendPoint;
|
||||
value: number;
|
||||
}>;
|
||||
|
||||
if (points.length === 0) {
|
||||
return <div className="empty-chart">No data</div>;
|
||||
}
|
||||
|
||||
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 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 (
|
||||
<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>
|
||||
);
|
||||
}
|
||||
|
||||
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);
|
||||
const [refreshing, setRefreshing] = useState(false);
|
||||
const [error, setError] = useState<string | null>(null);
|
||||
|
||||
async function load() {
|
||||
setLoading(true);
|
||||
setError(null);
|
||||
try {
|
||||
const [summaryData, trendData] = await Promise.all([
|
||||
fetchSummary(days, modelFilter || undefined, gpuFilter || undefined, sourceFilter || undefined),
|
||||
fetchTrends(days, modelFilter || undefined, gpuFilter || undefined, sourceFilter || undefined)
|
||||
]);
|
||||
setSummary(summaryData);
|
||||
setTrends(trendData.groups);
|
||||
} catch (err) {
|
||||
setError(err instanceof Error ? err.message : String(err));
|
||||
} finally {
|
||||
setLoading(false);
|
||||
}
|
||||
}
|
||||
|
||||
async function refresh() {
|
||||
setRefreshing(true);
|
||||
setError(null);
|
||||
try {
|
||||
await refreshData();
|
||||
await load();
|
||||
} catch (err) {
|
||||
setError(err instanceof Error ? err.message : String(err));
|
||||
} finally {
|
||||
setRefreshing(false);
|
||||
}
|
||||
}
|
||||
|
||||
useEffect(() => {
|
||||
load();
|
||||
const interval = window.setInterval(load, 5 * 60 * 1000);
|
||||
return () => window.clearInterval(interval);
|
||||
}, [days, modelFilter, gpuFilter, sourceFilter]);
|
||||
|
||||
const models = useMemo(() => {
|
||||
const values = new Set(summary?.rows.map((row) => row.model_id) ?? []);
|
||||
trends.forEach((trend) => values.add(trend.model_id));
|
||||
return [...values].sort();
|
||||
}, [summary, trends]);
|
||||
|
||||
const gpus = useMemo(() => {
|
||||
const values = new Set(summary?.rows.map((row) => row.gpu_type) ?? []);
|
||||
trends.forEach((trend) => values.add(trend.gpu_type));
|
||||
return [...values].sort();
|
||||
}, [summary, trends]);
|
||||
|
||||
const latestRows = summary?.rows ?? [];
|
||||
const totalRuns = trends.reduce((total, group) => total + group.points.length, 0);
|
||||
const sync = summary?.sync;
|
||||
|
||||
return (
|
||||
<main className="dashboard">
|
||||
<header className="topbar">
|
||||
<div>
|
||||
<p className="eyebrow">FastVideo CI</p>
|
||||
<h1>Performance Dashboard</h1>
|
||||
</div>
|
||||
<button className="refresh-button" onClick={refresh} disabled={refreshing || loading}>
|
||||
{refreshing ? "Refreshing" : "Refresh"}
|
||||
</button>
|
||||
</header>
|
||||
|
||||
<section className="filters" aria-label="Filters">
|
||||
<label>
|
||||
Days
|
||||
<input
|
||||
type="number"
|
||||
min="1"
|
||||
max="3650"
|
||||
value={days}
|
||||
onChange={(event) => setDays(Number(event.target.value) || 90)}
|
||||
/>
|
||||
</label>
|
||||
<label>
|
||||
Model
|
||||
<select value={modelFilter} onChange={(event) => setModelFilter(event.target.value)}>
|
||||
<option value="">All models</option>
|
||||
{models.map((model) => (
|
||||
<option key={model} value={model}>
|
||||
{model}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</label>
|
||||
<label>
|
||||
GPU
|
||||
<select value={gpuFilter} onChange={(event) => setGpuFilter(event.target.value)}>
|
||||
<option value="">All GPUs</option>
|
||||
{gpus.map((gpu) => (
|
||||
<option key={gpu} value={gpu}>
|
||||
{gpu}
|
||||
</option>
|
||||
))}
|
||||
</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>}
|
||||
{loading && <div className="notice">Loading performance data</div>}
|
||||
|
||||
<section className="cards" aria-label="Overview">
|
||||
<div className="stat">
|
||||
<span>Groups</span>
|
||||
<strong>{summary?.count ?? 0}</strong>
|
||||
</div>
|
||||
<div className="stat">
|
||||
<span>Failing</span>
|
||||
<strong>{summary?.status_counts.fail ?? 0}</strong>
|
||||
</div>
|
||||
<div className="stat">
|
||||
<span>Runs</span>
|
||||
<strong>{totalRuns}</strong>
|
||||
</div>
|
||||
<div className="stat wide">
|
||||
<span>Last sync</span>
|
||||
<strong>{formatTime(sync?.last_sync_at)}</strong>
|
||||
<small>{sync?.repo_id ?? "FastVideo/performance-tracking"}</small>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section className="panel">
|
||||
<div className="panel-header">
|
||||
<h2>Latest Status</h2>
|
||||
<span>{latestRows.length} model/GPU groups</span>
|
||||
</div>
|
||||
{latestRows.length === 0 ? (
|
||||
<div className="empty">No records match the selected filters.</div>
|
||||
) : (
|
||||
<div className="table-wrap">
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Stored Status</th>
|
||||
<th>Recomputed</th>
|
||||
<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>
|
||||
<th>Memory</th>
|
||||
<th>Worst</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{latestRows.map((row) => (
|
||||
<tr key={`${row.model_id}-${row.gpu_type}`}>
|
||||
<td>
|
||||
<span className={`badge ${row.status}`}>{row.status}</span>
|
||||
</td>
|
||||
<td>
|
||||
<span className={`badge muted ${row.computed_regression_status}`}>
|
||||
{row.computed_regression_status}
|
||||
</span>
|
||||
</td>
|
||||
<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>
|
||||
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
|
||||
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
)}
|
||||
</section>
|
||||
|
||||
<section className="panel">
|
||||
<div className="panel-header">
|
||||
<h2>Trends</h2>
|
||||
<span>{days} day window</span>
|
||||
</div>
|
||||
<div className="trend-grid">
|
||||
{trends.length === 0 ? (
|
||||
<div className="empty full-width">
|
||||
No trend records found in the selected time window. Increase the day range or refresh after new CI
|
||||
performance records are uploaded.
|
||||
</div>
|
||||
) : (
|
||||
trends.map((group) =>
|
||||
METRIC_KEYS.map((metricKey) => (
|
||||
<article className="trend-card" key={`${group.model_id}-${group.gpu_type}-${metricKey}`}>
|
||||
<div>
|
||||
<h3>{metricLabel(metricKey)}</h3>
|
||||
<p>
|
||||
{group.model_id} | {group.gpu_type}
|
||||
</p>
|
||||
</div>
|
||||
<TrendChart group={group} metricKey={metricKey} />
|
||||
</article>
|
||||
))
|
||||
)
|
||||
)}
|
||||
</div>
|
||||
</section>
|
||||
</main>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
export type MetricValue = {
|
||||
current: number | null;
|
||||
baseline: number | null;
|
||||
regression_pct: number | null;
|
||||
label: string;
|
||||
lower_is_better: boolean;
|
||||
precision: number;
|
||||
};
|
||||
|
||||
export type SummaryRow = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
timestamp: string | null;
|
||||
commit_sha: string | null;
|
||||
success: boolean;
|
||||
baseline_n: number;
|
||||
worst_regression_pct: number | null;
|
||||
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;
|
||||
status_counts: {
|
||||
pass: number;
|
||||
fail: number;
|
||||
};
|
||||
filters: {
|
||||
days: number | null;
|
||||
trend_window_days?: number;
|
||||
model_id: string | null;
|
||||
gpu_type: string | null;
|
||||
run_source: string | null;
|
||||
};
|
||||
sync: SyncState;
|
||||
};
|
||||
|
||||
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>;
|
||||
};
|
||||
|
||||
export type TrendGroup = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
points: TrendPoint[];
|
||||
};
|
||||
|
||||
export type TrendsResponse = {
|
||||
groups: TrendGroup[];
|
||||
count: number;
|
||||
sync: SyncState;
|
||||
};
|
||||
|
||||
export type SyncState = {
|
||||
ok: boolean;
|
||||
repo_id: string;
|
||||
tracking_root: string;
|
||||
last_sync_at: string | null;
|
||||
last_sync_error: string | null;
|
||||
};
|
||||
|
||||
const jsonHeaders = {
|
||||
Accept: "application/json"
|
||||
};
|
||||
|
||||
function params(values: Record<string, string | number | null | undefined>) {
|
||||
const out = new URLSearchParams();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
if (value !== null && value !== undefined && value !== "") {
|
||||
out.set(key, String(value));
|
||||
}
|
||||
}
|
||||
return out.toString();
|
||||
}
|
||||
|
||||
async function getJson<T>(path: string): Promise<T> {
|
||||
const response = await fetch(path, { headers: jsonHeaders });
|
||||
if (!response.ok) {
|
||||
throw new Error(`${response.status} ${response.statusText}`);
|
||||
}
|
||||
return response.json() as Promise<T>;
|
||||
}
|
||||
|
||||
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, run_source: runSource })}`
|
||||
);
|
||||
}
|
||||
|
||||
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, run_source: runSource })}`
|
||||
);
|
||||
}
|
||||
|
||||
export async function refreshData() {
|
||||
const response = await fetch("/api/performance/refresh", { method: "POST", headers: jsonHeaders });
|
||||
if (!response.ok) {
|
||||
throw new Error(`${response.status} ${response.statusText}`);
|
||||
}
|
||||
return response.json() as Promise<SyncState>;
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
import React from "react";
|
||||
import { createRoot } from "react-dom/client";
|
||||
|
||||
import App from "./App";
|
||||
import "./styles.css";
|
||||
|
||||
createRoot(document.getElementById("root") as HTMLElement).render(
|
||||
<React.StrictMode>
|
||||
<App />
|
||||
</React.StrictMode>
|
||||
);
|
||||
|
||||
@@ -0,0 +1,414 @@
|
||||
:root {
|
||||
color: #1f2933;
|
||||
background: #eef2f5;
|
||||
font-family:
|
||||
Inter, ui-sans-serif, system-ui, -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif;
|
||||
line-height: 1.4;
|
||||
}
|
||||
|
||||
* {
|
||||
box-sizing: border-box;
|
||||
}
|
||||
|
||||
body {
|
||||
margin: 0;
|
||||
min-width: 320px;
|
||||
}
|
||||
|
||||
button,
|
||||
input,
|
||||
select {
|
||||
font: inherit;
|
||||
}
|
||||
|
||||
.dashboard {
|
||||
width: min(1440px, 100%);
|
||||
margin: 0 auto;
|
||||
padding: 28px;
|
||||
}
|
||||
|
||||
.topbar {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
gap: 20px;
|
||||
margin-bottom: 22px;
|
||||
}
|
||||
|
||||
.eyebrow {
|
||||
margin: 0 0 4px;
|
||||
color: #607080;
|
||||
font-size: 0.78rem;
|
||||
font-weight: 700;
|
||||
letter-spacing: 0;
|
||||
text-transform: uppercase;
|
||||
}
|
||||
|
||||
h1,
|
||||
h2,
|
||||
h3,
|
||||
p {
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
h1 {
|
||||
color: #12202f;
|
||||
font-size: clamp(2rem, 4vw, 3.5rem);
|
||||
letter-spacing: 0;
|
||||
}
|
||||
|
||||
h2 {
|
||||
color: #182736;
|
||||
font-size: 1.08rem;
|
||||
}
|
||||
|
||||
h3 {
|
||||
color: #233242;
|
||||
font-size: 0.94rem;
|
||||
}
|
||||
|
||||
.refresh-button {
|
||||
min-height: 42px;
|
||||
border: 1px solid #0f6b8f;
|
||||
border-radius: 6px;
|
||||
padding: 0 18px;
|
||||
color: #ffffff;
|
||||
background: #0f6b8f;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.refresh-button:disabled {
|
||||
cursor: default;
|
||||
opacity: 0.6;
|
||||
}
|
||||
|
||||
.filters {
|
||||
display: grid;
|
||||
grid-template-columns: 120px minmax(200px, 1fr) minmax(200px, 1fr) minmax(180px, 0.8fr);
|
||||
gap: 14px;
|
||||
margin-bottom: 18px;
|
||||
}
|
||||
|
||||
.filters label {
|
||||
display: grid;
|
||||
gap: 6px;
|
||||
color: #425466;
|
||||
font-size: 0.82rem;
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.filters input,
|
||||
.filters select {
|
||||
width: 100%;
|
||||
min-height: 40px;
|
||||
border: 1px solid #c8d2dc;
|
||||
border-radius: 6px;
|
||||
padding: 0 10px;
|
||||
color: #17212b;
|
||||
background: #ffffff;
|
||||
}
|
||||
|
||||
.notice {
|
||||
border: 1px solid #c8d2dc;
|
||||
border-radius: 6px;
|
||||
margin-bottom: 16px;
|
||||
padding: 12px 14px;
|
||||
background: #ffffff;
|
||||
}
|
||||
|
||||
.notice.error {
|
||||
border-color: #d94f4f;
|
||||
color: #8a1f1f;
|
||||
background: #fff4f4;
|
||||
}
|
||||
|
||||
.cards {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, minmax(120px, 1fr)) minmax(260px, 1.8fr);
|
||||
gap: 12px;
|
||||
margin-bottom: 18px;
|
||||
}
|
||||
|
||||
.stat,
|
||||
.panel,
|
||||
.trend-card {
|
||||
border: 1px solid #d5dde5;
|
||||
border-radius: 8px;
|
||||
background: #ffffff;
|
||||
}
|
||||
|
||||
.stat {
|
||||
min-height: 96px;
|
||||
padding: 16px;
|
||||
}
|
||||
|
||||
.stat span,
|
||||
.panel-header span,
|
||||
.trend-card p {
|
||||
color: #607080;
|
||||
font-size: 0.82rem;
|
||||
}
|
||||
|
||||
.stat strong {
|
||||
display: block;
|
||||
margin-top: 8px;
|
||||
color: #132232;
|
||||
font-size: 1.9rem;
|
||||
}
|
||||
|
||||
.stat.wide strong {
|
||||
font-size: 1rem;
|
||||
}
|
||||
|
||||
.stat small {
|
||||
display: block;
|
||||
margin-top: 8px;
|
||||
color: #607080;
|
||||
}
|
||||
|
||||
.panel {
|
||||
margin-bottom: 18px;
|
||||
overflow: hidden;
|
||||
}
|
||||
|
||||
.panel-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
gap: 12px;
|
||||
border-bottom: 1px solid #d5dde5;
|
||||
padding: 14px 16px;
|
||||
}
|
||||
|
||||
.table-wrap {
|
||||
overflow-x: auto;
|
||||
}
|
||||
|
||||
table {
|
||||
width: 100%;
|
||||
min-width: 1120px;
|
||||
border-collapse: collapse;
|
||||
}
|
||||
|
||||
th,
|
||||
td {
|
||||
border-bottom: 1px solid #e6ebf0;
|
||||
padding: 11px 12px;
|
||||
text-align: left;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
th {
|
||||
color: #536171;
|
||||
font-size: 0.78rem;
|
||||
text-transform: uppercase;
|
||||
}
|
||||
|
||||
td {
|
||||
color: #1b2836;
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
|
||||
.badge {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
min-width: 52px;
|
||||
border-radius: 999px;
|
||||
padding: 4px 9px;
|
||||
font-size: 0.75rem;
|
||||
font-weight: 800;
|
||||
text-transform: uppercase;
|
||||
}
|
||||
|
||||
.badge.pass {
|
||||
color: #166534;
|
||||
background: #dcfce7;
|
||||
}
|
||||
|
||||
.badge.fail {
|
||||
color: #991b1b;
|
||||
background: #fee2e2;
|
||||
}
|
||||
|
||||
.badge.muted {
|
||||
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;
|
||||
}
|
||||
|
||||
.empty.full-width {
|
||||
grid-column: 1 / -1;
|
||||
}
|
||||
|
||||
.trend-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fill, minmax(340px, 1fr));
|
||||
gap: 12px;
|
||||
padding: 14px;
|
||||
}
|
||||
|
||||
.trend-card {
|
||||
display: grid;
|
||||
gap: 10px;
|
||||
min-width: 0;
|
||||
padding: 14px;
|
||||
}
|
||||
|
||||
.trend-chart {
|
||||
width: 100%;
|
||||
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 {
|
||||
display: grid;
|
||||
min-height: 96px;
|
||||
place-items: center;
|
||||
color: #72808f;
|
||||
background: #f5f7f9;
|
||||
}
|
||||
|
||||
@media (max-width: 760px) {
|
||||
.dashboard {
|
||||
padding: 18px;
|
||||
}
|
||||
|
||||
.topbar,
|
||||
.panel-header {
|
||||
align-items: flex-start;
|
||||
flex-direction: column;
|
||||
}
|
||||
|
||||
.refresh-button {
|
||||
width: 100%;
|
||||
}
|
||||
|
||||
.filters,
|
||||
.cards {
|
||||
grid-template-columns: 1fr;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
{
|
||||
"compilerOptions": {
|
||||
"target": "ES2022",
|
||||
"useDefineForClassFields": true,
|
||||
"lib": ["DOM", "DOM.Iterable", "ES2022"],
|
||||
"allowJs": false,
|
||||
"skipLibCheck": true,
|
||||
"esModuleInterop": true,
|
||||
"allowSyntheticDefaultImports": true,
|
||||
"strict": true,
|
||||
"forceConsistentCasingInFileNames": true,
|
||||
"module": "ESNext",
|
||||
"moduleResolution": "Node",
|
||||
"resolveJsonModule": true,
|
||||
"isolatedModules": true,
|
||||
"noEmit": true,
|
||||
"jsx": "react-jsx"
|
||||
},
|
||||
"include": ["src"],
|
||||
"references": []
|
||||
}
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
import react from "@vitejs/plugin-react";
|
||||
import { defineConfig } from "vite";
|
||||
|
||||
export default defineConfig({
|
||||
plugins: [react()],
|
||||
server: {
|
||||
port: 5173,
|
||||
proxy: {
|
||||
"/api": "http://127.0.0.1:8000"
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
+1
-1
@@ -220,7 +220,7 @@ follow_imports = "silent"
|
||||
# ``*/_vendored/*`` matches upstream-provenance files vendored under any
|
||||
# ``_vendored/`` subdir (project-wide convention; mirrors the
|
||||
# ``_``-prefixed auto-discovery skip).
|
||||
skip = "./data,./wandb,ui/package-lock.json,*/_vendored/*"
|
||||
skip = "./data,./wandb,ui/package-lock.json,performance_dashboard/frontend/package-lock.json,*/_vendored/*"
|
||||
# "tread" matches daVinci-MagiHuman's acronym "TReAD" (Token Routing and
|
||||
# Early Drop). codespell lowercases ignore-words entries, so the single
|
||||
# lowercase form silences all case variants.
|
||||
|
||||
Reference in New Issue
Block a user