diff --git a/benchmarks/conftest.py b/benchmarks/conftest.py index d55b0a0c8..d3c1044d1 100644 --- a/benchmarks/conftest.py +++ b/benchmarks/conftest.py @@ -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`. @@ -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. diff --git a/benchmarks/hardware/memory/hbm_bandwidth.py b/benchmarks/hardware/memory/hbm_bandwidth.py index 8088a69cc..bc64de467 100644 --- a/benchmarks/hardware/memory/hbm_bandwidth.py +++ b/benchmarks/hardware/memory/hbm_bandwidth.py @@ -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 = [ @@ -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. diff --git a/benchmarks/ops/bench_fused_gated.py b/benchmarks/ops/bench_fused_gated.py index b144e2a0f..e05c66654 100644 --- a/benchmarks/ops/bench_fused_gated.py +++ b/benchmarks/ops/bench_fused_gated.py @@ -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``.""" @@ -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. @@ -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, diff --git a/benchmarks/ops/bench_pool.py b/benchmarks/ops/bench_pool.py index f41677cfd..827adbbda 100644 --- a/benchmarks/ops/bench_pool.py +++ b/benchmarks/ops/bench_pool.py @@ -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, @@ -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) @@ -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. diff --git a/benchmarks/tests/test_benchmark_boundaries.py b/benchmarks/tests/test_benchmark_boundaries.py index 3e1b0e7ff..5a54f21fb 100644 --- a/benchmarks/tests/test_benchmark_boundaries.py +++ b/benchmarks/tests/test_benchmark_boundaries.py @@ -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 @@ -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})" diff --git a/scripts/lint/tilelang_idioms_lint.py b/scripts/lint/tilelang_idioms_lint.py index 907057bae..99155f50d 100755 --- a/scripts/lint/tilelang_idioms_lint.py +++ b/scripts/lint/tilelang_idioms_lint.py @@ -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 = [] @@ -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. @@ -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: diff --git a/scripts/nightly_report.py b/scripts/nightly_report.py index 5f21e2a26..77da964ad 100755 --- a/scripts/nightly_report.py +++ b/scripts/nightly_report.py @@ -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 = {} @@ -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)) @@ -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. @@ -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 @@ -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. diff --git a/scripts/validate_manifest.py b/scripts/validate_manifest.py index a2922fe9a..c31adaac1 100755 --- a/scripts/validate_manifest.py +++ b/scripts/validate_manifest.py @@ -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], @@ -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): @@ -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") @@ -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): diff --git a/scripts/validate_roofline_bytes.py b/scripts/validate_roofline_bytes.py index 8fac541a0..84fa84b0a 100755 --- a/scripts/validate_roofline_bytes.py +++ b/scripts/validate_roofline_bytes.py @@ -48,12 +48,6 @@ COLD_CACHE_PREMISE = "cold-cache replay (ncu --cache-control all)" -def _op_class(op_name: str, entry: dict): - from tileops.manifest.registry import op_class - - return op_class(op_name, entry) - - # The calls whose read half is not a lower bound: op name -> (condition over the call's ``ix``, # reason). READ_BOUND_EXCEPTIONS: dict = { @@ -68,6 +62,12 @@ def _op_class(op_name: str, entry: dict): } +def _op_class(op_name: str, entry: dict): + from tileops.manifest.registry import op_class + + return op_class(op_name, entry) + + def _call(op_name: str, entry: dict, row: dict, case: dict): """One row and dtype case of an entry, instantiated.""" from tileops.manifest import load_adts diff --git a/src/tileops/__init__.py b/src/tileops/__init__.py index dcecc266d..182c1a350 100644 --- a/src/tileops/__init__.py +++ b/src/tileops/__init__.py @@ -42,6 +42,7 @@ quantization, reduction, rope, + sampling, sequence_modeling, ) from .ops.op_base import Op @@ -57,6 +58,7 @@ "convolution", "fft", "moe", + "sampling", "rope", "attention", "linear_attention", diff --git a/src/tileops/backend/protocol.py b/src/tileops/backend/protocol.py index f3dbb71c1..cb63c0fa8 100644 --- a/src/tileops/backend/protocol.py +++ b/src/tileops/backend/protocol.py @@ -6,36 +6,33 @@ import torch - -class TensorSpec(NamedTuple): - """What one tensor is, without the tensor. Handed to ``build_kernel``.""" - - device: torch.device - dtype: torch.dtype - shape: tuple[int, ...] - - @staticmethod - def of(tensor: torch.Tensor) -> "TensorSpec": - """Describe *tensor*.""" - return TensorSpec(tensor.device, tensor.dtype, tuple(tensor.shape)) - - # One call's result. A purely mutating op returns ``None``: ``torch.library.custom_op`` # cannot express a return value aliasing an input. KernelResult = Union[torch.Tensor, tuple[torch.Tensor, ...], None] - # Called ``build_kernel(*inputs, **params)``: a `TensorSpec` per input in # ``signature.inputs`` order — ``None`` for an ``optional: true`` input the call did not # pass, so presence is read off the slot rather than off how many slots there are — then # ``signature.params`` by keyword. Both lists are per-op, which the type system cannot # express, hence ``...``. BuildKernel = Callable[..., Callable[..., KernelResult]] - # "Is this the kind of device my kernels are written for" — ``False``, not an exception, # for the rest. Per-call support is ``build_kernel``'s answer; it sees the dtypes too. DetectFn = Callable[[torch.device], bool] +class TensorSpec(NamedTuple): + """What one tensor is, without the tensor. Handed to ``build_kernel``.""" + + device: torch.device + dtype: torch.dtype + shape: tuple[int, ...] + + @staticmethod + def of(tensor: torch.Tensor) -> "TensorSpec": + """Describe *tensor*.""" + return TensorSpec(tensor.device, tensor.dtype, tuple(tensor.shape)) + + class _Builtin: """The type of :data:`BUILTIN`. One instance, compared by identity.""" diff --git a/src/tileops/kernels/attention/call_spec.py b/src/tileops/kernels/attention/call_spec.py index d49eec102..ec0d61e1e 100644 --- a/src/tileops/kernels/attention/call_spec.py +++ b/src/tileops/kernels/attention/call_spec.py @@ -35,6 +35,15 @@ WS_ARCH = 90 +# Tile heights the warp-specialized paged decode kernel can pick from. A tile +# divides the page size, so one tile never straddles two pages, and it splits +# evenly across the four consumer warps. +_WS_DECODE_TILES = (16, 32, 64, 128) +# Head dims that map onto one warp: the score reduction is a shuffle chain over +# 32 lanes, so a lane owns ``dim / 32`` elements of the head vector. +_WS_DECODE_LANES = 32 + + def fp8_dtype() -> Optional[torch.dtype]: """Return ``torch.float8_e4m3fn`` when the torch build carries it.""" return getattr(torch, "float8_e4m3fn", None) @@ -81,15 +90,6 @@ def uses_sliding_window(call: AttentionCall) -> bool: return call.window_size_left != -1 or call.window_size_right != -1 -# Tile heights the warp-specialized paged decode kernel can pick from. A tile -# divides the page size, so one tile never straddles two pages, and it splits -# evenly across the four consumer warps. -_WS_DECODE_TILES = (16, 32, 64, 128) -# Head dims that map onto one warp: the score reduction is a shuffle chain over -# 32 lanes, so a lane owns ``dim / 32`` elements of the head vector. -_WS_DECODE_LANES = 32 - - def paged_decode_ws_region(call: AttentionCall) -> bool: """The paged-decode region the warp-specialized MHA kernel serves. diff --git a/src/tileops/kernels/attention/gqa_dense.py b/src/tileops/kernels/attention/gqa_dense.py index 3ac31a980..7c1359ade 100644 --- a/src/tileops/kernels/attention/gqa_dense.py +++ b/src/tileops/kernels/attention/gqa_dense.py @@ -41,6 +41,31 @@ ] +# Causal warp-specialized Dense attention. +BLOCK_M = 128 +BLOCK_N = 128 +NSK = 2 +NSV = 2 +THREADS = 384 +NMMA = 256 +_pc = { + tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True, + tilelang.PassConfigKey.TL_DISABLE_THREAD_STORAGE_SYNC: True, +} +_cf = [ + "-O3", + "--use_fast_math", + "-Wno-deprecated-declarations", + "-U__CUDA_NO_HALF_OPERATORS__", + "-U__CUDA_NO_HALF_CONVERSIONS__", + "-U__CUDA_NO_HALF2_OPERATORS__", + "-U__CUDA_NO_BFLOAT16_CONVERSIONS__", + "--expt-relaxed-constexpr", + "--expt-extended-lambda", + "-DNDEBUG", +] + + @functools.lru_cache(maxsize=32) @tilelang.jit(out_idx=[4, 5], pass_configs=_PASS_CONFIGS, compile_flags=_COMPILE_FLAGS) def _gqa_dense_rope_qk_kernel( @@ -201,32 +226,6 @@ def make_dense_qk_rope_preprocessor( ) -# Causal warp-specialized Dense attention. -BLOCK_M = 128 -BLOCK_N = 128 -NSK = 2 -NSV = 2 -THREADS = 384 -NMMA = 256 - -_pc = { - tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True, - tilelang.PassConfigKey.TL_DISABLE_THREAD_STORAGE_SYNC: True, -} -_cf = [ - "-O3", - "--use_fast_math", - "-Wno-deprecated-declarations", - "-U__CUDA_NO_HALF_OPERATORS__", - "-U__CUDA_NO_HALF_CONVERSIONS__", - "-U__CUDA_NO_HALF2_OPERATORS__", - "-U__CUDA_NO_BFLOAT16_CONVERSIONS__", - "--expt-relaxed-constexpr", - "--expt-extended-lambda", - "-DNDEBUG", -] - - @functools.lru_cache(maxsize=32) @tilelang.jit(out_idx=[3], pass_configs=_pc, compile_flags=_cf) def _gqa_dense_ws_kernel( diff --git a/src/tileops/kernels/attention/gqa_fwd.py b/src/tileops/kernels/attention/gqa_fwd.py index 42895b501..3be46c23a 100644 --- a/src/tileops/kernels/attention/gqa_fwd.py +++ b/src/tileops/kernels/attention/gqa_fwd.py @@ -25,14 +25,6 @@ ] -def _tile_stage_thread_configs() -> list[dict]: - """The default GQA search space: block_m x block_n x num_stages x threads.""" - return [ - {"block_m": bm, "block_n": bn, "num_stages": ns, "threads": th} - for bm, bn, ns, th in itertools.product((32, 64, 128), (32, 64, 128), (1, 2, 3), (128, 256)) - ] - - _FAST_COMPILE_FLAGS = [ "-O3", "--use_fast_math", @@ -47,6 +39,14 @@ def _tile_stage_thread_configs() -> list[dict]: ] +def _tile_stage_thread_configs() -> list[dict]: + """The default GQA search space: block_m x block_n x num_stages x threads.""" + return [ + {"block_m": bm, "block_n": bn, "num_stages": ns, "threads": th} + for bm, bn, ns, th in itertools.product((32, 64, 128), (32, 64, 128), (1, 2, 3), (128, 256)) + ] + + def _make_apply_softcap_no_mask_guard(score_scale, softcap, accum_dtype, block_rows, block_cols): @T.macro def apply_softcap(acc_s): diff --git a/src/tileops/kernels/elementwise/_broadcast.py b/src/tileops/kernels/elementwise/_broadcast.py index 23025adf9..91684a4a1 100644 --- a/src/tileops/kernels/elementwise/_broadcast.py +++ b/src/tileops/kernels/elementwise/_broadcast.py @@ -4,6 +4,9 @@ import torch +# CUDA caps grid.y at 65535, and one grid axis carries the rows. +_CUDA_MAX_GRID_Y = 65535 + def _flat(t): """The flat view every PrimFunc here takes.""" @@ -102,10 +105,6 @@ def _is_contiguous_same_shape(coalesced_shape, a_strides, b_strides): ) -# CUDA caps grid.y at 65535, and one grid axis carries the rows. -_CUDA_MAX_GRID_Y = 65535 - - def row_broadcast_split(coalesced_shape, a_strides, b_strides): """``(rows, inner)`` when the innermost coalesced dim reads at stride 0 or 1. diff --git a/src/tileops/kernels/elementwise/_dtype.py b/src/tileops/kernels/elementwise/_dtype.py index ca83efb32..b4370a385 100644 --- a/src/tileops/kernels/elementwise/_dtype.py +++ b/src/tileops/kernels/elementwise/_dtype.py @@ -9,11 +9,6 @@ BOOL_STORAGE_DTYPE = "int8" -def log_for_output_precision(value, wide): - """Return ``log(wide)`` computed to the precision *value*'s dtype can keep.""" - return T.log(wide) if value.dtype == "float32" else T.__log(wide) - - _BITWISE_DTYPES = ( torch.bool, torch.uint8, @@ -22,33 +17,28 @@ def log_for_output_precision(value, wide): torch.int32, torch.int64, ) - - # The dtypes every elementwise kernel refuses. _FP8_DTYPES = ( torch.float8_e4m3fn, torch.float8_e5m2, ) - - _FLOAT_DTYPES = ( torch.float16, torch.bfloat16, torch.float32, ) - - _LOGICAL_DTYPES = _BITWISE_DTYPES + _FLOAT_DTYPES - - _BINARY_FULL_DTYPES = _BITWISE_DTYPES + ( torch.float16, torch.bfloat16, torch.float32, ) +_BINARY_NO_BOOL_DTYPES = tuple(dt for dt in _BINARY_FULL_DTYPES if dt is not torch.bool) -_BINARY_NO_BOOL_DTYPES = tuple(dt for dt in _BINARY_FULL_DTYPES if dt is not torch.bool) +def log_for_output_precision(value, wide): + """Return ``log(wide)`` computed to the precision *value*'s dtype can keep.""" + return T.log(wide) if value.dtype == "float32" else T.__log(wide) def _torch_dtype_nbytes(dtype: torch.dtype) -> int: diff --git a/src/tileops/kernels/fft.py b/src/tileops/kernels/fft.py index 59321753b..03ba6da5e 100644 --- a/src/tileops/kernels/fft.py +++ b/src/tileops/kernels/fft.py @@ -20,6 +20,89 @@ __all__ = ["FFTC2CCall", "FFTC2CDecomposedKernel", "FFTC2COneCTAKernel"] +# n -> the radices of the one-CTA plan's passes. +_RADIX_PLAN = { + 2: (2,), + 4: (4,), + 8: (8,), + 16: (16,), + 32: (32,), + 64: (8, 8), + 128: (8, 8, 2), + 256: (16, 16), + 512: (8, 8, 8), + 1024: (16, 16, 4), + 2048: (16, 16, 8), + 4096: (16, 16, 16), + 8192: (16, 16, 16, 2), + 16384: (16, 16, 16, 4), +} +# (n, dtype) -> four-step factors, outermost first, one kernel each. +_FOUR_STEP_PLAN = { + (1 << 14, "complex128"): (256, 64), + (1 << 15, "complex64"): (256, 128), + (1 << 15, "complex128"): (256, 128), + (1 << 16, "complex64"): (256, 256), + (1 << 16, "complex128"): (256, 256), + (1 << 17, "complex64"): (128, 1024), + (1 << 17, "complex128"): (128, 1024), + (1 << 18, "complex64"): (256, 1024), + (1 << 18, "complex128"): (256, 1024), + (1 << 19, "complex64"): (256, 2048), + (1 << 19, "complex128"): (256, 2048), + (1 << 20, "complex64"): (1024, 1024), + (1 << 20, "complex128"): (1024, 1024), + (1 << 21, "complex64"): (1024, 2048), + (1 << 21, "complex128"): (1024, 2048), + (1 << 22, "complex64"): (2048, 2048), + (1 << 22, "complex128"): (2048, 2048), + (1 << 23, "complex64"): (2048, 4096), + (1 << 23, "complex128"): (2048, 4096), + (1 << 24, "complex64"): (4096, 4096), + (1 << 24, "complex128"): (4096, 4096), + (1 << 25, "complex64"): (256, 512, 256), + (1 << 25, "complex128"): (256, 512, 256), + (1 << 26, "complex64"): (512, 512, 256), + (1 << 26, "complex128"): (512, 512, 256), + (1 << 27, "complex64"): (256, 512, 1024), + (1 << 27, "complex128"): (256, 512, 1024), + (1 << 28, "complex64"): (512, 512, 1024), + (1 << 28, "complex128"): (512, 512, 1024), +} +# (n, dtype) -> transforms per CTA for each kernel of the plan. +_FOUR_STEP_TILE = { + (1 << 14, "complex128"): (4, 4), + (1 << 15, "complex64"): (16, 4), + (1 << 15, "complex128"): (8, 4), + (1 << 16, "complex64"): (16, 8), + (1 << 16, "complex128"): (8, 4), + (1 << 17, "complex64"): (16, 4), + (1 << 17, "complex128"): (16, 4), + (1 << 18, "complex64"): (16, 4), + (1 << 18, "complex128"): (8, 4), + (1 << 19, "complex64"): (32, 4), + (1 << 19, "complex128"): (8, 2), + (1 << 20, "complex64"): (8, 4), + (1 << 20, "complex128"): (4, 4), + (1 << 21, "complex64"): (8, 4), + (1 << 21, "complex128"): (4, 2), + (1 << 22, "complex64"): (8, 4), + (1 << 22, "complex128"): (4, 4), + (1 << 23, "complex64"): (8, 4), + (1 << 23, "complex128"): (4, 2), + (1 << 24, "complex64"): (4, 4), + (1 << 24, "complex128"): (2, 2), + (1 << 25, "complex64"): (32, 16, 8), + (1 << 25, "complex128"): (8, 4, 8), + (1 << 26, "complex64"): (16, 16, 8), + (1 << 26, "complex128"): (4, 4, 8), + (1 << 27, "complex64"): (32, 16, 8), + (1 << 27, "complex128"): (32, 4, 4), + (1 << 28, "complex64"): (16, 16, 8), + (1 << 28, "complex128"): (4, 4, 4), +} + + @dataclasses.dataclass(frozen=True) class FFTC2CCall(CallSpec): """What selects the C2C kernel; the batch is symbolic in every kernel, so it is absent.""" @@ -1317,91 +1400,6 @@ def main( return _func -# n -> the radices of the one-CTA plan's passes. -_RADIX_PLAN = { - 2: (2,), - 4: (4,), - 8: (8,), - 16: (16,), - 32: (32,), - 64: (8, 8), - 128: (8, 8, 2), - 256: (16, 16), - 512: (8, 8, 8), - 1024: (16, 16, 4), - 2048: (16, 16, 8), - 4096: (16, 16, 16), - 8192: (16, 16, 16, 2), - 16384: (16, 16, 16, 4), -} - -# (n, dtype) -> four-step factors, outermost first, one kernel each. -_FOUR_STEP_PLAN = { - (1 << 14, "complex128"): (256, 64), - (1 << 15, "complex64"): (256, 128), - (1 << 15, "complex128"): (256, 128), - (1 << 16, "complex64"): (256, 256), - (1 << 16, "complex128"): (256, 256), - (1 << 17, "complex64"): (128, 1024), - (1 << 17, "complex128"): (128, 1024), - (1 << 18, "complex64"): (256, 1024), - (1 << 18, "complex128"): (256, 1024), - (1 << 19, "complex64"): (256, 2048), - (1 << 19, "complex128"): (256, 2048), - (1 << 20, "complex64"): (1024, 1024), - (1 << 20, "complex128"): (1024, 1024), - (1 << 21, "complex64"): (1024, 2048), - (1 << 21, "complex128"): (1024, 2048), - (1 << 22, "complex64"): (2048, 2048), - (1 << 22, "complex128"): (2048, 2048), - (1 << 23, "complex64"): (2048, 4096), - (1 << 23, "complex128"): (2048, 4096), - (1 << 24, "complex64"): (4096, 4096), - (1 << 24, "complex128"): (4096, 4096), - (1 << 25, "complex64"): (256, 512, 256), - (1 << 25, "complex128"): (256, 512, 256), - (1 << 26, "complex64"): (512, 512, 256), - (1 << 26, "complex128"): (512, 512, 256), - (1 << 27, "complex64"): (256, 512, 1024), - (1 << 27, "complex128"): (256, 512, 1024), - (1 << 28, "complex64"): (512, 512, 1024), - (1 << 28, "complex128"): (512, 512, 1024), -} - -# (n, dtype) -> transforms per CTA for each kernel of the plan. -_FOUR_STEP_TILE = { - (1 << 14, "complex128"): (4, 4), - (1 << 15, "complex64"): (16, 4), - (1 << 15, "complex128"): (8, 4), - (1 << 16, "complex64"): (16, 8), - (1 << 16, "complex128"): (8, 4), - (1 << 17, "complex64"): (16, 4), - (1 << 17, "complex128"): (16, 4), - (1 << 18, "complex64"): (16, 4), - (1 << 18, "complex128"): (8, 4), - (1 << 19, "complex64"): (32, 4), - (1 << 19, "complex128"): (8, 2), - (1 << 20, "complex64"): (8, 4), - (1 << 20, "complex128"): (4, 4), - (1 << 21, "complex64"): (8, 4), - (1 << 21, "complex128"): (4, 2), - (1 << 22, "complex64"): (8, 4), - (1 << 22, "complex128"): (4, 4), - (1 << 23, "complex64"): (8, 4), - (1 << 23, "complex128"): (4, 2), - (1 << 24, "complex64"): (4, 4), - (1 << 24, "complex128"): (2, 2), - (1 << 25, "complex64"): (32, 16, 8), - (1 << 25, "complex128"): (8, 4, 8), - (1 << 26, "complex64"): (16, 16, 8), - (1 << 26, "complex128"): (4, 4, 8), - (1 << 27, "complex64"): (32, 16, 8), - (1 << 27, "complex128"): (32, 4, 4), - (1 << 28, "complex64"): (16, 16, 8), - (1 << 28, "complex128"): (4, 4, 4), -} - - def _plan_table() -> Dict[tuple, FFTPlan]: """One record per served (length, dtype); every one-CTA builder takes (row, grp).""" records = {} diff --git a/src/tileops/kernels/gemm/dense.py b/src/tileops/kernels/gemm/dense.py index d138b53cd..9e53f6411 100644 --- a/src/tileops/kernels/gemm/dense.py +++ b/src/tileops/kernels/gemm/dense.py @@ -42,6 +42,10 @@ _FP8_WS_BLOCK_K = 128 +_TILE_K = 8 +_SMEM_CAP = 224 * 1024 + + def _tma_misalignment( m: int, n: int, k: int, dtype: torch.dtype, trans_a: bool, trans_b: bool ) -> Optional[str]: @@ -2904,10 +2908,6 @@ def _gemm_small_batch_main( return _gemm_small_batch_func -_TILE_K = 8 -_SMEM_CAP = 224 * 1024 - - def _bandwidth_autotune_grid(rts: tuple, bns: tuple, nss: tuple) -> list[dict]: """Config grid for the bandwidth-mode kernels, guarded by thread and SMEM caps.""" return [ diff --git a/src/tileops/kernels/gemm/heuristics.py b/src/tileops/kernels/gemm/heuristics.py index c9c92cbe3..56e0e112a 100644 --- a/src/tileops/kernels/gemm/heuristics.py +++ b/src/tileops/kernels/gemm/heuristics.py @@ -67,6 +67,11 @@ _NS_CAP = {"basic": 4, "splitk": 4, "coop2": 4, "coop2_splitk": 4} +#: The row count the small-M split-K band was fitted at, by calibration key. A board +#: without an entry has no band. +_SMALL_M_SPLITK_M = {"h200": 32} + + @dataclass(frozen=True) class _Calibration: """The scorer's ranking constants for one board. @@ -483,11 +488,6 @@ def small_batch_config(n: int, k: int, sm_count: int) -> dict: return cfg -#: The row count the small-M split-K band was fitted at, by calibration key. A board -#: without an entry has no band. -_SMALL_M_SPLITK_M = {"h200": 32} - - def small_m_splitk_config( m: int, n: int, k: int, sm_count: int, device_name: str ) -> Optional[dict]: diff --git a/src/tileops/kernels/gemm/w4a16.py b/src/tileops/kernels/gemm/w4a16.py index d54f476bf..bc95f7a00 100644 --- a/src/tileops/kernels/gemm/w4a16.py +++ b/src/tileops/kernels/gemm/w4a16.py @@ -24,6 +24,10 @@ ) +# What identifies a tile shape, as opposed to how its K loop is sliced. +_TILE_KEYS = ("block_m", "block_n", "block_k", "num_stages", "threads") + + @dataclass(frozen=True) class _Layout: """Constants fixed by the packed-weight ABI, not tuning parameters.""" @@ -275,10 +279,6 @@ class _TileBuffers(NamedTuple): out_shared: Any -# What identifies a tile shape, as opposed to how its K loop is sliced. -_TILE_KEYS = ("block_m", "block_n", "block_k", "num_stages", "threads") - - def _select_config(m: int, n: int, k: int, group_size: int, sms: int) -> dict: """Choose a tile shape, then its lowest-cost whole-K, split-K, or stream-K variant.""" legal = list(_legal_configs(m, n, k, group_size, sms)) diff --git a/src/tileops/kernels/grouped_gemm/heuristics.py b/src/tileops/kernels/grouped_gemm/heuristics.py index f4a11d678..cc44ae277 100644 --- a/src/tileops/kernels/grouped_gemm/heuristics.py +++ b/src/tileops/kernels/grouped_gemm/heuristics.py @@ -25,6 +25,12 @@ ] +# Gated activations the epilogue can fuse: B stacks gate and up along N; a tile's B +# half-loads block_n / 2 gate columns and the matching up columns, one accumulator +# holds both, and the epilogue stores act(gate) * up, so C has N / 2 columns. +ACTIVATIONS = ("none", "silu_and_mul", "gelu_and_mul") + + @dataclasses.dataclass(frozen=True) class _HeuristicPolicy: """The constants the selector reads, in three kinds a reader must tell apart. @@ -148,12 +154,6 @@ class GemmType(str, enum.Enum): _FLAT_LIKE_TYPES = (GemmType.DENSE, GemmType.BATCHED, GemmType.K_GROUPED_CONTIGUOUS) -# Gated activations the epilogue can fuse: B stacks gate and up along N; a tile's B -# half-loads block_n / 2 gate columns and the matching up columns, one accumulator -# holds both, and the epilogue stores act(gate) * up, so C has N / 2 columns. -ACTIVATIONS = ("none", "silu_and_mul", "gelu_and_mul") - - class Major(str, enum.Enum): """Which logical dim is contiguous in memory for an operand.""" diff --git a/src/tileops/kernels/norm/fused_add_norm.py b/src/tileops/kernels/norm/fused_add_norm.py index 8c226b9a6..a61ede64e 100644 --- a/src/tileops/kernels/norm/fused_add_norm.py +++ b/src/tileops/kernels/norm/fused_add_norm.py @@ -37,6 +37,10 @@ # Fused Add + LayerNorm kernel +# This kernel serves 16-bit dtypes only, so one 16-byte access moves eight elements. +_VEC = VECTOR_ACCESS_BYTES // 2 + + @functools.lru_cache(maxsize=32) def _fused_add_layer_norm_kernel(M, N, eps, dtype): N_padded = align_up(N, ALIGNMENT) @@ -229,9 +233,6 @@ def forward( # Fused Add + RMSNorm kernel -# This kernel serves 16-bit dtypes only, so one 16-byte access moves eight elements. -_VEC = VECTOR_ACCESS_BYTES // 2 - @functools.lru_cache(maxsize=32) def _fused_add_rms_norm_kernel(M, N, eps, dtype, splits): diff --git a/src/tileops/kernels/pool/common.py b/src/tileops/kernels/pool/common.py index 0aff54713..5d665004f 100644 --- a/src/tileops/kernels/pool/common.py +++ b/src/tileops/kernels/pool/common.py @@ -7,6 +7,10 @@ from tileops.kernels.constants import STATIC_SHARED_BYTES, VECTOR_ACCESS_BYTES from tileops.kernels.kernel_base import Kernel +# Window sums promote to fp32 and cast back at the store: a narrow accumulator loses the +# low bits of a window this wide. +ACCUM_DTYPE = "float" + def dtype_itemsize(dtype: str) -> int: """Bytes one element of *dtype* takes, over the dtypes these kernels accept.""" @@ -103,11 +107,6 @@ def pool_output_dim( return max(out, 0) -# Window sums promote to fp32 and cast back at the store: a narrow accumulator loses the -# low bits of a window this wide. -ACCUM_DTYPE = "float" - - class AvgPoolWindow(NamedTuple): """One average-pooling problem, and the extents and facts that follow from it. diff --git a/src/tileops/kernels/reduction/_primitives.py b/src/tileops/kernels/reduction/_primitives.py index b00980d79..e5d887464 100644 --- a/src/tileops/kernels/reduction/_primitives.py +++ b/src/tileops/kernels/reduction/_primitives.py @@ -75,6 +75,11 @@ FRAGMENT_ELEMS_PER_THREAD: int = 64 +# Largest integer count fp32 carries exactly; a statistic folded through +# fp32 counts or weights is trusted only below it. +FP32_EXACT_INT_LIMIT = 1 << 24 + + def ceildiv_int(x: int, y: int) -> int: """Return ``ceil(x / y)`` for positive integer dimensions.""" return -(-x // y) @@ -856,11 +861,6 @@ class _LeadingAxisReducePolicy: _LEADING_POLICY = _LeadingAxisReducePolicy() -# Largest integer count fp32 carries exactly; a statistic folded through -# fp32 counts or weights is trusted only below it. -FP32_EXACT_INT_LIMIT = 1 << 24 - - def edge_axis_plan( shape: "tuple[int, ...]", k: int, diff --git a/src/tileops/kernels/reduction/call_spec.py b/src/tileops/kernels/reduction/call_spec.py index 809ce9736..836687910 100644 --- a/src/tileops/kernels/reduction/call_spec.py +++ b/src/tileops/kernels/reduction/call_spec.py @@ -15,6 +15,12 @@ ] +# The fused pass runs one block per kept column and has no other parallelism, so +# it takes over only where that alone is enough: the fewest kept columns that fill the +# device, per calibrated board. A board without an entry uses the general implementation. +_EDGE_FUSED_MIN_KEPT = {"h200": 32} + + @dataclasses.dataclass(frozen=True) class LogicalReduceCall(CallSpec): """Semantic and shape facts used to select a logical reduction implementation.""" @@ -37,12 +43,6 @@ def logical_reduce_region(call: LogicalReduceCall) -> bool: return call.op_kind in {"any", "all", "count_nonzero"} -# The fused pass runs one block per kept column and has no other parallelism, so -# it takes over only where that alone is enough: the fewest kept columns that fill the -# device, per calibrated board. A board without an entry uses the general implementation. -_EDGE_FUSED_MIN_KEPT = {"h200": 32} - - def logical_edge_fused_region(call: LogicalReduceCall) -> bool: """The edge-axis logical reduction region a calibrated board serves with the fused pass.""" diff --git a/src/tileops/kernels/reduction/logical_reduce.py b/src/tileops/kernels/reduction/logical_reduce.py index 307307a19..2cee9fa3e 100644 --- a/src/tileops/kernels/reduction/logical_reduce.py +++ b/src/tileops/kernels/reduction/logical_reduce.py @@ -70,6 +70,15 @@ _UNSUPPORTED_STORAGE_DTYPES = _BYTE_REINTERPRETED_DTYPES | _WIDENED_STORAGE_DTYPES +# Elements a lane folds in the edge-fused pass. One block holds a `trail`-wide +# fp32 fragment, so this is what fixes its register footprint per lane rather +# than letting it grow with the row. Eight is the flat optimum at every width +# the manifest asks for. +_FUSED_EDGE_ELEMS_PER_LANE = 8 +_FUSED_EDGE_MIN_THREADS = 64 +_FUSED_EDGE_MAX_THREADS = 1024 + + def storage_dtype_for(dtype: torch.dtype) -> torch.dtype: """The dtype the prim_func declares for an input of *dtype*.""" if dtype in _BYTE_REINTERPRETED_DTYPES: @@ -108,15 +117,6 @@ def _logical_out_dtype(op_kind: str, partial: bool) -> str: return "float32" if partial else "int64" -# Elements a lane folds in the edge-fused pass. One block holds a `trail`-wide -# fp32 fragment, so this is what fixes its register footprint per lane rather -# than letting it grow with the row. Eight is the flat optimum at every width -# the manifest asks for. -_FUSED_EDGE_ELEMS_PER_LANE = 8 -_FUSED_EDGE_MIN_THREADS = 64 -_FUSED_EDGE_MAX_THREADS = 1024 - - def fused_edge_threads(trail: int) -> int: """The thread width the edge-fused pass runs a ``trail``-wide row at.""" lanes = ceildiv_int(trail, _FUSED_EDGE_ELEMS_PER_LANE) diff --git a/src/tileops/kernels/reduction/reduce.py b/src/tileops/kernels/reduction/reduce.py index 719cfb787..516b5200f 100644 --- a/src/tileops/kernels/reduction/reduce.py +++ b/src/tileops/kernels/reduction/reduce.py @@ -45,6 +45,9 @@ } +_LEADING_AXIS_KINDS = frozenset({"sum", "mean", "amax", "amin"}) + + @dataclass(frozen=True) class ProductReducePolicy: """Launch heuristics for product reductions.""" @@ -62,9 +65,6 @@ class ProductReducePolicy: # Simple reduce kernel -_LEADING_AXIS_KINDS = frozenset({"sum", "mean", "amax", "amin"}) - - class ReduceKernel(Kernel): """Unified reduce kernel supporting sum/mean/amin/amax/prod/std/var/var_mean. diff --git a/src/tileops/manifest/dtype_rules.py b/src/tileops/manifest/dtype_rules.py index e3a85948b..108acaa58 100644 --- a/src/tileops/manifest/dtype_rules.py +++ b/src/tileops/manifest/dtype_rules.py @@ -1,28 +1,37 @@ -"""The dtype registry: every dtype name a manifest may write, with its bits per element. +"""The dtype registry: every dtype name a manifest may write, with its bits per element and +its category. This package depends on nothing beyond the standard library and PyYAML, so dtypes are names. """ from __future__ import annotations -__all__ = ["DTYPE_BITS"] +__all__ = ["DTYPE_BITS", "DTYPE_CATEGORY", "FLOAT8_DTYPES"] -DTYPE_BITS: dict[str, int] = { - "bool": 8, - "uint8": 8, - "int8": 8, - "int16": 16, - "int32": 32, - "int64": 64, - "float16": 16, - "bfloat16": 16, - "float32": 32, - "float64": 64, - "complex64": 64, - "complex128": 128, - "float8_e4m3fn": 8, - "float8_e5m2": 8, - "float8_e4m3": 8, - "float8_e5m2fnuz": 8, - "float8_e4m3fnuz": 8, +# name: (bits per element, category) +_DTYPES: dict[str, tuple[int, str]] = { + "bool": (8, "bool"), + "uint8": (8, "int"), + "int8": (8, "int"), + "int16": (16, "int"), + "int32": (32, "int"), + "int64": (64, "int"), + "float16": (16, "float"), + "bfloat16": (16, "float"), + "float32": (32, "float"), + "float64": (64, "float"), + "complex64": (64, "complex"), + "complex128": (128, "complex"), + "float8_e4m3fn": (8, "float"), + "float8_e5m2": (8, "float"), + "float8_e4m3": (8, "float"), + "float8_e5m2fnuz": (8, "float"), + "float8_e4m3fnuz": (8, "float"), } + +DTYPE_BITS: dict[str, int] = {name: bits for name, (bits, _) in _DTYPES.items()} +# 'bool', 'int', 'float' or 'complex'. +DTYPE_CATEGORY: dict[str, str] = {name: kind for name, (_, kind) in _DTYPES.items()} +FLOAT8_DTYPES: frozenset[str] = frozenset( + name for name, (bits, kind) in _DTYPES.items() if bits == 8 and kind == "float" +) diff --git a/src/tileops/manifest/expr.py b/src/tileops/manifest/expr.py index ba7cbf545..8e588f466 100644 --- a/src/tileops/manifest/expr.py +++ b/src/tileops/manifest/expr.py @@ -54,17 +54,8 @@ ] -class SignatureError(ValueError): - """A declaration outside the schema; the message names it.""" - - -class EvaluationError(SignatureError): - """An expression that failed to evaluate; the message names its declaration.""" - - # The value of an expression that reads more than a point fixes. OPEN = object() - _COMPREHENSION_CALLEES = frozenset({"all", "sum", "max", "min"}) _NODES = ( ast.BoolOp, @@ -102,6 +93,15 @@ class EvaluationError(SignatureError): ast.comprehension, ast.keyword, ) +_UNFOLDED = (ast.Constant, ast.Name, ast.List, ast.Starred, ast.Slice, ast.GeneratorExp) + + +class SignatureError(ValueError): + """A declaration outside the schema; the message names it.""" + + +class EvaluationError(SignatureError): + """An expression that failed to evaluate; the message names its declaration.""" # ---------------------------------------------------------------- parsing and the language @@ -644,7 +644,6 @@ def visit_Call(self, node): _NAMESPACE = namespace() -_UNFOLDED = (ast.Constant, ast.Name, ast.List, ast.Starred, ast.Slice, ast.GeneratorExp) def bind(node: ast.expr, point: dict) -> ast.expr: diff --git a/src/tileops/manifest/primitives.py b/src/tileops/manifest/primitives.py index cf491fc47..1c8571026 100644 --- a/src/tileops/manifest/primitives.py +++ b/src/tileops/manifest/primitives.py @@ -10,7 +10,7 @@ import numbers from types import SimpleNamespace -from .dtype_rules import DTYPE_BITS +from .dtype_rules import DTYPE_BITS, DTYPE_CATEGORY # The seed both conftests give the global RNG; every private workload RNG derives from it. WORKLOAD_SEED = 1235 @@ -32,6 +32,23 @@ ] +# The largest finite value of each floating dtype; the lowest is its negation. +_FLOAT_MAX = { + "float16": 65504.0, + "bfloat16": 3.3895313892515355e38, + "float32": 3.4028234663852886e38, + "float64": 1.7976931348623157e308, + "float8_e4m3fn": 448.0, + "float8_e4m3": 240.0, + "float8_e5m2": 57344.0, + "float8_e4m3fnuz": 240.0, + "float8_e5m2fnuz": 57344.0, +} +_COMPLEX_PART = {"complex64": "float32", "complex128": "float64"} +# Floating formats without an infinity. +_NO_INF = frozenset({"float8_e4m3fn", "float8_e4m3fnuz", "float8_e5m2fnuz"}) + + def normalize_axis(axis: int, rank: int) -> int: """At rank 0, `0` and `-1` name the scalar axis; otherwise an axis lies in [-rank, rank).""" lo, hi = (-1, 1) if rank == 0 else (-rank, rank) @@ -163,7 +180,7 @@ def moe_capacity(layout, rows, experts): def promote_int_to_float(dtype): """float32 for an integral dtype, else the dtype itself.""" - return "float32" if dtype in ("uint8", "int8", "int16", "int32", "int64") else dtype + return "float32" if category(dtype) == "int" else dtype def coalesce_dtype(value, dtype): @@ -193,29 +210,10 @@ def repeat(value, count): return [value] * count -# The largest finite value of each floating dtype; the lowest is its negation. -_FLOAT_MAX = { - "float16": 65504.0, - "bfloat16": 3.3895313892515355e38, - "float32": 3.4028234663852886e38, - "float64": 1.7976931348623157e308, - "float8_e4m3fn": 448.0, - "float8_e4m3": 240.0, - "float8_e5m2": 57344.0, - "float8_e4m3fnuz": 240.0, - "float8_e5m2fnuz": 57344.0, -} -_COMPLEX_PART = {"complex64": "float32", "complex128": "float64"} - - def category(x): """`'bool'`, `'int'`, `'float'` or `'complex'`: the category of a number or a dtype name.""" if isinstance(x, str): - if x == "bool": - return "bool" - if x in _COMPLEX_PART: - return "complex" - return "float" if x in _FLOAT_MAX else "int" + return DTYPE_CATEGORY.get(x, "int") if isinstance(x, bool): return "bool" if isinstance(x, numbers.Integral): @@ -223,10 +221,6 @@ def category(x): return "float" if isinstance(x, numbers.Real) else "complex" -# Floating formats without an infinity. -_NO_INF = frozenset({"float8_e4m3fn", "float8_e4m3fnuz", "float8_e5m2fnuz"}) - - def _fits_float(v, dtype): if math.isinf(v): return dtype not in _NO_INF @@ -449,6 +443,24 @@ def causal_topk_indices(rng, batch, seq, heads_kv, k, extent, start, stride): return rows +def sparse_topk_positions(rng, lengths, queries, k): + """Per request of length `n` and query `s`, up to `k` distinct positions in `[0, n - queries + s]`, padded with -1.""" + if queries <= 0 or k <= 0 or any(n < queries for n in lengths): + raise ValueError( + f"attn.sparse_topk_positions needs S_q > 0, K > 0 and every length >= S_q, " + f"got lengths={lengths}, S_q={queries}, K={k}" + ) + rows = [] + for n in lengths: + per_query = [] + for s in range(queries): + visible = n - queries + s + 1 + picked = rng.sample(range(visible), min(k, visible)) + per_query.append(picked + [-1] * (k - len(picked))) + rows.append(per_query) + return rows + + def key_windows(lengths, first, count, side): """Per query position, the first key of its sequence (`start`) or one past itself (`end`).""" if not lengths or any(n <= 0 for n in lengths): @@ -523,6 +535,7 @@ def moe_layout_metadata(layout, rows, experts): "sample_indices": sample_indices, "moe.layout_metadata": moe_layout_metadata, "causal_topk_indices": causal_topk_indices, + "attn.sparse_topk_positions": sparse_topk_positions, "key_windows": key_windows, "full": full, } @@ -543,6 +556,7 @@ def moe_layout_metadata(layout, rows, experts): "sample_indices": (("Int", "Int"), "Value"), "moe.layout_metadata": (("ADT", "Int", "Int"), "Value"), "causal_topk_indices": (("Int", "Int", "Int", "Int", "Int", "Int", "Int"), "Value"), + "attn.sparse_topk_positions": (("Seq[Int]", "Int", "Int"), "Value"), "key_windows": (("Seq[Int]", "Int", "Int", "'start' | 'end'"), "Value"), "full": (("Seq[Int]", "Int"), "Value"), } @@ -563,6 +577,7 @@ def moe_layout_metadata(layout, rows, experts): "sample_indices": 1, "moe.layout_metadata": 1, "causal_topk_indices": 4, + "attn.sparse_topk_positions": 3, "key_windows": 1, } # The shape of each generator's result, from its arguments. @@ -589,6 +604,7 @@ def moe_layout_metadata(layout, rows, experts): heads, k, ), + "attn.sparse_topk_positions": lambda L, queries, k: (len(L), queries, k), "key_windows": lambda L, first, count, side: (count,), "full": lambda shape, value: tuple(shape), } @@ -600,6 +616,7 @@ def moe_layout_metadata(layout, rows, experts): "topk_ids", "sample_indices", "causal_topk_indices", + "attn.sparse_topk_positions", } ) diff --git a/src/tileops/manifest/spec/attention.yaml b/src/tileops/manifest/spec/attention.yaml index c62b43c02..12f62576f 100644 --- a/src/tileops/manifest/spec/attention.yaml +++ b/src/tileops/manifest/spec/attention.yaml @@ -493,3 +493,242 @@ NSAVarlenFwdOp: roofline: # The blocks the selection keeps decide the scores and the key/value rows read. func: "tileops.perf.formulas.nsa_fwd_varlen_roofline" + +# --------------------------------------------------------------------------- +# attention — paged KV cache and latent-cache operators +# --------------------------------------------------------------------------- +# A paged cache is [NP, PS, ...]: NP pages of PS token rows. Slot s addresses row +# s % PS of page s // PS. The Multi-Head Latent Attention (MLA) latent cache is one +# [NP, PS, R] tensor with no head axis: every query head reads the same latent row. + +MultiHeadLatentAttentionPagedFwdOp: + # Absorbed-form MLA decode over a paged latent cache: key = the whole cache row + # (kv_c ‖ k_pe), value = its first kv_lora_rank columns. The S_q query tokens are the + # last S_q positions of each request's cache_seqlens. sm_scale defaults to DK ** -0.5. + family: attention + status: spec-only + signature: + forall: {B: Dim, S_q: Dim, H: Dim, DK: Dim, NP: Dim, PS: Dim, W: Dim, T: "DType[float16 | bfloat16]", KV: "DType[float16 | bfloat16 | float8_e4m3fn]", cache_lens: "Seq[Int]"} + params: + kv_lora_rank: {type: int} + is_causal: {type: bool, default: true} + sm_scale: {type: "float | None", default: null} + inputs: + q: {dtype: T, shape: "[B, S_q, H, DK]"} + kv_cache: {dtype: KV, shape: "[NP, PS, DK]", contiguous: true} + block_table: {dtype: int32, shape: "[B, W]", values: "paged_block_table(B, W, NP)", requires: ["in_range(0, NP)"]} + cache_seqlens: {dtype: int32, shape: "[B]", values: "as_tensor(cache_lens)", requires: ["in_range(S_q, W * PS + 1)"]} + # Per-tensor dequantization scale of an FP8 cache. + kv_scale: {dtype: float32, shape: "[1]", optional: true} + outputs: + o: {dtype: T, shape: "[B, S_q, H, kv_lora_rank]"} + lse: {dtype: float32, shape: "[B, S_q, H]"} + dtype_combos: + - {T: float16, KV: float16} + - {T: bfloat16, KV: bfloat16} + - {T: float16, KV: float8_e4m3fn} + - {T: bfloat16, KV: float8_e4m3fn} + shape_rules: + - "S_q > 0 and PS > 0" + - "0 < kv_lora_rank < DK" + - "(KV == 'float8_e4m3fn') == present(kv_scale)" + workloads: + - {S_q: 1, H: 128, DK: 576, NP: 4096, PS: 64, W: 64, kv_lora_rank: 512, cache_lens: "repeat(4096, 64)", dtype_cases: [{T: bfloat16, KV: bfloat16}, {T: float16, KV: float16}], label: ds-v3-4k} + - {S_q: 1, H: 128, DK: 576, NP: 32768, PS: 64, W: 512, kv_lora_rank: 512, cache_lens: "repeat(32768, 64)", dtype_cases: [{T: bfloat16, KV: bfloat16}], label: ds-v3-32k} + - {S_q: 2, H: 128, DK: 576, NP: 4096, PS: 64, W: 128, kv_lora_rank: 512, cache_lens: "repeat(8192, 32)", dtype_cases: [{T: bfloat16, KV: bfloat16}], label: ds-v3-mtp} + - {S_q: 1, H: 64, DK: 576, NP: 32768, PS: 64, W: 256, kv_lora_rank: 512, cache_lens: "repeat(16384, 128)", some: [kv_scale], dtype_cases: [{T: bfloat16, KV: float8_e4m3fn}], label: kimi-k2-fp8} + roofline: + # The scores each query row sees, and the cache rows the block table reaches. + func: "tileops.perf.formulas.mla_paged_fwd_roofline" + +MultiHeadLatentAttentionVarlenFwdOp: + # MLA prefill over packed requests after the latent is decompressed: the key of head h + # is k_nope[:, h] ‖ k_pe, with k_pe shared by every head; queries and keys are the same + # tokens. sm_scale defaults to (DN + PE) ** -0.5. + family: attention + status: spec-only + signature: + forall: {B: Dim, T_q: Dim, H: Dim, DN: Dim, PE: Dim, DV: Dim, T: "DType[float16 | bfloat16]", seq_lens: "Seq[Int]"} + params: + is_causal: {type: bool, default: true} + sm_scale: {type: "float | None", default: null} + inputs: + q: {dtype: T, shape: "[T_q, H, DN + PE]"} + k_nope: {dtype: T, shape: "[T_q, H, DN]"} + k_pe: {dtype: T, shape: "[T_q, PE]"} + v: {dtype: T, shape: "[T_q, H, DV]"} + cu_seqlens: {dtype: int32, shape: "[B + 1]", values: "prefix_sum(seq_lens)", requires: ["prefix_offsets(T_q)"]} + outputs: + o: {dtype: T, shape: "[T_q, H, DV]"} + lse: {dtype: float32, shape: "[T_q, H]"} + shape_rules: + - "DN + PE > 0" + workloads: + - {T_q: 32768, H: 128, DN: 128, PE: 64, DV: 128, seq_lens: "repeat(4096, 8)", dtype_cases: [{T: bfloat16}, {T: float16}], label: ds-v3-8x4k} + - {T_q: 32256, H: 128, DN: 128, PE: 64, DV: 128, seq_lens: [512, 1024, 2048, 4096, 8192, 16384], dtype_cases: [{T: bfloat16}], label: ds-v3-mixed} + - {T_q: 32768, H: 64, DN: 128, PE: 64, DV: 128, seq_lens: "repeat(8192, 4)", dtype_cases: [{T: bfloat16}], label: kimi-k2-4x8k} + roofline: + # The scores each request sees under its mask; every tensor moves once. + func: "tileops.perf.formulas.mla_varlen_fwd_roofline" + +PagedKVCacheWriteFwdOp: + # Scatters token n's k and v rows into slot slot_mapping[n] of k_pages and v_pages; -1 + # skips the token, and distinct non-negative slots are the caller's obligation. An FP8 + # cache stores (k / k_scale) and (v / v_scale) with a saturating cast. + family: attention + status: spec-only + signature: + forall: {N: Dim, H_kv: Dim, D: Dim, NP: Dim, PS: Dim, T: "DType[float16 | bfloat16]", KV: "DType[float16 | bfloat16 | float8_e4m3fn]"} + inputs: + k: {dtype: T, shape: "[N, H_kv, D]"} + v: {dtype: T, shape: "[N, H_kv, D]"} + k_pages: {dtype: KV, shape: "[NP, PS, H_kv, D]", mutated: true, contiguous: true} + v_pages: {dtype: KV, shape: "[NP, PS, H_kv, D]", mutated: true, contiguous: true} + slot_mapping: {dtype: int64, shape: "[N]", values: "sample_indices(N, NP * PS)", requires: ["in_range(-1, NP * PS)"]} + k_scale: {dtype: float32, shape: "[1]", optional: true} + v_scale: {dtype: float32, shape: "[1]", optional: "present(k_scale)"} + outputs: {} + dtype_combos: + - {T: float16, KV: float16} + - {T: bfloat16, KV: bfloat16} + - {T: float16, KV: float8_e4m3fn} + - {T: bfloat16, KV: float8_e4m3fn} + shape_rules: + - "PS > 0" + - "(KV == 'float8_e4m3fn') == present(k_scale)" + workloads: + - {N: 64, H_kv: 8, D: 128, NP: 4096, PS: 16, dtype_cases: [{T: bfloat16, KV: bfloat16}, {T: float16, KV: float16}], label: llama-8b-decode} + - {N: 8192, H_kv: 8, D: 128, NP: 4096, PS: 16, dtype_cases: [{T: bfloat16, KV: bfloat16}], label: llama-8b-prefill} + - {N: 4096, H_kv: 4, D: 128, NP: 4096, PS: 16, some: [k_scale], dtype_cases: [{T: bfloat16, KV: float8_e4m3fn}], label: qwen3-235b-fp8} + - {N: 256, H_kv: 8, D: 64, NP: 4096, PS: 16, dtype_cases: [{T: bfloat16, KV: bfloat16}], label: gpt-oss-120b} + roofline: + # Only the tokens with a slot are read and written. + func: "tileops.perf.formulas.paged_kv_cache_write_roofline" + +MultiHeadLatentAttentionKVCacheWriteFwdOp: + # Writes kv_c ‖ k_pe of token n into slot slot_mapping[n] of the latent cache; -1 skips + # the token, and distinct non-negative slots are the caller's obligation. With + # fuse_rope, k_pe is rotated at positions[n] by cos_sin_cache (cos ‖ sin halves) as it + # is written; k_pe itself is not modified. An FP8 cache stores the row divided by scale + # with a saturating cast. + family: attention + status: spec-only + signature: + forall: {N: Dim, DC: Dim, PE: Dim, NP: Dim, PS: Dim, P: Dim, T: "DType[float16 | bfloat16]", KV: "DType[float16 | bfloat16 | float8_e4m3fn]", C: "DType[float16 | bfloat16 | float32]", seq_lens: "Seq[Int]"} + params: + fuse_rope: {type: bool, default: false} + rope_layout: {type: "'neox' | 'interleaved'", default: neox} + inputs: + kv_c: {dtype: T, shape: "[N, DC]"} + k_pe: {dtype: T, shape: "[N, PE]"} + kv_cache: {dtype: KV, shape: "[NP, PS, DC + PE]", mutated: true, contiguous: true} + slot_mapping: {dtype: int64, shape: "[N]", values: "sample_indices(N, NP * PS)", requires: ["in_range(-1, NP * PS)"]} + scale: {dtype: float32, shape: "[1]", optional: true} + positions: {dtype: int64, shape: "[N]", optional: fuse_rope, values: "packed_positions(seq_lens)", requires: ["in_range(0, P)"]} + cos_sin_cache: {dtype: C, shape: "[P, PE]", optional: fuse_rope} + outputs: {} + dtype_combos: + - {T: float16, KV: float16} + - {T: bfloat16, KV: bfloat16} + - {T: float16, KV: float8_e4m3fn} + - {T: bfloat16, KV: float8_e4m3fn} + shape_rules: + - "PS > 0" + - "not fuse_rope or PE % 2 == 0" + - "(KV == 'float8_e4m3fn') == present(scale)" + workloads: + - {N: 64, DC: 512, PE: 64, NP: 2048, PS: 64, dtype_cases: [{T: bfloat16, KV: bfloat16}], label: ds-v3-decode} + - {N: 16384, DC: 512, PE: 64, NP: 2048, PS: 64, dtype_cases: [{T: bfloat16, KV: bfloat16}, {T: float16, KV: float16}], label: ds-v3-prefill} + - {DC: 512, PE: 64, NP: 2048, PS: 64, P: 163840, seq_lens: "repeat(1024, 4)", fuse_rope: true, rope_layout: interleaved, some: [scale], dtype_cases: [{T: bfloat16, KV: float8_e4m3fn, C: float32}], label: ds-v3-fp8-rope} + - {DC: 512, PE: 64, NP: 2048, PS: 64, P: 131072, seq_lens: [128], fuse_rope: true, rope_layout: interleaved, dtype_cases: [{T: bfloat16, KV: bfloat16, C: bfloat16}], label: kimi-k2-rope} + roofline: + # Only the tokens with a slot are read and written; the rotation reads the cos/sin + # rows their positions name. + func: "tileops.perf.formulas.mla_kv_cache_write_roofline" + +MergeAttentionStatesFwdOp: + # Merges two partial attention results over disjoint key sets by their log-sum-exp: + # v = (v_a * e^s_a + v_b * e^s_b) / (e^s_a + e^s_b), s = log(e^s_a + e^s_b), computed + # against max(s_a, s_b) (FlashInfer merge_state). A -inf state contributes nothing; two + # merge to s = -inf and v = 0. + family: attention + status: spec-only + signature: + forall: {T_q: Dim, H: Dim, D: Dim, T: "DType[float16 | bfloat16]"} + inputs: + v_a: {dtype: T, shape: "[T_q, H, D]"} + s_a: {dtype: float32, shape: "[T_q, H]"} + v_b: {dtype: T, shape: "[T_q, H, D]"} + s_b: {dtype: float32, shape: "[T_q, H]"} + outputs: + v: {dtype: T, shape: "[T_q, H, D]"} + s: {dtype: float32, shape: "[T_q, H]"} + workloads: + - {T_q: 8192, H: 128, D: 128, dtype_cases: [{T: bfloat16}, {T: float16}], label: ds-v3-chunked-prefill} + - {T_q: 4096, H: 64, D: 128, dtype_cases: [{T: bfloat16}], label: llama-70b-cascade} + - {T_q: 256, H: 64, D: 128, dtype_cases: [{T: bfloat16}], label: qwen3-235b-dcp} + roofline: + # Per row: the max, two subtracts, two exps, the sum, log and add, and the two weights' + # divides; per element: two multiplies and an add. + flops: "T_q * H * (3 * D + 10)" + +DeepSeekSparseAttentionPagedFwdOp: + # DeepSeek Sparse Attention: MLA decode over the key positions indices names, read from + # a paged cache in the FlashMLA DeepSeek-V3.2 FP8 row format. A 656-byte row holds 512 + # float8_e4m3fn latent values, 4 float32 scales (one per 128 of them) and 64 bfloat16 + # rope values; the value is the dequantized latent. -1 pads indices; every other entry + # below cache_seqlens[b] is the caller's obligation. sm_scale defaults to 576 ** -0.5. + family: attention + status: spec-only + signature: + forall: {B: Dim, S_q: Dim, H: Dim, K: Dim, NP: Dim, PS: Dim, W: Dim, cache_lens: "Seq[Int]"} + params: + sm_scale: {type: "float | None", default: null} + inputs: + q: {dtype: bfloat16, shape: "[B, S_q, H, 576]"} + kv_cache: {dtype: uint8, shape: "[NP, PS, 656]", contiguous: true} + block_table: {dtype: int32, shape: "[B, W]", values: "paged_block_table(B, W, NP)", requires: ["in_range(0, NP)"]} + cache_seqlens: {dtype: int32, shape: "[B]", values: "as_tensor(cache_lens)", requires: ["in_range(S_q, W * PS + 1)"]} + indices: {dtype: int32, shape: "[B, S_q, K]", values: "attn.sparse_topk_positions(cache_lens, S_q, K)", requires: ["in_range(-1, W * PS)"]} + outputs: + o: {dtype: bfloat16, shape: "[B, S_q, H, 512]"} + lse: {dtype: float32, shape: "[B, S_q, H]"} + shape_rules: + - "S_q > 0 and PS > 0" + workloads: + - {S_q: 1, H: 128, K: 2048, NP: 32768, PS: 64, W: 512, cache_lens: "repeat(32768, 64)", label: ds-v32-32k} + - {S_q: 2, H: 128, K: 2048, NP: 32768, PS: 64, W: 2048, cache_lens: "repeat(131072, 16)", label: ds-v32-mtp-128k} + roofline: + # The selected scores, and the cache rows they resolve to through the block table. + func: "tileops.perf.formulas.dsa_paged_fwd_roofline" + +PagedKVCacheGatherFwdOp: + # Gathers request b's cache positions [start_b, start_b + len_b) into rows + # cu_seq_lens[b] .. cu_seq_lens[b + 1] of dst, start_b = seq_starts[b] or 0; an FP8 + # cache is dequantized by scale into out_dtype. dst carries the length cu_seq_lens[-1], + # which no input shape states. + family: attention + status: spec-only + signature: + forall: {B: Dim, T_q: Dim, NP: Dim, PS: Dim, W: Dim, E: Shape, KV: "DType[float16 | bfloat16 | float8_e4m3fn]", seq_lens: "Seq[Int]", starts: "Seq[Int]"} + params: + out_dtype: {type: "float16 | bfloat16 | None", default: null} + inputs: + dst: {dtype: "coalesce_dtype(out_dtype, KV)", shape: "[T_q, *E]", mutated: true, write_only: true} + cache: {dtype: KV, shape: "[NP, PS, *E]", contiguous: true} + block_table: {dtype: int32, shape: "[B, W]", values: "paged_block_table(B, W, NP)", requires: ["in_range(0, NP)"]} + cu_seq_lens: {dtype: int32, shape: "[B + 1]", values: "prefix_sum(seq_lens)", requires: ["prefix_offsets(T_q)", "max_segment(W * PS)"]} + seq_starts: {dtype: int32, shape: "[B]", optional: true, values: "as_tensor(starts)", requires: ["in_range(0, W * PS)", "attn.paged_fits(cu_seq_lens, W * PS)"]} + scale: {dtype: float32, shape: "[1]", optional: true} + outputs: {} + shape_rules: + - "PS > 0" + - "(KV == 'float8_e4m3fn') == present(scale)" + - "present(out_dtype) if KV == 'float8_e4m3fn' else (not present(out_dtype) or out_dtype.value == KV)" + workloads: + - {T_q: 131072, NP: 2048, PS: 64, W: 256, E: [576], seq_lens: "repeat(16384, 8)", out_dtype: bfloat16, some: [scale], dtype_cases: [{KV: float8_e4m3fn}], label: ds-v3-chunk-fp8} + - {T_q: 129024, NP: 6144, PS: 64, W: 1024, E: [576], seq_lens: [2048, 4096, 8192, 16384, 32768, 65536], dtype_cases: [{KV: bfloat16}], label: ds-v3-mixed} + - {T_q: 32768, NP: 4096, PS: 16, W: 512, E: [8, 128], seq_lens: "repeat(4096, 8)", starts: "repeat(2048, 8)", some: [seq_starts], dtype_cases: [{KV: bfloat16}], label: llama-70b-cp} + roofline: + # The cache rows each request's range reaches through the block table. + func: "tileops.perf.formulas.paged_kv_cache_gather_roofline" diff --git a/src/tileops/manifest/spec/linear_attention.yaml b/src/tileops/manifest/spec/linear_attention.yaml index b0ed018f4..aaf0594bd 100644 --- a/src/tileops/manifest/spec/linear_attention.yaml +++ b/src/tileops/manifest/spec/linear_attention.yaml @@ -99,6 +99,75 @@ GatedDeltaNetFwdOp: roofline: flops: "B * S * HV * 7 * K * V" +KimiDeltaAttentionFwdOp: + # Kimi Delta Attention (KDA): the gated delta rule with a per-key-channel decay g + # (log space), one contract for equal-length, packed-varlen and single-token calls, with + # GatedDeltaNetFwdOp's state layout; final_state is always returned. With + # use_gate_in_kernel, g is the raw gate input and the decay is + # -exp(A_log) * softplus(g + dt_bias), or lower_bound * sigmoid(exp(A_log) * (g + dt_bias)) + # when lower_bound is passed (FLA fused_recurrent_kda). + family: linear_attention + status: spec-only + signature: + types: + State: + params: {packed: Bool, vf: Bool, B: Dim, N: Dim, HV: Dim, K: Dim, V: Dim} + match: [packed, vf] + cases: + - {when: [false, false], is: "[B, HV, K, V]"} + - {when: [false, true], is: "[B, HV, V, K]"} + - {when: [true, false], is: "[N, HV, K, V]"} + - {when: [true, true], is: "[N, HV, V, K]"} + forall: {B: Dim, S: Dim, H: Dim, HV: Dim, K: Dim, V: Dim, N: Dim, T: "DType[float16 | bfloat16]", seq_lens: "Seq[Int]"} + params: + scale: {type: "float | None", default: null} + use_qk_l2norm_in_kernel: {type: bool, default: false} + use_beta_sigmoid_in_kernel: {type: bool, default: false} + allow_neg_eigval: {type: bool, default: false} + state_v_first: {type: bool, default: false} + use_gate_in_kernel: {type: bool, default: false} + lower_bound: {type: "float | None", default: null} + inputs: + q: {dtype: T, shape: "[B, S, H, K]"} + k: {dtype: T, shape: "[B, S, H, K]"} + v: {dtype: T, shape: "[B, S, HV, V]"} + g: {dtype: T, shape: "[B, S, HV, K]"} + beta: {dtype: T, shape: "[B, S, HV]"} + initial_state: {dtype: float32, shape: "State[present(cu_seqlens), state_v_first, B, N, HV, K, V]", optional: true} + cu_seqlens: {dtype: int64, shape: "[N + 1]", optional: true, values: "prefix_sum(seq_lens)", requires: ["prefix_offsets(S)"]} + cu_seqlens_cpu: {dtype: int64, shape: "[N + 1]", optional: true, device: cpu, values: "prefix_sum(seq_lens)", requires: ["prefix_offsets(S)"]} + A_log: {dtype: float32, shape: "[HV]", optional: use_gate_in_kernel} + dt_bias: {dtype: float32, shape: "[HV * K]", optional: true} + outputs: + o: {dtype: T, shape: "[B, S, HV, V]"} + final_state: {dtype: float32, shape: "State[present(cu_seqlens), state_v_first, B, N, HV, K, V]"} + shape_rules: + - "H > 0 and HV % H == 0" + - "K > 0 and V > 0" + - "not present(cu_seqlens_cpu) or present(cu_seqlens)" + - "not present(cu_seqlens) or B == 1" + - "not allow_neg_eigval or use_beta_sigmoid_in_kernel" + - "not present(dt_bias) or use_gate_in_kernel" + - "not present(lower_bound) or use_gate_in_kernel" + workloads: + - {B: 1, S: 4096, H: 32, HV: 32, K: 128, V: 128, seq_lens: "repeat(1024, 4)", use_qk_l2norm_in_kernel: true, some: [cu_seqlens], dtype_cases: [{T: bfloat16}, {T: float16}], label: kimi-linear-4k} + - {B: 1, S: 32768, H: 32, HV: 32, K: 128, V: 128, seq_lens: "repeat(8192, 4)", use_qk_l2norm_in_kernel: true, some: [cu_seqlens], dtype_cases: [{T: bfloat16}], label: kimi-linear-32k} + - {B: 1, S: 1, H: 32, HV: 32, K: 128, V: 128, use_qk_l2norm_in_kernel: true, some: [initial_state], dtype_cases: [{T: bfloat16}], label: kimi-linear-decode-b1} + - {B: 64, S: 1, H: 32, HV: 32, K: 128, V: 128, use_qk_l2norm_in_kernel: true, some: [initial_state], dtype_cases: [{T: bfloat16}], label: kimi-linear-decode-b64} + roofline: + # The per-token recurrence (the signature carries no chunk size). Per token and value + # head: exp(g), the decay, k^T S, v - k^T S and beta, the rank-1 update and the q^T S + # readout; per query head the q scale, shared by its value heads. Flag terms: the q/k + # l2 norms; the in-kernel gate per channel (softplus form 3, sigmoid form 6, the bias + # add 1) and its per-head factor (-exp(A_log) 2, exp(A_log) 1); the beta sigmoid and its + # doubling. + flops: >- + B * S * (HV * (7 * K * V + K + 2 * V) + H * K + + (H * (6 * K + 4) if use_qk_l2norm_in_kernel else 0) + + (HV * K * ((6 if present(lower_bound) else 3) + (1 if present(dt_bias) else 0)) if use_gate_in_kernel else 0) + + (HV * (4 + (1 if allow_neg_eigval else 0)) if use_beta_sigmoid_in_kernel else 0)) + + ((1 if present(lower_bound) else 2) * HV if use_gate_in_kernel else 0) + DeltaNetDecodeFwdOp: # Single-token ungated DeltaNet decode. State ownership is functional: state # is read and new_state is returned. diff --git a/src/tileops/manifest/spec/norm.yaml b/src/tileops/manifest/spec/norm.yaml index f1d19f422..518d1d97f 100644 --- a/src/tileops/manifest/spec/norm.yaml +++ b/src/tileops/manifest/spec/norm.yaml @@ -276,3 +276,37 @@ InstanceNormFwdOp: # Input statistics: mean (1), centered variance (3), normalize the centered value (1); # running statistics: normalize (2). One more per present affine tensor. flops: "((5 if use_input_stats else 2) + (1 if present(weight) else 0) + (1 if present(bias) else 0)) * B * C * prod(L)" + +FusedQKNormRopeFwdOp: + # In place on a packed projection qkv = [q heads | k heads | v heads] of width D each: + # RMSNorm over each q head (q_weight) and each k head (k_weight), then RoPE at + # positions[n] on the first R columns of each q and k head, from cos_sin_cache (cos ‖ sin + # halves). The v columns are untouched. + family: norm + status: spec-only + signature: + forall: {N: Dim, D: Dim, P: Dim, R: Dim, T: "DType[float16 | bfloat16]", C: "DType[float16 | bfloat16 | float32]", seq_lens: "Seq[Int]"} + params: + num_heads: {type: int} + num_kv_heads: {type: int} + eps: {type: float, default: 1.0e-6} + rope_layout: {type: "'neox' | 'interleaved'", default: neox} + inputs: + qkv: {dtype: T, shape: "[N, (num_heads + 2 * num_kv_heads) * D]", mutated: true} + q_weight: {dtype: T, shape: "[D]"} + k_weight: {dtype: T, shape: "[D]"} + cos_sin_cache: {dtype: C, shape: "[P, R]"} + positions: {dtype: int64, shape: "[N]", values: "packed_positions(seq_lens)", requires: ["in_range(0, P)"]} + outputs: {} + shape_rules: + - "num_heads > 0 and num_kv_heads > 0" + - "R > 0 and R % 2 == 0 and R <= D" + workloads: + - {D: 128, P: 40960, R: 128, num_heads: 32, num_kv_heads: 8, seq_lens: [64], dtype_cases: [{T: bfloat16, C: bfloat16}, {T: float16, C: float32}], label: qwen3-8b-t64} + - {D: 128, P: 40960, R: 128, num_heads: 32, num_kv_heads: 8, seq_lens: "repeat(1024, 8)", dtype_cases: [{T: bfloat16, C: bfloat16}], label: qwen3-8b-t8192} + - {D: 128, P: 40960, R: 128, num_heads: 64, num_kv_heads: 4, seq_lens: "repeat(1024, 4)", dtype_cases: [{T: bfloat16, C: bfloat16}], label: qwen3-235b-t4096} + - {D: 128, P: 131072, R: 64, num_heads: 96, num_kv_heads: 8, seq_lens: "repeat(1024, 4)", dtype_cases: [{T: bfloat16, C: float32}], label: glm-4.5-partial-rope} + roofline: + # The q and k columns are read and written, the v columns not touched, and only the + # cos/sin rows the positions name are read. + func: "tileops.perf.formulas.fused_qk_norm_rope_roofline" diff --git a/src/tileops/manifest/spec/quantization.yaml b/src/tileops/manifest/spec/quantization.yaml index f3a00512c..da1d7d8da 100644 --- a/src/tileops/manifest/spec/quantization.yaml +++ b/src/tileops/manifest/spec/quantization.yaml @@ -1,7 +1,8 @@ -# Narrow quantization helper operators. +# Quantization and dequantization operators. # -# These entries describe existing helper surfaces only. They do not claim the -# full INT8/INT4/NF4/FP8 quantization family from the older release plan. +# INT8 is symmetric in [-127, 127]; scales are float32; rounding is half-to-even +# (torch.round); an all-zero group gets scale 1.0, so dequantization never divides +# by zero. A scale is the dequantization multiplier: x ~= q * scale. FP8QuantFwdOp: family: quantization @@ -25,3 +26,199 @@ FP8QuantFwdOp: # Per element: abs, the row-max comparison, the scaling and the clamp; per row: the amax # floor and the scale. flops: "4 * batch * seq_len_kv * kv_group * index_dim + 2 * batch * seq_len_kv * kv_group" + +INT8QuantPerTensorFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + x: {dtype: T, shape: "[M, K]"} + outputs: + q: {dtype: int8, shape: "[M, K]"} + # amax(|x|) / 127 + scale: {dtype: float32, shape: "[1]"} + shape_rules: + - "M > 0 and K > 0" + workloads: + - {M: 4096, K: 4096, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-prefill} + - {M: 32, K: 4096, dtype_cases: [{T: bfloat16}], label: llama-8b-decode} + - {M: 8192, K: 5120, dtype_cases: [{T: bfloat16}, {T: float32}], label: qwen3-32b-prefill} + roofline: + # Per element: abs, the max, the scaling, round and clamp; once: the scale and its zero guard. + flops: "5 * M * K + 2" + +INT8QuantPerChannelFwdOp: + family: quantization + status: spec-only + signature: + forall: {N: Dim, K: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + w: {dtype: T, shape: "[N, K]"} + outputs: + q: {dtype: int8, shape: "[N, K]"} + # One per output channel (row): amax(|w[n, :]|) / 127. + scale: {dtype: float32, shape: "[N]"} + shape_rules: + - "K > 0" + workloads: + - {N: 14336, K: 4096, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-mlp-up} + - {N: 8192, K: 28672, dtype_cases: [{T: bfloat16}], label: llama-70b-mlp-down} + - {N: 5120, K: 8192, dtype_cases: [{T: bfloat16}, {T: float32}], label: qwen3-32b-o-proj} + roofline: + # Per element: abs, the max, the scaling, round and clamp; per row: the scale and its zero guard. + flops: "5 * N * K + 2 * N" + +INT8QuantPerBlockFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + x: {dtype: T, shape: "[M, K]"} + outputs: + q: {dtype: int8, shape: "[M, K]"} + # One per 128 contiguous elements along K; the last block of a row may be partial. + scale: {dtype: float32, shape: "[M, ceil_div(K, 128)]"} + workloads: + - {M: 4096, K: 7168, dtype_cases: [{T: float16}, {T: bfloat16}], label: ds-v3-prefill} + - {M: 64, K: 7168, dtype_cases: [{T: bfloat16}], label: ds-v3-decode} + - {M: 4096, K: 4096, dtype_cases: [{T: bfloat16}, {T: float32}], label: qwen3-235b-prefill} + roofline: + # Per element: abs, the max, the scaling, round and clamp; per block: the scale and its zero guard. + flops: "5 * M * K + 2 * M * ceil_div(K, 128)" + +INT4QuantPerGroupFwdOp: + # Asymmetric INT4 with a scale and a zero point per group of group_size elements along + # K: group_size == K is per channel. The outputs are GemmW4A16FwdOp's weight operands. + family: quantization + status: spec-only + signature: + forall: {N: Dim, K: Dim, T: "DType[float16]"} + params: + group_size: {type: int, default: 128} + inputs: + w: {dtype: T, shape: "[N, K]"} + outputs: + # Two INT4 values per byte, in the order GemmW4A16FwdOp.repack produces. + packed_weight: {dtype: uint8, shape: "[N, K // 2]"} + # A constant group gets a nonzero scale. + weight_scale: {dtype: T, shape: "[N, K // group_size]"} + weight_zero: {dtype: uint8, shape: "[N, K // group_size]"} + shape_rules: + - "group_size > 0" + - "K % 2 == 0 and K % group_size == 0" + workloads: + - {N: 14336, K: 4096, dtype_cases: [{T: float16}], label: llama-8b-mlp-up} + - {N: 14336, K: 4096, group_size: 4096, dtype_cases: [{T: float16}], label: llama-8b-mlp-up-chan} + - {N: 8192, K: 28672, dtype_cases: [{T: float16}], label: llama-70b-mlp-down} + - {N: 5120, K: 8192, dtype_cases: [{T: float16}], label: qwen3-32b-o-proj} + roofline: + # Per element: the min, the max, the scaling, the zero-point add, round and clamp; per + # group: the range, its divide and zero guard, and the zero point (negate, divide, round, clamp). + flops: "6 * N * K + 7 * N * (K // group_size)" + +SmoothQuantFwdOp: + # q = quantize_per_row(x / smooth): per-channel smoothing, then symmetric INT8 per row. + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim, T: "DType[float16 | bfloat16]"} + inputs: + x: {dtype: T, shape: "[M, K]"} + smooth: {dtype: float32, shape: "[K]"} + outputs: + q: {dtype: int8, shape: "[M, K]"} + scale: {dtype: float32, shape: "[M]"} + shape_rules: + - "K > 0" + workloads: + - {M: 4096, K: 4096, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-prefill} + - {M: 32, K: 4096, dtype_cases: [{T: bfloat16}], label: llama-8b-decode} + - {M: 4096, K: 8192, dtype_cases: [{T: bfloat16}], label: llama-70b-prefill} + roofline: + # Per element: the smoothing divide, abs, the max, the scaling, round and clamp; per row: + # the scale and its zero guard. + flops: "6 * M * K + 2 * M" + +INT8DequantPerTensorFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim} + params: + out_dtype: {type: "float16 | bfloat16 | float32"} + inputs: + q: {dtype: int8, shape: "[M, K]"} + scale: {dtype: float32, shape: "[1]"} + outputs: + x: {dtype: out_dtype, shape: "[M, K]"} + workloads: + - {M: 4096, K: 4096, out_dtype: bfloat16, label: llama-8b-prefill} + - {M: 32, K: 4096, out_dtype: bfloat16, label: llama-8b-decode} + - {M: 8192, K: 5120, out_dtype: float16, label: qwen3-32b-prefill} + - {M: 8192, K: 5120, out_dtype: float32, label: qwen3-32b-prefill} + roofline: + # One multiply per element, in float32; the cast is not counted. + flops: "M * K" + +INT8DequantPerChannelFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim} + params: + out_dtype: {type: "float16 | bfloat16 | float32"} + inputs: + q: {dtype: int8, shape: "[M, K]"} + scale: {dtype: float32, shape: "[M]"} + outputs: + x: {dtype: out_dtype, shape: "[M, K]"} + workloads: + - {M: 14336, K: 4096, out_dtype: bfloat16, label: llama-8b-mlp-up} + - {M: 14336, K: 4096, out_dtype: float16, label: llama-8b-mlp-up} + - {M: 8192, K: 28672, out_dtype: bfloat16, label: llama-70b-mlp-down} + roofline: + flops: "M * K" + +INT8DequantPerBlockFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim} + params: + out_dtype: {type: "float16 | bfloat16 | float32"} + inputs: + q: {dtype: int8, shape: "[M, K]"} + # One per 128 contiguous elements along K. + scale: {dtype: float32, shape: "[M, ceil_div(K, 128)]"} + outputs: + x: {dtype: out_dtype, shape: "[M, K]"} + workloads: + - {M: 4096, K: 7168, out_dtype: bfloat16, label: ds-v3-prefill} + - {M: 64, K: 7168, out_dtype: bfloat16, label: ds-v3-decode} + - {M: 4096, K: 4096, out_dtype: float16, label: qwen3-235b-prefill} + roofline: + flops: "M * K" + +FP8QuantPerBlockFwdOp: + # The 128x128 block-scaled FP8 weight format of DeepSeek-V3 checkpoints; scale is the + # dequantization multiplier stored as weight_scale_inv. Edge tiles may be partial. + family: quantization + status: spec-only + signature: + forall: {N: Dim, K: Dim, T: "DType[bfloat16 | float16 | float32]"} + inputs: + w: {dtype: T, shape: "[N, K]"} + outputs: + q: {dtype: float8_e4m3fn, shape: "[N, K]"} + # amax(|tile|) / 448 per 128x128 tile. + scale: {dtype: float32, shape: "[ceil_div(N, 128), ceil_div(K, 128)]"} + workloads: + - {N: 7168, K: 18432, dtype_cases: [{T: bfloat16}], label: ds-v3-mlp-down} + - {N: 4096, K: 7168, dtype_cases: [{T: bfloat16}, {T: float32}], label: ds-v3-expert-gate-up} + - {N: 9216, K: 4096, dtype_cases: [{T: bfloat16}, {T: float16}], label: qwen3-235b-qkv} + roofline: + # Per element: abs, the max, the scaling and the saturating cast; per tile: the scale and + # its zero guard. + flops: "4 * N * K + 2 * ceil_div(N, 128) * ceil_div(K, 128)" diff --git a/src/tileops/manifest/spec/sampling.yaml b/src/tileops/manifest/spec/sampling.yaml new file mode 100644 index 000000000..431fd72cc --- /dev/null +++ b/src/tileops/manifest/spec/sampling.yaml @@ -0,0 +1,167 @@ +# Logit filters and token draws of a decode step. +# +# Logits are [B, V]; a filter returns logits of the same shape and dtype with the removed +# entries set to -inf, so filters compose. Sampling parameters are per-row tensors. The +# random ops take their Philox state as (seed, offset) tensors, so a fixed pair gives a +# fixed result. + +TopKMaskFwdOp: + # Keeps the k[b] largest logits of row b, ties at the threshold kept; a row with + # k[b] >= V is unchanged (FlashInfer top_k_mask_logits). + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim, T: "DType[float16 | bfloat16 | float32]", k_list: "Seq[Int]"} + inputs: + logits: {dtype: T, shape: "[B, V]"} + k: {dtype: int32, shape: "[B]", values: "as_tensor(k_list)", requires: ["in_range(1, 2147483648)"]} + outputs: + masked_logits: {dtype: T, shape: "[B, V]"} + workloads: + - {V: 128256, k_list: "repeat(50, 1)", dtype_cases: [{T: bfloat16}, {T: float32}], label: llama-8b-b1} + - {V: 128256, k_list: "repeat(50, 64)", dtype_cases: [{T: bfloat16}], label: llama-8b-b64} + - {V: 128256, k_list: "repeat(50, 256)", dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-b256} + - {V: 151936, k_list: "repeat(20, 128)", dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} + - {V: 201088, k_list: "repeat(50, 64)", dtype_cases: [{T: bfloat16}], label: gpt-oss-120b-b64} + roofline: + # The threshold search and the mask on each row k restricts; a row with k >= V is the identity. + func: "tileops.perf.formulas.top_k_mask_roofline" + +MinPMaskFwdOp: + # Masks every logit whose probability is below min_p[b] * max(prob[b]), i.e. every + # logit below max_logit + log(min_p[b]). 0 < min_p <= 1 is the caller's obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + logits: {dtype: T, shape: "[B, V]"} + min_p: {dtype: float32, shape: "[B]"} + outputs: + masked_logits: {dtype: T, shape: "[B, V]"} + shape_rules: + - "V > 0" + workloads: + - {B: 1, V: 128256, dtype_cases: [{T: bfloat16}, {T: float32}], label: llama-8b-b1} + - {B: 64, V: 128256, dtype_cases: [{T: bfloat16}], label: llama-8b-b64} + - {B: 256, V: 128256, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-b256} + - {B: 128, V: 151936, dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} + roofline: + # Per element: the row max and the masking compare-select; per row: the threshold (log, add). + flops: "2 * B * V + 2 * B" + +TopPMaskFwdOp: + # Keeps the smallest set of highest-probability tokens whose cumulative probability + # reaches p[b]: a token survives while the exclusive cumulative sum before it is below + # p[b]. 0 < p < 1 is the caller's obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + logits: {dtype: T, shape: "[B, V]"} + p: {dtype: float32, shape: "[B]"} + outputs: + masked_logits: {dtype: T, shape: "[B, V]"} + shape_rules: + - "V > 0" + workloads: + - {B: 1, V: 128256, dtype_cases: [{T: bfloat16}, {T: float32}], label: llama-8b-b1} + - {B: 64, V: 128256, dtype_cases: [{T: bfloat16}], label: llama-8b-b64} + - {B: 256, V: 128256, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-b256} + - {B: 128, V: 151936, dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} + - {B: 64, V: 201088, dtype_cases: [{T: bfloat16}], label: gpt-oss-120b-b64} + roofline: + # Per element: the softmax numerator (max, subtract, exp, sum), a weighted selection of + # the threshold (compare, accumulate) and the mask; per row: p times the sum. + flops: "7 * B * V + B" + +TopKTopPMaskFwdOp: + # TopPMaskFwdOp(TopKMaskFwdOp(logits, k), p): top-p's probabilities are renormalized + # over the tokens top-k keeps. 0 < p < 1 is the caller's obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim, T: "DType[float16 | bfloat16 | float32]", k_list: "Seq[Int]"} + inputs: + logits: {dtype: T, shape: "[B, V]"} + k: {dtype: int32, shape: "[B]", values: "as_tensor(k_list)", requires: ["in_range(1, 2147483648)"]} + p: {dtype: float32, shape: "[B]"} + outputs: + masked_logits: {dtype: T, shape: "[B, V]"} + shape_rules: + - "V > 0" + workloads: + - {V: 128256, k_list: "repeat(50, 1)", dtype_cases: [{T: bfloat16}, {T: float32}], label: llama-8b-b1} + - {V: 128256, k_list: "repeat(50, 64)", dtype_cases: [{T: bfloat16}], label: llama-8b-b64} + - {V: 128256, k_list: "repeat(50, 256)", dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-b256} + - {V: 151936, k_list: "repeat(20, 128)", dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} + roofline: + # Top-k on the rows k restricts, top-p over each row's survivors, the final mask. + func: "tileops.perf.formulas.top_k_top_p_mask_roofline" + +SamplingFromProbsFwdOp: + # Draws one index per row with probability proportional to probs[b]; a zero-weight + # token is never drawn. Each row finite, non-negative and with a positive total is the + # caller's obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim} + inputs: + probs: {dtype: float32, shape: "[B, V]"} + seed: {dtype: int64, shape: "[1]"} + offset: {dtype: int64, shape: "[1]"} + outputs: + samples: {dtype: int32, shape: "[B]"} + workloads: + - {B: 1, V: 128256, label: llama-8b-b1} + - {B: 64, V: 128256, label: llama-8b-b64} + - {B: 256, V: 128256, label: llama-8b-b256} + - {B: 128, V: 151936, label: qwen3-235b-b128} + - {B: 64, V: 201088, label: gpt-oss-120b-b64} + roofline: + # Per element: the row total, then the running sum and its compare against the draw; + # per row: the draw scaled by the total. + flops: "3 * B * V + B" + +ChainSpeculativeSamplingFwdOp: + # Verifies N draft tokens per request. Draft token i is accepted with probability + # min(1, target / draft); at the first rejection a token is drawn from + # normalize(max(0, target - draft)) and the rest of the row is -1; when every draft is + # accepted a bonus token is drawn from target row N. num_accepted is the accepted + # prefix length, the bonus excluded. Every probability row finite, non-negative and + # normalized, and each draft token of positive draft probability, is the caller's + # obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, N: Dim, V: Dim} + inputs: + draft_probs: {dtype: float32, shape: "[B, N, V]"} + draft_token_ids: {dtype: int32, shape: "[B, N]", values: "topk_ids(B, N, V)", requires: ["in_range(0, V)"]} + target_probs: {dtype: float32, shape: "[B, N + 1, V]"} + seed: {dtype: int64, shape: "[1]"} + offset: {dtype: int64, shape: "[1]"} + outputs: + output_token_ids: {dtype: int32, shape: "[B, N + 1]"} + num_accepted: {dtype: int32, shape: "[B]"} + shape_rules: + - "N > 0" + workloads: + - {B: 16, N: 4, V: 128256, label: llama-8b-b16-n4} + - {B: 64, N: 4, V: 128256, label: llama-8b-b64-n4} + - {B: 64, N: 1, V: 129280, label: ds-v3-mtp-n1} + - {B: 64, N: 3, V: 129280, label: ds-v3-mtp-n3} + - {B: 32, N: 5, V: 151936, label: qwen3-235b-eagle-n5} + roofline: + # The work depends on where the chain stops, which the random draws decide, so both + # counts are the cheaper of the two extreme outcomes. All accepted: N ratio tests (a + # multiply and a compare), a draw from target row N (3 per token, 1 per row), reading + # the N token ids, 2 N probabilities and one target row. Rejected at the first: one + # test, the residual max(0, t - d) (2 per token) and a draw from it, reading one id and + # the two rows. + flops: "B * (2 * N + 3 * V + 1 if 2 * N + 3 * V + 1 < 5 * V + 3 else 5 * V + 3)" + bytes: >- + bytes(seed) + bytes(offset) + bytes(output_token_ids) + bytes(num_accepted) + + B * (12 * N + 4 * V if 12 * N + 4 * V < 4 + 8 * V else 4 + 8 * V) diff --git a/src/tileops/ops/_signature_codegen.py b/src/tileops/ops/_signature_codegen.py index bc762e05c..9a42b16c4 100644 --- a/src/tileops/ops/_signature_codegen.py +++ b/src/tileops/ops/_signature_codegen.py @@ -51,6 +51,9 @@ __all__ = ["CheckError", "SignatureCall", "install", "maybe_install_signature"] +_SCHEMA_TYPES = {int: "SymInt", float: "float", bool: "bool", str: "str"} + + def operator_name(family: str, class_name: str) -> str: """``("norm", "RMSNormFwdOp")`` -> ``"norm_rms_norm_fwd"``; a class whose own name already opens with the family, such as ``MoePrePermuteFwdOp``, names it once.""" @@ -1234,9 +1237,6 @@ def binder(self, cls: type): return scope["_call_boundary"] -_SCHEMA_TYPES = {int: "SymInt", float: "float", bool: "bool", str: "str"} - - def _schema_type(cls: type, parameter: inspect.Parameter) -> str: """The operator-schema type of one execution parameter, from its annotation.""" annotation = parameter.annotation diff --git a/src/tileops/ops/elementwise/_base.py b/src/tileops/ops/elementwise/_base.py index 22a7dce0d..3f6ad5db5 100644 --- a/src/tileops/ops/elementwise/_base.py +++ b/src/tileops/ops/elementwise/_base.py @@ -23,6 +23,15 @@ from ..op_base import Op +_MANIFEST_INT_DTYPES = ( + torch.uint8, + torch.int8, + torch.int16, + torch.int32, + torch.int64, +) +_PREDICATE_FALLBACK_DTYPES = _MANIFEST_INT_DTYPES + (torch.bool,) + class _PerDtypeKernels: """The family's one way to reach a kernel: ``self._kernel(inputs, dtype, *dims)``. @@ -333,23 +342,11 @@ def _build_kernel_instance(self, tune, dtype, impl, a_shape, b_shape): return impl(a_shape, b_shape, dtype, tune=tune, alpha=self.alpha) -_MANIFEST_INT_DTYPES = ( - torch.uint8, - torch.int8, - torch.int16, - torch.int32, - torch.int64, -) - - def _int_identity(input: torch.Tensor) -> torch.Tensor: """The default integer answer: the op leaves such a value unchanged.""" return input.clone() -_PREDICATE_FALLBACK_DTYPES = _MANIFEST_INT_DTYPES + (torch.bool,) - - class _IntFallbackCall: """What ``_IntIdentityUnaryOp`` builds for a dtype the shipped kernels do not serve. diff --git a/src/tileops/ops/elementwise/arithmetic.py b/src/tileops/ops/elementwise/arithmetic.py index 635d9c3ff..5e5a1f53c 100644 --- a/src/tileops/ops/elementwise/arithmetic.py +++ b/src/tileops/ops/elementwise/arithmetic.py @@ -24,6 +24,8 @@ from ..op_base import Op from ._base import BinaryOp, _AlphaScaledBinaryOp, _PerDtypeKernels +_DIV_KEY_BY_ROUNDING_MODE = {None: "div", "trunc": "div_trunc", "floor": "floor_divide"} + class AddFwdOp(_AlphaScaledBinaryOp): """Element-wise addition with broadcast: y = input + alpha * other. @@ -53,9 +55,6 @@ class MulFwdOp(BinaryOp): kernel_types = {"mul": MulFwdKernel} -_DIV_KEY_BY_ROUNDING_MODE = {None: "div", "trunc": "div_trunc", "floor": "floor_divide"} - - class DivFwdOp(BinaryOp): """Element-wise division with broadcast: y = input / other. diff --git a/src/tileops/ops/op_base.py b/src/tileops/ops/op_base.py index 44a40bb59..993d4bcfd 100644 --- a/src/tileops/ops/op_base.py +++ b/src/tileops/ops/op_base.py @@ -38,6 +38,11 @@ _Entry = TypeVar("_Entry") +# Every dispatch key a created op class declares in ``kernel_types``. Constructing an op imports +# it and every sub-op it builds, so every key that can replace something in that op is here. +_DISPATCH_KEYS: set[str] = set() + + class _Unresolved: """The type of :data:`_UNRESOLVED`, so a traceback says what it is.""" @@ -52,10 +57,6 @@ def __repr__(self) -> str: _UNRESOLVED = _Unresolved() -# Every dispatch key a created op class declares in ``kernel_types``. Constructing an op imports -# it and every sub-op it builds, so every key that can replace something in that op is here. -_DISPATCH_KEYS: set[str] = set() - # The calls in progress on this thread, innermost last, each with the checked calls completed # inside it: what a composite's call collects from its sub-ops. _OPEN_CALLS = threading.local() diff --git a/src/tileops/ops/pool.py b/src/tileops/ops/pool.py index 7cd08b7bf..f8e3c56df 100644 --- a/src/tileops/ops/pool.py +++ b/src/tileops/ops/pool.py @@ -45,6 +45,17 @@ ] +# Per-axis name suffixes, indexed by spatial dimensionality. +_POOL_DIM_NAMES: Dict[int, Tuple[str, ...]] = {1: ("l",), 2: ("h", "w"), 3: ("d", "h", "w")} +# Kernel-kwarg suffixes for kernel_size/stride/padding(/dilation). +# Why: the 1d max-pool kernels name their pooling axis `w`, not `l`. +_MAX_POOL_PARAM_SUFFIXES: Dict[int, Tuple[str, ...]] = { + 1: ("w",), + 2: ("h", "w"), + 3: ("d", "h", "w"), +} + + def _per_axis(value: "int | Sequence[int]", ndim: int) -> tuple[int, ...]: """A pooling parameter as one value per spatial axis, as ``per_axis`` reads it.""" return (value,) * ndim if isinstance(value, int) else tuple(value) @@ -292,17 +303,6 @@ def _check_ragged( raise ValueError("indices must name each chunk offsets implies exactly once") -# Per-axis name suffixes, indexed by spatial dimensionality. -_POOL_DIM_NAMES: Dict[int, Tuple[str, ...]] = {1: ("l",), 2: ("h", "w"), 3: ("d", "h", "w")} -# Kernel-kwarg suffixes for kernel_size/stride/padding(/dilation). -# Why: the 1d max-pool kernels name their pooling axis `w`, not `l`. -_MAX_POOL_PARAM_SUFFIXES: Dict[int, Tuple[str, ...]] = { - 1: ("w",), - 2: ("h", "w"), - 3: ("d", "h", "w"), -} - - class _AvgPoolFwdOpBase(Op): """Generic average-pooling forward, parametrized by class-attribute ``ndim``. diff --git a/src/tileops/perf/formulas.py b/src/tileops/perf/formulas.py index 3c4800fad..c807c0a91 100644 --- a/src/tileops/perf/formulas.py +++ b/src/tileops/perf/formulas.py @@ -13,6 +13,8 @@ from math import prod from typing import TYPE_CHECKING +from tileops.manifest.dtype_rules import FLOAT8_DTYPES + if TYPE_CHECKING: from tileops.manifest.workload import CallView @@ -22,11 +24,13 @@ "conv_roofline", "dsa_decode_roofline", "dsa_distinct_kv_rows", + "dsa_paged_fwd_roofline", "dsa_selected_keys", "fft_c2c_roofline", "fp8_lightning_indexer_roofline", "fused_moe_fwd_roofline", "fused_moe_shared_expert_fwd_roofline", + "fused_qk_norm_rope_roofline", "gqa_dense_fwd_roofline", "gqa_paged_cache_rows", "gqa_paged_fwd_roofline", @@ -34,6 +38,9 @@ "gqa_prefill_paged_with_kv_cache_fwd_roofline", "gqa_varlen_fwd_roofline", "lightning_indexer_scored_keys", + "mla_kv_cache_write_roofline", + "mla_paged_fwd_roofline", + "mla_varlen_fwd_roofline", "moe_expert_mlp_roofline", "moe_grouped_gemm_roofline", "moe_layout_active_experts", @@ -47,8 +54,12 @@ "nsa_topk_varlen_roofline", "paged_decode_cache_rows", "paged_decode_roofline", + "paged_kv_cache_gather_roofline", + "paged_kv_cache_write_roofline", "paged_rows", "pool_roofline", + "top_k_mask_roofline", + "top_k_top_p_mask_roofline", "topk_selector_roofline", "topk_selector_window_scores", "visible_score_rows", @@ -56,6 +67,17 @@ ] +# Per gated element: the activation of the gate (silu: 5; the erf gelu: 5) and the multiply +# by the up projection. +_GATED_ACTIVATION = 6 +# Per score: the scale, the running max, the subtraction, the exp and the sum of a softmax. +_SOFTMAX_PER_SCORE = 5 +# Per score: the divide, tanh and multiply of a logit softcap. +_SOFTCAP_PER_SCORE = 3 +# Per head score: the relu, the weight multiply and the add into the sum over heads. +_INDEXER_EPILOGUE_PER_SCORE = 3 + + def _distribute_total(total: int, batch: int, max_len: int) -> list[int]: lengths = [0] * batch remaining = total @@ -88,11 +110,6 @@ def _expert_weight_bytes(call) -> int: return (call.bytes("w_gate_up") + call.bytes("w_down")) // experts -# Per gated element: the activation of the gate (silu: 5; the erf gelu: 5) and the multiply -# by the up projection. -_GATED_ACTIVATION = 6 - - def _routing_flops(call) -> int: """Routing as FusedTopKFwdOp prices it, per token: scoring (sigmoid 4 per logit; softmax 3, plus the row sum and a divide per kept weight unless renormalizing), top_k @@ -397,12 +414,6 @@ def visible_scores(q_len: int, kv_len: int, is_causal: bool, left: int, right: i return visible_score_rows(q_len, kv_len, is_causal, left, right)[0] -# Per score: the scale, the running max, the subtraction, the exp and the sum of a softmax. -_SOFTMAX_PER_SCORE = 5 -# Per score: the divide, tanh and multiply of a logit softcap. -_SOFTCAP_PER_SCORE = 3 - - def attention_flops( heads: int, scores: int, rows: int, qk_dim: int, v_dim: int, softcap: bool = False ) -> int: @@ -722,10 +733,6 @@ def lightning_indexer_scored_keys(call: "CallView") -> int: ) -# Per head score: the relu, the weight multiply and the add into the sum over heads. -_INDEXER_EPILOGUE_PER_SCORE = 3 - - def fp8_lightning_indexer_roofline(call: "CallView") -> tuple[int, int]: """Lightning indexer: per query head and windowed key, a D-long contraction and the relu, weight and head-sum epilogue; each tensor moves once, ``logits`` written whole.""" @@ -759,3 +766,179 @@ def gqa_prefill_paged_cache_rows(call: "CallView") -> int: # A request with no new token attends nothing and reads none of its cache. read_lens = [c if q else 0 for q, c in zip(q_lens, cache_lens, strict=True)] return paged_rows(call.values("block_table"), read_lens, call.ix["page_size"])[0] + + +# ---------------------------------------------------------------- paged caches and MLA + + +def _elem_bytes(call: "CallView", name: str) -> int: + """Bytes of one element of tensor *name*.""" + return call.bytes(name) // max(1, prod(call.tensors[name][0])) + + +def _paged_cache_read_bytes( + call: "CallView", cache: str, table: str, ends: list, starts: "list | None" = None +) -> int: + """Bytes of the distinct rows of the paged *cache* ``[NP, PS, ...]`` that the requests' + ranges ``[starts[b], ends[b])`` reach through *table*, plus the table entries consulted.""" + shape = call.tensors[cache][0] + rows, consulted = paged_rows(call.values(table), ends, shape[1], starts) + row_bytes = call.bytes(cache) // max(1, shape[0] * shape[1]) + return rows * row_bytes + consulted * _elem_bytes(call, table) + + +def mla_paged_fwd_roofline(call: "CallView") -> tuple[int, int]: + """Paged MLA decode: QK over the whole latent row, PV over its first ``kv_lora_rank`` + columns, for the scores each query row sees in its request's cache. + + An FP8 cache's per-tensor scale folds into the score scale and, once per query row and + head, into the softmax denominator. The cache is read at the rows the block table reaches; + every other tensor moves once. + """ + ix = call.ix + lengths = call.values("cache_seqlens") + pairs = [visible_score_rows(ix["S_q"], c, ix["is_causal"], -1, -1) for c in lengths] + scores, rows = sum(p[0] for p in pairs), sum(p[1] for p in pairs) + flops = attention_flops(ix["H"], scores, rows, ix["DK"], ix["kv_lora_rank"]) + if call.tensors["kv_cache"][1] in FLOAT8_DTYPES: + flops += ix["H"] * rows + moved = _derived_bytes(call) - call.bytes("kv_cache") - call.bytes("block_table") + moved += _paged_cache_read_bytes(call, "kv_cache", "block_table", lengths) + return flops, moved + + +def mla_varlen_fwd_roofline(call: "CallView") -> tuple[int, int]: + """Packed MLA prefill: QK over ``DN + PE``, PV over ``DV``, for the scores each request + sees under its mask; every tensor moves once.""" + ix = call.ix + pairs = [ + visible_score_rows(n, n, ix["is_causal"], -1, -1) for n in _segments(call, "cu_seqlens") + ] + scores, rows = sum(p[0] for p in pairs), sum(p[1] for p in pairs) + flops = attention_flops(ix["H"], scores, rows, ix["DN"] + ix["PE"], ix["DV"]) + return flops, _derived_bytes(call) + + +# FlashMLA DeepSeek-V3.2 FP8 cache row: the QK width and the value (latent) width. +_DSA_QK_DIM, _DSA_V_DIM = 576, 512 + + +def dsa_paged_fwd_roofline(call: "CallView") -> tuple[int, int]: + """Paged DeepSeek sparse attention over an FP8 latent cache. + + Every index slot ``0 <= j < cache_seqlens[b]`` is a score, a repeated slot once per + occurrence. The latent of each distinct cache row is dequantized once (a multiply per + value) and read once; the block table is read at the pages the slots resolve to; ``q`` is + read at the query rows with a score; every other tensor moves once. + """ + ix = call.ix + lengths, table = call.values("cache_seqlens"), call.values("block_table") + page_size = call.tensors["kv_cache"][0][1] + scores = rows = 0 + distinct: set = set() + pages: set = set() + for b, per_query in enumerate(call.values("indices")): + for slots in per_query: + valid = [j for j in slots if 0 <= j < lengths[b]] + scores += len(valid) + rows += bool(valid) + for j in valid: + pages.add((b, j // page_size)) + distinct.add((table[b][j // page_size], j % page_size)) + flops = attention_flops(ix["H"], scores, rows, _DSA_QK_DIM, _DSA_V_DIM) + flops += _DSA_V_DIM * len(distinct) + shape = call.tensors["kv_cache"][0] + moved = _derived_bytes(call) - call.bytes("q") - call.bytes("kv_cache") + moved -= call.bytes("block_table") + moved += rows * ix["H"] * _DSA_QK_DIM * _elem_bytes(call, "q") + moved += len(distinct) * shape[2] + len(pages) * _elem_bytes(call, "block_table") + return flops, moved + + +def _written_slots(call: "CallView") -> int: + """Tokens a cache write stores: the ``slot_mapping`` entries other than -1.""" + return sum(1 for s in call.values("slot_mapping") if s >= 0) + + +def paged_kv_cache_write_roofline(call: "CallView") -> tuple[int, int]: + """Paged K/V cache write: each token with a slot reads its k and v rows and writes them + into both caches; an FP8 cache scales and saturates each value (2 per value).""" + shape = call.tensors["k"][0] + tokens, row = _written_slots(call), prod(shape[1:]) + flops = 4 * tokens * row if call.tensors["k_pages"][1] in FLOAT8_DTYPES else 0 + moved = 2 * tokens * row * (_elem_bytes(call, "k") + _elem_bytes(call, "k_pages")) + moved += call.bytes("slot_mapping") + if tokens * row: + moved += sum(call.bytes(s) for s in ("k_scale", "v_scale") if call.present(s)) + return flops, moved + + +def mla_kv_cache_write_roofline(call: "CallView") -> tuple[int, int]: + """MLA latent-cache write: each token with a slot reads kv_c and k_pe and writes the + concatenated row; an FP8 cache scales and saturates each value (2 per value), and a fused + RoPE rotates k_pe (3 per value) from the cos/sin rows the tokens' positions name.""" + ix = call.ix + width, pe = ix["DC"] + ix["PE"], ix["PE"] + slots = call.values("slot_mapping") + tokens = _written_slots(call) + flops = 2 * tokens * width if call.tensors["kv_cache"][1] in FLOAT8_DTYPES else 0 + moved = tokens * width * (_elem_bytes(call, "kv_c") + _elem_bytes(call, "kv_cache")) + moved += call.bytes("slot_mapping") + moved += call.bytes("scale") if tokens * width and call.present("scale") else 0 + if ix["fuse_rope"]: + flops += 3 * tokens * pe + positions = [p for p, s in zip(call.values("positions"), slots, strict=True) if s >= 0] + moved += len(positions) * _elem_bytes(call, "positions") + moved += len(set(positions)) * pe * _elem_bytes(call, "cos_sin_cache") + return flops, moved + + +def paged_kv_cache_gather_roofline(call: "CallView") -> tuple[int, int]: + """Paged cache gather: the cache rows each request's range reaches are read once and + written to ``dst``; an FP8 cache is dequantized with a multiply per value.""" + lengths = _segments(call, "cu_seq_lens") + starts = call.values("seq_starts") if call.present("seq_starts") else [0] * len(lengths) + ends = [s + n for s, n in zip(starts, lengths, strict=True)] + fp8 = call.tensors["cache"][1] in FLOAT8_DTYPES + flops = prod(call.tensors["dst"][0]) if fp8 else 0 + moved = _derived_bytes(call) - call.bytes("cache") - call.bytes("block_table") + moved += _paged_cache_read_bytes(call, "cache", "block_table", ends, starts) + if call.present("scale") and not prod(call.tensors["dst"][0]): + moved -= call.bytes("scale") # nothing is dequantized + return flops, moved + + +def fused_qk_norm_rope_roofline(call: "CallView") -> tuple[int, int]: + """Q/K RMSNorm and RoPE in place: per q and k value, RMSNorm with its weight (4), and 3 + per rotated value. The q and k columns are read and written, the v columns untouched, + and the cos/sin rows the positions name read once.""" + ix = call.ix + heads = ix["num_heads"] + ix["num_kv_heads"] + values = ix["N"] * heads * ix["D"] + flops = ix["N"] * heads * (4 * ix["D"] + 3 * ix["R"]) + moved = 2 * values * _elem_bytes(call, "qkv") + moved += sum(call.bytes(t) for t in ("q_weight", "k_weight", "positions")) + moved += len(set(call.values("positions"))) * ix["R"] * _elem_bytes(call, "cos_sin_cache") + return flops, moved + + +# ---------------------------------------------------------------- sampling + + +def top_k_mask_roofline(call: "CallView") -> tuple[int, int]: + """Top-k logit mask: a threshold selection (1 per logit) and the mask (1 per logit) on + each row ``k`` restricts; a row with ``k >= V`` is the identity. Every tensor moves once.""" + vocab = call.ix["V"] + return 2 * vocab * sum(1 for k in call.values("k") if k < vocab), _derived_bytes(call) + + +def top_k_top_p_mask_roofline(call: "CallView") -> tuple[int, int]: + """Top-k then top-p logit mask, per row: the top-k threshold (1 per logit) where ``k`` + restricts the row; over the ``min(k, V)`` survivors the softmax numerator (max, + subtract, exp, sum) and the weighted threshold selection (compare, accumulate); ``p`` + times the sum; the final mask (1 per logit). Every tensor moves once.""" + vocab = call.ix["V"] + flops = sum( + (vocab if k < vocab else 0) + 6 * min(k, vocab) + 1 + vocab for k in call.values("k") + ) + return flops, _derived_bytes(call) diff --git a/src/tileops/perf/profile.py b/src/tileops/perf/profile.py index 8c3e7766f..b8d919fd7 100644 --- a/src/tileops/perf/profile.py +++ b/src/tileops/perf/profile.py @@ -20,6 +20,18 @@ ) +# Tensor-core dtype keys, by the dtype the contraction consumes. fp32 maps to +# tf32 because that is the unit an fp32 contraction runs on when tensor cores +# serve it. Encode side of the roof-key format; ``resolve_roof`` is the decode. +_TENSOR_CORE_DTYPE_KEYS = { + "float16": "fp16", + "bfloat16": "bf16", + "float32": "tf32", + "float8_e4m3fn": "fp8", + "float8_e5m2": "fp8", +} + + def get_profile_path(gpu_name: str) -> Path: """Return the path to a GPU profile YAML. @@ -71,18 +83,6 @@ def _inject_effective(profile): section["effective"] = section["theoretical"] * section["calibration"] -# Tensor-core dtype keys, by the dtype the contraction consumes. fp32 maps to -# tf32 because that is the unit an fp32 contraction runs on when tensor cores -# serve it. Encode side of the roof-key format; ``resolve_roof`` is the decode. -_TENSOR_CORE_DTYPE_KEYS = { - "float16": "fp16", - "bfloat16": "bf16", - "float32": "tf32", - "float8_e4m3fn": "fp8", - "float8_e5m2": "fp8", -} - - def tensor_core_roof(dtype) -> str: """Tensor-core roof key for a contraction computing at *dtype*. diff --git a/src/tileops/sampling.py b/src/tileops/sampling.py new file mode 100644 index 000000000..2803fe20e --- /dev/null +++ b/src/tileops/sampling.py @@ -0,0 +1,3 @@ +"""The sampling ops, at the public path ``tileops.sampling``.""" + +__all__: list[str] = [] diff --git a/src/tileops/trace/record.py b/src/tileops/trace/record.py index 77949364f..34f9ca8e2 100644 --- a/src/tileops/trace/record.py +++ b/src/tileops/trace/record.py @@ -21,40 +21,36 @@ ] -class EventKind(IntEnum): - """Event kind packed into ``w1`` bits 24..27.""" - - RANGE_BEGIN = 0 - RANGE_END = 1 - INSTANT = 2 - # Reserved: ``trace.dag`` is now a build-time declaration (no runtime record), - # so no DAG record is ever emitted. Kept to keep the enum value stable. - DAG = 3 - - # Per-slot record capacity (config default; callers may override). MAX_EVENTS_DEFAULT = 768 - # Field widths and bit offsets within w1. _EVENT_ID_BITS = 24 _KIND_BITS = 4 _LANE_BITS = 4 _PAYLOAD_BITS = 32 - _EVENT_ID_SHIFT = 0 _KIND_SHIFT = 24 _LANE_SHIFT = 28 _PAYLOAD_SHIFT = 32 - _EVENT_ID_MASK = (1 << _EVENT_ID_BITS) - 1 _KIND_MASK = (1 << _KIND_BITS) - 1 _LANE_MASK = (1 << _LANE_BITS) - 1 _PAYLOAD_MASK = (1 << _PAYLOAD_BITS) - 1 - # Max distinct lanes that fit the 4-bit lane field. MAX_LANES = 1 << _LANE_BITS +class EventKind(IntEnum): + """Event kind packed into ``w1`` bits 24..27.""" + + RANGE_BEGIN = 0 + RANGE_END = 1 + INSTANT = 2 + # Reserved: ``trace.dag`` is now a build-time declaration (no runtime record), + # so no DAG record is ever emitted. Kept to keep the enum value stable. + DAG = 3 + + def pack_w1(event_id: int, kind: int, lane: int, payload: int) -> int: """Pack the four ``w1`` fields into a single unsigned 64-bit word. diff --git a/src/tileops/trace/ui.py b/src/tileops/trace/ui.py index d04cf3fd5..6a2c6b688 100644 --- a/src/tileops/trace/ui.py +++ b/src/tileops/trace/ui.py @@ -74,6 +74,59 @@ _INK = "#191a16" +# Plotly config: horizontal-only zoom + pan. With yaxis.fixedrange set, scrollZoom +# stretches only x and pan moves only x; the vertical / box / autoscale buttons are +# stripped so the y axis can never be rescaled. +_CONFIG = { + "scrollZoom": True, + "displaylogo": False, + "responsive": True, + "modeBarButtonsToRemove": [ + "zoom2d", + "select2d", + "lasso2d", + "zoomIn2d", + "zoomOut2d", + "autoScale2d", + ], +} +_PLOTLY_CDN = "https://cdn.plot.ly/plotly-2.35.2.min.js" +_HTML_TEMPLATE = """
+ +