Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 5 additions & 7 deletions aiter/aot/flydsl/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ the JIT path hits the cache instead of compiling again.
| `grouped_moe.py` | `GROUPED_MOE` | gfx1250 grouped MoE GEMM kernels |
| `chunk_gdn_h.py` | `CHUNK_GDN_H` | chunk-gdn-h opt (K5) kernels |
| `mega_moe.py` | `MEGA_MOE` | MegaMoE A8W4 profile bundles for MTPR 8192/16384/32768 |
| `fmha_fp8.py` | `FMHA_FP8` | gfx950 FP8 flash-attention forward, per head shape listed in `configs/fmha_fp8_aot.csv` + `configs/model_configs/*_fmha_fp8_aot.csv` |
| `fmha_fp8.py` | `FMHA_FP8` | gfx950 FP8 flash-attention forward, per head shape in `DEFAULT_SHAPES` |
| `common.py` | — | Shared job collection, the deadlock-free fork pool, and cache-hit checking logic |

---
Expand Down Expand Up @@ -64,16 +64,14 @@ python -m aiter.aot.flydsl.mega_moe --experts-per-rank 48
The FP8 flash-attention kernel picks its tile, rescale threshold and split-K
factor per call, so a serving process JIT-compiles each new combination on the
request path. `fmha_fp8` enumerates every combination the wrapper's heuristics
can reach for each head shape listed in `AITER_CONFIG_FMHA_FP8_AOT_FILE`: the
header-only `configs/fmha_fp8_aot.csv` merged with
`configs/model_configs/*_fmha_fp8_aot.csv`. Add a model by adding a row
(`cu_num,num_heads,num_kv_heads,head_dim,head_dim_v,layout,model`, heads per rank,
layout `varlen_cross` for the chunked-prefill serving path).
can reach for each head shape in its `DEFAULT_SHAPES` table. Add a model by
adding an entry (`num_heads, num_kv_heads, head_dim, head_dim_v, layout`, heads
per rank, layout `varlen_cross` for the chunked-prefill serving path).

```bash
# every configured shape (Kimi-K3 at TP8: 12:12:192:128 -> 92 kernels)
python -m aiter.aot.flydsl.fmha_fp8
# ad hoc shape, without touching the CSVs
# ad hoc shape, without editing DEFAULT_SHAPES
python -m aiter.aot.flydsl.fmha_fp8 --shape 12:12:192:128
```

Expand Down
6 changes: 4 additions & 2 deletions aiter/aot/flydsl/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,6 +142,10 @@ def _collect_aot_jobs_for(kind: OpKind) -> list[dict[str, Any]]:
if kind is OpKind.MEGA_MOE:
from .mega_moe import default_jobs

return default_jobs()
if kind is OpKind.FMHA_FP8:
from .fmha_fp8 import default_jobs

return default_jobs()
if kind is OpKind.MOE:
from .moe import DEFAULT_CSVS, parse_csv
Expand All @@ -155,8 +159,6 @@ def _collect_aot_jobs_for(kind: OpKind) -> list[dict[str, Any]]:
from .grouped_moe import DEFAULT_CSVS, parse_csv
elif kind is OpKind.CHUNK_GDN_H:
from .chunk_gdn_h import DEFAULT_CSVS, parse_csv
elif kind is OpKind.FMHA_FP8:
from .fmha_fp8 import DEFAULT_CSVS, parse_csv
else:
raise ValueError(f"unknown FlyDSL AOT kind: {kind!r}")
return collect_aot_jobs(DEFAULT_CSVS, parse_csv)
Expand Down
121 changes: 37 additions & 84 deletions aiter/aot/flydsl/fmha_fp8.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,59 +5,46 @@

"""AOT pre-compilation for the gfx950 FlyDSL FP8 flash-attention forward.

``flydsl_flash_attn_fp8_func`` picks its compile-time configuration per call:
the rescale threshold from the KV length, BLOCK_M from the KV tile count and
occupancy, and the split-K factor from a makespan model. Every new combination
is a JIT compile on the serving path, where it stalls a request for seconds.

The launcher's cache key is fixed by the ``_build_fp8`` parameters alone:
runtime ints/floats (batch, sequence lengths, strides, scale) enter it by type
only and the flattened tensors by dtype only. So the full set a model can hit is
finite and is enumerated here from the heuristics' own module constants, which
keeps the job list in step with them.

Head shapes come from ``AITER_CONFIGS.AITER_CONFIG_FMHA_FP8_AOT_FILE``: the
header-only ``aiter/configs/fmha_fp8_aot.csv`` merged with every
``model_configs/*_fmha_fp8_aot.csv`` (columns ``cu_num, num_heads,
num_kv_heads, head_dim, head_dim_v, layout, model``), so ``setup.py`` builds
them with the other FlyDSL AOT kinds. ``model`` is informational, but it is
part of the merge key: two models sharing a head shape stay distinct rows
(the merge rejects exact duplicates) and their jobs dedupe by kernel name.

python -m aiter.aot.flydsl.fmha_fp8 # config CSVs
``flydsl_flash_attn_fp8_func`` picks its rescale threshold, BLOCK_M and split-K
factor per call, so each new combination is a JIT compile on the serving path.
The launcher's cache key is fixed by the ``_build_fp8`` parameters alone -- the
runtime ints/floats enter it by type only and the flattened tensors by dtype
only -- so the set a head shape can reach is finite. It is enumerated here from
the heuristics' own module constants, which keeps the job list in step with
them.

Usage:
python -m aiter.aot.flydsl.fmha_fp8 # DEFAULT_SHAPES
python -m aiter.aot.flydsl.fmha_fp8 --shape 12:12:192:128 # ad hoc

``layout`` is ``varlen_cross`` (packed Q/KV with independent lengths, the
chunked-prefill serving path), ``varlen``, ``dense`` or ``dense_cross``.
"""

