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
7 changes: 3 additions & 4 deletions benchmarks/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,9 @@
import benchmarks.baselines # noqa: F401
from benchmarks.report import BenchmarkReport, _bench_results

# What a row carries besides its measurements.
_NOT_A_MEASUREMENT = frozenset({"tag", "op", "op_module", "ops", "params", "run_config", "result"})


def pytest_make_parametrize_id(config, val, argname):
"""Render the values pytest would otherwise collect as `shape0`, `dtype0`.
Expand All @@ -30,10 +33,6 @@ def pytest_make_parametrize_id(config, val, argname):
return None


# What a row carries besides its measurements.
_NOT_A_MEASUREMENT = frozenset({"tag", "op", "op_module", "ops", "params", "run_config", "result"})


def _prop(value) -> str:
"""Format one measurement for the XML.

Expand Down
6 changes: 3 additions & 3 deletions benchmarks/hardware/memory/hbm_bandwidth.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,9 @@
_CU_SRC = Path(__file__).parent / "hbm_saturation.cu"


_MIXES = ("copy", "triad", "read", "write")


def _compile(cu_path, binary_path, arch="sm_90"):
"""Compile the CUDA source. Raises on failure."""
cmd = [
Expand Down Expand Up @@ -57,9 +60,6 @@ def _run(binary_path, size_mb, theo_peak_gbs):
return result.stdout.strip().splitlines()


_MIXES = ("copy", "triad", "read", "write")


def _parse_peaks(lines):
"""Best bandwidth (GB/s) per access mix from the CSV output.

Expand Down
35 changes: 16 additions & 19 deletions benchmarks/ops/bench_fused_gated.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,22 @@
)
from workloads.workload_base import FixtureBase

# Scenario -> (tokens, width).
_STRATEGY_SHAPES = {
"llama-hidden-1k-tokens": (1024, 4096),
"llama-7b-ffn-1k-tokens": (1024, 11008),
"llama-hidden-4k-tokens": (4096, 4096),
}
_STRATEGY_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
_STRATEGY_KERNELS = [
("silu_and_mul", SiluAndMulFwdKernel),
("gelu_and_mul", GeluAndMulFwdKernel),
("gelu_tanh_and_mul", GeluTanhAndMulFwdKernel),
]
# How far behind the fastest strategy the default may sit before the choice is
# stale. Wide enough to clear run-to-run spread, narrow enough to flag a flip.
_STRATEGY_MARGIN = 1.25


class FusedGatedBenchmark(BenchmarkBase[FusedGatedBenchCase]):
"""Times the strategy decision; it records no row, so both metrics are ``None``."""
Expand Down Expand Up @@ -82,20 +98,6 @@ def test_gelu_tanh_and_mul_bench(call) -> None:
_profile_fused_gated(GeluTanhAndMulFwdOp, call, "gelu_tanh_and_mul")


# Scenario -> (tokens, width).
_STRATEGY_SHAPES = {
"llama-hidden-1k-tokens": (1024, 4096),
"llama-7b-ffn-1k-tokens": (1024, 11008),
"llama-hidden-4k-tokens": (4096, 4096),
}
_STRATEGY_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
_STRATEGY_KERNELS = [
("silu_and_mul", SiluAndMulFwdKernel),
("gelu_and_mul", GeluAndMulFwdKernel),
("gelu_tanh_and_mul", GeluTanhAndMulFwdKernel),
]


def _strategy_params():
"""Default-strategy sentinel: shape and dtype axes on the first kernel, plus
one reference-point direct-vs-explicit sentinel per remaining kernel.
Expand Down Expand Up @@ -127,11 +129,6 @@ class FusedGatedStrategyBenchFixture(FixtureBase):
PARAMS = [("op_name, M, N, dtype, kernel_cls", _strategy_params())]


# How far behind the fastest strategy the default may sit before the choice is
# stale. Wide enough to clear run-to-run spread, narrow enough to flag a flip.
_STRATEGY_MARGIN = 1.25


@FusedGatedStrategyBenchFixture
def test_fused_gated_default_strategy_is_the_fast_one(
op_name: str,
Expand Down
29 changes: 13 additions & 16 deletions benchmarks/ops/bench_pool.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,19 @@
from workloads.pool import MeanPoolingCallWorkload, MeanPoolingWorkload
from workloads.workload_base import CallWorkload

# Which library serves an op, and the pooling kind and rank its adapter needs. An op absent
# here has none: no library covers 1D, adaptive pooling, or 3D max-pool indices. Every row
# is also timed against torch, eager and compiled, so this table is not the whole baseline.
_BASELINE: dict[str, tuple[str, str, int]] = {
"AvgPool2dFwdOp": (FLAGGEMS_TAG, "avg", 2),
"AvgPool3dFwdOp": ("cudnn", "avg", 3),
"MaxPool2dFwdOp": (FLAGGEMS_TAG, "max", 2),
"MaxPool2dIndicesFwdOp": (FLAGGEMS_TAG, "max", 2),
"MaxPool3dFwdOp": ("cudnn", "max", 3),
}
# Autotuning is a bench-run policy; manifest workloads do not carry it.
_TUNE = True


def flaggems_pool_fn(
kind: str,
Expand Down Expand Up @@ -176,18 +189,6 @@ def run_max(x: torch.Tensor):
return None


# Which library serves an op, and the pooling kind and rank its adapter needs. An op absent
# here has none: no library covers 1D, adaptive pooling, or 3D max-pool indices. Every row
# is also timed against torch, eager and compiled, so this table is not the whole baseline.
_BASELINE: dict[str, tuple[str, str, int]] = {
"AvgPool2dFwdOp": (FLAGGEMS_TAG, "avg", 2),
"AvgPool3dFwdOp": ("cudnn", "avg", 3),
"MaxPool2dFwdOp": (FLAGGEMS_TAG, "max", 2),
"MaxPool2dIndicesFwdOp": (FLAGGEMS_TAG, "max", 2),
"MaxPool3dFwdOp": ("cudnn", "max", 3),
}


def _as_tuple(value, ndim: int) -> tuple:
if isinstance(value, (tuple, list)):
return tuple(value)
Expand Down Expand Up @@ -342,10 +343,6 @@ def test_adaptive_max_pool2d_indices_bench(call) -> None:
# MeanPoolingFwdOp, the chunked sequence mean.


# Autotuning is a bench-run policy; manifest workloads do not carry it.
_TUNE = True


def _torch_view_mean(workload: MeanPoolingWorkload):
"""The same mean over a reshaped view, or None where the chunks are ragged.

Expand Down
20 changes: 10 additions & 10 deletions benchmarks/tests/test_benchmark_boundaries.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,16 @@
BENCHMARK_DIRS = ("benchmarks/ops",)


# A benchmark takes (flops, bytes) from its op — docs/design/roofline.md §4.2. An entry
# here declares the two methods for a reason the name below states. An entry whose
# subject is an op goes as soon as that op gains a manifest entry; an entry whose
# subject is not an op stays, because a manifest entry is something only an op can have.
_ROOFLINE_OF_ITS_OWN = {
"FusedGatedBenchmark": "times a forced kernel strategy, which no op can request and "
"no report has a row for; both metrics return None",
}


def _benchmark_files() -> list[Path]:
return [
path
Expand Down Expand Up @@ -72,16 +82,6 @@ def test_benchmarks_do_not_author_gen_inputs() -> None:
assert _scan(_defines_gen_inputs) == {}


# A benchmark takes (flops, bytes) from its op — docs/design/roofline.md §4.2. An entry
# here declares the two methods for a reason the name below states. An entry whose
# subject is an op goes as soon as that op gains a manifest entry; an entry whose
# subject is not an op stays, because a manifest entry is something only an op can have.
_ROOFLINE_OF_ITS_OWN = {
"FusedGatedBenchmark": "times a forced kernel strategy, which no op can request and "
"no report has a row for; both metrics return None",
}


def _writes_its_own_roofline(tree: ast.AST) -> list[str]:
return [
f"{node.name}.{fn.name} (line {fn.lineno})"
Expand Down
36 changes: 16 additions & 20 deletions scripts/lint/tilelang_idioms_lint.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,22 @@
_DTYPE_NAME = re.compile(r"^(u?int[0-9]+|b?float[0-9]+|float8[a-z0-9_]*|bool|handle)$")


_NONSCALAR_KINDS = {
ast.List: "list",
ast.ListComp: "list",
ast.Dict: "dict",
ast.DictComp: "dict",
ast.Set: "set",
ast.SetComp: "set",
ast.Tuple: "tuple",
ast.GeneratorExp: "generator",
ast.Lambda: "function",
}
_SCOPE = (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)
_Func = ast.FunctionDef | ast.AsyncFunctionDef
_SCALAR_ANNOTATIONS = frozenset({"int", "float", "str", "bool", "None", "NoneType"})


def _attr_path(node: ast.AST) -> str | None:
"""Dotted name of an attribute chain, e.g. ``T.reinterpret``; None otherwise."""
parts = []
Expand Down Expand Up @@ -150,23 +166,6 @@ def _arg(call: ast.Call, pos: int, name: str) -> ast.AST | None:
return next((k.value for k in call.keywords if k.arg == name), None)


_NONSCALAR_KINDS = {
ast.List: "list",
ast.ListComp: "list",
ast.Dict: "dict",
ast.DictComp: "dict",
ast.Set: "set",
ast.SetComp: "set",
ast.Tuple: "tuple",
ast.GeneratorExp: "generator",
ast.Lambda: "function",
}

_SCOPE = (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)

_Func = ast.FunctionDef | ast.AsyncFunctionDef


def _jit_names(tree: ast.Module) -> tuple[set[str], set[str]]:
"""What this file binds to ``tilelang.jit``: module aliases, then bare names.

Expand Down Expand Up @@ -239,9 +238,6 @@ def _function_tables(top: symtable.SymbolTable) -> dict[tuple[str, int], symtabl
return tables


_SCALAR_ANNOTATIONS = frozenset({"int", "float", "str", "bool", "None", "NoneType"})


def _string_annotation_kind(annotation: str) -> str | None:
"""Classify a quoted annotation."""
try:
Expand Down
38 changes: 16 additions & 22 deletions scripts/nightly_report.py
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,22 @@
# ---------------------------------------------------------------------------


# A dtype name as a case id spells it: the id ends in its dtype cases and dtype parameters.
_DTYPE_TOKEN = re.compile(r"bfloat16|bool|u?int\d+|float\d+(?:_[a-z0-9]+)?|complex\d+")
# Verdicts are drawn on device execution time, not on the span that also covers the
# gaps between a call's kernels.
_CONCLUSION_KEY = "device_busy_ms"
#: Absorbs the rounding a count recovered from ``tflops`` carries.
_DERIVED_COUNT_RTOL = 1e-3
# A kernel file below this share of executed lines was never constructed by any
# test. The lowest genuinely-built kernel file measures 36.9%, so the threshold
# has room before it starts catching built kernels.
_KERNEL_BUILT_PCT = 25
# Below this, one statement swings the percentage too far to read.
_COVERAGE_MIN_STMTS = 20
_COVERAGE_WORST_N = 15 # rows in the least-covered file list


def _get_properties(testcase: ET.Element) -> dict[str, str]:
"""Extract user properties from a JUnit testcase element."""
props = {}
Expand Down Expand Up @@ -301,10 +317,6 @@ def _case_id(name: str) -> str:
return name


# A dtype name as a case id spells it: the id ends in its dtype cases and dtype parameters.
_DTYPE_TOKEN = re.compile(r"bfloat16|bool|u?int\d+|float\d+(?:_[a-z0-9]+)?|complex\d+")


def _dtypes_of(name: str) -> tuple[str, ...]:
"""The dtype names in a row's case id, in order."""
return tuple(t for t in _case_id(name).split("-") if _DTYPE_TOKEN.fullmatch(t))
Expand Down Expand Up @@ -353,11 +365,6 @@ def history_window(runs: list[dict], retention_days: int = HISTORY_RETENTION_DAY
return sorted(kept.values(), key=lambda r: r["date"])


# Verdicts are drawn on device execution time, not on the span that also covers the
# gaps between a call's kernels.
_CONCLUSION_KEY = "device_busy_ms"


def _conclusion(cfg: dict) -> tuple[float | None, str]:
"""The reading a verdict is drawn on, and which key it came from.

Expand All @@ -376,10 +383,6 @@ def _conclusion_ms(cfg: dict) -> float | None:
return _conclusion(cfg)[0]


#: Absorbs the rounding a count recovered from ``tflops`` carries.
_DERIVED_COUNT_RTOL = 1e-3


class _WorkCounts(NamedTuple):
flops: float | None
nbytes: float | None
Expand Down Expand Up @@ -1131,15 +1134,6 @@ def _pct(hit: int, total: int) -> str:
return f"{100 * hit / total:.1f}%" if total else "-"


# A kernel file below this share of executed lines was never constructed by any
# test. The lowest genuinely-built kernel file measures 36.9%, so the threshold
# has room before it starts catching built kernels.
_KERNEL_BUILT_PCT = 25
# Below this, one statement swings the percentage too far to read.
_COVERAGE_MIN_STMTS = 20
_COVERAGE_WORST_N = 15 # rows in the least-covered file list


def _coverage_signals(files: list[dict]) -> dict:
"""Reduce per-file coverage to the three numbers worth acting on.

Expand Down
38 changes: 17 additions & 21 deletions scripts/validate_manifest.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,23 @@
_STAGE_KEYS = {"name", "op", "kernel", "optional"}


_ENTRY_KEYS = {
"family": str,
"status": str,
"signature": dict,
"workloads": list,
"roofline": dict,
"ref_api": str,
"composition": dict,
}
_REQUIRED = ("family", "status", "signature", "workloads", "roofline")
_OP_KEY = re.compile(r"[A-Z][A-Za-z0-9]*(Fwd|Bwd)Op")
# Execution-policy parameters every op takes, in order with their defaults, and the reserved one
# it may take (docs/design/manifest.md § Signature).
_POLICY_PARAMETERS = {"target": None, "kernel_map": None, "tune": False}
_RESERVED_POLICY = "config"


def _key_format_errors(
op_name: str,
all_op_names: Collection[str],
Expand Down Expand Up @@ -232,18 +249,6 @@ def _check_bench_files(repo_root: Path) -> list[str]:
return errors


_ENTRY_KEYS = {
"family": str,
"status": str,
"signature": dict,
"workloads": list,
"roofline": dict,
"ref_api": str,
"composition": dict,
}
_REQUIRED = ("family", "status", "signature", "workloads", "roofline")


def _schema_errors(op_name: str, entry: dict, all_op_names) -> list[str]:
"""Top-level fields of an entry (docs/design/manifest.md § Top-Level Fields)."""
if not isinstance(op_name, str):
Expand Down Expand Up @@ -273,9 +278,6 @@ def _schema_errors(op_name: str, entry: dict, all_op_names) -> list[str]:
return errors


_OP_KEY = re.compile(r"[A-Z][A-Za-z0-9]*(Fwd|Bwd)Op")


def _family_errors(op_name: str, entry: dict) -> list[str]:
"""`family` names a public module; an implemented op is exported from it by its key."""
family = entry.get("family")
Expand Down Expand Up @@ -317,12 +319,6 @@ def _ref_api_errors(op_name: str, ref: str) -> list[str]:
return [f"[schema] {op_name}: ref_api {ref!r}: no prefix of it is an importable module"]


# Execution-policy parameters every op takes, in order with their defaults, and the reserved one
# it may take (docs/design/manifest.md § Signature).
_POLICY_PARAMETERS = {"target": None, "kernel_map": None, "tune": False}
_RESERVED_POLICY = "config"


def _normal_default(value):
"""A default as the manifest writes it: a dtype by name, a tuple as a list."""
if isinstance(value, tuple):
Expand Down
Loading
Loading