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 = """ + +{title} + + +
{tab_buttons}
+
+""" + + def _lane_label(gid: int, lane: int, group_id_to_name: dict, lane_id_to_name: dict) -> str: """Build a lane's y-axis label ``" / "``. @@ -348,61 +401,6 @@ def _figure_for_cta( return {"data": data, "layout": layout} -# 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 = """ - -{title} - - -
{tab_buttons}
-
-""" - - def export_timeline_html( events: list, path: str, diff --git a/src/tileops/utils/utils.py b/src/tileops/utils/utils.py index f4d585ef9..39d25c1f6 100644 --- a/src/tileops/utils/utils.py +++ b/src/tileops/utils/utils.py @@ -21,6 +21,17 @@ # `get_device_name` string scan would run on every forward. +# Spin cycles queued before a device_busy_of measurement: tens of milliseconds +# on any supported clock, ample to enqueue every timed call first. +_BUSY_TIMING_SPIN_CYCLES = 50_000_000 + + +# Calibrated boards: the key selection tables use -> the name fragment CUDA reports. +# All SKUs of a board share its key. GPU profiles match the full name instead +# (:func:`tileops.perf.find_profile`): a speed-of-light reading is not shared. +_CALIBRATION_BOARDS = {"h200": "H200"} + + @functools.lru_cache(maxsize=16) def _device_name(index: int) -> str: return torch.cuda.get_device_name(index) @@ -32,12 +43,6 @@ def _sm_version(index: int) -> int: return major * 10 + minor -# Calibrated boards: the key selection tables use -> the name fragment CUDA reports. -# All SKUs of a board share its key. GPU profiles match the full name instead -# (:func:`tileops.perf.find_profile`): a speed-of-light reading is not shared. -_CALIBRATION_BOARDS = {"h200": "H200"} - - def calibration_key(device_name: str) -> "str | None": """The key of the calibrated board *device_name* belongs to, or ``None``. @@ -102,11 +107,6 @@ def forget_device_properties() -> None: _device_facts.cache_clear() -# Spin cycles queued before a device_busy_of measurement: tens of milliseconds -# on any supported clock, ample to enqueue every timed call first. -_BUSY_TIMING_SPIN_CYCLES = 50_000_000 - - def device_busy_of(call, device: "torch.device", warmup: int = 5, rep: int = 20) -> float: """Mean device time of *call* in milliseconds with host gaps excluded. diff --git a/tests/roofline_binder.py b/tests/roofline_binder.py index 5a3574473..adfe37919 100644 --- a/tests/roofline_binder.py +++ b/tests/roofline_binder.py @@ -1,6 +1,8 @@ """Build a bytes-oracle case for each workload row of an op from its manifest entry alone. Each row is instantiated, the op constructed from it and its call checked on meta tensors. +An implemented entry's op is its class; a spec-only entry's is a class carrying only what its +signature generates, since the recount needs no implementation. In parallel the oracle counts the traffic the checked call implies -- one read per input it binds, one write per output, both for a written input -- and the caller requires the two to be equal. The `roofline` block is never read. @@ -19,15 +21,36 @@ from tileops.manifest.plan import entry_plan from tileops.manifest.registry import op_class from tileops.manifest.workload import instantiate +from tileops.ops._signature_codegen import install +from tileops.ops.op_base import Op -__all__ = ["manifest_cases"] +__all__ = ["manifest_cases", "signature_class"] + + +def signature_class(op_name: str, entry: dict) -> type: + """An `Op` subclass with *entry*'s generated methods and no kernel.""" + + def construct(self, **params): + vars(self).update(params) + self.dispatch_kernel(None) + + body = {"__init__": construct, "default_kernel_map": property(lambda self: {})} + body["forward"] = body["_eager_forward"] = lambda self, *args: None + cls = type(f"Signature{op_name}", (Op,), body) + if not install(cls, entry): + raise ValueError(f"{op_name}: the signature does not generate") + return cls def manifest_cases(op_name: str): """Yield ``(label, dtype case, op, oracle bytes, oracle read bytes)`` per row and dtype case.""" entry = load_manifest()[op_name] plan = entry_plan(op_name, entry, load_adts()) - cls = op_class(op_name, entry) + cls = ( + op_class(op_name, entry) + if entry["status"] == "implemented" + else signature_class(op_name, entry) + ) for row in entry["workloads"]: for case in row.get("dtype_cases") or [{}]: call = instantiate(plan, row, case) diff --git a/tests/test_manifest_generators.py b/tests/test_manifest_generators.py index 1994bbe2e..09e8fef38 100644 --- a/tests/test_manifest_generators.py +++ b/tests/test_manifest_generators.py @@ -102,3 +102,19 @@ def test_sample_indices_draws_distinct_values_in_range(): assert sorted(GENERATORS["sample_indices"](random.Random(0), 5, 5)) == list(range(5)) with pytest.raises(ValueError): GENERATORS["sample_indices"](random.Random(0), 3, 2) + + +def test_sparse_topk_positions_draw_visible_positions_and_pad(): + # Request lengths 3 and 5, two queries: query s of a length-n request sees n - 1 + s positions. + rows = GENERATORS["attn.sparse_topk_positions"](random.Random(0), [3, 5], 2, 4) + assert (len(rows), len(rows[0]), len(rows[0][0])) == (2, 2, 4) + for n, per_query in zip([3, 5], rows, strict=True): + for s, slots in enumerate(per_query): + visible = n - 1 + s + picked = [j for j in slots if j != -1] + assert len(picked) == len(set(picked)) == min(4, visible) + assert all(0 <= j < visible for j in picked) and slots[len(picked) :] == [-1] * ( + 4 - len(picked) + ) + with pytest.raises(ValueError): + GENERATORS["attn.sparse_topk_positions"](random.Random(0), [1], 2, 4) diff --git a/tests/test_public_api.py b/tests/test_public_api.py index f19f01797..4e6fc01dd 100644 --- a/tests/test_public_api.py +++ b/tests/test_public_api.py @@ -112,10 +112,13 @@ def test_each_public_name_has_exactly_one_family(): @pytest.mark.smoke def test_public_surface_is_the_manifest(): - """An op reachable as `tileops..` has a manifest entry, and every entry - is reachable. Abstract bases are not public at all.""" + """An op reachable as `tileops..` has a manifest entry, and every + implemented entry is reachable; a spec-only entry may have no class yet. Abstract + bases are not public at all.""" public = {name for family in FAMILIES for name in _family_module(family).__all__} - assert public == set(load_manifest()) + manifest = load_manifest() + assert public <= set(manifest) + assert {name for name, entry in manifest.items() if entry["status"] == "implemented"} <= public @pytest.mark.smoke diff --git a/tests/test_roofline_oracle.py b/tests/test_roofline_oracle.py index 3a718ffd9..72aaed731 100644 --- a/tests/test_roofline_oracle.py +++ b/tests/test_roofline_oracle.py @@ -521,7 +521,268 @@ def test_grouped_gemm_does_not_charge_the_padding_offsets_it_ignores(self): assert self._priced(GroupedGemmFwdOp(), tensors)[1] == oracle -# Coverage levels. Every implemented op sits at +def _evaluated(op_name: str, row: dict, case: dict, **values): + """``(flops, bytes)`` the generated evaluator prices for a row of a spec-only entry, and the + call; *values* replace the named metadata tensors' generated contents.""" + import dataclasses + + from tests.roofline_binder import signature_class + from tileops.manifest import load_adts, load_manifest + from tileops.manifest.plan import entry_plan + from tileops.manifest.workload import instantiate + + entry = load_manifest()[op_name] + plan = entry_plan(op_name, entry, load_adts()) + call = instantiate(plan, {**row, "label": "recount"}, case) + specs = { + **call.specs, + **{n: dataclasses.replace(call.specs[n], values=v) for n, v in values.items()}, + } + call = dataclasses.replace(call, specs=specs) + tensors = call.materialize("meta") + cls = signature_class(op_name, entry) + op = cls(**call.arguments(tensors)) + checked = cls._signature.check(op, {t: tensors[t] for t in plan.sig.inputs}) + metadata = {n: torch.tensor(call.values(n)) for n in checked.metadata} + op._signature_call = dataclasses.replace(checked, metadata=metadata) + return op.eval_roofline(), call + + +def _attention_flops(heads, qk, v, scores, rows): + # Per score two contractions and the softmax (5); per output element its divide. + return heads * (scores * (2 * qk + 2 * v + 5) + rows * v) + + +_BF16, _F32, _FP8, _I32, _I64 = ( + torch.bfloat16, + torch.float32, + torch.float8_e4m3fn, + torch.int32, + torch.int64, +) + + +class TestSpecOnlyRecounts: + """Spec-only entries whose traffic or arithmetic follows their metadata values, recounted + by walking the call. Each case picks metadata that reaches the branches the values decide.""" + + @pytest.mark.parametrize( + "row,case", + [ + # Two requests with disjoint pages; two causal query rows each. + ( + {"S_q": 2, "NP": 8, "W": 4, "cache_lens": [5, 9]}, + {"T": "bfloat16", "KV": "bfloat16"}, + ), + # A pool smaller than the requests' pages, so requests share rows; an FP8 cache. + ( + {"S_q": 1, "NP": 3, "W": 3, "cache_lens": [5, 9, 12], "some": ["kv_scale"]}, + {"T": "bfloat16", "KV": "float8_e4m3fn"}, + ), + ], + ) + def test_mla_paged_reads_the_rows_its_block_table_reaches(self, row, case): + name = "MultiHeadLatentAttentionPagedFwdOp" + row = {"H": 3, "DK": 12, "PS": 4, "kv_lora_rank": 8, **row} + (flops, moved), call = _evaluated(name, row, case) + heads, dk, rank, page, s_q = row["H"], row["DK"], row["kv_lora_rank"], row["PS"], row["S_q"] + table, lengths = call.values("block_table"), call.values("cache_seqlens") + fp8 = case["KV"] == "float8_e4m3fn" + scores = rows = 0 + for c in lengths: + for i in range(s_q): + seen = c - s_q + i + 1 + scores, rows = scores + seen, rows + 1 + cache_rows = { + (table[b][j // page], j % page) for b, c in enumerate(lengths) for j in range(c) + } + pages = {(b, j // page) for b, c in enumerate(lengths) for j in range(c)} + batch = len(lengths) + assert moved == _ledger( + name, + q=((batch, s_q, heads, dk), _BF16), + kv_cache=((len(cache_rows), dk), _FP8 if fp8 else _BF16), + block_table=((len(pages),), _I32), + cache_seqlens=((batch,), _I32), + kv_scale=((1,), _F32) if fp8 else None, + o=((batch, s_q, heads, rank), _BF16), + lse=((batch, s_q, heads), _F32), + ) + assert flops == _attention_flops(heads, dk, rank, scores, rows) + ( + heads * rows if fp8 else 0 + ) + + def test_dsa_paged_scores_each_valid_slot_and_reads_the_rows_they_name(self): + name = "DeepSeekSparseAttentionPagedFwdOp" + row = {"S_q": 2, "H": 2, "K": 4, "NP": 6, "PS": 4, "W": 3, "cache_lens": [5, 10]} + # Request 0: query 0 selects nothing, query 1 three positions on two pages. + # Request 1: a repeated slot, then nothing. + indices = [[[-1, -1, -1, -1], [0, 4, 3, -1]], [[3, 3, -1, -1], [-1, -1, -1, -1]]] + (flops, moved), call = _evaluated(name, row, {}, indices=indices) + table, lengths, page = call.values("block_table"), call.values("cache_seqlens"), row["PS"] + valid = [ + [[j for j in slots if 0 <= j < lengths[b]] for slots in per_q] + for b, per_q in enumerate(indices) + ] + scores = sum(len(v) for per_q in valid for v in per_q) + rows = sum(1 for per_q in valid for v in per_q if v) + cache_rows = { + (table[b][j // page], j % page) + for b, per_q in enumerate(valid) + for v in per_q + for j in v + } + pages = {(b, j // page) for b, per_q in enumerate(valid) for v in per_q for j in v} + assert moved == _ledger( + name, + q=((rows, row["H"], 576), _BF16), # the query rows that score something + kv_cache=((len(cache_rows), 656), torch.uint8), + block_table=((len(pages),), _I32), + cache_seqlens=((2,), _I32), + indices=((2, 2, 4), _I32), + o=((2, 2, row["H"], 512), _BF16), + lse=((2, 2, row["H"]), _F32), + ) + # Each distinct row's 512 latent values are dequantized once. + assert flops == _attention_flops(row["H"], 576, 512, scores, rows) + 512 * len(cache_rows) + + def test_paged_kv_cache_write_moves_only_the_tokens_with_a_slot(self): + name = "PagedKVCacheWriteFwdOp" + row = {"N": 5, "H_kv": 2, "D": 4, "NP": 4, "PS": 3, "some": ["k_scale"]} + slots = [7, -1, 2, -1, 11] + (flops, moved), _call = _evaluated( + name, row, {"T": "bfloat16", "KV": "float8_e4m3fn"}, slot_mapping=slots + ) + written = ((3, 2, 4), _FP8) + assert moved == _ledger( + name, + k=((3, 2, 4), _BF16), + v=((3, 2, 4), _BF16), + k_pages_unread=True, + k_pages_write=written, + v_pages_unread=True, + v_pages_write=written, + slot_mapping=((5,), _I64), + k_scale=((1,), _F32), + v_scale=((1,), _F32), + ) + # Per stored FP8 value, the scale and the saturating cast. + assert flops == 2 * 2 * 3 * 2 * 4 + + def test_mla_kv_cache_write_rotates_and_moves_only_the_tokens_with_a_slot(self): + name = "MultiHeadLatentAttentionKVCacheWriteFwdOp" + row = { + "DC": 6, + "PE": 4, + "NP": 4, + "PS": 3, + "P": 16, + "seq_lens": [3, 2], + "fuse_rope": True, + "some": ["scale"], + } + slots = [5, -1, 0, 8, -1] # tokens 0, 2 and 3, at positions 0, 2 and 0 + case = {"T": "bfloat16", "KV": "float8_e4m3fn", "C": "float32"} + (flops, moved), _call = _evaluated(name, row, case, slot_mapping=slots) + assert moved == _ledger( + name, + kv_c=((3, 6), _BF16), + k_pe=((3, 4), _BF16), + kv_cache_unread=True, + kv_cache_write=((3, 10), _FP8), + slot_mapping=((5,), _I64), + scale=((1,), _F32), + positions=((3,), _I64), + cos_sin_cache=((2, 4), _F32), # positions 0 and 2 + ) + assert flops == 3 * (2 * 10 + 3 * 4) + + @pytest.mark.parametrize( + "row,case", + [ + ({"seq_lens": [3, 4], "starts": [2, 5], "some": ["seq_starts"]}, {"KV": "bfloat16"}), + ( + {"seq_lens": [3, 4], "out_dtype": "float16", "some": ["scale"]}, + {"KV": "float8_e4m3fn"}, + ), + ], + ) + def test_paged_kv_cache_gather_reads_the_rows_each_range_reaches(self, row, case): + name = "PagedKVCacheGatherFwdOp" + row = {"T_q": 7, "NP": 6, "PS": 4, "W": 3, "E": [2, 3], **row} + (flops, moved), call = _evaluated(name, row, case) + table, page = call.values("block_table"), row["PS"] + starts = row.get("starts", [0, 0]) + ranges = [range(s, s + n) for s, n in zip(starts, row["seq_lens"], strict=True)] + cache_rows = {(table[b][j // page], j % page) for b, r in enumerate(ranges) for j in r} + pages = {(b, j // page) for b, r in enumerate(ranges) for j in r} + fp8 = case["KV"] == "float8_e4m3fn" + assert moved == _ledger( + name, + dst_unread=True, + dst_write=((7, 2, 3), torch.float16 if fp8 else _BF16), + cache=((len(cache_rows), 2, 3), _FP8 if fp8 else _BF16), + block_table=((len(pages),), _I32), + cu_seq_lens=((3,), _I32), + seq_starts=None if fp8 else ((2,), _I32), + scale=((1,), _F32) if fp8 else None, + ) + assert flops == (7 * 2 * 3 if fp8 else 0) + + def test_fused_qk_norm_rope_touches_the_q_and_k_columns_and_the_named_rows(self): + name = "FusedQKNormRopeFwdOp" + row = {"D": 8, "P": 16, "R": 4, "num_heads": 3, "num_kv_heads": 1, "seq_lens": [3, 2]} + (flops, moved), _call = _evaluated(name, row, {"T": "bfloat16", "C": "float32"}) + qk = ((5, 4 * 8), _BF16) # 5 tokens, 3 q heads and 1 k head of width 8 + assert moved == _ledger( + name, + qkv=qk, + qkv_write=qk, + q_weight=((8,), _BF16), + k_weight=((8,), _BF16), + cos_sin_cache=((3, 4), _F32), # positions 0, 1, 2 + positions=((5,), _I64), + ) + assert flops == 5 * 4 * (4 * 8 + 3 * 4) + + def test_chain_speculative_sampling_prices_the_cheaper_outcome(self): + name = "ChainSpeculativeSamplingFwdOp" + batch, n, vocab = 2, 3, 10 + (flops, moved), _call = _evaluated(name, {"B": batch, "N": n, "V": vocab}, {}) + # With N < V, every draft accepted is cheaper than a rejection at the first: the N + # ratio tests and a draw from target row N; the ids, 2 N probabilities, one row. + assert moved == _ledger( + name, + draft_probs=((batch * n,), _F32), + draft_token_ids=((batch, n), _I32), + target_probs=((batch * (n + vocab),), _F32), + seed=((1,), _I64), + offset=((1,), _I64), + output_token_ids=((batch, n + 1), _I32), + num_accepted=((batch,), _I32), + ) + assert flops == batch * (2 * n + 3 * vocab + 1) + + def test_top_k_masks_pay_nothing_on_a_row_k_leaves_whole(self): + vocab, ks = 10, [3, 10, 12, 1] + row = {"V": vocab, "k_list": ks} + (flops, _moved), _call = _evaluated("TopKMaskFwdOp", row, {"T": "float32"}) + assert flops == sum(2 * vocab for k in ks if k < vocab) + (flops, _moved), _call = _evaluated("TopKTopPMaskFwdOp", row, {"T": "float32"}) + # Top-k where k restricts; over the survivors max, subtract, exp, sum, compare and + # accumulate; p times the sum; the final mask. + assert flops == sum((vocab if k < vocab else 0) + 6 * min(k, vocab) + 1 + vocab for k in ks) + + def test_mla_varlen_scores_each_request_under_its_causal_mask(self): + row = {"T_q": 7, "H": 2, "DN": 4, "PE": 2, "DV": 3, "seq_lens": [3, 4]} + (flops, _moved), _call = _evaluated( + "MultiHeadLatentAttentionVarlenFwdOp", row, {"T": "bfloat16"} + ) + scores = sum(i + 1 for n in row["seq_lens"] for i in range(n)) + assert flops == _attention_flops(2, 6, 3, scores, 7) + + +# Coverage levels. Every op, implemented or spec-only, sits at # exactly one, and the level says what an independent recount rests on. # # one The binder builds the case from the manifest: signature, one workload @@ -554,6 +815,14 @@ def test_grouped_gemm_does_not_charge_the_padding_offsets_it_ignores(self): "NSAVarlenFwdOp": "how much it reads follows the values in `block_counts`", "NSATopkVarlenFwdOp": "`lse_in` is passed and the kernel recomputes the lse instead of reading it", "IndexedExpertMLPFwdOp": "the routed weight reads follow the values in `topk_ids`", + "GroupedQueryAttentionPagedFwdOp": "it reads the rows its page table names, not the pool", + "MultiHeadLatentAttentionPagedFwdOp": "it reads the cache rows its block table reaches, not the pool", + "DeepSeekSparseAttentionPagedFwdOp": "it reads the cache rows its valid index slots name", + "PagedKVCacheWriteFwdOp": "only the tokens `slot_mapping` gives a slot are read and written", + "MultiHeadLatentAttentionKVCacheWriteFwdOp": "only the tokens `slot_mapping` gives a slot are read and written", + "PagedKVCacheGatherFwdOp": "it reads the cache rows each request's range reaches, not the pool", + "FusedQKNormRopeFwdOp": "it leaves the v columns untouched and reads only the named cos/sin rows", + "ChainSpeculativeSamplingFwdOp": "where the chain stops is drawn at run time, so it prices the cheaper outcome", } # Level three: no independent recount is available. Empty, and an entry here has @@ -561,12 +830,11 @@ def test_grouped_gemm_does_not_charge_the_padding_offsets_it_ignores(self): NOT_RECOUNTABLE: dict[str, str] = {} -def _implemented_ops() -> list[str]: +def _entries() -> list[str]: + """Every entry: a spec-only one is recounted from its signature, needing no implementation.""" from tileops.manifest import load_manifest - return sorted( - name for name, entry in load_manifest().items() if entry.get("status") == "implemented" - ) + return sorted(load_manifest()) def _draws_metadata(op_name: str) -> bool: @@ -606,13 +874,13 @@ def _binder_agrees(op_name: str) -> bool: class TestCoverageLevels: - """Every implemented op sits at exactly one level, and the level is the truth.""" + """Every op sits at exactly one level, and the level is the truth.""" def test_a_generated_case_equals_its_op(self): from tests.roofline_binder import manifest_cases checked = 0 - for op_name in _implemented_ops(): + for op_name in _entries(): if op_name in HAND_WRITTEN or op_name in NOT_RECOUNTABLE: continue for label, dtype, op, oracle, _reads in manifest_cases(op_name): @@ -627,7 +895,7 @@ def test_a_generated_case_agrees_on_the_read_half(self): from tests.roofline_binder import manifest_cases checked = 0 - for op_name in _implemented_ops(): + for op_name in _entries(): if op_name in HAND_WRITTEN or op_name in NOT_RECOUNTABLE: continue for label, dtype, op, _oracle, reads in manifest_cases(op_name): @@ -637,11 +905,11 @@ def test_a_generated_case_agrees_on_the_read_half(self): checked += 1 assert checked > 0 - def test_every_implemented_op_sits_at_one_level(self): + def test_every_op_sits_at_one_level(self): both = sorted(set(HAND_WRITTEN) & set(NOT_RECOUNTABLE)) assert not both, f"declared at two levels: {both}" - unknown = sorted((set(HAND_WRITTEN) | set(NOT_RECOUNTABLE)) - set(_implemented_ops())) - assert not unknown, f"declared but not implemented: {unknown}" + unknown = sorted((set(HAND_WRITTEN) | set(NOT_RECOUNTABLE)) - set(_entries())) + assert not unknown, f"declared but not in the manifest: {unknown}" def test_a_declared_op_is_one_the_manifest_does_not_already_check(self): """Level two and three are for ops the manifest cannot recount, not a queue. diff --git a/tests/test_spec_reference.py b/tests/test_spec_reference.py new file mode 100644 index 000000000..c53a26fc3 --- /dev/null +++ b/tests/test_spec_reference.py @@ -0,0 +1,541 @@ +"""Reference conformance of spec-only entries that have no implementation yet. + +Each entry's reference is the torch expression its issue names, or the library reference +where one exists. Every workload row is instantiated; data tensors live on ``meta`` and +metadata tensors on the CPU with the row's generated values, so the reference runs as +written at the row's full size and costs nothing. Checked against the signature: the +reference's outputs have the inferred names, shapes and dtypes, and it writes exactly the +inputs the call's effects mark written. One call the signature rejects is rejected by the +reference too. +""" + +from __future__ import annotations + +import math + +import pytest +import torch +import torch.nn.functional as F +from torch.utils._python_dispatch import TorchDispatchMode + +from tests.roofline_binder import signature_class +from tileops.manifest import load_adts, load_manifest +from tileops.manifest.plan import entry_plan +from tileops.manifest.workload import instantiate + +pytestmark = pytest.mark.smoke + +_INF = float("inf") + + +# ---------------------------------------------------------------- quantization + + +def _int8_scale(amax): + """``amax / 127``, and 1.0 for an all-zero group.""" + return torch.where(amax > 0, amax / 127, torch.ones_like(amax)) + + +def _int8_per_tensor(p, t): + x = t["x"] + _m, _k = x.shape + xf = x.float() + scale = _int8_scale(xf.abs().amax()).reshape(1) + q = torch.round(xf / scale).clamp(-127, 127).to(torch.int8) + return {"q": q, "scale": scale} + + +def _int8_per_channel(p, t): + w = t["w"] + _n, _k = w.shape + wf = w.float() + scale = _int8_scale(wf.abs().amax(dim=1)) + return {"q": torch.round(wf / scale[:, None]).clamp(-127, 127).to(torch.int8), "scale": scale} + + +def _blocks(xf, block=128): + m, k = xf.shape + nb = -(-k // block) + return F.pad(xf, (0, nb * block - k)).view(m, nb, block), nb + + +def _int8_per_block(p, t): + xf = t["x"].float() + m, k = xf.shape + blocks, _nb = _blocks(xf) + scale = _int8_scale(blocks.abs().amax(-1)) + q = torch.round(xf / scale.repeat_interleave(128, 1)[:, :k]).clamp(-127, 127) + return {"q": q.to(torch.int8), "scale": scale} + + +def _int4_per_group(p, t): + w, g = t["w"], p["group_size"] + n, k = w.shape + wg = w.float().view(n, k // g, g) + lo, hi = wg.amin(-1), wg.amax(-1) + scale = torch.where(hi > lo, (hi - lo) / 15, torch.ones_like(hi)) + zero = torch.round(-lo / scale).clamp(0, 15) + q = torch.round(wg / scale[..., None] + zero[..., None]).clamp(0, 15).to(torch.uint8) + # Two values per byte in row order; the byte order GemmW4A16FwdOp consumes is a permutation + # of these bytes that only its repack kernel states, accepted by the GEMM round trip. + q = q.view(n, k // 2, 2) + packed = q[..., 0] | (q[..., 1] << 4) + return { + "packed_weight": packed, + "weight_scale": scale.to(w.dtype), + "weight_zero": zero.to(torch.uint8), + } + + +def _smooth_quant(p, t): + x, smooth = t["x"], t["smooth"] + _m, _k = x.shape + xs = x.float() / smooth + scale = _int8_scale(xs.abs().amax(dim=1)) + return {"q": torch.round(xs / scale[:, None]).clamp(-127, 127).to(torch.int8), "scale": scale} + + +def _dequant(expand): + def reference(p, t): + q, scale = t["q"], t["scale"] + m, k = q.shape + return {"x": (q.float() * expand(scale, m, k)).to(p["out_dtype"])} + + return reference + + +def _fp8_per_block(p, t): + w = t["w"] + n, k = w.shape + nn, nk = -(-n // 128), -(-k // 128) + wf = F.pad(w.float(), (0, nk * 128 - k, 0, nn * 128 - n)).view(nn, 128, nk, 128) + amax = wf.abs().amax(dim=(1, 3)) + scale = torch.where(amax > 0, amax / 448, torch.ones_like(amax)) + full = scale.repeat_interleave(128, 0).repeat_interleave(128, 1)[:n, :k] + return {"q": (w.float() / full).clamp(-448, 448).to(torch.float8_e4m3fn), "scale": scale} + + +# ---------------------------------------------------------------- sampling + + +def _top_k(logits, k): + batch, vocab = logits.shape + k = k.view(batch) + kth = ( + logits.float().sort(-1, descending=True).values.gather(1, (k.clamp(max=vocab) - 1)[:, None]) + ) + kept = (logits.float() >= kth) | (k >= vocab).to(logits.device)[:, None] + return logits.masked_fill(~kept, -_INF) + + +def _top_p(logits, p): + probs = logits.float().softmax(-1) + p = p.view(logits.shape[0]) + sorted_probs, order = probs.sort(-1, descending=True) + exclusive = sorted_probs.cumsum(-1) - sorted_probs + removed = torch.zeros_like(probs, dtype=torch.bool).scatter( + 1, order, exclusive >= p.view(-1, 1) + ) + return logits.masked_fill(removed, -_INF) + + +def _top_k_mask(p, t): + return {"masked_logits": _top_k(t["logits"], t["k"])} + + +def _min_p_mask(p, t): + logits, min_p = t["logits"], t["min_p"] + probs = logits.float().softmax(-1) + removed = probs < min_p.view(logits.shape[0], 1) * probs.amax(-1, keepdim=True) + return {"masked_logits": logits.masked_fill(removed, -_INF)} + + +def _top_p_mask(p, t): + return {"masked_logits": _top_p(t["logits"], t["p"])} + + +def _top_k_top_p_mask(p, t): + return {"masked_logits": _top_p(_top_k(t["logits"], t["k"]), t["p"])} + + +def _sampling_from_probs(p, t): + return {"samples": torch.multinomial(t["probs"], 1).squeeze(1).to(torch.int32)} + + +def _chain_speculative_sampling(p, t): + draft, target, ids = t["draft_probs"], t["target_probs"], t["draft_token_ids"].long() + batch, n, vocab = draft.shape + d = draft.gather(2, ids[..., None]).squeeze(-1) + q = target[:, :n].gather(2, ids[..., None]).squeeze(-1) + accepted = torch.rand(batch, n, device=draft.device) * d < q + num = accepted.int().cumprod(1).sum(1) + padded_draft = torch.cat([draft, torch.zeros_like(draft[:, :1])], 1) + pick = num[:, None, None].expand(batch, 1, vocab) + residual = (target.gather(1, pick) - padded_draft.gather(1, pick)).squeeze(1).clamp_min(0) + resampled = torch.multinomial(residual, 1) + position = torch.arange(n + 1, device=draft.device)[None] + drafts = torch.cat([ids.to(draft.device), torch.full_like(resampled, -1)], 1) + tokens = torch.where( + position < num[:, None], drafts, torch.where(position == num[:, None], resampled, -1) + ) + return {"output_token_ids": tokens.to(torch.int32), "num_accepted": num.to(torch.int32)} + + +# ---------------------------------------------------------------- attention and caches + + +def _attend(q, k, v, scale, visible): + """Float32 softmax attention of ``q [S, H, Dk]`` over ``k [N, Dk]``/``v [N, Dv]`` shared by + every head, or per head when ``k`` is ``[N, H, Dk]``; ``visible [S, N]`` on the CPU.""" + kh = k if k.dim() == 3 else k[:, None].expand(-1, q.shape[1], -1) + vh = v if v.dim() == 3 else v[:, None].expand(-1, q.shape[1], -1) + scores = torch.einsum("shd,nhd->hsn", q.float(), kh.float()) * scale + scores = scores.masked_fill(~visible.to(q.device)[None], -_INF) + return torch.einsum("hsn,nhd->shd", scores.softmax(-1), vh.float()), scores.logsumexp(-1).T + + +def _causal(queries, keys, is_causal): + """Bottom-right aligned visibility of ``queries`` rows over ``keys``.""" + rows = torch.arange(queries)[:, None] + keys - queries + return ( + torch.arange(keys)[None] <= rows + if is_causal + else torch.ones(queries, keys, dtype=torch.bool) + ) + + +def _paged_rows(cache, table, start, end): + """Rows ``[start, end)`` of one request, read through its block table.""" + page_size = cache.shape[1] + pages = cache[table[: -(-end // page_size)].long()] + return pages.flatten(0, 1)[start:end] + + +def _mla_paged(p, t): + q, cache, table, lens = t["q"], t["kv_cache"], t["block_table"], t["cache_seqlens"] + batch, s_q, _h, dk = q.shape + rank = p["kv_lora_rank"] + scale = p["sm_scale"] if p["sm_scale"] is not None else dk**-0.5 + outs, lses = [], [] + for b in range(batch): + kv = _paged_rows(cache, table[b], 0, int(lens[b])).float() + if t.get("kv_scale") is not None: + kv = kv * t["kv_scale"] + o, lse = _attend(q[b], kv, kv[:, :rank], scale, _causal(s_q, kv.shape[0], p["is_causal"])) + outs.append(o) + lses.append(lse) + return {"o": torch.stack(outs).to(q.dtype), "lse": torch.stack(lses)} + + +def _mla_varlen(p, t): + q, k_nope, k_pe, v, cu = t["q"], t["k_nope"], t["k_pe"], t["v"], t["cu_seqlens"].tolist() + heads = q.shape[1] + scale = p["sm_scale"] if p["sm_scale"] is not None else q.shape[-1] ** -0.5 + outs, lses = [], [] + for a, e in zip(cu, cu[1:], strict=False): + k = torch.cat([k_nope[a:e], k_pe[a:e, None].expand(-1, heads, -1)], -1) + o, lse = _attend(q[a:e], k, v[a:e], scale, _causal(e - a, e - a, p["is_causal"])) + outs.append(o) + lses.append(lse) + return {"o": torch.cat(outs).to(q.dtype), "lse": torch.cat(lses)} + + +def _stored(rows, cache, scale): + """Rows as ``cache`` stores them, one cache row each: divided by the scale when it has one.""" + rows = rows.reshape(rows.shape[0], *cache.shape[2:]) + return (rows.float() / scale if scale is not None else rows).to(cache.dtype) + + +def _paged_kv_cache_write(p, t): + slots = t["slot_mapping"] + tokens = (slots >= 0).nonzero().squeeze(1) + for name, pages, scale in (("k", "k_pages", "k_scale"), ("v", "v_pages", "v_scale")): + cache = t[pages] + cache.view(-1, *cache.shape[2:])[slots[tokens]] = _stored( + t[name][tokens], cache, t.get(scale) + ) + return {} + + +def _rope(x, cos_sin, positions, layout): + """Rotate ``x [N, ..., R']`` on its first ``R`` columns by ``cos_sin [P, R]`` at ``positions``.""" + half = cos_sin.shape[1] // 2 + cos, sin = cos_sin[positions].float().chunk(2, -1) + shape = (x.shape[0],) + (1,) * (x.dim() - 2) + (half,) + cos, sin = cos.view(shape), sin.view(shape) + rot, rest = x[..., : 2 * half].float(), x[..., 2 * half :] + if layout == "neox": + a, b = rot[..., :half], rot[..., half:] + rotated = torch.cat([a * cos - b * sin, b * cos + a * sin], -1) + else: + a, b = rot[..., 0::2], rot[..., 1::2] + rotated = torch.stack([a * cos - b * sin, b * cos + a * sin], -1).flatten(-2) + return torch.cat([rotated.to(x.dtype), rest], -1) + + +def _mla_kv_cache_write(p, t): + slots, cache = t["slot_mapping"], t["kv_cache"] + tokens = (slots >= 0).nonzero().squeeze(1) + k_pe = t["k_pe"][tokens] + if p["fuse_rope"]: + k_pe = _rope(k_pe, t["cos_sin_cache"], t["positions"][tokens], p["rope_layout"]) + rows = torch.cat([t["kv_c"][tokens], k_pe], -1) + cache.view(-1, cache.shape[-1])[slots[tokens]] = _stored(rows, cache, t.get("scale")) + return {} + + +def _fused_qk_norm_rope(p, t): + qkv, eps = t["qkv"], p["eps"] + heads, kv_heads = p["num_heads"], p["num_kv_heads"] + dim = t["q_weight"].shape[0] + tokens = qkv.shape[0] + for first, count, weight in ((0, heads, t["q_weight"]), (heads, kv_heads, t["k_weight"])): + view = qkv[:, first * dim : (first + count) * dim].view(tokens, count, dim) + normed = F.rms_norm(view.float(), (dim,), weight.float(), eps).to(qkv.dtype) + view.copy_(_rope(normed, t["cos_sin_cache"], t["positions"], p["rope_layout"])) + return {} + + +def _merge_attention_states(p, t): + s_a, s_b = t["s_a"], t["s_b"] + top = torch.maximum(s_a, s_b) + empty = top == -_INF + shift = torch.where(empty, torch.zeros_like(top), top) + w_a, w_b = (s_a - shift).exp(), (s_b - shift).exp() + total = torch.where(empty, torch.ones_like(top), w_a + w_b) + v = t["v_a"].float() * (w_a / total)[..., None] + t["v_b"].float() * (w_b / total)[..., None] + return {"v": v.to(t["v_a"].dtype), "s": torch.where(empty, top, top + total.log())} + + +def _dsa_paged(p, t): + q, cache, table, lens, indices = ( + t["q"], + t["kv_cache"], + t["block_table"], + t["cache_seqlens"], + t["indices"], + ) + batch, s_q, _h, dk = q.shape + page_size = cache.shape[1] + scale = p["sm_scale"] if p["sm_scale"] is not None else dk**-0.5 + flat = cache.view(-1, cache.shape[-1]) + outs, lses = [], [] + for b in range(batch): + for s in range(s_q): + pos = indices[b, s].long() + pos = pos[(pos >= 0) & (pos < lens[b])] + rows = flat[table[b][pos // page_size].long() * page_size + pos % page_size] + latent = rows[:, :512].view(torch.float8_e4m3fn).float() + scales = rows[:, 512:528].view(torch.float32).repeat_interleave(128, 1) + value = latent * scales + key = torch.cat([value, rows[:, 528:].view(torch.bfloat16).float()], -1) + visible = torch.ones(1, key.shape[0], dtype=torch.bool) + o, lse = _attend(q[b, s : s + 1], key, value, scale, visible) + outs.append(o[0]) + lses.append(lse[0]) + shape = (batch, s_q) + return { + "o": torch.stack(outs).view(*shape, *outs[0].shape).to(q.dtype), + "lse": torch.stack(lses).view(*shape, -1), + } + + +def _paged_kv_cache_gather(p, t): + dst, cache, table = t["dst"], t["cache"], t["block_table"] + cu = t["cu_seq_lens"].tolist() + starts = t["seq_starts"].tolist() if t.get("seq_starts") is not None else [0] * (len(cu) - 1) + for b, (a, e) in enumerate(zip(cu, cu[1:], strict=False)): + rows = _paged_rows(cache, table[b], starts[b], starts[b] + e - a) + if t.get("scale") is not None: + rows = rows.float() * t["scale"] + dst[a:e] = rows.to(dst.dtype) + return {} + + +# ---------------------------------------------------------------- linear attention + + +def _kda(p, t): + naive = pytest.importorskip( + "fla.ops.kda.naive", + reason="KimiDeltaAttentionFwdOp unverified: its reference, FLA (package `fla`), is not installed", + ).naive_recurrent_kda + q, k, v, g, beta = t["q"], t["k"], t["v"], t["g"], t["beta"] + if p["use_qk_l2norm_in_kernel"]: + q, k = ( + F.normalize(q.float(), dim=-1).to(q.dtype), + F.normalize(k.float(), dim=-1).to(k.dtype), + ) + if p["use_gate_in_kernel"]: + heads = v.shape[2] + a = t["A_log"].exp()[:, None] + x = g.float() + (t["dt_bias"].view(heads, -1) if t.get("dt_bias") is not None else 0) + lower = p["lower_bound"] + g = lower * torch.sigmoid(a * x) if lower is not None else -a * F.softplus(x) + if p["use_beta_sigmoid_in_kernel"]: + beta = beta.float().sigmoid() * (2 if p["allow_neg_eigval"] else 1) + cu = t["cu_seqlens"].tolist() if t.get("cu_seqlens") is not None else None + spans = list(zip(cu, cu[1:], strict=False)) if cu else [(0, q.shape[1])] + init = t.get("initial_state") + if init is not None and p["state_v_first"]: + init = init.transpose(-1, -2) + outs, states = [], [] + for i, (a, e) in enumerate(spans): + sel = slice(a, e) + h0 = init[i : i + 1] if cu and init is not None else init + o, s = naive(q[:, sel], k[:, sel], v[:, sel], g[:, sel], beta[:, sel], p["scale"], h0, True) + outs.append(o) + states.append(s) + state = torch.cat(states) if cu else states[0] + if p["state_v_first"]: + state = state.transpose(-1, -2) + return {"o": torch.cat(outs, 1).to(v.dtype), "final_state": state.float()} + + +# ---------------------------------------------------------------- entries + + +def _narrow(name, axis=-1): + """A rejected call: tensor *name* one element shorter along *axis*.""" + + def edit(tensors): + x = tensors[name] + tensors[name] = x.narrow(axis, 0, x.shape[axis] - 1) + + return edit + + +def _unsqueeze(name): + def edit(tensors): + tensors[name] = tensors[name][None] + + return edit + + +def _empty_rows(name): + def edit(tensors): + tensors[name] = tensors[name][:0] + + return edit + + +REFERENCES = { + # op: (reference, the edit of the first row's tensors that the signature rejects) + "INT8QuantPerTensorFwdOp": (_int8_per_tensor, _empty_rows("x")), + "INT8QuantPerChannelFwdOp": (_int8_per_channel, _unsqueeze("w")), + "INT8QuantPerBlockFwdOp": (_int8_per_block, _unsqueeze("x")), + "INT4QuantPerGroupFwdOp": (_int4_per_group, _narrow("w")), + "SmoothQuantFwdOp": (_smooth_quant, _narrow("smooth")), + "INT8DequantPerTensorFwdOp": (_dequant(lambda s, m, k: s), _unsqueeze("q")), + "INT8DequantPerChannelFwdOp": (_dequant(lambda s, m, k: s[:, None]), _narrow("scale")), + "INT8DequantPerBlockFwdOp": ( + _dequant(lambda s, m, k: s.repeat_interleave(128, 1)[:, :k]), + _narrow("scale", 0), + ), + "FP8QuantPerBlockFwdOp": (_fp8_per_block, _unsqueeze("w")), + "TopKMaskFwdOp": (_top_k_mask, _narrow("k")), + "MinPMaskFwdOp": (_min_p_mask, _narrow("min_p", 0)), + "TopPMaskFwdOp": (_top_p_mask, _narrow("p", 0)), + "TopKTopPMaskFwdOp": (_top_k_top_p_mask, _narrow("p", 0)), + "SamplingFromProbsFwdOp": (_sampling_from_probs, _unsqueeze("probs")), + "ChainSpeculativeSamplingFwdOp": (_chain_speculative_sampling, _narrow("draft_token_ids")), + "MultiHeadLatentAttentionPagedFwdOp": (_mla_paged, _narrow("kv_cache")), + "MultiHeadLatentAttentionVarlenFwdOp": (_mla_varlen, _narrow("k_nope")), + "PagedKVCacheWriteFwdOp": (_paged_kv_cache_write, _narrow("v")), + "FusedQKNormRopeFwdOp": (_fused_qk_norm_rope, _narrow("k_weight")), + "MultiHeadLatentAttentionKVCacheWriteFwdOp": (_mla_kv_cache_write, _narrow("kv_cache")), + "MergeAttentionStatesFwdOp": (_merge_attention_states, _narrow("v_b")), + "DeepSeekSparseAttentionPagedFwdOp": (_dsa_paged, _narrow("kv_cache")), + "PagedKVCacheGatherFwdOp": (_paged_kv_cache_gather, _narrow("dst")), + "KimiDeltaAttentionFwdOp": (_kda, _narrow("g")), +} + + +class _Writes(TorchDispatchMode): + """The storages the dispatched ops write, by their schema's mutable arguments.""" + + def __init__(self): + super().__init__() + self.storages = set() + + def __torch_dispatch__(self, func, types, args=(), kwargs=None): + kwargs = kwargs or {} + for i, arg in enumerate(func._schema.arguments): + value = args[i] if i < len(args) else kwargs.get(arg.name) + if ( + arg.alias_info is not None + and arg.alias_info.is_write + and isinstance(value, torch.Tensor) + ): + self.storages.add(value.untyped_storage()._cdata) + return func(*args, **kwargs) + + +def _calls(name): + """Per row and dtype case: the instantiated call, its meta tensors and the reference's + tensors (metadata on the CPU with the row's values).""" + entry = load_manifest()[name] + plan = entry_plan(name, entry, load_adts()) + for row in entry["workloads"]: + for case in row.get("dtype_cases") or [{}]: + call = instantiate(plan, row, case) + meta = call.materialize("meta") + host = dict(meta) + for n, spec in call.specs.items(): + if spec is not None and spec.values is not None and n in meta: + host[n] = torch.tensor(spec.values, dtype=getattr(torch, spec.dtype)).reshape( + spec.shape + ) + yield row["label"], case, plan, call, meta, host + + +@pytest.mark.parametrize("name", sorted(REFERENCES)) +def test_reference_agrees_with_the_signature(name): + entry = load_manifest()[name] + reference, _reject = REFERENCES[name] + cls = signature_class(name, entry) + for label, case, plan, call, meta, host in _calls(name): + where = f"{name} {label} {case}" + op = cls(**call.arguments(meta)) + checked = cls._signature.check(op, {n: meta[n] for n in plan.sig.inputs}) + inputs = {n: host[n] for n in plan.sig.inputs} + with _Writes() as writes: + outputs = reference(call.arguments(meta), inputs) + declared = {n: call.tensors[n] for n in plan.sig.outputs} + got = {n: (tuple(v.shape), str(v.dtype).removeprefix("torch.")) for n, v in outputs.items()} + assert got == {n: (tuple(s), d) for n, (s, d) in declared.items()}, where + written = { + n + for n, v in inputs.items() + if v is not None and v.untyped_storage()._cdata in writes.storages + } + assert written == set(checked.written), where + + +@pytest.mark.parametrize("name", sorted(REFERENCES)) +def test_a_call_the_signature_rejects_the_reference_rejects(name): + reference, reject = REFERENCES[name] + label, case, plan, call, meta, host = next(_calls(name)) + cls = signature_class(name, load_manifest()[name]) + op = cls(**call.arguments(meta)) + bad_meta = {n: meta[n] for n in plan.sig.inputs} + bad_host = {n: host[n] for n in plan.sig.inputs} + reject(bad_meta) + reject(bad_host) + with pytest.raises((ValueError, TypeError)): + cls._signature.check(op, bad_meta) + with pytest.raises((RuntimeError, ValueError, IndexError, TypeError)): + reference(call.arguments(meta), bad_host) + + +def test_a_write_only_buffer_is_written_whole(): + """`PagedKVCacheGatherFwdOp`'s `dst` is write-only: on a small call with real values the + reference leaves no row of it unwritten.""" + name = "PagedKVCacheGatherFwdOp" + plan = entry_plan(name, load_manifest()[name], load_adts()) + row = {"T_q": 7, "NP": 6, "PS": 4, "W": 3, "E": [2, 3], "seq_lens": [3, 4], "starts": [2, 5]} + call = instantiate(plan, {**row, "some": ["seq_starts"], "label": "s"}, {"KV": "bfloat16"}) + tensors = call.materialize("cpu") + tensors["dst"].fill_(math.nan) + _paged_kv_cache_gather(call.arguments(tensors), {n: tensors[n] for n in plan.sig.inputs}) + assert not tensors["dst"].isnan().any() diff --git a/workloads/elementwise.py b/workloads/elementwise.py index 1582efe43..b284987ab 100644 --- a/workloads/elementwise.py +++ b/workloads/elementwise.py @@ -141,7 +141,7 @@ def gen_inputs(self) -> tuple[torch.Tensor]: if self.dtype == torch.uint8: x = torch.randint(0, 8, (self.n_total,), device=run_device(), dtype=self.dtype) - elif self.dtype in (torch.int8, torch.int16, torch.int32, torch.int64): + elif not (self.dtype.is_floating_point or self.dtype.is_complex) and self.dtype.is_signed: x = torch.randint(-4, 4, (self.n_total,), device=run_device(), dtype=self.dtype) else: x = torch.randn(self.n_total, device=run_device(), dtype=self.dtype) diff --git a/workloads/gemm.py b/workloads/gemm.py index b774d7220..adff6e46d 100644 --- a/workloads/gemm.py +++ b/workloads/gemm.py @@ -11,6 +11,9 @@ W4A16_GROUP_SIZE = 128 +_FP8_INIT_SCALE: float = 0.25 + + class GemmWorkload(WorkloadBase): def __init__( self, @@ -334,9 +337,6 @@ def ref_program(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: return torch.bmm(a, b) -_FP8_INIT_SCALE: float = 0.25 - - class BmmFp8Workload(WorkloadBase): """Workload for batched FP8 GEMM.