from __future__ import annotations

import argparse
import csv
import os
import sys
import time

from aiter.aot.flydsl.common import (
_CU_NUM_TO_ARCH,
collect_aot_jobs,
compile_only_env,
cu_num_to_arch,
dedupe_jobs,
override_env,
run_jobs_parallel,
)
from aiter.jit.core import AITER_CONFIGS

DEFAULT_CSVS = [AITER_CONFIGS.AITER_CONFIG_FMHA_FP8_AOT_FILE]
AOT_ARCH = "gfx950"
LAYOUTS = {
"varlen_cross": (True, True),
"varlen": (True, False),
"dense": (False, False),
"dense_cross": (False, True),
}
# Head shapes to pre-compile, per model: (num_heads, num_kv_heads, head_dim,
# head_dim_v, layout) with the head counts per rank. ``varlen_cross`` is the
# chunked-prefill serving path. Jobs dedupe by kernel name, so two models may
# share a head shape.
DEFAULT_SHAPES = {
"kimi_k3_tp8": [(12, 12, 192, 128, "varlen_cross")],
}
# Flags the public wrapper defaults to; they are part of the compile key.
_DEFAULT_FLAGS = {
"daz": True,
Expand All @@ -67,11 +54,8 @@
}


def _layout(value) -> str:
# get_config_file fills a column a model file lacks with 0, and pandas writes
# a blank cell back as "" or "nan": all of those mean "not given".
layout = str(value if value is not None else "").strip()
layout = "varlen_cross" if layout in ("", "0", "nan") else layout
def _layout(value: str | None) -> str:
layout = str(value or "").strip() or "varlen_cross"
if layout not in LAYOUTS:
raise ValueError(
f"unknown layout {layout!r}; expected one of {sorted(LAYOUTS)}"
Expand All @@ -88,8 +72,8 @@ def parse_shape(spec: str) -> tuple[int, int, int, int, str]:

def unsupported_reason(h: int, hkv: int, d: int, dv: int) -> str | None:
"""Why the kernel cannot serve this head shape, or None. The runtime gate
(flydsl_flash_attn_fp8_supported) uses the same check, so a rejected shape
is never dispatched and its kernels would only fail the build."""
uses the same check, so a rejected shape is never dispatched and its
kernels would only fail the build."""
from aiter.ops.flydsl.kernels.flash_attn_func_fp8_gfx950 import (
_fp8_config_reason,
)
Expand Down Expand Up @@ -174,45 +158,24 @@ def jobs_for_shape(h: int, hkv: int, d: int, dv: int, layout: str) -> list[dict]
return jobs


def parse_csv(csv_path: str) -> list[dict]:
"""One row per head shape -> every variant that shape can reach.
def default_jobs(shapes: dict[str, list] = DEFAULT_SHAPES) -> list[dict]:
"""Every variant ``DEFAULT_SHAPES`` can reach.

A row that is malformed, not gfx950, or a shape the kernel rejects is
skipped with a warning rather than failing the whole build.
A shape the kernel's own gate rejects is skipped with a warning rather
than failing the whole build.
"""
jobs = []
with open(csv_path, "r", encoding="utf-8", newline="") as f:
rows = csv.DictReader(line for line in f if not line.lstrip().startswith("#"))
for row in rows:
try:
cu_num = int(float(row["cu_num"]))
h, hkv, d, dv = (
int(float(row[k]))
for k in ("num_heads", "num_kv_heads", "head_dim", "head_dim_v")
)
layout = _layout(row.get("layout"))
except (KeyError, TypeError, ValueError) as error:
print(f" [WARN] {csv_path}: malformed row {dict(row)}: {error}")
continue
if cu_num not in _CU_NUM_TO_ARCH:
print(
f" [WARN] {csv_path}: unknown cu_num={cu_num}, assuming {AOT_ARCH}"
)
arch = cu_num_to_arch(cu_num)
if arch != AOT_ARCH:
print(
f" [WARN] {csv_path}: cu_num={cu_num} is {arch}; "
f"the FP8 FMHA kernel is {AOT_ARCH}-only, skipping"
)
continue
for model, entries in shapes.items():
for h, hkv, d, dv, layout in entries:
reason = unsupported_reason(h, hkv, d, dv)
if reason:
print(
f" [WARN] {csv_path}: {h}:{hkv}:{d}:{dv} unsupported ({reason}), skipping"
f" [WARN] {model}: {h}:{hkv}:{d}:{dv} unsupported "
f"({reason}), skipping"
)
continue
jobs.extend(jobs_for_shape(h, hkv, d, dv, layout))
return jobs
jobs.extend(jobs_for_shape(h, hkv, d, dv, _layout(layout)))
return dedupe_jobs(jobs)


def _compile_to_cache(
Expand Down Expand Up @@ -250,10 +213,10 @@ def _compile_to_cache(
**_DEFAULT_FLAGS,
)

# Mirror the runtime call exactly (flash_attn_func_fp8_gfx950.py: the exe(...)
# at the end of flydsl_flash_attn_fp8_func): flat tensors, the same optional
# kwargs present or absent, the same dtypes. Values and sizes never reach
# the key; a missing or extra kwarg, or a wrong dtype, silently would.
# Mirror the runtime exe(...) at the end of flydsl_flash_attn_fp8_func: flat
# tensors, the same optional kwargs present or absent, the same dtypes.
# Values and sizes never reach the key; a missing or extra kwarg, or a wrong
# dtype, silently would.
cpu = torch.device("cpu")

def t(n, dtype):
Expand Down Expand Up @@ -310,17 +273,11 @@ def compile_one_config(**job) -> dict:

def main():
parser = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
parser.add_argument(
"--csv",
nargs="+",
default=DEFAULT_CSVS,
help="head-shape CSV(s) (default: AITER_CONFIG_FMHA_FP8_AOT_FILE)",
)
parser.add_argument(
"--shape",
action="append",
default=[],
help="H:Hkv:D:Dv[@layout], repeatable; replaces --csv",
help="H:Hkv:D:Dv[@layout], repeatable; replaces DEFAULT_SHAPES",
)
parser.add_argument("--list", action="store_true", help="print the jobs and exit")
args = parser.parse_args()
Expand All @@ -335,11 +292,7 @@ def main():
jobs.extend(jobs_for_shape(h, hkv, d, dv, layout))
jobs = dedupe_jobs(jobs)
else:
for csv_path in args.csv:
if not os.path.isfile(csv_path):
print(f"Error: CSV file not found: {csv_path}")
sys.exit(1)
jobs = collect_aot_jobs(args.csv, parse_csv)
jobs = default_jobs()
if args.list:
for job in jobs:
print(job["kernel_name"])
Expand Down
1 change: 0 additions & 1 deletion aiter/configs/fmha_fp8_aot.csv

This file was deleted.

2 changes: 0 additions & 2 deletions aiter/configs/model_configs/kimik3_fmha_fp8_aot.csv

This file was deleted.

17 changes: 0 additions & 17 deletions aiter/jit/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,15 +223,6 @@ def mp_lock(
f"{AITER_ROOT_DIR}/aiter/configs/chunk_gdn_h_opt_tuned.csv",
)

# Head shapes to AOT-compile the gfx950 FlyDSL FP8 flash-attention forward for
# (aiter/aot/flydsl/fmha_fp8.py). Per-model rows live under model_configs/
# (*_fmha_fp8_aot.csv) and get merged into this canonical file, which ships
# header-only like chunk_gdn_h_opt_tuned.csv.
AITER_CONFIG_FMHA_FP8_AOT = os.getenv(
"AITER_CONFIG_FMHA_FP8_AOT",
f"{AITER_ROOT_DIR}/aiter/configs/fmha_fp8_aot.csv",
)

AITER_CONFIG_DISPATCH_COMBINE_INTRANODE = os.getenv(
"AITER_CONFIG_DISPATCH_COMBINE_INTRANODE",
f"{AITER_ROOT_DIR}/aiter/configs/tuned_dispatch_combine_intranode.csv",
Expand Down Expand Up @@ -381,14 +372,6 @@ def AITER_CONFIG_GDN_K5_OPT_FILE(self):
"chunk_gdn_h_opt_tuned",
)

@property
def AITER_CONFIG_FMHA_FP8_AOT_FILE(self):
return self.get_config_file(
"AITER_CONFIG_FMHA_FP8_AOT",
AITER_CONFIG_FMHA_FP8_AOT,
"fmha_fp8_aot",
)

@property
def AITER_CONFIG_DISPATCH_COMBINE_INTRANODE_FILE(self):
return self.get_config_file(
Expand Down
Loading